use crate::WaldSe;
use bobyqa::{Bobyqa, Config, Status};
use faer::{Mat, MatRef};
use crate::lmm::{LmmGroupings, PIN_THETA, RHO_BEGIN, RHO_END, THETA0, THETA_TRUTH_FLOOR};
pub const PIRLS_MAX_ITERS: usize = 50;
pub const PIRLS_TOL_REL: f64 = 1e-6;
pub const BETA_BOX: f64 = crate::glm::BETA_CAP;
pub const FD_STEP_REL: f64 = 1e-2;
pub struct GlmmFit {
pub converged: bool,
pub boundary_hit: u8,
pub pinned_components: u32,
pub n_eval: usize,
pub tau_squared_hat: f64,
pub joint_t_sq: f64,
pub hessian_fallback: bool,
}
pub struct GlmmWorkspace {
pub groupings: LmmGroupings, pub k: usize, pub p: usize, pub n_theta: usize,
pub z: Mat<f64>, pub m: Mat<f64>, pub solver: Bobyqa, pub params: Vec<f64>, pub lower: Vec<f64>,
pub upper: Vec<f64>,
pub theta_truth: Vec<f64>, pub eta: Vec<f64>,
pub prob: Vec<f64>,
pub w: Vec<f64>,
pub eta_fixed: Vec<f64>, pub m_buf: Vec<f64>, pub z_buf: Vec<f64>, pub mu: Vec<f64>, pub u: Vec<f64>,
pub u_seed: Vec<f64>, pub a: Mat<f64>, pub wm: Mat<f64>, pub a_rhs: Vec<f64>, pub a_blocks: Vec<f64>, pub core_blocks: Vec<f64>, pub coupling: Vec<f64>, pub schur_blk: Vec<f64>, pub lam: Vec<f64>, pub m_core_buf: Vec<f64>, pub cross_val: Vec<f64>, pub cross_col: Vec<u32>, pub n_cross: Vec<u8>, pub xtwx: Mat<f64>, pub xtwm: Mat<f64>, pub ainv_mtwx: Mat<f64>, pub schur: Mat<f64>, pub betas: Vec<f64>, pub var_diag: Vec<f64>, pub t_sq: Vec<f64>, pub fwd_solve: Vec<f64>, pub joint_k_inv: Mat<f64>,
pub joint_sigma_t_chol: Mat<f64>,
pub joint_rhs: Vec<f64>,
pub hess_scratch: Mat<f64>, pub fd_saved: Vec<f64>, pub fd_steps: Vec<f64>, pub warm_seed_active: bool,
}
impl GlmmWorkspace {
pub fn for_cluster_spec(
p: usize,
cluster: &crate::ModelSpec,
max_n: usize,
slope_cols: &[usize],
) -> Self {
let groupings = LmmGroupings::from_cluster_spec(cluster, max_n, slope_cols);
let k = groupings.k_total;
let n_theta = groupings.n_theta();
let q = groupings.primary_q;
let n_primary = groupings.n_primary;
let q_core = q + groupings.nested_per_parent;
let e_crossed = groupings.k_crossed();
let theta_truth = crate::lmm::cluster_theta_truth(cluster);
let (theta0, mut lower, mut upper) = groupings.blind_theta_and_bounds();
let mut params = theta0;
params.extend(std::iter::repeat_n(0.0, p)); lower.extend(std::iter::repeat_n(-BETA_BOX, p));
upper.extend(std::iter::repeat_n(BETA_BOX, p));
let min_diag = groupings
.diagonal_theta()
.iter()
.map(|&i| theta_truth[i].max(THETA_TRUTH_FLOOR))
.fold(f64::INFINITY, f64::min);
let rho_begin = (0.1 * min_diag).min(RHO_BEGIN);
let config = Config {
rho_begin,
rho_end: RHO_END,
..Config::new(n_theta + p)
};
GlmmWorkspace {
groupings,
k,
p,
n_theta,
z: Mat::zeros(max_n, k.max(1)),
m: Mat::zeros(max_n, k.max(1)),
solver: Bobyqa::new(n_theta + p, config)
.expect("BOBYQA config constants are valid by construction"),
params,
lower,
upper,
theta_truth,
eta: vec![0.0; max_n],
prob: vec![0.0; max_n],
w: vec![0.0; max_n],
eta_fixed: vec![0.0; max_n],
m_buf: vec![0.0; max_n * q],
z_buf: vec![0.0; max_n * (q - 1)],
mu: vec![0.0; max_n],
u: vec![0.0; k.max(1)],
u_seed: vec![0.0; k.max(1)],
a: Mat::zeros(k.max(1), k.max(1)),
wm: Mat::zeros(max_n, k.max(1)),
a_rhs: vec![0.0; k.max(1)],
a_blocks: vec![0.0; (q * q * n_primary).max(1)],
core_blocks: vec![0.0; (q_core * q_core * n_primary).max(1)],
coupling: vec![0.0; (q_core * n_primary * e_crossed).max(1)],
schur_blk: vec![0.0; (e_crossed * e_crossed).max(1)],
lam: vec![0.0; q * q],
m_core_buf: vec![0.0; (max_n * q_core).max(1)],
cross_val: vec![0.0; (max_n * crate::lmm::MAX_EXTRA_GROUPINGS).max(1)],
cross_col: vec![0u32; (max_n * crate::lmm::MAX_EXTRA_GROUPINGS).max(1)],
n_cross: vec![0u8; max_n.max(1)],
xtwx: Mat::zeros(p, p),
xtwm: Mat::zeros(p, k.max(1)),
ainv_mtwx: Mat::zeros(k.max(1), p),
schur: Mat::zeros(p, p),
betas: vec![0.0; p],
var_diag: vec![0.0; p],
t_sq: vec![0.0; p],
fwd_solve: vec![0.0; p],
joint_k_inv: Mat::zeros(p, p),
joint_sigma_t_chol: Mat::zeros(p, p),
joint_rhs: vec![0.0; p],
hess_scratch: Mat::zeros((n_theta + p).max(1), (n_theta + p).max(1)),
fd_saved: vec![0.0; n_theta + p],
fd_steps: vec![0.0; n_theta + p],
warm_seed_active: false,
}
}
}
fn glmm_block_chol(blk: &mut [f64], q: usize) -> bool {
for j in 0..q {
let mut d = blk[j * q + j];
for k in 0..j {
d -= blk[j * q + k] * blk[j * q + k];
}
if !(d.is_finite() && d > 0.0) {
return false;
}
let l = d.sqrt();
blk[j * q + j] = l;
for i in (j + 1)..q {
let mut v = blk[i * q + j];
for k in 0..j {
v -= blk[i * q + k] * blk[j * q + k];
}
blk[i * q + j] = v / l;
}
}
true
}
fn glmm_block_solve(l: &[f64], q: usize, b: &mut [f64]) {
for r in 0..q {
let mut v = b[r];
for c in 0..r {
v -= l[r * q + c] * b[c];
}
b[r] = v / l[r * q + r];
}
for r in (0..q).rev() {
let mut v = b[r];
for c in (r + 1)..q {
v -= l[c * q + r] * b[c];
}
b[r] = v / l[r * q + r];
}
}
pub fn build_z(
ws: &mut GlmmWorkspace,
x: MatRef<f64>,
cluster_ids: &[u32],
extra_ids: &[Vec<u32>],
n: usize,
) {
let g = &ws.groupings;
let q = g.primary_q;
for c in 0..ws.k {
for i in 0..n {
ws.z[(i, c)] = 0.0;
}
}
for i in 0..n {
let lvl = cluster_ids[i] as usize;
let base = lvl * q;
ws.z[(i, base)] = 1.0; for d in 0..q - 1 {
ws.z[(i, base + 1 + d)] = x[(i, g.primary_slope_cols[d])];
}
}
for (e, ids) in extra_ids.iter().enumerate() {
let off = g.extra_offsets[e]; #[allow(clippy::needless_range_loop)]
for i in 0..n {
ws.z[(i, off + ids[i] as usize)] = 1.0;
}
}
}
fn fill_z_f64(g: &LmmGroupings, x: MatRef<f64>, z_buf: &mut [f64], n: usize) {
let q = g.primary_q;
for i in 0..n {
for d in 0..q - 1 {
z_buf[i * (q - 1) + d] = x[(i, g.primary_slope_cols[d])];
}
}
}
pub(crate) fn apply_lambda(
groupings: &LmmGroupings,
params: &[f64],
z: MatRef<f64>,
m: &mut Mat<f64>,
lam: &mut [f64],
n: usize,
) {
let q = groupings.primary_q;
let s = groupings.n_primary;
crate::lmm::primary_lambda(¶ms[..groupings.n_theta()], q, lam);
for lvl in 0..s {
let base = lvl * q;
for i in 0..n {
for c in 0..q {
let mut acc = 0.0;
for r in c..q {
acc += z[(i, base + r)] * lam[r * q + c];
}
m[(i, base + c)] = acc;
}
}
}
let base_theta = q * (q + 1) / 2;
debug_assert!(!groupings.extra_slopes_any);
for (e, &off) in groupings.extra_offsets.iter().enumerate() {
let theta_e = params[base_theta + e];
let width = if groupings.nested.map(|nf| nf.vech_start) == Some(base_theta + e) {
s * groupings.nested_per_parent
} else {
groupings
.crossed
.iter()
.find(|cf| cf.vech_start == base_theta + e)
.map(|cf| cf.n_levels)
.expect("an extra grouping is either nested or crossed")
};
for col in off..off + width {
for i in 0..n {
m[(i, col)] = z[(i, col)] * theta_e;
}
}
}
}
#[allow(clippy::too_many_arguments)]
fn build_packed_m(
g: &LmmGroupings,
params: &[f64],
z: MatRef<f64>,
lam: &mut [f64],
cluster_ids: &[u32],
m_core_buf: &mut [f64],
cross_val: &mut [f64],
cross_col: &mut [u32],
n_cross: &mut [u8],
n: usize,
) {
let q = g.primary_q;
let s = g.n_primary;
let np = g.nested_per_parent;
let qc = q + np;
let prim_width = q * s;
let k_family = qc * s;
let base_theta = q * (q + 1) / 2;
let g_cap = crate::lmm::MAX_EXTRA_GROUPINGS;
debug_assert!(!g.extra_slopes_any);
crate::lmm::primary_lambda(¶ms[..g.n_theta()], q, lam);
let theta_nested = g.nested.map(|nf| params[nf.vech_start]).unwrap_or(0.0);
for i in 0..n {
let f = cluster_ids[i] as usize;
for c in 0..q {
let mut acc = 0.0;
for r in c..q {
acc += z[(i, f * q + r)] * lam[r * q + c];
}
m_core_buf[i * qc + c] = acc;
}
for j in 0..np {
let col = prim_width + f * np + j;
m_core_buf[i * qc + q + j] = z[(i, col)] * theta_nested;
}
let mut cnt = 0usize;
for cf in &g.crossed {
let theta = params[cf.vech_start];
if theta == 0.0 {
continue;
}
let off = g.extra_offsets[cf.vech_start - base_theta];
for col in off..off + cf.n_levels {
let zv = z[(i, col)];
if zv != 0.0 {
cross_col[i * g_cap + cnt] = (col - k_family) as u32;
cross_val[i * g_cap + cnt] = zv * theta;
cnt += 1;
break;
}
}
}
n_cross[i] = cnt as u8;
}
}
#[allow(clippy::too_many_arguments)]
fn pirls_solve(
k: usize,
p: usize,
m: MatRef<f64>,
x: MatRef<f64>,
y: &[f64],
beta: &[f64],
eta: &mut [f64],
prob: &mut [f64],
w: &mut [f64],
u: &mut [f64],
eta_fixed: &mut [f64],
mu: &mut [f64],
wm: &mut Mat<f64>,
a: &mut Mat<f64>,
a_rhs: &mut [f64],
n: usize,
) -> (f64, f64, f64, bool) {
use faer::linalg::matmul::triangular::BlockStructure;
use faer::linalg::matmul::{matmul, triangular};
use faer::linalg::solvers::Solve;
use faer::{Accum, MatMut, Par};
let m = m.subrows(0, n);
for i in 0..n {
let mut e = 0.0;
for j in 0..p {
e += x[(i, j)] * beta[j];
}
eta_fixed[i] = e;
}
let mut pen_prev = f64::INFINITY;
let mut converged = false;
let mut dev = f64::NAN;
let mut pen = f64::NAN;
let mut logdet = 0.0;
for _ in 0..PIRLS_MAX_ITERS {
matmul(
MatMut::from_column_major_slice_mut(&mut mu[..n], n, 1),
Accum::Replace,
m,
MatRef::from_column_major_slice(&u[..k], k, 1),
1.0,
Par::Seq,
);
let mut yeta = 0.0;
for i in 0..n {
let e = eta_fixed[i] + mu[i];
eta[i] = e;
yeta += y[i] * e;
}
let lp_sum =
crate::simd_transcendental::pw_and_log1pexp_sum(&eta[..n], &mut prob[..n], &mut w[..n]);
dev = 2.0 * (lp_sum - yeta);
for c in 0..k {
for i in 0..n {
wm[(i, c)] = w[i] * m[(i, c)];
}
}
triangular::matmul(
a.as_mut(),
BlockStructure::TriangularLower,
Accum::Replace,
m.transpose(),
BlockStructure::Rectangular,
wm.as_ref().subrows(0, n),
BlockStructure::Rectangular,
1.0,
Par::Seq,
);
for r in 0..k {
a[(r, r)] += 1.0;
}
for i in 0..n {
mu[i] = w[i] * mu[i] + (y[i] - prob[i]);
}
matmul(
MatMut::from_column_major_slice_mut(&mut a_rhs[..k], k, 1),
Accum::Replace,
m.transpose(),
MatRef::from_column_major_slice(&mu[..n], n, 1),
1.0,
Par::Seq,
);
let ac = match a.as_ref().llt(faer::Side::Lower) {
Ok(c) => c,
Err(_) => return (f64::NAN, f64::NAN, f64::NAN, false),
};
let rhs = MatMut::from_column_major_slice_mut(&mut a_rhs[..k], k, 1usize);
ac.solve_in_place(rhs);
pen = 0.0;
for c in 0..k {
u[c] = a_rhs[c];
pen += u[c] * u[c];
}
let penalized = dev + pen;
if (penalized - pen_prev).abs() < PIRLS_TOL_REL * (1.0 + penalized.abs()) {
converged = true;
for r in 0..k {
logdet += ac.L()[(r, r)].ln();
}
break;
}
pen_prev = penalized;
}
(dev, pen, logdet, converged)
}
#[allow(clippy::too_many_arguments)]
fn pirls_solve_blocked(
g: &crate::lmm::LmmGroupings,
cluster_ids: &[u32],
x: MatRef<f64>,
y: &[f64],
beta: &[f64],
lam: &[f64],
z_buf: &[f64],
m_buf: &mut [f64],
eta: &mut [f64],
prob: &mut [f64],
w: &mut [f64],
u: &mut [f64],
eta_fixed: &mut [f64],
a_blocks: &mut [f64],
a_rhs: &mut [f64],
n: usize,
) -> (f64, f64, f64, bool) {
let q = g.primary_q;
let s = g.n_primary;
let k = q * s;
let p = beta.len();
for i in 0..n {
let mut e = 0.0;
for j in 0..p {
e += x[(i, j)] * beta[j];
}
eta_fixed[i] = e;
}
for i in 0..n {
for c in 0..q {
let mut acc = 0.0;
for r in c..q {
let zr = if r == 0 {
1.0
} else {
z_buf[i * (q - 1) + (r - 1)]
};
acc += zr * lam[r * q + c];
}
m_buf[i * q + c] = acc;
}
}
let mut pen_prev = f64::INFINITY;
let mut converged = false;
let (mut dev, mut pen, mut logdet) = (f64::NAN, f64::NAN, 0.0);
for _ in 0..PIRLS_MAX_ITERS {
for v in a_blocks[..s * q * q].iter_mut() {
*v = 0.0;
}
for v in a_rhs[..k].iter_mut() {
*v = 0.0;
}
let mut yeta = 0.0;
for i in 0..n {
let m_row = &m_buf[i * q..i * q + q];
let ubase = cluster_ids[i] as usize * q;
let mut e = eta_fixed[i];
for c in 0..q {
e += m_row[c] * u[ubase + c];
}
eta[i] = e;
yeta += y[i] * e;
}
let lp_sum =
crate::simd_transcendental::pw_and_log1pexp_sum(&eta[..n], &mut prob[..n], &mut w[..n]);
dev = 2.0 * (lp_sum - yeta);
for i in 0..n {
let m_row = &m_buf[i * q..i * q + q];
let f = cluster_ids[i] as usize;
let ubase = f * q;
let ablk = f * q * q;
let wi = w[i];
let resid = y[i] - prob[i];
for r in 0..q {
a_rhs[ubase + r] += m_row[r] * resid;
let wr = wi * m_row[r];
for c in 0..=r {
a_blocks[ablk + r * q + c] += wr * m_row[c];
}
}
}
logdet = 0.0;
pen = 0.0;
for f in 0..s {
let ablk = f * q * q;
let ubase = f * q;
for r in 0..q {
let mut acc = a_rhs[ubase + r];
for c in 0..q {
let (hi, lo) = if r >= c { (r, c) } else { (c, r) };
acc += a_blocks[ablk + hi * q + lo] * u[ubase + c];
}
a_rhs[ubase + r] = acc;
}
for r in 0..q {
a_blocks[ablk + r * q + r] += 1.0;
}
if !glmm_block_chol(&mut a_blocks[ablk..ablk + q * q], q) {
return (f64::NAN, f64::NAN, f64::NAN, false);
}
for r in 0..q {
logdet += a_blocks[ablk + r * q + r].ln();
}
u[ubase..ubase + q].copy_from_slice(&a_rhs[ubase..ubase + q]);
glmm_block_solve(&a_blocks[ablk..ablk + q * q], q, &mut u[ubase..ubase + q]);
for r in 0..q {
pen += u[ubase + r] * u[ubase + r];
}
}
let penalized = dev + pen;
if (penalized - pen_prev).abs() < PIRLS_TOL_REL * (1.0 + penalized.abs()) {
converged = true;
break;
}
pen_prev = penalized;
}
(dev, pen, logdet, converged)
}
fn structured_factor(
g: &crate::lmm::LmmGroupings,
core_blocks: &mut [f64],
coupling: &[f64],
schur_blk: &mut [f64],
) -> Option<f64> {
use crate::lmm::MAX_PRIMARY_Q;
let qc = g.primary_q + g.nested_per_parent;
let s = g.n_primary;
let e = g.k_crossed();
let mut logdet = 0.0;
for f in 0..s {
let cb = f * qc * qc;
if !glmm_block_chol(&mut core_blocks[cb..cb + qc * qc], qc) {
return None;
}
for r in 0..qc {
logdet += core_blocks[cb + r * qc + r].ln();
}
let coup = f * qc * e;
let mut ycol = [0.0_f64; MAX_PRIMARY_Q];
for b in 0..e {
for local in 0..qc {
ycol[local] = coupling[coup + local * e + b];
}
glmm_block_solve(&core_blocks[cb..cb + qc * qc], qc, &mut ycol[..qc]);
#[allow(clippy::needless_range_loop)]
for a in b..e {
let mut acc = 0.0;
for local in 0..qc {
acc += coupling[coup + local * e + a] * ycol[local];
}
schur_blk[a * e + b] -= acc;
}
}
}
if e > 0 {
if !glmm_block_chol(&mut schur_blk[..e * e], e) {
return None;
}
for b in 0..e {
logdet += schur_blk[b * e + b].ln();
}
}
Some(logdet)
}
fn structured_ainv_solve(
g: &crate::lmm::LmmGroupings,
core_blocks: &[f64],
coupling: &[f64],
schur_blk: &[f64],
a_rhs: &mut [f64],
) {
use crate::lmm::MAX_PRIMARY_Q;
let qc = g.primary_q + g.nested_per_parent;
let s = g.n_primary;
let e = g.k_crossed();
let k_family = qc * s;
for f in 0..s {
let cb = f * qc * qc;
let gcb = f * qc;
let coup = f * qc * e;
glmm_block_solve(
&core_blocks[cb..cb + qc * qc],
qc,
&mut a_rhs[gcb..gcb + qc],
);
for b in 0..e {
let mut acc = 0.0;
for local in 0..qc {
acc += coupling[coup + local * e + b] * a_rhs[gcb + local];
}
a_rhs[k_family + b] -= acc;
}
}
if e == 0 {
return;
}
glmm_block_solve(&schur_blk[..e * e], e, &mut a_rhs[k_family..k_family + e]);
for f in 0..s {
let cb = f * qc * qc;
let gcb = f * qc;
let coup = f * qc * e;
let mut v = [0.0_f64; MAX_PRIMARY_Q];
#[allow(clippy::needless_range_loop)]
for local in 0..qc {
let mut acc = 0.0;
for b in 0..e {
acc += coupling[coup + local * e + b] * a_rhs[k_family + b];
}
v[local] = acc;
}
glmm_block_solve(&core_blocks[cb..cb + qc * qc], qc, &mut v[..qc]);
for local in 0..qc {
a_rhs[gcb + local] -= v[local];
}
}
}
#[allow(clippy::too_many_arguments)]
fn pirls_solve_blocked_extras(
g: &crate::lmm::LmmGroupings,
cluster_ids: &[u32],
m_core_buf: &[f64],
cross_val: &[f64],
cross_col: &[u32],
n_cross: &[u8],
x: MatRef<f64>,
y: &[f64],
beta: &[f64],
eta: &mut [f64],
prob: &mut [f64],
w: &mut [f64],
u: &mut [f64],
eta_fixed: &mut [f64],
mu: &mut [f64],
core_blocks: &mut [f64],
coupling: &mut [f64],
schur_blk: &mut [f64],
a_rhs: &mut [f64],
n: usize,
) -> (f64, f64, f64, bool) {
let g_cap = crate::lmm::MAX_EXTRA_GROUPINGS;
let q = g.primary_q;
let np = g.nested_per_parent;
let qc = q + np; let s = g.n_primary;
let prim_width = q * s;
let k_family = qc * s; let e = g.k_crossed(); let k = k_family + e; let p = beta.len();
let core_col = |f: usize, local: usize| -> usize {
if local < q {
f * q + local
} else {
prim_width + f * np + (local - q)
}
};
for i in 0..n {
let mut ef = 0.0;
for j in 0..p {
ef += x[(i, j)] * beta[j];
}
eta_fixed[i] = ef;
}
let mut pen_prev = f64::INFINITY;
let mut converged = false;
let (mut dev, mut pen, mut logdet) = (f64::NAN, f64::NAN, 0.0);
for _ in 0..PIRLS_MAX_ITERS {
let mut yeta = 0.0;
for i in 0..n {
let f = cluster_ids[i] as usize;
let m_core = &m_core_buf[i * qc..i * qc + qc];
let mut mui = 0.0;
for local in 0..qc {
mui += m_core[local] * u[core_col(f, local)];
}
let cbase = i * g_cap;
for z in 0..n_cross[i] as usize {
let b = cross_col[cbase + z] as usize;
mui += cross_val[cbase + z] * u[k_family + b];
}
eta[i] = eta_fixed[i] + mui;
mu[i] = mui; yeta += y[i] * eta[i];
}
let lp_sum =
crate::simd_transcendental::pw_and_log1pexp_sum(&eta[..n], &mut prob[..n], &mut w[..n]);
dev = 2.0 * (lp_sum - yeta);
for i in 0..n {
mu[i] = w[i] * mu[i] + (y[i] - prob[i]);
}
for v in core_blocks[..s * qc * qc].iter_mut() {
*v = 0.0;
}
for v in coupling[..s * qc * e].iter_mut() {
*v = 0.0;
}
for v in schur_blk[..e * e].iter_mut() {
*v = 0.0;
}
for v in a_rhs[..k].iter_mut() {
*v = 0.0;
}
for i in 0..n {
let f = cluster_ids[i] as usize;
let wi = w[i];
let ri = mu[i]; let m_core = &m_core_buf[i * qc..i * qc + qc];
let cbase = i * g_cap;
let ncz = n_cross[i] as usize;
let cb = f * qc * qc;
let gcb = f * qc;
let coup = f * qc * e;
for r in 0..qc {
let mr = m_core[r];
a_rhs[gcb + r] += mr * ri;
let wmr = wi * mr;
for c in 0..=r {
core_blocks[cb + r * qc + c] += wmr * m_core[c];
}
for z in 0..ncz {
coupling[coup + r * e + cross_col[cbase + z] as usize] +=
wmr * cross_val[cbase + z];
}
}
for z in 0..ncz {
let b = cross_col[cbase + z] as usize;
let vb = cross_val[cbase + z];
a_rhs[k_family + b] += vb * ri;
let wvb = wi * vb;
for z2 in 0..ncz {
let b2 = cross_col[cbase + z2] as usize;
if b2 <= b {
schur_blk[b * e + b2] += wvb * cross_val[cbase + z2];
}
}
}
}
for f in 0..s {
let cb = f * qc * qc;
for r in 0..qc {
core_blocks[cb + r * qc + r] += 1.0;
}
}
for b in 0..e {
schur_blk[b * e + b] += 1.0;
}
logdet = match structured_factor(g, core_blocks, coupling, schur_blk) {
Some(ld) => ld,
None => return (f64::NAN, f64::NAN, f64::NAN, false),
};
structured_ainv_solve(g, core_blocks, coupling, schur_blk, a_rhs);
pen = 0.0;
for f in 0..s {
let gcb = f * qc;
for local in 0..qc {
let val = a_rhs[gcb + local];
u[core_col(f, local)] = val;
pen += val * val;
}
}
for b in 0..e {
let val = a_rhs[k_family + b];
u[k_family + b] = val;
pen += val * val;
}
let penalized = dev + pen;
if (penalized - pen_prev).abs() < PIRLS_TOL_REL * (1.0 + penalized.abs()) {
converged = true;
break;
}
pen_prev = penalized;
}
(dev, pen, logdet, converged)
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn laplace_deviance(
groupings: &LmmGroupings,
params: &[f64],
z: MatRef<f64>,
m: &mut Mat<f64>,
lam: &mut [f64],
z_buf: &[f64],
m_buf: &mut [f64],
x: MatRef<f64>,
y: &[f64],
cluster_ids: &[u32],
eta: &mut [f64],
prob: &mut [f64],
w: &mut [f64],
u: &mut [f64],
eta_fixed: &mut [f64],
mu: &mut [f64],
wm: &mut Mat<f64>,
a: &mut Mat<f64>,
a_rhs: &mut [f64],
a_blocks: &mut [f64],
core_blocks: &mut [f64],
coupling: &mut [f64],
schur_blk: &mut [f64],
m_core_buf: &mut [f64],
cross_val: &mut [f64],
cross_col: &mut [u32],
n_cross: &mut [u8],
p: usize,
n: usize,
) -> f64 {
let k = groupings.k_total;
let n_theta = groupings.n_theta();
let (dev, pen, logdet, conv) = if groupings.extra_offsets.is_empty() {
crate::lmm::primary_lambda(¶ms[..n_theta], groupings.primary_q, lam);
pirls_solve_blocked(
groupings,
cluster_ids,
x,
y,
¶ms[n_theta..n_theta + p],
lam,
z_buf,
m_buf,
eta,
prob,
w,
u,
eta_fixed,
a_blocks,
a_rhs,
n,
)
} else if groupings.structured_extras_eligible() {
build_packed_m(
groupings,
params,
z,
lam,
cluster_ids,
m_core_buf,
cross_val,
cross_col,
n_cross,
n,
);
pirls_solve_blocked_extras(
groupings,
cluster_ids,
m_core_buf,
cross_val,
cross_col,
n_cross,
x,
y,
¶ms[n_theta..n_theta + p],
eta,
prob,
w,
u,
eta_fixed,
mu,
core_blocks,
coupling,
schur_blk,
a_rhs,
n,
)
} else {
apply_lambda(groupings, params, z, m, lam, n);
pirls_solve(
k,
p,
m.as_ref(),
x,
y,
¶ms[n_theta..n_theta + p],
eta,
prob,
w,
u,
eta_fixed,
mu,
wm,
a,
a_rhs,
n,
)
};
if !conv || !dev.is_finite() {
return f64::INFINITY;
}
dev + pen + 2.0 * logdet
}
pub(crate) fn laplace_deviance_at(
ws: &mut GlmmWorkspace,
x: MatRef<f64>,
y: &[f64],
cluster_ids: &[u32],
n: usize,
) -> f64 {
let kk = ws.k.max(1);
if ws.warm_seed_active {
ws.u[..kk].copy_from_slice(&ws.u_seed[..kk]);
} else {
for v in ws.u[..kk].iter_mut() {
*v = 0.0;
}
}
let GlmmWorkspace {
groupings,
params: prm,
p,
z,
m,
lam,
z_buf,
m_buf,
eta,
prob,
w,
u,
eta_fixed,
mu,
wm,
a,
a_rhs,
a_blocks,
core_blocks,
coupling,
schur_blk,
m_core_buf,
cross_val,
cross_col,
n_cross,
..
} = ws;
laplace_deviance(
groupings,
&prm[..],
z.as_ref(),
m,
lam,
z_buf,
m_buf,
x,
y,
cluster_ids,
eta,
prob,
w,
u,
eta_fixed,
mu,
wm,
a,
a_rhs,
a_blocks,
core_blocks,
coupling,
schur_blk,
m_core_buf,
cross_val,
cross_col,
n_cross,
*p,
n,
)
}
#[cfg(test)]
pub(crate) fn glmm_laplace_deviance(
params: &[f64],
ws: &mut GlmmWorkspace,
x: MatRef<f64>,
y: &[f64],
cluster_ids: &[u32],
n: usize,
) -> f64 {
ws.params[..params.len()].copy_from_slice(params);
fill_z_f64(&ws.groupings, x, &mut ws.z_buf, n);
laplace_deviance_at(ws, x, y, cluster_ids, n)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum FdHessianStatus {
Ok,
NonPdFellBackToRx,
}
fn fd_eval(
ws: &mut GlmmWorkspace,
coords: &[usize],
deltas: &[f64],
x: MatRef<f64>,
y: &[f64],
cluster_ids: &[u32],
n: usize,
) -> f64 {
let m = ws.fd_saved.len();
ws.params[..m].copy_from_slice(&ws.fd_saved[..m]);
for (&c, &d) in coords.iter().zip(deltas) {
ws.params[c] += d;
}
laplace_deviance_at(ws, x, y, cluster_ids, n)
}
pub(crate) fn rx_cov_into(
ws: &mut GlmmWorkspace,
x: MatRef<f64>,
cluster_ids: &[u32],
p: usize,
n: usize,
out_cov: &mut Mat<f64>,
) -> bool {
use faer::linalg::solvers::Solve;
let inf_ok = if ws.groupings.extra_offsets.is_empty() {
blocked_schur_fill(ws, x, cluster_ids, n)
} else if ws.groupings.structured_extras_eligible() {
structured_schur_fill(ws, x, cluster_ids, n)
} else {
dense_schur_fill(ws, x, n)
};
if !inf_ok {
return false;
}
let chol = match ws.schur.as_ref().llt(faer::Side::Lower) {
Ok(c) => c,
Err(_) => return false,
};
let mut inv = Mat::<f64>::identity(p, p);
chol.solve_in_place(inv.as_mut());
for a in 0..p {
for b in 0..p {
out_cov[(a, b)] = inv[(a, b)];
}
}
true
}
pub fn fd_hessian_cov(
ws: &mut GlmmWorkspace,
x: MatRef<f64>,
y: &[f64],
cluster_ids: &[u32],
p: usize,
n: usize,
out_cov: &mut Mat<f64>,
) -> FdHessianStatus {
use faer::linalg::solvers::Solve;
let m = ws.params.len();
let n_theta = ws.n_theta;
ws.fd_saved[..m].copy_from_slice(&ws.params[..m]);
for k in 0..m {
ws.fd_steps[k] = FD_STEP_REL * ws.fd_saved[k].abs().max(1.0);
}
if ws.groupings.extra_offsets.is_empty() {
let GlmmWorkspace {
groupings, z_buf, ..
} = &mut *ws;
fill_z_f64(groupings, x, z_buf, n);
}
macro_rules! fallback {
() => {{
let _ = fd_eval(ws, &[], &[], x, y, cluster_ids, n);
let ok = rx_cov_into(ws, x, cluster_ids, p, n, out_cov);
debug_assert!(ok, "RX fallback Schur must be PD at a converged fit");
if !ok {
for a in 0..p {
for b in 0..p {
out_cov[(a, b)] = f64::NAN;
}
}
}
ws.params[..m].copy_from_slice(&ws.fd_saved[..m]);
ws.warm_seed_active = false; return FdHessianStatus::NonPdFellBackToRx;
}};
}
let f0 = fd_eval(ws, &[], &[], x, y, cluster_ids, n);
if !f0.is_finite() {
fallback!();
}
let kk = ws.k.max(1);
ws.u_seed[..kk].copy_from_slice(&ws.u[..kk]);
ws.warm_seed_active = true;
macro_rules! second_diff {
($k:expr, $s:expr) => {{
let s = $s;
let fp = fd_eval(ws, &[$k], &[s], x, y, cluster_ids, n);
let fm = fd_eval(ws, &[$k], &[-s], x, y, cluster_ids, n);
if !(fp.is_finite() && fm.is_finite()) {
fallback!();
}
(fp - 2.0 * f0 + fm) / (s * s)
}};
}
macro_rules! mixed_diff {
($i:expr, $j:expr, $si:expr, $sj:expr) => {{
let (si, sj) = ($si, $sj);
let fpp = fd_eval(ws, &[$i, $j], &[si, sj], x, y, cluster_ids, n);
let fpm = fd_eval(ws, &[$i, $j], &[si, -sj], x, y, cluster_ids, n);
let fmp = fd_eval(ws, &[$i, $j], &[-si, sj], x, y, cluster_ids, n);
let fmm = fd_eval(ws, &[$i, $j], &[-si, -sj], x, y, cluster_ids, n);
if !(fpp.is_finite() && fpm.is_finite() && fmp.is_finite() && fmm.is_finite()) {
fallback!();
}
(fpp - fpm - fmp + fmm) / (4.0 * si * sj)
}};
}
for i in 0..m {
let hi = ws.fd_steps[i];
let d_full = second_diff!(i, hi);
let d_half = second_diff!(i, hi * 0.5);
let hii = (4.0 * d_half - d_full) / 3.0;
ws.hess_scratch[(i, i)] = hii;
for j in (i + 1)..m {
let hj = ws.fd_steps[j];
let mx_full = mixed_diff!(i, j, hi, hj);
let mx_half = mixed_diff!(i, j, hi * 0.5, hj * 0.5);
let hij = (4.0 * mx_half - mx_full) / 3.0;
ws.hess_scratch[(i, j)] = hij;
ws.hess_scratch[(j, i)] = hij;
}
}
let chol = match ws.hess_scratch.as_ref().llt(faer::Side::Lower) {
Ok(c) => c,
Err(_) => fallback!(),
};
let mut inv = Mat::<f64>::identity(m, m);
chol.solve_in_place(inv.as_mut());
for a in 0..p {
for b in 0..p {
out_cov[(a, b)] = 2.0 * inv[(n_theta + a, n_theta + b)];
}
}
ws.params[..m].copy_from_slice(&ws.fd_saved[..m]);
ws.warm_seed_active = false; FdHessianStatus::Ok
}
#[allow(clippy::too_many_arguments)]
pub fn fit_glmm(
ws: &mut GlmmWorkspace,
x: MatRef<f64>,
y: &[f64],
cluster_ids: &[u32],
target_indices: &[u32],
theta_start: Option<&[f64]>,
beta_start: &[f64],
n: usize,
wald_se: WaldSe,
) -> GlmmFit {
let (k, p, n_theta) = (ws.k, ws.p, ws.n_theta);
match theta_start {
Some(ts) => {
for (t, &v) in ws.params[..n_theta].iter_mut().zip(ts) {
*t = v.max(THETA_TRUTH_FLOOR);
}
}
None => {
for t in ws.params[..n_theta].iter_mut() {
*t = THETA0;
}
}
}
for (j, &b) in beta_start.iter().enumerate().take(p) {
ws.params[n_theta + j] = b.clamp(-BETA_BOX, BETA_BOX);
}
for v in ws.u_seed[..k].iter_mut() {
*v = 0.0;
}
let GlmmWorkspace {
solver,
params,
lower,
upper,
groupings,
z,
m,
eta,
prob,
w,
u,
u_seed,
eta_fixed,
mu,
wm,
a,
a_rhs,
a_blocks,
core_blocks,
coupling,
schur_blk,
lam,
z_buf,
m_buf,
m_core_buf,
cross_val,
cross_col,
n_cross,
p: pf,
..
} = ws;
if groupings.extra_offsets.is_empty() {
fill_z_f64(groupings, x, z_buf, n);
}
let mut best_obj = f64::INFINITY;
let out = solver.minimize(
|gamma| {
u[..k].copy_from_slice(&u_seed[..k]);
let obj = laplace_deviance(
groupings,
gamma,
z.as_ref(),
m,
lam,
z_buf,
m_buf,
x,
y,
cluster_ids,
eta,
prob,
w,
u,
eta_fixed,
mu,
wm,
a,
a_rhs,
a_blocks,
core_blocks,
coupling,
schur_blk,
m_core_buf,
cross_val,
cross_col,
n_cross,
*pf,
n,
);
if obj < best_obj {
best_obj = obj;
u_seed[..k].copy_from_slice(&u[..k]);
}
obj
},
params,
lower,
upper,
);
debug_assert!(out.status != Status::InvalidArgs);
let ok = matches!(out.status, Status::Converged);
let diag = ws.groupings.diagonal_theta();
let mut pinned_components = 0u32;
let mut pinned = false;
if ok {
for (kk, &ti) in diag.iter().enumerate() {
if ws.params[ti] <= PIN_THETA {
ws.params[ti] = 0.0;
pinned = true;
pinned_components |= 1 << kk;
}
}
}
if ok {
ws.u[..k].copy_from_slice(&ws.u_seed[..k]);
let GlmmWorkspace {
groupings,
params,
p,
z,
m,
lam,
z_buf,
m_buf,
eta,
prob,
w,
u,
eta_fixed,
mu,
wm,
a,
a_rhs,
a_blocks,
core_blocks,
coupling,
schur_blk,
m_core_buf,
cross_val,
cross_col,
n_cross,
..
} = ws;
let _ = laplace_deviance(
groupings,
¶ms[..],
z.as_ref(),
m,
lam,
z_buf,
m_buf,
x,
y,
cluster_ids,
eta,
prob,
w,
u,
eta_fixed,
mu,
wm,
a,
a_rhs,
a_blocks,
core_blocks,
coupling,
schur_blk,
m_core_buf,
cross_val,
cross_col,
n_cross,
*p,
n,
);
}
if !ok {
return nan_fit(ws, target_indices, out.n_eval);
}
for j in 0..p {
ws.betas[j] = ws.params[n_theta + j];
}
let mut hessian_fallback = false;
let joint_t_sq = match wald_se {
WaldSe::Rx => {
let inf_ok = if ws.groupings.extra_offsets.is_empty() {
blocked_schur_fill(ws, x, cluster_ids, n)
} else if ws.groupings.structured_extras_eligible() {
structured_schur_fill(ws, x, cluster_ids, n)
} else {
dense_schur_fill(ws, x, n)
};
if !inf_ok {
return nan_fit(ws, target_indices, out.n_eval);
}
let sc = match ws.schur.as_ref().llt(faer::Side::Lower) {
Ok(c) => c,
Err(_) => return nan_fit(ws, target_indices, out.n_eval),
};
let lschur = sc.L();
for &tj in target_indices {
let tj = tj as usize;
for i in 0..p {
let mut acc = if i == tj { 1.0 } else { 0.0 };
for kk in 0..i {
acc -= lschur[(i, kk)] * ws.fwd_solve[kk];
}
ws.fwd_solve[i] = acc / lschur[(i, i)];
}
let vd: f64 = ws.fwd_solve[..p].iter().map(|v| v * v).sum();
ws.var_diag[tj] = vd;
ws.t_sq[tj] = if vd.is_finite() && vd > 0.0 {
ws.betas[tj] * ws.betas[tj] / vd
} else {
f64::NAN
};
}
if target_indices.is_empty() {
f64::NAN
} else {
crate::lme::joint_wald_chi_sq(
ws.schur.as_ref(),
&ws.betas,
1.0,
target_indices,
ws.joint_k_inv.as_mut(),
ws.joint_sigma_t_chol.as_mut(),
&mut ws.joint_rhs,
)
}
}
WaldSe::Hessian => {
let mut cov = Mat::<f64>::zeros(p, p);
let status = fd_hessian_cov(ws, x, y, cluster_ids, p, n, &mut cov);
if !cov[(0, 0)].is_finite() {
return nan_fit(ws, target_indices, out.n_eval);
}
hessian_fallback = matches!(status, FdHessianStatus::NonPdFellBackToRx);
for &tj in target_indices {
let tj = tj as usize;
let vd = cov[(tj, tj)];
ws.var_diag[tj] = vd;
ws.t_sq[tj] = if vd.is_finite() && vd > 0.0 {
ws.betas[tj] * ws.betas[tj] / vd
} else {
f64::NAN
};
}
if target_indices.is_empty() {
f64::NAN
} else {
use faer::linalg::solvers::Solve;
match cov.as_ref().llt(faer::Side::Lower) {
Ok(chol) => {
let mut inv = Mat::<f64>::identity(p, p);
chol.solve_in_place(inv.as_mut());
for a in 0..p {
for b in 0..p {
ws.schur[(a, b)] = inv[(a, b)];
}
}
crate::lme::joint_wald_chi_sq(
ws.schur.as_ref(),
&ws.betas,
1.0,
target_indices,
ws.joint_k_inv.as_mut(),
ws.joint_sigma_t_chol.as_mut(),
&mut ws.joint_rhs,
)
}
Err(_) => f64::NAN,
}
}
}
};
crate::lmm::primary_lambda(&ws.params[..n_theta], ws.groupings.primary_q, &mut ws.lam);
let q = ws.groupings.primary_q;
let mut d00 = 0.0;
for r in 0..q {
d00 += ws.lam[r] * ws.lam[r];
}
GlmmFit {
converged: true,
boundary_hit: u8::from(pinned),
pinned_components,
n_eval: out.n_eval,
tau_squared_hat: d00,
joint_t_sq,
hessian_fallback,
}
}
fn nan_fit(ws: &mut GlmmWorkspace, targets: &[u32], n_eval: usize) -> GlmmFit {
for v in ws.betas.iter_mut() {
*v = f64::NAN;
}
for &t in targets {
ws.var_diag[t as usize] = f64::NAN;
ws.t_sq[t as usize] = f64::NAN;
}
GlmmFit {
converged: false,
boundary_hit: 2,
pinned_components: 0,
n_eval,
tau_squared_hat: f64::NAN,
joint_t_sq: f64::NAN,
hessian_fallback: false,
}
}
fn dense_schur_fill(ws: &mut GlmmWorkspace, x: MatRef<f64>, n: usize) -> bool {
use faer::linalg::solvers::Solve;
let (k, p) = (ws.k, ws.p);
for r in 0..p {
for c in 0..=r {
let mut s = 0.0;
for i in 0..n {
s += x[(i, r)] * ws.w[i] * x[(i, c)];
}
ws.xtwx[(r, c)] = s;
ws.xtwx[(c, r)] = s;
}
}
for r in 0..p {
for c in 0..k {
let mut s = 0.0;
for i in 0..n {
s += x[(i, r)] * ws.w[i] * ws.m[(i, c)];
}
ws.xtwm[(r, c)] = s;
}
}
let ac = match ws.a.as_ref().llt(faer::Side::Lower) {
Ok(c) => c,
Err(_) => return false,
};
for r in 0..k {
for c in 0..p {
ws.ainv_mtwx[(r, c)] = ws.xtwm[(c, r)];
}
}
ac.solve_in_place(ws.ainv_mtwx.as_mut());
for r in 0..p {
for c in 0..p {
let mut s = ws.xtwx[(r, c)];
for j in 0..k {
s -= ws.xtwm[(r, j)] * ws.ainv_mtwx[(j, c)];
}
ws.schur[(r, c)] = s;
}
}
true
}
fn blocked_schur_fill(
ws: &mut GlmmWorkspace,
x: MatRef<f64>,
cluster_ids: &[u32],
n: usize,
) -> bool {
let (p, q, s) = (ws.p, ws.groupings.primary_q, ws.groupings.n_primary);
let k = q * s;
for r in 0..p {
for c in 0..=r {
let mut sm = 0.0;
for i in 0..n {
sm += x[(i, r)] * ws.w[i] * x[(i, c)];
}
ws.xtwx[(r, c)] = sm;
ws.xtwx[(c, r)] = sm;
}
}
for r in 0..p {
for c in 0..k {
ws.xtwm[(r, c)] = 0.0;
}
}
for i in 0..n {
let f = cluster_ids[i] as usize;
let mut m_row = [0.0_f64; crate::lmm::MAX_PRIMARY_Q];
#[allow(clippy::needless_range_loop)]
for c in 0..q {
let mut acc = 0.0;
for rr in c..q {
let zr = if rr == 0 {
1.0
} else {
x[(i, ws.groupings.primary_slope_cols[rr - 1])]
};
acc += zr * ws.lam[rr * q + c];
}
m_row[c] = acc;
}
let wi = ws.w[i];
for r in 0..p {
let xw = x[(i, r)] * wi;
#[allow(clippy::needless_range_loop)]
for c in 0..q {
ws.xtwm[(r, f * q + c)] += xw * m_row[c];
}
}
}
for f in 0..s {
let ablk = f * q * q;
for col in 0..p {
let mut rhs = [0.0_f64; crate::lmm::MAX_PRIMARY_Q];
#[allow(clippy::needless_range_loop)]
for c in 0..q {
rhs[c] = ws.xtwm[(col, f * q + c)];
}
glmm_block_solve(&ws.a_blocks[ablk..ablk + q * q], q, &mut rhs[..q]);
#[allow(clippy::needless_range_loop)]
for c in 0..q {
ws.ainv_mtwx[(f * q + c, col)] = rhs[c];
}
}
}
for r in 0..p {
for c in 0..p {
let mut sm = ws.xtwx[(r, c)];
for j in 0..k {
sm -= ws.xtwm[(r, j)] * ws.ainv_mtwx[(j, c)];
}
ws.schur[(r, c)] = sm;
}
}
true
}
fn structured_schur_fill(
ws: &mut GlmmWorkspace,
x: MatRef<f64>,
cluster_ids: &[u32],
n: usize,
) -> bool {
let p = ws.p;
let g = &ws.groupings;
let (q, np, s) = (g.primary_q, g.nested_per_parent, g.n_primary);
let qc = q + np;
let e = g.k_crossed();
let prim_width = q * s;
let k_family = qc * s;
let k = ws.k;
let g_cap = crate::lmm::MAX_EXTRA_GROUPINGS;
let core_col = |f: usize, local: usize| -> usize {
if local < q {
f * q + local
} else {
prim_width + f * np + (local - q)
}
};
for r in 0..p {
for c in 0..=r {
let mut sm = 0.0;
for i in 0..n {
sm += x[(i, r)] * ws.w[i] * x[(i, c)];
}
ws.xtwx[(r, c)] = sm;
ws.xtwx[(c, r)] = sm;
}
}
for r in 0..p {
for c in 0..k {
ws.xtwm[(r, c)] = 0.0;
}
}
for i in 0..n {
let f = cluster_ids[i] as usize;
let wi = ws.w[i];
let cbase = i * g_cap;
let ncz = ws.n_cross[i] as usize;
for r in 0..p {
let xw = x[(i, r)] * wi;
for local in 0..qc {
ws.xtwm[(r, core_col(f, local))] += xw * ws.m_core_buf[i * qc + local];
}
for z in 0..ncz {
let b = ws.cross_col[cbase + z] as usize;
ws.xtwm[(r, k_family + b)] += xw * ws.cross_val[cbase + z];
}
}
}
for c in 0..p {
for f in 0..s {
for local in 0..qc {
ws.a_rhs[f * qc + local] = ws.xtwm[(c, core_col(f, local))];
}
}
for b in 0..e {
ws.a_rhs[k_family + b] = ws.xtwm[(c, k_family + b)];
}
structured_ainv_solve(
&ws.groupings,
&ws.core_blocks,
&ws.coupling,
&ws.schur_blk,
&mut ws.a_rhs,
);
for f in 0..s {
for local in 0..qc {
ws.ainv_mtwx[(core_col(f, local), c)] = ws.a_rhs[f * qc + local];
}
}
for b in 0..e {
ws.ainv_mtwx[(k_family + b, c)] = ws.a_rhs[k_family + b];
}
}
for r in 0..p {
for c in 0..p {
let mut sm = ws.xtwx[(r, c)];
for j in 0..k {
sm -= ws.xtwm[(r, j)] * ws.ainv_mtwx[(j, c)];
}
ws.schur[(r, c)] = sm;
}
}
true
}
#[cfg(test)]
mod tests {
use super::*;
use crate::test_support::{intercept_only_spec, TestWs};
use crate::{Estimator, Grouping, GroupingRelation, ModelSpec, Sizing, SlopeTerm, WaldSe};
use faer::linalg::solvers::Solve;
fn lcg(s: &mut u64) -> f64 {
*s = s
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
((*s >> 11) as f64) / ((1u64 << 53) as f64) - 0.5
}
fn glmm_intercept_dataset_layout(contiguous: bool) -> (Mat<f64>, Vec<f64>, Vec<u32>) {
let (n, nc) = (80usize, 8usize);
let mut st = 7u64;
let u0: Vec<f64> = (0..nc).map(|_| 0.6 * lcg(&mut st)).collect();
let mut x = Mat::<f64>::zeros(n, 2);
let mut y = vec![0.0f64; n];
let mut ids = vec![0u32; n];
for i in 0..n {
let c = if contiguous { i / (n / nc) } else { i % nc };
ids[i] = c as u32;
let x1 = lcg(&mut st);
x[(i, 0)] = 1.0;
x[(i, 1)] = x1;
let eta = 0.2 + 0.8 * x1 + u0[c];
let p = 1.0 / (1.0 + (-eta).exp());
y[i] = if lcg(&mut st) + 0.5 < p { 1.0 } else { 0.0 };
}
(x, y, ids)
}
fn glmm_intercept_dataset() -> (Mat<f64>, Vec<f64>, Vec<u32>) {
glmm_intercept_dataset_layout(false)
}
fn glmm_slope_crossed_dataset() -> (Mat<f64>, Vec<f64>, Vec<u32>, Vec<u32>, ModelSpec) {
let (n, n_prim, n_crossed) = (96usize, 8usize, 4usize);
let mut st = 13u64;
let u0: Vec<f64> = (0..n_prim).map(|_| 0.6 * lcg(&mut st)).collect();
let u1: Vec<f64> = (0..n_prim).map(|_| 0.4 * lcg(&mut st)).collect();
let uc: Vec<f64> = (0..n_crossed).map(|_| 0.5 * lcg(&mut st)).collect();
let mut x = Mat::<f64>::zeros(n, 2);
let mut y = vec![0.0f64; n];
let mut ids = vec![0u32; n];
let mut crossed = vec![0u32; n];
for i in 0..n {
let c = i % n_prim;
let cc = i % n_crossed;
ids[i] = c as u32;
crossed[i] = cc as u32;
let x1 = lcg(&mut st);
x[(i, 0)] = 1.0;
x[(i, 1)] = x1;
let eta = 0.2 + 0.8 * x1 + u0[c] + u1[c] * x1 + uc[cc];
let p = 1.0 / (1.0 + (-eta).exp());
y[i] = if lcg(&mut st) + 0.5 < p { 1.0 } else { 0.0 };
}
let cluster = ModelSpec {
sizing: Sizing::FixedClusters {
n_clusters: n_prim as u32,
},
tau_squared: 0.25,
slopes: vec![SlopeTerm {
column: 0,
variance: 0.16,
corr_with_intercept: 0.0,
corr_with: vec![],
}],
extra_groupings: vec![Grouping {
relation: GroupingRelation::Crossed {
n_clusters: n_crossed as u32,
},
tau_squared: 0.25,
slopes: vec![],
}],
estimator: Estimator::Glm,
wald_se: WaldSe::Hessian,
};
(x, y, ids, crossed, cluster)
}
fn brute_force_intercept_laplace(
theta0: f64,
beta: &[f64],
x: &Mat<f64>,
y: &[f64],
ids: &[u32],
nc: usize,
) -> f64 {
let (n, p) = (x.nrows(), x.ncols());
let mut m = Mat::<f64>::zeros(n, nc);
for i in 0..n {
m[(i, ids[i] as usize)] = theta0;
}
let mut u = vec![0.0f64; nc];
let pen_dev = |u: &[f64], w_out: Option<&mut [f64]>| -> f64 {
let mut d = 0.0;
let mut pen = 0.0;
let mut wbuf = vec![0.0; n];
for i in 0..n {
let mut eta = 0.0;
for j in 0..p {
eta += x[(i, j)] * beta[j];
}
for c in 0..nc {
eta += m[(i, c)] * u[c];
}
let pi = 1.0 / (1.0 + (-eta).exp());
d += if eta > 0.0 {
eta + (-eta).exp().ln_1p()
} else {
eta.exp().ln_1p()
} - y[i] * eta;
wbuf[i] = (pi * (1.0 - pi)).max(1e-6);
let _ = pi;
}
for &uc in u {
pen += uc * uc;
}
if let Some(w) = w_out {
w.copy_from_slice(&wbuf);
}
2.0 * d + pen
};
let mut w = vec![0.0; n];
for _ in 0..50 {
let mut eta = vec![0.0; n];
let mut pvec = vec![0.0; n];
for i in 0..n {
let mut e = 0.0;
for j in 0..p {
e += x[(i, j)] * beta[j];
}
for c in 0..nc {
e += m[(i, c)] * u[c];
}
eta[i] = e;
let pi = 1.0 / (1.0 + (-e).exp());
pvec[i] = pi;
w[i] = (pi * (1.0 - pi)).max(1e-6);
}
let mut g = vec![0.0; nc];
for c in 0..nc {
let mut s = 0.0;
for i in 0..n {
s += m[(i, c)] * (y[i] - pvec[i]);
}
g[c] = 2.0 * u[c] - 2.0 * s;
}
let mut h = Mat::<f64>::zeros(nc, nc);
for a in 0..nc {
for b in 0..nc {
let mut s = 0.0;
for i in 0..n {
s += m[(i, a)] * w[i] * m[(i, b)];
}
h[(a, b)] = 2.0 * (s + if a == b { 1.0 } else { 0.0 });
}
}
let hc = h.as_ref().llt(faer::Side::Lower).unwrap();
let mut step = Mat::<f64>::zeros(nc, 1);
for c in 0..nc {
step[(c, 0)] = g[c];
}
hc.solve_in_place(step.as_mut());
let mut max = 0.0f64;
for c in 0..nc {
u[c] -= step[(c, 0)];
max = max.max(step[(c, 0)].abs());
}
if max < 1e-10 {
break;
}
}
let _ = pen_dev(&u, Some(&mut w));
let mut a = Mat::<f64>::zeros(nc, nc);
for r in 0..nc {
for c in 0..nc {
let mut s = 0.0;
for i in 0..n {
s += m[(i, r)] * w[i] * m[(i, c)];
}
a[(r, c)] = s + if r == c { 1.0 } else { 0.0 };
}
}
let ac = a.as_ref().llt(faer::Side::Lower).unwrap();
let mut logdet = 0.0;
for r in 0..nc {
logdet += ac.L()[(r, r)].ln();
}
pen_dev(&u, None) + 2.0 * logdet
}
#[test]
fn laplace_deviance_matches_brute_force_intercept() {
let (xf64, y, ids) = glmm_intercept_dataset();
let beta = [0.2_f64, 0.8];
let want = brute_force_intercept_laplace(0.5, &beta, &xf64, &y, &ids, 8);
let cluster = intercept_only_spec(Sizing::FixedClusters { n_clusters: 8 }, 0.25);
let mut ws = GlmmWorkspace::for_cluster_spec(2, &cluster, 80, &[]);
build_z(&mut ws, xf64.as_ref(), &ids, &[], 80);
ws.params[0] = 0.5;
ws.params[1] = beta[0];
ws.params[2] = beta[1];
let got = glmm_laplace_deviance(&ws.params.clone(), &mut ws, xf64.as_ref(), &y, &ids, 80);
assert!(
(got - want).abs() < 1e-6,
"laplace dev: got {got}, want {want}"
);
}
#[test]
fn laplace_deviance_collapses_to_glm_at_theta_zero() {
let (xf64, y, ids) = glmm_intercept_dataset();
let beta = [0.2_f64, 0.8];
let cluster = intercept_only_spec(Sizing::FixedClusters { n_clusters: 8 }, 0.25);
let mut ws = GlmmWorkspace::for_cluster_spec(2, &cluster, 80, &[]);
build_z(&mut ws, xf64.as_ref(), &ids, &[], 80);
ws.params[0] = 0.0;
ws.params[1] = beta[0];
ws.params[2] = beta[1];
let got = glmm_laplace_deviance(&ws.params.clone(), &mut ws, xf64.as_ref(), &y, &ids, 80);
let mut d = 0.0;
for i in 0..80 {
let eta = beta[0] + beta[1] * xf64[(i, 1)];
d += if eta > 0.0 {
eta + (-eta).exp().ln_1p()
} else {
eta.exp().ln_1p()
} - y[i] * eta;
}
let want = 2.0 * d;
assert!(
(got - want).abs() < 1e-9,
"collapse: got {got}, want {want}"
);
}
#[test]
fn build_z_width_general_populates_all_columns() {
let (xf64, _y, ids, crossed_ids, cluster) = glmm_slope_crossed_dataset();
let n = ids.len();
let mut ws = GlmmWorkspace::for_cluster_spec(2, &cluster, n, &[1]);
build_z(
&mut ws,
xf64.as_ref(),
&ids,
std::slice::from_ref(&crossed_ids),
n,
);
let mut touched = vec![false; ws.k];
#[allow(clippy::needless_range_loop)]
for c in 0..ws.k {
for i in 0..n {
if ws.z[(i, c)] != 0.0 {
touched[c] = true;
}
}
}
assert!(
touched.iter().all(|&t| t),
"every RE column must be populated — offset wiring"
);
}
#[test]
fn apply_lambda_handles_nonmonotonic_extra_offsets() {
let (n_prim, n_crossed, n_per_parent) = (4usize, 3usize, 2usize);
let cluster = ModelSpec {
sizing: Sizing::FixedClusters {
n_clusters: n_prim as u32,
},
tau_squared: 0.25,
slopes: vec![],
extra_groupings: vec![
Grouping {
relation: GroupingRelation::Crossed {
n_clusters: n_crossed as u32,
},
tau_squared: 0.16,
slopes: vec![],
},
Grouping {
relation: GroupingRelation::NestedWithin {
n_per_parent: n_per_parent as u32,
},
tau_squared: 0.09,
slopes: vec![],
},
],
estimator: Estimator::Glm,
wald_se: WaldSe::Hessian,
};
let n = 8usize;
let ws = GlmmWorkspace::for_cluster_spec(2, &cluster, n, &[]);
let g = &ws.groupings;
assert!(
g.extra_offsets[0] > g.extra_offsets[1],
"fixture must produce non-monotonic offsets, got {:?}",
g.extra_offsets
);
let base_theta = 1usize;
let (theta_crossed, theta_nested) = (2.0_f64, 3.0_f64);
let mut params = vec![0.0; g.n_theta()];
params[0] = 1.0; params[base_theta] = theta_crossed;
params[base_theta + 1] = theta_nested;
let mut z = Mat::<f64>::zeros(n, g.k_total);
for i in 0..n {
for c in 0..g.k_total {
z[(i, c)] = 1.0;
}
}
let mut m = Mat::<f64>::zeros(n, g.k_total);
let mut lam = vec![0.0; g.primary_q * g.primary_q];
apply_lambda(g, ¶ms, z.as_ref(), &mut m, &mut lam, n);
let coff = g.extra_offsets[0];
for c in coff..coff + n_crossed {
assert!(
(m[(0, c)] - theta_crossed).abs() < 1e-12,
"crossed col {c} = {}, want {theta_crossed}",
m[(0, c)]
);
}
let noff = g.extra_offsets[1];
for c in noff..noff + n_prim * n_per_parent {
assert!(
(m[(0, c)] - theta_nested).abs() < 1e-12,
"nested col {c} = {}, want {theta_nested}",
m[(0, c)]
);
}
}
#[test]
fn fit_glmm_recovers_direction_and_finite_inference() {
let (xf64, y, ids) = glmm_intercept_dataset();
let cluster = intercept_only_spec(Sizing::FixedClusters { n_clusters: 8 }, 0.25);
let mut ws = GlmmWorkspace::for_cluster_spec(2, &cluster, 80, &[]);
build_z(&mut ws, xf64.as_ref(), &ids, &[], 80);
let targets = [1u32];
let beta_truth = [0.2_f64, 0.8];
let fit = fit_glmm(
&mut ws,
xf64.as_ref(),
&y,
&ids,
&targets,
Some(&[0.5]),
&beta_truth,
80,
WaldSe::Rx,
);
assert!(fit.converged);
assert!(
ws.betas[1] > 0.3,
"β̂₁ should be positive (truth 0.8), got {}",
ws.betas[1]
);
assert!(
ws.t_sq[1].is_finite() && ws.t_sq[1] > 0.0,
"t²[1] = {} must be finite and strictly positive",
ws.t_sq[1]
);
assert!(fit.tau_squared_hat.is_finite() && fit.tau_squared_hat >= 0.0);
}
#[derive(serde::Deserialize)]
struct HessianFixture {
n: usize,
x: Vec<Vec<f64>>,
y: Vec<f64>,
cluster_ids: Vec<u32>,
theta: f64,
beta: Vec<f64>,
vcov_hessian: Vec<Vec<f64>>,
#[allow(dead_code)]
vcov_rx: Vec<Vec<f64>>,
}
fn load_hessian_fixture() -> HessianFixture {
let path = concat!(
env!("CARGO_MANIFEST_DIR"),
"/tests/fixtures/glmm_hessian_vcov.json"
);
let s = std::fs::read_to_string(path).expect("read hessian fixture");
serde_json::from_str(&s).expect("parse hessian fixture")
}
#[test]
fn fd_hessian_cov_matches_glmer_use_hessian_true() {
let fx = load_hessian_fixture();
let n = fx.n;
let p = fx.beta.len();
let n_clusters = fx.cluster_ids.iter().max().unwrap() + 1;
let cluster = intercept_only_spec(Sizing::FixedClusters { n_clusters }, 0.25);
let mut ws = GlmmWorkspace::for_cluster_spec(p, &cluster, n, &[]);
let mut xf64 = Mat::<f64>::zeros(n, p);
for i in 0..n {
for j in 0..p {
xf64[(i, j)] = fx.x[i][j];
}
}
let y = fx.y.clone();
let ids = fx.cluster_ids.clone();
build_z(&mut ws, xf64.as_ref(), &ids, &[], n);
ws.params[0] = fx.theta;
for j in 0..p {
ws.params[1 + j] = fx.beta[j];
}
let _ = laplace_deviance_at(&mut ws, xf64.as_ref(), &y, &ids, n);
let mut rx = Mat::<f64>::zeros(p, p);
assert!(rx_cov_into(&mut ws, xf64.as_ref(), &ids, p, n, &mut rx));
for i in 0..p {
for j in 0..p {
let (got, want) = (rx[(i, j)], fx.vcov_rx[i][j]);
assert!(
(got - want).abs() < 1e-5,
"rx[{i}][{j}] got {got} want {want} (gap {})",
(got - want).abs()
);
}
}
let fit = fit_glmm(
&mut ws,
xf64.as_ref(),
&y,
&ids,
&[1u32],
None,
&vec![0.0; p],
n,
WaldSe::Rx,
);
assert!(fit.converged, "fit_glmm must converge on the fixture");
assert!(
(ws.params[0] - fx.theta).abs() / fx.theta < 0.01,
"our θ̂ {} vs lme4 θ̂ {} ({}% rel)",
ws.params[0],
fx.theta,
100.0 * (ws.params[0] - fx.theta).abs() / fx.theta
);
for j in 0..p {
assert!(
(ws.params[1 + j] - fx.beta[j]).abs() < 5e-3,
"our β̂[{j}] {} vs lme4 β̂[{j}] {} (gap {})",
ws.params[1 + j],
fx.beta[j],
(ws.params[1 + j] - fx.beta[j]).abs()
);
}
let our_theta = ws.params[0];
let mut cov = Mat::<f64>::zeros(p, p);
let status = fd_hessian_cov(&mut ws, xf64.as_ref(), &y, &ids, p, n, &mut cov);
assert_eq!(status, FdHessianStatus::Ok);
assert!((ws.params[0] - our_theta).abs() < 1e-15);
let tol = 2e-3;
for i in 0..p {
for j in 0..p {
let (got, want) = (cov[(i, j)], fx.vcov_hessian[i][j]);
assert!(
(got - want).abs() < tol,
"vcov[{i}][{j}] got {got} want {want} (gap {})",
(got - want).abs()
);
}
}
}
#[test]
fn hessian_mode_t_sq_uses_fd_hessian_cov() {
let fx = load_hessian_fixture();
let n = fx.n;
let p = fx.beta.len();
let n_clusters = fx.cluster_ids.iter().max().unwrap() + 1;
let cluster = intercept_only_spec(Sizing::FixedClusters { n_clusters }, 0.25);
let mut xf64 = Mat::<f64>::zeros(n, p);
for i in 0..n {
for j in 0..p {
xf64[(i, j)] = fx.x[i][j];
}
}
let y = fx.y.clone();
let ids = fx.cluster_ids.clone();
let t1 = 1usize;
let mut ws_h = GlmmWorkspace::for_cluster_spec(p, &cluster, n, &[]);
build_z(&mut ws_h, xf64.as_ref(), &ids, &[], n);
let fit_h = fit_glmm(
&mut ws_h,
xf64.as_ref(),
&y,
&ids,
&[1u32],
None,
&vec![0.0; p],
n,
WaldSe::Hessian,
);
assert!(fit_h.converged, "hessian-mode fit must converge");
let se_h = ws_h.var_diag[t1].sqrt();
let mut ws_rx = GlmmWorkspace::for_cluster_spec(p, &cluster, n, &[]);
build_z(&mut ws_rx, xf64.as_ref(), &ids, &[], n);
let fit_rx = fit_glmm(
&mut ws_rx,
xf64.as_ref(),
&y,
&ids,
&[1u32],
None,
&vec![0.0; p],
n,
WaldSe::Rx,
);
assert!(fit_rx.converged, "rx-mode fit must converge");
let se_rx = ws_rx.var_diag[t1].sqrt();
assert!(se_h > se_rx, "hessian SE {se_h} must exceed rx SE {se_rx}");
let want_var = fx.vcov_hessian[t1][t1];
assert!(
(ws_h.var_diag[t1] - want_var).abs() < 2e-3,
"hessian var {} must match fixture vcov_hessian diag {want_var}",
ws_h.var_diag[t1]
);
}
#[test]
fn fd_hessian_non_pd_falls_back_to_rx_and_counts() {
let (n, nc) = (80usize, 8usize);
let per = n / nc;
let mut st = 4242u64;
let mut xf64 = Mat::<f64>::zeros(n, 2);
let mut y = vec![0.0f64; n];
let mut ids = vec![0u32; n];
for i in 0..n {
let c = i / per; ids[i] = c as u32;
let u_c = 10.0 * (2.0 * (c as f64) / ((nc - 1) as f64) - 1.0); let x1 = lcg(&mut st); xf64[(i, 0)] = 1.0;
xf64[(i, 1)] = x1;
let eta = 0.0 + 0.8 * x1 + u_c;
let pr = 1.0 / (1.0 + (-eta).exp());
y[i] = if lcg(&mut st) + 0.5 < pr { 1.0 } else { 0.0 };
}
let p = 2usize;
let cluster = intercept_only_spec(
Sizing::FixedClusters {
n_clusters: nc as u32,
},
0.25,
);
let mut ws = GlmmWorkspace::for_cluster_spec(p, &cluster, n, &[]);
build_z(&mut ws, xf64.as_ref(), &ids, &[], n);
let fit = fit_glmm(
&mut ws,
xf64.as_ref(),
&y,
&ids,
&[1u32],
None,
&[0.0, 0.8],
n,
WaldSe::Rx,
);
assert!(fit.converged, "fixture fit must converge");
assert!(
ws.params[0] > 5.0,
"fit must reach the high-variance non-PD regime (θ̂ = {})",
ws.params[0]
);
let mut cov = Mat::<f64>::zeros(p, p);
let status = fd_hessian_cov(&mut ws, xf64.as_ref(), &y, &ids, p, n, &mut cov);
assert_eq!(
status,
FdHessianStatus::NonPdFellBackToRx,
"high-variance joint Hessian must be non-PD ⇒ RX fallback"
);
let m = ws.params.len();
for i in 0..m {
for j in 0..m {
assert!(
ws.hess_scratch[(i, j)].is_finite(),
"assembled Hessian must be finite (non-finite branch NOT taken): H[{i}][{j}]"
);
}
}
assert!(
ws.hess_scratch.as_ref().llt(faer::Side::Lower).is_err(),
"assembled joint Hessian must be non-PD (LLT must fail)"
);
let mut rx = Mat::<f64>::zeros(p, p);
assert!(
rx_cov_into(&mut ws, xf64.as_ref(), &ids, p, n, &mut rx),
"β-only Schur must stay PD (well-conditioned β design)"
);
for i in 0..p {
for j in 0..p {
assert!(
(cov[(i, j)] - rx[(i, j)]).abs() < 1e-10,
"cov[{i}][{j}] {} vs rx {}",
cov[(i, j)],
rx[(i, j)]
);
}
}
}
#[test]
fn fit_glmm_collapses_to_plain_irls_when_tau_negligible() {
use crate::glm::{glm_irls_fit, GlmScratch};
let (n, nc) = (200usize, 10usize);
let mut st = 99u64;
let mut x = Mat::<f64>::zeros(n, 2);
let mut y = vec![0.0f64; n];
let mut ids = vec![0u32; n];
for i in 0..n {
ids[i] = (i % nc) as u32;
let x1 = lcg(&mut st);
x[(i, 0)] = 1.0;
x[(i, 1)] = x1;
let eta = 0.1 + 0.7 * x1; let p = 1.0 / (1.0 + (-eta).exp());
let v = if lcg(&mut st) + 0.5 < p {
1.0f64
} else {
0.0f64
};
y[i] = v;
}
let mut sw = TestWs::new(n, 2, 0);
let irls = {
let s = GlmScratch {
irls_eta: &mut sw.irls_eta[..n],
irls_p: &mut sw.irls_p[..n],
irls_w: &mut sw.irls_w[..n],
irls_z: &mut sw.irls_z[..n],
irls_betas: &mut sw.irls_betas[..2],
irls_betas_new: &mut sw.irls_betas_new[..2],
irls_var_diag: &mut sw.irls_var_diag[..1],
irls_t_sq: &mut sw.irls_t_sq[..1],
irls_u_scratch: &mut sw.irls_u_scratch[..2],
irls_xtwx: sw.irls_xtwx.as_mut().submatrix_mut(0, 0, 2, 2),
irls_xtwz: &mut sw.irls_xtwz[..2],
irls_l: sw.irls_l.as_mut().submatrix_mut(0, 0, 2, 2),
irls_wx: &mut sw.irls_wx[..n * 2],
};
let f = glm_irls_fit(x.as_ref(), &y, &[1], None, s);
(f.betas.to_vec(), f.t_sq.to_vec(), f.converged)
};
assert!(irls.2, "plain IRLS must converge");
let cluster = intercept_only_spec(
Sizing::FixedClusters {
n_clusters: nc as u32,
},
0.1,
);
let mut ws = GlmmWorkspace::for_cluster_spec(2, &cluster, n, &[]);
build_z(&mut ws, x.as_ref(), &ids, &[], n);
let fit = fit_glmm(
&mut ws,
x.as_ref(),
&y,
&ids,
&[1],
Some(&[0.05]),
&[0.1, 0.7],
n,
WaldSe::Rx,
);
assert!(fit.converged);
assert!(
(ws.betas[1] - irls.0[1]).abs() < 1e-2,
"β̂₁ glmm {} vs irls {}",
ws.betas[1],
irls.0[1]
);
assert!(
(ws.t_sq[1].sqrt() - irls.1[0].sqrt()).abs() < 5e-2,
"z glmm {} vs irls {}",
ws.t_sq[1].sqrt(),
irls.1[0].sqrt()
);
}
#[test]
fn fit_glmm_width_general_slope_and_crossed() {
let (xf64, y, ids, crossed_ids, cluster) = glmm_slope_crossed_dataset();
let n = y.len();
let mut ws = GlmmWorkspace::for_cluster_spec(2, &cluster, n, &[1]);
build_z(
&mut ws,
xf64.as_ref(),
&ids,
std::slice::from_ref(&crossed_ids),
n,
);
let fit = fit_glmm(
&mut ws,
xf64.as_ref(),
&y,
&ids,
&[1],
None,
&[0.2, 0.8],
n,
WaldSe::Rx,
);
assert!(fit.converged);
assert!(
ws.betas[1] > 0.5,
"slope must be strongly positive, got {}",
ws.betas[1]
);
assert!(
ws.t_sq[1].is_finite() && ws.t_sq[1] > 3.84,
"z² must clear the α=0.05 bar (3.84), got {}",
ws.t_sq[1]
);
assert!(
fit.tau_squared_hat.is_finite() && (0.0..5.0).contains(&fit.tau_squared_hat),
"τ̂² {}",
fit.tau_squared_hat
);
}
#[test]
#[ignore]
fn fit_glmm_warm_path_bounded_alloc() {
let (xf64, y, ids) = glmm_intercept_dataset();
let cluster = intercept_only_spec(Sizing::FixedClusters { n_clusters: 8 }, 0.25);
let mut ws = GlmmWorkspace::for_cluster_spec(2, &cluster, 80, &[]);
build_z(&mut ws, xf64.as_ref(), &ids, &[], 80);
let _ = fit_glmm(
&mut ws,
xf64.as_ref(),
&y,
&ids,
&[1],
Some(&[0.5]),
&[0.2, 0.8],
80,
WaldSe::Rx,
); let profiler = dhat::Profiler::builder().testing().build();
for _ in 0..20 {
let _ = fit_glmm(
&mut ws,
xf64.as_ref(),
&y,
&ids,
&[1],
Some(&[0.5]),
&[0.2, 0.8],
80,
WaldSe::Rx,
);
}
let stats = dhat::HeapStats::get();
drop(profiler);
const BOUND: u64 = 124;
assert!(
stats.total_blocks <= BOUND,
"warm-path alloc regressed: {} blocks across 20 fits (BOUND = {})",
stats.total_blocks,
BOUND
);
}
#[test]
#[ignore]
fn fit_glmm_structured_warm_path_bounded_alloc() {
let (xf64, y, ids, extra_ids, cluster) = glmm_extras_q1_dataset(2, 6);
let n = y.len();
let mut ws = GlmmWorkspace::for_cluster_spec(2, &cluster, n, &[]);
assert!(ws.groupings.structured_extras_eligible());
build_z(&mut ws, xf64.as_ref(), &ids, &extra_ids, n);
let theta = [0.5_f64, 0.4, 0.45];
let _ = fit_glmm(
&mut ws,
xf64.as_ref(),
&y,
&ids,
&[1],
Some(&theta),
&[0.2, 0.8],
n,
WaldSe::Rx,
); let profiler = dhat::Profiler::builder().testing().build();
for _ in 0..20 {
let _ = fit_glmm(
&mut ws,
xf64.as_ref(),
&y,
&ids,
&[1],
Some(&theta),
&[0.2, 0.8],
n,
WaldSe::Rx,
);
}
let stats = dhat::HeapStats::get();
drop(profiler);
const BOUND: u64 = 124;
assert!(
stats.total_blocks <= BOUND,
"structured warm-path alloc regressed: {} blocks across 20 fits (BOUND = {})",
stats.total_blocks,
BOUND
);
}
#[test]
fn block_chol_and_solve_match_faer() {
let q = 3usize;
let a = [4.0, 0.0, 0.0, 2.0, 5.0, 0.0, 1.0, 3.0, 6.0];
let b = [1.0_f64, -2.0, 0.5];
let mut af = Mat::<f64>::zeros(q, q);
for r in 0..q {
for c in 0..=r {
af[(r, c)] = a[r * q + c];
af[(c, r)] = a[r * q + c];
}
}
let ac = af.as_ref().llt(faer::Side::Lower).unwrap();
let mut rhs = Mat::<f64>::zeros(q, 1);
for r in 0..q {
rhs[(r, 0)] = b[r];
}
ac.solve_in_place(rhs.as_mut());
let mut blk = a;
assert!(super::glmm_block_chol(&mut blk, q), "block should be PD");
let mut x = b;
super::glmm_block_solve(&blk, q, &mut x);
for r in 0..q {
assert!(
(x[r] - rhs[(r, 0)]).abs() < 1e-12,
"x[{r}] = {}, faer {}",
x[r],
rhs[(r, 0)]
);
}
let logdet_helper: f64 = (0..q).map(|r| blk[r * q + r].ln()).sum::<f64>() * 2.0;
let logdet_faer: f64 = (0..q).map(|r| ac.L()[(r, r)].ln()).sum::<f64>() * 2.0;
assert!((logdet_helper - logdet_faer).abs() < 1e-12);
}
#[test]
fn block_chol_rejects_non_pd() {
let q = 2usize;
let mut blk = [1.0_f64, 0.0, 2.0, 1.0]; assert!(!super::glmm_block_chol(&mut blk, q));
}
fn glmm_slope_noextra_dataset_layout(contiguous: bool) -> (Mat<f64>, Vec<f64>, Vec<u32>) {
let (n, nc) = (96usize, 8usize);
let mut st = 21u64;
let u0: Vec<f64> = (0..nc).map(|_| 0.6 * lcg(&mut st)).collect();
let u1: Vec<f64> = (0..nc).map(|_| 0.4 * lcg(&mut st)).collect();
let mut x = Mat::<f64>::zeros(n, 2);
let mut y = vec![0.0f64; n];
let mut ids = vec![0u32; n];
for i in 0..n {
let c = if contiguous { i / (n / nc) } else { i % nc };
ids[i] = c as u32;
let x1 = lcg(&mut st);
x[(i, 0)] = 1.0;
x[(i, 1)] = x1;
let eta = 0.2 + 0.8 * x1 + u0[c] + u1[c] * x1;
let p = 1.0 / (1.0 + (-eta).exp());
y[i] = if lcg(&mut st) + 0.5 < p { 1.0 } else { 0.0 };
}
(x, y, ids)
}
fn glmm_slope_noextra_dataset() -> (Mat<f64>, Vec<f64>, Vec<u32>) {
glmm_slope_noextra_dataset_layout(false)
}
#[test]
fn blocked_laplace_matches_brute_force_intercept() {
let (xf64, y, ids) = glmm_intercept_dataset();
let beta = [0.2_f64, 0.8];
let want = brute_force_intercept_laplace(0.5, &beta, &xf64, &y, &ids, 8);
let cluster = intercept_only_spec(Sizing::FixedClusters { n_clusters: 8 }, 0.25);
let mut ws = GlmmWorkspace::for_cluster_spec(2, &cluster, 80, &[]);
assert!(
ws.groupings.extra_offsets.is_empty(),
"fixture must route blocked"
);
build_z(&mut ws, xf64.as_ref(), &ids, &[], 80);
ws.params[0] = 0.5;
ws.params[1] = beta[0];
ws.params[2] = beta[1];
let got = glmm_laplace_deviance(&ws.params.clone(), &mut ws, xf64.as_ref(), &y, &ids, 80);
assert!(
(got - want).abs() < 1e-6,
"blocked laplace: got {got}, want {want}"
);
}
#[test]
fn blocked_laplace_matches_brute_force_intercept_contiguous() {
let (xf64, y, ids) = glmm_intercept_dataset_layout(true);
let beta = [0.2_f64, 0.8];
let want = brute_force_intercept_laplace(0.5, &beta, &xf64, &y, &ids, 8);
let cluster = intercept_only_spec(Sizing::FixedClusters { n_clusters: 8 }, 0.25);
let mut ws = GlmmWorkspace::for_cluster_spec(2, &cluster, 80, &[]);
assert!(
ws.groupings.extra_offsets.is_empty(),
"fixture must route blocked"
);
build_z(&mut ws, xf64.as_ref(), &ids, &[], 80);
ws.params[0] = 0.5;
ws.params[1] = beta[0];
ws.params[2] = beta[1];
let got = glmm_laplace_deviance(&ws.params.clone(), &mut ws, xf64.as_ref(), &y, &ids, 80);
assert!(
(got - want).abs() < 1e-6,
"blocked laplace (contiguous): got {got}, want {want}"
);
}
#[test]
fn blocked_pirls_matches_dense_slope_noextra() {
let (xf64, y, ids) = glmm_slope_noextra_dataset();
let n = y.len();
let cluster = ModelSpec {
sizing: Sizing::FixedClusters { n_clusters: 8 },
tau_squared: 0.25,
slopes: vec![SlopeTerm {
column: 1,
variance: 0.16,
corr_with_intercept: 0.0,
corr_with: vec![],
}],
extra_groupings: vec![],
estimator: Estimator::Glm,
wald_se: WaldSe::Hessian,
};
let mut ws = GlmmWorkspace::for_cluster_spec(2, &cluster, n, &[1]);
assert!(ws.groupings.extra_offsets.is_empty());
build_z(&mut ws, xf64.as_ref(), &ids, &[], n);
let theta = [0.5_f64, 0.1, 0.4]; let beta = [0.2_f64, 0.8];
let (k, p, nt) = (ws.k, ws.p, ws.n_theta);
let mut params = vec![0.0; nt + p];
params[..nt].copy_from_slice(&theta);
params[nt..].copy_from_slice(&beta);
let GlmmWorkspace {
groupings,
z,
m,
lam,
eta,
prob,
w,
u,
eta_fixed,
mu,
wm,
a,
a_rhs,
..
} = &mut ws;
apply_lambda(groupings, ¶ms, z.as_ref(), m, lam, n);
let dense = pirls_solve(
k,
p,
m.as_ref(),
xf64.as_ref(),
&y,
¶ms[nt..],
eta,
prob,
w,
u,
eta_fixed,
mu,
wm,
a,
a_rhs,
n,
);
let mut ws2 = GlmmWorkspace::for_cluster_spec(2, &cluster, n, &[1]);
build_z(&mut ws2, xf64.as_ref(), &ids, &[], n);
crate::lmm::primary_lambda(&theta, ws2.groupings.primary_q, &mut ws2.lam);
fill_z_f64(&ws2.groupings, xf64.as_ref(), &mut ws2.z_buf, n);
let GlmmWorkspace {
groupings,
lam,
z_buf,
m_buf,
eta,
prob,
w,
u,
eta_fixed,
a_blocks,
a_rhs,
..
} = &mut ws2;
let blocked = pirls_solve_blocked(
groupings,
&ids,
xf64.as_ref(),
&y,
&beta,
lam,
z_buf,
m_buf,
eta,
prob,
w,
u,
eta_fixed,
a_blocks,
a_rhs,
n,
);
assert_eq!(dense.3, blocked.3, "convergence flag");
assert!(
(dense.0 - blocked.0).abs() < 1e-9,
"dev: dense {} blocked {}",
dense.0,
blocked.0
);
assert!(
(dense.1 - blocked.1).abs() < 1e-9,
"pen: dense {} blocked {}",
dense.1,
blocked.1
);
assert!(
(dense.2 - blocked.2).abs() < 1e-9,
"logdet: dense {} blocked {}",
dense.2,
blocked.2
);
for c in 0..k {
assert!(
(ws.u[c] - ws2.u[c]).abs() < 1e-7,
"u[{c}]: dense {} blocked {}",
ws.u[c],
ws2.u[c]
);
}
}
#[test]
fn blocked_pirls_matches_dense_slope_contiguous() {
let (xf64, y, ids) = glmm_slope_noextra_dataset_layout(true);
let n = y.len();
let cluster = ModelSpec {
sizing: Sizing::FixedClusters { n_clusters: 8 },
tau_squared: 0.25,
slopes: vec![SlopeTerm {
column: 1,
variance: 0.16,
corr_with_intercept: 0.0,
corr_with: vec![],
}],
extra_groupings: vec![],
estimator: Estimator::Glm,
wald_se: WaldSe::Hessian,
};
let mut ws = GlmmWorkspace::for_cluster_spec(2, &cluster, n, &[1]);
assert!(ws.groupings.extra_offsets.is_empty());
build_z(&mut ws, xf64.as_ref(), &ids, &[], n);
let theta = [0.5_f64, 0.1, 0.4]; let beta = [0.2_f64, 0.8];
let (k, p, nt) = (ws.k, ws.p, ws.n_theta);
let mut params = vec![0.0; nt + p];
params[..nt].copy_from_slice(&theta);
params[nt..].copy_from_slice(&beta);
let GlmmWorkspace {
groupings,
z,
m,
lam,
eta,
prob,
w,
u,
eta_fixed,
mu,
wm,
a,
a_rhs,
..
} = &mut ws;
apply_lambda(groupings, ¶ms, z.as_ref(), m, lam, n);
let dense = pirls_solve(
k,
p,
m.as_ref(),
xf64.as_ref(),
&y,
¶ms[nt..],
eta,
prob,
w,
u,
eta_fixed,
mu,
wm,
a,
a_rhs,
n,
);
let mut ws2 = GlmmWorkspace::for_cluster_spec(2, &cluster, n, &[1]);
build_z(&mut ws2, xf64.as_ref(), &ids, &[], n);
crate::lmm::primary_lambda(&theta, ws2.groupings.primary_q, &mut ws2.lam);
fill_z_f64(&ws2.groupings, xf64.as_ref(), &mut ws2.z_buf, n);
let GlmmWorkspace {
groupings,
lam,
z_buf,
m_buf,
eta,
prob,
w,
u,
eta_fixed,
a_blocks,
a_rhs,
..
} = &mut ws2;
let blocked = pirls_solve_blocked(
groupings,
&ids,
xf64.as_ref(),
&y,
&beta,
lam,
z_buf,
m_buf,
eta,
prob,
w,
u,
eta_fixed,
a_blocks,
a_rhs,
n,
);
assert_eq!(dense.3, blocked.3, "convergence flag");
assert!(
(dense.0 - blocked.0).abs() < 1e-9,
"dev: dense {} blocked {}",
dense.0,
blocked.0
);
assert!(
(dense.1 - blocked.1).abs() < 1e-9,
"pen: dense {} blocked {}",
dense.1,
blocked.1
);
assert!(
(dense.2 - blocked.2).abs() < 1e-9,
"logdet: dense {} blocked {}",
dense.2,
blocked.2
);
for c in 0..k {
assert!(
(ws.u[c] - ws2.u[c]).abs() < 1e-7,
"u[{c}]: dense {} blocked {}",
ws.u[c],
ws2.u[c]
);
}
}
#[test]
fn blocked_inference_matches_dense_slope_noextra() {
let (xf64, y, ids) = glmm_slope_noextra_dataset();
let n = y.len();
let cluster = ModelSpec {
sizing: Sizing::FixedClusters { n_clusters: 8 },
tau_squared: 0.25,
slopes: vec![SlopeTerm {
column: 1,
variance: 0.16,
corr_with_intercept: 0.0,
corr_with: vec![],
}],
extra_groupings: vec![],
estimator: Estimator::Glm,
wald_se: WaldSe::Hessian,
};
let mut ws = GlmmWorkspace::for_cluster_spec(2, &cluster, n, &[1]);
build_z(&mut ws, xf64.as_ref(), &ids, &[], n);
let fit = fit_glmm(
&mut ws,
xf64.as_ref(),
&y,
&ids,
&[1],
Some(&[0.5, 0.1, 0.4]),
&[0.2, 0.8],
n,
WaldSe::Rx,
);
assert!(fit.converged);
let (k, p, nt) = (ws.k, ws.p, ws.n_theta);
let beta_blocked: Vec<f64> = ws.betas[..p].to_vec();
let var_blocked = ws.var_diag[1];
let tsq_blocked = ws.t_sq[1];
crate::lmm::primary_lambda(&ws.params[..nt], ws.groupings.primary_q, &mut ws.lam);
{
let GlmmWorkspace {
groupings,
params,
z,
m,
lam,
..
} = &mut ws;
apply_lambda(groupings, ¶ms[..], z.as_ref(), m, lam, n);
}
let mut xtwx = Mat::<f64>::zeros(p, p);
for r in 0..p {
for c in 0..p {
let mut sm = 0.0;
for i in 0..n {
sm += xf64[(i, r)] * ws.w[i] * xf64[(i, c)];
}
xtwx[(r, c)] = sm;
}
}
let mut xtwm = Mat::<f64>::zeros(p, k);
for r in 0..p {
for c in 0..k {
let mut sm = 0.0;
for i in 0..n {
sm += xf64[(i, r)] * ws.w[i] * ws.m[(i, c)];
}
xtwm[(r, c)] = sm;
}
}
let mut a = Mat::<f64>::zeros(k, k);
for r in 0..k {
for c in 0..k {
let mut sm = if r == c { 1.0 } else { 0.0 };
for i in 0..n {
sm += ws.m[(i, r)] * ws.w[i] * ws.m[(i, c)];
}
a[(r, c)] = sm;
}
}
let ac = a.as_ref().llt(faer::Side::Lower).unwrap();
let mut ainv = Mat::<f64>::zeros(k, p);
for r in 0..k {
for c in 0..p {
ainv[(r, c)] = xtwm[(c, r)];
}
}
ac.solve_in_place(ainv.as_mut());
let mut schur = Mat::<f64>::zeros(p, p);
for r in 0..p {
for c in 0..p {
let mut sm = xtwx[(r, c)];
for j in 0..k {
sm -= xtwm[(r, j)] * ainv[(j, c)];
}
schur[(r, c)] = sm;
}
}
let sc = schur.as_ref().llt(faer::Side::Lower).unwrap();
let mut fwd = vec![0.0; p];
for i in 0..p {
let mut acc = if i == 1 { 1.0 } else { 0.0 };
#[allow(clippy::needless_range_loop)]
for kk in 0..i {
acc -= sc.L()[(i, kk)] * fwd[kk];
}
fwd[i] = acc / sc.L()[(i, i)];
}
let var_dense: f64 = fwd.iter().map(|v| v * v).sum();
assert!(
(var_blocked - var_dense).abs() < 1e-8,
"var: blocked {var_blocked} dense {var_dense}"
);
let tsq_dense = beta_blocked[1] * beta_blocked[1] / var_dense;
assert!(
(tsq_blocked - tsq_dense).abs() < 1e-6,
"z²: blocked {tsq_blocked} dense {tsq_dense}"
);
}
#[test]
fn warm_start_is_per_fit_deterministic() {
let (xf64, y, ids) = glmm_intercept_dataset();
let cluster = intercept_only_spec(Sizing::FixedClusters { n_clusters: 8 }, 0.25);
let mk = || {
let mut ws = GlmmWorkspace::for_cluster_spec(2, &cluster, 80, &[]);
build_z(&mut ws, xf64.as_ref(), &ids, &[], 80);
ws
};
let mut ws_ref = mk();
let _ = fit_glmm(
&mut ws_ref,
xf64.as_ref(),
&y,
&ids,
&[1],
Some(&[0.5]),
&[0.2, 0.8],
80,
WaldSe::Rx,
);
let ref_beta = ws_ref.betas[1].to_bits();
let mut ws = mk();
let _ = fit_glmm(
&mut ws,
xf64.as_ref(),
&y,
&ids,
&[1],
Some(&[0.5]),
&[0.2, 0.8],
80,
WaldSe::Rx,
);
let _ = fit_glmm(
&mut ws,
xf64.as_ref(),
&y,
&ids,
&[1],
Some(&[0.5]),
&[0.2, 0.8],
80,
WaldSe::Rx,
);
assert_eq!(
ws.betas[1].to_bits(),
ref_beta,
"re-fit β̂ must match the fresh cold fit bit-for-bit"
);
}
#[test]
fn warm_start_objective_is_seed_independent() {
let (xf64, y, ids) = glmm_intercept_dataset();
let cluster = intercept_only_spec(Sizing::FixedClusters { n_clusters: 8 }, 0.25);
let mut ws = GlmmWorkspace::for_cluster_spec(2, &cluster, 80, &[]);
build_z(&mut ws, xf64.as_ref(), &ids, &[], 80);
ws.params[0] = 0.5;
ws.params[1] = 0.2;
ws.params[2] = 0.8;
for v in ws.u.iter_mut() {
*v = 0.0;
}
let cold = glmm_laplace_deviance(&ws.params.clone(), &mut ws, xf64.as_ref(), &y, &ids, 80);
for (c, v) in ws.u.iter_mut().enumerate() {
*v = 0.05 * (c as f64 - 4.0);
}
let warm = glmm_laplace_deviance(&ws.params.clone(), &mut ws, xf64.as_ref(), &y, &ids, 80);
assert!(
(cold - warm).abs() < 1e-6,
"objective seed-dependent: cold {cold} warm {warm}"
);
}
#[allow(clippy::type_complexity)]
fn glmm_extras_q1_dataset(
np: usize,
n_crossed: usize,
) -> (Mat<f64>, Vec<f64>, Vec<u32>, Vec<Vec<u32>>, ModelSpec) {
let (n, n_prim) = (96usize, 8usize);
let mut st = 29u64;
let u0: Vec<f64> = (0..n_prim).map(|_| 0.6 * lcg(&mut st)).collect();
let un: Vec<f64> = (0..n_prim * np.max(1))
.map(|_| 0.4 * lcg(&mut st))
.collect();
let uc: Vec<f64> = (0..n_crossed.max(1)).map(|_| 0.5 * lcg(&mut st)).collect();
let mut x = Mat::<f64>::zeros(n, 2);
let mut y = vec![0.0f64; n];
let mut ids = vec![0u32; n];
let mut nested = vec![0u32; n];
let mut crossed = vec![0u32; n];
for i in 0..n {
let c = i % n_prim;
ids[i] = c as u32;
let x1 = lcg(&mut st);
x[(i, 0)] = 1.0;
x[(i, 1)] = x1;
let mut eta = 0.2 + 0.8 * x1 + u0[c];
if np > 0 {
let within = (i / n_prim) % np;
let gid = c * np + within;
nested[i] = gid as u32;
eta += un[gid];
}
if n_crossed > 0 {
let cc = i % n_crossed;
crossed[i] = cc as u32;
eta += uc[cc];
}
let p = 1.0 / (1.0 + (-eta).exp());
y[i] = if lcg(&mut st) + 0.5 < p { 1.0 } else { 0.0 };
}
let mut extra_groupings = Vec::new();
let mut extra_ids = Vec::new();
if np > 0 {
extra_groupings.push(Grouping {
relation: GroupingRelation::NestedWithin {
n_per_parent: np as u32,
},
tau_squared: 0.16,
slopes: vec![],
});
extra_ids.push(nested);
}
if n_crossed > 0 {
extra_groupings.push(Grouping {
relation: GroupingRelation::Crossed {
n_clusters: n_crossed as u32,
},
tau_squared: 0.25,
slopes: vec![],
});
extra_ids.push(crossed);
}
let cluster = ModelSpec {
sizing: Sizing::FixedClusters {
n_clusters: n_prim as u32,
},
tau_squared: 0.25,
slopes: vec![],
extra_groupings,
estimator: Estimator::Glm,
wald_se: WaldSe::Hessian,
};
(x, y, ids, extra_ids, cluster)
}
#[allow(clippy::too_many_arguments)]
fn brute_force_extras_laplace(
theta_p: f64,
theta_n: f64,
theta_c: f64,
beta: &[f64],
x: &Mat<f64>,
y: &[f64],
ids: &[u32],
nested: &[u32],
crossed: &[u32],
n_prim: usize,
np: usize,
n_crossed: usize,
) -> f64 {
let (n, p) = (x.nrows(), x.ncols());
let nc = n_prim + n_prim * np + n_crossed;
let mut m = Mat::<f64>::zeros(n, nc);
let nest_base = n_prim;
let cross_base = n_prim + n_prim * np;
for i in 0..n {
m[(i, ids[i] as usize)] = theta_p;
if np > 0 {
m[(i, nest_base + nested[i] as usize)] = theta_n;
}
if n_crossed > 0 {
m[(i, cross_base + crossed[i] as usize)] = theta_c;
}
}
let eta_of = |u: &[f64], i: usize| -> f64 {
let mut e = 0.0;
for j in 0..p {
e += x[(i, j)] * beta[j];
}
for c in 0..nc {
e += m[(i, c)] * u[c];
}
e
};
let mut u = vec![0.0f64; nc];
for _ in 0..80 {
let mut pvec = vec![0.0; n];
let mut w = vec![0.0; n];
for i in 0..n {
let e = eta_of(&u, i);
let pi = 1.0 / (1.0 + (-e).exp());
pvec[i] = pi;
w[i] = (pi * (1.0 - pi)).max(1e-6);
}
let mut g = vec![0.0; nc];
for c in 0..nc {
let mut s = 0.0;
for i in 0..n {
s += m[(i, c)] * (y[i] - pvec[i]);
}
g[c] = 2.0 * u[c] - 2.0 * s;
}
let mut h = Mat::<f64>::zeros(nc, nc);
for a in 0..nc {
for b in 0..nc {
let mut s = 0.0;
for i in 0..n {
s += m[(i, a)] * w[i] * m[(i, b)];
}
h[(a, b)] = 2.0 * (s + if a == b { 1.0 } else { 0.0 });
}
}
let hc = h.as_ref().llt(faer::Side::Lower).unwrap();
let mut step = Mat::<f64>::zeros(nc, 1);
for c in 0..nc {
step[(c, 0)] = g[c];
}
hc.solve_in_place(step.as_mut());
let mut max = 0.0f64;
for c in 0..nc {
u[c] -= step[(c, 0)];
max = max.max(step[(c, 0)].abs());
}
if max < 1e-11 {
break;
}
}
let mut d = 0.0;
let mut w = vec![0.0; n];
for i in 0..n {
let e = eta_of(&u, i);
let pi = 1.0 / (1.0 + (-e).exp());
w[i] = (pi * (1.0 - pi)).max(1e-6);
d += if e > 0.0 {
e + (-e).exp().ln_1p()
} else {
e.exp().ln_1p()
} - y[i] * e;
}
let pen: f64 = u.iter().map(|v| v * v).sum();
let mut a = Mat::<f64>::zeros(nc, nc);
for r in 0..nc {
for c in 0..nc {
let mut s = 0.0;
for i in 0..n {
s += m[(i, r)] * w[i] * m[(i, c)];
}
a[(r, c)] = s + if r == c { 1.0 } else { 0.0 };
}
}
let ac = a.as_ref().llt(faer::Side::Lower).unwrap();
let mut logdet = 0.0;
for r in 0..nc {
logdet += ac.L()[(r, r)].ln();
}
2.0 * d + pen + 2.0 * logdet
}
#[test]
fn structured_extras_laplace_matches_brute_force() {
for (np, ncr, label) in [
(0usize, 6usize, "crossed"),
(2, 0, "nested"),
(2, 6, "crossed_nested"),
] {
let (xf64, y, ids, extra_ids, cluster) = glmm_extras_q1_dataset(np, ncr);
let n = y.len();
let mut ws = GlmmWorkspace::for_cluster_spec(2, &cluster, n, &[]);
assert!(
!ws.groupings.extra_offsets.is_empty(),
"{label}: fixture must route through the extras path"
);
let extra_refs: Vec<&[u32]> = extra_ids.iter().map(|v| v.as_slice()).collect();
build_z(&mut ws, xf64.as_ref(), &ids, &extra_ids, n);
let theta_p = 0.5;
let theta_n = 0.4;
let theta_c = 0.45;
let beta = [0.2_f64, 0.8];
let nt = ws.n_theta;
ws.params[0] = theta_p;
let mut ti = 1;
if np > 0 {
ws.params[ti] = theta_n;
ti += 1;
}
if ncr > 0 {
ws.params[ti] = theta_c;
}
ws.params[nt] = beta[0];
ws.params[nt + 1] = beta[1];
let got =
glmm_laplace_deviance(&ws.params.clone(), &mut ws, xf64.as_ref(), &y, &ids, n);
let nested = if np > 0 { extra_refs[0] } else { &[][..] };
let crossed = if ncr > 0 {
extra_refs[extra_refs.len() - 1]
} else {
&[][..]
};
let want = brute_force_extras_laplace(
theta_p, theta_n, theta_c, &beta, &xf64, &y, &ids, nested, crossed, 8, np, ncr,
);
assert!(
(got - want).abs() < 1e-6,
"{label} dense laplace: got {got}, want {want}"
);
}
}
#[test]
fn structured_extras_matches_dense() {
for (np, ncr, label) in [
(0usize, 6usize, "crossed"),
(2, 0, "nested"),
(2, 6, "crossed_nested"),
] {
let (xf64, y, ids, extra_ids, cluster) = glmm_extras_q1_dataset(np, ncr);
let n = y.len();
let mut theta = vec![0.5_f64];
if np > 0 {
theta.push(0.4);
}
if ncr > 0 {
theta.push(0.45);
}
let beta = [0.2_f64, 0.8];
let mut ws = GlmmWorkspace::for_cluster_spec(2, &cluster, n, &[]);
assert!(
ws.groupings.structured_extras_eligible(),
"{label}: fixture must be structured-eligible"
);
build_z(&mut ws, xf64.as_ref(), &ids, &extra_ids, n);
let (k, p, nt) = (ws.k, ws.p, ws.n_theta);
let mut params = vec![0.0; nt + p];
params[..nt].copy_from_slice(&theta);
params[nt..].copy_from_slice(&beta);
let GlmmWorkspace {
groupings,
z,
m,
lam,
eta,
prob,
w,
u,
eta_fixed,
mu,
wm,
a,
a_rhs,
..
} = &mut ws;
apply_lambda(groupings, ¶ms, z.as_ref(), m, lam, n);
let dense = pirls_solve(
k,
p,
m.as_ref(),
xf64.as_ref(),
&y,
¶ms[nt..],
eta,
prob,
w,
u,
eta_fixed,
mu,
wm,
a,
a_rhs,
n,
);
let mut ws2 = GlmmWorkspace::for_cluster_spec(2, &cluster, n, &[]);
build_z(&mut ws2, xf64.as_ref(), &ids, &extra_ids, n);
{
let GlmmWorkspace {
groupings,
params: prm,
z,
lam,
m_core_buf,
cross_val,
cross_col,
n_cross,
..
} = &mut ws2;
prm[..nt].copy_from_slice(&theta);
prm[nt..nt + p].copy_from_slice(&beta);
build_packed_m(
groupings,
&prm[..],
z.as_ref(),
lam,
&ids,
m_core_buf,
cross_val,
cross_col,
n_cross,
n,
);
}
let structured = {
let GlmmWorkspace {
groupings,
m_core_buf,
cross_val,
cross_col,
n_cross,
eta,
prob,
w,
u,
eta_fixed,
mu,
core_blocks,
coupling,
schur_blk,
a_rhs,
..
} = &mut ws2;
pirls_solve_blocked_extras(
groupings,
&ids,
m_core_buf,
cross_val,
cross_col,
n_cross,
xf64.as_ref(),
&y,
&beta,
eta,
prob,
w,
u,
eta_fixed,
mu,
core_blocks,
coupling,
schur_blk,
a_rhs,
n,
)
};
assert_eq!(dense.3, structured.3, "{label}: convergence flag");
assert!(
(dense.0 - structured.0).abs() < 1e-9,
"{label} dev: dense {} structured {}",
dense.0,
structured.0
);
assert!(
(dense.1 - structured.1).abs() < 1e-9,
"{label} pen: dense {} structured {}",
dense.1,
structured.1
);
assert!(
(dense.2 - structured.2).abs() < 1e-9,
"{label} logdet: dense {} structured {}",
dense.2,
structured.2
);
for c in 0..k {
assert!(
(ws.u[c] - ws2.u[c]).abs() < 1e-7,
"{label} u[{c}]: dense {} structured {}",
ws.u[c],
ws2.u[c]
);
}
}
}
#[test]
fn structured_inference_matches_dense() {
for (np, ncr, label) in [
(0usize, 6usize, "crossed"),
(2, 0, "nested"),
(2, 6, "crossed_nested"),
] {
let (xf64, y, ids, extra_ids, cluster) = glmm_extras_q1_dataset(np, ncr);
let n = y.len();
let mut ws = GlmmWorkspace::for_cluster_spec(2, &cluster, n, &[]);
assert!(
ws.groupings.structured_extras_eligible(),
"{label}: eligible"
);
build_z(&mut ws, xf64.as_ref(), &ids, &extra_ids, n);
let fit = fit_glmm(
&mut ws,
xf64.as_ref(),
&y,
&ids,
&[1],
None,
&[0.2, 0.8],
n,
WaldSe::Rx,
);
assert!(fit.converged, "{label}: fit must converge");
let (k, p, nt) = (ws.k, ws.p, ws.n_theta);
let var_structured = ws.var_diag[1];
let tsq_structured = ws.t_sq[1];
let beta1 = ws.betas[1];
crate::lmm::primary_lambda(&ws.params[..nt], ws.groupings.primary_q, &mut ws.lam);
{
let GlmmWorkspace {
groupings,
params,
z,
m,
lam,
..
} = &mut ws;
apply_lambda(groupings, ¶ms[..], z.as_ref(), m, lam, n);
}
let mut xtwx = Mat::<f64>::zeros(p, p);
for r in 0..p {
for c in 0..p {
let mut sm = 0.0;
for i in 0..n {
sm += xf64[(i, r)] * ws.w[i] * xf64[(i, c)];
}
xtwx[(r, c)] = sm;
}
}
let mut xtwm = Mat::<f64>::zeros(p, k);
for r in 0..p {
for c in 0..k {
let mut sm = 0.0;
for i in 0..n {
sm += xf64[(i, r)] * ws.w[i] * ws.m[(i, c)];
}
xtwm[(r, c)] = sm;
}
}
let mut a = Mat::<f64>::zeros(k, k);
for r in 0..k {
for c in 0..k {
let mut sm = if r == c { 1.0 } else { 0.0 };
for i in 0..n {
sm += ws.m[(i, r)] * ws.w[i] * ws.m[(i, c)];
}
a[(r, c)] = sm;
}
}
let ac = a.as_ref().llt(faer::Side::Lower).unwrap();
let mut ainv = Mat::<f64>::zeros(k, p);
for r in 0..k {
for c in 0..p {
ainv[(r, c)] = xtwm[(c, r)];
}
}
ac.solve_in_place(ainv.as_mut());
let mut schur = Mat::<f64>::zeros(p, p);
for r in 0..p {
for c in 0..p {
let mut sm = xtwx[(r, c)];
for j in 0..k {
sm -= xtwm[(r, j)] * ainv[(j, c)];
}
schur[(r, c)] = sm;
}
}
let sc = schur.as_ref().llt(faer::Side::Lower).unwrap();
let mut fwd = vec![0.0; p];
for i in 0..p {
let mut acc = if i == 1 { 1.0 } else { 0.0 };
#[allow(clippy::needless_range_loop)]
for kk in 0..i {
acc -= sc.L()[(i, kk)] * fwd[kk];
}
fwd[i] = acc / sc.L()[(i, i)];
}
let var_dense: f64 = fwd.iter().map(|v| v * v).sum();
assert!(
(var_structured - var_dense).abs() < 1e-8,
"{label} var: structured {var_structured} dense {var_dense}"
);
let tsq_dense = beta1 * beta1 / var_dense;
assert!(
(tsq_structured - tsq_dense).abs() < 1e-6,
"{label} z²: structured {tsq_structured} dense {tsq_dense}"
);
}
}
}