use faer::Mat;
use crate::ols::{OlsScratch, OlsSuffStats, PANEL_ROWS};
use super::common::{fill_se_compact, nan_vcov, vcov_from_chol};
use super::{Fit, FitOptions};
pub(super) fn fit_ols(x: &[f64], y: &[f64], n: usize, p: usize, opts: &FitOptions) -> Fit {
let t = opts.target_indices.len();
let p1 = p.max(1); let mut fit_betas = vec![0.0f64; p1];
let mut fit_var_diag = vec![0.0f64; t.max(1)];
let mut fit_t_sq = vec![0.0f64; t.max(1)];
let mut fit_u_scratch = vec![0.0f64; p1];
let mut fit_factor = Mat::<f64>::zeros(p1, p1);
let mut fit_rhs = Mat::<f64>::zeros(p1, 1);
let mut suff_xtx = Mat::<f64>::zeros(p1, p1);
let mut suff_xty = vec![0.0f64; p1];
let mut suff_yty = 0.0f64;
let mut suff_sum_y = 0.0f64;
let mut suff_n_rows = 0usize;
let mut suff_xtx_work = Mat::<f64>::zeros(p1, p1);
let mut panel_x = vec![0.0f64; PANEL_ROWS * p1];
let mut panel_y = vec![0.0f64; PANEL_ROWS];
let sqrt_w: Option<Vec<f64>> = opts
.weights
.as_ref()
.map(|w| w.iter().map(|v| v.sqrt()).collect());
let mut x_mat = Mat::<f64>::zeros(n.max(1), p1);
for i in 0..n {
let s = sqrt_w.as_ref().map_or(1.0, |sw| sw[i]);
for j in 0..p {
x_mat[(i, j)] = s * x[i * p + j];
}
}
let y_scaled: Vec<f64>;
let y_eff: &[f64] = match &sqrt_w {
Some(sw) => {
y_scaled = y.iter().zip(sw).map(|(&yi, &s)| yi * s).collect();
&y_scaled
}
None => y,
};
{
let mut suff = OlsSuffStats {
xtx: suff_xtx.as_mut(),
xty: &mut suff_xty,
yty: &mut suff_yty,
sum_y: &mut suff_sum_y,
n_rows: &mut suff_n_rows,
panel_x: &mut panel_x,
panel_y: &mut panel_y,
};
if n > 0 && p > 0 {
suff.add_rows(x_mat.as_ref().subrows(0, n), y_eff);
}
}
let view = {
let scratch = OlsScratch {
fit_betas: &mut fit_betas,
fit_var_diag: &mut fit_var_diag,
fit_t_sq: &mut fit_t_sq,
fit_u_scratch: &mut fit_u_scratch,
fit_factor: fit_factor.as_mut(),
fit_rhs: fit_rhs.as_mut(),
};
crate::ols::fit_suff_stats_t_sq(
suff_xtx.as_ref(),
&suff_xty,
suff_yty,
suff_sum_y,
suff_n_rows,
&opts.target_indices,
1e-12,
suff_xtx_work.as_mut(),
scratch,
)
};
let beta = view.betas.to_vec();
let converged = view.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)
};
Fit {
beta,
se,
vcov,
tau2: vec![],
dispersion: view.sigma_sq,
converged,
varcorr: vec![],
stddev_se: vec![],
aliased: vec![false; p],
n_eval: 0,
deviance: f64::NAN,
singular: false,
}
}