use crate::lmm::LmmGroupings;
use bobyqa::Status;
use faer::dyn_stack::{MemBuffer, MemStack};
use faer::linalg::cholesky::llt::factor::{
cholesky_in_place, cholesky_in_place_scratch, LltRegularization,
};
use faer::mat::AsMatMut;
use faer::sparse::linalg::cholesky::{
factorize_symbolic_cholesky, CholeskySymbolicParams, SymbolicCholesky,
};
use faer::sparse::linalg::SupernodalThreshold;
use faer::sparse::{SparseColMat, Triplet};
use faer::{Conj, Mat, MatRef, Par, Side, Spec};
mod glmm;
#[cfg(test)]
mod tests;
pub(crate) use glmm::{fit_glmm_nb_sparse, fit_glmm_sparse};
#[cfg(test)]
use glmm::{sparse_glmm_deviance, SparseGlmmWorkspace};
#[allow(clippy::too_many_arguments)]
pub(crate) fn fit_mle_sparse(
x: &[f64],
y: &[f64],
n: usize,
p: usize,
model: &crate::ModelSpec,
cluster_ids: &[u32],
extra_ids: &[Vec<u32>],
start: Option<&crate::StartValues>,
opts: &crate::FitOptions,
) -> crate::Fit {
let re = model
.re
.as_ref()
.expect("fit_mle_sparse requires a mixed model (re: Some)");
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();
let g = LmmGroupings::from_cluster_spec_ext(model, n, &slope_cols, &extra_slope_cols);
let xm = MatRef::from_row_major_slice(x, n, p);
let sqrt_w: Option<Vec<f64>> = opts
.weights
.as_ref()
.map(|w| w.iter().map(|v| v.sqrt()).collect());
let mut ws =
SparseLmmWorkspace::new(&g, xm, cluster_ids, extra_ids, y, n, p, sqrt_w.as_deref());
let (mut solver, mut theta, lower, upper) = crate::lmm::sparse_lmm_seed(&g);
match start {
Some(s) => {
debug_assert_eq!(s.theta.len(), theta.len());
for (t, &v) in theta.iter_mut().zip(&s.theta) {
*t = v.max(crate::lmm::THETA_TRUTH_FLOOR);
}
}
None => {
for t in theta.iter_mut() {
*t = 0.0;
}
for &i in g.diagonal_theta() {
theta[i] = crate::lmm::THETA0;
}
}
}
let out = solver.minimize(
|xs| sparse_reml_deviance(xs, &mut ws),
&mut theta,
&lower,
&upper,
);
debug_assert!(out.status != Status::InvalidArgs);
let converged_status = matches!(out.status, Status::Converged);
let has_endpoint = matches!(out.status, Status::Converged | Status::MaxFunReached);
let mut pinned = false;
if has_endpoint {
for &ti in g.diagonal_theta() {
if theta[ti] <= crate::lmm::PIN_THETA {
theta[ti] = 0.0;
if converged_status {
pinned = true;
}
}
}
}
let factor_ok = has_endpoint && sparse_schur_factor(&theta, &mut ws).is_some();
let degenerate = if factor_ok {
crate::ols::chol_rank_deficient(ws.factor.as_ref(), p, crate::lmm::EPS_RANK)
} else {
true
};
let converged = converged_status && !degenerate;
let has_recovery = has_endpoint && !degenerate;
if !has_recovery {
return crate::Fit {
beta: vec![f64::NAN; p],
se: vec![f64::NAN; p],
vcov: crate::fit::nan_vcov(p),
tau2: theta.iter().map(|_| f64::NAN).collect(),
dispersion: f64::NAN,
converged: false,
varcorr: vec![],
stddev_se: vec![],
aliased: vec![false; p],
n_eval: out.n_eval,
deviance: f64::NAN,
singular: false,
};
}
let dev = sparse_reml_deviance(&theta, &mut ws);
let l = &ws.factor;
let sigma_sq = {
let lyy = l[(p, p)];
lyy * lyy / ((n - p) as f64)
};
let mut beta = vec![0.0f64; p];
for j in (0..p).rev() {
let mut acc = l[(p, j)];
for k in (j + 1)..p {
acc -= l[(k, j)] * beta[k];
}
beta[j] = acc / l[(j, j)];
}
let mut se = vec![f64::NAN; p];
let mut u = vec![0.0f64; p];
for &tj in &opts.target_indices {
let tj = tj as usize;
for v in u.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 -= l[(i, k)] * u[k];
}
u[i] = acc / l[(i, i)];
}
let norm_sq: f64 = u.iter().map(|v| v * v).sum();
let vd = sigma_sq * norm_sq;
if vd.is_finite() && vd >= 0.0 {
se[tj] = vd.sqrt();
}
}
let vcov = crate::fit::vcov_from_chol(l.as_ref(), p, &opts.target_indices, sigma_sq);
let tau2: Vec<f64> = theta.iter().map(|&t| t * t * sigma_sq).collect();
let varcorr = crate::fit::assemble_varcorr(&theta, &g, sigma_sq);
let dev = match &opts.weights {
Some(w) => dev - w.iter().map(|v| v.ln()).sum::<f64>(),
None => dev,
};
let mut fit = crate::Fit {
beta,
se,
vcov,
tau2,
dispersion: sigma_sq,
converged,
varcorr,
stddev_se: vec![],
aliased: vec![false; p],
n_eval: out.n_eval,
deviance: dev,
singular: pinned,
};
fit.singular = fit.singular || fit.has_negligible_component();
fit
}
pub(crate) fn logdet_llt(symbolic: &SymbolicCholesky<usize>, l_values: &[f64]) -> f64 {
use faer::sparse::linalg::cholesky::{supernodal::SupernodalLltRef, SymbolicCholeskyRaw};
let mut acc = 0.0f64;
let mut push = |ljj: f64| -> bool {
if ljj <= 0.0 || !ljj.is_finite() {
return false;
}
acc += ljj.ln();
true
};
match symbolic.raw() {
SymbolicCholeskyRaw::Simplicial(simp) => {
let col_ptr = simp.col_ptr();
let row_idx = simp.row_idx();
let n = col_ptr.len() - 1;
for j in 0..n {
let mut ljj = f64::NAN;
for k in col_ptr[j]..col_ptr[j + 1] {
if row_idx[k] == j {
ljj = l_values[k];
break;
}
}
if !push(ljj) {
return f64::INFINITY;
}
}
}
SymbolicCholeskyRaw::Supernodal(sup) => {
let llt = SupernodalLltRef::new(sup, l_values);
for si in 0..sup.n_supernodes() {
let panel = llt.supernode(si).val();
for j in 0..panel.ncols() {
if !push(panel[(j, j)]) {
return f64::INFINITY;
}
}
}
}
}
2.0 * acc
}
pub(crate) const TAIL_SPARSE_MIN: usize = 128;
#[cfg(test)]
thread_local! {
pub(crate) static FORCE_SPARSE_TAIL: std::cell::Cell<bool> =
const { std::cell::Cell::new(false) };
}
pub(crate) struct SparseTail {
pub(crate) symbolic: SymbolicCholesky<usize>,
pub(crate) axx: SparseColMat<usize, f64>,
pub(crate) l_values: Vec<f64>,
pub(crate) fac_mem: MemBuffer,
pub(crate) solve_mem: MemBuffer,
pub(crate) diag_slots: Vec<u32>,
pub(crate) a22_slots: Vec<u32>,
pub(crate) fam_dd_slots: Vec<u32>,
pub(crate) fam_dd_off: Vec<usize>,
pub(crate) panel: Vec<f64>,
pub(crate) dd_temp: Vec<f64>,
pub(crate) b2_temp: Vec<f64>,
pub(crate) x2: Mat<f64>,
}
pub(crate) struct SparseLmmWorkspace {
pub(crate) g: LmmGroupings,
pub(crate) ztxy: Vec<f64>,
pub(crate) cxy: Mat<f64>,
pub(crate) pk_fam: Vec<f64>,
pub(crate) pk_a21: Vec<f64>,
pub(crate) pk_a22: Vec<f64>,
pub(crate) a21_blk: Vec<u32>,
pub(crate) a21_off: Vec<usize>,
pub(crate) pk_a21_off: Vec<usize>,
pub(crate) a22_pairs: Vec<[u32; 2]>,
pub(crate) lam_blocks: Vec<LamBlock>,
pub(crate) lam_small: Vec<f64>,
pub(crate) fam_w: usize,
pub(crate) fam_a: Vec<f64>,
pub(crate) l21: Vec<f64>,
pub(crate) u1: Vec<f64>,
pub(crate) s22: Mat<f64>,
pub(crate) u2: Vec<f64>,
pub(crate) tail: Option<SparseTail>,
pub(crate) factor: Mat<f64>,
pub(crate) tail_llt_mem: MemBuffer,
pub(crate) factor_llt_mem: MemBuffer,
pub(crate) m: usize,
pub(crate) p: usize,
pub(crate) n: usize,
}
pub(crate) struct LamBlock {
start: usize,
stride: usize,
q: usize,
lam_off: usize,
}
impl SparseLmmWorkspace {
#[allow(clippy::too_many_arguments)] pub(crate) fn new(
g: &LmmGroupings,
x: MatRef<f64>,
cluster_ids: &[u32],
extra_ids: &[Vec<u32>],
y: &[f64],
n: usize,
p: usize,
sqrt_w: Option<&[f64]>,
) -> Self {
let m = p + 1;
let q_p = g.primary_q;
let n_prim = g.n_primary;
let mut lam_blocks: Vec<LamBlock> = Vec::new();
let mut lam_off = 0usize;
for f in 0..n_prim {
lam_blocks.push(LamBlock {
start: f,
stride: n_prim,
q: q_p,
lam_off,
});
}
lam_off += q_p * q_p;
if let Some(nf) = g.nested {
let q_n = nf.q;
let np = g.nested_per_parent;
let prim_width = q_p * n_prim;
for f in 0..n_prim {
for c in 0..np {
let ic = prim_width + (f * np + c) * q_n;
lam_blocks.push(LamBlock {
start: ic,
stride: 1,
q: q_n,
lam_off,
});
}
}
lam_off += q_n * q_n;
}
for cf in &g.crossed {
let off = g.extra_offsets[cf.decl];
for c in 0..cf.n_levels {
let ic = off + c * cf.q;
lam_blocks.push(LamBlock {
start: ic,
stride: 1,
q: cf.q,
lam_off,
});
}
lam_off += cf.q * cf.q;
}
let lam_small = vec![0.0f64; lam_off.max(1)];
let q_nested = g.nested.map(|nf| nf.q).unwrap_or(0);
let np = g.nested_per_parent;
let fam_w = q_p + np * q_nested;
let kf = g.k_family();
let e = g.k_crossed();
#[cfg(test)]
let force_sparse = FORCE_SPARSE_TAIL.with(|c| c.get());
#[cfg(not(test))]
let force_sparse = false;
let sparse_tail = e > 0 && (e > TAIL_SPARSE_MIN || force_sparse);
let tail_llt_mem = MemBuffer::new(cholesky_in_place_scratch::<f64>(
if sparse_tail { 0 } else { e },
Par::Seq,
Spec::default(),
));
let q_n = q_nested;
let cb0 = n_prim + n_prim * np; let mut cf_base = Vec::with_capacity(g.crossed.len());
{
let mut b = cb0;
for cf in &g.crossed {
cf_base.push(b);
b += cf.n_levels;
}
}
let n_extras = extra_ids.len();
let mut seg_off = vec![0usize; n_extras];
{
let mut o = q_p;
for (ei, so) in seg_off.iter_mut().enumerate() {
*so = o;
o += g.extra_q[ei];
}
}
let nested_decl = g.nested.map(|nf| nf.decl);
let mut decl_base = vec![usize::MAX; n_extras];
for (ci, cf) in g.crossed.iter().enumerate() {
decl_base[cf.decl] = cf_base[ci];
}
let mut fam_crossed: Vec<Vec<u32>> = vec![Vec::new(); n_prim];
let mut pair_seen = std::collections::HashSet::<(u32, u32)>::new();
let mut a22_pairs: Vec<[u32; 2]> = Vec::new();
let mut row_cb: Vec<u32> = Vec::with_capacity(g.crossed.len());
for i in 0..n {
let f = cluster_ids[i] as usize;
row_cb.clear();
for cf in &g.crossed {
row_cb.push((decl_base[cf.decl] + extra_ids[cf.decl][i] as usize) as u32);
}
for (ai, &bi) in row_cb.iter().enumerate() {
fam_crossed[f].push(bi);
for &bj in &row_cb[..=ai] {
let key = if bi >= bj { (bi, bj) } else { (bj, bi) };
if pair_seen.insert(key) {
a22_pairs.push([key.0, key.1]);
}
}
}
}
drop(pair_seen);
for v in fam_crossed.iter_mut() {
v.sort_unstable();
v.dedup();
}
a22_pairs.sort_unstable_by_key(|pr| (pr[1], pr[0]));
let fam_len = q_p * q_p + np * (q_n * q_p + q_n * q_n);
let mut pk_fam = vec![0.0f64; n_prim * fam_len];
let mut a21_blk: Vec<u32> = Vec::new();
let mut a21_off = Vec::with_capacity(n_prim + 1);
let mut pk_a21_off = Vec::with_capacity(n_prim + 1);
let mut a21_slab_off: Vec<usize> = Vec::new();
let mut pk_a21_len = 0usize;
for list in &fam_crossed {
a21_off.push(a21_blk.len());
pk_a21_off.push(pk_a21_len);
for &bi in list {
a21_blk.push(bi);
a21_slab_off.push(pk_a21_len);
pk_a21_len += lam_blocks[bi as usize].q * fam_w;
}
}
a21_off.push(a21_blk.len());
pk_a21_off.push(pk_a21_len);
let mut pk_a21 = vec![0.0f64; pk_a21_len];
let mut a22_off_map = std::collections::HashMap::<(u32, u32), usize>::new();
let mut pk_a22_len = 0usize;
for pr in &a22_pairs {
a22_off_map.insert((pr[0], pr[1]), pk_a22_len);
pk_a22_len += lam_blocks[pr[0] as usize].q * lam_blocks[pr[1] as usize].q;
}
let mut pk_a22 = vec![0.0f64; pk_a22_len];
let mut ztxy = vec![0.0f64; g.k_total * m];
let mut cxy = Mat::<f64>::zeros(m, m);
let mut row: Vec<(usize, f64)> =
Vec::with_capacity(g.primary_q + g.extra_q.iter().sum::<usize>());
for i in 0..n {
let sw = sqrt_w.map_or(1.0, |w| w[i]);
row.clear();
for_each_z_entry(g, x, cluster_ids, extra_ids, i, sqrt_w, |col, v| {
row.push((col, v))
});
for &(ca, va) in &row {
for j in 0..p {
ztxy[ca * m + j] += va * (sw * x[(i, j)]);
}
ztxy[ca * m + p] += va * (sw * y[i]);
}
for a in 0..m {
let wa = if a < p { sw * x[(i, a)] } else { sw * y[i] };
for b in 0..m {
let wb = if b < p { sw * x[(i, b)] } else { sw * y[i] };
cxy[(a, b)] += wa * wb;
}
}
let f = cluster_ids[i] as usize;
let vp = &row[..q_p];
let fam = &mut pk_fam[f * fam_len..(f + 1) * fam_len];
for a in 0..q_p {
for b in 0..q_p {
fam[a * q_p + b] += vp[a].1 * vp[b].1;
}
}
let child = nested_decl.map(|d| {
let gc = extra_ids[d][i] as usize;
debug_assert!(
gc >= f * np && gc < (f + 1) * np,
"nested ids are parent-padded"
);
(d, gc - f * np)
});
if let Some((d, c)) = child {
let vc = &row[seg_off[d]..seg_off[d] + q_n];
let b2 = q_p * q_p + c * (q_n * q_p + q_n * q_n);
for a in 0..q_n {
for b in 0..q_p {
fam[b2 + a * q_p + b] += vc[a].1 * vp[b].1;
}
}
let b3 = b2 + q_n * q_p;
for a in 0..q_n {
for b in 0..q_n {
fam[b3 + a * q_n + b] += vc[a].1 * vc[b].1;
}
}
}
let fam_list = &fam_crossed[f];
for (ci, cf) in g.crossed.iter().enumerate() {
let bi = (cf_base[ci] + extra_ids[cf.decl][i] as usize) as u32;
let q = cf.q;
let vx = &row[seg_off[cf.decl]..seg_off[cf.decl] + q];
let idx = fam_list
.binary_search(&bi)
.expect("row's crossed block is in its family clique");
let slab0 = a21_slab_off[a21_off[f] + idx];
for a in 0..q {
for b in 0..q_p {
pk_a21[slab0 + a * q_p + b] += vx[a].1 * vp[b].1;
}
}
if let Some((d, c)) = child {
let vc = &row[seg_off[d]..seg_off[d] + q_n];
let base = slab0 + q * q_p + c * q * q_n;
for a in 0..q {
for b in 0..q_n {
pk_a21[base + a * q_n + b] += vx[a].1 * vc[b].1;
}
}
}
let off = a22_off_map[&(bi, bi)];
for a in 0..q {
for b in 0..q {
pk_a22[off + a * q + b] += vx[a].1 * vx[b].1;
}
}
for (cj, cf2) in g.crossed.iter().enumerate().take(ci) {
let bj = (cf_base[cj] + extra_ids[cf2.decl][i] as usize) as u32;
let vx2 = &row[seg_off[cf2.decl]..seg_off[cf2.decl] + cf2.q];
let (bh, vh, qh, bl, vl, ql) = if bi >= bj {
(bi, vx, q, bj, vx2, cf2.q)
} else {
(bj, vx2, cf2.q, bi, vx, q)
};
let off = a22_off_map[&(bh, bl)];
for a in 0..qh {
for b in 0..ql {
pk_a22[off + a * ql + b] += vh[a].1 * vl[b].1;
}
}
}
}
}
drop(a22_off_map);
let tail = if sparse_tail {
Some(build_sparse_tail(
&lam_blocks,
&fam_crossed,
&a22_pairs,
kf,
e,
m,
fam_w,
))
} else {
None
};
Self {
g: g.clone(),
ztxy,
cxy,
pk_fam,
pk_a21,
pk_a22,
a21_blk,
a21_off,
pk_a21_off,
a22_pairs,
lam_blocks,
lam_small,
fam_w,
fam_a: vec![0.0f64; fam_w * fam_w],
l21: if sparse_tail {
Vec::new()
} else {
vec![0.0f64; e * kf]
},
u1: vec![0.0f64; kf * m],
s22: if sparse_tail {
Mat::zeros(0, 0)
} else {
Mat::zeros(e, e)
},
u2: vec![0.0f64; e * m],
tail,
factor: Mat::zeros(m, m),
tail_llt_mem,
factor_llt_mem: MemBuffer::new(cholesky_in_place_scratch::<f64>(
m,
Par::Seq,
Spec::default(),
)),
m,
p,
n,
}
}
}
fn build_sparse_tail(
lam_blocks: &[LamBlock],
fam_crossed: &[Vec<u32>],
a22_pairs: &[[u32; 2]],
kf: usize,
e: usize,
m: usize,
fam_w: usize,
) -> SparseTail {
let mut blk_pairs: Vec<(u32, u32)> = Vec::new();
for list in fam_crossed {
for (ii, &bi) in list.iter().enumerate() {
for &bj in &list[..=ii] {
blk_pairs.push((bi, bj)); }
}
}
for pr in a22_pairs {
blk_pairs.push((pr[0], pr[1]));
}
let mut seen = std::collections::HashSet::<(usize, usize)>::new();
let mut trips = Vec::<Triplet<usize, usize, f64>>::new();
for t in 0..e {
trips.push(Triplet::new(t, t, 0.0));
seen.insert((t, t));
}
for &(bi, bj) in &blk_pairs {
let (ri, rj) = (&lam_blocks[bi as usize], &lam_blocks[bj as usize]);
let (ti0, tj0) = (ri.start - kf, rj.start - kf);
for il in 0..ri.q {
for jl in 0..rj.q {
let (a, b) = (ti0 + il, tj0 + jl);
let key = if a >= b { (a, b) } else { (b, a) };
if seen.insert(key) {
trips.push(Triplet::new(key.0, key.1, 0.0));
}
}
}
}
let axx = SparseColMat::<usize, f64>::try_new_from_triplets(e, e, &trips)
.expect("S22 pattern triplets well-formed");
let symbolic = factorize_symbolic_cholesky(
axx.symbolic(),
Side::Lower,
Default::default(), CholeskySymbolicParams {
supernodal_flop_ratio_threshold: SupernodalThreshold::AUTO,
..Default::default()
},
)
.expect("S22 symbolic factorization");
let mut slot = std::collections::HashMap::<(usize, usize), u32>::new();
{
let sym = axx.symbolic();
let col_ptr = sym.col_ptr();
let row_idx = sym.row_idx();
for j in 0..e {
for (k, &ri) in row_idx
.iter()
.enumerate()
.take(col_ptr[j + 1])
.skip(col_ptr[j])
{
slot.insert((ri, j), k as u32);
}
}
}
let slot_of = |t: (usize, usize)| -> u32 {
*slot
.get(&t)
.unwrap_or_else(|| panic!("S22 pattern missing entry ({}, {})", t.0, t.1))
};
let diag_slots: Vec<u32> = (0..e).map(|t| slot_of((t, t))).collect();
let mut a22_slots = Vec::new();
for pr in a22_pairs {
let (ri, rj) = (&lam_blocks[pr[0] as usize], &lam_blocks[pr[1] as usize]);
let (ti0, tj0) = (ri.start - kf, rj.start - kf);
for jl in 0..rj.q {
for il in 0..ri.q {
if ti0 + il >= tj0 + jl {
a22_slots.push(slot_of((ti0 + il, tj0 + jl)));
}
}
}
}
let mut fam_dd_slots = Vec::new();
let mut fam_dd_off = Vec::with_capacity(fam_crossed.len() + 1);
fam_dd_off.push(0);
let mut cols: Vec<usize> = Vec::new();
let mut max_ef = 0usize;
for list in fam_crossed {
cols.clear();
for &bi in list {
let b = &lam_blocks[bi as usize];
let t0 = b.start - kf;
cols.extend(t0..t0 + b.q);
}
max_ef = max_ef.max(cols.len());
for bl in 0..cols.len() {
for a in bl..cols.len() {
fam_dd_slots.push(slot_of((cols[a], cols[bl])));
}
}
fam_dd_off.push(fam_dd_slots.len());
}
let l_values = vec![0.0f64; symbolic.len_val()];
let fac_mem =
MemBuffer::new(symbolic.factorize_numeric_llt_scratch::<f64>(Par::Seq, Spec::default()));
let solve_mem = MemBuffer::new(symbolic.solve_in_place_scratch::<f64>(m, Par::Seq));
SparseTail {
symbolic,
axx,
l_values,
fac_mem,
solve_mem,
diag_slots,
a22_slots,
fam_dd_slots,
fam_dd_off,
panel: vec![0.0f64; fam_w * max_ef],
dd_temp: vec![0.0f64; max_ef * max_ef],
b2_temp: vec![0.0f64; max_ef * m],
x2: Mat::zeros(e, m),
}
}
#[inline]
fn for_each_z_entry(
g: &LmmGroupings,
x: MatRef<f64>,
cluster_ids: &[u32],
extra_ids: &[Vec<u32>],
i: usize,
sqrt_w: Option<&[f64]>,
mut emit: impl FnMut(usize, f64),
) {
let sw = sqrt_w.map_or(1.0, |w| w[i]);
let f = cluster_ids[i] as usize;
emit(f, sw);
for (k, &col) in g.primary_slope_cols.iter().enumerate() {
emit((k + 1) * g.n_primary + f, sw * x[(i, col)]);
}
for (e, ids_e) in extra_ids.iter().enumerate() {
let q_g = g.extra_q[e];
let off = g.extra_offsets[e] + ids_e[i] as usize * q_g;
emit(off, sw);
for (c, &col) in g.extra_slope_cols[e].iter().enumerate() {
emit(off + 1 + c, sw * x[(i, col)]);
}
}
}
#[cfg(test)]
pub(crate) fn build_sparse_z(
g: &LmmGroupings,
x: MatRef<f64>,
cluster_ids: &[u32],
extra_ids: &[Vec<u32>],
n: usize,
) -> SparseColMat<usize, f64> {
let mut trips: Vec<Triplet<usize, usize, f64>> =
Vec::with_capacity(n * (g.primary_q + extra_ids.len()));
for i in 0..n {
for_each_z_entry(g, x, cluster_ids, extra_ids, i, None, |col, v| {
trips.push(Triplet::new(i, col, v));
});
}
SparseColMat::try_new_from_triplets(n, g.k_total, &trips).expect("Z triplets well-formed")
}
pub(super) fn fill_lambda_small(theta: &[f64], g: &LmmGroupings, lam_small: &mut [f64]) {
let q_p = g.primary_q;
let mut off = 0usize;
crate::lmm::primary_lambda(theta, q_p, &mut lam_small[off..off + q_p * q_p]);
off += q_p * q_p;
if let Some(nf) = g.nested {
let q_n = nf.q;
crate::lmm::primary_lambda(
&theta[nf.vech_start..],
q_n,
&mut lam_small[off..off + q_n * q_n],
);
off += q_n * q_n;
}
for cf in &g.crossed {
let q = cf.q;
crate::lmm::primary_lambda(&theta[cf.vech_start..], q, &mut lam_small[off..off + q * q]);
off += q * q;
}
}
#[inline]
#[allow(clippy::too_many_arguments)]
fn fold_packed_col(
gblk: &[f64],
q_r: usize,
q_c: usize,
lam: &[f64],
lo_r: usize,
lo_c: usize,
jl: usize,
mut sink: impl FnMut(usize, f64),
) {
for il in 0..q_r {
let mut acc = 0.0;
for a in il..q_r {
let la = lam[lo_r + a * q_r + il];
let mut inner = 0.0;
for b in jl..q_c {
inner += gblk[a * q_c + b] * lam[lo_c + b * q_c + jl];
}
acc += la * inner;
}
sink(il, acc);
}
}
pub(crate) fn sparse_reml_deviance(theta: &[f64], ws: &mut SparseLmmWorkspace) -> f64 {
let p = ws.p;
let log_lzz_sq = match sparse_schur_factor(theta, ws) {
Some(v) => v,
None => return f64::INFINITY, };
let l = &ws.factor;
let mut log_lxx_sq = 0.0_f64;
for j in 0..p {
let ljj = l[(j, j)];
if !(ljj.is_finite() && ljj > 0.0) {
return f64::INFINITY;
}
log_lxx_sq += ljj.ln();
}
log_lxx_sq *= 2.0;
let lyy = l[(p, p)];
let df = (ws.n - p) as f64;
let sigma_sq = lyy * lyy / df;
if !(sigma_sq.is_finite() && sigma_sq > 0.0) {
return f64::INFINITY;
}
log_lzz_sq + log_lxx_sq + df * sigma_sq.ln()
}
fn schur_phase_b(buf: &mut [f64], col_len: usize, w: usize, fam_a: &[f64]) {
for c in 0..w {
let (done, rest) = buf.split_at_mut(c * col_len);
let col_c = &mut rest[..col_len];
for k in 0..c {
let l_ck = fam_a[c * w + k];
let col_k = &done[k * col_len..(k + 1) * col_len];
for (x, &yv) in col_c.iter_mut().zip(col_k) {
*x -= l_ck * yv;
}
}
let inv_cc = 1.0 / fam_a[c * w + c];
for v in col_c.iter_mut() {
*v *= inv_cc;
}
}
}
fn sparse_schur_factor(theta: &[f64], ws: &mut SparseLmmWorkspace) -> Option<f64> {
fill_lambda_small(theta, &ws.g, &mut ws.lam_small);
let SparseLmmWorkspace {
g,
ztxy,
cxy,
pk_fam,
pk_a21,
pk_a22,
a21_blk,
a21_off,
pk_a21_off,
a22_pairs,
lam_blocks,
lam_small,
fam_w,
fam_a,
l21,
u1,
s22,
u2,
tail,
factor,
tail_llt_mem,
factor_llt_mem,
m,
..
} = ws;
let m = *m;
let w = *fam_w;
let n_prim = g.n_primary;
let q_p = g.primary_q;
let np = g.nested_per_parent;
let q_n = g.nested.map(|nf| nf.q).unwrap_or(0);
let kf = g.k_family();
let e = g.k_crossed();
let fam_len = q_p * q_p + np * (q_n * q_p + q_n * q_n);
let cb0 = n_prim + n_prim * np;
if e > 0 {
for br in &lam_blocks[cb0..] {
let (q, t0) = (br.q, br.start - kf);
for il in 0..q {
for c in 0..m {
u2[c * e + t0 + il] = 0.0;
}
for a in il..q {
let la = lam_small[br.lam_off + a * q + il];
let zr = (br.start + a) * m;
for c in 0..m {
u2[c * e + t0 + il] += la * ztxy[zr + c];
}
}
}
}
if let Some(tail) = tail.as_mut() {
let SparseTail {
axx,
a22_slots,
diag_slots,
..
} = tail;
let (_, vals) = axx.parts_mut();
vals.fill(0.0);
let mut cur = 0usize;
let mut si = 0usize;
for pr in a22_pairs.iter() {
let bri = &lam_blocks[pr[0] as usize];
let bcj = &lam_blocks[pr[1] as usize];
let (ti0, tj0) = (bri.start - kf, bcj.start - kf);
let blk = &pk_a22[cur..cur + bri.q * bcj.q];
cur += bri.q * bcj.q;
for jl in 0..bcj.q {
fold_packed_col(
blk,
bri.q,
bcj.q,
lam_small,
bri.lam_off,
bcj.lam_off,
jl,
|il, v| {
if ti0 + il >= tj0 + jl {
vals[a22_slots[si] as usize] += v;
si += 1;
}
},
);
}
}
debug_assert_eq!(si, a22_slots.len(), "a22 slot replay exhausted");
for &s in diag_slots.iter() {
vals[s as usize] += 1.0;
}
}
}
let mut log_lzz_half = 0.0_f64;
for f in 0..n_prim {
let fam_gram = &pk_fam[f * fam_len..(f + 1) * fam_len];
let pb = &lam_blocks[f];
let lo_p = pb.lam_off;
for jl in 0..q_p {
fold_packed_col(
&fam_gram[..q_p * q_p],
q_p,
q_p,
lam_small,
lo_p,
lo_p,
jl,
|il, v| {
if il >= jl {
fam_a[il * w + jl] = if il == jl { v + 1.0 } else { v };
}
},
);
}
for c in 0..np {
let lo_n = lam_blocks[n_prim + f * np + c].lam_off;
let b2 = q_p * q_p + c * (q_n * q_p + q_n * q_n);
let rc = q_p + c * q_n;
for jl in 0..q_p {
fold_packed_col(
&fam_gram[b2..b2 + q_n * q_p],
q_n,
q_p,
lam_small,
lo_n,
lo_p,
jl,
|il, v| {
fam_a[(rc + il) * w + jl] = v;
},
);
}
for c2 in 0..c {
let rc2 = q_p + c2 * q_n;
for jl in 0..q_n {
for il in 0..q_n {
fam_a[(rc + il) * w + (rc2 + jl)] = 0.0;
}
}
}
let b3 = b2 + q_n * q_p;
for jl in 0..q_n {
fold_packed_col(
&fam_gram[b3..b3 + q_n * q_n],
q_n,
q_n,
lam_small,
lo_n,
lo_n,
jl,
|il, v| {
if il >= jl {
fam_a[(rc + il) * w + (rc + jl)] = if il == jl { v + 1.0 } else { v };
}
},
);
}
}
for j in 0..w {
let mut d = fam_a[j * w + j];
for k in 0..j {
let v = fam_a[j * w + k];
d -= v * v;
}
if !(d.is_finite() && d > 0.0) {
return None;
}
let l = d.sqrt();
fam_a[j * w + j] = l;
log_lzz_half += l.ln();
for i in (j + 1)..w {
let mut v = fam_a[i * w + j];
for k in 0..j {
v -= fam_a[i * w + k] * fam_a[j * w + k];
}
fam_a[i * w + j] = v / l;
}
}
let fb = f * w;
for r in 0..w {
let (blkr, il) = if r < q_p {
(pb, r)
} else {
let rr = r - q_p;
(&lam_blocks[n_prim + f * np + rr / q_n], rr % q_n)
};
for c in 0..m {
u1[c * kf + fb + r] = 0.0;
}
for a in il..blkr.q {
let la = lam_small[blkr.lam_off + a * blkr.q + il];
let zr = (blkr.start + a * blkr.stride) * m;
for c in 0..m {
u1[c * kf + fb + r] += la * ztxy[zr + c];
}
}
}
for r in 0..w {
for k in 0..r {
let lrk = fam_a[r * w + k];
for c in 0..m {
u1[c * kf + fb + r] -= lrk * u1[c * kf + fb + k];
}
}
let lrr = fam_a[r * w + r];
for c in 0..m {
u1[c * kf + fb + r] /= lrr;
}
}
let fam_blks = &a21_blk[a21_off[f]..a21_off[f + 1]];
let a21_fam = &pk_a21[pk_a21_off[f]..pk_a21_off[f + 1]];
let fold_a21 = |out: &mut [f64], stride: usize, col_off: usize, panel_layout: bool| {
let mut cur = 0usize;
let mut loc0 = 0usize;
for &bi in fam_blks {
let br = &lam_blocks[bi as usize];
let (q, lo_x) = (br.q, br.lam_off);
let row_base = if panel_layout { loc0 } else { br.start - kf };
for a in 0..q {
for cb in 0..=np {
let (q_c, lo_c, c0, roff) = if cb == 0 {
(q_p, lo_p, 0, cur + a * q_p)
} else {
let ch = cb - 1;
(
q_n,
lam_blocks[n_prim + f * np + ch].lam_off,
q_p + ch * q_n,
cur + q * q_p + ch * q * q_n + a * q_n,
)
};
for jl in 0..q_c {
let mut h = 0.0;
for b in jl..q_c {
h += a21_fam[roff + b] * lam_small[lo_c + b * q_c + jl];
}
let colbase = (col_off + c0 + jl) * stride;
for il in 0..=a {
out[colbase + row_base + il] += lam_small[lo_x + a * q + il] * h;
}
}
}
}
cur += q * w;
loc0 += q;
}
};
match tail.as_mut() {
None => {
l21[fb * e..(fb + w) * e].fill(0.0);
fold_a21(&mut l21[..], e, fb, false);
schur_phase_b(&mut l21[fb * e..(fb + w) * e], e, w, fam_a);
}
Some(SparseTail {
axx,
fam_dd_slots,
fam_dd_off,
panel,
dd_temp,
b2_temp,
..
}) => {
let e_f: usize = fam_blks.iter().map(|&bi| lam_blocks[bi as usize].q).sum();
let panel = &mut panel[..e_f * w];
panel.fill(0.0);
fold_a21(panel, e_f, 0, true);
schur_phase_b(panel, e_f, w, fam_a);
let dd = &mut dd_temp[..e_f * e_f];
if w == 1 {
let p = &panel[..e_f];
for b in 0..e_f {
let pb = p[b];
let col = &mut dd[b * e_f..b * e_f + e_f];
for a in b..e_f {
col[a] = p[a] * pb;
}
}
} else {
let panel_ref = MatRef::from_column_major_slice(&panel[..e_f * w], e_f, w);
faer::linalg::matmul::triangular::matmul(
faer::MatMut::from_column_major_slice_mut(dd, e_f, e_f),
faer::linalg::matmul::triangular::BlockStructure::TriangularLower,
faer::Accum::Replace,
panel_ref,
faer::linalg::matmul::triangular::BlockStructure::Rectangular,
panel_ref.transpose(),
faer::linalg::matmul::triangular::BlockStructure::Rectangular,
1.0,
Par::Seq,
);
}
let (_, vals) = axx.parts_mut();
let mut rest = &fam_dd_slots[fam_dd_off[f]..fam_dd_off[f + 1]];
for b in 0..e_f {
let (col, tail_slots) = rest.split_at(e_f - b);
rest = tail_slots;
for (&s, &d) in col.iter().zip(&dd[b * e_f + b..b * e_f + e_f]) {
vals[s as usize] -= d;
}
}
debug_assert!(rest.is_empty(), "family downdate slot replay exhausted");
let panel_ref = MatRef::from_column_major_slice(&panel[..e_f * w], e_f, w);
let u1_sub = MatRef::from_column_major_slice_with_stride(&u1[fb..], w, m, kf);
faer::linalg::matmul::matmul(
faer::MatMut::from_column_major_slice_mut(&mut b2_temp[..e_f * m], e_f, m),
faer::Accum::Replace,
panel_ref,
u1_sub,
1.0,
Par::Seq,
);
let b2 = &b2_temp[..e_f * m];
let mut loc0 = 0usize;
for &bi in fam_blks {
let br = &lam_blocks[bi as usize];
let t0 = br.start - kf;
for a in 0..br.q {
for cm in 0..m {
u2[cm * e + t0 + a] -= b2[loc0 + a + cm * e_f];
}
}
loc0 += br.q;
}
}
}
}
if e > 0 {
match tail.as_mut() {
None => {
for tj in 0..e {
for ti in tj..e {
s22[(ti, tj)] = 0.0;
}
}
let mut cur = 0usize;
for pr in a22_pairs.iter() {
let bri = &lam_blocks[pr[0] as usize];
let bcj = &lam_blocks[pr[1] as usize];
let (ti0, tj0) = (bri.start - kf, bcj.start - kf);
let blk = &pk_a22[cur..cur + bri.q * bcj.q];
cur += bri.q * bcj.q;
for jl in 0..bcj.q {
fold_packed_col(
blk,
bri.q,
bcj.q,
lam_small,
bri.lam_off,
bcj.lam_off,
jl,
|il, v| {
let (ti, tj) = (ti0 + il, tj0 + jl);
if ti >= tj {
s22[(ti, tj)] = v;
}
},
);
}
}
for t in 0..e {
s22[(t, t)] += 1.0;
}
let l21_ref = MatRef::from_column_major_slice(&l21[..e * kf], e, kf);
faer::linalg::matmul::triangular::matmul(
s22.as_mat_mut(),
faer::linalg::matmul::triangular::BlockStructure::TriangularLower,
faer::Accum::Add,
l21_ref,
faer::linalg::matmul::triangular::BlockStructure::Rectangular,
l21_ref.transpose(),
faer::linalg::matmul::triangular::BlockStructure::Rectangular,
-1.0,
Par::Seq,
);
cholesky_in_place(
s22.as_mat_mut(),
LltRegularization::default(),
Par::Seq,
MemStack::new(tail_llt_mem),
Spec::default(),
)
.ok()?; for t in 0..e {
let ltt = s22[(t, t)];
if !(ltt.is_finite() && ltt > 0.0) {
return None;
}
log_lzz_half += ltt.ln();
}
faer::linalg::matmul::matmul(
faer::MatMut::from_column_major_slice_mut(&mut u2[..e * m], e, m),
faer::Accum::Add,
l21_ref,
MatRef::from_column_major_slice(&u1[..kf * m], kf, m),
-1.0,
Par::Seq,
);
for k in 0..e {
let s22k = s22.col(k).try_as_col_major().unwrap().as_slice();
for c in 0..m {
let col = &mut u2[c * e..(c + 1) * e];
let xk = col[k] / s22k[k];
col[k] = xk;
for (x, &s) in col[k + 1..].iter_mut().zip(&s22k[k + 1..]) {
*x -= s * xk;
}
}
}
}
Some(tail) => {
let llt = tail
.symbolic
.factorize_numeric_llt(
&mut tail.l_values,
tail.axx.as_ref(),
Side::Lower,
LltRegularization::default(),
Par::Seq,
MemStack::new(&mut tail.fac_mem),
Spec::default(),
)
.ok()?;
for c in 0..m {
for t in 0..e {
tail.x2[(t, c)] = u2[c * e + t];
}
}
llt.solve_in_place_with_conj(
Conj::No,
tail.x2.as_mat_mut(),
Par::Seq,
MemStack::new(&mut tail.solve_mem),
);
let _ = llt; let log_s22 = logdet_llt(&tail.symbolic, &tail.l_values);
if !log_s22.is_finite() {
return None;
}
log_lzz_half += 0.5 * log_s22;
}
}
}
for c in 0..m {
for r in 0..m {
factor[(r, c)] = 0.0;
}
}
for c in 0..m {
for r in c..m {
factor[(r, c)] = cxy[(r, c)];
}
}
let u1_ref = MatRef::from_column_major_slice(&u1[..kf * m], kf, m);
faer::linalg::matmul::triangular::matmul(
factor.as_mat_mut(),
faer::linalg::matmul::triangular::BlockStructure::TriangularLower,
faer::Accum::Add,
u1_ref.transpose(),
faer::linalg::matmul::triangular::BlockStructure::Rectangular,
u1_ref,
faer::linalg::matmul::triangular::BlockStructure::Rectangular,
-1.0,
Par::Seq,
);
let u2_ref = MatRef::from_column_major_slice(&u2[..e * m], e, m);
let rhs2 = match tail.as_ref() {
None => u2_ref,
Some(tail) => tail.x2.as_ref(),
};
faer::linalg::matmul::triangular::matmul(
factor.as_mat_mut(),
faer::linalg::matmul::triangular::BlockStructure::TriangularLower,
faer::Accum::Add,
u2_ref.transpose(),
faer::linalg::matmul::triangular::BlockStructure::Rectangular,
rhs2,
faer::linalg::matmul::triangular::BlockStructure::Rectangular,
-1.0,
Par::Seq,
);
cholesky_in_place(
factor.as_mat_mut(),
LltRegularization::default(),
Par::Seq,
MemStack::new(factor_llt_mem),
Spec::default(),
)
.ok()?;
Some(2.0 * log_lzz_half)
}
#[cfg(test)]
pub(crate) fn test_lcg(state: &mut u64) -> f64 {
*state = state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
(((*state >> 11) as f64) / ((1u64 << 53) as f64)) * 2.0 - 1.0
}