Skip to main content

qualia_core_db/solvers/learning/glm/
multinomial.rs

1//! Multinomial logistic regression / softmax classifier (ISL ch 4.3.5, PRML ch 4).
2//!
3//! `P(y=c | x) = softmax(W_c·φ(x))`. Fit by gradient ascent on the regularized
4//! log-likelihood (the multinomial cross-entropy is convex, so this converges to
5//! the global optimum); a small L2 term keeps the solution finite under separation.
6//! Kernel-class `DenseLinear` (the logits) — scalar fit loop is CPU.
7
8use crate::solvers::learning::LearningError;
9
10/// A fitted softmax classifier. `weights` is `n_classes × k` row-major
11/// (`k = p + intercept`).
12#[derive(Debug, Clone)]
13pub struct MultinomialLogistic {
14    pub classes: Vec<usize>,
15    weights: Vec<f64>,
16    fit_intercept: bool,
17    p: usize,
18}
19
20fn design_row(x_row: &[f64], fit_intercept: bool, out: &mut [f64]) {
21    if fit_intercept {
22        out[0] = 1.0;
23        out[1..].copy_from_slice(x_row);
24    } else {
25        out.copy_from_slice(x_row);
26    }
27}
28
29/// Numerically-stable softmax of `logits` into `out`.
30fn softmax(logits: &[f64], out: &mut [f64]) {
31    let m = logits.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
32    let mut s = 0.0;
33    for (o, &l) in out.iter_mut().zip(logits) {
34        *o = (l - m).exp();
35        s += *o;
36    }
37    if s > 0.0 {
38        for o in out.iter_mut() {
39            *o /= s;
40        }
41    }
42}
43
44impl MultinomialLogistic {
45    /// Fit by gradient ascent. `lr` learning rate, `l2` ridge penalty (≥ 0),
46    /// `max_iter` iterations. Fails closed on shape mismatch / a single class.
47    pub fn fit(
48        x: &[f64],
49        y: &[usize],
50        n: usize,
51        p: usize,
52        fit_intercept: bool,
53        lr: f64,
54        l2: f64,
55        max_iter: usize,
56    ) -> Result<Self, LearningError> {
57        if n == 0 || p == 0 || x.len() != n * p || y.len() != n {
58            return Err(LearningError::InvalidDimension);
59        }
60        if !(lr > 0.0) || l2 < 0.0 {
61            return Err(LearningError::InsufficientData);
62        }
63        let mut classes: Vec<usize> = y.to_vec();
64        classes.sort_unstable();
65        classes.dedup();
66        let c = classes.len();
67        if c < 2 {
68            return Err(LearningError::InsufficientData);
69        }
70        let cls_idx = |label: usize| classes.iter().position(|&v| v == label).unwrap();
71        let k = p + usize::from(fit_intercept);
72
73        // Design matrix.
74        let mut d = vec![0.0; n * k];
75        for i in 0..n {
76            design_row(
77                &x[i * p..(i + 1) * p],
78                fit_intercept,
79                &mut d[i * k..(i + 1) * k],
80            );
81        }
82
83        let mut w = vec![0.0; c * k];
84        let mut logits = vec![0.0; c];
85        let mut probs = vec![0.0; c];
86        let mut grad = vec![0.0; c * k];
87
88        for _ in 0..max_iter.max(1) {
89            grad.iter_mut().for_each(|g| *g = 0.0);
90            for i in 0..n {
91                let di = &d[i * k..(i + 1) * k];
92                for cc in 0..c {
93                    logits[cc] = (0..k).map(|j| w[cc * k + j] * di[j]).sum();
94                }
95                softmax(&logits, &mut probs);
96                let yi = cls_idx(y[i]);
97                for cc in 0..c {
98                    let err = (if cc == yi { 1.0 } else { 0.0 }) - probs[cc];
99                    for j in 0..k {
100                        grad[cc * k + j] += err * di[j];
101                    }
102                }
103            }
104            // Gradient step with L2 shrinkage (don't penalize the intercept column).
105            for cc in 0..c {
106                for j in 0..k {
107                    let mut g = grad[cc * k + j] / n as f64;
108                    if !(fit_intercept && j == 0) {
109                        g -= l2 * w[cc * k + j];
110                    }
111                    w[cc * k + j] += lr * g;
112                }
113            }
114        }
115
116        Ok(Self {
117            classes,
118            weights: w,
119            fit_intercept,
120            p,
121        })
122    }
123
124    /// Class probabilities for one row, aligned with `classes`.
125    pub fn predict_proba_row(&self, x_row: &[f64]) -> Vec<f64> {
126        let c = self.classes.len();
127        let k = self.p + usize::from(self.fit_intercept);
128        let mut di = vec![0.0; k];
129        design_row(x_row, self.fit_intercept, &mut di);
130        let logits: Vec<f64> = (0..c)
131            .map(|cc| (0..k).map(|j| self.weights[cc * k + j] * di[j]).sum())
132            .collect();
133        let mut probs = vec![0.0; c];
134        softmax(&logits, &mut probs);
135        probs
136    }
137
138    /// Predicted class label (argmax probability).
139    pub fn predict_row(&self, x_row: &[f64]) -> usize {
140        let probs = self.predict_proba_row(x_row);
141        let mut best = 0;
142        for cc in 1..probs.len() {
143            if probs[cc] > probs[best] {
144                best = cc;
145            }
146        }
147        self.classes[best]
148    }
149}
150
151#[cfg(test)]
152mod tests {
153    use super::*;
154
155    #[test]
156    fn classifies_three_separated_clusters() {
157        // class 0 near (0,0), 1 near (10,0), 2 near (0,10).
158        let mut x = Vec::new();
159        let mut y = Vec::new();
160        for &(cx, cy, lbl) in &[(0.0, 0.0, 0usize), (10.0, 0.0, 1), (0.0, 10.0, 2)] {
161            for d in 0..5 {
162                x.push(cx + (d as f64 - 2.0) * 0.1);
163                x.push(cy + (d as f64 - 2.0) * 0.1);
164                y.push(lbl);
165            }
166        }
167        let n = 15;
168        let m = MultinomialLogistic::fit(&x, &y, n, 2, true, 0.5, 1e-4, 500).unwrap();
169        assert_eq!(m.classes, vec![0, 1, 2]);
170        assert_eq!(m.predict_row(&[0.0, 0.0]), 0);
171        assert_eq!(m.predict_row(&[10.0, 0.0]), 1);
172        assert_eq!(m.predict_row(&[0.0, 10.0]), 2);
173        // Probabilities sum to 1.
174        let s: f64 = m.predict_proba_row(&[0.0, 0.0]).iter().sum();
175        assert!((s - 1.0).abs() < 1e-9);
176    }
177
178    #[test]
179    fn guards() {
180        assert_eq!(
181            MultinomialLogistic::fit(&[1.0, 2.0], &[0], 1, 2, true, 0.1, 0.0, 10).unwrap_err(),
182            LearningError::InsufficientData
183        );
184    }
185}