use faer::Mat;
use crate::glm::{glm_irls_fit, GlmFitView, GlmScratch};
use crate::{Family, NegBinomialLink};
use super::common::{fill_se_compact, nan_vcov, to_col_major, vcov_from_chol, FitDiagnostics};
use super::{Diagnostics, Fit, FitOptions};
impl GlmFitView<'_> {
pub(crate) fn diagnostics(&self) -> FitDiagnostics {
FitDiagnostics {
pivot: self.pivot,
pivot_col: self.pivot_col,
ill_conditioned: self.pivot < crate::ols::PIVOT_MIN,
..FitDiagnostics::fixed_only(self.converged)
}
}
}
fn fit_unsupported_family(p: usize) -> Fit {
Fit {
beta: vec![f64::NAN; p],
se: vec![f64::NAN; p],
vcov: nan_vcov(p),
tau2: vec![],
dispersion: f64::NAN,
diagnostics: Diagnostics::from_flags(false, false, p),
varcorr: vec![],
stddev_se: vec![],
n_eval: 0,
#[cfg(feature = "counters")]
counters: crate::counters::EvalCounters::new(),
deviance: f64::NAN,
loglik: f64::NAN,
df: 0,
reml: false,
fitted: vec![],
ranef: vec![],
ranef_levels: vec![],
}
}
pub(crate) struct GlmScratchBuf {
irls_eta: Vec<f64>,
irls_p: Vec<f64>,
irls_w: Vec<f64>,
irls_z: Vec<f64>,
irls_betas: Vec<f64>,
irls_betas_new: Vec<f64>,
irls_var_diag: Vec<f64>,
irls_t_sq: Vec<f64>,
irls_u_scratch: Vec<f64>,
irls_xtwx: Mat<f64>,
irls_xtwz: Vec<f64>,
irls_l: Mat<f64>,
irls_wx: Vec<f64>,
}
impl GlmScratchBuf {
pub(crate) fn new(n: usize, p: usize, t: usize) -> Self {
let (n1, p1, t1) = (n.max(1), p.max(1), t.max(1));
GlmScratchBuf {
irls_eta: vec![0.0f64; n1],
irls_p: vec![0.0f64; n1],
irls_w: vec![0.0f64; n1],
irls_z: vec![0.0f64; n1],
irls_betas: vec![0.0f64; p1],
irls_betas_new: vec![0.0f64; p1],
irls_var_diag: vec![0.0f64; t1],
irls_t_sq: vec![0.0f64; t1],
irls_u_scratch: vec![0.0f64; p1],
irls_xtwx: Mat::<f64>::zeros(p1, p1),
irls_xtwz: vec![0.0f64; p1],
irls_l: Mat::<f64>::zeros(p1, p1),
irls_wx: vec![0.0f64; n1 * p1], }
}
fn as_scratch(&mut self) -> GlmScratch<'_> {
GlmScratch {
irls_eta: &mut self.irls_eta,
irls_p: &mut self.irls_p,
irls_w: &mut self.irls_w,
irls_z: &mut self.irls_z,
irls_betas: &mut self.irls_betas,
irls_betas_new: &mut self.irls_betas_new,
irls_var_diag: &mut self.irls_var_diag,
irls_t_sq: &mut self.irls_t_sq,
irls_u_scratch: &mut self.irls_u_scratch,
irls_xtwx: self.irls_xtwx.as_mut(),
irls_xtwz: &mut self.irls_xtwz,
irls_l: self.irls_l.as_mut(),
irls_wx: &mut self.irls_wx,
}
}
}
#[cfg(test)]
pub(super) fn fit_glm(
family: Family,
nb_theta: f64,
x: &[f64],
y: &[f64],
n: usize,
p: usize,
opts: &FitOptions,
) -> Fit {
let mut buf = GlmScratchBuf::new(n, p, opts.target_indices.len());
let x_mat = to_col_major(x, n, p);
let view = fit_glm_prebuilt(
family,
nb_theta,
x_mat.as_ref().subrows(0, n),
y,
opts,
&mut buf,
);
glm_view_to_fit(&view, y, family, nb_theta, n, p, opts)
}
pub(crate) fn fit_glm_prebuilt<'a>(
family: Family,
nb_theta: f64,
x_mat: faer::MatRef<'_, f64>,
y: &[f64],
opts: &FitOptions,
buf: &'a mut GlmScratchBuf,
) -> GlmFitView<'a> {
glm_irls_fit(
family,
nb_theta,
x_mat,
y,
&opts.target_indices,
None,
opts.weights.as_deref(),
opts.offset.as_deref(),
buf.as_scratch(),
)
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn glm_view_to_fit(
view: &GlmFitView<'_>,
y: &[f64],
family: Family,
nb_theta: f64,
n: usize,
p: usize,
opts: &FitOptions,
) -> Fit {
let beta = view.betas.to_vec();
let diag = view.diagnostics();
let converged = diag.converged;
let irls_deviance = view.deviance;
let mut se = vec![f64::NAN; p];
fill_se_compact(view.var_diag, &opts.target_indices, &mut se);
let mut vcov = if converged {
vcov_from_chol(view.l, p, &opts.target_indices, 1.0)
} else {
nan_vcov(p)
};
let dispersion = if !converged {
f64::NAN
} else {
match family {
Family::Gamma { .. } | Family::InverseGaussian { .. } => {
let phi = match opts.dispersion {
Some(v) => v,
None => crate::family::pearson_dispersion(
y,
view.mu,
family,
nb_theta,
n,
p,
opts.weights.as_deref(),
),
};
let sqrt_phi = phi.sqrt();
for v in se.iter_mut() {
if v.is_finite() {
*v *= sqrt_phi;
}
}
phi
}
_ => 1.0,
}
};
for row in vcov.iter_mut() {
for v in row.iter_mut() {
*v *= dispersion;
}
}
let (fitted, loglik) = if converged {
let mu = view.mu.to_vec();
let ll = match family {
Family::Gamma { .. } => {
-0.5 * (crate::family::gamma_aic(y, &mu, irls_deviance, n, opts.weights.as_deref())
- 2.0)
}
Family::InverseGaussian { .. } => {
-0.5 * (crate::family::inv_gaussian_aic(
y,
irls_deviance,
n,
opts.weights.as_deref(),
) - 2.0)
}
_ => {
-0.5 * irls_deviance
+ crate::family::saturated_loglik(family, nb_theta, y, opts.weights.as_deref())
}
};
(mu, ll)
} else {
(vec![], f64::NAN)
};
Fit {
beta,
se,
vcov,
tau2: vec![],
dispersion,
diagnostics: super::common::materialize_diagnostics(&diag, p, &[]),
varcorr: vec![],
stddev_se: vec![],
n_eval: 0,
#[cfg(feature = "counters")]
counters: crate::counters::EvalCounters::new(),
deviance: f64::NAN,
loglik,
df: if converged {
super::common::model_df(family, p, 0, opts.dispersion.is_some())
} else {
0
},
reml: false,
fitted,
ranef: vec![],
ranef_levels: vec![],
}
}
pub(super) const NB_MAX_OUTER: usize = 25;
const NB_THETA_TOL: f64 = 1e-6;
pub(crate) const NB_THETA_LO: f64 = 1e-3;
pub(crate) const NB_THETA_HI: f64 = 1e4;
pub(crate) fn nb_profile_loglik(y: &[f64], mu: &[f64], theta: f64, weights: Option<&[f64]>) -> f64 {
let mut ll = 0.0;
for (i, (&yi, &mi)) in y.iter().zip(mu.iter()).enumerate() {
let mut s = 0.0;
for k in 0..(yi.round() as u64) {
s += (theta + k as f64).ln();
}
if mi > 0.0 {
s += theta * (theta / (theta + mi)).ln() + yi * (mi / (theta + mi)).ln();
}
ll += weights.map_or(1.0, |w| w[i]) * s;
}
ll
}
pub(crate) fn golden_max_ln_theta(mut g: impl FnMut(f64) -> f64) -> f64 {
const INV_PHI: f64 = 0.618_033_988_749_894_9; let (mut a, mut b) = (NB_THETA_LO.ln(), NB_THETA_HI.ln());
let mut c = b - (b - a) * INV_PHI;
let mut d = a + (b - a) * INV_PHI;
let (mut fc, mut fd) = (g(c), g(d));
for _ in 0..200 {
if fc > fd {
b = d;
d = c;
fd = fc;
c = b - (b - a) * INV_PHI;
fc = g(c);
} else {
a = c;
c = d;
fc = fd;
d = a + (b - a) * INV_PHI;
fd = g(d);
}
if (b - a).abs() < 1e-4 {
break;
}
}
(0.5 * (a + b)).exp()
}
fn optimize_nb_theta(y: &[f64], mu: &[f64], weights: Option<&[f64]>) -> f64 {
golden_max_ln_theta(|t| nb_profile_loglik(y, mu, t.exp(), weights))
}
pub(crate) fn nb_theta_moment_seed(y: &[f64], n: usize) -> f64 {
let ybar = y.iter().sum::<f64>() / n as f64;
let var = y.iter().map(|&yi| (yi - ybar).powi(2)).sum::<f64>() / (n.max(2) - 1) as f64;
(ybar * ybar / (var - ybar).max(1e-6)).clamp(NB_THETA_LO, NB_THETA_HI)
}
pub(super) fn fit_glm_nb(
x: &[f64],
y: &[f64],
n: usize,
p: usize,
theta_seed: Option<f64>,
opts: &FitOptions,
) -> (Fit, f64) {
fit_glm_nb_capped(x, y, n, p, theta_seed, opts, NB_MAX_OUTER)
}
#[allow(clippy::too_many_arguments)]
pub(super) fn fit_glm_nb_capped(
x: &[f64],
y: &[f64],
n: usize,
p: usize,
theta_seed: Option<f64>,
opts: &FitOptions,
max_outer: usize,
) -> (Fit, f64) {
let mut theta = theta_seed.unwrap_or_else(|| nb_theta_moment_seed(y, n));
let family = Family::NegativeBinomial {
link: NegBinomialLink::Log,
};
let x_mat = to_col_major(x, n, p);
let x_ref = x_mat.as_ref().subrows(0, n);
let mut buf = GlmScratchBuf::new(n, p, opts.target_indices.len());
let mut mu = vec![0.0f64; n];
let mut fit_result = fit_unsupported_family(p);
for _ in 0..max_outer {
let view = fit_glm_prebuilt(family, theta, x_ref, y, opts, &mut buf);
fit_result = glm_view_to_fit(&view, y, family, theta, n, p, opts);
if !fit_result.converged() {
break;
}
for (i, mi) in mu.iter_mut().enumerate() {
let mut eta: f64 = (0..p).map(|j| x[i * p + j] * fit_result.beta[j]).sum();
if let Some(o) = &opts.offset {
eta += o[i];
}
*mi = crate::family::link_inv(family, eta);
}
let new_theta = optimize_nb_theta(y, &mu, opts.weights.as_deref());
let converged = (new_theta - theta).abs() / theta < NB_THETA_TOL;
theta = new_theta;
if converged {
let view = fit_glm_prebuilt(family, theta, x_ref, y, opts, &mut buf);
fit_result = glm_view_to_fit(&view, y, family, theta, n, p, opts);
break;
}
}
if fit_result.converged() {
fit_result.dispersion = theta;
}
(fit_result, theta)
}