use super::*;
use crate::dual::{Dual, HyperDual};
use crate::glmm::{unpack_hessian, DerivStatus};
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<T: Scalar>(
suff: &LmmSuffStats,
fit: &mut LmmFitScratch<T>,
) -> 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.collapse_stage[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.collapse_stage, &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>) -> 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)];
}
}
if !<f64 as Scalar>::chol_lower(&fit.blocked_p, dim, &mut fit.blocked_l) {
return f64::INFINITY;
}
let l = faer::MatRef::from_column_major_slice(&fit.blocked_l[..dim * dim], dim, dim);
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<T: Scalar + 'static>(
theta: &[T],
suff: &LmmSuffStats,
fit: &mut LmmFitScratch<T>,
) -> T {
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 T::from_f64(f64::INFINITY);
}
if g.extra_slopes_any {
assert!(
T::IS_F64,
"reml_deviance_blocked (crossed/nested slopes) is f64-only"
);
let mut theta_f64 = [0.0f64; crate::consts::MAX_THETA];
for (dst, src) in theta_f64.iter_mut().zip(theta.iter()) {
*dst = src.value();
}
let fit64: &mut LmmFitScratch<f64> = (fit as &mut dyn std::any::Any)
.downcast_mut()
.expect("T::IS_F64 asserted above");
return T::from_f64(reml_deviance_blocked(
&theta_f64[..theta.len()],
suff,
fit64,
));
}
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(T::ZERO);
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(T::ZERO);
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] * T::from_f64(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] = T::ONE + lam * lam * T::from_f64(suff.counts[gcol]);
for j in 0..m {
tcol[kx + j] = lam * T::from_f64(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];
for (dst, &src) in tcol[kx + j..kx + m].iter_mut().zip(&ccol[j..m]) {
*dst = T::from_f64(src);
}
}
let collapse = !slope && fit.collapse_n_active > 0;
let mut log_lzz_half = T::ZERO; if collapse {
let n_active = fit.collapse_n_active;
let n_f = T::from_f64(suff.counts[0]);
fit.fam_a[0] = T::ONE + th_p * th_p * n_f;
for c in 0..np {
let n_c = T::from_f64(suff.counts[g.n_primary + c]);
for c2 in 0..np {
fit.fam_a[(1 + c) * w + (1 + c2)] = T::ZERO;
}
fit.fam_a[(1 + c) * w] = th_p * th_n * n_c;
fit.fam_a[(1 + c) * w + (1 + c)] = T::ONE + th_n * th_n * n_c;
}
let mut log_l_half = T::ZERO;
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.value().is_finite() && d.value() > 0.0) {
return T::from_f64(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 = T::from_f64(n_active as f64) * log_l_half;
for r in 0..w {
for i in 0..w {
let mut acc = if i == r { T::ONE } else { T::ZERO };
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(T::ZERO);
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.value() != 0.0 {
if r == rp {
for j in 0..t_dim {
for i in j..t_dim {
comb[j * t_dim + i] += coeff * T::from_f64(gblk[j * t_dim + i]);
}
}
} else {
for j in 0..t_dim {
for i in j..t_dim {
comb[j * t_dim + i] +=
coeff * T::from_f64(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 { T::ONE };
for i in j..t_dim {
let di = if i < kx { fit.lam_x[i] } else { T::ONE };
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 T::from_f64(f64::INFINITY);
}
let mut fam_prod = T::ONE;
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 LmmFitScratch {
bt,
tail,
syrk_scratch,
..
} = &mut *fit;
T::syrk_lower_sub(
&bt[..t_dim * w_tot],
t_dim,
w_tot,
&mut tail[..t_dim * t_dim],
syrk_scratch,
);
}
if !T::chol_lower(
&fit.tail[..t_dim * t_dim],
t_dim,
&mut fit.tail_l[..t_dim * t_dim],
) {
return T::from_f64(f64::INFINITY);
}
for b in 0..kx {
let lbb = fit.tail_l[b * t_dim + b];
if !(lbb.value().is_finite() && lbb.value() > 0.0) {
return T::from_f64(f64::INFINITY);
}
log_lzz_half += lbb.ln();
}
let log_lzz_sq = T::from_f64(2.0) * log_lzz_half;
for j in 0..m {
for i in 0..m {
fit.factor[(i, j)] = if i >= j {
fit.tail_l[(kx + j) * t_dim + (kx + i)].value()
} else {
0.0
};
}
}
let mut log_lxx_sq = T::ZERO;
for j in 0..p {
let ljj = fit.tail_l[(kx + j) * t_dim + (kx + j)];
if !(ljj.value().is_finite() && ljj.value() > 0.0) {
return T::from_f64(f64::INFINITY);
}
log_lxx_sq += ljj.ln();
}
log_lxx_sq *= T::from_f64(2.0);
let lyy = fit.tail_l[(kx + p) * t_dim + (kx + p)];
let r_sq = lyy * lyy;
let df = T::from_f64((suff.n_rows - p) as f64);
let sigma_sq = r_sq / df;
if !(sigma_sq.value().is_finite() && sigma_sq.value() > 0.0) {
return T::from_f64(f64::INFINITY);
}
fit.sigma_sq = sigma_sq;
log_lzz_sq + log_lxx_sq + df * sigma_sq.ln()
}
pub(crate) enum LmmDualScratch {
D4(LmmFitScratch<Dual<4>>),
D5(LmmFitScratch<Dual<5>>),
D6(LmmFitScratch<Dual<6>>),
D8(LmmFitScratch<Dual<8>>),
D12(LmmFitScratch<Dual<12>>),
}
impl LmmDualScratch {
pub(crate) fn for_groupings(n_theta: usize, p: usize, g: &LmmGroupings) -> Option<Self> {
Some(match n_theta {
0..=4 => LmmDualScratch::D4(LmmFitScratch::with_groupings(p, g)),
5 => LmmDualScratch::D5(LmmFitScratch::with_groupings(p, g)),
6 => LmmDualScratch::D6(LmmFitScratch::with_groupings(p, g)),
7..=8 => LmmDualScratch::D8(LmmFitScratch::with_groupings(p, g)),
_ => LmmDualScratch::D12(LmmFitScratch::with_groupings(p, g)),
})
}
}
#[allow(clippy::large_enum_variant)]
pub(crate) enum LmmHyperScratch {
H4(LmmFitScratch<HyperDual<4, 10>>),
H5(LmmFitScratch<HyperDual<5, 15>>),
H6(LmmFitScratch<HyperDual<6, 21>>),
H8(LmmFitScratch<HyperDual<8, 36>>),
H12(LmmFitScratch<HyperDual<12, 78>>),
}
impl LmmHyperScratch {
pub(crate) fn for_groupings(n_theta: usize, p: usize, g: &LmmGroupings) -> Option<Self> {
Some(match n_theta {
0..=4 => LmmHyperScratch::H4(LmmFitScratch::with_groupings(p, g)),
5 => LmmHyperScratch::H5(LmmFitScratch::with_groupings(p, g)),
6 => LmmHyperScratch::H6(LmmFitScratch::with_groupings(p, g)),
7..=8 => LmmHyperScratch::H8(LmmFitScratch::with_groupings(p, g)),
9..=12 => LmmHyperScratch::H12(LmmFitScratch::with_groupings(p, g)),
_ => return None,
})
}
}
fn run_reml_gradient<const N: usize>(
theta: &[f64],
suff: &LmmSuffStats,
fit: &mut LmmFitScratch<Dual<N>>,
grad: &mut [f64],
) -> DerivStatus {
let n_theta = theta.len();
let zero = Dual::<N> {
v: 0.0,
d: [0.0; N],
};
let mut last = f64::NAN;
let mut theta_full = [zero; crate::consts::MAX_THETA];
for (j, &tj) in theta.iter().enumerate() {
theta_full[j] = Dual { v: tj, d: [0.0; N] };
}
for offset in (0..n_theta).step_by(N) {
let width = N.min(n_theta - offset);
for t in theta_full[..n_theta].iter_mut() {
t.d = [0.0; N];
}
for j in 0..width {
theta_full[offset + j].d[j] = 1.0; }
let dev = reml_deviance(&theta_full[..n_theta], suff, fit);
if !dev.value().is_finite() {
return DerivStatus::NotConverged;
}
grad[offset..offset + width].copy_from_slice(&dev.d[..width]);
last = dev.value();
}
DerivStatus::Ok(last)
}
pub(crate) fn reml_gradient(
theta: &[f64],
suff: &LmmSuffStats,
scratch: &mut LmmDualScratch,
grad: &mut [f64],
) -> DerivStatus {
debug_assert_eq!(theta.len(), suff.groupings.n_theta());
if suff.groupings.extra_slopes_any {
return DerivStatus::Unsupported;
}
match scratch {
LmmDualScratch::D4(fit) => {
precompute_balanced_collapse(suff, fit);
run_reml_gradient::<4>(theta, suff, fit, grad)
}
LmmDualScratch::D5(fit) => {
precompute_balanced_collapse(suff, fit);
run_reml_gradient::<5>(theta, suff, fit, grad)
}
LmmDualScratch::D6(fit) => {
precompute_balanced_collapse(suff, fit);
run_reml_gradient::<6>(theta, suff, fit, grad)
}
LmmDualScratch::D8(fit) => {
precompute_balanced_collapse(suff, fit);
run_reml_gradient::<8>(theta, suff, fit, grad)
}
LmmDualScratch::D12(fit) => {
precompute_balanced_collapse(suff, fit);
run_reml_gradient::<12>(theta, suff, fit, grad)
}
}
}
fn run_reml_hessian<const N: usize, const H: usize>(
theta: &[f64],
suff: &LmmSuffStats,
fit: &mut LmmFitScratch<HyperDual<N, H>>,
grad: &mut [f64],
hess: &mut Mat<f64>,
) -> DerivStatus {
let n_theta = theta.len();
if n_theta > N {
return DerivStatus::Unsupported;
}
let zero = HyperDual::<N, H> {
v: 0.0,
d: [0.0; N],
h: [0.0; H],
};
let mut theta_t = [zero; N];
for (j, &tj) in theta.iter().enumerate() {
let mut d = [0.0; N];
d[j] = 1.0;
theta_t[j] = HyperDual {
v: tj,
d,
h: [0.0; H],
};
}
let dev = reml_deviance(&theta_t[..n_theta], suff, fit);
if !dev.value().is_finite() {
return DerivStatus::NotConverged;
}
grad[..n_theta].copy_from_slice(&dev.d[..n_theta]);
let hlen = n_theta * (n_theta + 1) / 2;
unpack_hessian(hess, &dev.h[..hlen], n_theta);
DerivStatus::Ok(dev.value())
}
pub(crate) fn reml_hessian(
theta: &[f64],
suff: &LmmSuffStats,
scratch: &mut LmmHyperScratch,
grad: &mut [f64],
hess: &mut Mat<f64>,
) -> DerivStatus {
debug_assert_eq!(theta.len(), suff.groupings.n_theta());
if suff.groupings.extra_slopes_any {
return DerivStatus::Unsupported;
}
match scratch {
LmmHyperScratch::H4(fit) => {
precompute_balanced_collapse(suff, fit);
run_reml_hessian::<4, 10>(theta, suff, fit, grad, hess)
}
LmmHyperScratch::H5(fit) => {
precompute_balanced_collapse(suff, fit);
run_reml_hessian::<5, 15>(theta, suff, fit, grad, hess)
}
LmmHyperScratch::H6(fit) => {
precompute_balanced_collapse(suff, fit);
run_reml_hessian::<6, 21>(theta, suff, fit, grad, hess)
}
LmmHyperScratch::H8(fit) => {
precompute_balanced_collapse(suff, fit);
run_reml_hessian::<8, 36>(theta, suff, fit, grad, hess)
}
LmmHyperScratch::H12(fit) => {
precompute_balanced_collapse(suff, fit);
run_reml_hessian::<12, 78>(theta, suff, fit, grad, hess)
}
}
}