use faer::Mat;
use crate::glmm::{build_z, GlmmWorkspace, StructuredSchur};
use crate::lmm::LmmWorkspace;
use crate::{BinomialLink, Family, GroupIds, GroupingRelation, ModelSpec, StartValues};
use super::common::{
assert_group_ids, assert_model_shape, fill_col_major, unpermute_fit, warm_theta,
FitDiagnostics, Perm,
};
use super::glm::GlmScratchBuf;
use super::lmm::LmmResultView;
use super::ols::OlsWorkspace;
use super::{classify_design, Fit, FitOptions, Solver};
pub struct FitView<'a> {
kind: FitViewKind<'a>,
perm: Perm,
theta_declared: Vec<f64>,
}
#[allow(clippy::large_enum_variant)]
enum FitViewKind<'a> {
Ols(crate::ols::OlsFitView<'a>),
Glm(crate::glm::GlmFitView<'a>),
Lmm(LmmResultView<'a>),
Glmm(super::glmm::GlmmResultView<'a>),
Prebuilt {
fit: Fit,
#[allow(dead_code)]
t_sq: Vec<f64>,
#[allow(dead_code)]
var_diag: Vec<f64>,
},
}
#[allow(dead_code)]
impl FitView<'_> {
pub fn t_sq(&self) -> &[f64] {
match &self.kind {
FitViewKind::Ols(v) => v.t_sq,
FitViewKind::Glm(v) => v.t_sq,
FitViewKind::Lmm(v) => v.t_sq(),
FitViewKind::Glmm(v) => v.t_sq(),
FitViewKind::Prebuilt { t_sq, .. } => t_sq,
}
}
pub fn diagnostics(&self) -> FitDiagnostics {
match &self.kind {
FitViewKind::Ols(v) => v.diagnostics(),
FitViewKind::Glm(v) => v.diagnostics(),
FitViewKind::Lmm(v) => v.diagnostics(),
FitViewKind::Glmm(v) => v.diagnostics(),
FitViewKind::Prebuilt { fit, .. } => FitDiagnostics {
boundary_hit: fit.singular() as u8,
..FitDiagnostics::fixed_only(fit.converged())
},
}
}
pub fn converged(&self) -> bool {
self.diagnostics().converged
}
pub fn betas(&self) -> &[f64] {
match &self.kind {
FitViewKind::Ols(v) => v.betas,
FitViewKind::Glm(v) => v.betas,
FitViewKind::Lmm(v) => v.betas(),
FitViewKind::Glmm(v) => v.betas(),
FitViewKind::Prebuilt { fit, .. } => &fit.beta,
}
}
pub fn var_diag(&self) -> &[f64] {
match &self.kind {
FitViewKind::Ols(v) => v.var_diag,
FitViewKind::Glm(v) => v.var_diag,
FitViewKind::Lmm(v) => v.var_diag(),
FitViewKind::Glmm(v) => v.var_diag(),
FitViewKind::Prebuilt { var_diag, .. } => var_diag,
}
}
pub fn joint_t_sq(&self) -> f64 {
match &self.kind {
FitViewKind::Ols(_) | FitViewKind::Glm(_) => f64::NAN,
FitViewKind::Lmm(v) => v.joint_t_sq(),
FitViewKind::Glmm(v) => v.joint_t_sq(),
FitViewKind::Prebuilt { .. } => f64::NAN,
}
}
pub fn n_eval(&self) -> usize {
match &self.kind {
FitViewKind::Ols(_) | FitViewKind::Glm(_) => 0,
FitViewKind::Lmm(v) => v.n_eval(),
FitViewKind::Glmm(v) => v.n_eval(),
FitViewKind::Prebuilt { fit, .. } => fit.n_eval,
}
}
pub fn dispersion(&self) -> f64 {
match &self.kind {
FitViewKind::Ols(_) | FitViewKind::Glm(_) => f64::NAN,
FitViewKind::Lmm(v) => v.dispersion(),
FitViewKind::Glmm(v) => v.dispersion(),
FitViewKind::Prebuilt { fit, .. } => fit.dispersion,
}
}
pub fn theta(&self) -> &[f64] {
if self.theta_declared.is_empty() {
self.kernel_theta()
} else {
&self.theta_declared
}
}
fn kernel_groupings(&self) -> Option<&crate::lmm::LmmGroupings> {
match &self.kind {
FitViewKind::Lmm(v) => Some(v.groupings()),
FitViewKind::Glmm(v) => Some(v.groupings()),
FitViewKind::Ols(_) | FitViewKind::Glm(_) | FitViewKind::Prebuilt { .. } => None,
}
}
fn kernel_theta(&self) -> &[f64] {
match &self.kind {
FitViewKind::Ols(_) | FitViewKind::Glm(_) => &[],
FitViewKind::Lmm(v) => v.theta(),
FitViewKind::Glmm(v) => v.theta(),
FitViewKind::Prebuilt { .. } => &[],
}
}
#[allow(clippy::too_many_arguments)] pub fn into_fit(
self,
x: &[f64],
y: &[f64],
ids: &GroupIds,
n: usize,
p: usize,
model: &ModelSpec,
opts: &FitOptions,
) -> Fit {
let perm = self.perm;
let mut fit = match self.kind {
FitViewKind::Ols(v) => super::ols::ols_view_to_fit(&v, x, y, n, p, opts),
FitViewKind::Glm(v) => {
super::glm::glm_view_to_fit(&v, y, model.family, f64::NAN, n, p, opts)
}
FitViewKind::Lmm(v) => super::lmm::lmm_view_to_fit(&v, x, ids, n, p, opts),
FitViewKind::Glmm(v) => super::glmm::glmm_view_to_fit(&v, y, n, p, model, opts).0,
FitViewKind::Prebuilt { fit, .. } => fit,
};
unpermute_fit(perm, &mut fit);
fit
}
}
impl<'a> FitView<'a> {
fn new(kind: FitViewKind<'a>, perm: Perm) -> Self {
let mut view = FitView {
kind,
perm,
theta_declared: Vec::new(),
};
let scales = view
.kernel_groupings()
.filter(|g| g.any_slope_scaled())
.map(|g| g.theta_row_scales());
if scales.is_some() || !perm.is_identity() {
let mut theta = view.kernel_theta().to_vec();
if let Some(s) = &scales {
for (t, &sc) in theta.iter_mut().zip(s.iter()) {
*t /= sc;
}
}
perm.swap_slots(&mut theta);
view.theta_declared = theta;
}
view
}
}
fn prebuilt_stats(fit: &Fit) -> (Vec<f64>, Vec<f64>) {
let p = fit.beta.len();
let mut t_sq = vec![f64::NAN; p];
let mut var_diag = vec![f64::NAN; p];
for j in 0..p {
let se = fit.se[j];
if se.is_finite() {
var_diag[j] = se * se;
t_sq[j] = (fit.beta[j] / se).powi(2);
}
}
(t_sq, var_diag)
}
pub struct FitWorkspace {
n_max: usize,
p: usize,
sized: ModelSpec,
perm: Perm,
build_primary_levels: usize,
build_extra_capacity: Vec<usize>,
nagq: u8,
n_targets: usize,
has_weights: bool,
has_offset: bool,
parallel_inner: bool,
kind: FitKind,
}
#[allow(clippy::large_enum_variant)]
enum FitKind {
Ols {
ws: OlsWorkspace,
x_mat: Mat<f64>,
},
Glm {
buf: GlmScratchBuf,
x_mat: Mat<f64>,
},
LmmDense {
ws: LmmWorkspace,
x_mat: Mat<f64>,
y_shifted: Vec<f64>,
},
GlmmDense {
ws: GlmmWorkspace,
x_mat: Mat<f64>,
},
Prebuilt(PrebuiltRoute),
}
#[derive(Clone, Copy)]
enum PrebuiltRoute {
GlmNb,
GlmmNbDense,
LmmSparse,
GlmmSparse,
GlmmNbSparse,
}
#[cfg(test)]
impl FitWorkspace {
pub(crate) fn is_ols(&self) -> bool {
matches!(self.kind, FitKind::Ols { .. })
}
pub(crate) fn is_glm(&self) -> bool {
matches!(self.kind, FitKind::Glm { .. })
}
pub(crate) fn is_lmm_dense(&self) -> bool {
matches!(self.kind, FitKind::LmmDense { .. })
}
pub(crate) fn is_glmm_dense(&self) -> bool {
matches!(self.kind, FitKind::GlmmDense { .. })
}
pub(crate) fn is_prebuilt(&self) -> bool {
matches!(self.kind, FitKind::Prebuilt(_))
}
}
pub fn build_workspace(
sized: &ModelSpec,
perm: Perm,
n_max: usize,
p: usize,
opts: &FitOptions,
) -> FitWorkspace {
assert_model_shape(sized, p, opts.nagq);
let t = opts.target_indices.len();
let kind = match (&sized.family, sized.re.as_ref()) {
(Family::Gaussian, None) => FitKind::Ols {
ws: OlsWorkspace::new(n_max, p, t, opts.weights.is_some()),
x_mat: Mat::<f64>::zeros(n_max.max(1), p.max(1)),
},
(Family::NegativeBinomial { .. }, None) => FitKind::Prebuilt(PrebuiltRoute::GlmNb),
(
Family::Poisson { .. }
| Family::Gamma { .. }
| Family::Binomial {
link: BinomialLink::Probit | BinomialLink::Logit,
},
None,
) => FitKind::Glm {
buf: GlmScratchBuf::new(n_max, p, t),
x_mat: Mat::<f64>::zeros(n_max.max(1), p.max(1)),
},
(family, Some(re)) => match classify_design(sized, opts.nagq) {
Solver::NoZ => match family {
Family::Gaussian => {
let slope_cols: Vec<usize> = re.slopes.iter().map(|&c| c as usize).collect();
let extra_slope_cols: Vec<Vec<usize>> = re
.extra_groupings
.iter()
.map(|g| g.slopes.iter().map(|&c| c as usize).collect())
.collect();
FitKind::LmmDense {
ws: LmmWorkspace::for_cluster_spec_ext(
p,
sized,
n_max,
&slope_cols,
&extra_slope_cols,
),
x_mat: Mat::<f64>::zeros(n_max.max(1), p.max(1)),
y_shifted: vec![0.0f64; n_max.max(1)],
}
}
Family::NegativeBinomial { .. } => FitKind::Prebuilt(PrebuiltRoute::GlmmNbDense),
_ => {
let slope_cols: Vec<usize> = re.slopes.iter().map(|&c| c as usize).collect();
let ws =
GlmmWorkspace::for_cluster_spec(p, sized, n_max, &slope_cols, opts.nagq);
let x_mat = Mat::<f64>::zeros(n_max.max(1), p.max(1));
FitKind::GlmmDense { ws, x_mat }
}
},
Solver::Sparse => match family {
Family::Gaussian => FitKind::Prebuilt(PrebuiltRoute::LmmSparse),
Family::NegativeBinomial { .. } => FitKind::Prebuilt(PrebuiltRoute::GlmmNbSparse),
_ => FitKind::Prebuilt(PrebuiltRoute::GlmmSparse),
},
},
};
let build_primary_levels = sized
.re
.as_ref()
.map(|re| re.sizing.n_clusters_at(n_max))
.unwrap_or(0);
let build_extra_capacity: Vec<usize> = sized
.re
.as_ref()
.map(|re| {
re.extra_groupings
.iter()
.map(|g| match g.relation {
GroupingRelation::Crossed { n_clusters } => n_clusters.max(1) as usize,
GroupingRelation::NestedWithin { n_per_parent } => {
build_primary_levels * n_per_parent.max(1) as usize
}
})
.collect()
})
.unwrap_or_default();
FitWorkspace {
n_max,
p,
sized: sized.clone(),
perm,
build_primary_levels,
build_extra_capacity,
nagq: opts.nagq,
n_targets: t,
has_weights: opts.weights.is_some(),
has_offset: opts.offset.is_some(),
parallel_inner: opts.parallel_inner,
kind,
}
}
pub fn fit_on<'a>(
ws: &'a mut FitWorkspace,
x: &[f64],
y: &[f64],
ids: &GroupIds,
start: Option<&StartValues>,
opts: &FitOptions,
) -> FitView<'a> {
let n = y.len();
let p = ws.p;
assert!(
n <= ws.n_max,
"fit_on: shape mismatch — n={n} exceeds build n_max={}",
ws.n_max
);
assert_eq!(
x.len(),
n * p,
"fit_on: shape mismatch — x.len()={} vs n*p={}",
x.len(),
n * p
);
assert_eq!(opts.nagq, ws.nagq, "fit_on: nagq is frozen at build");
assert_eq!(
opts.target_indices.len(),
ws.n_targets,
"fit_on: target count is frozen at build"
);
assert_eq!(
opts.weights.is_some(),
ws.has_weights,
"fit_on: weights presence is frozen at build"
);
assert_eq!(
opts.offset.is_some(),
ws.has_offset,
"fit_on: offset presence is frozen at build"
);
assert_eq!(
opts.parallel_inner, ws.parallel_inner,
"fit_on: parallel_inner is frozen at build"
);
let mixed = ws.sized.re.is_some();
if mixed {
let re = ws.sized.re.as_ref().unwrap();
assert_group_ids(re, ids, n);
let call_primary = ids
.primary
.iter()
.copied()
.max()
.map(|m| m as usize + 1)
.unwrap_or(1);
assert_eq!(
call_primary, ws.build_primary_levels,
"fit_on: shape mismatch — primary level count {call_primary} != build {}",
ws.build_primary_levels
);
for (g, e) in ids.extra.iter().enumerate() {
let levels = e.iter().copied().max().map(|m| m as usize + 1).unwrap_or(0);
assert!(
levels <= ws.build_extra_capacity[g],
"fit_on: shape mismatch — extra grouping {g} needs {levels} levels, build capacity {}",
ws.build_extra_capacity[g]
);
}
}
let perm = ws.perm;
let permuted_start;
let start = match start {
Some(s) if !perm.is_identity() && !s.theta.is_empty() => {
let mut s = s.clone();
perm.swap_slots(&mut s.theta);
permuted_start = s;
Some(&permuted_start)
}
unchanged => unchanged,
};
match &mut ws.kind {
FitKind::Ols { ws: ols_ws, x_mat } => {
fill_col_major(x_mat, x, n, p);
let v =
super::ols::fit_ols_prebuilt(ols_ws, x_mat.as_ref().subrows(0, n), y, n, p, opts);
FitView::new(FitViewKind::Ols(v), perm)
}
FitKind::Glm { buf, x_mat } => {
fill_col_major(x_mat, x, n, p);
let v = super::glm::fit_glm_prebuilt(
ws.sized.family,
f64::NAN,
x_mat.as_ref().subrows(0, n),
y,
opts,
buf,
);
FitView::new(FitViewKind::Glm(v), perm)
}
FitKind::LmmDense {
ws: lmm_ws,
x_mat,
y_shifted,
} => {
fill_col_major(x_mat, x, n, p);
let y_eff: &[f64] = match &opts.offset {
Some(o) => {
for i in 0..n {
y_shifted[i] = y[i] - o[i];
}
&y_shifted[..n]
}
None => y,
};
super::lmm::accumulate_lmm_rows(
lmm_ws,
x_mat.as_ref().subrows(0, n),
y_eff,
n,
p,
&ids.primary,
&ids.extra,
opts.weights.as_deref(),
);
let v = super::lmm::lmm_run_on(lmm_ws, &opts.target_indices, warm_theta(start));
FitView::new(FitViewKind::Lmm(v), perm)
}
FitKind::GlmmDense { ws: glmm_ws, x_mat } => {
fill_col_major(x_mat, x, n, p);
glmm_ws.parallel_inner = opts.parallel_inner;
if let Some(w) = &opts.weights {
glmm_ws.prior_w[..n].copy_from_slice(w);
glmm_ws.weighted = true;
} else {
glmm_ws.weighted = false;
}
glmm_ws.offset = opts.offset.clone();
glmm_ws
.groupings
.set_slope_scales(x_mat.as_ref().subrows(0, n), opts.weights.as_deref());
build_z(
glmm_ws,
x_mat.as_ref().subrows(0, n),
&ids.primary,
&ids.extra,
n,
);
glmm_ws.structured_schur = if glmm_ws.groupings.structured_extras_eligible() {
StructuredSchur::new(&glmm_ws.groupings, &ids.primary, &ids.extra, n)
} else {
None
};
let v = super::glmm::run_glmm_on(
glmm_ws,
x_mat.as_ref().subrows(0, n),
y,
n,
p,
&ws.sized,
&ids.primary,
&ids.extra,
f64::NAN,
start,
opts,
);
FitView::new(FitViewKind::Glmm(v), perm)
}
FitKind::Prebuilt(route) => {
let route = *route;
let sized = &ws.sized;
let fit = match route {
PrebuiltRoute::GlmNb => super::glm::fit_glm_nb(x, y, n, p, None, opts),
PrebuiltRoute::GlmmNbDense => super::glmm::fit_glmm_nb(
x,
y,
n,
p,
sized,
&ids.primary,
&ids.extra,
start,
opts,
),
PrebuiltRoute::LmmSparse => crate::sparse::fit_mle_sparse(
x,
y,
n,
p,
sized,
&ids.primary,
&ids.extra,
start,
opts,
),
PrebuiltRoute::GlmmSparse => {
crate::sparse::fit_glmm_sparse(
x,
y,
n,
p,
sized,
&ids.primary,
&ids.extra,
f64::NAN,
start,
opts,
)
.0
}
PrebuiltRoute::GlmmNbSparse => crate::sparse::fit_glmm_nb_sparse(
x,
y,
n,
p,
sized,
&ids.primary,
&ids.extra,
start,
opts,
),
};
let (t_sq, var_diag) = prebuilt_stats(&fit);
FitView::new(
FitViewKind::Prebuilt {
fit,
t_sq,
var_diag,
},
perm,
)
}
}
}
#[cfg(test)]
#[path = "core_tests.rs"]
mod core_tests;