qualia_core_db/solvers/learning/classification/
discriminant.rs1use crate::solvers::learning::LearningError;
7use crate::solvers::linear_algebra::cholesky::{
8 cholesky_determinant, cholesky_factor, cholesky_solve,
9};
10
11struct ClassStats {
13 classes: Vec<usize>,
14 priors: Vec<f64>,
15 means: Vec<f64>, counts: Vec<usize>,
17}
18
19fn class_stats(x: &[f64], y: &[usize], n: usize, p: usize) -> Result<ClassStats, LearningError> {
20 if n == 0 || p == 0 || x.len() != n * p || y.len() != n {
21 return Err(LearningError::InvalidDimension);
22 }
23 let mut classes: Vec<usize> = y.to_vec();
24 classes.sort_unstable();
25 classes.dedup();
26 let c = classes.len();
27 let mut priors = vec![0.0; c];
28 let mut means = vec![0.0; c * p];
29 let mut counts = vec![0usize; c];
30 let cls_idx = |cls: usize| classes.iter().position(|&v| v == cls).unwrap();
31 for i in 0..n {
32 let ci = cls_idx(y[i]);
33 counts[ci] += 1;
34 for j in 0..p {
35 means[ci * p + j] += x[i * p + j];
36 }
37 }
38 for ci in 0..c {
39 if counts[ci] == 0 {
40 return Err(LearningError::InsufficientData);
41 }
42 priors[ci] = counts[ci] as f64 / n as f64;
43 for j in 0..p {
44 means[ci * p + j] /= counts[ci] as f64;
45 }
46 }
47 Ok(ClassStats {
48 classes,
49 priors,
50 means,
51 counts,
52 })
53}
54
55#[derive(Debug, Clone)]
60pub struct LdaModel {
61 pub classes: Vec<usize>,
62 w: Vec<f64>, b: Vec<f64>, p: usize,
65}
66
67impl LdaModel {
68 pub fn fit(x: &[f64], y: &[usize], n: usize, p: usize) -> Result<Self, LearningError> {
70 let cs = class_stats(x, y, n, p)?;
71 let c = cs.classes.len();
72 if n <= c {
73 return Err(LearningError::InsufficientData);
74 }
75 let mut cov = vec![0.0; p * p];
77 let cls_idx = |cls: usize| cs.classes.iter().position(|&v| v == cls).unwrap();
78 for i in 0..n {
79 let ci = cls_idx(y[i]);
80 for a in 0..p {
81 let da = x[i * p + a] - cs.means[ci * p + a];
82 for bb in 0..p {
83 let db = x[i * p + bb] - cs.means[ci * p + bb];
84 cov[a * p + bb] += da * db;
85 }
86 }
87 }
88 let denom = (n - c) as f64;
89 for v in cov.iter_mut() {
90 *v /= denom;
91 }
92 let l = cholesky_of(&cov, p)?;
93
94 let mut w = vec![0.0; c * p];
95 let mut b = vec![0.0; c];
96 let mut mu = vec![0.0; p];
97 let mut wc = vec![0.0; p];
98 for ci in 0..c {
99 mu.copy_from_slice(&cs.means[ci * p..(ci + 1) * p]);
100 cholesky_solve(p, &l, &mu, &mut wc)?; w[ci * p..(ci + 1) * p].copy_from_slice(&wc);
102 let quad: f64 = mu.iter().zip(wc.iter()).map(|(m, w)| m * w).sum();
103 b[ci] = -0.5 * quad + cs.priors[ci].ln();
104 }
105 Ok(Self {
106 classes: cs.classes,
107 w,
108 b,
109 p,
110 })
111 }
112
113 pub fn predict_row(&self, q: &[f64]) -> usize {
114 let mut best = 0;
115 let mut best_s = f64::NEG_INFINITY;
116 for ci in 0..self.classes.len() {
117 let s: f64 = q
118 .iter()
119 .zip(&self.w[ci * self.p..(ci + 1) * self.p])
120 .map(|(x, w)| x * w)
121 .sum::<f64>()
122 + self.b[ci];
123 if s > best_s {
124 best_s = s;
125 best = ci;
126 }
127 }
128 self.classes[best]
129 }
130
131 pub fn predict(&self, x: &[f64], m: usize) -> Vec<usize> {
132 (0..m)
133 .map(|i| self.predict_row(&x[i * self.p..(i + 1) * self.p]))
134 .collect()
135 }
136}
137
138#[derive(Debug, Clone)]
143pub struct QdaModel {
144 pub classes: Vec<usize>,
145 means: Vec<f64>, chol: Vec<f64>, log_det: Vec<f64>, log_prior: Vec<f64>,
149 p: usize,
150}
151
152impl QdaModel {
153 pub fn fit(x: &[f64], y: &[usize], n: usize, p: usize) -> Result<Self, LearningError> {
155 let cs = class_stats(x, y, n, p)?;
156 let c = cs.classes.len();
157 let cls_idx = |cls: usize| cs.classes.iter().position(|&v| v == cls).unwrap();
158 let mut chol = vec![0.0; c * p * p];
159 let mut log_det = vec![0.0; c];
160 let mut log_prior = vec![0.0; c];
161 for ci in 0..c {
162 if cs.counts[ci] <= p {
163 return Err(LearningError::InsufficientData); }
165 let mut cov = vec![0.0; p * p];
167 for i in 0..n {
168 if cls_idx(y[i]) != ci {
169 continue;
170 }
171 for a in 0..p {
172 let da = x[i * p + a] - cs.means[ci * p + a];
173 for bb in 0..p {
174 let db = x[i * p + bb] - cs.means[ci * p + bb];
175 cov[a * p + bb] += da * db;
176 }
177 }
178 }
179 let denom = (cs.counts[ci] - 1) as f64;
180 for v in cov.iter_mut() {
181 *v /= denom;
182 }
183 let l = cholesky_of(&cov, p)?;
184 log_det[ci] = cholesky_determinant(p, &l).max(1e-300).ln();
185 chol[ci * p * p..(ci + 1) * p * p].copy_from_slice(&l);
186 log_prior[ci] = cs.priors[ci].ln();
187 }
188 Ok(Self {
189 classes: cs.classes,
190 means: cs.means,
191 chol,
192 log_det,
193 log_prior,
194 p,
195 })
196 }
197
198 pub fn predict_row(&self, q: &[f64]) -> usize {
199 let p = self.p;
200 let mut diff = vec![0.0; p];
201 let mut sol = vec![0.0; p];
202 let mut best = 0;
203 let mut best_s = f64::NEG_INFINITY;
204 for ci in 0..self.classes.len() {
205 for j in 0..p {
206 diff[j] = q[j] - self.means[ci * p + j];
207 }
208 let l = &self.chol[ci * p * p..(ci + 1) * p * p];
210 if cholesky_solve(p, l, &diff, &mut sol).is_err() {
211 continue;
212 }
213 let maha: f64 = diff.iter().zip(sol.iter()).map(|(d, s)| d * s).sum();
214 let s = -0.5 * self.log_det[ci] - 0.5 * maha + self.log_prior[ci];
215 if s > best_s {
216 best_s = s;
217 best = ci;
218 }
219 }
220 self.classes[best]
221 }
222
223 pub fn predict(&self, x: &[f64], m: usize) -> Vec<usize> {
224 (0..m)
225 .map(|i| self.predict_row(&x[i * self.p..(i + 1) * self.p]))
226 .collect()
227 }
228}
229
230fn cholesky_of(cov: &[f64], p: usize) -> Result<Vec<f64>, LearningError> {
232 let mut l = vec![0.0; p * p];
233 cholesky_factor(p, cov, &mut l).map_err(|_| LearningError::Singular)?;
234 Ok(l)
235}
236
237#[cfg(test)]
238mod tests {
239 use super::*;
240
241 fn data() -> (Vec<f64>, Vec<usize>, usize) {
243 let mut x = Vec::new();
244 let mut y = Vec::new();
245 for &(a, b) in &[
247 (0.0, 0.0),
248 (0.5, -0.3),
249 (-0.4, 0.2),
250 (0.2, 0.4),
251 (-0.3, -0.2),
252 ] {
253 x.push(a);
254 x.push(b);
255 y.push(0);
256 }
257 for &(a, b) in &[(5.0, 5.0), (5.4, 4.7), (4.6, 5.3), (5.1, 4.9), (4.8, 5.2)] {
259 x.push(a);
260 x.push(b);
261 y.push(1);
262 }
263 (x, y, 10)
264 }
265
266 #[test]
267 fn lda_separates_classes() {
268 let (x, y, n) = data();
269 let m = LdaModel::fit(&x, &y, n, 2).unwrap();
270 assert_eq!(m.classes, vec![0, 1]);
271 assert_eq!(m.predict_row(&[0.1, 0.1]), 0);
272 assert_eq!(m.predict_row(&[5.0, 5.0]), 1);
273 let preds = m.predict(&x, n);
275 assert!(preds.iter().zip(&y).all(|(a, b)| a == b));
276 }
277
278 #[test]
279 fn qda_separates_classes() {
280 let (x, y, n) = data();
281 let m = QdaModel::fit(&x, &y, n, 2).unwrap();
282 assert_eq!(m.predict_row(&[0.0, 0.0]), 0);
283 assert_eq!(m.predict_row(&[5.2, 4.8]), 1);
284 let preds = m.predict(&x, n);
285 assert!(preds.iter().zip(&y).all(|(a, b)| a == b));
286 }
287
288 #[test]
289 fn qda_fails_closed_on_too_few_per_class() {
290 let x = [0.0, 0.0, 1.0, 0.0, 0.0, 1.0, 5.0, 5.0, 6.0, 5.0];
293 let y = [0, 0, 0, 1, 1];
294 assert_eq!(
295 QdaModel::fit(&x, &y, 5, 2).unwrap_err(),
296 LearningError::InsufficientData
297 );
298 }
299}