use crate::consts::{MAX_EXTRA_GROUPINGS, MAX_EXTRA_Q, MAX_PRIMARY_Q};
use crate::{BinomialLink, Family, GroupIds, GroupingRelation, ModelSpec, StartValues, WaldSe};
mod common;
mod glm;
mod glmm;
mod lmm;
mod loop_advanced_seam;
mod ols;
use common::{
assert_group_ids, assert_model_shape, detect_aliased, fit_rank_deficient, spec_sized_from_ids,
theta_width,
};
use glm::{fit_glm, fit_glm_nb};
use glmm::{fit_glmm, fit_glmm_nb};
use lmm::fit_mle;
use ols::fit_ols;
pub struct Fit {
pub beta: Vec<f64>,
pub se: Vec<f64>,
pub vcov: Vec<Vec<f64>>,
pub tau2: Vec<f64>,
pub dispersion: f64,
pub converged: bool,
pub varcorr: Vec<Vec<f64>>,
pub stddev_se: Vec<f64>,
pub aliased: Vec<bool>,
pub n_eval: usize,
pub deviance: f64,
pub singular: bool,
pub loglik: f64,
pub df: usize,
pub reml: bool,
pub fitted: Vec<f64>,
pub ranef: Vec<f64>,
pub ranef_levels: Vec<usize>,
}
const SINGULAR_REL_TOL: f64 = 1e-3;
impl Fit {
pub fn stddev_corr(&self, group_idx: usize) -> (Vec<f64>, Vec<Vec<f64>>) {
let vech = &self.varcorr[group_idx];
let len = vech.len();
let q = (((1 + 8 * len) as f64).sqrt() as usize - 1) / 2;
debug_assert_eq!(
q * (q + 1) / 2,
len,
"varcorr[{group_idx}] is not a valid vech"
);
let idx = |r: usize, c: usize| -> usize { c * q - (c * c - c) / 2 + (r - c) };
let stddev: Vec<f64> = (0..q).map(|i| vech[idx(i, i)].sqrt()).collect();
let mut corr = vec![vec![0.0; q]; q];
#[allow(clippy::needless_range_loop)]
for i in 0..q {
corr[i][i] = 1.0;
}
for c in 0..q {
for r in (c + 1)..q {
let rho = vech[idx(r, c)] / (stddev[r] * stddev[c]);
corr[r][c] = rho;
corr[c][r] = rho;
}
}
(stddev, corr)
}
pub(crate) fn has_negligible_component(&self) -> bool {
if self.varcorr.is_empty() {
return false;
}
let stddevs: Vec<Vec<f64>> = (0..self.varcorr.len())
.map(|g| self.stddev_corr(g).0)
.collect();
let max_sd = stddevs.iter().flatten().cloned().fold(0.0_f64, f64::max);
if max_sd <= 0.0 {
return false;
}
stddevs
.iter()
.flatten()
.any(|&s| s <= SINGULAR_REL_TOL * max_sd)
}
}
#[derive(Clone)]
pub struct FitOptions {
pub target_indices: Vec<u32>,
pub wald_se: WaldSe,
pub nagq: u8,
pub dispersion: Option<f64>,
pub weights: Option<Vec<f64>>,
pub offset: Option<Vec<f64>>,
pub parallel_inner: bool,
}
impl Default for FitOptions {
fn default() -> Self {
FitOptions {
target_indices: vec![],
wald_se: WaldSe::Hessian,
nagq: 1,
dispersion: None,
weights: None,
offset: None,
parallel_inner: false,
}
}
}
pub fn fit_cold(
x: &[f64],
y: &[f64],
n: usize,
p: usize,
model: &ModelSpec,
ids: &GroupIds,
opts: &FitOptions,
) -> Fit {
fit_warm(x, y, n, p, model, ids, None, opts)
}
#[allow(clippy::too_many_arguments)] pub fn fit_warm(
x: &[f64],
y: &[f64],
n: usize,
p: usize,
model: &ModelSpec,
ids: &GroupIds,
start: Option<&StartValues>,
opts: &FitOptions,
) -> Fit {
assert_eq!(
x.len(),
n * p,
"x must have n*p elements in row-major layout"
);
assert_eq!(y.len(), n, "y must have n elements");
assert_model_shape(model, p, opts.nagq);
if let Some(s) = start {
assert_eq!(s.beta.len(), p, "StartValues.beta must have p elements");
assert_eq!(
s.theta.len(),
theta_width(model.re.as_ref()),
"StartValues.theta must have n_theta elements for this RE structure"
);
}
if let Some(w) = &opts.weights {
assert_eq!(w.len(), n, "FitOptions.weights must have n elements");
assert!(
w.iter().all(|&v| v.is_finite() && v > 0.0),
"FitOptions.weights must be finite and > 0"
);
}
if let Some(o) = &opts.offset {
assert_eq!(o.len(), n, "FitOptions.offset must have n elements");
assert!(
o.iter().all(|&v| v.is_finite()),
"FitOptions.offset must be finite"
);
}
if n > 0 && p > 0 {
let aliased = detect_aliased(x, n, p);
if aliased.iter().any(|&a| a) {
return fit_rank_deficient(x, y, n, p, model, ids, start, opts, &aliased);
}
}
match (&model.family, model.re.as_ref()) {
(Family::Gaussian, None) => fit_ols(x, y, n, p, opts),
(
Family::Poisson { .. }
| Family::Gamma { .. }
| Family::Binomial {
link: BinomialLink::Probit | BinomialLink::Logit,
},
None,
) => fit_glm(model.family, f64::NAN, x, y, n, p, opts),
(Family::NegativeBinomial { .. }, None) => fit_glm_nb(x, y, n, p, None, opts),
(family, Some(re)) => {
assert_group_ids(re, ids, n);
let sized = spec_sized_from_ids(model, ids);
match classify_design(&sized, opts.nagq) {
Solver::NoZ => match family {
Family::Gaussian => {
fit_mle(x, y, n, p, &sized, &ids.primary, &ids.extra, start, opts)
}
Family::NegativeBinomial { .. } => {
fit_glmm_nb(x, y, n, p, &sized, &ids.primary, &ids.extra, start, opts)
}
Family::Binomial { .. } | Family::Poisson { .. } | Family::Gamma { .. } => {
fit_glmm(
x,
y,
n,
p,
&sized,
&ids.primary,
&ids.extra,
f64::NAN,
start,
opts,
)
.0
}
},
Solver::Sparse => match family {
Family::Gaussian => crate::sparse::fit_mle_sparse(
x,
y,
n,
p,
&sized,
&ids.primary,
&ids.extra,
start,
opts,
),
Family::NegativeBinomial { .. } => crate::sparse::fit_glmm_nb_sparse(
x,
y,
n,
p,
&sized,
&ids.primary,
&ids.extra,
start,
opts,
),
Family::Binomial { .. } | Family::Poisson { .. } | Family::Gamma { .. } => {
crate::sparse::fit_glmm_sparse(
x,
y,
n,
p,
&sized,
&ids.primary,
&ids.extra,
f64::NAN,
start,
opts,
)
.0
}
},
}
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum Solver {
NoZ,
Sparse,
}
pub(crate) fn classify_design(model: &ModelSpec, _nagq: u8) -> Solver {
let Some(re) = model.re.as_ref() else {
return Solver::NoZ; };
let q_p = 1 + re.slopes.len();
let over = re.extra_groupings.len() > MAX_EXTRA_GROUPINGS
|| q_p > MAX_PRIMARY_Q
|| re
.extra_groupings
.iter()
.any(|g| 1 + g.slopes.len() > MAX_EXTRA_Q);
let slope_extras = re.extra_groupings.iter().any(|g| !g.slopes.is_empty());
let crossed_levels: usize = re
.extra_groupings
.iter()
.map(|g| match g.relation {
GroupingRelation::Crossed { n_clusters } => n_clusters as usize,
GroupingRelation::NestedWithin { .. } => 0,
})
.sum();
if over || slope_extras || crossed_levels > crate::consts::MAX_CROSSED_LEVELS {
Solver::Sparse
} else {
Solver::NoZ
}
}
#[cfg(test)]
pub(crate) fn classify_design_pub(model: &ModelSpec, nagq: u8) -> Solver {
classify_design(model, nagq)
}
pub(crate) use common::{
assemble_ranef_sparse, assemble_varcorr, glmm_loglik, lmm_loglik, model_df, nan_vcov,
ranef_level_counts, vcov_from_chol,
};
#[cfg(test)]
pub(crate) use common::{assert_model_shape_pub, spec_sized_from_ids_pub};
pub(crate) use glm::{golden_max_ln_theta, nb_profile_loglik};
pub(crate) use glmm::glm_warm_start_beta;
#[cfg(test)]
pub(crate) use lmm::fit_mle_noz_pub;
#[cfg(feature = "loop_advanced")]
pub use loop_advanced_seam::{
build_lmm_seam_ws, build_lmm_workspace, lmm_objective_at, lmm_sweep_fit, lmm_sweep_fit_on,
refit_lmm, LmmSeamWs, LmmSweepOutcome,
};
#[cfg(test)]
mod common_tests;
#[cfg(test)]
mod glm_tests;
#[cfg(test)]
mod glmm_tests;
#[cfg(test)]
mod lmm_tests;
#[cfg(test)]
mod ols_tests;