qualia_core_db/solvers/learning/glm/
multinomial.rs1use crate::solvers::learning::LearningError;
9
10#[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
29fn 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 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 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 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 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 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 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 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}