use ndarray::{Array1, Array2, ArrayView1, ArrayView2};
use statrs::distribution::{ContinuousCDF, Normal};
use crate::error::{RegressionError, Result};
use crate::linalg::dmatrix_from_rows;
const PROB_EPS: f64 = 1e-12;
#[derive(Debug, Clone)]
pub struct MultinomialFit {
x: Array2<f64>,
y: Array1<f64>,
coefficients: Array2<f64>,
probabilities: Array2<f64>,
cov: Array2<f64>,
log_likelihood: f64,
intercept_col: Option<usize>,
iterations: usize,
n: usize,
p: usize,
k: usize,
}
impl MultinomialFit {
pub fn new(x: Array2<f64>, y: Array1<f64>) -> Result<Self> {
Self::with_options(x, y, 100, 1e-10)
}
pub fn with_options(x: Array2<f64>, y: Array1<f64>, max_iter: usize, tol: f64) -> Result<Self> {
let n = x.nrows();
let p = x.ncols();
if n == 0 || p == 0 {
return Err(RegressionError::EmptyInput { what: "X" });
}
if y.len() != n {
return Err(RegressionError::ShapeMismatch {
what: "y length vs X rows",
expected: n,
got: y.len(),
});
}
let k = validate_labels(&y)?;
let m = (k - 1) * p;
let mut beta = Array2::<f64>::zeros((k - 1, p));
let mut probs = Array2::<f64>::zeros((n, k));
let mut cov = Array2::<f64>::zeros((m, m));
let mut iterations = 0usize;
let mut converged = false;
while iterations < max_iter {
iterations += 1;
fill_probabilities(&x, &beta, &mut probs);
let mut grad = Array1::<f64>::zeros(m);
let mut info = Array2::<f64>::zeros((m, m));
for kk in 1..k {
let bk = kk - 1;
for a in 0..p {
let mut g = 0.0;
for i in 0..n {
let yik = if y[i] as usize == kk { 1.0 } else { 0.0 };
g += x[(i, a)] * (yik - probs[(i, kk)]);
}
grad[bk * p + a] = g;
}
}
for kk in 1..k {
for ll in 1..k {
let bk = kk - 1;
let bl = ll - 1;
let delta = if kk == ll { 1.0 } else { 0.0 };
for a in 0..p {
for b in 0..p {
let mut s = 0.0;
for i in 0..n {
let w = probs[(i, kk)] * (delta - probs[(i, ll)]);
s += x[(i, a)] * w * x[(i, b)];
}
info[(bk * p + a, bl * p + b)] = s;
}
}
}
}
let info_dm = dmatrix_from_rows(m, m, info.as_standard_layout().as_slice().unwrap());
let inv = info_dm.try_inverse().ok_or(RegressionError::RankDeficient)?;
let inv_arr = Array2::from_shape_fn((m, m), |(i, j)| inv[(i, j)]);
let delta = inv_arr.dot(&grad);
for kk in 1..k {
let bk = kk - 1;
for a in 0..p {
beta[(bk, a)] += delta[bk * p + a];
}
}
cov = inv_arr;
let step = delta.iter().fold(0.0_f64, |mx, v| mx.max(v.abs()));
if !beta.iter().all(|v| v.is_finite()) || beta.iter().any(|v| v.abs() > 1e8) {
return Err(RegressionError::NotConverged {
iterations,
msg: "coefficients diverging (likely separation)".into(),
});
}
if step < tol {
converged = true;
break;
}
}
if !converged {
return Err(RegressionError::NotConverged {
iterations,
msg: "Newton iteration did not reach tolerance".into(),
});
}
fill_probabilities(&x, &beta, &mut probs);
let log_likelihood = (0..n)
.map(|i| probs[(i, y[i] as usize)].max(PROB_EPS).ln())
.sum();
let intercept_col = detect_constant_column(&x);
Ok(Self {
x,
y,
coefficients: beta,
probabilities: probs,
cov,
log_likelihood,
intercept_col,
iterations,
n,
p,
k,
})
}
pub fn n_observations(&self) -> usize {
self.n
}
pub fn n_features(&self) -> usize {
self.p
}
pub fn n_classes(&self) -> usize {
self.k
}
pub fn n_parameters(&self) -> usize {
(self.k - 1) * self.p
}
pub fn has_intercept(&self) -> bool {
self.intercept_col.is_some()
}
pub fn iterations(&self) -> usize {
self.iterations
}
pub fn design_matrix(&self) -> ArrayView2<'_, f64> {
self.x.view()
}
pub fn response(&self) -> ArrayView1<'_, f64> {
self.y.view()
}
pub fn coefficients(&self) -> ArrayView2<'_, f64> {
self.coefficients.view()
}
pub fn fitted_probabilities(&self) -> ArrayView2<'_, f64> {
self.probabilities.view()
}
pub fn covariance(&self) -> ArrayView2<'_, f64> {
self.cov.view()
}
pub fn log_likelihood(&self) -> f64 {
self.log_likelihood
}
pub fn coefficient_standard_errors(&self) -> Array2<f64> {
Array2::from_shape_fn((self.k - 1, self.p), |(bk, a)| {
let idx = bk * self.p + a;
self.cov[(idx, idx)].max(0.0).sqrt()
})
}
pub fn z_values(&self) -> Array2<f64> {
let se = self.coefficient_standard_errors();
Array2::from_shape_fn((self.k - 1, self.p), |(bk, a)| {
if se[(bk, a)] > 0.0 {
self.coefficients[(bk, a)] / se[(bk, a)]
} else {
f64::NAN
}
})
}
pub fn p_values(&self) -> Array2<f64> {
let z = self.z_values();
let normal = Normal::new(0.0, 1.0).expect("standard normal");
Array2::from_shape_fn((self.k - 1, self.p), |(bk, a)| {
let zv = z[(bk, a)];
if zv.is_finite() {
2.0 * (1.0 - normal.cdf(zv.abs()))
} else {
f64::NAN
}
})
}
pub fn residual_deviance(&self) -> f64 {
-2.0 * self.log_likelihood
}
pub fn null_deviance(&self) -> f64 {
-2.0 * self.null_log_likelihood()
}
fn null_log_likelihood(&self) -> f64 {
let n = self.n as f64;
let mut counts = vec![0.0_f64; self.k];
for &yi in self.y.iter() {
counts[yi as usize] += 1.0;
}
counts
.iter()
.filter(|&&c| c > 0.0)
.map(|&c| c * (c / n).ln())
.sum()
}
pub fn mcfadden_r2(&self) -> f64 {
let ll0 = self.null_log_likelihood();
if ll0 != 0.0 {
1.0 - self.log_likelihood / ll0
} else {
f64::NAN
}
}
pub fn aic(&self) -> f64 {
self.residual_deviance() + 2.0 * self.n_parameters() as f64
}
pub fn bic(&self) -> f64 {
self.residual_deviance() + (self.n as f64).ln() * self.n_parameters() as f64
}
pub fn predict_proba(&self, x: ArrayView2<'_, f64>) -> Array2<f64> {
let rows = x.nrows();
let mut out = Array2::<f64>::zeros((rows, self.k));
let xo = x.to_owned();
fill_probabilities(&xo, &self.coefficients, &mut out);
out
}
}
pub fn deviance_residuals(fit: &MultinomialFit) -> Array1<f64> {
let y = fit.response();
let p = fit.fitted_probabilities();
Array1::from_shape_fn(fit.n_observations(), |i| {
let pi = p[(i, y[i] as usize)].max(PROB_EPS);
(-2.0 * pi.ln()).max(0.0).sqrt()
})
}
fn fill_probabilities(x: &Array2<f64>, beta: &Array2<f64>, probs: &mut Array2<f64>) {
let n = x.nrows();
let p = x.ncols();
let k = beta.nrows() + 1;
for i in 0..n {
let mut eta = vec![0.0_f64; k];
let mut maxe = 0.0_f64;
for kk in 1..k {
let mut e = 0.0;
for a in 0..p {
e += x[(i, a)] * beta[(kk - 1, a)];
}
eta[kk] = e;
if e > maxe {
maxe = e;
}
}
let mut denom = 0.0;
for e in eta.iter_mut() {
*e = (*e - maxe).exp();
denom += *e;
}
for kk in 0..k {
probs[(i, kk)] = eta[kk] / denom;
}
}
}
fn validate_labels(y: &Array1<f64>) -> Result<usize> {
let mut max_label = 0usize;
for &v in y.iter() {
if !v.is_finite() || v < 0.0 || v.fract() != 0.0 {
return Err(RegressionError::InvalidResponse {
msg: format!("class labels must be non-negative integers, found {v}"),
});
}
max_label = max_label.max(v as usize);
}
let k = max_label + 1;
if k < 2 {
return Err(RegressionError::InvalidResponse {
msg: "multinomial response needs at least two classes".into(),
});
}
let mut present = vec![false; k];
for &v in y.iter() {
present[v as usize] = true;
}
if let Some(missing) = present.iter().position(|&b| !b) {
return Err(RegressionError::InvalidResponse {
msg: format!("class {missing} has no observations; labels must be 0..K-1 with all present"),
});
}
Ok(k)
}
fn detect_constant_column(x: &Array2<f64>) -> Option<usize> {
for (j, col) in x.columns().into_iter().enumerate() {
let first = col[0];
let scale = first.abs().max(1.0);
if col.iter().all(|&v| (v - first).abs() <= 1e-12 * scale) {
return Some(j);
}
}
None
}