use super::*;
pub struct LmmSuffStats {
pub m: usize,
pub n_rows: usize,
pub n_clusters: usize,
pub groupings: LmmGroupings,
pub c: Mat<f64>,
pub s: Mat<f64>,
pub counts: Vec<f64>,
pub zx: Mat<f64>,
pub zx_slope: Mat<f64>,
pub w_buf: Vec<f64>,
}
impl LmmSuffStats {
pub fn new(p: usize, max_clusters: usize) -> Self {
Self::with_groupings(p, LmmGroupings::single(max_clusters))
}
pub fn with_groupings(p: usize, groupings: LmmGroupings) -> Self {
let m = p + 1;
let k = groupings.k_total;
let kx = groupings.k_crossed();
LmmSuffStats {
m,
n_rows: 0,
n_clusters: 0,
c: Mat::zeros(m, m),
s: Mat::zeros(m, k),
counts: vec![0.0; k],
zx: Mat::zeros(if kx > 0 { k } else { 0 }, kx),
zx_slope: Mat::zeros(if kx > 0 { k } else { 0 }, kx),
w_buf: vec![0.0; m],
groupings,
}
}
pub fn reset(&mut self) {
let m = self.m;
for j in 0..m {
for i in 0..m {
self.c[(i, j)] = 0.0;
}
}
for a in 0..self.counts.len() {
for j in 0..m {
self.s[(j, a)] = 0.0;
}
self.counts[a] = 0.0;
}
let (zr, zc) = (self.zx.nrows(), self.zx.ncols());
for j in 0..zc {
for i in 0..zr {
self.zx[(i, j)] = 0.0;
self.zx_slope[(i, j)] = 0.0;
}
}
self.n_rows = 0;
self.n_clusters = 0;
}
pub fn add_rows(&mut self, x: MatRef<'_, f64>, y: &[f64], cluster_ids: &[u32]) {
self.add_rows_multi(x, y, cluster_ids, &[], None);
}
pub fn add_rows_multi(
&mut self,
x: MatRef<'_, f64>,
y: &[f64],
cluster_ids: &[u32],
extra_ids: &[Vec<u32>],
weights: Option<&[f64]>,
) {
debug_assert_eq!(x.nrows(), y.len());
debug_assert_eq!(x.nrows(), cluster_ids.len());
debug_assert_eq!(extra_ids.len(), self.groupings.extra_offsets.len());
debug_assert!(weights.is_none_or(|w| w.len() == x.nrows()));
let p = self.m - 1;
debug_assert_eq!(x.ncols(), p);
let kf = self.groupings.k_family();
let n_g = 1 + extra_ids.len();
let mut gid = [0usize; 1 + MAX_EXTRA_GROUPINGS];
for row in 0..x.nrows() {
let wi = weights.map_or(1.0, |w| w[row]);
let zw = wi.sqrt();
gid[0] = cluster_ids[row] as usize;
for (e, ids) in extra_ids.iter().enumerate() {
gid[1 + e] =
self.groupings.extra_offsets[e] + ids[row] as usize * self.groupings.extra_q[e];
}
debug_assert!(gid[..n_g].iter().all(|&a| a < self.counts.len()));
for &a in &gid[..n_g] {
self.counts[a] += wi;
}
for j in 0..p {
self.w_buf[j] = x[(row, j)];
}
self.w_buf[p] = y[row];
for wj in &mut self.w_buf[..self.m] {
*wj *= zw;
}
for &a in &gid[..n_g] {
let scol = self
.s
.col_mut(a)
.try_as_col_major_mut()
.unwrap()
.as_slice_mut();
#[allow(clippy::needless_range_loop)]
for j in 0..self.m {
scol[j] += zw * self.w_buf[j];
}
}
for j in 0..self.m {
let wj = self.w_buf[j];
let ccol = self
.c
.col_mut(j)
.try_as_col_major_mut()
.unwrap()
.as_slice_mut();
#[allow(clippy::needless_range_loop)]
for i in j..self.m {
ccol[i] += self.w_buf[i] * wj;
}
}
if self.groupings.k_crossed() > 0 && !self.groupings.extra_slopes_any {
let slope = self.groupings.primary_q > 1;
let n_prim = self.groupings.n_primary;
for bi in 0..n_g {
let b = gid[bi];
if b >= kf {
let bl = b - kf;
#[allow(clippy::needless_range_loop)]
for ai in 0..n_g {
if ai != bi {
self.zx[(gid[ai], bl)] += wi;
}
}
if slope {
for (d, &sc) in self.groupings.primary_slope_cols.iter().enumerate() {
let z = self.w_buf[sc] / self.groupings.primary_slope_scales[d];
let scol = (d + 1) * n_prim + gid[0];
self.zx_slope[(scol, bl)] += z * zw;
}
}
}
}
} else if self.groupings.k_crossed() > 0 {
let g = &self.groupings;
let n_prim = g.n_primary;
for bi in 0..n_g {
let b = gid[bi];
if b < kf {
continue; }
let bl = b - kf;
let q_b = if bi == 0 {
g.primary_q
} else {
g.extra_q[bi - 1]
};
for db in 0..q_b {
let z_b = if db == 0 {
zw } else if bi == 0 {
self.w_buf[g.primary_slope_cols[db - 1]]
/ g.primary_slope_scales[db - 1]
} else {
self.w_buf[g.extra_slope_cols[bi - 1][db - 1]]
/ g.extra_slope_scales[bi - 1][db - 1]
};
let b_local = bl + db;
for ai in 0..n_g {
if ai == bi {
continue;
}
let q_a = if ai == 0 {
g.primary_q
} else {
g.extra_q[ai - 1]
};
for da in 0..q_a {
let (a_col, z_a) = if da == 0 {
(gid[ai], zw) } else if ai == 0 {
(
da * n_prim + gid[0],
self.w_buf[g.primary_slope_cols[da - 1]]
/ g.primary_slope_scales[da - 1],
)
} else {
(
gid[ai] + da,
self.w_buf[g.extra_slope_cols[ai - 1][da - 1]]
/ g.extra_slope_scales[ai - 1][da - 1],
)
};
self.zx[(a_col, b_local)] += z_a * z_b;
}
}
}
}
}
if self.groupings.primary_q > 1 {
let n_prim = self.groupings.n_primary;
for (k, &sc) in self.groupings.primary_slope_cols.iter().enumerate() {
let z = self.w_buf[sc] / self.groupings.primary_slope_scales[k];
let scol = (k + 1) * n_prim + gid[0];
let scol_mut = self
.s
.col_mut(scol)
.try_as_col_major_mut()
.unwrap()
.as_slice_mut();
#[allow(clippy::needless_range_loop)]
for j in 0..self.m {
scol_mut[j] += z * self.w_buf[j];
}
}
}
if self.groupings.extra_slopes_any {
for e in 0..self.groupings.extra_slope_cols.len() {
let gintercept = gid[1 + e];
let n_d = self.groupings.extra_slope_cols[e].len();
for d in 0..n_d {
let sc = self.groupings.extra_slope_cols[e][d];
let z = self.w_buf[sc] / self.groupings.extra_slope_scales[e][d];
let scol = gintercept + 1 + d;
let scol_mut = self
.s
.col_mut(scol)
.try_as_col_major_mut()
.unwrap()
.as_slice_mut();
#[allow(clippy::needless_range_loop)]
for j in 0..self.m {
scol_mut[j] += z * self.w_buf[j];
}
}
}
}
if gid[0] + 1 > self.n_clusters {
self.n_clusters = gid[0] + 1;
}
}
self.n_rows += x.nrows();
}
}
pub(crate) fn precompute_balanced_collapse(suff: &LmmSuffStats, fit: &mut LmmFitScratch) -> bool {
let g = &suff.groupings;
fit.collapse_n_active = 0;
if g.primary_q != 1 || g.n_primary == 0 || suff.n_rows == 0 {
return false;
}
let np = g.nested_per_parent;
let w = 1 + np;
let kx = g.k_crossed();
let m = suff.m;
let t_dim = kx + m;
let n0 = suff.counts[0];
if n0 == 0.0 {
return false;
}
let mut n_active = 1;
while n_active < g.n_primary && suff.counts[n_active] == n0 {
n_active += 1;
}
if suff.counts[n_active..g.n_primary].iter().any(|&c| c != 0.0) {
return false; }
for c in 0..np {
let c0 = suff.counts[g.n_primary + c]; for f in 0..g.n_primary {
let cc = suff.counts[g.n_primary + f * np + c];
if (f < n_active && cc != c0) || (f >= n_active && cc != 0.0) {
return false;
}
}
}
let blk = t_dim * t_dim;
let npairs = w * (w + 1) / 2;
fit.fam_gram[..npairs * blk].fill(0.0);
for f in 0..n_active {
for r in 0..w {
let gcol = if r == 0 {
f
} else {
g.n_primary + f * np + (r - 1)
};
let dst = &mut fit.bt[r * t_dim..(r + 1) * t_dim];
for (b, slot) in dst[..kx].iter_mut().enumerate() {
*slot = suff.zx[(gcol, b)];
}
let scol = suff.s.col(gcol).try_as_col_major().unwrap().as_slice();
dst[kx..kx + m].copy_from_slice(scol);
}
let (bt, gram) = (&fit.bt, &mut fit.fam_gram);
let mut pidx = 0;
for r in 0..w {
for rp in r..w {
let gblk = &mut gram[pidx * blk..(pidx + 1) * blk];
for j in 0..t_dim {
let vj = bt[rp * t_dim + j];
if vj != 0.0 {
for i in 0..t_dim {
gblk[j * t_dim + i] += bt[r * t_dim + i] * vj;
}
}
}
pidx += 1;
}
}
}
fit.collapse_n_active = n_active;
true
}
fn reml_deviance_blocked(theta: &[f64], suff: &LmmSuffStats, fit: &mut LmmFitScratch) -> f64 {
let g = &suff.groupings;
let m = suff.m;
let p = m - 1;
let k = g.k_total;
let dim = k + m;
let n_prim = g.n_primary;
let q_p = g.primary_q;
let np = g.nested_per_parent;
let prim_width = q_p * n_prim;
let kf = g.k_family();
fit.blocked_lam[..k * k].fill(0.0);
primary_lambda(theta, q_p, &mut fit.prim_lam);
for f in 0..n_prim {
for dc in 0..q_p {
for dr in dc..q_p {
let row = dr * n_prim + f;
let col = dc * n_prim + f;
fit.blocked_lam[col * k + row] = fit.prim_lam[dr * q_p + dc];
}
}
}
if let Some(nf) = g.nested {
let q_n = nf.q;
let mut lam_n = [0.0f64; MAX_EXTRA_Q * MAX_EXTRA_Q];
primary_lambda(&theta[nf.vech_start..], q_n, &mut lam_n);
for f in 0..n_prim {
for c in 0..np {
let ic = prim_width + (f * np + c) * q_n;
for dc in 0..q_n {
for dr in dc..q_n {
fit.blocked_lam[(ic + dc) * k + (ic + dr)] = lam_n[dr * q_n + dc];
}
}
}
}
}
let mut lam_g = [0.0f64; MAX_EXTRA_Q * MAX_EXTRA_Q];
for cf in &g.crossed {
let q = cf.q;
primary_lambda(&theta[cf.vech_start..], q, &mut lam_g);
let off = g.extra_offsets[cf.decl];
for c in 0..cf.n_levels {
let ic = off + c * q;
for dc in 0..q {
for dr in dc..q {
fit.blocked_lam[(ic + dc) * k + (ic + dr)] = lam_g[dr * q + dc];
}
}
}
}
fit.blocked_g[..k * k].fill(0.0);
for b in kf..k {
let bl = b - kf;
for a in 0..k {
let v = suff.zx[(a, bl)];
fit.blocked_g[b * k + a] = v;
fit.blocked_g[a * k + b] = v;
}
}
for f in 0..n_prim {
primary_gram(suff, g, f, q_p, &mut fit.prim_gram);
for dr in 0..q_p {
for dc in 0..q_p {
fit.blocked_g[(dc * n_prim + f) * k + (dr * n_prim + f)] =
fit.prim_gram[dr * q_p + dc];
}
}
}
if let Some(nf) = g.nested {
let q_n = nf.q;
let nscols = &g.extra_slope_cols[nf.decl];
let nssc = &g.extra_slope_scales[nf.decl];
for f in 0..n_prim {
for c in 0..np {
let ic = prim_width + (f * np + c) * q_n;
let n_c = suff.counts[ic];
for dr in 0..q_n {
for dc in 0..q_n {
let v = if dr == 0 && dc == 0 {
n_c
} else if dr == 0 {
suff.s[(nscols[dc - 1], ic)] / nssc[dc - 1] } else if dc == 0 {
suff.s[(nscols[dr - 1], ic)] / nssc[dr - 1] } else {
suff.s[(nscols[dr - 1], ic + dc)] / nssc[dr - 1] };
fit.blocked_g[(ic + dc) * k + (ic + dr)] = v;
}
}
for da in 0..q_p {
let prow = da * n_prim + f;
for dc in 0..q_n {
let ccol = ic + dc;
let v = if da == 0 && dc == 0 {
n_c
} else if dc == 0 {
suff.s[(g.primary_slope_cols[da - 1], ic)]
/ g.primary_slope_scales[da - 1] } else if da == 0 {
suff.s[(nscols[dc - 1], ic)] / nssc[dc - 1] } else {
suff.s[(g.primary_slope_cols[da - 1], ic + dc)]
/ g.primary_slope_scales[da - 1] };
fit.blocked_g[ccol * k + prow] = v;
fit.blocked_g[prow * k + ccol] = v;
}
}
}
}
}
for cf in &g.crossed {
let q = cf.q;
let off = g.extra_offsets[cf.decl];
let scols = &g.extra_slope_cols[cf.decl];
let ssc = &g.extra_slope_scales[cf.decl];
for c in 0..cf.n_levels {
let ic = off + c * q;
let n_c = suff.counts[ic];
for dr in 0..q {
for dc in 0..q {
let v = if dr == 0 && dc == 0 {
n_c
} else if dr == 0 {
suff.s[(scols[dc - 1], ic)] / ssc[dc - 1] } else if dc == 0 {
suff.s[(scols[dr - 1], ic)] / ssc[dr - 1] } else {
suff.s[(scols[dr - 1], ic + dc)] / ssc[dr - 1] };
fit.blocked_g[(ic + dc) * k + (ic + dr)] = v;
}
}
}
}
fit.blocked_tmp[..k * k].fill(0.0);
for bp in 0..k {
for a in 0..k {
let mut acc = 0.0;
for ap in 0..k {
let l = fit.blocked_lam[a * k + ap]; if l != 0.0 {
acc += l * fit.blocked_g[bp * k + ap]; }
}
fit.blocked_tmp[bp * k + a] = acc;
}
}
for b in 0..k {
for a in b..k {
let mut acc = 0.0;
for bp in 0..k {
let l = fit.blocked_lam[b * k + bp]; if l != 0.0 {
acc += fit.blocked_tmp[bp * k + a] * l;
}
}
if a == b {
acc += 1.0; }
fit.blocked_p[b * dim + a] = acc;
}
}
for a in 0..k {
for j in 0..m {
let mut acc = 0.0;
for ap in 0..k {
let l = fit.blocked_lam[a * k + ap];
if l != 0.0 {
acc += l * suff.s[(j, ap)];
}
}
fit.blocked_p[a * dim + (k + j)] = acc;
}
}
for j in 0..m {
for i in j..m {
fit.blocked_p[(k + j) * dim + (k + i)] = suff.c[(i, j)];
}
}
let pref = faer::MatRef::from_column_major_slice(&fit.blocked_p[..dim * dim], dim, dim);
let chol = match pref.llt(faer::Side::Lower) {
Ok(c) => c,
Err(_) => return f64::INFINITY,
};
let l = chol.L();
let mut log_lzz_half = 0.0_f64;
for i in 0..k {
let lii = l[(i, i)];
if !(lii.is_finite() && lii > 0.0) {
return f64::INFINITY;
}
log_lzz_half += lii.ln();
}
let mut log_lxx_sq = 0.0_f64;
for j in 0..p {
let ljj = l[(k + j, k + j)];
if !(ljj.is_finite() && ljj > 0.0) {
return f64::INFINITY;
}
log_lxx_sq += ljj.ln();
}
log_lxx_sq *= 2.0;
for j in 0..m {
for i in 0..m {
fit.factor[(i, j)] = if i >= j { l[(k + i, k + j)] } else { 0.0 };
}
}
let lyy = fit.factor[(p, p)];
let r_sq = lyy * lyy;
let df = (suff.n_rows - p) as f64;
let sigma_sq = r_sq / df;
if !(sigma_sq.is_finite() && sigma_sq > 0.0) {
return f64::INFINITY;
}
fit.sigma_sq = sigma_sq;
2.0 * log_lzz_half + log_lxx_sq + df * sigma_sq.ln()
}
pub fn reml_deviance(theta: &[f64], suff: &LmmSuffStats, fit: &mut LmmFitScratch) -> f64 {
let g = &suff.groupings;
debug_assert_eq!(theta.len(), g.n_theta());
let m = suff.m;
let p = m - 1;
if suff.n_rows <= p || p == 0 {
return f64::INFINITY;
}
if g.extra_slopes_any {
return reml_deviance_blocked(theta, suff, fit);
}
let kf = g.k_family();
let kx = g.k_crossed();
let t_dim = kx + m;
let np = g.nested_per_parent;
let w = g.primary_q + np; let th_p = theta[0];
let th_n = g.nested.map(|nf| theta[nf.vech_start]).unwrap_or(0.0);
let slope = g.primary_q > 1;
if slope {
primary_lambda(theta, g.primary_q, &mut fit.prim_lam);
}
debug_assert!(!g.extra_slopes_any);
{
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;
}
}
}
fit.tail[..t_dim * t_dim].fill(0.0);
for b in 0..kx {
let lam = fit.lam_x[b];
let gcol = kf + b;
let zxb = suff.zx.col(b).try_as_col_major().unwrap().as_slice();
for a in 0..b {
fit.tail[a * t_dim + b] = lam * fit.lam_x[a] * zxb[kf + a];
}
let scol = suff.s.col(gcol).try_as_col_major().unwrap().as_slice();
let tcol = &mut fit.tail[b * t_dim..(b + 1) * t_dim];
tcol[b] = 1.0 + lam * lam * suff.counts[gcol];
for j in 0..m {
tcol[kx + j] = lam * scol[j];
}
}
for j in 0..m {
let ccol = suff.c.col(j).try_as_col_major().unwrap().as_slice();
let tcol = &mut fit.tail[(kx + j) * t_dim..(kx + j + 1) * t_dim];
tcol[kx + j..kx + m].copy_from_slice(&ccol[j..m]);
}
let collapse = !slope && fit.collapse_n_active > 0;
let mut log_lzz_half = 0.0_f64; if collapse {
let n_active = fit.collapse_n_active;
let n_f = suff.counts[0];
fit.fam_a[0] = 1.0 + th_p * th_p * n_f;
for c in 0..np {
let n_c = suff.counts[g.n_primary + c];
for c2 in 0..np {
fit.fam_a[(1 + c) * w + (1 + c2)] = 0.0;
}
fit.fam_a[(1 + c) * w] = th_p * th_n * n_c;
fit.fam_a[(1 + c) * w + (1 + c)] = 1.0 + th_n * th_n * n_c;
}
let mut log_l_half = 0.0_f64;
for j in 0..w {
let mut d = fit.fam_a[j * w + j];
for k in 0..j {
let v = fit.fam_a[j * w + k];
d -= v * v;
}
if !(d.is_finite() && d > 0.0) {
return f64::INFINITY;
}
let l = d.sqrt();
fit.fam_a[j * w + j] = l;
log_l_half += l.ln();
for i in (j + 1)..w {
let mut v = fit.fam_a[i * w + j];
for k in 0..j {
v -= fit.fam_a[i * w + k] * fit.fam_a[j * w + k];
}
fit.fam_a[i * w + j] = v / l;
}
}
log_lzz_half = (n_active as f64) * log_l_half;
for r in 0..w {
for i in 0..w {
let mut acc = if i == r { 1.0 } else { 0.0 };
for k in 0..i {
acc -= fit.fam_a[i * w + k] * fit.comb[k];
}
fit.comb[i] = acc / fit.fam_a[i * w + i];
}
for i in (0..w).rev() {
let mut acc = fit.comb[i];
for k in (i + 1)..w {
acc -= fit.fam_a[k * w + i] * fit.a_inv[k * w + r];
}
fit.a_inv[i * w + r] = acc / fit.fam_a[i * w + i];
}
}
let t2 = t_dim * t_dim;
fit.comb[..t2].fill(0.0);
let (comb, gram) = (&mut fit.comb, &fit.fam_gram);
let mut pidx = 0;
for r in 0..w {
let sr = if r == 0 { th_p } else { th_n };
for rp in r..w {
let srp = if rp == 0 { th_p } else { th_n };
let coeff = sr * srp * fit.a_inv[r * w + rp];
let gblk = &gram[pidx * t2..(pidx + 1) * t2];
if coeff != 0.0 {
if r == rp {
for j in 0..t_dim {
for i in j..t_dim {
comb[j * t_dim + i] += coeff * gblk[j * t_dim + i];
}
}
} else {
for j in 0..t_dim {
for i in j..t_dim {
comb[j * t_dim + i] +=
coeff * (gblk[j * t_dim + i] + gblk[i * t_dim + j]);
}
}
}
}
pidx += 1;
}
}
for j in 0..t_dim {
let dj = if j < kx { fit.lam_x[j] } else { 1.0 };
for i in j..t_dim {
let di = if i < kx { fit.lam_x[i] } else { 1.0 };
fit.tail[j * t_dim + i] -= di * dj * fit.comb[j * t_dim + i];
}
}
} else {
for f in 0..g.n_primary {
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 f64::INFINITY;
}
let mut fam_prod = 1.0_f64;
for j in 0..w {
fam_prod *= fit.fam_a[j * w + j];
}
log_lzz_half += fam_prod.ln();
let fb = f * w;
let bt_fam = &mut fit.bt[fb * t_dim..(fb + 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);
}
let w_tot = g.n_primary * w;
{
let bt = faer::MatRef::from_column_major_slice(&fit.bt[..t_dim * w_tot], t_dim, w_tot);
let tail = faer::MatMut::from_column_major_slice_mut(
&mut fit.tail[..t_dim * t_dim],
t_dim,
t_dim,
);
faer::linalg::matmul::triangular::matmul(
tail,
faer::linalg::matmul::triangular::BlockStructure::TriangularLower,
faer::Accum::Add,
bt,
faer::linalg::matmul::triangular::BlockStructure::Rectangular,
bt.transpose(),
faer::linalg::matmul::triangular::BlockStructure::Rectangular,
-1.0,
faer::Par::Seq,
);
}
}
let tail_ref = faer::MatRef::from_column_major_slice(&fit.tail[..t_dim * t_dim], t_dim, t_dim);
let chol = match tail_ref.llt(faer::Side::Lower) {
Ok(c) => c,
Err(_) => return f64::INFINITY,
};
let l = chol.L();
for b in 0..kx {
let lbb = l[(b, b)];
if !(lbb.is_finite() && lbb > 0.0) {
return f64::INFINITY;
}
log_lzz_half += lbb.ln();
}
let log_lzz_sq = 2.0 * log_lzz_half;
for j in 0..m {
let lcol = l.col(kx + j).try_as_col_major().unwrap().as_slice();
for i in 0..m {
fit.factor[(i, j)] = if i >= j { lcol[kx + i] } else { 0.0 };
}
}
let mut log_lxx_sq = 0.0_f64;
for j in 0..p {
let ljj = fit.factor[(j, j)];
if !(ljj.is_finite() && ljj > 0.0) {
return f64::INFINITY;
}
log_lxx_sq += ljj.ln();
}
log_lxx_sq *= 2.0;
let lyy = fit.factor[(p, p)];
let r_sq = lyy * lyy;
let df = (suff.n_rows - p) as f64;
let sigma_sq = r_sq / df;
if !(sigma_sq.is_finite() && sigma_sq > 0.0) {
return f64::INFINITY;
}
fit.sigma_sq = sigma_sq;
log_lzz_sq + log_lxx_sq + df * sigma_sq.ln()
}