use bobyqa::{Bobyqa, Config, RestartConfig, Status};
use faer::reborrow::ReborrowMut;
use faer::{Mat, MatMut, MatRef};
use std::sync::OnceLock;
use crate::scalar::Scalar;
use crate::FLOAT_NEAR_ZERO;
mod kernel;
#[cfg(test)]
mod tests;
pub(crate) use kernel::precompute_balanced_collapse;
pub use kernel::{reml_deviance, LmmSuffStats};
pub(crate) use kernel::{reml_gradient, reml_hessian, LmmDualScratch, LmmHyperScratch};
pub const THETA0: f64 = 1.0;
pub const THETA_HI: f64 = 1e3;
pub const RHO_BEGIN: f64 = 0.5;
pub const RHO_END: f64 = 1e-6;
pub const GLMM_RHO_END: f64 = 3e-6;
pub const THETA_TRUTH_FLOOR: f64 = 0.01;
pub const PIN_THETA: f64 = 1e-4;
pub const PIVOT_MIN: f64 = 1e-12;
#[cfg(test)]
pub fn bobyqa_config(n_theta: usize) -> Config {
let mut config = Config::new(n_theta);
config.rho_begin = RHO_BEGIN;
config.rho_end = RHO_END;
apply_campaign_overrides(&mut config, n_theta);
config
}
pub(crate) fn eval_formula(formula: &str, n: usize) -> Option<usize> {
let (mult, add) = formula.split_once('n')?;
let mult: f64 = mult.parse().ok()?;
let add: usize = add.parse().ok()?;
Some((mult * n as f64).ceil() as usize + add)
}
pub(crate) fn npt_from_formula(formula: &str, n: usize) -> Option<usize> {
eval_formula(formula, n).map(|v| v.clamp(n + 2, (n + 1) * (n + 2) / 2))
}
fn env_formula(var: &'static str, cell: &'static OnceLock<Option<String>>) -> Option<String> {
cell.get_or_init(|| std::env::var(var).ok().filter(|s| !s.is_empty()))
.clone()
}
pub(crate) fn npt_override(n: usize) -> Option<usize> {
static V: OnceLock<Option<String>> = OnceLock::new();
npt_from_formula(&env_formula("LMM_NPT_FORMULA", &V)?, n)
}
pub(crate) fn max_fun_override(n: usize) -> Option<usize> {
static V: OnceLock<Option<String>> = OnceLock::new();
eval_formula(&env_formula("LMM_MAX_FUN_FORMULA", &V)?, n)
}
pub(crate) fn apply_campaign_overrides(config: &mut Config, n: usize) {
if let Some(npt) = npt_override(n) {
config.npt = npt;
}
if let Some(mf) = max_fun_override(n) {
config.max_fun = mf.max(config.npt + 1);
}
let mut restart = RestartConfig::new();
restart.cycle_budget_frac = 0.125;
restart.max_restarts = 1;
restart.improve_rel_tol = 1e-6;
restart.stall_reductions = 0;
config.restart = Some(restart);
}
fn two_stage_enabled() -> bool {
static V: OnceLock<bool> = OnceLock::new();
*V.get_or_init(|| std::env::var("LMM_TWO_STAGE").is_ok_and(|v| v == "1"))
}
fn two_stage_minimize(
suff: &LmmSuffStats,
fit: &mut LmmFitScratch,
theta: &mut [f64],
lower: &[f64],
upper: &[f64],
finite_evals: &mut usize,
) -> bobyqa::Outcome {
let n = theta.len();
let c1 = {
let mut c = Config::new(n);
c.npt = n + 2;
c.rho_begin = RHO_BEGIN;
c.rho_end = 1e-3;
c
};
let mut s1 = Bobyqa::new(n, c1).expect("stage-1 config valid");
let out1 = s1.minimize(
|xs| {
let d = reml_deviance(xs, suff, fit);
if d.is_finite() {
*finite_evals += 1;
}
d
},
theta,
lower,
upper,
);
let theta1 = theta.to_vec();
let min_diag = suff
.groupings
.diagonal_theta()
.iter()
.map(|&i| theta[i])
.fold(f64::INFINITY, f64::min);
let rho_begin2 = (0.1 * min_diag).clamp(10.0 * RHO_END, RHO_BEGIN);
let c2 = {
let mut c = Config::new(n);
c.npt = 2 * n + 1;
c.rho_begin = rho_begin2;
c.rho_end = RHO_END;
c
};
let mut s2 = Bobyqa::new(n, c2).expect("stage-2 config valid");
let out2 = s2.minimize(
|xs| {
let d = reml_deviance(xs, suff, fit);
if d.is_finite() {
*finite_evals += 1;
}
d
},
theta,
lower,
upper,
);
if std::env::var("LMM_STAGE_PROBE").is_ok_and(|v| v == "1") {
let dist = theta1
.iter()
.zip(theta.iter())
.map(|(a, b)| (a - b).powi(2))
.sum::<f64>()
.sqrt();
eprintln!(
"stage_evals={},{} stage1_dist={dist:.6e}",
out1.n_eval, out2.n_eval
);
}
bobyqa::Outcome {
n_eval: out1.n_eval + out2.n_eval,
..out2
}
}
pub(crate) fn sparse_lmm_seed(groupings: &LmmGroupings) -> (Bobyqa, Vec<f64>, Vec<f64>, Vec<f64>) {
let n_theta = groupings.n_theta();
let blind_theta = vec![THETA0; n_theta];
let rho_begin = (0.1
* groupings
.diagonal_theta()
.iter()
.map(|&i| blind_theta[i])
.fold(f64::INFINITY, f64::min))
.min(RHO_BEGIN);
let npt = if n_theta >= 3 {
(3 * n_theta).div_ceil(2) + 1
} else {
2 * n_theta + 1
};
let mut config = Config::new(n_theta);
config.rho_begin = rho_begin;
config.rho_end = RHO_END;
config.npt = npt;
apply_campaign_overrides(&mut config, n_theta);
let (theta, lower, upper) = groupings.blind_theta_and_bounds();
let solver =
Bobyqa::new(n_theta, config).expect("BOBYQA config constants are valid by construction");
(solver, theta, lower, upper)
}
pub use crate::consts::{MAX_EXTRA_GROUPINGS, MAX_EXTRA_Q, MAX_PRIMARY_Q};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct CrossedFactor {
pub vech_start: usize,
pub q: usize,
pub n_levels: usize,
pub decl: usize,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct NestedFactor {
pub vech_start: usize,
pub q: usize,
pub decl: usize,
}
#[derive(Clone)]
pub struct LmmGroupings {
pub n_primary: usize,
pub nested_per_parent: usize,
pub nested: Option<NestedFactor>,
pub crossed: Vec<CrossedFactor>,
pub extra_offsets: Vec<usize>,
pub k_total: usize,
pub primary_q: usize,
pub primary_slope_cols: Vec<usize>,
pub diagonal_theta: Vec<usize>,
pub diagonal_run_len: Vec<usize>,
pub extra_slopes_any: bool,
pub extra_q: Vec<usize>,
pub extra_slope_cols: Vec<Vec<usize>>,
pub primary_slope_scales: Vec<f64>,
pub extra_slope_scales: Vec<Vec<f64>>,
}
pub fn rms_column_scale(x: MatRef<'_, f64>, col: usize, weights: Option<&[f64]>) -> f64 {
let mut num = 0.0f64;
let mut den = 0.0f64;
for i in 0..x.nrows() {
let w = weights.map_or(1.0, |w| w[i]);
let v = x[(i, col)];
num += w * v * v;
den += w;
}
if den <= 0.0 {
return 1.0;
}
let s = (num / den).sqrt();
if s.is_finite() && s > 0.0 {
s
} else {
1.0
}
}
fn compute_diagonal_theta(primary_q: usize, extra_qs: &[usize]) -> Vec<usize> {
let mut idx = Vec::with_capacity(primary_q + extra_qs.len());
let mut start = 0usize;
for &q in std::iter::once(&primary_q).chain(extra_qs.iter()) {
let mut off = start;
for d in 0..q {
idx.push(off);
off += q - d; }
start += q * (q + 1) / 2; }
idx
}
fn compute_diagonal_run_len(primary_q: usize, extra_qs: &[usize]) -> Vec<usize> {
let mut run_len = Vec::with_capacity(primary_q + extra_qs.len());
for &q in std::iter::once(&primary_q).chain(extra_qs.iter()) {
for d in 0..q {
run_len.push(q - d);
}
}
run_len
}
impl LmmGroupings {
pub fn single(max_clusters: usize) -> Self {
LmmGroupings {
n_primary: max_clusters,
nested_per_parent: 0,
nested: None,
crossed: vec![],
extra_offsets: vec![],
k_total: max_clusters,
primary_q: 1,
primary_slope_cols: vec![],
diagonal_theta: compute_diagonal_theta(1, &[]), diagonal_run_len: compute_diagonal_run_len(1, &[]), extra_slopes_any: false,
extra_q: vec![],
extra_slope_cols: vec![],
primary_slope_scales: vec![],
extra_slope_scales: vec![],
}
}
pub fn from_cluster_spec(
cluster: &crate::ModelSpec,
max_n: usize,
slope_cols: &[usize],
) -> Self {
Self::from_cluster_spec_ext(cluster, max_n, slope_cols, &[])
}
pub fn from_cluster_spec_ext(
cluster: &crate::ModelSpec,
max_n: usize,
slope_cols: &[usize],
extra_slope_cols: &[Vec<usize>],
) -> Self {
use crate::GroupingRelation;
let re = cluster
.re
.as_ref()
.expect("LmmGroupings::from_cluster_spec_ext requires re: Some (mixed model)");
let n_primary = re.sizing.n_clusters_at(max_n);
let q_p = 1 + slope_cols.len();
let prim_width = q_p * n_primary;
let n_extras = re.extra_groupings.len();
let base_theta = q_p * (q_p + 1) / 2;
let extra_qs: Vec<usize> = re
.extra_groupings
.iter()
.map(|gs| 1 + gs.slopes.len())
.collect();
let mut vech_starts = vec![0usize; n_extras];
let mut cursor = base_theta;
for (g, &q_g) in extra_qs.iter().enumerate() {
vech_starts[g] = cursor;
cursor += q_g * (q_g + 1) / 2;
}
let mut nested_per_parent = 0usize;
let mut nested = None;
let mut extra_offsets = vec![0usize; n_extras];
for (g, gs) in re.extra_groupings.iter().enumerate() {
if let GroupingRelation::NestedWithin { n_per_parent } = gs.relation {
nested_per_parent = (n_per_parent).max(1) as usize;
nested = Some(NestedFactor {
vech_start: vech_starts[g],
q: extra_qs[g],
decl: g,
});
extra_offsets[g] = prim_width; }
}
let q_nested = nested.map(|nf| nf.q).unwrap_or(0);
let mut off = prim_width + n_primary * nested_per_parent * q_nested;
let mut crossed = Vec::new();
for (g, gs) in re.extra_groupings.iter().enumerate() {
if let GroupingRelation::Crossed { n_clusters } = gs.relation {
let k = (n_clusters).max(1) as usize;
let q_g = extra_qs[g];
crossed.push(CrossedFactor {
vech_start: vech_starts[g],
q: q_g,
n_levels: k,
decl: g,
});
extra_offsets[g] = off;
off += k * q_g; }
}
let extra_slopes_any = extra_qs.iter().any(|&q| q > 1);
let extra_slope_cols: Vec<Vec<usize>> = (0..n_extras)
.map(|g| {
let v = extra_slope_cols.get(g).cloned().unwrap_or_default();
debug_assert!(v.is_empty() || v.len() == extra_qs[g] - 1);
v
})
.collect();
LmmGroupings {
n_primary,
nested_per_parent,
nested,
crossed,
extra_offsets,
k_total: off,
primary_q: q_p,
primary_slope_cols: slope_cols.to_vec(),
diagonal_theta: compute_diagonal_theta(q_p, &extra_qs),
diagonal_run_len: compute_diagonal_run_len(q_p, &extra_qs),
extra_slopes_any,
extra_q: extra_qs,
primary_slope_scales: vec![1.0; slope_cols.len()],
extra_slope_scales: extra_slope_cols
.iter()
.map(|v| vec![1.0; v.len()])
.collect(),
extra_slope_cols,
}
}
pub fn set_slope_scales(&mut self, x: MatRef<'_, f64>, weights: Option<&[f64]>) {
for d in 0..self.primary_slope_cols.len() {
self.primary_slope_scales[d] = rms_column_scale(x, self.primary_slope_cols[d], weights);
}
for e in 0..self.extra_slope_cols.len() {
for d in 0..self.extra_slope_cols[e].len() {
self.extra_slope_scales[e][d] =
rms_column_scale(x, self.extra_slope_cols[e][d], weights);
}
}
}
pub fn block_row_scale(&self, b: usize, r: usize) -> f64 {
if r == 0 {
return 1.0;
}
if b == 0 {
self.primary_slope_scales[r - 1]
} else {
self.extra_slope_scales[b - 1][r - 1]
}
}
pub fn theta_row_scales(&self) -> Vec<f64> {
let mut out = vec![0.0; self.n_theta()];
self.fill_theta_row_scales(&mut out);
out
}
pub fn fill_theta_row_scales(&self, out: &mut [f64]) {
debug_assert_eq!(out.len(), self.n_theta());
let mut i = 0;
for (b, &q) in std::iter::once(&self.primary_q)
.chain(self.extra_q.iter())
.enumerate()
{
for c in 0..q {
for r in c..q {
out[i] = self.block_row_scale(b, r);
i += 1;
}
}
}
}
pub fn any_slope_scaled(&self) -> bool {
self.primary_slope_scales.iter().any(|&s| s != 1.0)
|| self
.extra_slope_scales
.iter()
.any(|v| v.iter().any(|&s| s != 1.0))
}
pub fn n_theta(&self) -> usize {
let prim = self.primary_q * (self.primary_q + 1) / 2;
let nested = self.nested.map(|nf| nf.q * (nf.q + 1) / 2).unwrap_or(0);
let crossed: usize = self.crossed.iter().map(|cf| cf.q * (cf.q + 1) / 2).sum();
prim + nested + crossed
}
pub fn k_family(&self) -> usize {
let q_nested = self.nested.map(|nf| nf.q).unwrap_or(0);
self.n_primary * self.primary_q + self.n_primary * self.nested_per_parent * q_nested
}
pub fn k_crossed(&self) -> usize {
self.k_total - self.k_family()
}
pub fn structured_extras_eligible(&self) -> bool {
!self.extra_offsets.is_empty() && self.primary_q + self.nested_per_parent <= MAX_PRIMARY_Q
}
pub fn diagonal_theta(&self) -> &[usize] {
&self.diagonal_theta
}
pub fn diagonal_has_nonzero_below(&self, k: usize, theta: &[f64]) -> bool {
let run_len = self.diagonal_run_len[k];
if run_len <= 1 {
return false;
}
let ti = self.diagonal_theta[k];
theta[ti + 1..ti + run_len].iter().any(|&v| v != 0.0)
}
pub fn blind_theta_and_bounds(&self) -> (Vec<f64>, Vec<f64>, Vec<f64>) {
let n = self.n_theta();
let mut theta = vec![THETA0; n];
let mut lower = vec![0.0; n];
let upper = vec![THETA_HI; n];
let diag = self.diagonal_theta();
for i in 0..n {
if !diag.contains(&i) {
theta[i] = 0.0; lower[i] = -THETA_HI; }
}
(theta, lower, upper)
}
}
pub struct LmmFitScratch<T = f64> {
pub fam_a: Vec<T>,
pub bt: Vec<T>,
pub tail: Vec<T>,
pub tail_l: Vec<T>,
pub lam_x: Vec<T>,
pub prim_lam: Vec<T>,
pub prim_gram: Vec<T>,
pub fam_gram: Vec<f64>,
pub collapse_stage: Vec<f64>,
pub comb: Vec<T>,
pub a_inv: Vec<T>,
pub collapse_n_active: usize,
pub syrk_scratch: Vec<f64>,
pub factor: Mat<f64>,
pub betas: Vec<f64>,
pub var_diag: Vec<f64>,
pub t_sq: Vec<f64>,
pub u: Vec<f64>,
pub ranef_u: Vec<f64>,
pub ranef_ok: bool,
pub ranef_ux: Vec<f64>,
pub ranef_rhs: Vec<f64>,
pub sigma_sq: T,
pub joint_xtvix: Mat<f64>,
pub joint_k_inv: Mat<f64>,
pub joint_sigma_t_chol: Mat<f64>,
pub joint_rhs: Vec<f64>,
pub blocked_lam: Vec<f64>,
pub blocked_g: Vec<f64>,
pub blocked_tmp: Vec<f64>,
pub blocked_p: Vec<f64>,
pub blocked_l: Vec<f64>,
}
impl<T: Scalar> LmmFitScratch<T> {
pub fn new(p: usize, max_clusters: usize) -> Self {
Self::with_groupings(p, &LmmGroupings::single(max_clusters))
}
pub fn with_groupings(p: usize, g: &LmmGroupings) -> Self {
let m = p + 1;
let w = g.primary_q + g.nested_per_parent; let t_dim = g.k_crossed() + m;
let q2 = if g.primary_q > 1 || g.extra_slopes_any {
g.primary_q * g.primary_q
} else {
0
};
let npairs = if g.primary_q == 1 { w * (w + 1) / 2 } else { 0 };
let blocked_kk = if g.extra_slopes_any {
g.k_total * g.k_total
} else {
0
};
let blocked_dim = if g.extra_slopes_any { g.k_total + m } else { 0 };
LmmFitScratch {
fam_a: vec![T::ZERO; w * w],
bt: vec![T::ZERO; g.n_primary * w * t_dim],
tail: vec![T::ZERO; t_dim * t_dim],
tail_l: vec![T::ZERO; t_dim * t_dim],
lam_x: vec![T::ZERO; g.k_crossed()],
prim_lam: vec![T::ZERO; q2],
prim_gram: vec![T::ZERO; q2],
fam_gram: vec![0.0; npairs * t_dim * t_dim],
collapse_stage: vec![0.0; w * t_dim],
comb: vec![
T::ZERO;
if npairs > 0 {
(t_dim * t_dim).max(w)
} else {
0
}
],
a_inv: vec![T::ZERO; if npairs > 0 { w * w } else { 0 }],
collapse_n_active: 0,
syrk_scratch: if T::IS_F64 {
Vec::new()
} else {
vec![0.0; 2 * g.n_primary * w * t_dim + t_dim * t_dim]
},
factor: Mat::zeros(m, m),
betas: vec![0.0; p],
var_diag: vec![0.0; p],
t_sq: vec![0.0; p],
u: vec![0.0; p],
ranef_u: vec![0.0; g.k_total],
ranef_ok: false,
ranef_ux: vec![0.0; g.k_crossed()],
ranef_rhs: vec![0.0; w],
sigma_sq: T::from_f64(f64::NAN),
joint_xtvix: Mat::zeros(p, p),
joint_k_inv: Mat::zeros(p, p),
joint_sigma_t_chol: Mat::zeros(p, p),
joint_rhs: vec![0.0; p],
blocked_lam: vec![0.0; blocked_kk],
blocked_g: vec![0.0; blocked_kk],
blocked_tmp: vec![0.0; blocked_kk],
blocked_p: vec![0.0; blocked_dim * blocked_dim],
blocked_l: vec![0.0; blocked_dim * blocked_dim],
}
}
}
pub struct LmmWorkspace {
pub suff: LmmSuffStats,
pub fit: LmmFitScratch,
pub solver: Bobyqa,
pub theta: Vec<f64>,
pub lower: Vec<f64>,
pub upper: Vec<f64>,
pub(crate) dual_scratch: Option<Box<LmmDualScratch>>,
pub(crate) hyper_scratch: Option<Box<LmmHyperScratch>>,
}
impl LmmWorkspace {
#[cfg(test)]
pub fn new(p: usize, max_clusters: usize) -> Self {
Self::with_groupings(p, LmmGroupings::single(max_clusters))
}
#[cfg(test)]
pub fn for_cluster_spec(
p: usize,
cluster: &crate::ModelSpec,
max_n: usize,
slope_cols: &[usize],
) -> Self {
Self::for_cluster_spec_ext(p, cluster, max_n, slope_cols, &[])
}
pub fn for_cluster_spec_ext(
p: usize,
cluster: &crate::ModelSpec,
max_n: usize,
slope_cols: &[usize],
extra_slope_cols: &[Vec<usize>],
) -> Self {
let groupings =
LmmGroupings::from_cluster_spec_ext(cluster, max_n, slope_cols, extra_slope_cols);
let n_theta = groupings.n_theta();
let blind_theta = vec![THETA0; n_theta];
let rho_begin = (0.1
* groupings
.diagonal_theta()
.iter()
.map(|&i| blind_theta[i])
.fold(f64::INFINITY, f64::min))
.min(RHO_BEGIN);
let npt = if n_theta >= 3 {
(3 * n_theta).div_ceil(2) + 1
} else {
2 * n_theta + 1
};
let mut config = Config::new(n_theta);
config.rho_begin = rho_begin;
config.rho_end = RHO_END;
config.npt = npt;
apply_campaign_overrides(&mut config, n_theta);
let fit = LmmFitScratch::with_groupings(p, &groupings);
let (theta, lower, upper) = groupings.blind_theta_and_bounds();
LmmWorkspace {
suff: LmmSuffStats::with_groupings(p, groupings),
fit,
solver: Bobyqa::new(n_theta, config)
.expect("BOBYQA config constants are valid by construction"),
theta,
lower,
upper,
dual_scratch: None,
hyper_scratch: None,
}
}
#[cfg(test)]
pub fn with_groupings(p: usize, groupings: LmmGroupings) -> Self {
let n_theta = groupings.n_theta();
let fit = LmmFitScratch::with_groupings(p, &groupings);
let (theta, lower, upper) = groupings.blind_theta_and_bounds();
LmmWorkspace {
suff: LmmSuffStats::with_groupings(p, groupings),
fit,
solver: Bobyqa::new(n_theta, bobyqa_config(n_theta))
.expect("BOBYQA config constants are valid by construction"),
theta,
lower,
upper,
dual_scratch: None,
hyper_scratch: None,
}
}
}
pub fn primary_lambda<T: Scalar>(theta: &[T], q: usize, lam: &mut [T]) {
for v in lam[..q * q].iter_mut() {
*v = T::ZERO;
}
let mut t = 0;
for c in 0..q {
for r in c..q {
lam[r * q + c] = theta[t];
t += 1;
}
}
}
pub(crate) fn canonicalize_pinned_blocks(g: &LmmGroupings, theta: &mut [f64]) -> bool {
let mut changed = false;
let mut start = 0usize;
for &q in std::iter::once(&g.primary_q).chain(g.extra_q.iter()) {
let len = q * (q + 1) / 2;
if (2..=MAX_PRIMARY_Q).contains(&q) && canonicalize_block(&mut theta[start..start + len], q)
{
changed = true;
}
start += len;
}
changed
}
fn canonicalize_block(vech: &mut [f64], q: usize) -> bool {
let mut off = 0usize;
let mut any_pinned = false;
for d in 0..q {
any_pinned |= vech[off] <= PIN_THETA;
off += q - d;
}
if !any_pinned {
return false;
}
let mut lam = [0.0_f64; MAX_PRIMARY_Q * MAX_PRIMARY_Q];
primary_lambda(vech, q, &mut lam); let mut sig = [0.0_f64; MAX_PRIMARY_Q * MAX_PRIMARY_Q];
let mut trace = 0.0;
for r in 0..q {
for c in 0..=r {
let mut s = 0.0;
for k in 0..=c {
s += lam[r * q + k] * lam[c * q + k];
}
sig[r * q + c] = s;
}
trace += sig[r * q + r];
}
let tol = q as f64 * f64::EPSILON * trace;
let mut l = [0.0_f64; MAX_PRIMARY_Q * MAX_PRIMARY_Q];
for c in 0..q {
let mut d = sig[c * q + c];
for k in 0..c {
d -= l[c * q + k] * l[c * q + k];
}
if d <= tol || d.is_nan() {
continue; }
let lcc = d.sqrt();
l[c * q + c] = lcc;
for r in (c + 1)..q {
let mut s = sig[r * q + c];
for k in 0..c {
s -= l[r * q + k] * l[c * q + k];
}
l[r * q + c] = s / lcc;
}
}
let mut changed = false;
let mut t = 0usize;
for c in 0..q {
for r in c..q {
changed |= l[r * q + c].to_bits() != vech[t].to_bits();
vech[t] = l[r * q + c];
t += 1;
}
}
changed
}
fn primary_gram<T: Scalar>(
suff: &LmmSuffStats,
g: &LmmGroupings,
f: usize,
q: usize,
gram: &mut [T],
) {
let n_prim = g.n_primary;
for v in gram[..q * q].iter_mut() {
*v = T::ZERO;
}
gram[0] = T::from_f64(suff.counts[f]); for a in 1..q {
let s_a = g.primary_slope_scales[a - 1];
let sa = T::from_f64(suff.s[(g.primary_slope_cols[a - 1], f)] / s_a); gram[a] = sa;
gram[a * q] = sa;
for b in 1..=a {
let v = T::from_f64(suff.s[(g.primary_slope_cols[a - 1], b * n_prim + f)] / s_a);
gram[a * q + b] = v;
gram[b * q + a] = v;
}
}
}
fn assemble_primary_a<T: Scalar>(fam_a: &mut [T], stride: usize, lam: &[T], gram: &[T], q: usize) {
let mut m_r = [T::ZERO; MAX_PRIMARY_Q];
for r in 0..q {
for (e, m_re) in m_r.iter_mut().enumerate().take(q) {
let mut acc = T::ZERO;
for d in r..q {
acc += lam[d * q + r] * gram[d * q + e];
}
*m_re = acc;
}
for c in 0..=r {
let mut s = T::ZERO;
for e in c..q {
s += m_r[e] * lam[e * q + c];
}
fam_a[r * stride + c] = if r == c { T::ONE + s } else { s };
}
}
}
#[allow(clippy::too_many_arguments)] fn assemble_fam_a<T: Scalar>(
fam_a: &mut [T],
prim_gram: &mut [T],
prim_lam: &[T],
suff: &LmmSuffStats,
f: usize,
w: usize,
th_p: T,
th_n: T,
slope: bool,
) {
let g = &suff.groupings;
let np = g.nested_per_parent;
if slope {
let q = g.primary_q;
primary_gram(suff, g, f, q, prim_gram);
assemble_primary_a(fam_a, w, prim_lam, prim_gram, q); for c in 0..np {
let gcol = g.n_primary * g.primary_q + f * np + c;
let n_c = T::from_f64(suff.counts[gcol]);
for c2 in 0..np {
fam_a[(q + c) * w + (q + c2)] = T::ZERO;
}
fam_a[(q + c) * w + (q + c)] = T::ONE + th_n * th_n * n_c;
for e in 0..q {
let mut acc = T::ZERO;
for d in e..q {
let graw_d = if d == 0 {
n_c
} else {
T::from_f64(
suff.s[(g.primary_slope_cols[d - 1], gcol)]
/ g.primary_slope_scales[d - 1],
)
};
acc += prim_lam[d * q + e] * graw_d;
}
fam_a[(q + c) * w + e] = th_n * acc;
}
}
} else {
let n_f = T::from_f64(suff.counts[f]);
fam_a[0] = T::ONE + th_p * th_p * n_f;
for c in 0..np {
let gcol = g.n_primary + f * np + c;
let n_c = T::from_f64(suff.counts[gcol]);
for c2 in 0..np {
fam_a[(1 + c) * w + (1 + c2)] = T::ZERO;
}
fam_a[(1 + c) * w] = th_p * th_n * n_c;
fam_a[(1 + c) * w + (1 + c)] = T::ONE + th_n * th_n * n_c;
}
}
}
#[allow(clippy::too_many_arguments)] fn assemble_fam_b<T: Scalar>(
bt_fam: &mut [T],
lam_x: &[T],
prim_lam: &[T],
suff: &LmmSuffStats,
f: usize,
t_dim: usize,
kx: usize,
slope: bool,
th_p: T,
th_n: T,
) {
let g = &suff.groupings;
let m = suff.m;
let np = g.nested_per_parent;
if slope {
let q = g.primary_q;
let n_prim = g.n_primary;
for b in 0..kx {
let lam_b = lam_x[b];
let zxb = suff.zx.col(b).try_as_col_major().unwrap().as_slice();
let zxsb = suff.zx_slope.col(b).try_as_col_major().unwrap().as_slice();
for r in 0..q {
let mut brb = T::ZERO;
for d in r..q {
let zeta = T::from_f64(if d == 0 { zxb[f] } else { zxsb[d * n_prim + f] });
brb += prim_lam[d * q + r] * zeta;
}
bt_fam[r * t_dim + b] = lam_b * brb;
}
}
let mut s_cols: [&[f64]; MAX_PRIMARY_Q] = [&[]; MAX_PRIMARY_Q];
for (d, sc) in s_cols.iter_mut().enumerate().take(q) {
*sc = suff
.s
.col(d * n_prim + f)
.try_as_col_major()
.unwrap()
.as_slice();
}
for r in 0..q {
let bcol = &mut bt_fam[r * t_dim + kx..r * t_dim + kx + m];
for j in 0..m {
let mut brj = T::ZERO;
#[allow(clippy::needless_range_loop)]
for d in r..q {
brj += prim_lam[d * q + r] * T::from_f64(s_cols[d][j]);
}
bcol[j] = brj;
}
}
for c in 0..np {
let gcol = n_prim * q + f * np + c; let off = (q + c) * t_dim;
for b in 0..kx {
bt_fam[off + b] = th_n * lam_x[b] * T::from_f64(suff.zx[(gcol, b)]);
}
let scol = suff.s.col(gcol).try_as_col_major().unwrap().as_slice();
let bcol = &mut bt_fam[off + kx..off + kx + m];
for j in 0..m {
bcol[j] = th_n * T::from_f64(scol[j]);
}
}
} else {
let s_f = suff.s.col(f).try_as_col_major().unwrap().as_slice();
for b in 0..kx {
bt_fam[b] = th_p * lam_x[b] * T::from_f64(suff.zx[(f, b)]);
}
{
let bcol = &mut bt_fam[kx..kx + m];
for j in 0..m {
bcol[j] = th_p * T::from_f64(s_f[j]);
}
}
for c in 0..np {
let gcol = g.n_primary + f * np + c;
let off = (1 + c) * t_dim;
for b in 0..kx {
bt_fam[off + b] = th_n * lam_x[b] * T::from_f64(suff.zx[(gcol, b)]);
}
let scol = suff.s.col(gcol).try_as_col_major().unwrap().as_slice();
let bcol = &mut bt_fam[off + kx..off + kx + m];
for j in 0..m {
bcol[j] = th_n * T::from_f64(scol[j]);
}
}
}
}
fn fam_forward_solve<T: Scalar>(bt_fam: &mut [T], t_dim: usize, w: usize, fam_a: &[T]) {
for r in 0..w {
let (done, rest) = bt_fam.split_at_mut(r * t_dim);
let col_r = &mut rest[..t_dim];
for k in 0..r {
let l_rk = fam_a[r * w + k];
let col_k = &done[k * t_dim..(k + 1) * t_dim];
for t in 0..t_dim {
col_r[t] -= l_rk * col_k[t];
}
}
let l_rr = fam_a[r * w + r];
#[allow(clippy::needless_range_loop)]
for t in 0..t_dim {
col_r[t] /= l_rr;
}
}
}
pub(crate) fn recover_ranef(theta: &[f64], suff: &LmmSuffStats, fit: &mut LmmFitScratch) {
fit.ranef_ok = false;
let k = suff.groupings.k_total;
if k == 0 || suff.n_rows == 0 {
return;
}
fit.ranef_u[..k].fill(0.0);
fit.ranef_ok = if suff.groupings.extra_slopes_any {
recover_ranef_blocked(theta, suff, fit)
} else {
recover_ranef_family(theta, suff, fit)
};
}
fn recover_ranef_blocked(theta: &[f64], suff: &LmmSuffStats, fit: &mut LmmFitScratch) -> bool {
let _ = theta; let g = &suff.groupings;
let m = suff.m;
let p = m - 1;
let k = g.k_total;
let dim = k + m;
let pref = faer::MatRef::from_column_major_slice(&fit.blocked_p[..dim * dim], dim, dim);
let Ok(chol) = pref.llt(faer::Side::Lower) else {
return false;
};
let l = chol.L();
let u = &mut fit.ranef_u;
for a in 0..k {
let mut acc = l[(k + p, a)];
for j in 0..p {
acc -= l[(k + j, a)] * fit.betas[j];
}
u[a] = acc;
}
for a in (0..k).rev() {
let mut acc = u[a];
for i in (a + 1)..k {
acc -= l[(i, a)] * u[i];
}
let laa = l[(a, a)];
if !(laa.is_finite() && laa > 0.0) {
return false;
}
u[a] = acc / laa;
}
true
}
fn recover_ranef_family(theta: &[f64], suff: &LmmSuffStats, fit: &mut LmmFitScratch) -> bool {
let g = &suff.groupings;
let m = suff.m;
let p = m - 1;
let kf = g.k_family();
let kx = g.k_crossed();
let t_dim = kx + m;
let np = g.nested_per_parent;
let q_p = g.primary_q;
let w = q_p + np;
let n_prim = g.n_primary;
let th_p = theta[0];
let th_n = g.nested.map(|nf| theta[nf.vech_start]).unwrap_or(0.0);
let slope = q_p > 1;
if slope {
primary_lambda(theta, q_p, &mut fit.prim_lam);
}
{
let mut b = 0usize;
for cf in &g.crossed {
for _ in 0..cf.n_levels {
fit.lam_x[b] = theta[cf.vech_start];
b += 1;
}
}
}
let u_x = &mut fit.ranef_ux;
if kx > 0 {
let tail_ref =
faer::MatRef::from_column_major_slice(&fit.tail[..t_dim * t_dim], t_dim, t_dim);
let Ok(chol) = tail_ref.llt(faer::Side::Lower) else {
return false;
};
let l = chol.L();
for b in 0..kx {
let mut acc = l[(kx + p, b)];
for j in 0..p {
acc -= l[(kx + j, b)] * fit.betas[j];
}
u_x[b] = acc;
}
for b in (0..kx).rev() {
let mut acc = u_x[b];
for i in (b + 1)..kx {
acc -= l[(i, b)] * u_x[i];
}
let lbb = l[(b, b)];
if !(lbb.is_finite() && lbb > 0.0) {
return false;
}
u_x[b] = acc / lbb;
}
for (b, &v) in u_x.iter().enumerate() {
fit.ranef_u[kf + b] = v;
}
}
let rhs = &mut fit.ranef_rhs;
for f in 0..n_prim {
assemble_fam_a(
&mut fit.fam_a,
&mut fit.prim_gram,
&fit.prim_lam,
suff,
f,
w,
th_p,
th_n,
slope,
);
if !crate::linalg::block_chol(&mut fit.fam_a[..w * w], w) {
return false;
}
let bt_fam = &mut fit.bt[..w * t_dim];
assemble_fam_b(
bt_fam,
&fit.lam_x,
&fit.prim_lam,
suff,
f,
t_dim,
kx,
slope,
th_p,
th_n,
);
fam_forward_solve(bt_fam, t_dim, w, &fit.fam_a);
for r in 0..w {
let col = &bt_fam[r * t_dim..(r + 1) * t_dim];
let mut acc = col[kx + p];
for j in 0..p {
acc -= col[kx + j] * fit.betas[j];
}
for (b, &ux) in u_x.iter().enumerate() {
acc -= col[b] * ux;
}
rhs[r] = acc;
}
for r in (0..w).rev() {
let mut acc = rhs[r];
for (i, &solved) in rhs.iter().enumerate().take(w).skip(r + 1) {
acc -= fit.fam_a[i * w + r] * solved;
}
rhs[r] = acc / fit.fam_a[r * w + r];
}
for (r, &v) in rhs.iter().enumerate().take(q_p) {
fit.ranef_u[r * n_prim + f] = v;
}
for c in 0..np {
fit.ranef_u[q_p * n_prim + f * np + c] = rhs[q_p + c];
}
}
true
}
pub struct LmmFit {
pub sigma_sq: f64,
pub converged: bool,
pub boundary_hit: u8,
pub n_eval: usize,
#[cfg(feature = "counters")]
pub counters: crate::counters::EvalCounters,
pub joint_t_sq: f64,
pub pinned_components: u64,
pub deviance: f64,
pub pivot: f64,
pub pivot_col: u32,
}
pub(crate) fn joint_wald_chi_sq(
xtvix: MatRef<'_, f64>,
betas: &[f64],
sigma_sq: f64,
target_indices: &[u32],
mut joint_k_inv: MatMut<'_, f64>,
mut joint_sigma_t_chol: MatMut<'_, f64>,
joint_rhs: &mut [f64],
) -> f64 {
use faer::linalg::solvers::Solve;
use faer::reborrow::Reborrow;
let p = betas.len();
let k = target_indices.len();
if k == 0 || p == 0 || !(sigma_sq.is_finite() && sigma_sq > FLOAT_NEAR_ZERO) {
return f64::NAN;
}
let chol = match xtvix.llt(faer::Side::Lower) {
Ok(c) => c,
Err(_) => return f64::NAN,
};
for j in 0..p {
for i in 0..p {
joint_k_inv[(i, j)] = if i == j { 1.0 } else { 0.0 };
}
}
chol.solve_in_place(joint_k_inv.rb_mut());
for (a, &ti) in target_indices.iter().enumerate() {
let ti = ti as usize;
if ti >= p {
return f64::NAN;
}
for (b, &tj) in target_indices.iter().enumerate() {
let tj = tj as usize;
if tj >= p {
return f64::NAN;
}
joint_sigma_t_chol[(a, b)] = joint_k_inv[(ti, tj)];
}
}
let sigma_t_view = joint_sigma_t_chol.rb().submatrix(0, 0, k, k);
let chol_t = match sigma_t_view.llt(faer::Side::Lower) {
Ok(c) => c,
Err(_) => return f64::NAN,
};
for (a, &ti) in target_indices.iter().enumerate() {
joint_rhs[a] = betas[ti as usize];
}
{
let mut rhs = MatMut::from_column_major_slice_mut(&mut joint_rhs[..k], k, 1usize);
chol_t.solve_in_place(rhs.rb_mut());
}
let mut w_raw = 0.0_f64;
for (a, &ti) in target_indices.iter().enumerate() {
w_raw += betas[ti as usize] * joint_rhs[a];
}
if !w_raw.is_finite() {
return f64::NAN;
}
w_raw / sigma_sq
}
pub fn fit_lmm(
ws: &mut LmmWorkspace,
target_indices: &[u32],
theta_start: Option<&[f64]>,
) -> LmmFit {
fit_lmm_impl(ws, target_indices, theta_start, two_stage_enabled())
}
#[cfg(test)]
pub(crate) fn fit_lmm_two_stage(
ws: &mut LmmWorkspace,
target_indices: &[u32],
theta_start: Option<&[f64]>,
) -> LmmFit {
fit_lmm_impl(ws, target_indices, theta_start, true)
}
fn fit_lmm_impl(
ws: &mut LmmWorkspace,
target_indices: &[u32],
theta_start: Option<&[f64]>,
two_stage: bool,
) -> LmmFit {
let LmmWorkspace {
suff,
fit,
solver,
theta,
lower,
upper,
..
} = ws;
let p = suff.m - 1;
precompute_balanced_collapse(suff, fit);
match theta_start {
Some(ts) => {
debug_assert_eq!(ts.len(), theta.len());
let s = suff.groupings.theta_row_scales();
for ((t, &v), &sc) in theta.iter_mut().zip(ts).zip(s.iter()) {
*t = v * sc;
}
for &i in suff.groupings.diagonal_theta() {
theta[i] = theta[i].max(THETA_TRUTH_FLOOR);
}
}
None => {
for t in theta.iter_mut() {
*t = 0.0;
}
for &i in suff.groupings.diagonal_theta() {
theta[i] = THETA0;
}
}
}
let mut counters = crate::counters::EvalCounters::new();
let mut finite_evals = 0usize;
let out = if two_stage {
two_stage_minimize(suff, fit, theta, lower, upper, &mut finite_evals)
} else {
solver.minimize(
|xs| {
let d = reml_deviance(xs, suff, fit);
if d.is_finite() {
finite_evals += 1;
}
counters.record_eval(crate::counters::Stage::Two, d);
d
},
theta,
lower,
upper,
)
};
debug_assert!(out.status != Status::InvalidArgs);
let converged = matches!(out.status, Status::Converged) && finite_evals >= 2;
let has_endpoint = matches!(out.status, Status::Converged | Status::MaxFunReached);
let diag = suff.groupings.diagonal_theta();
let mut pinned = false;
let mut pinned_components = 0u64;
if has_endpoint {
for (k, &ti) in diag.iter().enumerate() {
if theta[ti] <= PIN_THETA {
theta[ti] = 0.0;
if converged {
pinned = true;
if k < u64::BITS as usize {
pinned_components |= 1u64 << k;
}
}
}
}
if canonicalize_pinned_blocks(&suff.groupings, theta) {
pinned = false;
pinned_components = 0;
for (k, &ti) in diag.iter().enumerate() {
if theta[ti] <= PIN_THETA {
theta[ti] = 0.0;
if converged {
pinned = true;
if k < u64::BITS as usize {
pinned_components |= 1u64 << k;
}
}
}
}
}
}
let dev = if has_endpoint {
reml_deviance(theta, suff, fit)
} else {
f64::INFINITY
};
let (pivot, pivot_col) = if dev.is_finite() {
crate::ols::min_pivot_ratio(fit.factor.as_ref(), p)
} else {
(f64::NAN, 0)
};
if !has_endpoint || !dev.is_finite() {
fit.ranef_ok = false;
for v in fit.betas.iter_mut() {
*v = f64::NAN;
}
for &t in target_indices {
fit.var_diag[t as usize] = f64::NAN;
fit.t_sq[t as usize] = f64::NAN;
}
return LmmFit {
sigma_sq: f64::NAN,
converged: false,
boundary_hit: 2,
n_eval: out.n_eval,
#[cfg(feature = "counters")]
counters,
joint_t_sq: f64::NAN,
pinned_components: 0,
deviance: f64::NAN,
pivot,
pivot_col: pivot_col as u32,
};
}
for j in (0..p).rev() {
let mut acc = fit.factor[(p, j)];
for k in (j + 1)..p {
acc -= fit.factor[(k, j)] * fit.betas[k];
}
fit.betas[j] = acc / fit.factor[(j, j)];
}
let sigma_sq = fit.sigma_sq;
for &tj in target_indices {
let tj = tj as usize;
for v in fit.u[..p].iter_mut() {
*v = 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 -= fit.factor[(i, k)] * fit.u[k];
}
fit.u[i] = acc / fit.factor[(i, i)];
}
let norm_sq: f64 = fit.u[..p].iter().map(|v| v * v).sum();
let vd = sigma_sq * norm_sq;
fit.var_diag[tj] = vd;
fit.t_sq[tj] = if vd.is_finite() && vd > 0.0 {
(fit.betas[tj] * fit.betas[tj]) / vd
} else {
f64::NAN
};
}
let joint_t_sq = if target_indices.is_empty() {
f64::NAN
} else {
for j in 0..p {
for i in 0..p {
let mut acc = 0.0;
for k in 0..=i.min(j) {
acc += fit.factor[(i, k)] * fit.factor[(j, k)];
}
fit.joint_xtvix[(i, j)] = acc;
}
}
joint_wald_chi_sq(
fit.joint_xtvix.as_ref(),
&fit.betas,
sigma_sq,
target_indices,
fit.joint_k_inv.as_mut(),
fit.joint_sigma_t_chol.as_mut(),
&mut fit.joint_rhs,
)
};
recover_ranef(theta, suff, fit);
LmmFit {
sigma_sq,
converged,
boundary_hit: if converged { u8::from(pinned) } else { 2 },
n_eval: out.n_eval,
#[cfg(feature = "counters")]
counters,
joint_t_sq,
pinned_components,
deviance: dev,
pivot,
pivot_col: pivot_col as u32,
}
}