Skip to main content

qualia_core_db/solvers/learning/glm/
mod.rs

1//! Generalized linear models (ISL ch 4) — logistic and Poisson regression by
2//! iteratively reweighted least squares (IRLS).
3//!
4//! Each IRLS step is a weighted least-squares solve of `(DᵀWD)β = DᵀWz`, done with
5//! the engine's `linear_algebra::cholesky` (no re-implemented solver); the Wald
6//! standard errors come from `(DᵀWD)⁻¹` at convergence and the p-values from the
7//! Normal CDF in `statistics::distributions`. Kernel-class: `DenseLinear` per step
8//! (dispatch-ready); the IRLS loop itself is scalar CPU.
9
10pub 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/// A fitted GLM. `coefficients[0]` is the intercept when `fit_intercept`.
22#[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    /// Wald z-statistics (coefficient / std error).
29    pub z_values: Vec<f64>,
30    /// Two-sided Wald p-values from the Normal CDF.
31    pub p_values: Vec<f64>,
32    pub n_iter: usize,
33    pub converged: bool,
34    /// Residual deviance `Σ unit_deviance(yᵢ, μ̂ᵢ)`.
35    pub deviance: f64,
36    pub n: usize,
37}
38
39impl GlmModel {
40    /// Linear predictor `η` for one predictor row (length `p`).
41    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    /// Predicted mean `μ = g⁻¹(η)` for one predictor row — a probability for
51    /// logistic, an expected count for Poisson.
52    pub fn predict_row(&self, x_row: &[f64]) -> f64 {
53        self.family.inv_link(self.eta_row(x_row))
54    }
55
56    /// Predicted means for a row-major `n × p` matrix.
57    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
67/// Fit a GLM of `y` (length `n`) on a row-major `n × p` predictor matrix by IRLS.
68/// `fit_intercept` prepends a constant column. Fails closed: `InvalidDimension`,
69/// `InsufficientData` (`n ≤ params`), `Singular` (collinear / perfectly separated),
70/// `NotConverged`.
71pub 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    // Design matrix D (n × k).
88    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    // Initialise η from a safe starting mean.
101    let mut eta: Vec<f64> = y
102        .iter()
103        .map(|&yi| {
104            let mu = family.start_mu(yi);
105            // η₀ = link(μ₀); use a tiny IRLS-friendly init via inverse of inv_link.
106            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        // Working weights and response.
125        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        // A = DᵀWD, b = DᵀWz.
133        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        // Update η and check convergence on the coefficient change.
152        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    // Wald standard errors from (DᵀWD)⁻¹ at the final weights (a, l already hold
168    // the last factor). Deviance from the fitted means.
169    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
209/// Convenience: logistic regression (Bernoulli `y ∈ {0,1}`).
210pub 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
220/// Convenience: Poisson regression (count `y ≥ 0`).
221pub 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        // Non-separable (so the MLE is finite) but with a clear upward trend:
238        // mostly 0 at low x, mostly 1 at high x, with overlap in the middle.
239        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        // Predicted probability is higher for a large x than a small one.
249        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        // A stronger, larger non-separable signal → significant positive slope.
257        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        // y ≈ exp(0.5 + 0.3 x); fit should recover a positive slope near 0.3.
274        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}