use faer::linalg::matmul::triangular::BlockStructure;
use faer::linalg::matmul::{matmul, triangular};
use faer::reborrow::{IntoConst, Reborrow, ReborrowMut};
use faer::{Accum, MatMut, MatRef, Par};
use crate::ols::nan_fill_ols_scratch;
use crate::spec::{BinomialLink, Family};
use crate::FLOAT_NEAR_ZERO;
pub const MAX_IRLS_ITERS: u32 = 50;
pub const DEVIANCE_TOL: f64 = 1e-8;
pub const BETA_CAP: f64 = 30.0;
pub const WEIGHT_CLAMP: f64 = 1e-6;
pub const SATURATION_W: f64 = 1e-5;
pub const SATURATION_FRAC: f64 = 0.5;
pub struct GlmFitView<'a> {
pub betas: &'a [f64],
pub var_diag: &'a [f64],
pub t_sq: &'a [f64],
pub l: MatRef<'a, f64>,
pub n_iter: u32,
pub converged: bool,
pub deviance: f64,
pub deviance_null: f64,
}
pub struct GlmScratch<'w> {
pub irls_eta: &'w mut [f64],
pub irls_p: &'w mut [f64],
pub irls_w: &'w mut [f64],
pub irls_z: &'w mut [f64],
pub irls_betas: &'w mut [f64],
pub irls_betas_new: &'w mut [f64],
pub irls_var_diag: &'w mut [f64],
pub irls_t_sq: &'w mut [f64],
pub irls_u_scratch: &'w mut [f64],
pub irls_xtwx: MatMut<'w, f64>,
pub irls_xtwz: &'w mut [f64],
pub irls_l: MatMut<'w, f64>,
pub irls_wx: &'w mut [f64],
}
#[cfg_attr(not(test), allow(dead_code))]
#[inline]
pub fn sigmoid_stable(eta: f64) -> f64 {
if eta >= 0.0 {
let z = (-eta).exp();
1.0 / (1.0 + z)
} else {
let z = eta.exp();
z / (1.0 + z)
}
}
#[allow(clippy::too_many_arguments)] pub fn glm_irls_fit<'a>(
family: Family,
nb_theta: f64,
x: MatRef<'_, f64>,
y: &[f64],
target_indices: &[u32],
beta_start: Option<&[f64]>,
prior_w: Option<&[f64]>,
offset: Option<&[f64]>,
scratch: GlmScratch<'a>,
) -> GlmFitView<'a> {
let n = x.nrows();
let p = x.ncols();
let t = target_indices.len();
debug_assert_eq!(n, y.len(), "glm_irls_fit: y length must match X.nrows()");
let GlmScratch {
irls_eta,
irls_p,
irls_w,
irls_z,
irls_betas,
irls_betas_new,
irls_var_diag,
irls_t_sq,
irls_u_scratch,
mut irls_xtwx,
irls_xtwz,
mut irls_l,
irls_wx,
} = scratch;
debug_assert!(p <= irls_betas.len(), "scratch sized for fewer predictors");
debug_assert!(t <= irls_var_diag.len());
debug_assert!(n <= irls_eta.len());
nan_fill_ols_scratch(irls_betas, irls_var_diag, irls_t_sq, p, t);
if n <= p || p == 0 {
return GlmFitView {
betas: &irls_betas[..p],
var_diag: &irls_var_diag[..t],
t_sq: &irls_t_sq[..t],
l: irls_l.into_const(),
n_iter: 0,
converged: false,
deviance: f64::NAN,
deviance_null: f64::NAN,
};
}
let mut y_sum = 0.0;
for &yi in &y[..n] {
y_sum += yi;
}
if matches!(family, Family::Binomial { .. }) && (y_sum <= 0.0 || y_sum >= n as f64) {
return GlmFitView {
betas: &irls_betas[..p],
var_diag: &irls_var_diag[..t],
t_sq: &irls_t_sq[..t],
l: irls_l.into_const(),
n_iter: 0,
converged: false,
deviance: f64::NAN,
deviance_null: f64::NAN,
};
}
let y_bar = match prior_w {
Some(w) => {
let (mut wy_sum, mut w_sum) = (0.0, 0.0);
for (i, &yi) in y[..n].iter().enumerate() {
wy_sum += w[i] * yi;
w_sum += w[i];
}
wy_sum / w_sum
}
None => y_sum / n as f64,
};
let deviance_null = match family {
Family::Binomial {
link: BinomialLink::Logit,
} if prior_w.is_none() => {
-2.0 * (y_sum * y_bar.ln() + (n as f64 - y_sum) * (1.0 - y_bar).ln())
}
other => {
let mu0 = crate::family::clamp_mu(other, y_bar);
let mut d = 0.0;
for (i, &yi) in y[..n].iter().enumerate() {
let pw = prior_w.map_or(1.0, |w| w[i]);
d += pw * crate::family::dev_resid(other, nb_theta, yi, mu0);
}
d
}
};
match beta_start {
Some(b0) => {
debug_assert_eq!(b0.len(), p, "beta_start length must match X.ncols()");
irls_betas[..p].copy_from_slice(&b0[..p]);
irls_eta[..n].fill(0.0);
for j in 0..p {
let b_j = irls_betas[j];
for i in 0..n {
irls_eta[i] += x[(i, j)] * b_j;
}
}
if let Some(o) = offset {
for i in 0..n {
irls_eta[i] += o[i];
}
}
}
None => {
irls_betas[..p].fill(0.0);
match family {
Family::Gamma {
link: crate::spec::GammaLink::Inverse,
..
} => {
for i in 0..n {
irls_eta[i] = 1.0 / crate::family::clamp_mu(family, y[i]);
}
}
Family::Poisson { .. } | Family::NegativeBinomial { .. } => {
let ybar = y[..n].iter().sum::<f64>() / n.max(1) as f64;
irls_eta[..n].fill((ybar + 0.1).ln());
}
_ => irls_eta[..n].fill(0.0),
}
}
}
let mut last_chol = None;
let mut deviance_prev = f64::INFINITY;
let mut deviance_final = f64::NAN;
let mut converged = false;
let mut had_pd_failure = false;
let mut n_iter: u32 = 0;
for iter in 0..=MAX_IRLS_ITERS {
let deviance = match family {
Family::Binomial {
link: BinomialLink::Logit,
} if prior_w.is_none() => {
let lp_sum = crate::simd_transcendental::pw_and_log1pexp_sum(
&irls_eta[..n],
&mut irls_p[..n],
&mut irls_w[..n],
);
let mut yeta = 0.0;
for i in 0..n {
let yi = y[i];
yeta += yi * irls_eta[i];
irls_z[i] = irls_eta[i] + (yi - irls_p[i]) / irls_w[i];
}
if let Some(o) = offset {
for i in 0..n {
irls_z[i] -= o[i];
}
}
2.0 * (lp_sum - yeta)
}
other => {
let mut dev = 0.0;
for i in 0..n {
let e = crate::family::clamp_eta(other, irls_eta[i]);
let (mu, w_raw, r) =
crate::family::irls_weight_and_resid(other, nb_theta, y[i], e);
let pw = prior_w.map_or(1.0, |w| w[i]);
irls_p[i] = mu;
irls_w[i] = (pw * w_raw).max(WEIGHT_CLAMP);
irls_z[i] = e + r;
dev += pw * crate::family::dev_resid(other, nb_theta, y[i], mu);
}
if let Some(o) = offset {
for i in 0..n {
irls_z[i] -= o[i];
}
}
dev
}
};
if iter > 0 {
deviance_final = deviance;
if (deviance - deviance_prev).abs() < DEVIANCE_TOL {
converged = true;
break;
}
deviance_prev = deviance;
}
if iter == MAX_IRLS_ITERS {
break;
}
n_iter = iter + 1;
{
let wx = &mut irls_wx[..n * p];
for j in 0..p {
let wxj = &mut wx[j * n..(j + 1) * n];
for i in 0..n {
wxj[i] = irls_w[i] * x[(i, j)];
}
}
}
let wx_ref = MatRef::from_column_major_slice(&irls_wx[..n * p], n, p);
triangular::matmul(
irls_xtwx.rb_mut(),
BlockStructure::TriangularLower,
Accum::Replace,
x.transpose(),
BlockStructure::Rectangular,
wx_ref,
BlockStructure::Rectangular,
1.0,
Par::Seq,
);
matmul(
MatMut::from_column_major_slice_mut(&mut irls_xtwz[..p], p, 1),
Accum::Replace,
wx_ref.transpose(),
MatRef::from_column_major_slice(&irls_z[..n], n, 1),
1.0,
Par::Seq,
);
let chol = match irls_xtwx.rb().llt(faer::Side::Lower) {
Ok(c) => c,
Err(_) => {
had_pd_failure = true;
break;
}
};
{
use faer::linalg::solvers::Solve;
let mut rhs = MatMut::from_column_major_slice_mut(irls_xtwz, p, 1usize);
chol.solve_in_place(rhs.rb_mut());
}
irls_betas_new[..p].copy_from_slice(&irls_xtwz[..p]);
last_chol = Some(chol);
let mut all_finite = true;
for &b in &irls_betas_new[..p] {
if !b.is_finite() {
all_finite = false;
break;
}
}
if !all_finite {
break;
}
irls_betas[..p].copy_from_slice(&irls_betas_new[..p]);
irls_eta[..n].fill(0.0);
for j in 0..p {
let b_j = irls_betas[j];
for i in 0..n {
irls_eta[i] += x[(i, j)] * b_j;
}
}
if let Some(o) = offset {
for i in 0..n {
irls_eta[i] += o[i];
}
}
if iter >= 3 {
let mut max_abs: f64 = 0.0;
for &b in &irls_betas[..p] {
let ab = b.abs();
if ab > max_abs {
max_abs = ab;
}
}
if max_abs > BETA_CAP {
break;
}
}
}
if had_pd_failure {
converged = false;
}
if converged {
let saturated = (0..n)
.filter(|&i| {
let pw = prior_w.map_or(1.0, |w| w[i]);
irls_w[i] < SATURATION_W * pw
})
.count();
if (saturated as f64) / (n as f64) > SATURATION_FRAC {
converged = false;
}
}
if !converged {
irls_var_diag[..t].fill(f64::NAN);
irls_t_sq[..t].fill(f64::NAN);
return GlmFitView {
betas: &irls_betas[..p],
var_diag: &irls_var_diag[..t],
t_sq: &irls_t_sq[..t],
l: irls_l.into_const(),
n_iter,
converged: false,
deviance: f64::NAN,
deviance_null: f64::NAN,
};
}
if let Some(chol) = last_chol {
let l = chol.L();
for j in 0..p {
for i in 0..p {
irls_l[(i, j)] = if i >= j { l[(i, j)] } else { 0.0 };
}
}
}
for (out_idx, &tj) in target_indices.iter().enumerate() {
let tj = tj as usize;
if tj >= p {
continue;
}
irls_u_scratch[..p].fill(0.0);
for i in 0..p {
let b_i = if i == tj { 1.0 } else { 0.0 };
let mut acc = b_i;
for k in 0..i {
acc -= irls_l[(i, k)] * irls_u_scratch[k];
}
let l_ii = irls_l[(i, i)];
irls_u_scratch[i] = if l_ii.abs() < FLOAT_NEAR_ZERO {
f64::NAN
} else {
acc / l_ii
};
}
let mut norm_sq = 0.0;
for &v in &irls_u_scratch[..p] {
norm_sq += v * v;
}
irls_var_diag[out_idx] = norm_sq;
if norm_sq > FLOAT_NEAR_ZERO && norm_sq.is_finite() {
let beta_j = irls_betas[tj];
irls_t_sq[out_idx] = (beta_j * beta_j) / norm_sq;
} else {
irls_t_sq[out_idx] = f64::NAN;
}
}
GlmFitView {
betas: &irls_betas[..p],
var_diag: &irls_var_diag[..t],
t_sq: &irls_t_sq[..t],
l: irls_l.into_const(),
n_iter,
converged: true,
deviance: deviance_final,
deviance_null,
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::test_support::TestWs;
use faer::Mat;
fn glm_scratch(ws: &mut TestWs) -> GlmScratch<'_> {
GlmScratch {
irls_eta: &mut ws.irls_eta,
irls_p: &mut ws.irls_p,
irls_w: &mut ws.irls_w,
irls_z: &mut ws.irls_z,
irls_betas: &mut ws.irls_betas,
irls_betas_new: &mut ws.irls_betas_new,
irls_var_diag: &mut ws.irls_var_diag,
irls_t_sq: &mut ws.irls_t_sq,
irls_u_scratch: &mut ws.irls_u_scratch,
irls_xtwx: ws.irls_xtwx.as_mut(),
irls_xtwz: &mut ws.irls_xtwz,
irls_l: ws.irls_l.as_mut(),
irls_wx: &mut ws.irls_wx,
}
}
#[test]
fn glm_all_zero_y_short_circuits() {
let n = 100;
let p = 2;
let mut x = Mat::<f64>::zeros(n, p);
for i in 0..n {
x[(i, 0)] = 1.0;
x[(i, 1)] = (i as f64) / (n as f64) - 0.5;
}
let y = vec![0.0f64; n];
let mut ws = TestWs::new(n, p, 0);
let targets: Vec<u32> = vec![0, 1];
let fit = glm_irls_fit(
crate::Family::Binomial {
link: crate::BinomialLink::Logit,
},
f64::NAN,
x.as_ref(),
&y,
&targets,
None,
None,
None,
glm_scratch(&mut ws),
);
assert!(!fit.converged);
assert_eq!(fit.n_iter, 0);
}
#[test]
fn glm_all_one_y_short_circuits() {
let n = 100;
let p = 2;
let mut x = Mat::<f64>::zeros(n, p);
for i in 0..n {
x[(i, 0)] = 1.0;
x[(i, 1)] = (i as f64) / (n as f64) - 0.5;
}
let y = vec![1.0f64; n];
let mut ws = TestWs::new(n, p, 0);
let targets: Vec<u32> = vec![0, 1];
let fit = glm_irls_fit(
crate::Family::Binomial {
link: crate::BinomialLink::Logit,
},
f64::NAN,
x.as_ref(),
&y,
&targets,
None,
None,
None,
glm_scratch(&mut ws),
);
assert!(!fit.converged);
assert_eq!(fit.n_iter, 0);
}
#[test]
fn glm_rank_deficient_design() {
let n = 100;
let p = 3;
let mut x = Mat::<f64>::zeros(n, p);
for i in 0..n {
x[(i, 0)] = 1.0;
x[(i, 1)] = ((i as f64) / (n as f64)) - 0.5;
x[(i, 2)] = x[(i, 1)];
}
let y: Vec<f64> = (0..n).map(|i| if i % 2 == 0 { 1.0 } else { 0.0 }).collect();
let mut ws = TestWs::new(n, p, 0);
let targets: Vec<u32> = vec![1, 2];
let fit = glm_irls_fit(
crate::Family::Binomial {
link: crate::BinomialLink::Logit,
},
f64::NAN,
x.as_ref(),
&y,
&targets,
None,
None,
None,
glm_scratch(&mut ws),
);
assert!(
!fit.converged,
"rank-deficient design must report non-converged"
);
}
#[test]
fn glm_z_sq_nan_on_non_converged() {
let n = 100;
let p = 2;
let mut x = Mat::<f64>::zeros(n, p);
for i in 0..n {
x[(i, 0)] = 1.0;
x[(i, 1)] = (i as f64) / (n as f64) - 0.5;
}
let y = vec![0.0f64; n];
let mut ws = TestWs::new(n, p, 0);
let targets: Vec<u32> = vec![0, 1];
let fit = glm_irls_fit(
crate::Family::Binomial {
link: crate::BinomialLink::Logit,
},
f64::NAN,
x.as_ref(),
&y,
&targets,
None,
None,
None,
glm_scratch(&mut ws),
);
assert!(!fit.converged);
for &t in fit.t_sq.iter() {
assert!(t.is_nan(), "z² must be NaN on non-converged fit, got {t}");
}
}
#[test]
fn glm_deviance_nan_on_non_converged() {
let n = 100;
let p = 2;
let mut x = Mat::<f64>::zeros(n, p);
for i in 0..n {
x[(i, 0)] = 1.0;
x[(i, 1)] = (i as f64) / (n as f64) - 0.5;
}
let y = vec![0.0f64; n];
let mut ws = TestWs::new(n, p, 0);
let targets: Vec<u32> = vec![0, 1];
let fit = glm_irls_fit(
crate::Family::Binomial {
link: crate::BinomialLink::Logit,
},
f64::NAN,
x.as_ref(),
&y,
&targets,
None,
None,
None,
glm_scratch(&mut ws),
);
assert!(!fit.converged);
assert!(
fit.deviance.is_nan(),
"deviance must be NaN on non-converged"
);
assert!(
fit.deviance_null.is_nan(),
"deviance_null must be NaN on all-0 short-circuit (sum_y = 0)"
);
}
#[test]
fn glm_separation_marks_non_converged() {
let n = 200;
let p = 2;
let mut x = Mat::<f64>::zeros(n, p);
let mut y = vec![0.0f64; n];
for i in 0..n {
x[(i, 0)] = 1.0;
let xi = (i as f64 - n as f64 / 2.0) / 10.0; x[(i, 1)] = xi;
y[i] = if xi > 0.0 { 1.0 } else { 0.0 };
}
let mut ws = TestWs::new(n, p, 0);
let targets: Vec<u32> = vec![1];
let fit = glm_irls_fit(
crate::Family::Binomial {
link: crate::BinomialLink::Logit,
},
f64::NAN,
x.as_ref(),
&y,
&targets,
None,
None,
None,
glm_scratch(&mut ws),
);
assert!(
!fit.converged,
"fully separated data must report non-converged"
);
}
#[test]
fn glm_deviance_null_golden_value() {
let n = 100;
let p = 2;
let mut x = Mat::<f64>::zeros(n, p);
for i in 0..n {
x[(i, 0)] = 1.0;
x[(i, 1)] = (i as f64) / (n as f64) - 0.5; }
let mut y = vec![0.0f64; n];
for (i, v) in y.iter_mut().enumerate() {
if i % 5 < 2 {
*v = 1.0;
}
}
let mut ws = TestWs::new(n, p, 0);
let targets: Vec<u32> = vec![1];
let fit = glm_irls_fit(
crate::Family::Binomial {
link: crate::BinomialLink::Logit,
},
f64::NAN,
x.as_ref(),
&y,
&targets,
None,
None,
None,
glm_scratch(&mut ws),
);
assert!(fit.converged, "must converge: non-separated y pattern");
let expected = 134.6023334f64;
let abs_err = (fit.deviance_null - expected).abs();
assert!(
abs_err < 0.001,
"deviance_null = {}, expected {expected}, err = {abs_err}",
fit.deviance_null
);
}
#[test]
fn glm_weighted_deviance_null_golden_value() {
let x1: [f64; 40] = [
1.371, -0.5647, 0.3631, 0.6329, 0.4043, -0.1061, 1.5115, -0.0947, 2.0184, -0.0627,
1.3049, 2.2866, -1.3889, -0.2788, -0.1333, 0.636, -0.2843, -2.6565, -2.4405, 1.3201,
-0.3066, -1.7813, -0.1719, 1.2147, 1.8952, -0.4305, -0.2573, -1.7632, 0.4601, -0.64,
0.4555, 0.7048, 1.0351, -0.6089, 0.505, -1.717, -0.7845, -0.8509, -2.4142, 0.0361,
];
let w: [f64; 40] = [
4.0, 1.0, 2.0, 1.0, 1.0, 4.0, 4.0, 1.0, 3.0, 3.0, 1.0, 4.0, 1.0, 4.0, 4.0, 2.0, 1.0,
4.0, 2.0, 2.0, 2.0, 4.0, 1.0, 2.0, 1.0, 2.0, 4.0, 3.0, 4.0, 1.0, 4.0, 1.0, 4.0, 3.0,
2.0, 2.0, 3.0, 1.0, 1.0, 2.0,
];
let y: [f64; 40] = [
2.421196, 0.850101, 1.188318, 0.917668, 1.895064, 2.717167, 4.391082, 0.266883,
1.853922, 1.838375, 5.959549, 19.008523, 0.121882, 1.544704, 1.422566, 0.758422,
1.264496, 0.147806, 0.06751, 2.907132, 0.3538, 0.223494, 0.297625, 5.273375, 12.534684,
0.514577, 1.473477, 0.485665, 0.962023, 1.043896, 1.771311, 1.926229, 7.592099,
1.298714, 0.675125, 0.201756, 1.814679, 1.104297, 0.434436, 0.470596,
];
const REF_DEV_NULL: f64 = 140.09428224081;
const REF_DEV: f64 = 39.1211374203115;
let n = 40;
let p = 2;
let mut x = Mat::<f64>::zeros(n, p);
for i in 0..n {
x[(i, 0)] = 1.0;
x[(i, 1)] = x1[i];
}
let mut ws = TestWs::new(n, p, 0);
let targets: Vec<u32> = vec![0, 1];
let fit = glm_irls_fit(
crate::Family::Gamma {
link: crate::GammaLink::Log,
},
f64::NAN,
x.as_ref(),
&y,
&targets,
None,
Some(&w),
None,
glm_scratch(&mut ws),
);
assert!(fit.converged, "weighted gamma GLM must converge");
assert!(
(fit.deviance_null - REF_DEV_NULL).abs() / REF_DEV_NULL < 1e-8,
"deviance_null = {} vs R {REF_DEV_NULL}",
fit.deviance_null
);
assert!(
(fit.deviance - REF_DEV).abs() / REF_DEV < 1e-6,
"deviance = {} vs R {REF_DEV}",
fit.deviance
);
}
}