use crate::spec::{BinomialLink, Family, GammaLink};
const ETA_MAX: f64 = 700.0;
const MU_FLOOR: f64 = 1e-10;
const PROB_EPS: f64 = 1e-12;
const FRAC_1_SQRT_2PI: f64 = 0.398_942_280_401_432_7;
pub(crate) fn clamp_eta(family: Family, eta: f64) -> f64 {
match family {
Family::Gamma {
link: GammaLink::Inverse,
..
} => eta.clamp(MU_FLOOR, ETA_MAX),
Family::Poisson { .. }
| Family::Gamma {
link: GammaLink::Log,
..
}
| Family::NegativeBinomial { .. } => eta.clamp(-ETA_MAX, ETA_MAX),
Family::Binomial { .. } | Family::Gaussian => eta,
}
}
pub(crate) fn clamp_mu(family: Family, mu: f64) -> f64 {
match family {
Family::Binomial { .. } => mu.clamp(PROB_EPS, 1.0 - PROB_EPS),
Family::Poisson { .. } | Family::Gamma { .. } | Family::NegativeBinomial { .. } => {
mu.max(MU_FLOOR)
}
Family::Gaussian => mu,
}
}
pub(crate) fn link_inv(family: Family, eta: f64) -> f64 {
let eta = clamp_eta(family, eta);
let mu = match family {
Family::Gaussian => eta,
Family::Binomial {
link: BinomialLink::Logit,
} => crate::glm::sigmoid_stable(eta),
Family::Binomial {
link: BinomialLink::Probit,
} => crate::simd_transcendental::phi_hp(eta),
Family::Poisson { .. }
| Family::Gamma {
link: GammaLink::Log,
..
}
| Family::NegativeBinomial { .. } => eta.exp(),
Family::Gamma {
link: GammaLink::Inverse,
..
} => 1.0 / eta,
};
clamp_mu(family, mu)
}
pub(crate) fn mu_eta(family: Family, eta: f64) -> f64 {
let eta = clamp_eta(family, eta);
match family {
Family::Gaussian => 1.0,
Family::Binomial {
link: BinomialLink::Logit,
} => {
let mu = crate::glm::sigmoid_stable(eta);
mu * (1.0 - mu)
}
Family::Binomial {
link: BinomialLink::Probit,
} => FRAC_1_SQRT_2PI * (-0.5 * eta * eta).exp(),
Family::Poisson { .. }
| Family::Gamma {
link: GammaLink::Log,
..
}
| Family::NegativeBinomial { .. } => eta.exp(),
Family::Gamma {
link: GammaLink::Inverse,
..
} => {
let mu = 1.0 / eta;
-mu * mu
}
}
}
pub(crate) fn variance(family: Family, nb_theta: f64, mu: f64) -> f64 {
match family {
Family::Gaussian => 1.0,
Family::Binomial { .. } => mu * (1.0 - mu),
Family::Poisson { .. } => mu,
Family::Gamma { .. } => mu * mu,
Family::NegativeBinomial { .. } => mu + mu * mu / nb_theta,
}
}
pub(crate) fn dev_resid(family: Family, nb_theta: f64, y: f64, mu: f64) -> f64 {
match family {
Family::Gaussian => {
let r = y - mu;
r * r
}
Family::Binomial { .. } => {
let a = if y > 0.0 { y * (y / mu).ln() } else { 0.0 };
let b = if y < 1.0 {
(1.0 - y) * ((1.0 - y) / (1.0 - mu)).ln()
} else {
0.0
};
2.0 * (a + b)
}
Family::Poisson { .. } => {
let t = if y > 0.0 { y * (y / mu).ln() } else { 0.0 };
2.0 * (t - (y - mu))
}
Family::Gamma { .. } => 2.0 * (-(y / mu).ln() + (y - mu) / mu),
Family::NegativeBinomial { .. } => {
let t = if y > 0.0 { y * (y / mu).ln() } else { 0.0 };
2.0 * (t - (y + nb_theta) * ((y + nb_theta) / (mu + nb_theta)).ln())
}
}
}
pub(crate) fn gamma_aic(y: &[f64], mu: &[f64], dev: f64, n: usize, prior_w: Option<&[f64]>) -> f64 {
let sum_w = prior_w.map_or(n as f64, |w| w[..n].iter().sum());
let disp = dev / sum_w;
let a = 1.0 / disp; let ln_gamma_a = crate::simd_transcendental::ln_gamma(a);
let mut s = 0.0;
for (i, (&yi, &mui)) in y.iter().zip(mu).take(n).enumerate() {
let scale = mui * disp; s += prior_w.map_or(1.0, |w| w[i])
* ((a - 1.0) * yi.ln() - yi / scale - a * scale.ln() - ln_gamma_a);
}
-2.0 * s + 2.0
}
pub(crate) fn glmm_sigma_sq(
family: Family,
y: &[f64],
mu: &[f64],
u: &[f64],
prior_w: Option<&[f64]>,
) -> f64 {
match family {
Family::Gamma { .. } => {
let mut wrss = 0.0;
for (i, (&yi, &mui)) in y.iter().zip(mu).enumerate() {
let r = (yi - mui) / mui;
wrss += prior_w.map_or(1.0, |w| w[i]) * r * r;
}
let usq: f64 = u.iter().map(|&v| v * v).sum();
(wrss + usq) / y.len() as f64
}
_ => 1.0,
}
}
pub(crate) fn pearson_dispersion(
y: &[f64],
mu: &[f64],
family: Family,
nb_theta: f64,
n: usize,
p: usize,
prior_w: Option<&[f64]>,
) -> f64 {
let mut s = 0.0;
for i in 0..n {
let r = (y[i] - mu[i]) / variance(family, nb_theta, mu[i]).sqrt();
let pw = prior_w.map_or(1.0, |w| w[i]);
s += pw * r * r;
}
s / (n - p) as f64
}
pub(crate) fn is_canonical(family: Family) -> bool {
matches!(
family,
Family::Binomial {
link: BinomialLink::Logit
} | Family::Poisson { .. }
)
}
pub(crate) fn irls_weight_and_resid(
family: Family,
nb_theta: f64,
y: f64,
eta: f64,
) -> (f64, f64, f64) {
let mu = link_inv(family, eta);
let v = variance(family, nb_theta, mu);
if is_canonical(family) {
(mu, v, (y - mu) / v)
} else {
let dm = mu_eta(family, eta);
(mu, dm * dm / v, (y - mu) / dm)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{BinomialLink, Family, GammaLink, NegBinomialLink, PoissonLink};
#[test]
fn poisson_log_canonical_quantities() {
let f = Family::Poisson {
link: PoissonLink::Log,
};
let eta = 0.5_f64;
let mu = link_inv(f, eta);
assert!((mu - eta.exp()).abs() < 1e-12); assert!((variance(f, f64::NAN, mu) - mu).abs() < 1e-12); let (m, w, r) = irls_weight_and_resid(f, f64::NAN, 3.0, eta);
assert!((m - mu).abs() < 1e-12 && (w - mu).abs() < 1e-12);
assert!((r - (3.0 - mu) / mu).abs() < 1e-12);
}
#[test]
fn gamma_log_noncanonical_weight() {
let f = Family::Gamma {
link: GammaLink::Log,
};
let eta = 0.2_f64;
let mu = eta.exp();
let (_m, w, r) = irls_weight_and_resid(f, f64::NAN, 1.0, eta);
assert!((w - 1.0).abs() < 1e-12, "w={w}");
assert!((r - (1.0 - mu) / mu).abs() < 1e-12); }
#[test]
fn gamma_inverse_residual_sign() {
let f = Family::Gamma {
link: GammaLink::Inverse,
};
let eta = 0.5_f64; let mu = 1.0 / eta;
let (m, w, r) = irls_weight_and_resid(f, f64::NAN, 3.0, eta);
assert!((m - mu).abs() < 1e-12 && (w - mu * mu).abs() < 1e-12);
assert!((r - (-(3.0 - mu) / (mu * mu))).abs() < 1e-12, "r={r}");
}
#[test]
fn poisson_deviance_resid_zero_at_fit() {
let f = Family::Poisson {
link: PoissonLink::Log,
};
assert!(dev_resid(f, f64::NAN, 4.0, 4.0).abs() < 1e-10);
assert!(dev_resid(f, f64::NAN, 4.0, 2.0) > 0.0);
assert!((dev_resid(f, f64::NAN, 0.0, 1.0) - 2.0).abs() < 1e-10);
}
#[test]
fn nb_variance_uses_theta() {
let f = Family::NegativeBinomial {
link: NegBinomialLink::Log,
};
let mu = 3.0;
assert!((variance(f, 2.0, mu) - (mu + mu * mu / 2.0)).abs() < 1e-12);
}
#[test]
fn probit_mu_eta_is_normal_pdf() {
let f = Family::Binomial {
link: BinomialLink::Probit,
};
assert!((link_inv(f, 0.0) - 0.5).abs() < 1e-13);
assert!((mu_eta(f, 0.0) - (1.0 / (2.0 * std::f64::consts::PI).sqrt())).abs() < 1e-12);
}
}