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(all(test, feature = "formula"))]
mod fd_margin;
#[cfg(test)]
mod tests;
pub(crate) use glmm::{fit_glmm_nb_sparse, fit_glmm_sparse};
const PIVOT_MIN: f64 = 6e-10;
#[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 mut g = LmmGroupings::from_cluster_spec_ext(model, n, &slope_cols, &extra_slope_cols);
let xm = MatRef::from_row_major_slice(x, n, p);
g.set_slope_scales(xm, opts.weights.as_deref());
let g = g;
let sqrt_w: Option<Vec<f64>> = opts
.weights
.as_ref()
.map(|w| w.iter().map(|v| v.sqrt()).collect());
let y_shifted: Vec<f64>;
let y_eff: &[f64] = match &opts.offset {
Some(o) => {
y_shifted = y.iter().zip(o).map(|(&yi, &oi)| yi - oi).collect();
&y_shifted
}
None => y,
};
let mut ws = SparseLmmWorkspace::new(
&g,
xm,
cluster_ids,
extra_ids,
y_eff,
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());
let sc = g.theta_row_scales();
for ((t, &v), &f) in theta.iter_mut().zip(&s.theta).zip(sc.iter()) {
*t = v * f;
}
for &i in g.diagonal_theta() {
theta[i] = theta[i].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;
let mut pinned_components = 0u64;
if has_endpoint {
for (kk, &ti) in g.diagonal_theta().iter().enumerate() {
if theta[ti] <= crate::lmm::PIN_THETA {
theta[ti] = 0.0;
if converged_status {
pinned = true;
if kk < u64::BITS as usize {
pinned_components |= 1u64 << kk;
}
}
}
}
}
let factor_ok = has_endpoint && sparse_schur_factor(&theta, &mut ws).is_some();
let degenerate = if factor_ok {
crate::ols::min_pivot_ratio(ws.factor.as_ref(), p).0 < PIVOT_MIN
} 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,
diagnostics: crate::Diagnostics::from_flags(false, false, p),
varcorr: vec![],
stddev_se: vec![],
n_eval: out.n_eval,
deviance: f64::NAN,
loglik: f64::NAN,
df: 0,
reml: true,
fitted: vec![],
ranef: vec![],
ranef_levels: vec![],
};
}
ws.arm_recovery();
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 theta_scales = g.theta_row_scales();
let tau2: Vec<f64> = theta
.iter()
.zip(theta_scales.iter())
.map(|(&t, &s)| (t / s) * (t / s) * 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 pinned_grid = crate::fit::pinned_flags(pinned_components, &varcorr);
let (fitted, ranef, ranef_levels) = match converged.then(|| sparse_recover_u(&ws, &beta)) {
Some(Some(u)) => {
let ranef = crate::fit::assemble_ranef_sparse(&theta, &g, &u);
let fitted = crate::fit::lmm_fitted(
x,
n,
p,
&beta,
&ranef,
&g,
cluster_ids,
extra_ids,
opts.offset.as_deref(),
);
(fitted, ranef, crate::fit::ranef_level_counts(&g))
}
_ => (vec![], vec![], vec![]),
};
let mut fit = crate::Fit {
beta,
se,
vcov,
tau2,
dispersion: sigma_sq,
diagnostics: crate::Diagnostics {
pinned: pinned_grid,
..crate::Diagnostics::from_flags(converged, pinned, p)
},
varcorr,
stddev_se: vec![],
n_eval: out.n_eval,
deviance: dev,
loglik: crate::fit::lmm_loglik(dev, n, p),
df: p + theta.len() + 1,
reml: true,
fitted,
ranef,
ranef_levels,
};
fit.diagnostics.singular =
fit.diagnostics.singular || fit.has_negligible_component(&crate::fit::re_scale_grid(&g));
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;
const DD_DENSE_BETA: f64 = 0.5;
const DD_DENSE_MAX_BYTES: usize = 256 << 20;
#[cfg(test)]
thread_local! {
pub(crate) static FORCE_SPARSE_TAIL: std::cell::Cell<bool> =
const { std::cell::Cell::new(false) };
}
#[cfg(test)]
thread_local! {
pub(crate) static FORCE_DD_ROUTE: std::cell::Cell<Option<bool>> =
const { std::cell::Cell::new(None) };
}
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: FamDowndate,
pub(crate) panel: Vec<f64>,
pub(crate) dd_temp: Vec<f64>,
pub(crate) b2_temp: Vec<f64>,
pub(crate) x2: Mat<f64>,
}
pub(crate) enum FamDowndate {
Scatter { slots: Vec<u32>, off: Vec<usize> },
Dense {
rows: Vec<u32>,
off: Vec<usize>,
acc: Vec<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) rec: Option<SparseRecovery>,
}
pub(crate) struct SparseRecovery {
fam_l: Vec<f64>,
panel: Vec<f64>,
panel_off: Vec<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,
rec: None,
}
}
fn arm_recovery(&mut self) {
let w = self.fam_w;
let n_prim = self.g.n_primary;
let sparse_tail = self.tail.is_some();
let mut panel_off = Vec::with_capacity(n_prim + 1);
let mut acc = 0usize;
for f in 0..n_prim {
panel_off.push(acc);
if sparse_tail {
let e_f: usize = self.a21_blk[self.a21_off[f]..self.a21_off[f + 1]]
.iter()
.map(|&bi| self.lam_blocks[bi as usize].q)
.sum();
acc += e_f * w;
}
}
panel_off.push(acc);
self.rec = Some(SparseRecovery {
fam_l: vec![0.0f64; n_prim * w * w],
panel: vec![0.0f64; acc],
panel_off,
});
}
}
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 max_ef = 0usize;
let mut sum_pairs = 0.0f64;
for list in fam_crossed {
let ef: usize = list.iter().map(|&bi| lam_blocks[bi as usize].q).sum();
max_ef = max_ef.max(ef);
sum_pairs += (ef as f64) * (ef as f64 + 1.0) * 0.5;
}
let nnz = axx.symbolic().row_idx().len() as f64;
let cap_ok = e.saturating_mul(e).saturating_mul(8) <= DD_DENSE_MAX_BYTES;
let dense_route = cap_ok && nnz + (e as f64) * (e as f64) < DD_DENSE_BETA * sum_pairs;
#[cfg(test)]
let dense_route = match FORCE_DD_ROUTE.with(|c| c.get()) {
Some(want) => want && cap_ok,
None => dense_route,
};
let mut off = Vec::with_capacity(fam_crossed.len() + 1);
off.push(0);
let fam_dd = if dense_route {
let mut rows: Vec<u32> = Vec::new();
for list in fam_crossed {
for &bi in list {
let b = &lam_blocks[bi as usize];
let t0 = b.start - kf;
rows.extend((t0..t0 + b.q).map(|t| t as u32));
}
off.push(rows.len());
}
FamDowndate::Dense {
rows,
off,
acc: vec![0.0f64; e * e],
}
} else {
let mut slots = Vec::new();
let mut cols: Vec<usize> = Vec::new();
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);
}
for bl in 0..cols.len() {
for a in bl..cols.len() {
slots.push(slot_of((cols[a], cols[bl])));
}
}
off.push(slots.len());
}
FamDowndate::Scatter { slots, off }
};
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,
panel: vec![0.0f64; fam_w * max_ef],
dd_temp: vec![
0.0f64;
if dense_route && fam_w == 1 {
0
} else {
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)] / g.primary_slope_scales[k]),
);
}
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)] / g.extra_slope_scales[e][c]));
}
}
}
#[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,
rec,
..
} = 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,
fam_dd,
..
} = tail;
if let FamDowndate::Dense { acc, .. } = fam_dd {
acc.fill(0.0);
}
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;
}
}
if let Some(r) = rec.as_mut() {
r.fam_l[f * w * w..(f + 1) * w * w].copy_from_slice(&fam_a[..w * w]);
}
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,
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);
if let Some(r) = rec.as_mut() {
r.panel[r.panel_off[f]..r.panel_off[f + 1]].copy_from_slice(panel);
}
match fam_dd {
FamDowndate::Scatter { slots, off } => {
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 = &slots[off[f]..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");
}
FamDowndate::Dense { rows, off, acc } => {
let rmap = &rows[off[f]..off[f + 1]];
debug_assert_eq!(rmap.len(), e_f, "dense row map width");
if w == 1 {
let p = &panel[..e_f];
for b in 0..e_f {
let pb = p[b];
let col = &mut acc[rmap[b] as usize * e..];
for a in b..e_f {
col[rmap[a] as usize] += p[a] * pb;
}
}
} else {
let dd = &mut dd_temp[..e_f * e_f];
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,
);
for b in 0..e_f {
let col = &mut acc[rmap[b] as usize * e..];
for (&r, &d) in
rmap[b..].iter().zip(&dd[b * e_f + b..b * e_f + e_f])
{
col[r as usize] += d;
}
}
}
}
}
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) => {
if let SparseTail {
axx,
fam_dd: FamDowndate::Dense { acc, .. },
..
} = &mut *tail
{
let (sym, vals) = axx.parts_mut();
let col_ptr = sym.col_ptr();
let row_idx = sym.row_idx();
for j in 0..e {
let base = j * e;
for k in col_ptr[j]..col_ptr[j + 1] {
vals[k] -= acc[base + row_idx[k]];
}
}
}
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)
}
fn sparse_recover_u(ws: &SparseLmmWorkspace, beta: &[f64]) -> Option<Vec<f64>> {
let rec = ws.rec.as_ref()?;
let g = &ws.g;
let p = ws.p;
let w = ws.fam_w;
let kf = g.k_family();
let e = g.k_crossed();
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 dot_c = |row: &dyn Fn(usize) -> f64| -> f64 {
let mut acc = row(p);
for (j, &b) in beta.iter().enumerate().take(p) {
acc -= row(j) * b;
}
acc
};
let mut u2 = vec![0.0f64; e];
if e > 0 {
match ws.tail.as_ref() {
Some(tail) => {
for (t, slot) in u2.iter_mut().enumerate() {
*slot = dot_c(&|c| tail.x2[(t, c)]);
}
}
None => {
for (t, slot) in u2.iter_mut().enumerate() {
*slot = dot_c(&|c| ws.u2[c * e + t]);
}
for t in (0..e).rev() {
let mut acc = u2[t];
for (i, &solved) in u2.iter().enumerate().skip(t + 1) {
acc -= ws.s22[(i, t)] * solved;
}
let ltt = ws.s22[(t, t)];
if !(ltt.is_finite() && ltt > 0.0) {
return None;
}
u2[t] = acc / ltt;
}
}
}
}
let mut u = vec![0.0f64; g.k_total];
for (t, &v) in u2.iter().enumerate() {
u[kf + t] = v;
}
let mut rhs = vec![0.0f64; w];
for f in 0..n_prim {
let fb = f * w;
for (r, slot) in rhs.iter_mut().enumerate() {
*slot = dot_c(&|c| ws.u1[c * kf + fb + r]);
}
if e > 0 {
match ws.tail.as_ref() {
Some(_) => {
let panel = &rec.panel[rec.panel_off[f]..rec.panel_off[f + 1]];
let fam_blks = &ws.a21_blk[ws.a21_off[f]..ws.a21_off[f + 1]];
let e_f = panel.len().checked_div(w).unwrap_or(0);
let mut loc0 = 0usize;
for &bi in fam_blks {
let br = &ws.lam_blocks[bi as usize];
let t0 = br.start - kf;
for a in 0..br.q {
let ut = u2[t0 + a];
if ut != 0.0 {
for (r, slot) in rhs.iter_mut().enumerate() {
*slot -= panel[r * e_f + loc0 + a] * ut;
}
}
}
loc0 += br.q;
}
}
None => {
for (r, slot) in rhs.iter_mut().enumerate() {
let col = &ws.l21[(fb + r) * e..(fb + r + 1) * e];
for (t, &ut) in u2.iter().enumerate() {
*slot -= col[t] * ut;
}
}
}
}
}
let l = &rec.fam_l[fb * w..(fb + w) * w];
for r in (0..w).rev() {
let mut acc = rhs[r];
for (i, &solved) in rhs.iter().enumerate().skip(r + 1) {
acc -= l[i * w + r] * solved;
}
let lrr = l[r * w + r];
if !(lrr.is_finite() && lrr > 0.0) {
return None;
}
rhs[r] = acc / lrr;
}
for (r, &v) in rhs.iter().enumerate() {
let re_col = if r < q_p {
r * n_prim + f
} else {
let rr = r - q_p;
let br = &ws.lam_blocks[n_prim + f * np + rr / q_n];
br.start + (rr % q_n) * br.stride
};
u[re_col] = v;
}
}
Some(u)
}
#[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
}