use ndarray::{Array1, Array2, ArrayView1, ArrayView2};
use crate::error::{RegressionError, Result};
use crate::linalg::dmatrix_from_rows;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Method {
Reml,
Ml,
}
#[derive(Debug, Clone)]
pub struct LinearMixedModel {
coefficients: Array1<f64>,
cov_beta: Array2<f64>,
var_residual: f64,
var_group: f64,
lambda: f64,
blups: Array1<f64>,
log_likelihood: f64,
method: Method,
n: usize,
p: usize,
n_groups: usize,
}
impl LinearMixedModel {
pub fn new(x: Array2<f64>, y: Array1<f64>, groups: &[usize]) -> Result<Self> {
Self::with_method(x, y, groups, Method::Reml)
}
pub fn with_method(
x: Array2<f64>,
y: Array1<f64>,
groups: &[usize],
method: Method,
) -> Result<Self> {
let n = x.nrows();
let p = x.ncols();
if n == 0 || p == 0 {
return Err(RegressionError::EmptyInput { what: "X" });
}
if y.len() != n || groups.len() != n {
return Err(RegressionError::ShapeMismatch {
what: "y/groups length vs X rows",
expected: n,
got: y.len().min(groups.len()),
});
}
if n <= p {
return Err(RegressionError::NoResidualDegreesOfFreedom {
n,
p,
df: n as isize - p as isize,
});
}
let mut label_to_idx = std::collections::BTreeMap::new();
for &g in groups {
let next = label_to_idx.len();
label_to_idx.entry(g).or_insert(next);
}
let g = label_to_idx.len();
if g < 2 {
return Err(RegressionError::InvalidResponse {
msg: "a mixed model needs at least two groups".into(),
});
}
let mut group_rows: Vec<Vec<usize>> = vec![Vec::new(); g];
for (i, &lab) in groups.iter().enumerate() {
group_rows[label_to_idx[&lab]].push(i);
}
let xtx = x.t().dot(&x); let xty = x.t().dot(&y); let yty: f64 = y.iter().map(|v| v * v).sum();
let mut s = vec![vec![0.0f64; p]; g];
let mut t = vec![0.0f64; g];
let mut sizes = vec![0usize; g];
for (j, rows) in group_rows.iter().enumerate() {
sizes[j] = rows.len();
for &i in rows {
t[j] += y[i];
for a in 0..p {
s[j][a] += x[(i, a)];
}
}
}
let ctx = ProfileCtx {
xtx: &xtx,
xty: &xty,
yty,
s: &s,
t: &t,
sizes: &sizes,
n,
p,
g,
method,
};
let phi = (5.0_f64.sqrt() - 1.0) / 2.0;
let (mut lo, mut hi) = (0.0_f64, 1.0 - 1e-9);
let mut c = hi - phi * (hi - lo);
let mut d = lo + phi * (hi - lo);
let mut fc = ctx.objective(eta_to_lambda(c))?;
let mut fd = ctx.objective(eta_to_lambda(d))?;
for _ in 0..200 {
if fc < fd {
hi = d;
d = c;
fd = fc;
c = hi - phi * (hi - lo);
fc = ctx.objective(eta_to_lambda(c))?;
} else {
lo = c;
c = d;
fc = fd;
d = lo + phi * (hi - lo);
fd = ctx.objective(eta_to_lambda(d))?;
}
if (hi - lo) < 1e-10 {
break;
}
}
let eta_hat = 0.5 * (lo + hi);
let lambda = eta_to_lambda(eta_hat);
let sol = ctx.solve(lambda)?;
let dof = match method {
Method::Reml => (n - p) as f64,
Method::Ml => n as f64,
};
let var_residual = sol.rmr / dof;
let var_group = lambda * var_residual;
let cov_beta = &sol.xtmx_inv * var_residual;
let mut blups = Array1::<f64>::zeros(g);
let fitted_fixed = x.dot(&sol.beta);
for (j, rows) in group_rows.iter().enumerate() {
if rows.is_empty() {
continue;
}
let rbar: f64 = rows.iter().map(|&i| y[i] - fitted_fixed[i]).sum::<f64>()
/ rows.len() as f64;
let nj = rows.len() as f64;
blups[j] = (lambda * nj / (1.0 + lambda * nj)) * rbar;
}
let log_likelihood = -0.5 * ctx.objective(lambda)? - ctx.log_const();
Ok(Self {
coefficients: sol.beta,
cov_beta,
var_residual,
var_group,
lambda,
blups,
log_likelihood,
method,
n,
p,
n_groups: g,
})
}
pub fn n_observations(&self) -> usize {
self.n
}
pub fn n_parameters(&self) -> usize {
self.p
}
pub fn n_groups(&self) -> usize {
self.n_groups
}
pub fn method(&self) -> Method {
self.method
}
pub fn coefficients(&self) -> ArrayView1<'_, f64> {
self.coefficients.view()
}
pub fn covariance(&self) -> ArrayView2<'_, f64> {
self.cov_beta.view()
}
pub fn coefficient_standard_errors(&self) -> Array1<f64> {
Array1::from_shape_fn(self.p, |j| self.cov_beta[(j, j)].max(0.0).sqrt())
}
pub fn residual_variance(&self) -> f64 {
self.var_residual
}
pub fn group_variance(&self) -> f64 {
self.var_group
}
pub fn variance_ratio(&self) -> f64 {
self.lambda
}
pub fn icc(&self) -> f64 {
let total = self.var_group + self.var_residual;
if total > 0.0 {
self.var_group / total
} else {
f64::NAN
}
}
pub fn random_effects(&self) -> ArrayView1<'_, f64> {
self.blups.view()
}
pub fn log_likelihood(&self) -> f64 {
self.log_likelihood
}
pub fn aic(&self) -> f64 {
let k = match self.method {
Method::Ml => self.p as f64 + 2.0,
Method::Reml => 2.0,
};
-2.0 * self.log_likelihood + 2.0 * k
}
}
fn eta_to_lambda(eta: f64) -> f64 {
eta / (1.0 - eta)
}
struct ProfileCtx<'a> {
xtx: &'a Array2<f64>,
xty: &'a Array1<f64>,
yty: f64,
s: &'a [Vec<f64>],
t: &'a [f64],
sizes: &'a [usize],
n: usize,
p: usize,
g: usize,
method: Method,
}
struct Solve {
beta: Array1<f64>,
xtmx_inv: Array2<f64>,
rmr: f64,
}
impl ProfileCtx<'_> {
fn solve(&self, lambda: f64) -> Result<Solve> {
let p = self.p;
let mut xtmx = self.xtx.clone();
let mut xtmy = self.xty.clone();
let mut ytmy = self.yty;
for j in 0..self.g {
let nj = self.sizes[j] as f64;
let cj = lambda / (1.0 + lambda * nj);
if cj == 0.0 {
continue;
}
let sj = &self.s[j];
let tj = self.t[j];
for a in 0..p {
xtmy[a] -= cj * sj[a] * tj;
for b in 0..p {
xtmx[(a, b)] -= cj * sj[a] * sj[b];
}
}
ytmy -= cj * tj * tj;
}
let dm = dmatrix_from_rows(p, p, xtmx.as_standard_layout().as_slice().unwrap());
let inv = dm.try_inverse().ok_or(RegressionError::RankDeficient)?;
let xtmx_inv = Array2::from_shape_fn((p, p), |(i, j)| inv[(i, j)]);
let beta = xtmx_inv.dot(&xtmy);
let rmr = ytmy - beta.dot(&xtmy);
Ok(Solve {
beta,
xtmx_inv,
rmr,
})
}
fn objective(&self, lambda: f64) -> Result<f64> {
let sol = self.solve(lambda)?;
let dof = match self.method {
Method::Reml => (self.n - self.p) as f64,
Method::Ml => self.n as f64,
};
let sigma2 = (sol.rmr / dof).max(1e-300);
let ln_det_a: f64 = (0..self.g)
.map(|j| (1.0 + lambda * self.sizes[j] as f64).ln())
.sum();
let mut obj = dof * sigma2.ln() + ln_det_a;
if self.method == Method::Reml {
let p = self.p;
let mut xtmx = self.xtx.clone();
for j in 0..self.g {
let nj = self.sizes[j] as f64;
let cj = lambda / (1.0 + lambda * nj);
if cj == 0.0 {
continue;
}
let sj = &self.s[j];
for a in 0..p {
for b in 0..p {
xtmx[(a, b)] -= cj * sj[a] * sj[b];
}
}
}
let dm = dmatrix_from_rows(p, p, xtmx.as_standard_layout().as_slice().unwrap());
let det = dm.determinant();
obj += det.abs().max(1e-300).ln();
}
Ok(obj)
}
fn log_const(&self) -> f64 {
let dof = match self.method {
Method::Reml => (self.n - self.p) as f64,
Method::Ml => self.n as f64,
};
0.5 * dof * ((2.0 * std::f64::consts::PI).ln() + 1.0)
}
}