use faer::linalg::matmul::triangular::BlockStructure;
use faer::linalg::matmul::{matmul, triangular};
use faer::reborrow::{IntoConst, Reborrow, ReborrowMut};
use faer::{Accum, MatMut, MatRef, Par};
use crate::ols::{chol_rank_deficient, PANEL_ROWS};
use crate::FLOAT_NEAR_ZERO;
const MAX_BRENT_ITERS: u32 = 50;
const BRENT_REL_TOL: f64 = 1e-4;
const LOG_THETA_LOW: f64 = -9.210_340_371_976_184; #[expect(
clippy::approx_constant,
reason = "literal is ln(1e-1) = -ln(10), not a use of the std LN_10 constant"
)]
const LOG_THETA_MID: f64 = -2.302_585_092_994_046; const LOG_THETA_HIGH: f64 = 6.907_755_278_982_137; const BOUNDARY_LOG_SLACK: f64 = 0.1;
const TRUTH_BRACKET_HALF_WIDTH: f64 = 2.0;
pub struct LmeSuffStats<'w> {
pub xtx: MatMut<'w, f64>,
pub xty: &'w mut [f64],
pub yty: &'w mut f64,
pub sum_xc: MatMut<'w, f64>,
pub sum_yc: &'w mut [f64],
pub cluster_sizes: &'w mut [u32],
pub n_clusters_seen: &'w mut u32,
pub panel_x: &'w mut [f64],
pub panel_y: &'w mut [f64],
}
pub struct LmeScratch<'w> {
pub xtx: MatRef<'w, f64>,
pub xty: &'w [f64],
pub yty: f64,
pub sum_xc: MatRef<'w, f64>,
pub sum_yc: &'w [f64],
pub cluster_sizes: &'w [u32],
pub n_clusters: u32,
pub n_rows: u32,
pub xtvix: MatMut<'w, f64>,
pub xtviy: &'w mut [f64],
pub xtvix_factor: MatMut<'w, f64>,
pub v_diag_inv: &'w mut [f64],
pub betas: &'w mut [f64],
pub var_diag: &'w mut [f64],
pub t_sq: &'w mut [f64],
pub u_scratch: &'w mut [f64],
pub sigma_sq: f64,
pub brent_log_a: &'w mut f64,
pub brent_log_b: &'w mut f64,
pub brent_log_c: &'w mut f64,
pub brent_fa: &'w mut f64,
pub brent_fb: &'w mut f64,
pub brent_fc: &'w mut f64,
pub joint_sigma_t_chol: MatMut<'w, f64>,
pub joint_rhs: &'w mut [f64],
pub joint_k_inv: MatMut<'w, f64>,
}
pub struct LmeFitView<'a> {
pub betas: &'a [f64],
pub var_diag: &'a [f64],
pub t_sq: &'a [f64],
pub factor: MatRef<'a, f64>,
pub sigma_sq: f64,
pub tau_sq_hat: f64,
pub converged: bool,
pub boundary_hit: u8,
pub n_iter: u32,
pub n_evals: u32,
pub joint_t_sq: f64,
}
impl<'w> LmeSuffStats<'w> {
pub fn add_rows(
&mut self,
x_block: MatRef<'_, f64>,
y_block: &[f64],
cluster_ids_block: &[u32],
) {
debug_assert_eq!(x_block.nrows(), y_block.len());
debug_assert_eq!(x_block.nrows(), cluster_ids_block.len());
let p = self.xty.len();
debug_assert_eq!(x_block.ncols(), p);
let max_k = self.sum_yc.len();
let m = x_block.nrows();
debug_assert!(self.panel_x.len() >= PANEL_ROWS.min(m) * p);
debug_assert!(self.panel_y.len() >= PANEL_ROWS.min(m));
let mut off = 0;
while off < m {
let rows = (m - off).min(PANEL_ROWS);
for j in 0..p {
let col = &mut self.panel_x[j * rows..(j + 1) * rows];
for (i, v) in col.iter_mut().enumerate() {
*v = x_block[(off + i, j)];
}
}
for i in 0..rows {
let y_row = y_block[off + i];
let c = cluster_ids_block[off + i] as usize;
debug_assert!(
c < max_k,
"cluster_id {c} exceeds workspace's max_n_clusters {max_k}"
);
self.panel_y[i] = y_row;
self.sum_yc[c] += y_row;
self.cluster_sizes[c] += 1;
*self.yty += y_row * y_row;
let candidate = (c as u32).saturating_add(1);
if candidate > *self.n_clusters_seen {
*self.n_clusters_seen = candidate;
}
}
for j in 0..p {
let col = &self.panel_x[j * rows..(j + 1) * rows];
for (i, &v) in col.iter().enumerate() {
let c = cluster_ids_block[off + i] as usize;
self.sum_xc[(j, c)] += v;
}
}
let xp = MatRef::from_column_major_slice(&self.panel_x[..rows * p], rows, p);
triangular::matmul(
self.xtx.rb_mut(),
BlockStructure::TriangularLower,
Accum::Add,
xp.transpose(),
BlockStructure::Rectangular,
xp,
BlockStructure::Rectangular,
1.0,
Par::Seq,
);
matmul(
MatMut::from_column_major_slice_mut(&mut *self.xty, p, 1),
Accum::Add,
xp.transpose(),
MatRef::from_column_major_slice(&self.panel_y[..rows], rows, 1),
1.0,
Par::Seq,
);
off += rows;
}
}
}
pub fn profiled_deviance(theta: f64, scratch: &mut LmeScratch<'_>) -> f64 {
use faer::reborrow::Reborrow;
let p = scratch.betas.len();
let n_clusters = scratch.n_clusters as usize;
let n_rows = scratch.n_rows as usize;
let theta_sq = theta * theta;
if n_rows <= p || p == 0 {
return f64::INFINITY;
}
let mut log_det_v = 0.0_f64;
for c in 0..n_clusters {
let n_c = scratch.cluster_sizes[c] as f64;
let one_plus = 1.0 + theta_sq * n_c;
scratch.v_diag_inv[c] = 1.0 / one_plus;
log_det_v += one_plus.ln();
}
for j in 0..p {
for i in 0..p {
if i >= j {
scratch.xtvix[(i, j)] = scratch.xtx[(i, j)];
} else {
scratch.xtvix[(i, j)] = 0.0;
}
}
}
for c in 0..n_clusters {
let s_c = theta_sq * scratch.v_diag_inv[c];
for j in 0..p {
let s_xj = scratch.sum_xc[(j, c)];
for i in j..p {
let s_xi = scratch.sum_xc[(i, c)];
scratch.xtvix[(i, j)] -= s_c * s_xi * s_xj;
}
}
}
for j in 0..p {
scratch.xtviy[j] = scratch.xty[j];
}
for c in 0..n_clusters {
let s_c = theta_sq * scratch.v_diag_inv[c];
let s_yc = scratch.sum_yc[c];
let factor = s_c * s_yc;
for j in 0..p {
scratch.xtviy[j] -= factor * scratch.sum_xc[(j, c)];
}
}
let chol = match scratch.xtvix.rb().llt(faer::Side::Lower) {
Ok(c) => c,
Err(_) => return f64::INFINITY,
};
let l = chol.L();
for j in 0..p {
for i in 0..p {
scratch.xtvix_factor[(i, j)] = if i >= j { l[(i, j)] } else { 0.0 };
}
}
use faer::linalg::solvers::Solve;
scratch.betas[..p].copy_from_slice(&scratch.xtviy[..p]);
{
let mut rhs = MatMut::from_column_major_slice_mut(scratch.betas, p, 1usize);
chol.solve_in_place(rhs.rb_mut());
}
let mut ytviy = scratch.yty;
for c in 0..n_clusters {
let s_c = theta_sq * scratch.v_diag_inv[c];
let s_yc = scratch.sum_yc[c];
ytviy -= s_c * s_yc * s_yc;
}
let mut bty = 0.0_f64;
for j in 0..p {
bty += scratch.betas[j] * scratch.xtviy[j];
}
let r_sq = ytviy - bty;
if !r_sq.is_finite() || r_sq <= 0.0 {
return f64::INFINITY;
}
let df_resid = (n_rows - p) as f64;
let sigma_sq = r_sq / df_resid;
if !sigma_sq.is_finite() || sigma_sq <= 0.0 {
return f64::INFINITY;
}
scratch.sigma_sq = sigma_sq;
let mut log_det_xtvix = 0.0_f64;
for j in 0..p {
let ljj = scratch.xtvix_factor[(j, j)];
if !ljj.is_finite() || ljj <= 0.0 {
return f64::INFINITY;
}
log_det_xtvix += ljj.ln();
}
log_det_xtvix *= 2.0;
log_det_v + log_det_xtvix + df_resid * sigma_sq.ln()
}
const CGOLD: f64 = 0.381_966_011_250_105;
const ZEPS: f64 = 1.0e-10;
#[expect(
clippy::too_many_arguments,
reason = "Brent minimizer; args are the bracket triplet plus tuning knobs"
)]
pub fn brent_minimize<F: FnMut(f64) -> f64>(
a: &mut f64,
b: &mut f64,
c: &mut f64,
fa: &mut f64,
fb: &mut f64,
fc: &mut f64,
d: &mut f64,
e: &mut f64,
mut f: F,
rel_tol: f64,
max_iter: u32,
n_iter_out: &mut u32,
) -> (f64, f64) {
let (mut bracket_lo, mut bracket_hi) = if *a < *c { (*a, *c) } else { (*c, *a) };
let mut x = *b;
let mut fx = *fb;
let mut w = x;
let mut fw = fx;
let mut v = x;
let mut fv = fx;
*e = 0.0;
*d = 0.0;
for iter in 0..max_iter {
let xm = 0.5 * (bracket_lo + bracket_hi);
let tol1 = rel_tol * x.abs() + ZEPS;
let tol2 = 2.0 * tol1;
if (x - xm).abs() <= tol2 - 0.5 * (bracket_hi - bracket_lo) {
*n_iter_out = iter;
*a = bracket_lo;
*c = bracket_hi;
*b = x;
*fa = fv; *fb = fx;
*fc = fw; return (x, fx);
}
let u;
if (*e).abs() > tol1 {
let r = (x - w) * (fx - fv);
let q = (x - v) * (fx - fw);
let mut p_num = (x - v) * q - (x - w) * r;
let mut q_denom = 2.0 * (q - r);
if q_denom > 0.0 {
p_num = -p_num;
} else {
q_denom = -q_denom;
}
let p_step = p_num;
let q_step = q_denom;
if p_step.abs() < (0.5 * q_step * (*e).abs())
&& p_step > q_step * (bracket_lo - x)
&& p_step < q_step * (bracket_hi - x)
{
*e = *d;
*d = p_step / q_step;
let u_trial = x + *d;
if (u_trial - bracket_lo) < tol2 || (bracket_hi - u_trial) < tol2 {
*d = if xm >= x { tol1 } else { -tol1 };
}
u = x + *d;
} else {
*e = if x >= xm {
bracket_lo - x
} else {
bracket_hi - x
};
*d = CGOLD * (*e);
u = x + *d;
}
} else {
*e = if x >= xm {
bracket_lo - x
} else {
bracket_hi - x
};
*d = CGOLD * (*e);
u = x + *d;
}
let u_actual = if (*d).abs() >= tol1 {
u
} else if *d >= 0.0 {
x + tol1
} else {
x - tol1
};
let fu = f(u_actual);
if fu <= fx {
if u_actual >= x {
bracket_lo = x;
} else {
bracket_hi = x;
}
v = w;
fv = fw;
w = x;
fw = fx;
x = u_actual;
fx = fu;
} else {
if u_actual < x {
bracket_lo = u_actual;
} else {
bracket_hi = u_actual;
}
if fu <= fw || w == x {
v = w;
fv = fw;
w = u_actual;
fw = fu;
} else if fu <= fv || v == x || v == w {
v = u_actual;
fv = fu;
}
}
}
*n_iter_out = max_iter;
*a = bracket_lo;
*c = bracket_hi;
*b = x;
*fb = fx;
(x, fx)
}
pub(crate) fn joint_wald_chi_sq(
xtvix: MatRef<'_, f64>,
betas: &[f64],
sigma_sq: f64,
target_indices: &[u32],
mut joint_k_inv: MatMut<'_, f64>,
mut joint_sigma_t_chol: MatMut<'_, f64>,
joint_rhs: &mut [f64],
) -> f64 {
use faer::linalg::solvers::Solve;
use faer::reborrow::Reborrow;
let p = betas.len();
let k = target_indices.len();
if k == 0 || p == 0 || !(sigma_sq.is_finite() && sigma_sq > FLOAT_NEAR_ZERO) {
return f64::NAN;
}
let chol = match xtvix.llt(faer::Side::Lower) {
Ok(c) => c,
Err(_) => return f64::NAN,
};
for j in 0..p {
for i in 0..p {
joint_k_inv[(i, j)] = if i == j { 1.0 } else { 0.0 };
}
}
chol.solve_in_place(joint_k_inv.rb_mut());
for (a, &ti) in target_indices.iter().enumerate() {
let ti = ti as usize;
if ti >= p {
return f64::NAN;
}
for (b, &tj) in target_indices.iter().enumerate() {
let tj = tj as usize;
if tj >= p {
return f64::NAN;
}
joint_sigma_t_chol[(a, b)] = joint_k_inv[(ti, tj)];
}
}
let sigma_t_view = joint_sigma_t_chol.rb().submatrix(0, 0, k, k);
let chol_t = match sigma_t_view.llt(faer::Side::Lower) {
Ok(c) => c,
Err(_) => return f64::NAN,
};
for (a, &ti) in target_indices.iter().enumerate() {
joint_rhs[a] = betas[ti as usize];
}
{
let mut rhs = MatMut::from_column_major_slice_mut(&mut joint_rhs[..k], k, 1usize);
chol_t.solve_in_place(rhs.rb_mut());
}
let mut w_raw = 0.0_f64;
for (a, &ti) in target_indices.iter().enumerate() {
w_raw += betas[ti as usize] * joint_rhs[a];
}
if !w_raw.is_finite() {
return f64::NAN;
}
w_raw / sigma_sq
}
fn nan_fill_outputs(scratch: &mut LmeScratch<'_>, p: usize, target_indices: &[u32]) {
for v in scratch.betas[..p].iter_mut() {
*v = f64::NAN;
}
for &tj in target_indices.iter() {
let tj = tj as usize;
if tj < scratch.var_diag.len() {
scratch.var_diag[tj] = f64::NAN;
}
if tj < scratch.t_sq.len() {
scratch.t_sq[tj] = f64::NAN;
}
}
}
pub fn lme_fit<'a>(
_x: MatRef<'_, f64>,
_y: &[f64],
_cluster_ids: &[u32],
target_indices: &[u32],
theta_start: Option<f64>,
mut scratch: LmeScratch<'a>,
) -> LmeFitView<'a> {
let p = scratch.betas.len();
let mut n_evals: u32 = 0;
let mut brent_iters: u32 = 0;
let truth_bracket = theta_start.map(|theta0| {
let center = theta0.ln().clamp(
LOG_THETA_LOW + TRUTH_BRACKET_HALF_WIDTH,
LOG_THETA_HIGH - TRUTH_BRACKET_HALF_WIDTH,
);
(
center - TRUTH_BRACKET_HALF_WIDTH,
center,
center + TRUTH_BRACKET_HALF_WIDTH,
)
});
let (mut log_a, mut log_b, mut log_c) =
truth_bracket.unwrap_or((LOG_THETA_LOW, LOG_THETA_MID, LOG_THETA_HIGH));
let mut fa: f64;
let mut fb: f64;
let mut fc: f64;
let mut truth_retry_left = truth_bracket.is_some();
let boundary_pre_brent: Option<u8>;
loop {
fa = profiled_deviance(log_a.exp(), &mut scratch);
fb = profiled_deviance(log_b.exp(), &mut scratch);
fc = profiled_deviance(log_c.exp(), &mut scratch);
n_evals += 3;
let mut bracket_ok = fb < fa && fb < fc;
if !bracket_ok {
let log_mid = 0.5 * (log_a + log_c);
let f_mid = profiled_deviance(log_mid.exp(), &mut scratch);
n_evals += 1;
if f_mid < fa && f_mid < fc {
log_b = log_mid;
fb = f_mid;
bracket_ok = true;
}
}
if !bracket_ok {
const LN10: f64 = std::f64::consts::LN_10;
for _ in 0..2 {
if fa <= fc {
let new_log_a = log_a - LN10;
let new_fa = profiled_deviance(new_log_a.exp(), &mut scratch);
n_evals += 1;
if new_fa > fa {
log_c = log_b;
fc = fb;
log_b = log_a;
fb = fa;
log_a = new_log_a;
fa = new_fa;
bracket_ok = true;
break;
}
log_a = new_log_a;
fa = new_fa;
} else {
let new_log_c = log_c + LN10;
let new_fc = profiled_deviance(new_log_c.exp(), &mut scratch);
n_evals += 1;
if new_fc > fc {
log_a = log_b;
fa = fb;
log_b = log_c;
fb = fc;
log_c = new_log_c;
fc = new_fc;
bracket_ok = true;
break;
}
log_c = new_log_c;
fc = new_fc;
}
}
}
if bracket_ok {
boundary_pre_brent = None;
break;
}
if truth_retry_left {
truth_retry_left = false;
log_a = LOG_THETA_LOW;
log_b = LOG_THETA_MID;
log_c = LOG_THETA_HIGH;
continue;
}
boundary_pre_brent = Some(if fa <= fc { 1 } else { 2 });
break;
}
let mut log_theta_hat = log_b;
let mut brent_d: f64 = 0.0;
let mut brent_e: f64 = 0.0;
if boundary_pre_brent.is_none() {
let mut brent_n_evals: u32 = 0;
let (xmin, _fmin) = {
let f_closure = |log_x: f64| -> f64 {
brent_n_evals += 1;
profiled_deviance(log_x.exp(), &mut scratch)
};
brent_minimize(
&mut log_a,
&mut log_b,
&mut log_c,
&mut fa,
&mut fb,
&mut fc,
&mut brent_d,
&mut brent_e,
f_closure,
BRENT_REL_TOL,
MAX_BRENT_ITERS,
&mut brent_iters,
)
};
n_evals += brent_n_evals;
log_theta_hat = xmin;
}
let mut boundary_hit: u8 = boundary_pre_brent.unwrap_or(0);
if boundary_hit == 0 {
if log_theta_hat - LOG_THETA_LOW < BOUNDARY_LOG_SLACK {
boundary_hit = 1;
} else if LOG_THETA_HIGH - log_theta_hat < BOUNDARY_LOG_SLACK {
let mut log_a2 = (1e-1_f64).ln();
let mut log_b2 = (1e1_f64).ln();
let mut log_c2 = (1e3_f64).ln();
let mut fa2 = profiled_deviance(log_a2.exp(), &mut scratch);
let mut fb2 = profiled_deviance(log_b2.exp(), &mut scratch);
let mut fc2 = profiled_deviance(log_c2.exp(), &mut scratch);
n_evals += 3;
if fb2 < fa2 && fb2 < fc2 {
let mut brent_n_evals2: u32 = 0;
let mut bd2: f64 = 0.0;
let mut be2: f64 = 0.0;
let mut bi2: u32 = 0;
let (xmin2, _fmin2) = {
let f_closure = |log_x: f64| -> f64 {
brent_n_evals2 += 1;
profiled_deviance(log_x.exp(), &mut scratch)
};
brent_minimize(
&mut log_a2,
&mut log_b2,
&mut log_c2,
&mut fa2,
&mut fb2,
&mut fc2,
&mut bd2,
&mut be2,
f_closure,
BRENT_REL_TOL,
MAX_BRENT_ITERS,
&mut bi2,
)
};
n_evals += brent_n_evals2;
brent_iters = bi2;
log_theta_hat = xmin2;
if (1e3_f64).ln() - log_theta_hat < BOUNDARY_LOG_SLACK {
boundary_hit = 2;
nan_fill_outputs(&mut scratch, p, target_indices);
let factor = scratch.xtvix_factor.into_const();
return LmeFitView {
betas: &scratch.betas[..p],
var_diag: &scratch.var_diag[..p],
t_sq: &scratch.t_sq[..p],
factor,
sigma_sq: f64::NAN,
tau_sq_hat: f64::NAN,
converged: false,
boundary_hit,
n_iter: brent_iters,
n_evals,
joint_t_sq: f64::NAN,
};
}
boundary_hit = 0;
} else {
nan_fill_outputs(&mut scratch, p, target_indices);
let factor = scratch.xtvix_factor.into_const();
return LmeFitView {
betas: &scratch.betas[..p],
var_diag: &scratch.var_diag[..p],
t_sq: &scratch.t_sq[..p],
factor,
sigma_sq: f64::NAN,
tau_sq_hat: f64::NAN,
converged: false,
boundary_hit: 2,
n_iter: brent_iters,
n_evals,
joint_t_sq: f64::NAN,
};
}
}
}
let pin_theta = if boundary_hit == 1 {
LOG_THETA_LOW.exp()
} else {
log_theta_hat.exp()
};
let pin_dev = profiled_deviance(pin_theta, &mut scratch);
n_evals += 1;
let rank_deficient = chol_rank_deficient(scratch.xtvix_factor.rb(), p, crate::lmm::EPS_RANK);
if !pin_dev.is_finite() || rank_deficient {
nan_fill_outputs(&mut scratch, p, target_indices);
let factor = scratch.xtvix_factor.into_const();
return LmeFitView {
betas: &scratch.betas[..p],
var_diag: &scratch.var_diag[..p],
t_sq: &scratch.t_sq[..p],
factor,
sigma_sq: f64::NAN,
tau_sq_hat: f64::NAN,
converged: false,
boundary_hit,
n_iter: brent_iters,
n_evals,
joint_t_sq: f64::NAN,
};
}
let sigma_sq = scratch.sigma_sq;
for &tj in target_indices.iter() {
let tj = tj as usize;
if tj >= p {
continue;
}
for v in scratch.u_scratch[..p].iter_mut() {
*v = 0.0;
}
let mut bad_diag = false;
for i in 0..p {
let b_i = if i == tj { 1.0 } else { 0.0 };
let mut acc = b_i;
for k in 0..i {
acc -= scratch.xtvix_factor[(i, k)] * scratch.u_scratch[k];
}
let l_ii = scratch.xtvix_factor[(i, i)];
if l_ii.abs() < FLOAT_NEAR_ZERO {
scratch.u_scratch[i] = f64::NAN;
bad_diag = true;
} else {
scratch.u_scratch[i] = acc / l_ii;
}
}
let mut norm_sq = 0.0;
for &v in scratch.u_scratch[..p].iter() {
norm_sq += v * v;
}
let vd = sigma_sq * norm_sq;
scratch.var_diag[tj] = vd;
if !bad_diag && vd > FLOAT_NEAR_ZERO && vd.is_finite() {
let beta_j = scratch.betas[tj];
scratch.t_sq[tj] = (beta_j * beta_j) / vd;
} else {
scratch.t_sq[tj] = f64::NAN;
}
}
let joint_t_sq = joint_wald_chi_sq(
scratch.xtvix.rb(),
&scratch.betas[..p],
sigma_sq,
target_indices,
scratch.joint_k_inv.rb_mut(),
scratch.joint_sigma_t_chol.rb_mut(),
&mut scratch.joint_rhs[..p],
);
let converged = brent_iters < MAX_BRENT_ITERS;
let factor = scratch.xtvix_factor.into_const();
let tau_sq_hat = pin_theta * pin_theta * sigma_sq;
LmeFitView {
betas: &scratch.betas[..p],
var_diag: &scratch.var_diag[..p],
t_sq: &scratch.t_sq[..p],
factor,
sigma_sq,
tau_sq_hat,
converged,
boundary_hit,
n_iter: brent_iters,
n_evals,
joint_t_sq,
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::test_support::{build_lme_scratch, TestWs};
#[test]
fn brent_constants_log_spaced() {
assert!((LOG_THETA_LOW - (1e-4_f64).ln()).abs() < 1e-12);
assert!((LOG_THETA_HIGH - (1e3_f64).ln()).abs() < 1e-12);
#[expect(
clippy::assertions_on_constants,
reason = "compile-time ordering invariant on the log-theta brackets"
)]
{
assert!(LOG_THETA_LOW < LOG_THETA_MID && LOG_THETA_MID < LOG_THETA_HIGH);
}
}
#[test]
fn lme_add_rows_panel_matches_scalar_reference() {
let n = 611; let p = 5;
let k = 40; let mut x = faer::Mat::<f64>::zeros(n, p);
for i in 0..n {
for j in 0..p {
x[(i, j)] = ((((i * 13 + j * 29 + 3) % 47) as f64) / 11.0 - 2.0).sin();
}
}
let y: Vec<f64> = (0..n)
.map(|i| ((((i * 31 + 7) % 53) as f64) / 9.0 - 2.5).cos())
.collect();
let ids: Vec<u32> = (0..n).map(|i| (i / 16) as u32).collect();
let mut ref_xtx = faer::Mat::<f64>::zeros(p, p);
let mut ref_xty = vec![0.0f64; p];
let mut ref_yty = 0.0f64;
let mut ref_sum_xc = faer::Mat::<f64>::zeros(p, k);
let mut ref_sum_yc = vec![0.0f64; k];
let mut ref_sizes = vec![0u32; k];
let mut ref_seen = 0u32;
for row in 0..n {
let y_row = y[row];
let c = ids[row] as usize;
for j in 0..p {
let x_rj = x[(row, j)];
ref_sum_xc[(j, c)] += x_rj;
for i in j..p {
ref_xtx[(i, j)] += x[(row, i)] * x_rj;
}
ref_xty[j] += x_rj * y_row;
}
ref_sum_yc[c] += y_row;
ref_sizes[c] += 1;
ref_yty += y_row * y_row;
ref_seen = ref_seen.max(c as u32 + 1);
}
let mut ws = TestWs::new(n, p, k);
ws.reset_lme_suff_stats();
{
let mut s = LmeSuffStats {
xtx: ws.lme_xtx.as_mut(),
xty: &mut ws.lme_xty,
yty: &mut ws.lme_yty,
sum_xc: ws.lme_sum_xc.as_mut(),
sum_yc: &mut ws.lme_sum_yc,
cluster_sizes: &mut ws.lme_cluster_sizes,
n_clusters_seen: &mut ws.lme_n_clusters_seen,
panel_x: &mut ws.panel_x,
panel_y: &mut ws.panel_y,
};
s.add_rows(x.as_ref(), &y, &ids);
}
for j in 0..p {
for i in j..p {
let (got, want) = (ws.lme_xtx[(i, j)], ref_xtx[(i, j)]);
assert!(
(got - want).abs() <= 1e-12 * want.abs().max(1.0),
"xtx[{i},{j}] = {got}, scalar reference {want}"
);
}
assert!(
(ws.lme_xty[j] - ref_xty[j]).abs() <= 1e-12 * ref_xty[j].abs().max(1.0),
"xty[{j}] diverged"
);
for c in 0..k {
assert_eq!(
ws.lme_sum_xc[(j, c)].to_bits(),
ref_sum_xc[(j, c)].to_bits(),
"sum_xc[{j},{c}] must stay bit-identical"
);
}
}
for c in 0..k {
assert_eq!(ws.lme_sum_yc[c].to_bits(), ref_sum_yc[c].to_bits());
assert_eq!(ws.lme_cluster_sizes[c], ref_sizes[c]);
}
assert_eq!(ws.lme_yty.to_bits(), ref_yty.to_bits());
assert_eq!(ws.lme_n_clusters_seen, ref_seen);
}
#[test]
fn profiled_deviance_overwrites_state() {
use faer::Mat;
let n: usize = 6;
let p: usize = 2;
let n_clusters: usize = 2;
let cluster_ids: Vec<u32> = vec![0, 0, 0, 1, 1, 1];
let x1: [f64; 6] = [
0.1257302210933933,
-0.1321048632913019,
0.6404226504432821,
0.10490011715303971,
-0.535669373161111,
0.36159505490948474,
];
let y: [f64; 6] = [
0.7718630718197979,
-0.09922526468089643,
1.4699547603093044,
1.266192714345762,
-1.8688474280172964,
1.3141089963737067,
];
let mut x = Mat::<f64>::zeros(n, p);
for i in 0..n {
x[(i, 0)] = 1.0;
x[(i, 1)] = x1[i];
}
let mut ws = TestWs::new(n, p, n_clusters);
ws.reset_lme_suff_stats();
{
let mut s = LmeSuffStats {
xtx: ws.lme_xtx.as_mut(),
xty: &mut ws.lme_xty,
yty: &mut ws.lme_yty,
sum_xc: ws.lme_sum_xc.as_mut(),
sum_yc: &mut ws.lme_sum_yc,
cluster_sizes: &mut ws.lme_cluster_sizes,
n_clusters_seen: &mut ws.lme_n_clusters_seen,
panel_x: &mut ws.panel_x,
panel_y: &mut ws.panel_y,
};
s.add_rows(x.as_ref(), &y, &cluster_ids);
}
let mut scratch = build_lme_scratch(&mut ws, n as u32, n_clusters as u32);
let dev1_a = profiled_deviance(1.0, &mut scratch);
let beta1_a = [scratch.betas[0], scratch.betas[1]];
assert!(dev1_a.is_finite());
let _ = profiled_deviance(2.0, &mut scratch);
let dev1_b = profiled_deviance(1.0, &mut scratch);
let beta1_b = [scratch.betas[0], scratch.betas[1]];
assert_eq!(
dev1_a, dev1_b,
"deviance(θ=1) must be reproducible (no stale state)"
);
assert_eq!(
beta1_a, beta1_b,
"β̂(θ=1) must be reproducible (no stale state)"
);
}
#[test]
fn brent_finds_interior_minimum() {
let g = |x: f64| {
let dx = x - 0.7;
dx.powi(4) + 0.1 * dx.powi(2)
};
let a0 = -2.0_f64;
let c0 = 2.0_f64;
let mut a = a0;
let mut b = 0.5_f64;
let mut c = c0;
let mut fa = g(a);
let mut fb = g(b);
let mut fc = g(c);
let mut d = 0.0_f64;
let mut e = 0.0_f64;
let mut n_iter = 0_u32;
let (xmin, fmin) = brent_minimize(
&mut a,
&mut b,
&mut c,
&mut fa,
&mut fb,
&mut fc,
&mut d,
&mut e,
g,
BRENT_REL_TOL,
MAX_BRENT_ITERS,
&mut n_iter,
);
assert!(
xmin > a0 && xmin < c0,
"xmin {xmin} must stay inside the initial bracket"
);
assert!(
fmin <= g(a0) && fmin <= g(c0),
"fmin {fmin} must not exceed the endpoints"
);
assert!(
fmin <= g(xmin) + 1e-12,
"fmin must equal g at the returned minimiser"
);
assert!(
n_iter < MAX_BRENT_ITERS,
"n_iter {n_iter} must terminate below the cap"
);
}
#[test]
fn profiled_deviance_is_smooth_across_log_theta() {
use faer::Mat;
let n: usize = 6;
let p: usize = 2;
let n_clusters: usize = 2;
let cluster_ids: Vec<u32> = vec![0, 0, 0, 1, 1, 1];
let x1: [f64; 6] = [
0.1257302210933933,
-0.1321048632913019,
0.6404226504432821,
0.10490011715303971,
-0.535669373161111,
0.36159505490948474,
];
let y: [f64; 6] = [
0.7718630718197979,
-0.09922526468089643,
1.4699547603093044,
1.266192714345762,
-1.8688474280172964,
1.3141089963737067,
];
let mut x = Mat::<f64>::zeros(n, p);
for i in 0..n {
x[(i, 0)] = 1.0;
x[(i, 1)] = x1[i];
}
let mut ws = TestWs::new(n, p, n_clusters);
ws.reset_lme_suff_stats();
{
let mut s = LmeSuffStats {
xtx: ws.lme_xtx.as_mut(),
xty: &mut ws.lme_xty,
yty: &mut ws.lme_yty,
sum_xc: ws.lme_sum_xc.as_mut(),
sum_yc: &mut ws.lme_sum_yc,
cluster_sizes: &mut ws.lme_cluster_sizes,
n_clusters_seen: &mut ws.lme_n_clusters_seen,
panel_x: &mut ws.panel_x,
panel_y: &mut ws.panel_y,
};
s.add_rows(x.as_ref(), &y, &cluster_ids);
}
let n_samples = 11;
let log_lo = -4.0_f64;
let log_hi = 2.0_f64;
let mut devs = vec![0.0_f64; n_samples];
{
let mut scratch = build_lme_scratch(&mut ws, n as u32, n_clusters as u32);
for (idx, slot) in devs.iter_mut().enumerate() {
let frac = idx as f64 / (n_samples - 1) as f64;
let log_theta = log_lo + (log_hi - log_lo) * frac;
let theta = log_theta.exp();
*slot = profiled_deviance(theta, &mut scratch);
}
}
for (idx, d) in devs.iter().enumerate() {
assert!(d.is_finite(), "dev[{idx}] = {d} (not finite)");
}
let mut signs: Vec<i32> = Vec::with_capacity(n_samples - 2);
for i in 1..(n_samples - 1) {
let d2 = devs[i + 1] - 2.0 * devs[i] + devs[i - 1];
if d2.abs() > 1e-10 {
signs.push(if d2 > 0.0 { 1 } else { -1 });
}
}
let mut n_changes = 0;
for w in signs.windows(2) {
if w[0] != w[1] {
n_changes += 1;
}
}
assert!(
n_changes <= 1,
"profiled_deviance second-derivative sign changes = {n_changes} > 1; devs = {devs:?}"
);
}
#[test]
fn lme_fit_converges_with_valid_boundary() {
use faer::Mat;
const N: usize = 60;
const P: usize = 3;
const K: usize = 6;
let x1: [f64; 60] = [
0.30471707975443135,
-1.0399841062404955,
0.7504511958064572,
0.9405647163912139,
-1.9510351886538364,
-1.302179506862318,
0.12784040316728537,
-0.3162425923435822,
-0.016801157504288795,
-0.85304392757358,
0.8793979748628286,
0.7777919354289483,
0.06603069756121605,
1.1272412069680329,
0.4675093422520456,
-0.8592924628832382,
0.36875078408249884,
-0.9588826008289989,
0.8784503013072725,
-0.049925910986252896,
-0.18486236354526056,
-0.6809295444039414,
1.2225413386740303,
-0.15452948206880215,
-0.4283278221631072,
-0.3521335504882296,
0.5323091855533487,
0.36544406436407834,
0.4127326115959884,
0.43082100300788273,
2.1416476008704612,
-0.4064150163846156,
-0.5122427290715373,
-0.8137727282478777,
0.6159794225754956,
1.1289722927208916,
-0.11394745765487507,
-0.840156476962528,
-0.8244812156912396,
0.6505927878247011,
0.7432541712034423,
0.543154268305195,
-0.6655097072886943,
0.23216132306671977,
0.11668580914072822,
0.21868859672901295,
0.8714287779481898,
0.22359554877468227,
0.6789135630718949,
0.06757906948889146,
0.28911939868998415,
0.6312882258385404,
-1.4571558198556664,
-0.31967121635730134,
-0.4703726542927955,
-0.6388778482433419,
-0.27514225122668373,
1.4949413112343959,
-0.8658311156932432,
0.9682783545914808,
];
let x2: [f64; 60] = [
-1.6828697716158048,
-0.33488502998577485,
0.1627530651050056,
0.5862223313592781,
0.711226579792855,
0.7933472351999252,
-0.3487250722484376,
-0.46235179266456716,
0.8579758812571538,
-0.1913043248816149,
-1.2756863233379219,
-1.1332872140034806,
-0.9194522860016113,
0.49716074405376404,
0.14242573607056525,
0.6904853540677682,
-0.42725264633653426,
0.15853969107671423,
0.6255903939673367,
-0.3093465397202384,
0.45677523755741145,
-0.6619259410666513,
-0.3630538465650718,
-0.3817378939983291,
-1.1958396455890397,
0.4869724807855818,
-0.46940234020272387,
0.01249411872768743,
0.48074665890590895,
0.4465311760299441,
0.6653851089727862,
-0.09848548450942361,
-0.42329831204415375,
-0.07971821090639905,
-1.6873344339580298,
-1.4471124724230873,
-1.3226996123544024,
-0.9972468276014818,
0.3997742267234366,
-0.9054790553600608,
-0.3781625540393897,
1.2992282977860654,
-0.35626397106142593,
0.7375155684670865,
-0.933617680009877,
-0.20543755786763002,
-0.9500220549105812,
-0.3390330759005625,
0.8403081374573955,
-1.7273204231923487,
0.43442364354585733,
0.2377356023322779,
-0.5941499556967944,
-1.4460578543884546,
0.07212950771386951,
-0.5294927090638024,
0.23267621135470395,
0.02185214552344288,
1.6017788913209154,
-0.23935562747302427,
];
let y_data: [f64; 60] = [
0.4634838483807412,
-1.6958367493518591,
-0.07656189550801018,
0.5791988172301176,
-0.623788696202619,
-1.5640666548456659,
-0.05286302450777913,
-0.40566566398296267,
-1.082675989070649,
-0.6398666414265639,
0.8986776652224437,
1.1127414469896355,
1.1065939034406163,
3.740015159794447,
0.830029945030196,
-0.3524130442209007,
-1.5476759552488093,
0.06153714483417494,
-0.8333554776977836,
-0.4197403105019113,
-0.42835728975820775,
0.1797922600939008,
2.0965054734973476,
0.19096753379326917,
-1.3788672930282488,
-0.9522449167472798,
-1.4102109341814926,
0.5597880597592666,
0.824978479993657,
2.343004843474632,
0.7533972033686575,
0.8633702245648947,
-0.7432683368811597,
-0.7561198927806311,
-0.6351566537898222,
-0.3691685843204092,
1.0273694686701114,
-1.4272550988991226,
0.6853189837147633,
0.010351200249558046,
1.7179220596063591,
1.2720016651761046,
-1.2435885938521603,
0.4099566730250238,
-0.7373919750549822,
1.35341093351855,
0.2238950195023513,
-0.9898128365682515,
-1.333705002777748,
-0.19843480740416738,
1.9931161686261913,
1.406582108224644,
-0.4972904954045434,
-0.08212174931081978,
0.445137615643293,
-0.14552036838998333,
-0.33955737323366064,
2.7199407103663717,
1.0045452753534772,
2.1737935566316495,
];
let mut x = Mat::<f64>::zeros(N, P);
for i in 0..N {
x[(i, 0)] = 1.0;
x[(i, 1)] = x1[i];
x[(i, 2)] = x2[i];
}
let y: Vec<f64> = y_data.to_vec();
let cluster_ids: Vec<u32> = (0..N).map(|i| (i % K) as u32).collect();
let mut ws = TestWs::new(N, P, K);
ws.reset_lme_suff_stats();
{
let mut s = LmeSuffStats {
xtx: ws.lme_xtx.as_mut(),
xty: &mut ws.lme_xty,
yty: &mut ws.lme_yty,
sum_xc: ws.lme_sum_xc.as_mut(),
sum_yc: &mut ws.lme_sum_yc,
cluster_sizes: &mut ws.lme_cluster_sizes,
n_clusters_seen: &mut ws.lme_n_clusters_seen,
panel_x: &mut ws.panel_x,
panel_y: &mut ws.panel_y,
};
s.add_rows(x.as_ref(), &y, &cluster_ids);
}
let target_indices: Vec<u32> = vec![1, 2];
let fit_betas: [f64; P];
let fit_t_sq: [f64; P];
let fit_var_diag: [f64; P];
let fit_converged: bool;
let fit_boundary: u8;
{
let scratch = build_lme_scratch(&mut ws, N as u32, K as u32);
let fit = lme_fit(x.as_ref(), &y, &cluster_ids, &target_indices, None, scratch);
fit_betas = [fit.betas[0], fit.betas[1], fit.betas[2]];
fit_t_sq = [fit.t_sq[0], fit.t_sq[1], fit.t_sq[2]];
fit_var_diag = [fit.var_diag[0], fit.var_diag[1], fit.var_diag[2]];
fit_converged = fit.converged;
fit_boundary = fit.boundary_hit;
}
assert!(fit_converged, "fit failed to converge");
assert!(
fit_boundary == 0 || fit_boundary == 1,
"boundary_hit = {fit_boundary} (expected 0 or 1)"
);
for (j, b) in fit_betas.iter().enumerate() {
assert!(b.is_finite(), "β̂[{j}] must be finite on a converged fit");
}
for &tj in &target_indices {
let j = tj as usize;
assert!(
fit_t_sq[j].is_finite() && fit_t_sq[j] > 0.0,
"t²[{j}] must be finite and strictly positive on a converged fit"
);
assert!(
fit_var_diag[j] >= 0.0,
"var_diag[{j}] = {} must be non-negative",
fit_var_diag[j]
);
}
}
#[test]
fn lme_fit_detects_tau_zero_boundary() {
use faer::Mat;
const N: usize = 60;
const P: usize = 3;
const K: usize = 6;
let x1: [f64; 60] = [
0.30471707975443135,
-1.0399841062404955,
0.7504511958064572,
0.9405647163912139,
-1.9510351886538364,
-1.302179506862318,
0.12784040316728537,
-0.3162425923435822,
-0.016801157504288795,
-0.85304392757358,
0.8793979748628286,
0.7777919354289483,
0.06603069756121605,
1.1272412069680329,
0.4675093422520456,
-0.8592924628832382,
0.36875078408249884,
-0.9588826008289989,
0.8784503013072725,
-0.049925910986252896,
-0.18486236354526056,
-0.6809295444039414,
1.2225413386740303,
-0.15452948206880215,
-0.4283278221631072,
-0.3521335504882296,
0.5323091855533487,
0.36544406436407834,
0.4127326115959884,
0.43082100300788273,
2.1416476008704612,
-0.4064150163846156,
-0.5122427290715373,
-0.8137727282478777,
0.6159794225754956,
1.1289722927208916,
-0.11394745765487507,
-0.840156476962528,
-0.8244812156912396,
0.6505927878247011,
0.7432541712034423,
0.543154268305195,
-0.6655097072886943,
0.23216132306671977,
0.11668580914072822,
0.21868859672901295,
0.8714287779481898,
0.22359554877468227,
0.6789135630718949,
0.06757906948889146,
0.28911939868998415,
0.6312882258385404,
-1.4571558198556664,
-0.31967121635730134,
-0.4703726542927955,
-0.6388778482433419,
-0.27514225122668373,
1.4949413112343959,
-0.8658311156932432,
0.9682783545914808,
];
let x2: [f64; 60] = [
-1.6828697716158048,
-0.33488502998577485,
0.1627530651050056,
0.5862223313592781,
0.711226579792855,
0.7933472351999252,
-0.3487250722484376,
-0.46235179266456716,
0.8579758812571538,
-0.1913043248816149,
-1.2756863233379219,
-1.1332872140034806,
-0.9194522860016113,
0.49716074405376404,
0.14242573607056525,
0.6904853540677682,
-0.42725264633653426,
0.15853969107671423,
0.6255903939673367,
-0.3093465397202384,
0.45677523755741145,
-0.6619259410666513,
-0.3630538465650718,
-0.3817378939983291,
-1.1958396455890397,
0.4869724807855818,
-0.46940234020272387,
0.01249411872768743,
0.48074665890590895,
0.4465311760299441,
0.6653851089727862,
-0.09848548450942361,
-0.42329831204415375,
-0.07971821090639905,
-1.6873344339580298,
-1.4471124724230873,
-1.3226996123544024,
-0.9972468276014818,
0.3997742267234366,
-0.9054790553600608,
-0.3781625540393897,
1.2992282977860654,
-0.35626397106142593,
0.7375155684670865,
-0.933617680009877,
-0.20543755786763002,
-0.9500220549105812,
-0.3390330759005625,
0.8403081374573955,
-1.7273204231923487,
0.43442364354585733,
0.2377356023322779,
-0.5941499556967944,
-1.4460578543884546,
0.07212950771386951,
-0.5294927090638024,
0.23267621135470395,
0.02185214552344288,
1.6017788913209154,
-0.23935562747302427,
];
let mut x_f64 = Mat::<f64>::zeros(N, P);
for i in 0..N {
x_f64[(i, 0)] = 1.0;
x_f64[(i, 1)] = x1[i];
x_f64[(i, 2)] = x2[i];
}
let beta_true = [0.5, 0.4, -0.3];
let y_f64: Vec<f64> = (0..N)
.map(|i| {
let row_signal =
beta_true[0] + beta_true[1] * x_f64[(i, 1)] + beta_true[2] * x_f64[(i, 2)];
let noise = ((i as f64) * 0.0137).sin() * 0.05;
row_signal + noise
})
.collect();
let cluster_ids: Vec<u32> = (0..N).map(|i| (i % K) as u32).collect();
let mut xtx_ols = Mat::<f64>::zeros(P, P);
let mut xty_ols = [0.0_f64; P];
for i in 0..N {
for jj in 0..P {
xty_ols[jj] += x_f64[(i, jj)] * y_f64[i];
for ii in jj..P {
xtx_ols[(ii, jj)] += x_f64[(i, ii)] * x_f64[(i, jj)];
}
}
}
let chol = xtx_ols
.as_ref()
.llt(faer::Side::Lower)
.expect("OLS Cholesky failed");
let mut rhs = Mat::<f64>::zeros(P, 1);
for jj in 0..P {
rhs[(jj, 0)] = xty_ols[jj];
}
use faer::linalg::solvers::Solve;
chol.solve_in_place(rhs.as_mut());
let ols_beta = [rhs[(0, 0)], rhs[(1, 0)], rhs[(2, 0)]];
let mut x = Mat::<f64>::zeros(N, P);
for i in 0..N {
x[(i, 0)] = x_f64[(i, 0)];
x[(i, 1)] = x_f64[(i, 1)];
x[(i, 2)] = x_f64[(i, 2)];
}
let y: Vec<f64> = y_f64.clone();
let mut ws = TestWs::new(N, P, K);
ws.reset_lme_suff_stats();
{
let mut s = LmeSuffStats {
xtx: ws.lme_xtx.as_mut(),
xty: &mut ws.lme_xty,
yty: &mut ws.lme_yty,
sum_xc: ws.lme_sum_xc.as_mut(),
sum_yc: &mut ws.lme_sum_yc,
cluster_sizes: &mut ws.lme_cluster_sizes,
n_clusters_seen: &mut ws.lme_n_clusters_seen,
panel_x: &mut ws.panel_x,
panel_y: &mut ws.panel_y,
};
s.add_rows(x.as_ref(), &y, &cluster_ids);
}
let targets: Vec<u32> = vec![1, 2];
let fit_betas: [f64; P];
let fit_boundary: u8;
let fit_converged: bool;
{
let scratch = build_lme_scratch(&mut ws, N as u32, K as u32);
let fit = lme_fit(x.as_ref(), &y, &cluster_ids, &targets, None, scratch);
fit_betas = [fit.betas[0], fit.betas[1], fit.betas[2]];
fit_boundary = fit.boundary_hit;
fit_converged = fit.converged;
}
assert!(fit_converged, "lme_fit failed to converge on τ=0 case");
assert_eq!(
fit_boundary, 1,
"expected τ̂≈0 boundary_hit=1, got {fit_boundary}"
);
for j in 0..P {
let delta = (fit_betas[j] - ols_beta[j]).abs();
assert!(
delta < 1e-6,
"β̂[{j}] = {} vs OLS β̂ = {} (Δ = {})",
fit_betas[j],
ols_beta[j],
delta
);
}
}
fn make_clustered_fixture_interior_theta() -> (faer::Mat<f64>, Vec<f64>, Vec<u32>) {
use faer::Mat;
const N: usize = 60;
const P: usize = 3;
const K: usize = 6;
let offsets: [f64; K] = [0.9, -0.6, 0.3, -0.2, 0.65, -1.05];
let mut x = Mat::<f64>::zeros(N, P);
let mut y = vec![0.0_f64; N];
let mut cluster_ids = vec![0_u32; N];
let mut state = 0xD1B5_4A32_D192_ED03_u64;
for i in 0..N {
x[(i, 0)] = 1.0;
state = state.wrapping_mul(0x9E3779B97F4A7C15).wrapping_add(1);
let x1 = ((state >> 11) as f64 / (1u64 << 53) as f64) * 2.0 - 1.0;
x[(i, 1)] = x1;
state = state.wrapping_mul(0x9E3779B97F4A7C15).wrapping_add(1);
let x2 = ((state >> 11) as f64 / (1u64 << 53) as f64) * 2.0 - 1.0;
x[(i, 2)] = x2;
state = state.wrapping_mul(0x9E3779B97F4A7C15).wrapping_add(1);
let noise = ((state >> 11) as f64 / (1u64 << 53) as f64) * 2.0 - 1.0;
let c = i % K;
y[i] = 0.5 + 0.4 * x1 - 0.3 * x2 + offsets[c] + 0.3 * noise;
cluster_ids[i] = c as u32;
}
(x, y, cluster_ids)
}
#[test]
fn lme_truth_start_matches_cold_minimum_with_fewer_evals() {
let (x, y, cluster_ids) = make_clustered_fixture_interior_theta();
let n = y.len();
let p = x.ncols();
let k = *cluster_ids.iter().max().unwrap() as usize + 1;
let targets: Vec<u32> = vec![1, 2];
let mut ws = TestWs::new(n, p, k);
ws.reset_lme_suff_stats();
{
let mut s = LmeSuffStats {
xtx: ws.lme_xtx.as_mut(),
xty: &mut ws.lme_xty,
yty: &mut ws.lme_yty,
sum_xc: ws.lme_sum_xc.as_mut(),
sum_yc: &mut ws.lme_sum_yc,
cluster_sizes: &mut ws.lme_cluster_sizes,
n_clusters_seen: &mut ws.lme_n_clusters_seen,
panel_x: &mut ws.panel_x,
panel_y: &mut ws.panel_y,
};
s.add_rows(x.as_ref(), &y, &cluster_ids);
}
let (cold_betas, cold_t_sq, cold_boundary, cold_evals) = {
let scratch = build_lme_scratch(&mut ws, n as u32, k as u32);
let fit = lme_fit(x.as_ref(), &y, &cluster_ids, &targets, None, scratch);
assert!(fit.converged, "cold fit must converge");
(
fit.betas.to_vec(),
fit.t_sq.to_vec(),
fit.boundary_hit,
fit.n_evals,
)
};
assert_eq!(
cold_boundary, 0,
"fixture must have an interior minimum (cold boundary_hit = 0)"
);
let scratch = build_lme_scratch(&mut ws, n as u32, k as u32);
let warm = lme_fit(x.as_ref(), &y, &cluster_ids, &targets, Some(4.0), scratch);
assert!(warm.converged, "warm fit must converge");
assert_eq!(
warm.boundary_hit, cold_boundary,
"boundary flags must agree"
);
for (j, (&bw, &bc)) in warm.betas.iter().zip(&cold_betas).enumerate() {
assert!(
(bw - bc).abs() < 1e-5,
"β̂[{j}]: warm {bw} vs cold {bc} exceeds 1e-5"
);
}
for &tj in &targets {
let (tw, tc) = (warm.t_sq[tj as usize], cold_t_sq[tj as usize]);
assert!(
(tw.sqrt() - tc.sqrt()).abs() < 1e-4,
"z[{tj}]: warm {} vs cold {} exceeds 1e-4",
tw.sqrt(),
tc.sqrt()
);
}
assert!(
warm.n_evals < cold_evals,
"warm bracket must save deviance evals: warm {} vs cold {}",
warm.n_evals,
cold_evals
);
}
#[test]
fn lme_truth_start_zero_tau_keeps_boundary_semantics() {
let (x, y, cluster_ids) = make_lme_fixture_p3(0x5EED_CAFE);
let n = y.len();
let p = x.ncols();
let k = *cluster_ids.iter().max().unwrap() as usize + 1;
let targets: Vec<u32> = vec![1, 2];
let mut ws = TestWs::new(n, p, k);
ws.reset_lme_suff_stats();
{
let mut s = LmeSuffStats {
xtx: ws.lme_xtx.as_mut(),
xty: &mut ws.lme_xty,
yty: &mut ws.lme_yty,
sum_xc: ws.lme_sum_xc.as_mut(),
sum_yc: &mut ws.lme_sum_yc,
cluster_sizes: &mut ws.lme_cluster_sizes,
n_clusters_seen: &mut ws.lme_n_clusters_seen,
panel_x: &mut ws.panel_x,
panel_y: &mut ws.panel_y,
};
s.add_rows(x.as_ref(), &y, &cluster_ids);
}
let (cold_betas, cold_boundary, cold_converged) = {
let scratch = build_lme_scratch(&mut ws, n as u32, k as u32);
let fit = lme_fit(x.as_ref(), &y, &cluster_ids, &targets, None, scratch);
(fit.betas.to_vec(), fit.boundary_hit, fit.converged)
};
let scratch = build_lme_scratch(&mut ws, n as u32, k as u32);
let warm = lme_fit(x.as_ref(), &y, &cluster_ids, &targets, Some(0.0), scratch);
assert!(cold_converged && warm.converged, "both paths must converge");
assert_eq!(cold_boundary, 1, "τ=0 DGP must hit the τ̂≈0 boundary cold");
assert_eq!(warm.boundary_hit, 1, "θ₀ = 0 must keep boundary_hit = 1");
for (j, (&bw, &bc)) in warm.betas.iter().zip(&cold_betas).enumerate() {
assert_eq!(
bw.to_bits(),
bc.to_bits(),
"β̂[{j}]: warm {bw} != cold {bc} (both pin at LOG_THETA_LOW)"
);
}
}
#[test]
fn lme_fit_handles_rank_deficient_x() {
use faer::Mat;
const N: usize = 30;
const P: usize = 3;
const K: usize = 3;
let mut x = Mat::<f64>::zeros(N, P);
for i in 0..N {
x[(i, 0)] = 1.0;
x[(i, 1)] = 1.0;
x[(i, 2)] = ((i as f64) * 0.1) - 1.0;
}
let y: Vec<f64> = (0..N)
.map(|i| (i as f64) * 0.13 - 0.4 + ((i % 3) as f64) * 0.5)
.collect();
let cluster_ids: Vec<u32> = (0..N).map(|i| (i % K) as u32).collect();
let mut ws = TestWs::new(N, P, K);
ws.reset_lme_suff_stats();
{
let mut s = LmeSuffStats {
xtx: ws.lme_xtx.as_mut(),
xty: &mut ws.lme_xty,
yty: &mut ws.lme_yty,
sum_xc: ws.lme_sum_xc.as_mut(),
sum_yc: &mut ws.lme_sum_yc,
cluster_sizes: &mut ws.lme_cluster_sizes,
n_clusters_seen: &mut ws.lme_n_clusters_seen,
panel_x: &mut ws.panel_x,
panel_y: &mut ws.panel_y,
};
s.add_rows(x.as_ref(), &y, &cluster_ids);
}
let targets: Vec<u32> = vec![1, 2];
let fit_converged: bool;
let fit_betas: [f64; P];
let fit_var_diag: [f64; 2];
let fit_t_sq: [f64; 2];
{
let scratch = build_lme_scratch(&mut ws, N as u32, K as u32);
let fit = lme_fit(x.as_ref(), &y, &cluster_ids, &targets, None, scratch);
fit_converged = fit.converged;
fit_betas = [fit.betas[0], fit.betas[1], fit.betas[2]];
fit_var_diag = [fit.var_diag[1], fit.var_diag[2]];
fit_t_sq = [fit.t_sq[1], fit.t_sq[2]];
}
assert!(!fit_converged, "expected non-converged on singular X");
for (j, &b) in fit_betas.iter().enumerate() {
assert!(b.is_nan(), "β̂[{j}] = {b} should be NaN");
}
for v in &fit_var_diag {
assert!(v.is_nan(), "var_diag = {v} should be NaN");
}
for v in &fit_t_sq {
assert!(v.is_nan(), "t_sq = {v} should be NaN");
}
}
#[test]
fn lme_fit_bracket_repair_when_min_at_left_edge() {
use faer::Mat;
const N: usize = 30;
const P: usize = 2;
const K: usize = 3;
let mut x = Mat::<f64>::zeros(N, P);
for i in 0..N {
x[(i, 0)] = 1.0;
x[(i, 1)] = ((i as f64) * 0.07).sin();
}
let y: Vec<f64> = (0..N)
.map(|i| 0.5 + 0.4 * x[(i, 1)] + ((i as f64) * 0.011).cos() * 0.05)
.collect();
let cluster_ids: Vec<u32> = (0..N).map(|i| (i % K) as u32).collect();
let mut xtx_ols = Mat::<f64>::zeros(P, P);
let mut xty_ols = [0.0_f64; P];
for i in 0..N {
for jj in 0..P {
xty_ols[jj] += x[(i, jj)] * y[i];
for ii in jj..P {
xtx_ols[(ii, jj)] += x[(i, ii)] * x[(i, jj)];
}
}
}
let chol = xtx_ols
.as_ref()
.llt(faer::Side::Lower)
.expect("OLS Cholesky failed");
let mut rhs = Mat::<f64>::zeros(P, 1);
for jj in 0..P {
rhs[(jj, 0)] = xty_ols[jj];
}
use faer::linalg::solvers::Solve;
chol.solve_in_place(rhs.as_mut());
let ols_beta = [rhs[(0, 0)], rhs[(1, 0)]];
let mut ws = TestWs::new(N, P, K);
ws.reset_lme_suff_stats();
{
let mut s = LmeSuffStats {
xtx: ws.lme_xtx.as_mut(),
xty: &mut ws.lme_xty,
yty: &mut ws.lme_yty,
sum_xc: ws.lme_sum_xc.as_mut(),
sum_yc: &mut ws.lme_sum_yc,
cluster_sizes: &mut ws.lme_cluster_sizes,
n_clusters_seen: &mut ws.lme_n_clusters_seen,
panel_x: &mut ws.panel_x,
panel_y: &mut ws.panel_y,
};
s.add_rows(x.as_ref(), &y, &cluster_ids);
}
let targets: Vec<u32> = vec![1];
let fit_converged: bool;
let fit_boundary: u8;
let fit_betas: [f64; P];
{
let scratch = build_lme_scratch(&mut ws, N as u32, K as u32);
let fit = lme_fit(x.as_ref(), &y, &cluster_ids, &targets, None, scratch);
fit_converged = fit.converged;
fit_boundary = fit.boundary_hit;
fit_betas = [fit.betas[0], fit.betas[1]];
}
assert!(fit_converged, "bracket-repair path failed to converge");
assert_eq!(
fit_boundary, 1,
"expected τ̂≈0 boundary_hit=1, got {fit_boundary}"
);
for j in 0..P {
let delta = (fit_betas[j] - ols_beta[j]).abs();
assert!(
delta < 1e-6,
"β̂[{j}] = {}, OLS = {}, delta = {delta}",
fit_betas[j],
ols_beta[j]
);
}
}
#[cfg(feature = "alloc-tests")]
#[test]
#[ignore]
fn lme_fit_warm_path_bounded_alloc() {
let _serial = crate::test_support::alloc_test_guard();
use faer::Mat;
const N: usize = 60;
const P: usize = 3;
const K: usize = 6;
const N_CALLS: usize = 100;
const BOUND: u64 = 1800;
let x1: [f64; 60] = [
0.30471707975443135,
-1.0399841062404955,
0.7504511958064572,
0.9405647163912139,
-1.9510351886538364,
-1.302179506862318,
0.12784040316728537,
-0.3162425923435822,
-0.016801157504288795,
-0.85304392757358,
0.8793979748628286,
0.7777919354289483,
0.06603069756121605,
1.1272412069680329,
0.4675093422520456,
-0.8592924628832382,
0.36875078408249884,
-0.9588826008289989,
0.8784503013072725,
-0.049925910986252896,
-0.18486236354526056,
-0.6809295444039414,
1.2225413386740303,
-0.15452948206880215,
-0.4283278221631072,
-0.3521335504882296,
0.5323091855533487,
0.36544406436407834,
0.4127326115959884,
0.43082100300788273,
2.1416476008704612,
-0.4064150163846156,
-0.5122427290715373,
-0.8137727282478777,
0.6159794225754956,
1.1289722927208916,
-0.11394745765487507,
-0.840156476962528,
-0.8244812156912396,
0.6505927878247011,
0.7432541712034423,
0.543154268305195,
-0.6655097072886943,
0.23216132306671977,
0.11668580914072822,
0.21868859672901295,
0.8714287779481898,
0.22359554877468227,
0.6789135630718949,
0.06757906948889146,
0.28911939868998415,
0.6312882258385404,
-1.4571558198556664,
-0.31967121635730134,
-0.4703726542927955,
-0.6388778482433419,
-0.27514225122668373,
1.4949413112343959,
-0.8658311156932432,
0.9682783545914808,
];
let x2: [f64; 60] = [
-1.6828697716158048,
-0.33488502998577485,
0.1627530651050056,
0.5862223313592781,
0.711226579792855,
0.7933472351999252,
-0.3487250722484376,
-0.46235179266456716,
0.8579758812571538,
-0.1913043248816149,
-1.2756863233379219,
-1.1332872140034806,
-0.9194522860016113,
0.49716074405376404,
0.14242573607056525,
0.6904853540677682,
-0.42725264633653426,
0.15853969107671423,
0.6255903939673367,
-0.3093465397202384,
0.45677523755741145,
-0.6619259410666513,
-0.3630538465650718,
-0.3817378939983291,
-1.1958396455890397,
0.4869724807855818,
-0.46940234020272387,
0.01249411872768743,
0.48074665890590895,
0.4465311760299441,
0.6653851089727862,
-0.09848548450942361,
-0.42329831204415375,
-0.07971821090639905,
-1.6873344339580298,
-1.4471124724230873,
-1.3226996123544024,
-0.9972468276014818,
0.3997742267234366,
-0.9054790553600608,
-0.3781625540393897,
1.2992282977860654,
-0.35626397106142593,
0.7375155684670865,
-0.933617680009877,
-0.20543755786763002,
-0.9500220549105812,
-0.3390330759005625,
0.8403081374573955,
-1.7273204231923487,
0.43442364354585733,
0.2377356023322779,
-0.5941499556967944,
-1.4460578543884546,
0.07212950771386951,
-0.5294927090638024,
0.23267621135470395,
0.02185214552344288,
1.6017788913209154,
-0.23935562747302427,
];
let y_data: [f64; 60] = [
0.4634838483807412,
-1.6958367493518591,
-0.07656189550801018,
0.5791988172301176,
-0.623788696202619,
-1.5640666548456659,
-0.05286302450777913,
-0.40566566398296267,
-1.082675989070649,
-0.6398666414265639,
0.8986776652224437,
1.1127414469896355,
1.1065939034406163,
3.740015159794447,
0.830029945030196,
-0.3524130442209007,
-1.5476759552488093,
0.06153714483417494,
-0.8333554776977836,
-0.4197403105019113,
-0.42835728975820775,
0.1797922600939008,
2.0965054734973476,
0.19096753379326917,
-1.3788672930282488,
-0.9522449167472798,
-1.4102109341814926,
0.5597880597592666,
0.824978479993657,
2.343004843474632,
0.7533972033686575,
0.8633702245648947,
-0.7432683368811597,
-0.7561198927806311,
-0.6351566537898222,
-0.3691685843204092,
1.0273694686701114,
-1.4272550988991226,
0.6853189837147633,
0.010351200249558046,
1.7179220596063591,
1.2720016651761046,
-1.2435885938521603,
0.4099566730250238,
-0.7373919750549822,
1.35341093351855,
0.2238950195023513,
-0.9898128365682515,
-1.333705002777748,
-0.19843480740416738,
1.9931161686261913,
1.406582108224644,
-0.4972904954045434,
-0.08212174931081978,
0.445137615643293,
-0.14552036838998333,
-0.33955737323366064,
2.7199407103663717,
1.0045452753534772,
2.1737935566316495,
];
let mut x = Mat::<f64>::zeros(N, P);
for i in 0..N {
x[(i, 0)] = 1.0;
x[(i, 1)] = x1[i];
x[(i, 2)] = x2[i];
}
let y: Vec<f64> = y_data.to_vec();
let cluster_ids: Vec<u32> = (0..N).map(|i| (i % K) as u32).collect();
let targets: Vec<u32> = vec![1, 2];
let mut ws = TestWs::new(N, P, K);
ws.reset_lme_suff_stats();
{
let mut s = LmeSuffStats {
xtx: ws.lme_xtx.as_mut(),
xty: &mut ws.lme_xty,
yty: &mut ws.lme_yty,
sum_xc: ws.lme_sum_xc.as_mut(),
sum_yc: &mut ws.lme_sum_yc,
cluster_sizes: &mut ws.lme_cluster_sizes,
n_clusters_seen: &mut ws.lme_n_clusters_seen,
panel_x: &mut ws.panel_x,
panel_y: &mut ws.panel_y,
};
s.add_rows(x.as_ref(), &y, &cluster_ids);
}
{
let scratch = build_lme_scratch(&mut ws, N as u32, K as u32);
let _ = lme_fit(x.as_ref(), &y, &cluster_ids, &targets, None, scratch);
}
let profiler = dhat::Profiler::builder().testing().build();
for _ in 0..N_CALLS {
ws.reset_lme_suff_stats();
{
let mut s = LmeSuffStats {
xtx: ws.lme_xtx.as_mut(),
xty: &mut ws.lme_xty,
yty: &mut ws.lme_yty,
sum_xc: ws.lme_sum_xc.as_mut(),
sum_yc: &mut ws.lme_sum_yc,
cluster_sizes: &mut ws.lme_cluster_sizes,
n_clusters_seen: &mut ws.lme_n_clusters_seen,
panel_x: &mut ws.panel_x,
panel_y: &mut ws.panel_y,
};
s.add_rows(x.as_ref(), &y, &cluster_ids);
}
let scratch = build_lme_scratch(&mut ws, N as u32, K as u32);
let _ = lme_fit(x.as_ref(), &y, &cluster_ids, &targets, None, scratch);
}
let stats = dhat::HeapStats::get();
drop(profiler);
assert!(
stats.total_blocks <= BOUND,
"lme_fit allocated {} blocks across {} warm-path calls (BOUND = {})",
stats.total_blocks,
N_CALLS,
BOUND
);
}
fn make_lme_fixture_p3(seed: u64) -> (faer::Mat<f64>, Vec<f64>, Vec<u32>) {
use faer::Mat;
const N: usize = 15;
const P: usize = 3;
const K: usize = 3;
let mut x = Mat::<f64>::zeros(N, P);
let mut y = vec![0.0_f64; N];
let mut cluster_ids = vec![0_u32; N];
let mut state = seed.wrapping_mul(0x9E3779B97F4A7C15);
for i in 0..N {
x[(i, 0)] = 1.0;
state = state.wrapping_mul(0x9E3779B97F4A7C15).wrapping_add(1);
let x1 = ((state >> 11) as f64 / (1u64 << 53) as f64) * 2.0 - 1.0;
x[(i, 1)] = x1;
state = state.wrapping_mul(0x9E3779B97F4A7C15).wrapping_add(1);
let x2 = ((state >> 11) as f64 / (1u64 << 53) as f64) * 2.0 - 1.0;
x[(i, 2)] = x2;
state = state.wrapping_mul(0x9E3779B97F4A7C15).wrapping_add(1);
let noise = ((state >> 11) as f64 / (1u64 << 53) as f64) * 2.0 - 1.0;
y[i] = 0.5 + 0.4 * x1 - 0.3 * x2 + 0.3 * noise;
cluster_ids[i] = (i % K) as u32;
}
(x, y, cluster_ids)
}
#[test]
fn joint_wald_collapses_to_wald_z_sq_when_k_eq_1() {
let (x, y, cluster_ids) = make_lme_fixture_p3(0x1234_5678);
let target_indices = [1u32]; let n = y.len();
let p = x.ncols();
let n_clusters = *cluster_ids.iter().max().unwrap() as usize + 1;
let mut ws = TestWs::new(n, p, n_clusters);
ws.reset_lme_suff_stats();
{
let mut s = LmeSuffStats {
xtx: ws.lme_xtx.as_mut(),
xty: &mut ws.lme_xty,
yty: &mut ws.lme_yty,
sum_xc: ws.lme_sum_xc.as_mut(),
sum_yc: &mut ws.lme_sum_yc,
cluster_sizes: &mut ws.lme_cluster_sizes,
n_clusters_seen: &mut ws.lme_n_clusters_seen,
panel_x: &mut ws.panel_x,
panel_y: &mut ws.panel_y,
};
s.add_rows(x.as_ref(), &y, &cluster_ids);
}
let scratch = build_lme_scratch(&mut ws, n as u32, n_clusters as u32);
let fit = lme_fit(x.as_ref(), &y, &cluster_ids, &target_indices, None, scratch);
assert!(fit.converged, "fixture failed to converge");
assert!(fit.joint_t_sq.is_finite(), "joint_t_sq is NaN/inf");
let wald_z_sq = fit.t_sq[target_indices[0] as usize];
let diff = (fit.joint_t_sq - wald_z_sq).abs();
assert!(
diff < 1e-12,
"k=1 joint must equal Wald-z² within 1e-12: joint={}, wald={}, diff={}",
fit.joint_t_sq,
wald_z_sq,
diff
);
}
#[test]
fn joint_wald_ols_fallback_returns_finite_chi_sq() {
use faer::Mat;
const N: usize = 30;
const P: usize = 3;
const K: usize = 3;
let mut x = Mat::<f64>::zeros(N, P);
for i in 0..N {
x[(i, 0)] = 1.0;
x[(i, 1)] = ((i as f64) * 0.13).sin();
x[(i, 2)] = ((i as f64) * 0.07).cos();
}
let y: Vec<f64> = (0..N)
.map(|i| 0.5 + 0.4 * x[(i, 1)] - 0.3 * x[(i, 2)] + ((i as f64) * 0.011).sin() * 0.05)
.collect();
let cluster_ids: Vec<u32> = (0..N).map(|i| (i % K) as u32).collect();
let target_indices = [1u32, 2u32];
let mut ws = TestWs::new(N, P, K);
ws.reset_lme_suff_stats();
{
let mut s = LmeSuffStats {
xtx: ws.lme_xtx.as_mut(),
xty: &mut ws.lme_xty,
yty: &mut ws.lme_yty,
sum_xc: ws.lme_sum_xc.as_mut(),
sum_yc: &mut ws.lme_sum_yc,
cluster_sizes: &mut ws.lme_cluster_sizes,
n_clusters_seen: &mut ws.lme_n_clusters_seen,
panel_x: &mut ws.panel_x,
panel_y: &mut ws.panel_y,
};
s.add_rows(x.as_ref(), &y, &cluster_ids);
}
let scratch = build_lme_scratch(&mut ws, N as u32, K as u32);
let fit = lme_fit(x.as_ref(), &y, &cluster_ids, &target_indices, None, scratch);
assert!(fit.converged);
assert_eq!(
fit.boundary_hit, 1,
"expected boundary_hit=1 on zero-ICC fixture, got {}",
fit.boundary_hit
);
assert!(
fit.joint_t_sq.is_finite(),
"OLS-fallback joint must be finite, got {}",
fit.joint_t_sq
);
assert!(fit.joint_t_sq >= 0.0);
}
#[test]
fn lme_fit_golden_sigma_sq_and_betas() {
use faer::Mat;
const N: usize = 60;
const P: usize = 3;
const K: usize = 6;
let x1: [f64; 60] = [
0.30471707975443135,
-1.0399841062404955,
0.7504511958064572,
0.9405647163912139,
-1.9510351886538364,
-1.302179506862318,
0.12784040316728537,
-0.3162425923435822,
-0.016801157504288795,
-0.85304392757358,
0.8793979748628286,
0.7777919354289483,
0.06603069756121605,
1.1272412069680329,
0.4675093422520456,
-0.8592924628832382,
0.36875078408249884,
-0.9588826008289989,
0.8784503013072725,
-0.049925910986252896,
-0.18486236354526056,
-0.6809295444039414,
1.2225413386740303,
-0.15452948206880215,
-0.4283278221631072,
-0.3521335504882296,
0.5323091855533487,
0.36544406436407834,
0.4127326115959884,
0.43082100300788273,
2.1416476008704612,
-0.4064150163846156,
-0.5122427290715373,
-0.8137727282478777,
0.6159794225754956,
1.1289722927208916,
-0.11394745765487507,
-0.840156476962528,
-0.8244812156912396,
0.6505927878247011,
0.7432541712034423,
0.543154268305195,
-0.6655097072886943,
0.23216132306671977,
0.11668580914072822,
0.21868859672901295,
0.8714287779481898,
0.22359554877468227,
0.6789135630718949,
0.06757906948889146,
0.28911939868998415,
0.6312882258385404,
-1.4571558198556664,
-0.31967121635730134,
-0.4703726542927955,
-0.6388778482433419,
-0.27514225122668373,
1.4949413112343959,
-0.8658311156932432,
0.9682783545914808,
];
let x2: [f64; 60] = [
-1.6828697716158048,
-0.33488502998577485,
0.1627530651050056,
0.5862223313592781,
0.711226579792855,
0.7933472351999252,
-0.3487250722484376,
-0.46235179266456716,
0.8579758812571538,
-0.1913043248816149,
-1.2756863233379219,
-1.1332872140034806,
-0.9194522860016113,
0.49716074405376404,
0.14242573607056525,
0.6904853540677682,
-0.42725264633653426,
0.15853969107671423,
0.6255903939673367,
-0.3093465397202384,
0.45677523755741145,
-0.6619259410666513,
-0.3630538465650718,
-0.3817378939983291,
-1.1958396455890397,
0.4869724807855818,
-0.46940234020272387,
0.01249411872768743,
0.48074665890590895,
0.4465311760299441,
0.6653851089727862,
-0.09848548450942361,
-0.42329831204415375,
-0.07971821090639905,
-1.6873344339580298,
-1.4471124724230873,
-1.3226996123544024,
-0.9972468276014818,
0.3997742267234366,
-0.9054790553600608,
-0.3781625540393897,
1.2992282977860654,
-0.35626397106142593,
0.7375155684670865,
-0.933617680009877,
-0.20543755786763002,
-0.9500220549105812,
-0.3390330759005625,
0.8403081374573955,
-1.7273204231923487,
0.43442364354585733,
0.2377356023322779,
-0.5941499556967944,
-1.4460578543884546,
0.07212950771386951,
-0.5294927090638024,
0.23267621135470395,
0.02185214552344288,
1.6017788913209154,
-0.23935562747302427,
];
let y_data: [f64; 60] = [
0.4634838483807412,
-1.6958367493518591,
-0.07656189550801018,
0.5791988172301176,
-0.623788696202619,
-1.5640666548456659,
-0.05286302450777913,
-0.40566566398296267,
-1.082675989070649,
-0.6398666414265639,
0.8986776652224437,
1.1127414469896355,
1.1065939034406163,
3.740015159794447,
0.830029945030196,
-0.3524130442209007,
-1.5476759552488093,
0.06153714483417494,
-0.8333554776977836,
-0.4197403105019113,
-0.42835728975820775,
0.1797922600939008,
2.0965054734973476,
0.19096753379326917,
-1.3788672930282488,
-0.9522449167472798,
-1.4102109341814926,
0.5597880597592666,
0.824978479993657,
2.343004843474632,
0.7533972033686575,
0.8633702245648947,
-0.7432683368811597,
-0.7561198927806311,
-0.6351566537898222,
-0.3691685843204092,
1.0273694686701114,
-1.4272550988991226,
0.6853189837147633,
0.010351200249558046,
1.7179220596063591,
1.2720016651761046,
-1.2435885938521603,
0.4099566730250238,
-0.7373919750549822,
1.35341093351855,
0.2238950195023513,
-0.9898128365682515,
-1.333705002777748,
-0.19843480740416738,
1.9931161686261913,
1.406582108224644,
-0.4972904954045434,
-0.08212174931081978,
0.445137615643293,
-0.14552036838998333,
-0.33955737323366064,
2.7199407103663717,
1.0045452753534772,
2.1737935566316495,
];
let mut x = Mat::<f64>::zeros(N, P);
for i in 0..N {
x[(i, 0)] = 1.0;
x[(i, 1)] = x1[i];
x[(i, 2)] = x2[i];
}
let y: Vec<f64> = y_data.to_vec();
let cluster_ids: Vec<u32> = (0..N).map(|i| (i / 10) as u32).collect();
let target_indices: Vec<u32> = vec![1, 2];
let mut ws = TestWs::new(N, P, K);
ws.reset_lme_suff_stats();
{
let mut s = LmeSuffStats {
xtx: ws.lme_xtx.as_mut(),
xty: &mut ws.lme_xty,
yty: &mut ws.lme_yty,
sum_xc: ws.lme_sum_xc.as_mut(),
sum_yc: &mut ws.lme_sum_yc,
cluster_sizes: &mut ws.lme_cluster_sizes,
n_clusters_seen: &mut ws.lme_n_clusters_seen,
panel_x: &mut ws.panel_x,
panel_y: &mut ws.panel_y,
};
s.add_rows(x.as_ref(), &y, &cluster_ids);
}
let scratch = build_lme_scratch(&mut ws, N as u32, K as u32);
let fit = lme_fit(x.as_ref(), &y, &cluster_ids, &target_indices, None, scratch);
assert!(
fit.converged,
"must converge on the canonical 60-row fixture"
);
assert_eq!(
fit.boundary_hit, 0,
"expected interior optimum on the blocked fixture, got boundary_hit={}",
fit.boundary_hit
);
let r_sigma_sq = 0.9643488_f64;
let r_beta_x1 = 0.7513689_f64;
let r_tau_sq = 0.0828566_f64;
let rel_sigma = (fit.sigma_sq - r_sigma_sq).abs() / r_sigma_sq;
assert!(
rel_sigma < 1e-3,
"σ² = {}, R/lme4 = {r_sigma_sq}, rel = {rel_sigma}",
fit.sigma_sq
);
let rel_beta = (fit.betas[1] - r_beta_x1).abs() / r_beta_x1.abs().max(1e-9);
assert!(
rel_beta < 1e-3,
"β̂[x1] = {}, R/lme4 = {r_beta_x1}, rel = {rel_beta}",
fit.betas[1]
);
let rel_tau = (fit.tau_sq_hat - r_tau_sq).abs() / r_tau_sq;
assert!(
rel_tau < 1e-3,
"τ² = {}, R/lme4 = {r_tau_sq}, rel = {rel_tau}",
fit.tau_sq_hat
);
}
#[test]
fn profiled_deviance_value_and_curvature_at_theta_1() {
use faer::Mat;
let n: usize = 6;
let p: usize = 2;
let n_clusters: usize = 2;
let cluster_ids: Vec<u32> = vec![0, 0, 0, 1, 1, 1];
let x1: [f64; 6] = [
0.1257302210933933,
-0.1321048632913019,
0.6404226504432821,
0.10490011715303971,
-0.535669373161111,
0.36159505490948474,
];
let y: [f64; 6] = [
0.7718630718197979,
-0.09922526468089643,
1.4699547603093044,
1.266192714345762,
-1.8688474280172964,
1.3141089963737067,
];
let mut x = Mat::<f64>::zeros(n, p);
for i in 0..n {
x[(i, 0)] = 1.0;
x[(i, 1)] = x1[i];
}
let mut ws = TestWs::new(n, p, n_clusters);
ws.reset_lme_suff_stats();
{
let mut s = LmeSuffStats {
xtx: ws.lme_xtx.as_mut(),
xty: &mut ws.lme_xty,
yty: &mut ws.lme_yty,
sum_xc: ws.lme_sum_xc.as_mut(),
sum_yc: &mut ws.lme_sum_yc,
cluster_sizes: &mut ws.lme_cluster_sizes,
n_clusters_seen: &mut ws.lme_n_clusters_seen,
panel_x: &mut ws.panel_x,
panel_y: &mut ws.panel_y,
};
s.add_rows(x.as_ref(), &y, &cluster_ids);
}
let mut scratch = build_lme_scratch(&mut ws, n as u32, n_clusters as u32);
const DEVFUN_LME4_AT_THETA_1: f64 = 9.410566478412274;
let n_minus_p = (n - p) as f64; let reml_const = n_minus_p * (1.0 + (2.0 * std::f64::consts::PI).ln());
let expected_dev_at_theta_1 = DEVFUN_LME4_AT_THETA_1 - reml_const; let dev_at_theta_1 = profiled_deviance(1.0, &mut scratch);
assert!(
dev_at_theta_1.is_finite(),
"profiled_deviance(θ=1) must be finite, got {dev_at_theta_1}"
);
assert!(
(dev_at_theta_1 - expected_dev_at_theta_1).abs() < 1e-4,
"profiled_deviance(θ=1) = {dev_at_theta_1}, lme4-derived expected {expected_dev_at_theta_1}"
);
let dev_at_small_theta = profiled_deviance(1e-4, &mut scratch);
assert!(
dev_at_theta_1 > dev_at_small_theta,
"deviance at θ=1 ({dev_at_theta_1}) must exceed deviance at θ→0 ({dev_at_small_theta})"
);
}
}