use crate::model_types::EstimationError;
use crate::vector_response::VectorLikelihood;
use faer::Side;
use gam_linalg::faer_ndarray::{FaerArrayView, array2_to_matmut, factorize_symmetricwith_fallback};
use gam_problem::{
FixedLambdaCheckpoint, FixedLambdaResidualKind, FixedLambdaSolverStage, FixedLambdaStallReason,
FixedLambdaStationarityEvidence,
};
use gam_solve::pirls::dense_block_xtwx;
use ndarray::{Array1, Array2, ArrayView1, ArrayView2, ArrayView3};
use opt::{BacktrackConfig, RidgeSchedule, backtracking_line_search, escalate_ridge};
const BASE_RIDGE_FRACTION_OF_MAX_DIAG: f64 = 1.0e-10;
const MAX_RIDGE_ESCALATIONS: usize = 30;
const MAX_BACKTRACKS: usize = 8;
const LINE_SEARCH_SHRINK: f64 = 0.5;
const OBJECTIVE_DECREASE_SLACK: f64 = 1.0e-12;
const OPTIMALITY_GRAD_FRACTION: f64 = 1.0e-6;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum ClassPenaltyMetric {
#[default]
Diagonal,
Centered,
EquivariantPerClass,
}
impl ClassPenaltyMetric {
pub fn active_outputs(self, lambdas_len: usize) -> usize {
match self {
ClassPenaltyMetric::Diagonal | ClassPenaltyMetric::Centered => lambdas_len,
ClassPenaltyMetric::EquivariantPerClass => lambdas_len.saturating_sub(1),
}
}
}
pub(crate) fn equivariant_class_metric(lambdas: ArrayView1<'_, f64>, m: usize) -> Array2<f64> {
let k = (m + 1) as f64;
let total: f64 = lambdas.iter().sum();
let mut a_mat = Array2::<f64>::zeros((m, m));
for a in 0..m {
for b in 0..m {
let mut value = -(lambdas[a] + lambdas[b]) / k + total / (k * k);
if a == b {
value += lambdas[a];
}
a_mat[[a, b]] = value;
}
}
a_mat
}
pub struct PenalizedVectorGlmInputs<'a> {
pub design: ArrayView2<'a, f64>,
pub y: ArrayView2<'a, f64>,
pub penalty: ArrayView2<'a, f64>,
pub lambdas: ArrayView1<'a, f64>,
pub fisher_w_override: Option<ArrayView3<'a, f64>>,
pub max_iter: usize,
pub tol: f64,
pub class_penalty_metric: ClassPenaltyMetric,
pub resume_from: Option<VectorGlmResume<'a>>,
}
#[derive(Debug, Clone, Copy)]
pub struct VectorGlmResume<'a> {
pub coefficients: ArrayView2<'a, f64>,
pub completed_iterations: usize,
}
pub struct PenalizedVectorGlmOutputs {
pub coefficients: Array2<f64>,
pub eta: Array2<f64>,
pub iterations: usize,
pub log_likelihood: f64,
pub penalty_term: f64,
pub coefficient_covariance: Array2<f64>,
}
pub struct VectorGlmStall {
pub reason: VectorGlmStallReason,
pub coefficients: Array2<f64>,
pub eta: Array2<f64>,
pub iterations: usize,
pub log_likelihood: f64,
pub penalty_term: f64,
pub gradient_norm: f64,
pub gradient_bound: f64,
}
impl VectorGlmStall {
pub fn into_nonconvergence_error(
self,
stage: FixedLambdaSolverStage,
context: impl Into<String>,
) -> Result<EstimationError, EstimationError> {
let rows = self.coefficients.nrows();
let cols = self.coefficients.ncols();
let checkpoint = FixedLambdaCheckpoint::new(
stage,
self.coefficients.iter().copied().collect(),
rows,
cols,
self.iterations,
)
.map_err(|reason| {
EstimationError::InvalidInput(format!(
"fixed-lambda vector-GLM produced an invalid internal checkpoint: {reason}"
))
})?;
let reason = match self.reason {
VectorGlmStallReason::IterationBudgetExhausted => {
FixedLambdaStallReason::IterationBudgetExhausted
}
VectorGlmStallReason::LineSearchExhausted => {
FixedLambdaStallReason::LineSearchExhausted
}
VectorGlmStallReason::PostStepCertificateFailed => {
FixedLambdaStallReason::StationarityCertificateFailed
}
};
Ok(EstimationError::FixedLambdaNewtonDidNotConverge {
context: context.into(),
reason,
objective_value: -self.log_likelihood + self.penalty_term,
stationarity: FixedLambdaStationarityEvidence {
kind: FixedLambdaResidualKind::PenalizedGradientNorm,
residual: self.gradient_norm,
bound: self.gradient_bound,
},
checkpoint,
})
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum VectorGlmStallReason {
IterationBudgetExhausted,
LineSearchExhausted,
PostStepCertificateFailed,
}
pub enum VectorGlmSolve {
Converged(PenalizedVectorGlmOutputs),
Stalled(VectorGlmStall),
}
fn add_equivariant_penalty_blocks(
hessian: &mut Array2<f64>,
penalty: ArrayView2<'_, f64>,
lambdas: ArrayView1<'_, f64>,
p: usize,
m: usize,
) {
if m == 0 {
return;
}
let a_mat = equivariant_class_metric(lambdas, m);
for a in 0..m {
for b in 0..m {
let coef = a_mat[[a, b]];
if coef == 0.0 {
continue;
}
let (ba, bb) = (a * p, b * p);
for i in 0..p {
for j in 0..p {
hessian[[ba + i, bb + j]] += coef * penalty[[i, j]];
}
}
}
}
}
fn weighted_penalty_sum(
beta: &Array2<f64>,
penalty: ArrayView2<'_, f64>,
lambdas: ArrayView1<'_, f64>,
metric: ClassPenaltyMetric,
) -> f64 {
let (p, m) = beta.dim();
match metric {
ClassPenaltyMetric::Diagonal => {
let mut pen = 0.0_f64;
for a in 0..m {
let la = lambdas[a];
if la == 0.0 {
continue;
}
let beta_col = beta.column(a);
let mut quad = 0.0_f64;
for i in 0..p {
let mut s_beta_i = 0.0_f64;
for j in 0..p {
s_beta_i += penalty[[i, j]] * beta_col[j];
}
quad += beta_col[i] * s_beta_i;
}
pen += 0.5 * la * quad;
}
pen
}
ClassPenaltyMetric::Centered => {
if m == 0 {
return 0.0;
}
let lam = lambdas[0];
if lam == 0.0 {
return 0.0;
}
let k = (m + 1) as f64;
let mut g = vec![0.0_f64; p];
for a in 0..m {
let col = beta.column(a);
for i in 0..p {
g[i] += col[i];
}
}
let mut sum_quad = 0.0_f64;
for a in 0..m {
let col = beta.column(a);
for i in 0..p {
let mut s_beta_i = 0.0_f64;
for j in 0..p {
s_beta_i += penalty[[i, j]] * col[j];
}
sum_quad += col[i] * s_beta_i;
}
}
let mut g_quad = 0.0_f64;
for i in 0..p {
let mut s_g_i = 0.0_f64;
for j in 0..p {
s_g_i += penalty[[i, j]] * g[j];
}
g_quad += g[i] * s_g_i;
}
0.5 * lam * (sum_quad - g_quad / k)
}
ClassPenaltyMetric::EquivariantPerClass => {
if m == 0 {
return 0.0;
}
let a_mat = equivariant_class_metric(lambdas, m);
let mut s_beta = Array2::<f64>::zeros((p, m));
for b in 0..m {
let col = beta.column(b);
for i in 0..p {
let mut acc = 0.0_f64;
for j in 0..p {
acc += penalty[[i, j]] * col[j];
}
s_beta[[i, b]] = acc;
}
}
let mut pen = 0.0_f64;
for a in 0..m {
let col = beta.column(a);
for b in 0..m {
let coef = a_mat[[a, b]];
if coef == 0.0 {
continue;
}
let mut cross = 0.0_f64;
for i in 0..p {
cross += col[i] * s_beta[[i, b]];
}
pen += 0.5 * coef * cross;
}
}
pen
}
}
}
fn fill_penalized_gradient(
design: ArrayView2<'_, f64>,
residual: ArrayView2<'_, f64>,
beta: &Array2<f64>,
penalty: ArrayView2<'_, f64>,
lambdas: ArrayView1<'_, f64>,
metric: ClassPenaltyMetric,
out: &mut Array1<f64>,
) {
let (p, m) = beta.dim();
for a in 0..m {
for i in 0..p {
let mut acc = 0.0_f64;
for row in 0..design.nrows() {
acc += design[[row, i]] * residual[[row, a]];
}
out[a * p + i] = acc;
}
}
match metric {
ClassPenaltyMetric::Diagonal => {
for a in 0..m {
let la = lambdas[a];
if la == 0.0 {
continue;
}
let beta_col = beta.column(a);
for i in 0..p {
let mut s_beta_i = 0.0_f64;
for j in 0..p {
s_beta_i += penalty[[i, j]] * beta_col[j];
}
out[a * p + i] += la * s_beta_i;
}
}
}
ClassPenaltyMetric::Centered if m > 0 && lambdas[0] != 0.0 => {
let lam = lambdas[0];
let inv_k = 1.0 / ((m + 1) as f64);
let mut beta_bar = vec![0.0_f64; p];
for a in 0..m {
let col = beta.column(a);
for i in 0..p {
beta_bar[i] += col[i];
}
}
for value in &mut beta_bar {
*value *= inv_k;
}
for a in 0..m {
let beta_col = beta.column(a);
for i in 0..p {
let mut s_centered_i = 0.0_f64;
for j in 0..p {
s_centered_i += penalty[[i, j]] * (beta_col[j] - beta_bar[j]);
}
out[a * p + i] += lam * s_centered_i;
}
}
}
ClassPenaltyMetric::Centered => {}
ClassPenaltyMetric::EquivariantPerClass if m > 0 => {
let a_mat = equivariant_class_metric(lambdas, m);
let mut s_beta = Array2::<f64>::zeros((p, m));
for b in 0..m {
let col = beta.column(b);
for i in 0..p {
let mut acc = 0.0_f64;
for j in 0..p {
acc += penalty[[i, j]] * col[j];
}
s_beta[[i, b]] = acc;
}
}
for a in 0..m {
for i in 0..p {
let mut acc = 0.0_f64;
for b in 0..m {
acc += a_mat[[a, b]] * s_beta[[i, b]];
}
out[a * p + i] += acc;
}
}
}
ClassPenaltyMetric::EquivariantPerClass => {}
}
}
fn invert_symmetric_penalized_hessian(
hessian: &Array2<f64>,
dim: usize,
context: &str,
) -> Result<Array2<f64>, EstimationError> {
let max_diag = (0..dim).fold(0.0_f64, |acc, idx| acc.max(hessian[[idx, idx]].abs()));
let base_ridge = if max_diag.is_finite() && max_diag > 0.0 {
max_diag * BASE_RIDGE_FRACTION_OF_MAX_DIAG
} else {
BASE_RIDGE_FRACTION_OF_MAX_DIAG
};
let mut last_failure: Option<(f64, String)> = None;
let mut try_ridge = |ridge: f64| -> Option<Array2<f64>> {
let mut ridged = hessian.clone();
if ridge > 0.0 {
for idx in 0..dim {
ridged[[idx, idx]] += ridge;
}
}
let factor = match factorize_symmetricwith_fallback(
FaerArrayView::new(&ridged).as_ref(),
Side::Lower,
) {
Ok(factor) => factor,
Err(err) => {
last_failure = Some((ridge, err.to_string()));
return None;
}
};
let mut rhs = Array2::<f64>::eye(dim);
{
let rhs_view = array2_to_matmut(&mut rhs);
factor.solve_in_place(rhs_view);
}
if !rhs.iter().all(|v| v.is_finite()) {
last_failure = None;
return None;
}
let mut cov = Array2::<f64>::zeros((dim, dim));
for i in 0..dim {
for j in 0..dim {
cov[[i, j]] = 0.5 * (rhs[[i, j]] + rhs[[j, i]]);
}
}
Some(cov)
};
if let Some(cov) = try_ridge(0.0) {
return Ok(cov);
}
match escalate_ridge(
RidgeSchedule {
initial: base_ridge,
growth: 2.0,
max_escalations: MAX_RIDGE_ESCALATIONS,
},
&mut try_ridge,
) {
Ok(success) => Ok(success.value),
Err(_) => match last_failure {
Some((ridge, err)) => Err(EstimationError::InvalidInput(format!(
"{context}: covariance factorization failed even with ridge \
{ridge:.3e}: {err}"
))),
None => Err(EstimationError::InvalidInput(format!(
"{context}: covariance solve remained non-finite after {} ridge escalations \
(max_diag={max_diag:.3e})",
MAX_RIDGE_ESCALATIONS,
))),
},
}
}
pub fn fit_penalized_vector_glm<L: VectorLikelihood>(
inputs: PenalizedVectorGlmInputs<'_>,
likelihood: &L,
context: &str,
) -> Result<VectorGlmSolve, EstimationError> {
let PenalizedVectorGlmInputs {
design,
y,
penalty,
lambdas,
fisher_w_override,
max_iter,
tol,
class_penalty_metric,
resume_from,
} = inputs;
let n_obs = design.nrows();
let p = design.ncols();
if n_obs == 0 || p == 0 {
crate::bail_invalid_estim!("{context}: design must be nonempty (got {n_obs}x{p})");
}
let m = class_penalty_metric.active_outputs(lambdas.len());
if m == 0 {
crate::bail_invalid_estim!("{context}: need at least one active output (got M=0)");
}
if y.nrows() != n_obs {
crate::bail_invalid_estim!("{context}: y rows {} ≠ design rows {n_obs}", y.nrows());
}
if penalty.dim() != (p, p) {
crate::bail_invalid_estim!(
"{context}: penalty shape {:?} ≠ (P, P) = ({p}, {p})",
penalty.dim()
);
}
for (i, &v) in lambdas.iter().enumerate() {
if !(v.is_finite() && v >= 0.0) {
crate::bail_invalid_estim!("{context}: lambdas[{i}] must be finite and ≥ 0 (got {v})");
}
}
if let Some(fw) = fisher_w_override.as_ref() {
if fw.dim() != (n_obs, m, m) {
crate::bail_invalid_estim!(
"{context}: fisher_w_override shape {:?} ≠ (N, M, M) = ({n_obs}, {m}, {m})",
fw.dim()
);
}
}
for ((i, j), &v) in design.indexed_iter() {
if !v.is_finite() {
crate::bail_invalid_estim!("{context}: design[{i},{j}] must be finite (got {v})");
}
}
let (mut beta, completed_iterations) = match resume_from {
Some(resume) => {
if resume.coefficients.dim() != (p, m) {
crate::bail_invalid_estim!(
"{context}: resume checkpoint coefficient shape {:?} ≠ (P, M) = ({p}, {m})",
resume.coefficients.dim()
);
}
for ((i, a), &value) in resume.coefficients.indexed_iter() {
if !value.is_finite() {
crate::bail_invalid_estim!(
"{context}: resume checkpoint coefficient[{i},{a}] must be finite (got {value})"
);
}
}
(resume.coefficients.to_owned(), resume.completed_iterations)
}
None => (Array2::<f64>::zeros((p, m)), 0),
};
let mut eta = Array2::<f64>::zeros((n_obs, m));
let mut eta_objective_scratch = Array2::<f64>::zeros((n_obs, m));
let beta_flat_dim = p * m;
let mut grad_flat = Array1::<f64>::zeros(beta_flat_dim);
let mut iterations = completed_iterations;
let mut small_step_reached = false;
let mut stall_reason = VectorGlmStallReason::IterationBudgetExhausted;
let mut last_objective = f64::INFINITY;
let recompute_eta = |beta: &Array2<f64>, eta: &mut Array2<f64>| {
for a in 0..m {
let beta_col = beta.column(a);
for row in 0..n_obs {
let mut eta_val = 0.0_f64;
for i in 0..p {
eta_val += design[[row, i]] * beta_col[i];
}
eta[[row, a]] = eta_val;
}
}
};
let evaluate_objective =
|beta_trial: &Array2<f64>, eta_scratch: &mut Array2<f64>| -> Result<f64, EstimationError> {
recompute_eta(beta_trial, eta_scratch);
let ll = likelihood.log_lik(eta_scratch.view(), y)?;
let pen = weighted_penalty_sum(beta_trial, penalty, lambdas, class_penalty_metric);
Ok(-ll + pen)
};
for iter in 0..max_iter {
iterations = completed_iterations.checked_add(iter + 1).ok_or_else(|| {
EstimationError::InvalidInput(format!(
"{context}: resume checkpoint iteration count overflowed usize"
))
})?;
recompute_eta(&beta, &mut eta);
let analytic_fisher = match fisher_w_override.as_ref() {
Some(_) => None,
None => Some(likelihood.hess_block(eta.view(), y)?),
};
let fisher_blocks = match fisher_w_override.as_ref() {
Some(fw) => *fw,
None => analytic_fisher
.as_ref()
.expect("analytic Fisher computed when no override")
.view(),
};
let residual = likelihood.grad_eta(eta.view(), y)?.mapv(|v| -v);
let mut hessian = dense_block_xtwx(design, fisher_blocks, None)?;
if hessian.nrows() != beta_flat_dim || hessian.ncols() != beta_flat_dim {
crate::bail_invalid_estim!(
"{context}: assembled Hessian shape {:?} ≠ ({beta_flat_dim}, {beta_flat_dim})",
hessian.dim()
);
}
match class_penalty_metric {
ClassPenaltyMetric::Diagonal => {
for a in 0..m {
let la = lambdas[a];
if la == 0.0 {
continue;
}
let base = a * p;
for i in 0..p {
for j in 0..p {
hessian[[base + i, base + j]] += la * penalty[[i, j]];
}
}
}
}
ClassPenaltyMetric::Centered if m > 0 && lambdas[0] != 0.0 => {
let lam = lambdas[0];
let inv_k = 1.0 / ((m + 1) as f64);
for a in 0..m {
for b in 0..m {
let coef = lam * (if a == b { 1.0 } else { 0.0 } - inv_k);
let (ba, bb) = (a * p, b * p);
for i in 0..p {
for j in 0..p {
hessian[[ba + i, bb + j]] += coef * penalty[[i, j]];
}
}
}
}
}
ClassPenaltyMetric::Centered => {}
ClassPenaltyMetric::EquivariantPerClass => {
add_equivariant_penalty_blocks(&mut hessian, penalty, lambdas, p, m);
}
}
fill_penalized_gradient(
design,
residual.view(),
&beta,
penalty,
lambdas,
class_penalty_metric,
&mut grad_flat,
);
let max_diag =
(0..beta_flat_dim).fold(0.0_f64, |acc, idx| acc.max(hessian[[idx, idx]].abs()));
let base_ridge = if max_diag.is_finite() && max_diag > 0.0 {
max_diag * BASE_RIDGE_FRACTION_OF_MAX_DIAG
} else {
BASE_RIDGE_FRACTION_OF_MAX_DIAG
};
let mut last_factor_err: Option<(f64, String)> = None;
let delta = match escalate_ridge(
RidgeSchedule {
initial: base_ridge,
growth: 2.0,
max_escalations: MAX_RIDGE_ESCALATIONS + 1,
},
|ridge| {
let mut ridged = hessian.clone();
for idx in 0..beta_flat_dim {
ridged[[idx, idx]] += ridge;
}
let factor = match factorize_symmetricwith_fallback(
FaerArrayView::new(&ridged).as_ref(),
Side::Lower,
) {
Ok(factor) => factor,
Err(err) => {
last_factor_err = Some((ridge, err.to_string()));
return None;
}
};
last_factor_err = None;
let mut rhs = Array2::<f64>::zeros((beta_flat_dim, 1));
for i in 0..beta_flat_dim {
rhs[[i, 0]] = -grad_flat[i];
}
{
let rhs_view = array2_to_matmut(&mut rhs);
factor.solve_in_place(rhs_view);
}
(0..beta_flat_dim)
.all(|i| rhs[[i, 0]].is_finite())
.then(|| Array1::from_iter((0..beta_flat_dim).map(|i| rhs[[i, 0]])))
},
) {
Ok(success) => success.value,
Err(exhausted) => {
if let Some((ridge, err)) = last_factor_err {
return Err(EstimationError::InvalidInput(format!(
"{context}: Hessian factorization failed at iter {iter} \
even with ridge {ridge:.3e}: {err}"
)));
}
return Err(EstimationError::InvalidInput(format!(
"{context}: Newton step remained non-finite at iter {iter} after {} ridge \
escalations up to {:.3e}; the penalized Hessian is pathologically \
rank-deficient (grad_norm={:.3e}, max_diag={max_diag:.3e})",
MAX_RIDGE_ESCALATIONS,
exhausted.next_ridge,
grad_flat.iter().map(|v| v * v).sum::<f64>().sqrt(),
)));
}
};
let proposed_beta = |alpha: f64| -> Array2<f64> {
let mut out = beta.clone();
for a in 0..m {
for i in 0..p {
out[[i, a]] += alpha * delta[a * p + i];
}
}
out
};
if iter == 0 {
last_objective = evaluate_objective(&beta, &mut eta_objective_scratch)?;
if !last_objective.is_finite() {
crate::bail_invalid_estim!("{context}: non-finite objective at β = 0");
}
}
let accepted = backtracking_line_search::<_, EstimationError>(
BacktrackConfig {
contraction: LINE_SEARCH_SHRINK,
max_steps: MAX_BACKTRACKS + 1,
..BacktrackConfig::default()
},
|alpha| {
let candidate = proposed_beta(alpha);
let objective = evaluate_objective(&candidate, &mut eta_objective_scratch)?;
Ok(Some((objective, candidate)))
},
|_alpha, f| f.is_finite() && f <= last_objective + OBJECTIVE_DECREASE_SLACK,
)?;
let Some(accepted) = accepted else {
stall_reason = VectorGlmStallReason::LineSearchExhausted;
break;
};
let accepted_beta = accepted.payload;
let new_objective = accepted.value;
let mut step_norm_sq = 0.0_f64;
let mut beta_norm_sq = 0.0_f64;
for a in 0..m {
for i in 0..p {
let d = accepted_beta[[i, a]] - beta[[i, a]];
step_norm_sq += d * d;
let v = accepted_beta[[i, a]];
beta_norm_sq += v * v;
}
}
beta = accepted_beta;
last_objective = new_objective;
let step_norm = step_norm_sq.sqrt();
let beta_norm = beta_norm_sq.sqrt();
let grad_norm = grad_flat.iter().map(|v| v * v).sum::<f64>().sqrt();
let grad_optimal = grad_norm <= OPTIMALITY_GRAD_FRACTION * (1.0 + max_diag);
if step_norm <= tol * (1.0 + beta_norm) && grad_optimal {
small_step_reached = true;
break;
}
}
recompute_eta(&beta, &mut eta);
let log_likelihood = likelihood.log_lik(eta.view(), y)?;
let penalty_term = weighted_penalty_sum(&beta, penalty, lambdas, class_penalty_metric);
let analytic_fisher_final = match fisher_w_override.as_ref() {
Some(_) => None,
None => Some(likelihood.hess_block(eta.view(), y)?),
};
let fisher_blocks_final = match fisher_w_override.as_ref() {
Some(fw) => *fw,
None => analytic_fisher_final
.as_ref()
.expect("analytic Fisher computed when no override")
.view(),
};
let mut hessian_final = dense_block_xtwx(design, fisher_blocks_final, None)?;
match class_penalty_metric {
ClassPenaltyMetric::Diagonal => {
for a in 0..m {
let la = lambdas[a];
if la == 0.0 {
continue;
}
let base = a * p;
for i in 0..p {
for j in 0..p {
hessian_final[[base + i, base + j]] += la * penalty[[i, j]];
}
}
}
}
ClassPenaltyMetric::Centered if m > 0 && lambdas[0] != 0.0 => {
let lam = lambdas[0];
let inv_k = 1.0 / ((m + 1) as f64);
for a in 0..m {
for b in 0..m {
let coef = lam * (if a == b { 1.0 } else { 0.0 } - inv_k);
let (ba, bb) = (a * p, b * p);
for i in 0..p {
for j in 0..p {
hessian_final[[ba + i, bb + j]] += coef * penalty[[i, j]];
}
}
}
}
}
ClassPenaltyMetric::Centered => {}
ClassPenaltyMetric::EquivariantPerClass => {
add_equivariant_penalty_blocks(&mut hessian_final, penalty, lambdas, p, m);
}
}
let final_residual = likelihood.grad_eta(eta.view(), y)?.mapv(|value| -value);
fill_penalized_gradient(
design,
final_residual.view(),
&beta,
penalty,
lambdas,
class_penalty_metric,
&mut grad_flat,
);
let final_grad_norm = grad_flat
.iter()
.map(|value| value * value)
.sum::<f64>()
.sqrt();
let final_max_diag =
(0..beta_flat_dim).fold(0.0_f64, |acc, i| acc.max(hessian_final[[i, i]].abs()));
let final_grad_optimal = final_grad_norm <= OPTIMALITY_GRAD_FRACTION * (1.0 + final_max_diag);
if !(small_step_reached && final_grad_optimal) {
if small_step_reached {
stall_reason = VectorGlmStallReason::PostStepCertificateFailed;
}
return Ok(VectorGlmSolve::Stalled(VectorGlmStall {
reason: stall_reason,
coefficients: beta,
eta,
iterations,
log_likelihood,
penalty_term,
gradient_norm: final_grad_norm,
gradient_bound: OPTIMALITY_GRAD_FRACTION * (1.0 + final_max_diag),
}));
}
let coefficient_covariance =
invert_symmetric_penalized_hessian(&hessian_final, beta_flat_dim, context)?;
Ok(VectorGlmSolve::Converged(PenalizedVectorGlmOutputs {
coefficients: beta,
eta,
iterations,
log_likelihood,
penalty_term,
coefficient_covariance,
}))
}
#[cfg(test)]
mod parity_tests {
use super::{ClassPenaltyMetric, weighted_penalty_sum};
use crate::binomial_multi::{BinomialMultiFitInputs, fit_penalized_binomial_multi};
use crate::multinomial::{MultinomialFitInputs, fit_penalized_multinomial};
use gam_test_support::fd_checker::numerical_gradient_central_diff;
use ndarray::{Array1, Array2};
#[test]
fn centered_penalty_is_reference_class_invariant_1587() {
let s = ndarray::array![[2.0_f64, 0.5], [0.5, 1.0]];
let bt = [[1.0_f64, 0.5], [-0.3, 0.2], [-0.7, -0.7]];
for j in 0..2 {
let colsum: f64 = (0..3).map(|k| bt[k][j]).sum();
assert!(colsum.abs() < 1e-12, "test CLR set must sum to zero");
}
let mut symmetric = 0.0_f64;
for k in 0..3 {
for i in 0..2 {
for j in 0..2 {
symmetric += bt[k][i] * s[[i, j]] * bt[k][j];
}
}
}
let lambdas = Array1::from(vec![1.0_f64, 1.0]);
let mut centered_vals = Vec::new();
let mut diagonal_vals = Vec::new();
for r in 0..3 {
let others: Vec<usize> = (0..3).filter(|&k| k != r).collect();
let mut beta = Array2::<f64>::zeros((2, 2));
for (a, &o) in others.iter().enumerate() {
for i in 0..2 {
beta[[i, a]] = bt[o][i] - bt[r][i];
}
}
let c = weighted_penalty_sum(
&beta,
s.view(),
lambdas.view(),
ClassPenaltyMetric::Centered,
);
let d = weighted_penalty_sum(
&beta,
s.view(),
lambdas.view(),
ClassPenaltyMetric::Diagonal,
);
assert!(
(c - 0.5 * symmetric).abs() < 1e-12,
"ref {r}: Centered penalty {c} must equal ½·symmetric {}",
0.5 * symmetric
);
centered_vals.push(c);
diagonal_vals.push(d);
}
let cspread = centered_vals.iter().cloned().fold(f64::MIN, f64::max)
- centered_vals.iter().cloned().fold(f64::MAX, f64::min);
assert!(
cspread < 1e-12,
"Centered must be reference-invariant; got {centered_vals:?}"
);
let dspread = diagonal_vals.iter().cloned().fold(f64::MIN, f64::max)
- diagonal_vals.iter().cloned().fold(f64::MAX, f64::min);
assert!(
dspread > 1e-6,
"Diagonal is the non-invariant #1587 path; references must disagree, got {diagonal_vals:?}"
);
}
fn sigmoid(eta: f64) -> f64 {
if eta >= 0.0 {
1.0 / (1.0 + (-eta).exp())
} else {
let e = eta.exp();
e / (1.0 + e)
}
}
fn softmax_ref(eta_active: &[f64]) -> Vec<f64> {
let m = eta_active.len();
let mut out = vec![0.0_f64; m + 1];
let mut max_eta = 0.0_f64;
for &v in eta_active {
if v > max_eta {
max_eta = v;
}
}
let baseline = (-max_eta).exp();
let mut denom = baseline;
for (idx, &v) in eta_active.iter().enumerate() {
let e = (v - max_eta).exp();
out[idx] = e;
denom += e;
}
for v in out.iter_mut().take(m) {
*v /= denom;
}
out[m] = baseline / denom;
out
}
fn binomial_objective(
design: &Array2<f64>,
y: &Array2<f64>,
penalty: &Array2<f64>,
lambdas: &Array1<f64>,
beta: &Array2<f64>,
) -> f64 {
let (n, p) = design.dim();
let k = y.ncols();
let mut ll = 0.0_f64;
for row in 0..n {
for a in 0..k {
let mut eta = 0.0_f64;
for i in 0..p {
eta += design[[row, i]] * beta[[i, a]];
}
let mu = sigmoid(eta).clamp(1.0e-12, 1.0 - 1.0e-12);
let yv = y[[row, a]];
ll += yv * mu.ln() + (1.0 - yv) * (1.0 - mu).ln();
}
}
let mut pen = 0.0_f64;
for a in 0..k {
let la = lambdas[a];
for i in 0..p {
let mut sbi = 0.0_f64;
for j in 0..p {
sbi += penalty[[i, j]] * beta[[j, a]];
}
pen += 0.5 * la * beta[[i, a]] * sbi;
}
}
-ll + pen
}
fn multinomial_objective(
design: &Array2<f64>,
y_one_hot: &Array2<f64>,
penalty: &Array2<f64>,
lambdas: &Array1<f64>,
beta: &Array2<f64>,
) -> f64 {
let (n, p) = design.dim();
let k = y_one_hot.ncols();
let m = k - 1;
let mut ll = 0.0_f64;
let mut eta_active = vec![0.0_f64; m];
for row in 0..n {
for a in 0..m {
let mut eta = 0.0_f64;
for i in 0..p {
eta += design[[row, i]] * beta[[i, a]];
}
eta_active[a] = eta;
}
let probs = softmax_ref(&eta_active);
for c in 0..k {
let yc = y_one_hot[[row, c]];
if yc != 0.0 {
ll += yc * probs[c].max(1.0e-300).ln();
}
}
}
let kf = k as f64;
let mut pen = 0.0_f64;
let mut beta_bar = vec![0.0_f64; p];
for a in 0..m {
for i in 0..p {
beta_bar[i] += beta[[i, a]] / kf;
}
}
for c in 0..k {
let lc = lambdas[c];
if lc == 0.0 {
continue;
}
let gamma_i = |i: usize| -> f64 {
if c < m {
beta[[i, c]] - beta_bar[i]
} else {
-beta_bar[i]
}
};
for i in 0..p {
let mut s_gamma_i = 0.0_f64;
for j in 0..p {
s_gamma_i += penalty[[i, j]] * gamma_i(j);
}
pen += 0.5 * lc * gamma_i(i) * s_gamma_i;
}
}
-ll + pen
}
fn fd_grad<F: Fn(&Array2<f64>) -> f64>(beta: &Array2<f64>, f: F) -> f64 {
let (p, c) = beta.dim();
let flat = Array1::from_iter(beta.iter().copied());
let grad = numerical_gradient_central_diff(
|x: &Array1<f64>| {
let m = Array2::from_shape_vec((p, c), x.to_vec())
.expect("row-major reshape of coefficient vector");
f(&m)
},
&flat,
1.0e-6,
);
grad.iter().fold(0.0_f64, |acc, &g| acc.max(g.abs()))
}
fn binomial_fixture() -> (Array2<f64>, Array2<f64>, Array2<f64>, Array1<f64>) {
let n = 40;
let p = 3;
let k = 3;
let design = Array2::<f64>::from_shape_fn((n, p), |(i, j)| match j {
0 => 1.0,
1 => ((i + 1) as f64 * 0.37).sin(),
_ => ((i + 1) as f64 * 0.11).cos(),
});
let y = Array2::<f64>::from_shape_fn((n, k), |(i, a)| {
if ((i * 7 + a * 13 + 3) % 5) < 3 {
1.0
} else {
0.0
}
});
let penalty = Array2::<f64>::eye(p);
let lambdas = Array1::from(vec![0.3_f64, 1.2, 2.5]);
(design, y, penalty, lambdas)
}
fn multinomial_fixture() -> (Array2<f64>, Array2<f64>, Array2<f64>, Array1<f64>) {
let n = 45;
let p = 3;
let k = 4;
let design = Array2::<f64>::from_shape_fn((n, p), |(i, j)| match j {
0 => 1.0,
1 => ((i + 2) as f64 * 0.29).sin(),
_ => ((i + 2) as f64 * 0.17).cos(),
});
let mut y = Array2::<f64>::zeros((n, k));
for i in 0..n {
y[[i, (i * 3 + 1) % k]] = 1.0;
}
let penalty = Array2::<f64>::eye(p);
let lambdas = Array1::from(vec![0.5_f64, 1.0, 2.0, 0.8]);
(design, y, penalty, lambdas)
}
#[test]
fn binomial_engine_hits_optimum_and_is_self_consistent() {
let (design, y, penalty, lambdas) = binomial_fixture();
let fit = fit_penalized_binomial_multi(BinomialMultiFitInputs {
design: design.view(),
y: y.view(),
penalty: penalty.view(),
lambdas: lambdas.view(),
row_weights: None,
fisher_w_override: None,
max_iter: 100,
tol: 1.0e-12,
})
.expect("binomial fit must succeed");
let g = fd_grad(&fit.coefficients, |b| {
binomial_objective(&design, &y, &penalty, &lambdas, b)
});
assert!(
g < 1.0e-6,
"binomial penalized gradient at β̂ must vanish (max |∂F| = {g})"
);
let (n, p) = design.dim();
let k = y.ncols();
let mut log_lik = 0.0_f64;
for row in 0..n {
for a in 0..k {
let mut eta = 0.0_f64;
for i in 0..p {
eta += design[[row, i]] * fit.coefficients[[i, a]];
}
let mu = sigmoid(eta);
assert!(
(fit.fitted_probabilities[[row, a]] - mu).abs() < 1.0e-10,
"fitted probability must equal σ(X β̂)"
);
let muc = mu.clamp(1.0e-12, 1.0 - 1.0e-12);
let yv = y[[row, a]];
log_lik += yv * muc.ln() + (1.0 - yv) * (1.0 - muc).ln();
}
}
assert!(
(fit.deviance - (-2.0 * log_lik)).abs() < 1.0e-9,
"deviance must equal −2 log L"
);
}
#[test]
fn binomial_joint_solve_decouples_into_single_column_solves() {
let (design, y, penalty, lambdas) = binomial_fixture();
let joint = fit_penalized_binomial_multi(BinomialMultiFitInputs {
design: design.view(),
y: y.view(),
penalty: penalty.view(),
lambdas: lambdas.view(),
row_weights: None,
fisher_w_override: None,
max_iter: 100,
tol: 1.0e-12,
})
.expect("joint fit must succeed");
let k = y.ncols();
for a in 0..k {
let y_col = y.column(a).to_owned().insert_axis(ndarray::Axis(1));
let lam = Array1::from(vec![lambdas[a]]);
let single = fit_penalized_binomial_multi(BinomialMultiFitInputs {
design: design.view(),
y: y_col.view(),
penalty: penalty.view(),
lambdas: lam.view(),
row_weights: None,
fisher_w_override: None,
max_iter: 100,
tol: 1.0e-12,
})
.expect("single-column fit must succeed");
for i in 0..design.ncols() {
let dj = joint.coefficients[[i, a]];
let ds = single.coefficients[[i, 0]];
assert!(
(dj - ds).abs() < 1.0e-8,
"joint column {a} coef {i} ({dj}) must match single-column solve ({ds})"
);
}
}
}
#[test]
fn multinomial_engine_hits_optimum_and_is_self_consistent() {
let (design, y, penalty, lambdas) = multinomial_fixture();
let fit = fit_penalized_multinomial(MultinomialFitInputs {
design: design.view(),
y_one_hot: y.view(),
penalty: penalty.view(),
lambdas: lambdas.view(),
row_weights: None,
fisher_w_override: None,
max_iter: 100,
tol: 1.0e-12,
resume_from: None,
})
.expect("multinomial fit must succeed");
let g = fd_grad(&fit.coefficients_active, |b| {
multinomial_objective(&design, &y, &penalty, &lambdas, b)
});
assert!(
g < 1.0e-6,
"multinomial penalized gradient at β̂ must vanish (max |∂F| = {g})"
);
let (n, p) = design.dim();
let k = y.ncols();
let m = k - 1;
let mut log_lik = 0.0_f64;
let mut eta_active = vec![0.0_f64; m];
for row in 0..n {
for a in 0..m {
let mut eta = 0.0_f64;
for i in 0..p {
eta += design[[row, i]] * fit.coefficients_active[[i, a]];
}
eta_active[a] = eta;
}
let probs = softmax_ref(&eta_active);
let mut row_sum = 0.0_f64;
for c in 0..k {
assert!(
(fit.fitted_probabilities[[row, c]] - probs[c]).abs() < 1.0e-10,
"fitted probability must equal softmax(X β̂)"
);
row_sum += fit.fitted_probabilities[[row, c]];
let yc = y[[row, c]];
if yc != 0.0 {
log_lik += yc * probs[c].max(1.0e-300).ln();
}
}
assert!(
(row_sum - 1.0).abs() < 1.0e-10,
"fitted probabilities must sum to 1 per row"
);
}
assert!(
(fit.deviance - (-2.0 * log_lik)).abs() < 1.0e-9,
"deviance must equal −2 log L"
);
}
#[test]
fn multinomial_rank_deficient_block_recovers_via_ridge_not_crash() {
let n = 50;
let p = 4;
let k = 4;
let design = Array2::<f64>::from_shape_fn((n, p), |(i, j)| match j {
0 => 1.0,
1 => ((i + 1) as f64 * 0.23).sin(),
2 => ((i + 1) as f64 * 0.23).sin(), _ => ((i + 1) as f64 * 0.19).cos(),
});
let mut y = Array2::<f64>::zeros((n, k));
for i in 0..n {
y[[i, (i * 5 + 2) % k]] = 1.0;
}
let mut penalty = Array2::<f64>::zeros((p, p));
penalty[[3, 3]] = 1.0;
let lambdas = Array1::from(vec![1.0e-10_f64, 1.0e-10, 1.0e-10, 1.0e-10]);
let fit = fit_penalized_multinomial(MultinomialFitInputs {
design: design.view(),
y_one_hot: y.view(),
penalty: penalty.view(),
lambdas: lambdas.view(),
row_weights: None,
fisher_w_override: None,
max_iter: 200,
tol: 1.0e-10,
resume_from: None,
})
.expect("rank-deficient multinomial fit must NOT crash (#557): the ridge path recovers it");
for &c in fit.coefficients_active.iter() {
assert!(c.is_finite(), "coefficient must be finite, got {c}");
}
for &pr in fit.fitted_probabilities.iter() {
assert!(
pr.is_finite() && (-1.0e-9..=1.0 + 1.0e-9).contains(&pr),
"fitted probability must be a finite simplex entry, got {pr}"
);
}
let (nn, kk) = fit.fitted_probabilities.dim();
for row in 0..nn {
let s: f64 = (0..kk).map(|c| fit.fitted_probabilities[[row, c]]).sum();
assert!(
(s - 1.0).abs() < 1.0e-9,
"row {row} probabilities must sum to 1, got {s}"
);
}
let g = fd_grad(&fit.coefficients_active, |b| {
multinomial_objective(&design, &y, &penalty, &lambdas, b)
});
assert!(
g < 1.0e-4,
"penalized objective gradient at the ridge-recovered β̂ must (near-)vanish \
along identified directions (max |∂F| = {g})"
);
}
}