qualia_core_db/solvers/learning/glm/
mod.rs1pub mod family;
11pub mod multinomial;
12
13pub use family::Family;
14pub use multinomial::MultinomialLogistic;
15
16use crate::solvers::learning::LearningError;
17use crate::solvers::linear_algebra::cholesky::{cholesky_factor, cholesky_solve};
18use crate::solvers::linear_algebra::gemm::{matvec, Transpose};
19use crate::solvers::statistics::distributions::normal;
20
21#[derive(Debug, Clone)]
23pub struct GlmModel {
24 pub family: Family,
25 pub coefficients: Vec<f64>,
26 pub fit_intercept: bool,
27 pub std_errors: Vec<f64>,
28 pub z_values: Vec<f64>,
30 pub p_values: Vec<f64>,
32 pub n_iter: usize,
33 pub converged: bool,
34 pub deviance: f64,
36 pub n: usize,
37}
38
39impl GlmModel {
40 pub fn eta_row(&self, x_row: &[f64]) -> f64 {
42 let (b0, betas) = if self.fit_intercept {
43 (self.coefficients[0], &self.coefficients[1..])
44 } else {
45 (0.0, &self.coefficients[..])
46 };
47 b0 + betas.iter().zip(x_row).map(|(b, x)| b * x).sum::<f64>()
48 }
49
50 pub fn predict_row(&self, x_row: &[f64]) -> f64 {
53 self.family.inv_link(self.eta_row(x_row))
54 }
55
56 pub fn predict(&self, x: &[f64], n: usize, p: usize) -> Vec<f64> {
58 (0..n)
59 .map(|i| self.predict_row(&x[i * p..(i + 1) * p]))
60 .collect()
61 }
62}
63
64const MAX_ITER: usize = 100;
65const TOL: f64 = 1e-10;
66
67pub fn fit(
72 family: Family,
73 x: &[f64],
74 y: &[f64],
75 n: usize,
76 p: usize,
77 fit_intercept: bool,
78) -> Result<GlmModel, LearningError> {
79 if n == 0 || p == 0 || x.len() != n * p || y.len() != n {
80 return Err(LearningError::InvalidDimension);
81 }
82 let k = p + usize::from(fit_intercept);
83 if n <= k {
84 return Err(LearningError::InsufficientData);
85 }
86
87 let mut d = vec![0.0; n * k];
89 for i in 0..n {
90 let base = i * k;
91 if fit_intercept {
92 d[base] = 1.0;
93 d[base + 1..base + k].copy_from_slice(&x[i * p..(i + 1) * p]);
94 } else {
95 d[base..base + k].copy_from_slice(&x[i * p..(i + 1) * p]);
96 }
97 }
98
99 let mut beta = vec![0.0; k];
100 let mut eta: Vec<f64> = y
102 .iter()
103 .map(|&yi| {
104 let mu = family.start_mu(yi);
105 match family {
107 Family::Binomial => (mu / (1.0 - mu)).ln(),
108 Family::Poisson => mu.ln(),
109 }
110 })
111 .collect();
112
113 let mut a = vec![0.0; k * k];
114 let mut b = vec![0.0; k];
115 let mut l = vec![0.0; k * k];
116 let mut converged = false;
117 let mut iters = 0;
118
119 let mut w = vec![0.0; n];
120 let mut z = vec![0.0; n];
121
122 for it in 1..=MAX_ITER {
123 iters = it;
124 for i in 0..n {
126 let mu = family.inv_link(eta[i]);
127 let dmu = family.dmu_deta(mu).max(1e-12);
128 let var = family.variance(mu).max(1e-12);
129 w[i] = dmu * dmu / var;
130 z[i] = eta[i] + (y[i] - mu) / dmu;
131 }
132 for r in 0..k {
134 for c in 0..k {
135 let mut s = 0.0;
136 for i in 0..n {
137 s += d[i * k + r] * w[i] * d[i * k + c];
138 }
139 a[r * k + c] = s;
140 }
141 let mut s = 0.0;
142 for i in 0..n {
143 s += d[i * k + r] * w[i] * z[i];
144 }
145 b[r] = s;
146 }
147 cholesky_factor(k, &a, &mut l).map_err(|_| LearningError::Singular)?;
148 let mut beta_new = vec![0.0; k];
149 cholesky_solve(k, &l, &b, &mut beta_new)?;
150
151 let mut delta = 0.0;
153 for j in 0..k {
154 delta += (beta_new[j] - beta[j]).powi(2);
155 }
156 beta = beta_new;
157 matvec(Transpose::No, n, k, &d, &beta, &mut eta)?;
158 if delta.sqrt() < TOL {
159 converged = true;
160 break;
161 }
162 }
163 if !converged {
164 return Err(LearningError::NotConverged);
165 }
166
167 let mut std_errors = vec![0.0; k];
170 let mut z_values = vec![0.0; k];
171 let mut p_values = vec![0.0; k];
172 let mut ej = vec![0.0; k];
173 let mut cj = vec![0.0; k];
174 for j in 0..k {
175 ej.iter_mut().for_each(|v| *v = 0.0);
176 ej[j] = 1.0;
177 cholesky_solve(k, &l, &ej, &mut cj)?;
178 let se = if cj[j] > 0.0 { cj[j].sqrt() } else { 0.0 };
179 std_errors[j] = se;
180 if se > 0.0 {
181 let zv = beta[j] / se;
182 z_values[j] = zv;
183 p_values[j] = normal::two_sided_p(zv);
184 } else {
185 z_values[j] = if beta[j] == 0.0 { 0.0 } else { f64::INFINITY };
186 p_values[j] = if beta[j] == 0.0 { 1.0 } else { 0.0 };
187 }
188 }
189
190 let mut deviance = 0.0;
191 for i in 0..n {
192 deviance += family.unit_deviance(y[i], family.inv_link(eta[i]));
193 }
194
195 Ok(GlmModel {
196 family,
197 coefficients: beta,
198 fit_intercept,
199 std_errors,
200 z_values,
201 p_values,
202 n_iter: iters,
203 converged,
204 deviance,
205 n,
206 })
207}
208
209pub fn fit_logistic(
211 x: &[f64],
212 y: &[f64],
213 n: usize,
214 p: usize,
215 fit_intercept: bool,
216) -> Result<GlmModel, LearningError> {
217 fit(Family::Binomial, x, y, n, p, fit_intercept)
218}
219
220pub fn fit_poisson(
222 x: &[f64],
223 y: &[f64],
224 n: usize,
225 p: usize,
226 fit_intercept: bool,
227) -> Result<GlmModel, LearningError> {
228 fit(Family::Poisson, x, y, n, p, fit_intercept)
229}
230
231#[cfg(test)]
232mod tests {
233 use super::*;
234
235 #[test]
236 fn logistic_recovers_positive_trend() {
237 let x = [1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0];
240 let y = [0.0, 0.0, 0.0, 1.0, 0.0, 1.0, 0.0, 1.0, 1.0, 1.0];
241 let m = fit_logistic(&x, &y, 10, 1, true).unwrap();
242 assert!(m.converged);
243 assert!(
244 m.coefficients[1] > 0.0,
245 "slope should be positive: {}",
246 m.coefficients[1]
247 );
248 assert!(m.predict_row(&[10.0]) > m.predict_row(&[1.0]));
250 assert!(m.predict_row(&[10.0]) > 0.5 && m.predict_row(&[1.0]) < 0.5);
251 assert!(m.deviance.is_finite());
252 }
253
254 #[test]
255 fn logistic_significance_and_inference() {
256 let x = [
258 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0,
259 ];
260 let y = [0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0, 1.0];
261 let m = fit_logistic(&x, &y, 12, 1, true).unwrap();
262 assert!(m.coefficients[1] > 0.0);
263 assert!(
264 m.p_values[1] < 0.1,
265 "trend should be significant: p={}",
266 m.p_values[1]
267 );
268 assert!(m.std_errors[1] > 0.0 && m.std_errors[1].is_finite());
269 }
270
271 #[test]
272 fn poisson_recovers_log_linear_rate() {
273 let x: [f64; 8] = [0.0, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0];
275 let y: Vec<f64> = x.iter().map(|&xi| (0.5 + 0.3 * xi).exp().round()).collect();
276 let m = fit_poisson(&x, &y, 8, 1, true).unwrap();
277 assert!(m.converged);
278 assert!(
279 (m.coefficients[1] - 0.3).abs() < 0.1,
280 "slope {}",
281 m.coefficients[1]
282 );
283 assert!(
284 (m.coefficients[0] - 0.5).abs() < 0.2,
285 "intercept {}",
286 m.coefficients[0]
287 );
288 }
289
290 #[test]
291 fn guards() {
292 assert_eq!(
293 fit_logistic(&[1.0, 2.0], &[1.0], 2, 1, true).unwrap_err(),
294 LearningError::InvalidDimension
295 );
296 assert_eq!(
297 fit_logistic(&[1.0, 2.0], &[1.0, 0.0], 2, 1, true).unwrap_err(),
298 LearningError::InsufficientData
299 );
300 }
301}