use faer::Mat;
use crate::ols::{OlsFitView, OlsScratch, OlsSuffStats, PANEL_ROWS};
use super::common::{fill_se_compact, nan_vcov, vcov_from_chol, FitDiagnostics};
use super::{Fit, FitOptions};
impl OlsFitView<'_> {
pub(crate) fn diagnostics(&self) -> FitDiagnostics {
FitDiagnostics {
pivot: self.pivot,
pivot_col: self.pivot_col,
ill_conditioned: self.pivot < crate::ols::PIVOT_MIN,
..FitDiagnostics::fixed_only(self.converged)
}
}
}
pub(crate) struct OlsWorkspace {
fit_betas: Vec<f64>,
fit_var_diag: Vec<f64>,
fit_t_sq: Vec<f64>,
fit_u_scratch: Vec<f64>,
fit_factor: Mat<f64>,
fit_rhs: Mat<f64>,
suff_xtx: Mat<f64>,
suff_xty: Vec<f64>,
suff_xtx_work: Mat<f64>,
panel_x: Vec<f64>,
panel_y: Vec<f64>,
scaled_x: Mat<f64>,
scaled_y: Vec<f64>,
}
impl OlsWorkspace {
pub(crate) fn new(n_max: usize, p: usize, t: usize, has_weights: bool) -> Self {
let n1 = n_max.max(1);
let p1 = p.max(1); let t1 = t.max(1);
Self {
fit_betas: vec![0.0f64; p1],
fit_var_diag: vec![0.0f64; t1],
fit_t_sq: vec![0.0f64; t1],
fit_u_scratch: vec![0.0f64; p1],
fit_factor: Mat::<f64>::zeros(p1, p1),
fit_rhs: Mat::<f64>::zeros(p1, 1),
suff_xtx: Mat::<f64>::zeros(p1, p1),
suff_xty: vec![0.0f64; p1],
suff_xtx_work: Mat::<f64>::zeros(p1, p1),
panel_x: vec![0.0f64; PANEL_ROWS * p1],
panel_y: vec![0.0f64; PANEL_ROWS],
scaled_x: if has_weights {
Mat::<f64>::zeros(n1, p1)
} else {
Mat::<f64>::zeros(0, 0)
},
scaled_y: vec![0.0f64; n1],
}
}
}
pub(crate) fn fit_ols_prebuilt<'a>(
ws: &'a mut OlsWorkspace,
x_mat: faer::MatRef<'_, f64>,
y: &[f64],
n: usize,
p: usize,
opts: &FitOptions,
) -> OlsFitView<'a> {
ws.suff_xtx.fill(0.0);
ws.suff_xty.iter_mut().for_each(|z| *z = 0.0);
let mut suff_yty = 0.0f64;
let mut suff_sum_y = 0.0f64;
let mut suff_n_rows = 0usize;
let weighted = opts.weights.is_some();
let offset = opts.offset.is_some();
let x_eff = if weighted {
let w = opts.weights.as_ref().unwrap();
for i in 0..n {
let s = w[i].sqrt();
for j in 0..p {
ws.scaled_x[(i, j)] = s * x_mat[(i, j)];
}
}
ws.scaled_x.as_ref().subrows(0, n)
} else {
x_mat
};
let y_eff: &[f64] = if weighted || offset {
for i in 0..n {
let yi = y[i] - opts.offset.as_ref().map_or(0.0, |o| o[i]);
ws.scaled_y[i] = if weighted {
yi * opts.weights.as_ref().unwrap()[i].sqrt()
} else {
yi
};
}
&ws.scaled_y[..n]
} else {
y
};
{
let mut suff = OlsSuffStats {
xtx: ws.suff_xtx.as_mut(),
xty: &mut ws.suff_xty,
yty: &mut suff_yty,
sum_y: &mut suff_sum_y,
n_rows: &mut suff_n_rows,
panel_x: &mut ws.panel_x,
panel_y: &mut ws.panel_y,
};
if n > 0 && p > 0 {
suff.add_rows(x_eff, y_eff);
}
}
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(),
};
crate::ols::fit_suff_stats_t_sq(
ws.suff_xtx.as_ref(),
&ws.suff_xty,
suff_yty,
suff_sum_y,
suff_n_rows,
&opts.target_indices,
ws.suff_xtx_work.as_mut(),
scratch,
)
}
#[cfg(test)]
pub(super) fn fit_ols(x: &[f64], y: &[f64], n: usize, p: usize, opts: &FitOptions) -> Fit {
let mut ws = OlsWorkspace::new(n, p, opts.target_indices.len(), opts.weights.is_some());
let x_mat = super::common::to_col_major(x, n, p);
let view = fit_ols_prebuilt(&mut ws, x_mat.as_ref().subrows(0, n), y, n, p, opts);
ols_view_to_fit(&view, x, y, n, p, opts)
}
pub(crate) fn ols_view_to_fit(
view: &OlsFitView<'_>,
x: &[f64],
y: &[f64],
n: usize,
p: usize,
opts: &FitOptions,
) -> Fit {
let beta = view.betas.to_vec();
let diag = view.diagnostics();
let converged = diag.converged;
let mut se = vec![f64::NAN; p];
fill_se_compact(view.var_diag, &opts.target_indices, &mut se);
let vcov = if converged {
vcov_from_chol(view.factor, p, &opts.target_indices, view.sigma_sq)
} else {
nan_vcov(p)
};
let (fitted, loglik) = if converged && n > 0 {
let fitted: Vec<f64> = (0..n)
.map(|i| {
let o = opts.offset.as_ref().map_or(0.0, |o| o[i]);
o + (0..p).map(|j| x[i * p + j] * beta[j]).sum::<f64>()
})
.collect();
let rss: f64 = (0..n)
.map(|i| {
let r = y[i] - fitted[i];
opts.weights.as_ref().map_or(1.0, |w| w[i]) * r * r
})
.sum();
let sum_log_w = opts
.weights
.as_ref()
.map_or(0.0, |w| w.iter().map(|v| v.ln()).sum());
let nf = n as f64;
let ll =
0.5 * (sum_log_w - nf * ((2.0 * std::f64::consts::PI).ln() + 1.0 - nf.ln() + rss.ln()));
(fitted, ll)
} else {
(vec![], f64::NAN)
};
Fit {
beta,
se,
vcov,
tau2: vec![],
dispersion: view.sigma_sq,
diagnostics: super::common::materialize_diagnostics(&diag, p, &[]),
varcorr: vec![],
stddev_se: vec![],
n_eval: 0,
deviance: f64::NAN,
loglik,
df: if converged { p + 1 } else { 0 },
reml: false,
fitted,
ranef: vec![],
ranef_levels: vec![],
}
}