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::FLOAT_NEAR_ZERO;
pub const PANEL_ROWS: usize = 256;
pub struct OlsScratch<'w> {
pub fit_betas: &'w mut [f64],
pub fit_var_diag: &'w mut [f64],
pub fit_t_sq: &'w mut [f64],
pub fit_u_scratch: &'w mut [f64],
pub fit_factor: MatMut<'w, f64>,
pub fit_rhs: MatMut<'w, f64>,
}
pub struct OlsFitView<'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 df_resid: u32,
pub converged: bool,
pub rss: f64,
pub sst: f64,
pub pivot: f64,
pub pivot_col: u32,
}
pub struct OlsSuffStats<'w> {
pub xtx: MatMut<'w, f64>,
pub xty: &'w mut [f64],
pub yty: &'w mut f64,
pub sum_y: &'w mut f64,
pub n_rows: &'w mut usize,
pub panel_x: &'w mut [f64],
pub panel_y: &'w mut [f64],
}
impl<'w> OlsSuffStats<'w> {
pub fn add_rows(&mut self, x_block: MatRef<'_, f64>, y_block: &[f64]) {
debug_assert_eq!(x_block.nrows(), y_block.len());
let p = self.xty.len();
debug_assert_eq!(x_block.ncols(), p);
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];
self.panel_y[i] = y_row;
*self.yty += y_row * y_row;
*self.sum_y += y_row;
}
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;
}
*self.n_rows += m;
}
}
#[inline]
pub(crate) fn triangular_solve_norm_sq(
factor: MatRef<'_, f64>,
b: impl Fn(usize) -> f64,
scratch: &mut [f64],
p: usize,
upper: bool,
) -> f64 {
for v in &mut scratch[..p] {
*v = 0.0;
}
for i in 0..p {
let mut acc = b(i);
for k in 0..i {
acc -= if upper {
factor[(k, i)]
} else {
factor[(i, k)]
} * scratch[k];
}
let diag = factor[(i, i)];
if diag.abs() < FLOAT_NEAR_ZERO {
return f64::NAN;
}
scratch[i] = acc / diag;
}
let mut norm_sq = 0.0;
for &v in &scratch[..p] {
norm_sq += v * v;
}
norm_sq
}
#[inline]
pub(crate) fn nan_fill_ols_scratch(
betas: &mut [f64],
var_diag: &mut [f64],
t_sq: &mut [f64],
p: usize,
t: usize,
) {
betas[..p].fill(f64::NAN);
var_diag[..t].fill(f64::NAN);
t_sq[..t].fill(f64::NAN);
}
#[inline]
fn nonconverged_view<'a>(
betas: &'a [f64],
var_diag: &'a [f64],
t_sq: &'a [f64],
factor: MatRef<'a, f64>,
df_resid: u32,
) -> OlsFitView<'a> {
OlsFitView {
betas,
var_diag,
t_sq,
factor,
sigma_sq: f64::NAN,
df_resid,
converged: false,
rss: f64::NAN,
sst: f64::NAN,
pivot: f64::NAN,
pivot_col: 0,
}
}
pub(crate) const ALIAS_EPS: f64 = 1e-14;
pub(crate) fn aliased_columns(gram: MatRef<'_, f64>, p: usize, eps: f64) -> Vec<bool> {
let mut l = vec![0.0f64; p * p]; let mut aliased = vec![false; p];
for d in 0..p {
let g_dd = gram[(d, d)];
let mut piv = g_dd;
for j in 0..d {
if !aliased[j] {
piv -= l[d * p + j] * l[d * p + j];
}
}
if piv <= eps * g_dd {
aliased[d] = true; continue;
}
let ljj = piv.sqrt();
l[d * p + d] = ljj;
for i in (d + 1)..p {
let mut s = gram[(i, d)]; for j in 0..d {
if !aliased[j] {
s -= l[i * p + j] * l[d * p + j];
}
}
l[i * p + d] = s / ljj;
}
}
aliased
}
pub(crate) const PIVOT_MIN: f64 = 1e-12;
pub(crate) fn min_pivot_ratio(factor: MatRef<'_, f64>, p: usize) -> (f64, usize) {
let mut min_ratio = f64::INFINITY;
let mut min_col = 0usize;
for d in 0..p {
let mut g_dd = 0.0;
for k in 0..=d {
let l_dk = factor[(d, k)];
g_dd += l_dk * l_dk;
}
let l_dd = factor[(d, d)];
let pivot = l_dd * l_dd;
if !g_dd.is_finite() || g_dd <= 0.0 || !pivot.is_finite() {
return (0.0, d);
}
let ratio = pivot / g_dd;
if ratio < min_ratio {
min_ratio = ratio;
min_col = d;
}
}
(min_ratio, min_col)
}
#[expect(
clippy::too_many_arguments,
reason = "sufficient-statistics kernel; each arg is a distinct precomputed input"
)]
pub fn fit_suff_stats_t_sq<'a>(
xtx_lower: MatRef<'_, f64>,
xty: &[f64],
yty: f64,
sum_y: f64,
n_rows: usize,
target_indices: &[u32],
mut xtx_work: MatMut<'_, f64>,
scratch: OlsScratch<'a>,
) -> OlsFitView<'a> {
let p = xty.len();
let t = target_indices.len();
let n = n_rows;
debug_assert_eq!(xtx_lower.nrows(), p);
debug_assert_eq!(xtx_lower.ncols(), p);
debug_assert_eq!(xtx_work.nrows(), p);
debug_assert_eq!(xtx_work.ncols(), p);
let OlsScratch {
fit_betas,
fit_var_diag,
fit_t_sq,
fit_u_scratch,
mut fit_factor,
mut fit_rhs,
} = scratch;
debug_assert!(p <= fit_betas.len(), "scratch sized for fewer predictors");
debug_assert!(t <= fit_var_diag.len());
debug_assert!(p <= fit_rhs.nrows(), "fit_rhs must hold at least p rows");
nan_fill_ols_scratch(fit_betas, fit_var_diag, fit_t_sq, p, t);
if n <= p || p == 0 {
return nonconverged_view(
&fit_betas[..p],
&fit_var_diag[..t],
&fit_t_sq[..t],
fit_factor.into_const(),
n.saturating_sub(p) as u32,
);
}
for j in 0..p {
for i in j..p {
xtx_work[(i, j)] = xtx_lower[(i, j)];
}
}
let chol = match xtx_work.rb().llt(faer::Side::Lower) {
Ok(c) => c,
Err(_) => {
return nonconverged_view(
&fit_betas[..p],
&fit_var_diag[..t],
&fit_t_sq[..t],
fit_factor.into_const(),
(n - p) as u32,
);
}
};
let l = chol.L();
let (pivot, pivot_col) = min_pivot_ratio(l, p);
for j in 0..p {
for i in 0..p {
fit_factor[(i, j)] = if i >= j { l[(i, j)] } else { 0.0 };
}
}
for j in 0..p {
fit_rhs[(j, 0)] = xty[j];
}
use faer::linalg::solvers::Solve;
chol.solve_in_place(fit_rhs.rb_mut().subrows_mut(0, p));
for j in 0..p {
fit_betas[j] = fit_rhs[(j, 0)];
}
let mut bty = 0.0;
for j in 0..p {
bty += fit_betas[j] * xty[j];
}
let rss = yty - bty;
let df_resid = (n - p) as u32;
let sigma_sq = rss / df_resid as f64;
let sst = yty - (sum_y * sum_y) / n as f64;
for (out_idx, &tj) in target_indices.iter().enumerate() {
let tj = tj as usize;
if tj >= p {
continue;
}
let norm_sq = triangular_solve_norm_sq(
fit_factor.rb(),
|i| if i == tj { 1.0 } else { 0.0 },
fit_u_scratch,
p,
false, );
let vd = sigma_sq * norm_sq;
fit_var_diag[out_idx] = vd;
if vd > FLOAT_NEAR_ZERO && vd.is_finite() {
let beta_j = fit_betas[tj];
fit_t_sq[out_idx] = (beta_j * beta_j) / vd;
} else {
fit_t_sq[out_idx] = f64::NAN;
}
}
OlsFitView {
betas: &fit_betas[..p],
var_diag: &fit_var_diag[..t],
t_sq: &fit_t_sq[..t],
factor: fit_factor.into_const(),
sigma_sq,
df_resid,
converged: true,
rss,
sst,
pivot,
pivot_col: pivot_col as u32,
}
}
pub fn ols_contrast_t_sq(fit: &OlsFitView<'_>, p_col: u32, n_col: u32, scratch: &mut [f64]) -> f64 {
if !fit.converged {
return f64::NAN;
}
let p = fit.betas.len();
let pc = p_col as usize;
let nc = n_col as usize;
if pc >= p || nc >= p || scratch.len() < p {
return f64::NAN;
}
let norm_sq = triangular_solve_norm_sq(
fit.factor,
|i| {
if i == pc {
1.0
} else if i == nc {
-1.0
} else {
0.0
}
},
scratch,
p,
false, );
let var = fit.sigma_sq * norm_sq;
if var <= FLOAT_NEAR_ZERO || !var.is_finite() {
return f64::NAN;
}
let beta_diff = fit.betas[pc] - fit.betas[nc];
(beta_diff * beta_diff) / var
}
#[cfg(test)]
mod tests {
use super::*;
use crate::test_support::TestWs;
use faer::Mat;
fn suff_stats(ws: &mut TestWs) -> OlsSuffStats<'_> {
OlsSuffStats {
xtx: ws.suff_xtx.as_mut(),
xty: &mut ws.suff_xty,
yty: &mut ws.suff_yty,
sum_y: &mut ws.suff_sum_y,
n_rows: &mut ws.suff_n_rows,
panel_x: &mut ws.panel_x,
panel_y: &mut ws.panel_y,
}
}
fn build_x(n: usize, p: usize, mut fill: impl FnMut(usize, usize) -> f64) -> Mat<f64> {
let mut m = Mat::<f64>::zeros(n, p);
for i in 0..n {
for j in 0..p {
m[(i, j)] = fill(i, j);
}
}
m
}
fn gram_of(x: &[f64], n: usize, p: usize) -> faer::Mat<f64> {
let mut g = faer::Mat::<f64>::zeros(p, p);
for i in 0..n {
for a in 0..p {
let xa = x[i * p + a];
for b in 0..=a {
g[(a, b)] += xa * x[i * p + b];
}
}
}
g
}
#[test]
fn aliased_columns_flags_dependent_last() {
let n = 4;
let p = 3;
let x = vec![
1.0, 0.0, 1.0, 1.0, 1.0, 2.0, 1.0, 2.0, 3.0, 1.0, 3.0, 4.0, ];
let g = gram_of(&x, n, p);
let a = super::aliased_columns(g.as_ref(), p, super::ALIAS_EPS);
assert_eq!(a, vec![false, false, true]);
}
#[test]
fn aliased_columns_full_rank_none() {
let n = 4;
let p = 2;
let x = vec![1.0, 0.0, 1.0, 1.0, 1.0, 2.0, 1.0, 3.0];
let g = gram_of(&x, n, p);
let a = super::aliased_columns(g.as_ref(), p, super::ALIAS_EPS);
assert_eq!(a, vec![false, false]);
}
#[test]
fn aliased_columns_drops_later_duplicate() {
let n = 3;
let p = 2;
let x = vec![1.0, 1.0, 2.0, 2.0, 3.0, 3.0];
let g = gram_of(&x, n, p);
let a = super::aliased_columns(g.as_ref(), p, super::ALIAS_EPS);
assert_eq!(a, vec![false, true]);
}
#[test]
fn suff_stats_batch_split_invariance() {
let n = 5;
let p = 3;
let x = build_x(n, p, |i, j| ((i + 1) as f64).powi(j as i32));
let y = [1.0_f64, 2.0, 3.0, 4.0, 5.0];
let mut ws_split = TestWs::new(n, p, 0);
ws_split.reset_suff_stats();
{
let mut s = suff_stats(&mut ws_split);
s.add_rows(x.as_ref().subrows(0, 2), &y[0..2]);
s.add_rows(x.as_ref().subrows(2, 3), &y[2..5]);
}
let mut ws_full = TestWs::new(n, p, 0);
ws_full.reset_suff_stats();
{
let mut s = suff_stats(&mut ws_full);
s.add_rows(x.as_ref(), &y);
}
for j in 0..p {
for i in j..p {
assert!(
(ws_split.suff_xtx[(i, j)] - ws_full.suff_xtx[(i, j)]).abs() < 1e-12,
"xtx[{i},{j}] split-add != full-add"
);
}
}
for k in 0..p {
assert!(
(ws_split.suff_xty[k] - ws_full.suff_xty[k]).abs() < 1e-12,
"xty[{k}] split != full"
);
}
assert!(
(ws_split.suff_yty - ws_full.suff_yty).abs() < 1e-12,
"yty split != full"
);
assert!(
(ws_split.suff_sum_y - ws_full.suff_sum_y).abs() < 1e-12,
"sum_y split != full"
);
assert_eq!(ws_split.suff_n_rows, ws_full.suff_n_rows);
assert_eq!(ws_split.suff_n_rows, n);
}
#[test]
fn add_rows_panel_matches_scalar_reference() {
let n = 611;
let p = 7;
let x = build_x(n, p, |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 mut ref_xtx = Mat::<f64>::zeros(p, p);
let mut ref_xty = vec![0.0f64; p];
let (mut ref_yty, mut ref_sum_y) = (0.0f64, 0.0f64);
for row in 0..n {
let y_row = y[row];
for j in 0..p {
let x_rj = x[(row, j)] as f64;
for i in j..p {
ref_xtx[(i, j)] += x[(row, i)] as f64 * x_rj;
}
ref_xty[j] += x_rj * y_row;
}
ref_yty += y_row * y_row;
ref_sum_y += y_row;
}
let mut ws = TestWs::new(n, p, 0);
ws.reset_suff_stats();
{
let mut s = suff_stats(&mut ws);
s.add_rows(x.as_ref(), &y);
}
for j in 0..p {
for i in j..p {
let (got, want) = (ws.suff_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}"
);
}
}
#[allow(clippy::needless_range_loop)]
for k in 0..p {
assert!(
(ws.suff_xty[k] - ref_xty[k]).abs() <= 1e-12 * ref_xty[k].abs().max(1.0),
"xty[{k}] = {}, scalar reference {}",
ws.suff_xty[k],
ref_xty[k]
);
}
assert_eq!(
ws.suff_yty.to_bits(),
ref_yty.to_bits(),
"yty must stay bit-identical (scalar row-order pass)"
);
assert_eq!(
ws.suff_sum_y.to_bits(),
ref_sum_y.to_bits(),
"sum_y must stay bit-identical"
);
assert_eq!(ws.suff_n_rows, n);
}
#[test]
fn fit_suff_stats_golden_values() {
let n = 4;
let p = 2;
let x = build_x(n, p, |i, j| if j == 0 { 1.0 } else { i as f64 });
let y = [1.0_f64, 3.0, 4.0, 8.0];
let targets: Vec<u32> = vec![0, 1];
let mut ws = TestWs::new(n, p, 0);
ws.reset_suff_stats();
{
let mut s = suff_stats(&mut ws);
s.add_rows(x.as_ref(), &y);
}
let scratch = OlsScratch {
fit_betas: &mut ws.fit_betas,
fit_var_diag: &mut ws.fit_var_diag,
fit_t_sq: &mut ws.fit_t_sq,
fit_u_scratch: &mut ws.fit_u_scratch,
fit_factor: ws.fit_factor.as_mut(),
fit_rhs: ws.fit_rhs.as_mut(),
};
let res = fit_suff_stats_t_sq(
ws.suff_xtx.as_ref(),
&ws.suff_xty,
ws.suff_yty,
ws.suff_sum_y,
ws.suff_n_rows,
&targets,
ws.suff_xtx_work.as_mut(),
scratch,
);
assert!(res.converged);
assert_eq!(res.df_resid, 2);
let golden_betas = [0.7, 2.2];
let golden_var = [0.63, 0.18];
let golden_t_sq = [7.0 / 9.0, 242.0 / 9.0];
for (j, (&got, &want)) in res.betas.iter().zip(golden_betas.iter()).enumerate() {
assert!((got - want).abs() < 1e-9, "β̂[{j}] = {got}, golden {want}");
}
for k in 0..targets.len() {
assert!(
(res.var_diag[k] - golden_var[k]).abs() < 1e-9,
"var_diag[{k}] = {}, golden {}",
res.var_diag[k],
golden_var[k]
);
assert!(
(res.t_sq[k] - golden_t_sq[k]).abs() < 1e-9,
"t²[{k}] = {}, golden {}",
res.t_sq[k],
golden_t_sq[k]
);
}
assert!(
(res.rss - 1.8).abs() < 1e-9,
"rss = {}, golden 1.8",
res.rss
);
assert!(
(res.sst - 26.0).abs() < 1e-9,
"sst = {}, golden 26.0",
res.sst
);
assert!(
(res.sigma_sq - 0.9).abs() < 1e-9,
"σ̂² = {}, golden 0.9",
res.sigma_sq
);
}
#[test]
fn suff_stats_non_converged_when_n_le_p() {
let p = 4;
let n = 3; let x = build_x(n, p, |i, j| ((i + j) as f64).sin() + 1.0);
let y: Vec<f64> = (0..n).map(|i| i as f64).collect();
let mut ws = TestWs::new(p, p, 0); ws.reset_suff_stats();
{
let mut s = suff_stats(&mut ws);
s.add_rows(x.as_ref(), &y);
}
let scratch = OlsScratch {
fit_betas: &mut ws.fit_betas,
fit_var_diag: &mut ws.fit_var_diag,
fit_t_sq: &mut ws.fit_t_sq,
fit_u_scratch: &mut ws.fit_u_scratch,
fit_factor: ws.fit_factor.as_mut(),
fit_rhs: ws.fit_rhs.as_mut(),
};
let res = fit_suff_stats_t_sq(
ws.suff_xtx.as_ref(),
&ws.suff_xty,
ws.suff_yty,
ws.suff_sum_y,
ws.suff_n_rows,
&[1, 2, 3],
ws.suff_xtx_work.as_mut(),
scratch,
);
assert!(!res.converged, "n ≤ p must not converge");
for v in res.t_sq.iter() {
assert!(v.is_nan(), "t² must be NaN when n ≤ p");
}
for v in res.betas.iter() {
assert!(v.is_nan(), "β̂ must be NaN when n ≤ p");
}
assert!(res.rss.is_nan(), "rss must be NaN when n ≤ p");
assert!(res.sst.is_nan(), "sst must be NaN when n ≤ p");
}
#[test]
fn ols_contrast_t_sq_is_symmetric() {
let p = 3;
let mut factor = Mat::<f64>::zeros(p, p);
factor[(0, 0)] = 2.0;
factor[(1, 0)] = 1.0;
factor[(1, 1)] = 3.0;
factor[(2, 1)] = 1.0;
factor[(2, 2)] = 4.0;
let betas = [0.5_f64, 1.2, -0.7];
let var_diag = [0.0_f64; 3];
let t_sq_dummy = [0.0_f64; 3];
let fit = OlsFitView {
betas: &betas,
var_diag: &var_diag,
t_sq: &t_sq_dummy,
factor: factor.as_ref(),
sigma_sq: 0.4,
df_resid: 10,
converged: true,
rss: 0.0,
sst: 0.0,
pivot: 1.0,
pivot_col: 0,
};
let mut scratch = vec![0.0_f64; p];
let forward = ols_contrast_t_sq(&fit, 1, 2, &mut scratch);
let reversed = ols_contrast_t_sq(&fit, 2, 1, &mut scratch);
assert!(
forward.is_finite() && forward > 0.0,
"contrast t² must be positive finite"
);
assert!(
(forward - reversed).abs() / forward.abs().max(1.0) < 1e-12,
"contrast t² must be symmetric under p/n swap: {forward} vs {reversed}"
);
}
#[test]
fn ols_contrast_t_sq_golden_value() {
let n = 4;
let p = 2;
let x = build_x(n, p, |i, j| if j == 0 { 1.0 } else { i as f64 });
let y = [1.0_f64, 3.0, 4.0, 8.0];
let mut ws = TestWs::new(n, p, 0);
ws.reset_suff_stats();
{
let mut s = suff_stats(&mut ws);
s.add_rows(x.as_ref(), &y);
}
let scratch = OlsScratch {
fit_betas: &mut ws.fit_betas,
fit_var_diag: &mut ws.fit_var_diag,
fit_t_sq: &mut ws.fit_t_sq,
fit_u_scratch: &mut ws.fit_u_scratch,
fit_factor: ws.fit_factor.as_mut(),
fit_rhs: ws.fit_rhs.as_mut(),
};
let res = fit_suff_stats_t_sq(
ws.suff_xtx.as_ref(),
&ws.suff_xty,
ws.suff_yty,
ws.suff_sum_y,
ws.suff_n_rows,
&[0, 1],
ws.suff_xtx_work.as_mut(),
scratch,
);
assert!(res.converged);
let mut cscratch = vec![0.0_f64; p];
let t_sq = ols_contrast_t_sq(&res, 1, 0, &mut cscratch);
assert!(
(t_sq - 5.0 / 3.0).abs() < 1e-9,
"contrast t² = {t_sq}, golden 5/3"
);
let rev = ols_contrast_t_sq(&res, 0, 1, &mut cscratch);
assert!(
(rev - 5.0 / 3.0).abs() < 1e-9,
"reversed t² = {rev}, golden 5/3"
);
}
#[test]
fn contrast_t_sq_returns_nan_on_non_converged() {
let p = 2;
let factor = Mat::<f64>::zeros(p, p);
let betas = [0.0_f64, 1.0];
let var_diag = [0.0_f64; 2];
let t_sq_dummy = [0.0_f64; 2];
let fit = OlsFitView {
betas: &betas,
var_diag: &var_diag,
t_sq: &t_sq_dummy,
factor: factor.as_ref(),
sigma_sq: 1.0,
df_resid: 10,
converged: false,
rss: f64::NAN,
sst: f64::NAN,
pivot: f64::NAN,
pivot_col: 0,
};
let mut scratch = vec![0.0_f64; p];
let got = ols_contrast_t_sq(&fit, 0, 1, &mut scratch);
assert!(got.is_nan(), "non-converged fit must return NaN, got={got}");
}
#[test]
fn suff_stats_rank_deficiency_detected() {
let n = 50;
let p = 3;
let x = build_x(n, p, |i, j| match j {
0 => 1.0,
1 => (i as f64) * 0.1,
_ => 0.0,
});
let y: Vec<f64> = (0..n).map(|i| (i as f64) * 0.3).collect();
let targets: Vec<u32> = vec![1, 2];
let mut ws_su = TestWs::new(n, p, 0);
ws_su.reset_suff_stats();
{
let mut s = suff_stats(&mut ws_su);
s.add_rows(x.as_ref(), &y);
}
let xtx_ref = ws_su.suff_xtx.as_ref();
let xty_ref = ws_su.suff_xty.clone();
let yty_val = ws_su.suff_yty;
let sum_y_val = ws_su.suff_sum_y;
let n_rows_val = ws_su.suff_n_rows;
let scratch = OlsScratch {
fit_betas: &mut ws_su.fit_betas,
fit_var_diag: &mut ws_su.fit_var_diag,
fit_t_sq: &mut ws_su.fit_t_sq,
fit_u_scratch: &mut ws_su.fit_u_scratch,
fit_factor: ws_su.fit_factor.as_mut(),
fit_rhs: ws_su.fit_rhs.as_mut(),
};
let res_su = fit_suff_stats_t_sq(
xtx_ref,
&xty_ref,
yty_val,
sum_y_val,
n_rows_val,
&targets,
ws_su.suff_xtx_work.as_mut(),
scratch,
);
assert!(
!res_su.converged,
"Cholesky must reject near-collinear design"
);
assert!(res_su.rss.is_nan(), "rss must be NaN on rank-deficient");
assert!(res_su.sst.is_nan(), "sst must be NaN on rank-deficient");
for v in res_su.t_sq.iter() {
assert!(v.is_nan(), "rank-deficient t_sq must be NaN");
}
for v in res_su.betas.iter() {
assert!(v.is_nan(), "rank-deficient betas must be NaN");
}
}
#[cfg(feature = "alloc-tests")]
#[test]
#[ignore]
fn fit_suff_stats_warm_path_bounded_alloc() {
let _serial = crate::test_support::alloc_test_guard();
const FAER_LLT_BLOCKS_PER_FIT: usize = 2;
const N_CALLS: usize = 100;
const ONE_TIME: usize = 0;
const BOUND: usize = FAER_LLT_BLOCKS_PER_FIT * N_CALLS + ONE_TIME;
let n = 200;
let p = 6;
let x = build_x(n, p, |i, j| {
if j == 0 {
1.0
} else {
((i * 7 + j * 13 + 5) % 23) as f64 / 5.0 - 2.0
}
});
let y: Vec<f64> = (0..n).map(|i| ((i * 11) % 13) as f64 / 10.0).collect();
let targets: Vec<u32> = vec![1, 2];
let mut ws = TestWs::new(n, p, 0);
let run_fit = |ws: &mut TestWs| {
ws.reset_suff_stats();
{
let mut s = suff_stats(ws);
s.add_rows(x.as_ref(), &y);
}
let scratch = OlsScratch {
fit_betas: &mut ws.fit_betas,
fit_var_diag: &mut ws.fit_var_diag,
fit_t_sq: &mut ws.fit_t_sq,
fit_u_scratch: &mut ws.fit_u_scratch,
fit_factor: ws.fit_factor.as_mut(),
fit_rhs: ws.fit_rhs.as_mut(),
};
let fit = fit_suff_stats_t_sq(
ws.suff_xtx.as_ref(),
&ws.suff_xty,
ws.suff_yty,
ws.suff_sum_y,
ws.suff_n_rows,
&targets,
ws.suff_xtx_work.as_mut(),
scratch,
);
assert!(fit.converged);
};
run_fit(&mut ws);
crate::test_support::settle_background_allocs();
let profiler = dhat::Profiler::builder().testing().build();
for _ in 0..N_CALLS {
run_fit(&mut ws);
}
let stats = dhat::HeapStats::get();
drop(profiler);
assert!(
stats.total_blocks as usize <= BOUND,
"fit_suff_stats_t_sq allocated {} blocks across {} warm-path calls (expected ≤ {})",
stats.total_blocks,
N_CALLS,
BOUND
);
}
}