use ndarray::{Array1, Array2};
#[derive(Clone, Debug, PartialEq)]
pub struct LogisticModelWeights {
pub coefficients: Array2<f64>,
pub intercept: Array1<f64>,
pub n_classes: usize,
}
impl LogisticModelWeights {
pub fn new(
coefficients: Array2<f64>,
intercept: Array1<f64>,
n_classes: usize,
) -> crate::Result<Self> {
let scores = coefficients.nrows();
let valid = n_classes >= 2
&& intercept.len() == scores
&& ((n_classes == 2 && scores == 1) || (n_classes > 2 && scores == n_classes));
if !valid {
return Err(crate::Error::InvalidModel(format!(
"logistic shape mismatch: {scores} coefficient rows, {} intercepts, {n_classes} classes",
intercept.len()
)));
}
Ok(Self {
coefficients,
intercept,
n_classes,
})
}
#[must_use]
pub fn n_features(&self) -> usize {
self.coefficients.ncols()
}
}