use crate::custom_family::{
BlockwiseFitOptions, ParameterBlockSpec, ParameterBlockState, PenaltyMatrix,
fit_custom_family_with_rho_prior,
};
use crate::fit_orchestration::drivers::freeze_term_collection_from_design;
use crate::fit_orchestration::{
FitConfig, build_termspec_with_geometry_and_overrides, resolved_resource_policy,
};
use crate::model_types::EstimationError;
use crate::multinomial_reml::MultinomialFamily;
use crate::multinomial_posterior::{
MultinomialPosteriorIntegrationControl, integrate_multinomial_design_moments,
};
use crate::penalized_vector_glm::{
PenalizedVectorGlmInputs, VectorGlmResume, VectorGlmSolve, fit_penalized_vector_glm,
};
use crate::vector_response::{MultinomialLogitLikelihood, validate_multinomial_simplex};
use gam_data::ColumnKindTag;
use gam_data::EncodedDataset;
use gam_problem::{
FixedLambdaCheckpoint, FixedLambdaResidualKind, FixedLambdaSolverStage, FixedLambdaStallReason,
FixedLambdaStationarityEvidence, ResponseColumnKind,
};
use gam_runtime::resource::ProblemHints;
use gam_terms::inference::formula_dsl::parse_formula;
use gam_terms::smooth::{
PenaltyBlockInfo, TermCollectionDesign, TermCollectionSpec, build_term_collection_design,
};
use gam_terms::term_builder::resolve_role_col;
use ndarray::{Array1, Array2, ArrayView1, ArrayView2, ArrayView3};
use opt::{BacktrackConfig, backtracking_line_search};
use serde::{Deserialize, Serialize};
use std::convert::Infallible;
use std::sync::Arc;
const MULTINOMIAL_FORMULA_RIDGE_FLOOR: f64 = 1.0e-4;
const MULTINOMIAL_FORMULA_INNER_TOL: f64 = 1.0e-5;
fn multinomial_formula_penalty_scale(n_classes: usize) -> f64 {
let k = n_classes.max(2) as f64;
2.0 * (k - 1.0) / (k * k)
}
const MULTINOMIAL_EXACT_OUTER_HESSIAN_MAX_DIM: usize = 16;
fn multinomial_formula_use_outer_hessian(total_rho_dim: usize) -> bool {
total_rho_dim <= MULTINOMIAL_EXACT_OUTER_HESSIAN_MAX_DIM
}
const MULTINOMIAL_SEPARATION_ETA_THRESHOLD: f64 = 25.0;
const MULTINOMIAL_OUTER_REML_TOL: f64 = 1e-7;
const MULTINOMIAL_UNBIASED_PROBE_OUTER_MAX_ITER: usize = 20;
const MULTINOMIAL_FORMULA_FISHER_INFO_PER_OBS: f64 = 0.25;
const MULTINOMIAL_FORMULA_PRIOR_PSEUDO_OBS: f64 = 8.0e-4;
const MULTINOMIAL_FORMULA_SPARSE_REFERENCE_SUPPORT: f64 = 50.0;
const MULTINOMIAL_FORMULA_SPARSE_PRIOR_PSEUDO_OBS_MAX: f64 = 4.0e-3;
fn multinomial_formula_min_lambda(y_one_hot: ArrayView2<'_, f64>) -> f64 {
let base = MULTINOMIAL_FORMULA_PRIOR_PSEUDO_OBS * MULTINOMIAL_FORMULA_FISHER_INFO_PER_OBS;
let sparse =
MULTINOMIAL_FORMULA_SPARSE_PRIOR_PSEUDO_OBS_MAX * MULTINOMIAL_FORMULA_FISHER_INFO_PER_OBS;
let min_class_count = (0..y_one_hot.ncols())
.map(|class| y_one_hot.column(class).sum())
.fold(f64::INFINITY, f64::min);
if !min_class_count.is_finite() || min_class_count <= 0.0 {
return base;
}
let pseudo_obs_scale =
(MULTINOMIAL_FORMULA_SPARSE_REFERENCE_SUPPORT / min_class_count).max(1.0);
(base * pseudo_obs_scale).clamp(base, sparse)
}
fn max_abs_eta_location(eta: ArrayView2<'_, f64>) -> (f64, usize, usize) {
let mut best = (0.0_f64, 0usize, 0usize);
for ((row, active_class), &value) in eta.indexed_iter() {
let abs = value.abs();
if abs > best.0 {
best = (abs, row, active_class);
}
}
best
}
fn multinomial_formula_separation_diagnostic(
inner_cycles: usize,
outer_iterations: usize,
block_states: &[ParameterBlockState],
) -> Option<EstimationError> {
let mut nonfinite: Option<(f64, usize, usize)> = None;
for (active_class, state) in block_states.iter().enumerate() {
for (row, &value) in state.eta.iter().enumerate() {
if !value.is_finite() {
nonfinite = Some((value, row, active_class));
break;
}
}
if nonfinite.is_some() {
break;
}
}
nonfinite.map(|(value, row_index, active_class_index)| {
EstimationError::MultinomialSeparationDetected {
iteration: inner_cycles.max(outer_iterations),
max_abs_eta: value.abs(),
active_class_index,
row_index,
}
})
}
fn multinomial_formula_separation_evidence(block_states: &[ParameterBlockState]) -> Option<String> {
for (active_class, state) in block_states.iter().enumerate() {
for (row, &value) in state.eta.iter().enumerate() {
if !value.is_finite() {
return Some(format!(
"non-finite logit eta[row {row}, active class {active_class}] = {value}"
));
}
}
}
None
}
#[derive(Debug, Clone)]
pub struct MultinomialFitInputs<'a> {
pub design: ArrayView2<'a, f64>,
pub y_one_hot: ArrayView2<'a, f64>,
pub penalty: ArrayView2<'a, f64>,
pub lambdas: ArrayView1<'a, f64>,
pub row_weights: Option<ArrayView1<'a, f64>>,
pub fisher_w_override: Option<ArrayView3<'a, f64>>,
pub max_iter: usize,
pub tol: f64,
pub resume_from: Option<&'a FixedLambdaCheckpoint>,
}
#[derive(Debug, Clone)]
pub struct MultinomialFitOutputs {
pub coefficients_active: Array2<f64>,
pub fitted_probabilities: Array2<f64>,
pub iterations: usize,
pub penalized_neg_log_likelihood: f64,
pub deviance: f64,
pub coefficient_covariance: Array2<f64>,
}
impl MultinomialFitOutputs {
pub fn n_active_classes(&self) -> usize {
self.coefficients_active.ncols()
}
pub fn p_per_class(&self) -> usize {
self.coefficients_active.nrows()
}
pub fn predict_probabilities_with_se(
&self,
x_new: ArrayView2<'_, f64>,
) -> Result<(Array2<f64>, Array2<f64>), EstimationError> {
self.predict_probabilities_with_se_and_control(
x_new,
&MultinomialPosteriorIntegrationControl::default(),
)
}
pub fn predict_probabilities_with_se_and_control(
&self,
x_new: ArrayView2<'_, f64>,
control: &MultinomialPosteriorIntegrationControl,
) -> Result<(Array2<f64>, Array2<f64>), EstimationError> {
let moments = integrate_multinomial_design_moments(
self.coefficients_active.view(),
self.coefficient_covariance.view(),
x_new,
control,
)?;
Ok((moments.class_mean, moments.class_standard_deviation))
}
}
#[derive(Clone, Copy)]
struct FirthResume<'a> {
coefficients: ArrayView2<'a, f64>,
completed_iterations: usize,
}
fn fixed_lambda_checkpoint_coefficients(
checkpoint: &FixedLambdaCheckpoint,
expected_stage: FixedLambdaSolverStage,
p: usize,
m: usize,
) -> Result<Array2<f64>, EstimationError> {
checkpoint.validate().map_err(|reason| {
EstimationError::InvalidInput(format!(
"multinomial fixed-λ resume checkpoint is invalid: {reason}"
))
})?;
if checkpoint.stage() != expected_stage {
crate::bail_invalid_estim!(
"multinomial fixed-λ resume checkpoint stage is {}, expected {}",
checkpoint.stage(),
expected_stage,
);
}
if checkpoint.rows() != p || checkpoint.cols() != m {
crate::bail_invalid_estim!(
"multinomial fixed-λ resume checkpoint shape {}x{} does not match P x (K-1) = {p}x{m}",
checkpoint.rows(),
checkpoint.cols(),
);
}
Array2::from_shape_vec((p, m), checkpoint.values().to_vec()).map_err(|error| {
EstimationError::InvalidInput(format!(
"multinomial fixed-λ resume checkpoint could not be reshaped: {error}"
))
})
}
pub fn fit_penalized_multinomial(
inputs: MultinomialFitInputs<'_>,
) -> Result<MultinomialFitOutputs, EstimationError> {
let MultinomialFitInputs {
design,
y_one_hot,
penalty,
lambdas,
row_weights,
fisher_w_override,
max_iter,
tol,
resume_from,
} = inputs;
let n_obs = design.nrows();
let (y_rows, k) = y_one_hot.dim();
if y_rows != n_obs {
crate::bail_invalid_estim!(
"fit_penalized_multinomial: y rows {y_rows} ≠ design rows {n_obs}"
);
}
if k < 2 {
crate::bail_invalid_estim!(
"fit_penalized_multinomial: need at least 2 classes (got K={k})"
);
}
let m = k - 1;
if lambdas.len() != k {
crate::bail_invalid_estim!(
"fit_penalized_multinomial: lambdas length {} ≠ K = {k} (one λ per class, \
reference class included — the permutation-equivariant per-class contract, #2344)",
lambdas.len()
);
}
if let Some(fw) = fisher_w_override.as_ref() {
if fw.dim() != (n_obs, m, m) {
crate::bail_invalid_estim!(
"fit_penalized_multinomial: fisher_w_override shape {:?} ≠ (N, K-1, K-1) = ({n_obs}, {m}, {m})",
fw.dim()
);
}
}
if let Some(w) = row_weights.as_ref() {
if w.len() != n_obs {
crate::bail_invalid_estim!(
"fit_penalized_multinomial: row_weights length {} ≠ N = {n_obs}",
w.len()
);
}
for (i, &v) in w.iter().enumerate() {
if !(v.is_finite() && v >= 0.0) {
crate::bail_invalid_estim!(
"fit_penalized_multinomial: row_weights[{i}] must be finite and ≥ 0 (got {v})"
);
}
}
}
validate_multinomial_simplex(y_one_hot, "fit_penalized_multinomial")?;
let p = design.ncols();
let resumed_newton_coefficients = match resume_from {
Some(checkpoint) if checkpoint.stage() == FixedLambdaSolverStage::MultinomialFirth => {
let coefficients = fixed_lambda_checkpoint_coefficients(
checkpoint,
FixedLambdaSolverStage::MultinomialFirth,
p,
m,
)?;
return fit_penalized_multinomial_firth_fallback(
design,
y_one_hot,
penalty,
lambdas,
row_weights,
max_iter,
tol,
Some(FirthResume {
coefficients: coefficients.view(),
completed_iterations: checkpoint.completed_iterations(),
}),
);
}
Some(checkpoint) => Some(fixed_lambda_checkpoint_coefficients(
checkpoint,
FixedLambdaSolverStage::MultinomialNewton,
p,
m,
)?),
None => None,
};
let vector_resume = resumed_newton_coefficients
.as_ref()
.map(|coefficients| VectorGlmResume {
coefficients: coefficients.view(),
completed_iterations: resume_from
.map(FixedLambdaCheckpoint::completed_iterations)
.unwrap_or(0),
});
let mut likelihood = MultinomialLogitLikelihood::with_classes(k)?;
if let Some(w) = row_weights.as_ref() {
likelihood = likelihood.with_row_weights(w.to_owned())?;
}
let solve = fit_penalized_vector_glm(
PenalizedVectorGlmInputs {
design,
y: y_one_hot,
penalty,
lambdas,
fisher_w_override,
max_iter,
tol,
class_penalty_metric: crate::penalized_vector_glm::ClassPenaltyMetric::EquivariantPerClass,
resume_from: vector_resume,
},
&likelihood,
"fit_penalized_multinomial",
)?;
let fit = match solve {
VectorGlmSolve::Converged(fit) => fit,
VectorGlmSolve::Stalled(stall) => {
return handle_multinomial_fixed_lambda_stall(
stall,
design,
y_one_hot,
penalty,
lambdas,
row_weights,
max_iter,
tol,
);
}
};
let fitted_probabilities = likelihood.probabilities(fit.eta.view());
Ok(MultinomialFitOutputs {
coefficients_active: fit.coefficients,
fitted_probabilities,
iterations: fit.iterations,
penalized_neg_log_likelihood: -fit.log_likelihood + fit.penalty_term,
deviance: -2.0 * fit.log_likelihood,
coefficient_covariance: fit.coefficient_covariance,
})
}
fn handle_multinomial_fixed_lambda_stall(
stall: crate::penalized_vector_glm::VectorGlmStall,
design: ArrayView2<'_, f64>,
y_one_hot: ArrayView2<'_, f64>,
penalty: ArrayView2<'_, f64>,
lambdas: ArrayView1<'_, f64>,
row_weights: Option<ArrayView1<'_, f64>>,
max_iter: usize,
tol: f64,
) -> Result<MultinomialFitOutputs, EstimationError> {
let (max_abs_eta, row_index, active_class_index) = max_abs_eta_location(stall.eta.view());
if max_abs_eta >= MULTINOMIAL_SEPARATION_ETA_THRESHOLD {
let firth = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
fit_penalized_multinomial_firth_fallback(
design,
y_one_hot,
penalty,
lambdas,
row_weights,
max_iter,
tol,
None,
)
}));
match firth {
Ok(Ok(out)) => return Ok(out),
Ok(Err(err @ EstimationError::FixedLambdaNewtonDidNotConverge { .. })) => {
return Err(err);
}
Ok(Err(_)) | Err(_) => {
return Err(EstimationError::MultinomialSeparationDetected {
iteration: stall.iterations,
max_abs_eta,
active_class_index,
row_index,
});
}
}
}
Err(stall.into_nonconvergence_error(
FixedLambdaSolverStage::MultinomialNewton,
"fit_penalized_multinomial (fixed-λ softmax damped Newton)",
)?)
}
fn fit_penalized_multinomial_firth_fallback(
design: ArrayView2<'_, f64>,
y_one_hot: ArrayView2<'_, f64>,
penalty: ArrayView2<'_, f64>,
lambdas: ArrayView1<'_, f64>,
row_weights: Option<ArrayView1<'_, f64>>,
max_iter: usize,
tol: f64,
resume_from: Option<FirthResume<'_>>,
) -> Result<MultinomialFitOutputs, EstimationError> {
use faer::Side;
use gam_linalg::faer_ndarray::{
FaerArrayView, array1_to_col_matmut, array2_to_matmut, factorize_symmetricwith_fallback,
};
use gam_linalg::matrix::FactorizedSystem;
let n_obs = design.nrows();
let p = design.ncols();
let k = y_one_hot.ncols();
let m = k - 1;
let d = p * m;
let mut likelihood = MultinomialLogitLikelihood::with_classes(k)?;
if let Some(w) = row_weights.as_ref() {
likelihood = likelihood.with_row_weights(w.to_owned())?;
}
let weight = |row: usize| -> f64 { row_weights.as_ref().map_or(1.0, |w| w[row]) };
let tol_eff = if tol.is_finite() && tol > 0.0 {
tol
} else {
1e-8
};
let probs_at = |beta: &Array2<f64>| -> Array2<f64> {
let eta = design.dot(beta);
likelihood.probabilities(eta.view())
};
let assemble_info = |probs: &Array2<f64>| -> Array2<f64> {
let mut info = Array2::<f64>::zeros((d, d));
for row in 0..n_obs {
let w = weight(row);
if w == 0.0 {
continue;
}
for a in 0..m {
let pa = probs[[row, a]];
let ao = a * p;
for b in 0..m {
let pb = probs[[row, b]];
let wab = w * (if a == b { pa - pa * pb } else { -pa * pb });
if wab == 0.0 {
continue;
}
let bo = b * p;
for i in 0..p {
let xi = design[[row, i]];
if xi == 0.0 {
continue;
}
let cc = wab * xi;
for j in 0..p {
info[[ao + i, bo + j]] += cc * design[[row, j]];
}
}
}
}
}
info
};
let invert_spd = |mat: &Array2<f64>,
context: &str|
-> Result<(Array2<f64>, f64), EstimationError> {
let max_diag = (0..d).fold(0.0_f64, |acc, i| acc.max(mat[[i, i]].abs()));
let base = if max_diag.is_finite() && max_diag > 0.0 {
max_diag * 1e-10
} else {
1e-10
};
let mut ridge = 0.0_f64;
for _ in 0..=60 {
let mut ridged = mat.clone();
if ridge > 0.0 {
for i in 0..d {
ridged[[i, i]] += ridge;
}
}
if let Ok(factor) =
factorize_symmetricwith_fallback(FaerArrayView::new(&ridged).as_ref(), Side::Lower)
{
let logdet = factor.logdet();
if logdet.is_finite() {
let mut rhs = Array2::<f64>::eye(d);
{
let v = array2_to_matmut(&mut rhs);
factor.solve_in_place(v);
}
if rhs.iter().all(|x| x.is_finite()) {
let mut inv = Array2::<f64>::zeros((d, d));
for i in 0..d {
for j in 0..d {
inv[[i, j]] = 0.5 * (rhs[[i, j]] + rhs[[j, i]]);
}
}
return Ok((inv, logdet));
}
}
}
ridge = if ridge > 0.0 { ridge * 4.0 } else { base };
}
Err(EstimationError::InvalidInput(format!(
"multinomial Firth fallback: {context} not invertible (max_diag={max_diag:.3e})"
)))
};
let spd_logdet = |mat: &Array2<f64>| -> Option<f64> {
factorize_symmetricwith_fallback(FaerArrayView::new(mat).as_ref(), Side::Lower)
.ok()
.map(|factor| factor.logdet())
.filter(|ld| ld.is_finite())
};
let objective = |probs: &Array2<f64>, beta: &Array2<f64>, logdet_info: f64| -> f64 {
let mut ll = 0.0_f64;
for row in 0..n_obs {
let w = weight(row);
if w == 0.0 {
continue;
}
for c in 0..k {
let ycn = y_one_hot[[row, c]];
if ycn != 0.0 {
ll += w * ycn * probs[[row, c]].max(f64::MIN_POSITIVE).ln();
}
}
}
let a_mat = crate::penalized_vector_glm::equivariant_class_metric(lambdas, m);
let mut pen = 0.0_f64;
for a in 0..m {
let bcol = beta.column(a);
for b in 0..m {
let coef = a_mat[[a, b]];
if coef != 0.0 {
let sbeta = penalty.dot(&beta.column(b));
pen += 0.5 * coef * bcol.dot(&sbeta);
}
}
}
ll - pen + 0.5 * logdet_info
};
let firth_score =
|probs: &Array2<f64>, beta: &Array2<f64>, iinv: &Array2<f64>| -> Array1<f64> {
let mut u = Array1::<f64>::zeros(d);
let mut xn = vec![0.0_f64; p];
let mut pa = vec![0.0_f64; m];
let mut q = vec![0.0_f64; m * m];
for row in 0..n_obs {
let w = weight(row);
if w == 0.0 {
continue;
}
for i in 0..p {
xn[i] = design[[row, i]];
}
for a in 0..m {
pa[a] = probs[[row, a]];
}
for a in 0..m {
let resid = y_one_hot[[row, a]] - pa[a];
let ao = a * p;
for i in 0..p {
u[ao + i] += w * xn[i] * resid;
}
}
for a in 0..m {
let ao = a * p;
for b in 0..m {
let bo = b * p;
let mut s = 0.0_f64;
for i in 0..p {
let xi = xn[i];
if xi == 0.0 {
continue;
}
let mut inner = 0.0_f64;
for j in 0..p {
inner += iinv[[ao + i, bo + j]] * xn[j];
}
s += xi * inner;
}
q[a * m + b] = s;
}
}
for c in 0..m {
let pc = pa[c];
let mut h = 0.0_f64;
for a in 0..m {
for b in 0..m {
let dab = if a == b { 1.0 } else { 0.0 };
let dac = if a == c { 1.0 } else { 0.0 };
let dbc = if b == c { 1.0 } else { 0.0 };
let g =
dab * pa[a] * (dac - pc) - pa[a] * pa[b] * (dac + dbc - 2.0 * pc);
h += g * q[a * m + b];
}
}
let co = c * p;
for s in 0..p {
u[co + s] += 0.5 * w * h * xn[s];
}
}
}
let a_mat = crate::penalized_vector_glm::equivariant_class_metric(lambdas, m);
for b in 0..m {
let sbeta = penalty.dot(&beta.column(b));
for a in 0..m {
let coef = a_mat[[a, b]];
if coef == 0.0 {
continue;
}
let ao = a * p;
for i in 0..p {
u[ao + i] -= coef * sbeta[i];
}
}
}
u
};
let penalized_hessian = |info: &Array2<f64>| -> Array2<f64> {
let mut h = info.clone();
let a_mat = crate::penalized_vector_glm::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 (ao, bo) = (a * p, b * p);
for i in 0..p {
for j in 0..p {
h[[ao + i, bo + j]] += coef * penalty[[i, j]];
}
}
}
}
h
};
let solve_spd = |mat: &Array2<f64>,
rhs: &Array1<f64>|
-> Result<Array1<f64>, EstimationError> {
let max_diag = (0..d).fold(0.0_f64, |acc, i| acc.max(mat[[i, i]].abs()));
let base = if max_diag.is_finite() && max_diag > 0.0 {
max_diag * 1e-12
} else {
1e-12
};
let mut ridge = 0.0_f64;
for _ in 0..=60 {
let mut ridged = mat.clone();
if ridge > 0.0 {
for i in 0..d {
ridged[[i, i]] += ridge;
}
}
if let Ok(factor) =
factorize_symmetricwith_fallback(FaerArrayView::new(&ridged).as_ref(), Side::Lower)
{
let mut sol = rhs.clone();
{
let v = array1_to_col_matmut(&mut sol);
factor.solve_in_place(v);
}
if sol.iter().all(|x| x.is_finite()) {
return Ok(sol);
}
}
ridge = if ridge > 0.0 { ridge * 4.0 } else { base };
}
Err(EstimationError::InvalidInput(
"multinomial Firth fallback: penalized Hessian solve failed".to_string(),
))
};
let (mut beta, completed_iterations) = match resume_from {
Some(resume) => {
if resume.coefficients.dim() != (p, m) {
crate::bail_invalid_estim!(
"multinomial Firth resume coefficient shape {:?} does not match P x (K-1) = {p}x{m}",
resume.coefficients.dim(),
);
}
(resume.coefficients.to_owned(), resume.completed_iterations)
}
None => (Array2::<f64>::zeros((p, m)), 0),
};
let mut iterations = completed_iterations;
let mut stall_reason = FixedLambdaStallReason::IterationBudgetExhausted;
let mut small_step_reached = false;
for it in 0..max_iter {
iterations = completed_iterations.checked_add(it + 1).ok_or_else(|| {
EstimationError::InvalidInput(
"multinomial Firth resume iteration count overflowed usize".to_string(),
)
})?;
let probs = probs_at(&beta);
let info = assemble_info(&probs);
let (iinv, logdet_info) = invert_spd(&info, "Fisher information")?;
let u = firth_score(&probs, &beta, &iinv);
let hmat = penalized_hessian(&info);
let step_vec = solve_spd(&hmat, &u)?;
let decrement = u.dot(&step_vec);
if 0.5 * decrement.abs() < tol_eff {
break;
}
let mut delta = Array2::<f64>::zeros((p, m));
for a in 0..m {
let ao = a * p;
for i in 0..p {
delta[[i, a]] = step_vec[ao + i];
}
}
let o0 = objective(&probs, &beta, logdet_info);
let accepted_step = match backtracking_line_search::<_, Infallible>(
BacktrackConfig::default(),
|step| {
let cand = &beta + &(&delta * step);
let cand_probs = probs_at(&cand);
let cand_info = assemble_info(&cand_probs);
Ok(spd_logdet(&cand_info)
.map(|cand_logdet| (objective(&cand_probs, &cand, cand_logdet), cand)))
},
|_step, o1| o1 >= o0 - 1e-12,
) {
Ok(result) => result,
Err(never) => match never {},
};
let Some(accepted_step) = accepted_step else {
stall_reason = FixedLambdaStallReason::LineSearchExhausted;
break;
};
let step = accepted_step.step;
beta = accepted_step.payload;
let max_step = step * delta.iter().fold(0.0_f64, |acc, &v| acc.max(v.abs()));
let scale = 1.0 + beta.iter().fold(0.0_f64, |acc, &v| acc.max(v.abs()));
if max_step < tol_eff * scale {
small_step_reached = true;
break;
}
}
for (idx, &v) in beta.iter().enumerate() {
if !v.is_finite() {
crate::bail_invalid_estim!(
"multinomial Firth fallback: non-finite coefficient at flat index {idx} = {v}"
);
}
}
let coefficients_active = beta;
let mut log_likelihood = 0.0_f64;
let probs = probs_at(&coefficients_active);
for row in 0..n_obs {
let w = weight(row);
for c in 0..k {
let ycn = y_one_hot[[row, c]];
if ycn != 0.0 {
log_likelihood += w * ycn * probs[[row, c]].max(f64::MIN_POSITIVE).ln();
}
}
}
let a_mat = crate::penalized_vector_glm::equivariant_class_metric(lambdas, m);
let mut penalty_term = 0.0_f64;
for a in 0..m {
let beta_col = coefficients_active.column(a);
for b in 0..m {
let coef = a_mat[[a, b]];
if coef != 0.0 {
let sbeta = penalty.dot(&coefficients_active.column(b));
penalty_term += 0.5 * coef * beta_col.dot(&sbeta);
}
}
}
let info = assemble_info(&probs);
let (information_inverse, final_logdet_info) = invert_spd(&info, "final Fisher information")?;
let final_score = firth_score(&probs, &coefficients_active, &information_inverse);
let hmat = penalized_hessian(&info);
let final_step = solve_spd(&hmat, &final_score)?;
let final_decrement = 0.5 * final_score.dot(&final_step).abs();
if !(final_decrement.is_finite() && final_decrement < tol_eff) {
if small_step_reached {
stall_reason = FixedLambdaStallReason::StationarityCertificateFailed;
}
let checkpoint = FixedLambdaCheckpoint::new(
FixedLambdaSolverStage::MultinomialFirth,
coefficients_active.iter().copied().collect(),
p,
m,
iterations,
)
.map_err(|reason| {
EstimationError::InvalidInput(format!(
"multinomial Firth fallback produced an invalid internal checkpoint: {reason}"
))
})?;
return Err(EstimationError::FixedLambdaNewtonDidNotConverge {
context: "fit_penalized_multinomial (Firth/Jeffreys separation refit)".to_string(),
reason: stall_reason,
objective_value: -objective(&probs, &coefficients_active, final_logdet_info),
stationarity: FixedLambdaStationarityEvidence {
kind: FixedLambdaResidualKind::NewtonDecrement,
residual: final_decrement,
bound: tol_eff,
},
checkpoint,
});
}
let (coefficient_covariance, _) = invert_spd(&hmat, "penalized Hessian covariance")?;
Ok(MultinomialFitOutputs {
coefficients_active,
fitted_probabilities: probs,
iterations,
penalized_neg_log_likelihood: -log_likelihood + penalty_term,
deviance: -2.0 * log_likelihood,
coefficient_covariance,
})
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct MultinomialSavedModel {
pub formula: String,
pub class_levels: Vec<String>,
pub reference_class_index: usize,
pub resolved_termspec: TermCollectionSpec,
pub coefficients_flat: Vec<f64>,
pub p_per_class: usize,
pub n_active_classes: usize,
pub training_headers: Vec<String>,
pub training_table_kind: String,
pub lambdas: Vec<f64>,
pub lambdas_per_block: Vec<usize>,
pub iterations: usize,
pub penalized_neg_log_likelihood: f64,
pub deviance: f64,
#[serde(default)]
pub edf_per_class: Option<Vec<f64>>,
#[serde(default)]
pub edf_per_penalty: Option<Vec<f64>>,
pub coefficient_covariance_flat: Vec<f64>,
#[serde(default)]
pub coefficient_influence_flat: Option<Vec<f64>>,
#[serde(default)]
pub smooth_term_spans: Vec<MultinomialSmoothTermSpan>,
pub lambda_labels: Vec<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MultinomialSmoothTermSpan {
pub label: String,
pub col_start: usize,
pub col_end: usize,
pub nullspace_dim: usize,
}
fn penalty_component_label(info: Option<&PenaltyBlockInfo>, pen_idx: usize) -> String {
use gam_terms::basis::PenaltySource;
let term = info
.and_then(|i| i.termname.clone())
.unwrap_or_else(|| format!("s{pen_idx}"));
let role = match info.map(|i| &i.penalty.source) {
Some(PenaltySource::Primary) | None => None,
Some(PenaltySource::DoublePenaltyNullspace) => Some("null space".to_string()),
Some(PenaltySource::OperatorMass) => Some("mass".to_string()),
Some(PenaltySource::OperatorTension) => Some("tension".to_string()),
Some(PenaltySource::OperatorStiffness) => Some("stiffness".to_string()),
Some(PenaltySource::OperatorRelevance { axis }) => Some(format!("axis {axis}")),
Some(PenaltySource::TensorMarginal { dim }) => Some(format!("margin {dim}")),
Some(PenaltySource::TensorSeparable { penalized_margins }) => {
Some(format!("separable {penalized_margins:?}"))
}
Some(PenaltySource::TensorGlobalRidge) => Some("ridge".to_string()),
Some(PenaltySource::Other(s)) => Some(s.clone()),
};
match role {
Some(role) => format!("{term} [{role}]"),
None => term,
}
}
impl MultinomialSavedModel {
pub fn validate(&self) -> Result<(), EstimationError> {
if self.p_per_class == 0 || self.n_active_classes == 0 {
crate::bail_invalid_estim!(
"multinomial saved model dimensions must be nonzero, got P={} and K-1={}",
self.p_per_class,
self.n_active_classes,
);
}
if self.class_levels.len() != self.n_active_classes + 1 {
crate::bail_invalid_estim!(
"multinomial saved model has {} class levels but K-1={}",
self.class_levels.len(),
self.n_active_classes,
);
}
if self.reference_class_index != self.n_active_classes {
crate::bail_invalid_estim!(
"multinomial saved reference index {} does not equal the final class index {}",
self.reference_class_index,
self.n_active_classes,
);
}
let d = self
.p_per_class
.checked_mul(self.n_active_classes)
.ok_or_else(|| {
EstimationError::InvalidInput(
"multinomial saved coefficient dimension overflowed usize".to_string(),
)
})?;
if self.coefficients_flat.len() != d {
crate::bail_invalid_estim!(
"multinomial saved model has {} coefficient values, expected {d}",
self.coefficients_flat.len(),
);
}
if self.training_table_kind.trim().is_empty() {
crate::bail_invalid_estim!(
"multinomial saved model training_table_kind must be non-empty"
);
}
if self.lambdas_per_block.len() != self.n_active_classes {
crate::bail_invalid_estim!(
"multinomial saved model has {} lambda blocks, expected {}",
self.lambdas_per_block.len(),
self.n_active_classes,
);
}
let lambda_count = self
.lambdas_per_block
.iter()
.try_fold(0usize, |total, &count| total.checked_add(count))
.ok_or_else(|| {
EstimationError::InvalidInput(
"multinomial saved lambda count overflowed usize".to_string(),
)
})?;
if lambda_count != self.lambdas.len() {
crate::bail_invalid_estim!(
"multinomial saved model has {} lambdas but its blocks require {lambda_count}",
self.lambdas.len(),
);
}
if self
.lambdas_per_block
.iter()
.any(|&count| count != self.lambda_labels.len())
{
crate::bail_invalid_estim!(
"multinomial saved model has {} lambda labels but block sizes {:?}",
self.lambda_labels.len(),
self.lambdas_per_block,
);
}
if self.lambda_labels.iter().any(|label| label.trim().is_empty()) {
crate::bail_invalid_estim!("multinomial saved model lambda labels must be non-empty");
}
let covariance_len = d.checked_mul(d).ok_or_else(|| {
EstimationError::InvalidInput(
"multinomial saved covariance dimension overflowed usize".to_string(),
)
})?;
if self.coefficient_covariance_flat.len() != covariance_len {
crate::bail_invalid_estim!(
"multinomial saved model has {} covariance values, expected {covariance_len}",
self.coefficient_covariance_flat.len(),
);
}
if let Some((index, value)) = self
.coefficients_flat
.iter()
.chain(self.coefficient_covariance_flat.iter())
.copied()
.enumerate()
.find(|(_, value)| !value.is_finite())
{
crate::bail_invalid_estim!(
"multinomial saved numeric payload is non-finite at combined index {index}: {value}"
);
}
Ok(())
}
pub fn coefficients_active(&self) -> Result<Array2<f64>, EstimationError> {
Array2::from_shape_vec(
(self.p_per_class, self.n_active_classes),
self.coefficients_flat.clone(),
)
.map_err(|error| {
EstimationError::InvalidInput(format!(
"multinomial saved coefficient payload is inconsistent with P x (K-1): {error}"
))
})
}
pub fn coefficient_covariance(&self) -> Result<Array2<f64>, EstimationError> {
let d = self
.p_per_class
.checked_mul(self.n_active_classes)
.ok_or_else(|| {
EstimationError::InvalidInput(
"multinomial saved covariance dimension overflowed usize".to_string(),
)
})?;
Array2::from_shape_vec((d, d), self.coefficient_covariance_flat.clone()).map_err(|error| {
EstimationError::InvalidInput(format!(
"multinomial saved covariance payload is inconsistent with (P*(K-1)) squared: {error}"
))
})
}
pub fn coefficient_influence(&self) -> Option<Array2<f64>> {
let d = self.p_per_class.checked_mul(self.n_active_classes)?;
let flat = self.coefficient_influence_flat.as_ref()?;
Array2::from_shape_vec((d, d), flat.clone()).ok()
}
pub fn predict_probabilities(
&self,
x_new: ArrayView2<'_, f64>,
) -> Result<Array2<f64>, EstimationError> {
self.predict_probabilities_with_se(x_new)
.map(|(mean, _)| mean)
}
pub fn predict_probabilities_with_se(
&self,
x_new: ArrayView2<'_, f64>,
) -> Result<(Array2<f64>, Array2<f64>), EstimationError> {
self.predict_probabilities_with_se_and_control(
x_new,
&MultinomialPosteriorIntegrationControl::default(),
)
}
pub fn predict_probabilities_with_se_and_control(
&self,
x_new: ArrayView2<'_, f64>,
control: &MultinomialPosteriorIntegrationControl,
) -> Result<(Array2<f64>, Array2<f64>), EstimationError> {
let coefficients = self.coefficients_active()?;
let covariance = self.coefficient_covariance()?;
let moments = integrate_multinomial_design_moments(
coefficients.view(),
covariance.view(),
x_new,
control,
)?;
Ok((moments.class_mean, moments.class_standard_deviation))
}
pub fn smooth_significance(&self) -> Vec<MultinomialSmoothSignificance> {
let mut out = Vec::new();
let p = self.p_per_class;
let m = self.n_active_classes;
let Ok(cov) = self.coefficient_covariance() else {
return out;
};
if self.smooth_term_spans.is_empty() {
return out;
}
let Ok(beta) = self.coefficients_active() else {
return out;
};
let d = p * m;
let mut theta = Array1::<f64>::zeros(d);
for a in 0..m {
for i in 0..p {
theta[a * p + i] = beta[[i, a]];
}
}
let influence = self.coefficient_influence();
for a in 0..m {
let class_label = self
.class_levels
.get(a)
.cloned()
.unwrap_or_else(|| format!("class{a}"));
let base = a * p;
for span in &self.smooth_term_spans {
if span.col_end > p {
continue;
}
let start = base + span.col_start;
let end = base + span.col_end;
let block_len = (span.col_end - span.col_start) as f64;
let edf = influence
.as_ref()
.map(|f| (start..end).map(|i| f[[i, i]]).sum::<f64>())
.filter(|v| v.is_finite() && *v > 0.0)
.unwrap_or(block_len);
let result = gam_terms::inference::smooth_test::wood_smooth_test(
gam_terms::inference::smooth_test::SmoothTestInput {
beta: theta.view(),
covariance: &cov,
influence_matrix: influence.as_ref(),
whitening_gram: None,
coeff_range: start..end,
edf,
nullspace_dim: span.nullspace_dim,
residual_df: None,
scale: gam_terms::inference::smooth_test::SmoothTestScale::Known,
},
);
if let Some(res) = result {
out.push(MultinomialSmoothSignificance {
class_label: class_label.clone(),
term_label: span.label.clone(),
edf,
ref_df: res.ref_df,
statistic: res.statistic,
p_value: res.p_value,
});
}
}
}
out
}
pub fn sample_replicate_classes(
&self,
x_new: ArrayView2<'_, f64>,
n_draws: usize,
seed: u64,
) -> Result<Array2<u32>, EstimationError> {
use rand::{RngExt, SeedableRng};
let probs = self.predict_probabilities(x_new)?;
let n = probs.nrows();
let k = probs.ncols();
let mut out = Array2::<u32>::zeros((n_draws, n));
let mut rng = rand::rngs::StdRng::seed_from_u64(seed);
for d in 0..n_draws {
for row in 0..n {
let u: f64 = rng.random::<f64>();
let mut acc = 0.0_f64;
let mut chosen = k - 1; for c in 0..k {
acc += probs[[row, c]];
if u < acc {
chosen = c;
break;
}
}
out[[d, row]] = chosen as u32;
}
}
Ok(out)
}
}
pub const MULTINOMIAL_MODEL_CLASS: &str = "multinomial";
pub const MULTINOMIAL_MODEL_FORMAT_VERSION: u32 = 2;
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct MultinomialModelEnvelope {
pub model_class: String,
pub format_version: u32,
pub saved: MultinomialSavedModel,
}
impl MultinomialModelEnvelope {
pub fn new(saved: MultinomialSavedModel) -> Result<Self, EstimationError> {
saved.validate()?;
Ok(Self {
model_class: MULTINOMIAL_MODEL_CLASS.to_string(),
format_version: MULTINOMIAL_MODEL_FORMAT_VERSION,
saved,
})
}
pub fn to_json_bytes(&self) -> Result<Vec<u8>, EstimationError> {
self.saved.validate()?;
serde_json::to_vec(self).map_err(|err| {
EstimationError::InvalidInput(format!("failed to serialize multinomial model: {err}"))
})
}
pub fn from_json_bytes(bytes: &[u8]) -> Result<Self, EstimationError> {
#[derive(Deserialize)]
struct EnvelopeHeader {
#[serde(default)]
model_class: Option<String>,
#[serde(default)]
format_version: Option<u32>,
}
let header: EnvelopeHeader = serde_json::from_slice(bytes).map_err(|err| {
EstimationError::InvalidInput(format!("failed to deserialize multinomial model: {err}"))
})?;
match header.model_class.as_deref() {
Some(MULTINOMIAL_MODEL_CLASS) => {}
other => {
return Err(EstimationError::InvalidInput(format!(
"multinomial model: model_class = {other:?}, expected {MULTINOMIAL_MODEL_CLASS:?}",
)));
}
}
match header.format_version {
Some(MULTINOMIAL_MODEL_FORMAT_VERSION) => {}
Some(version) => {
return Err(EstimationError::InvalidInput(format!(
"multinomial model: format_version = {version}, expected {MULTINOMIAL_MODEL_FORMAT_VERSION}",
)));
}
None => {
return Err(EstimationError::InvalidInput(format!(
"multinomial model: format_version is absent (unversioned payload), expected {MULTINOMIAL_MODEL_FORMAT_VERSION}",
)));
}
}
let envelope: Self = serde_json::from_slice(bytes).map_err(|err| {
EstimationError::InvalidInput(format!("failed to deserialize multinomial model: {err}"))
})?;
envelope.saved.validate()?;
Ok(envelope)
}
}
#[cfg(test)]
mod multinomial_persistence_contract_tests {
use super::*;
#[test]
fn unversioned_payload_is_rejected() {
let payload = br#"{"model_class":"multinomial","saved":{}}"#;
let error = MultinomialModelEnvelope::from_json_bytes(payload)
.expect_err("unversioned multinomial persistence must not be guessed");
assert!(
error.to_string().contains("format_version"),
"unexpected persistence error: {error}"
);
}
}
#[derive(Debug, Clone)]
pub struct MultinomialSmoothSignificance {
pub class_label: String,
pub term_label: String,
pub edf: f64,
pub ref_df: f64,
pub statistic: f64,
pub p_value: f64,
}
fn one_hot_categorical_response(
data: &EncodedDataset,
y_col: usize,
response_name: &str,
) -> Result<(Array2<f64>, Vec<String>), EstimationError> {
let levels: Vec<String> = data
.schema
.columns
.get(y_col)
.map(|sc| sc.levels.clone())
.unwrap_or_default();
if levels.len() < 2 {
crate::bail_invalid_estim!(
"multinomial response '{response_name}' must have at least 2 categorical levels (got {})",
levels.len()
);
}
let n = data.values.nrows();
let k = levels.len();
let mut y_one_hot = Array2::<f64>::zeros((n, k));
for row in 0..n {
let encoded = data.values[[row, y_col]];
if !encoded.is_finite() {
crate::bail_invalid_estim!(
"multinomial response '{response_name}' row {row} is non-finite ({encoded})"
);
}
let class_idx = encoded.round() as i64;
if class_idx < 0 || (class_idx as usize) >= k {
crate::bail_invalid_estim!(
"multinomial response '{response_name}' row {row} encoded as {encoded} \
is outside the level range 0..{k}"
);
}
y_one_hot[[row, class_idx as usize]] = 1.0;
}
Ok((y_one_hot, levels))
}
fn build_formula_design_for_multinomial(
formula: &str,
data: &EncodedDataset,
config: &FitConfig,
) -> Result<
(
TermCollectionSpec,
TermCollectionDesign,
usize,
String,
ResponseColumnKind,
),
EstimationError,
> {
let parsed = parse_formula(formula).map_err(|err| {
EstimationError::InvalidInput(format!(
"multinomial fit: failed to parse formula {formula:?}: {err}"
))
})?;
let col_map = data.column_map();
let y_col = resolve_role_col(&col_map, &parsed.response, "response")
.map_err(|err| EstimationError::InvalidInput(format!("multinomial fit: {err}")))?;
let y_kind = crate::fit_orchestration::response_column_kind(data, y_col);
let policy = resolved_resource_policy(config, ProblemHints::default());
let mut inference_notes: Vec<String> = Vec::new();
let spec = build_termspec_with_geometry_and_overrides(
&parsed.terms,
data,
&col_map,
&mut inference_notes,
config.scale_dimensions,
&policy,
config.smooth_overrides.as_ref(),
None,
)
.map_err(|err| {
EstimationError::InvalidInput(format!("multinomial fit: build termspec: {err}"))
})?;
let design = build_term_collection_design(data.values.view(), &spec).map_err(|err| {
EstimationError::InvalidInput(format!("multinomial fit: build design: {err}"))
})?;
if design.affine_offset.iter().any(|value| *value != 0.0) {
crate::bail_invalid_estim!(
"multinomial fit does not support non-zero smooth anchors: the reference-coded \
softmax requires an explicit affine offset for every non-reference class"
);
}
Ok((spec, design, y_col, parsed.response, y_kind))
}
fn scale_multinomial_formula_penalty(penalty: PenaltyMatrix, scale: f64) -> PenaltyMatrix {
match penalty {
PenaltyMatrix::Dense(matrix) => PenaltyMatrix::Dense(matrix.mapv(|v| v * scale)),
PenaltyMatrix::KroneckerFactored { left, right } => PenaltyMatrix::KroneckerFactored {
left: left.mapv(|v| v * scale),
right,
},
PenaltyMatrix::Blockwise {
local,
col_range,
total_dim,
} => PenaltyMatrix::Blockwise {
local: local.mapv(|v| v * scale),
col_range,
total_dim,
},
PenaltyMatrix::Labeled { label, inner } => PenaltyMatrix::Labeled {
label,
inner: Box::new(scale_multinomial_formula_penalty(*inner, scale)),
},
PenaltyMatrix::Fixed { log_lambda, inner } => PenaltyMatrix::Fixed {
log_lambda,
inner: Box::new(scale_multinomial_formula_penalty(*inner, scale)),
},
}
}
#[derive(Clone, Copy)]
pub struct MultinomialFitRequest<'a> {
pub data: &'a EncodedDataset,
pub formula: &'a str,
pub config: &'a FitConfig,
pub init_lambda: f64,
pub max_iter: usize,
pub tol: f64,
}
impl<'a> MultinomialFitRequest<'a> {
pub fn new(data: &'a EncodedDataset, formula: &'a str, config: &'a FitConfig) -> Self {
Self {
data,
formula,
config,
init_lambda: 1.0,
max_iter: 50,
tol: 1.0e-7,
}
}
}
fn reject_unsupported_multinomial_config(config: &FitConfig) -> Result<(), EstimationError> {
if config.offset_column.is_some() || config.noise_offset_column.is_some() {
crate::bail_invalid_estim!(
"multinomial fit does not support offset columns: a single offset column has no \
canonical per-logit placement in the reference-coded softmax (offsets are per-class \
linear-predictor quantities); remove the offset or fit per-class models"
);
}
if config.noise_formula.is_some() {
crate::bail_invalid_estim!(
"noise_formula is not supported for the multinomial family: the softmax likelihood \
has no dispersion predictor"
);
}
if config.logslope_formula.is_some() || config.z_column.is_some() {
crate::bail_invalid_estim!(
"logslope_formula/z_column is not supported for the multinomial family"
);
}
if config.transformation_normal {
crate::bail_invalid_estim!(
"transformation_normal conflicts with the multinomial family"
);
}
if config.expectile_tau.is_some() {
crate::bail_invalid_estim!("expectile_tau requires the expectile family");
}
if config.firth {
crate::bail_invalid_estim!(
"manual firth is not accepted for the multinomial family: the Firth/Jeffreys \
separation stabilizer is armed automatically on separation evidence"
);
}
if !matches!(
config.frailty,
crate::survival::lognormal_kernel::FrailtySpec::None
) {
crate::bail_invalid_estim!("frailty is not supported for the multinomial family");
}
Ok(())
}
fn resolve_multinomial_row_weights(
data: &EncodedDataset,
config: &FitConfig,
) -> Result<Array1<f64>, EstimationError> {
let Some(name) = config.weight_column.as_deref() else {
return Ok(Array1::ones(data.values.nrows()));
};
let column = data.column_map().get(name).copied().ok_or_else(|| {
EstimationError::InvalidInput(format!(
"multinomial fit: weight column '{name}' not found in the dataset"
))
})?;
Ok(data.values.column(column).to_owned())
}
pub(crate) struct PenalizedMultinomialFormulaParts {
pub(crate) family: MultinomialFamily,
pub(crate) blocks: Vec<ParameterBlockSpec>,
pub(crate) options: BlockwiseFitOptions,
pub(crate) spec: TermCollectionSpec,
pub(crate) design: TermCollectionDesign,
pub(crate) class_levels: Vec<String>,
pub(crate) parametric_standardization: Vec<(usize, f64, f64)>,
pub(crate) penalties_arc: Arc<Vec<PenaltyMatrix>>,
}
pub(crate) fn penalized_multinomial_formula_parts(
request: &MultinomialFitRequest<'_>,
) -> Result<PenalizedMultinomialFormulaParts, EstimationError> {
let MultinomialFitRequest {
data,
formula,
config,
init_lambda,
max_iter,
tol,
} = *request;
if !(init_lambda.is_finite() && init_lambda > 0.0) {
crate::bail_invalid_estim!(
"multinomial fit: init_lambda must be finite and > 0 (got {init_lambda})"
);
}
reject_unsupported_multinomial_config(config)?;
let (raw_spec, design, y_col, response_name, y_kind) =
build_formula_design_for_multinomial(formula, data, config)?;
let spec = freeze_term_collection_from_design(&raw_spec, &design)?;
let class_levels = match y_kind {
ResponseColumnKind::Categorical { levels } => levels,
ResponseColumnKind::Binary => vec!["0".to_string(), "1".to_string()],
ResponseColumnKind::Numeric => {
crate::bail_invalid_estim!(
"multinomial fit: response '{response_name}' is numeric, not categorical; \
use family='gaussian'/'binomial'/... or convert the column to a categorical type"
);
}
};
if data.column_kinds.get(y_col) == Some(&ColumnKindTag::Binary) {
} else if data.column_kinds.get(y_col) != Some(&ColumnKindTag::Categorical) {
crate::bail_invalid_estim!(
"multinomial fit: response '{response_name}' must be a categorical column \
(got column kind {:?})",
data.column_kinds.get(y_col)
);
}
let (y_one_hot, _) = one_hot_categorical_response(data, y_col, &response_name)?;
let mut x_dense = design
.design
.try_to_dense_by_chunks("multinomial fit design")
.map_err(EstimationError::InvalidInput)?;
let parametric_standardization: Vec<(usize, f64, f64)> =
if design.coefficient_lower_bounds.is_some() || design.linear_constraints.is_some() {
Vec::new()
} else {
let p_total = x_dense.ncols();
let mut penalized = vec![false; p_total];
for bp in &design.penalties {
for col in bp.col_range.clone() {
if col < p_total {
penalized[col] = true;
}
}
}
let has_intercept = !design.intercept_range.is_empty();
let n_rows = x_dense.nrows().max(1) as f64;
let mut standardized = Vec::new();
for (_, range) in &design.linear_ranges {
for col in range.clone() {
if col >= p_total || penalized[col] {
continue;
}
let column = x_dense.column(col);
let mean = column.sum() / n_rows;
let var = column.iter().map(|v| (v - mean) * (v - mean)).sum::<f64>() / n_rows;
let scale = var.sqrt();
if !(scale.is_finite() && scale > 1e-8 * (mean.abs() + 1.0)) {
continue;
}
let center = if has_intercept { mean } else { 0.0 };
for v in x_dense.column_mut(col).iter_mut() {
*v = (*v - center) / scale;
}
standardized.push((col, center, scale));
}
}
standardized
};
let k = y_one_hot.ncols();
let m = k - 1;
let n_obs = y_one_hot.nrows();
let penalty_scale = multinomial_formula_penalty_scale(k);
let per_term_penalties: Vec<PenaltyMatrix> = design
.penalties_as_penalty_matrix()
.into_iter()
.map(|penalty| scale_multinomial_formula_penalty(penalty, penalty_scale))
.collect();
let design_arc = Arc::new(x_dense);
let penalties_arc = Arc::new(per_term_penalties);
let weights = resolve_multinomial_row_weights(data, config)?;
if weights.len() != n_obs {
crate::bail_invalid_estim!(
"multinomial fit: weight column length {} != N = {n_obs}",
weights.len()
);
}
let log_init = init_lambda.ln();
let family = MultinomialFamily::new(
y_one_hot.clone(),
weights,
k,
design_arc.clone(),
penalties_arc.clone(),
)
.map_err(EstimationError::InvalidInput)?
.with_joint_jeffreys_term(false)
.with_initial_log_lambda(log_init);
let mut blocks = family.build_block_specs();
for spec_block in blocks.iter_mut() {
for v in spec_block.initial_log_lambdas.iter_mut() {
*v = log_init;
}
}
let total_rho_dim = m.saturating_mul(penalties_arc.len());
let use_outer_hessian = multinomial_formula_use_outer_hessian(total_rho_dim);
let outer_max_iter = max_iter.max(1);
let outer_tol = if tol.is_finite() && tol > 0.0 {
tol.max(MULTINOMIAL_OUTER_REML_TOL)
} else {
MULTINOMIAL_OUTER_REML_TOL
};
let outer_rel_cost_tol = Some(BlockwiseFitOptions::default().outer_tol);
let inner_tol = MULTINOMIAL_FORMULA_INNER_TOL.max(tol.max(0.0));
let options = BlockwiseFitOptions {
inner_max_cycles: crate::custom_family::DEFAULT_CUSTOM_FAMILY_INNER_MAX_CYCLES,
inner_tol,
outer_max_iter,
outer_tol,
outer_rel_cost_tol,
rho_lower_bound: multinomial_formula_min_lambda(y_one_hot.view()).ln(),
ridge_floor: MULTINOMIAL_FORMULA_RIDGE_FLOOR,
ridge_policy: gam_problem::RidgePolicy::solver_only(),
use_outer_hessian,
screen_initial_rho: false,
compute_covariance: true,
..BlockwiseFitOptions::default()
};
Ok(PenalizedMultinomialFormulaParts {
family,
blocks,
options,
spec,
design,
class_levels,
parametric_standardization,
penalties_arc,
})
}
pub fn fit_penalized_multinomial_formula(
request: &MultinomialFitRequest<'_>,
) -> Result<MultinomialSavedModel, EstimationError> {
let PenalizedMultinomialFormulaParts {
family,
blocks,
options,
spec,
design,
class_levels,
parametric_standardization,
penalties_arc,
} = penalized_multinomial_formula_parts(request)?;
let MultinomialFitRequest {
data,
formula,
config,
..
} = *request;
let m = family.active_classes();
let mut unbiased_probe_options = options.clone();
unbiased_probe_options.outer_max_iter = unbiased_probe_options
.outer_max_iter
.min(MULTINOMIAL_UNBIASED_PROBE_OUTER_MAX_ITER);
let firth_refit_options = &options;
let run_firth_refit = |evidence: String| {
let firth_family = family.clone().with_joint_jeffreys_term(true);
fit_custom_family_with_rho_prior(
&firth_family,
&blocks,
firth_refit_options,
gam_problem::RhoPrior::Flat,
)
.map_err(|err| {
EstimationError::InvalidInput(format!(
"multinomial REML: Firth/Jeffreys-armed refit (separation evidence: \
{evidence}) failed: {err}"
))
})
};
let probe_attempt = fit_custom_family_with_rho_prior(
&family,
&blocks,
&unbiased_probe_options,
gam_problem::RhoPrior::Flat,
);
let fit = match probe_attempt {
Ok(probe_fit) => {
let separation = multinomial_formula_separation_evidence(&probe_fit.block_states);
if separation.is_none() {
probe_fit
} else {
let evidence = separation.expect("checked as present");
run_firth_refit(evidence)?
}
}
Err(err) => run_firth_refit(format!("unbiased-criterion REML solve failed: {err}"))?,
};
if let Some(err) = multinomial_formula_separation_diagnostic(
fit.inner_cycles,
fit.outer_iterations,
&fit.block_states,
) {
return Err(err);
}
if fit.blocks.len() != m {
crate::bail_invalid_estim!(
"multinomial REML: expected {m} fitted blocks (K-1), got {}",
fit.blocks.len()
);
}
let p_per_class = fit.blocks[0].beta.len();
let mut coefficients_active = Array2::<f64>::zeros((p_per_class, m));
for (a, block) in fit.blocks.iter().enumerate() {
if block.beta.len() != p_per_class {
crate::bail_invalid_estim!(
"multinomial REML: block {a} has {} coefs, expected {p_per_class}",
block.beta.len()
);
}
for i in 0..p_per_class {
coefficients_active[[i, a]] = block.beta[i];
}
}
if !parametric_standardization.is_empty() {
let intercept_col = design.intercept_range.clone().next();
for a in 0..m {
let mut intercept_adjust = 0.0;
for &(col, center, scale) in ¶metric_standardization {
if col < p_per_class {
let raw = coefficients_active[[col, a]] / scale;
coefficients_active[[col, a]] = raw;
intercept_adjust += raw * center;
}
}
if let Some(i0) = intercept_col
&& i0 < p_per_class
{
coefficients_active[[i0, a]] -= intercept_adjust;
}
}
}
let joint_recon = fit.artifacts.joint_log_lambdas.as_ref().and_then(|jll| {
let n_components = penalties_arc.len();
if n_components == 0 {
return None;
}
let joint_specs = family.equivariant_class_penalty_specs().ok()?;
if jll.len() != joint_specs.len() || joint_specs.len() % n_components != 0 {
return None;
}
let specs_per_term = joint_specs.len() / n_components;
let expected_joint = p_per_class.saturating_mul(m);
let hinv = fit
.covariance_conditional
.as_ref()
.filter(|c| c.nrows() == expected_joint && c.ncols() == expected_joint)?;
let lam: Vec<f64> = jll.iter().map(|&l| l.exp()).collect();
let mut hinv_st: Vec<Array2<f64>> = Vec::with_capacity(joint_specs.len());
for spec in &joint_specs {
if spec.matrix.nrows() != expected_joint || spec.matrix.ncols() != expected_joint {
return None;
}
hinv_st.push(hinv.dot(&spec.matrix));
}
let mut f = Array2::<f64>::eye(expected_joint);
for (s, hs) in hinv_st.iter().enumerate() {
f.scaled_add(-lam[s], hs);
}
let mut edf_per_class = Vec::with_capacity(m);
let mut edf_per_penalty = Vec::with_capacity(m * n_components);
for a in 0..m {
let base = a * p_per_class;
let mut class_trace = 0.0_f64;
for t in 0..n_components {
let mut tr_at = 0.0_f64;
for c in 0..specs_per_term {
let s = t * specs_per_term + c;
let mut tr = 0.0_f64;
for i in 0..p_per_class {
tr += hinv_st[s][[base + i, base + i]];
}
tr_at += lam[s] * tr;
}
class_trace += tr_at;
let spec0 = &joint_specs[t * specs_per_term];
let joint_rank = expected_joint - spec0.nullspace_dim;
let rank_t = if specs_per_term > 1 {
joint_rank as f64
} else {
(joint_rank as f64) / (m as f64)
};
edf_per_penalty.push((rank_t - tr_at).clamp(0.0, p_per_class as f64));
}
edf_per_class.push((p_per_class as f64 - class_trace).clamp(0.0, p_per_class as f64));
}
let mut lam_flat = Vec::with_capacity(m * n_components);
for a in 0..m {
for t in 0..n_components {
let s = if specs_per_term > 1 {
t * specs_per_term + a
} else {
t
};
lam_flat.push(lam[s]);
}
}
Some((f, edf_per_class, edf_per_penalty, n_components, lam_flat))
});
let (lambdas_per_block, lambdas_flat): (Vec<usize>, Vec<f64>) = match joint_recon.as_ref() {
Some((_, _, _, n_components, lam_flat)) => {
let per_block = vec![*n_components; m];
(per_block, lam_flat.clone())
}
None => {
let per_block: Vec<usize> = fit.blocks.iter().map(|b| b.lambdas.len()).collect();
let flat: Vec<f64> = fit
.blocks
.iter()
.flat_map(|b| b.lambdas.iter().copied())
.collect();
(per_block, flat)
}
};
let edf_per_class = joint_recon
.as_ref()
.map(|(_, epc, _, _, _)| epc.clone())
.or_else(|| {
fit.inference.as_ref().and_then(|info| {
let traces = &info.penalty_block_trace;
if traces.len() != lambdas_per_block.iter().sum::<usize>() {
return None;
}
let mut per_class = Vec::with_capacity(m);
let mut cursor = 0usize;
for &n_blocks in &lambdas_per_block {
let class_trace: f64 = traces[cursor..cursor + n_blocks].iter().sum();
per_class
.push((p_per_class as f64 - class_trace).clamp(0.0, p_per_class as f64));
cursor += n_blocks;
}
Some(per_class)
})
});
let edf_per_penalty = joint_recon
.as_ref()
.map(|(_, _, epp, _, _)| epp.clone())
.or_else(|| {
fit.inference.as_ref().and_then(|info| {
if info.edf_by_block.len() != lambdas_flat.len() {
return None;
}
Some(
info.edf_by_block
.iter()
.map(|&e| e.max(0.0))
.collect::<Vec<f64>>(),
)
})
});
let coefficients_flat: Vec<f64> = coefficients_active.iter().copied().collect();
let expected_joint = p_per_class.checked_mul(m).ok_or_else(|| {
EstimationError::InvalidInput(
"multinomial posterior covariance dimension overflowed usize".to_string(),
)
})?;
let intercept_col0 = design.intercept_range.clone().next();
let build_per_class_affine = |amat: &mut Array2<f64>| {
for &(col, center, scale) in ¶metric_standardization {
if col >= p_per_class {
continue;
}
amat[[col, col]] = 1.0 / scale;
if let Some(i0) = intercept_col0
&& i0 < p_per_class
{
amat[[i0, col]] = -center / scale;
}
}
};
let coefficient_covariance_flat = fit
.covariance_conditional
.as_ref()
.filter(|c| c.nrows() == expected_joint && c.ncols() == expected_joint)
.map(|cov_std| {
if parametric_standardization.is_empty() {
return cov_std.iter().copied().collect::<Vec<f64>>();
}
let mut a_joint = Array2::<f64>::eye(expected_joint);
let mut a_class = Array2::<f64>::eye(p_per_class);
build_per_class_affine(&mut a_class);
for a in 0..m {
let base = a * p_per_class;
for i in 0..p_per_class {
for j in 0..p_per_class {
a_joint[[base + i, base + j]] = a_class[[i, j]];
}
}
}
let cov_raw = a_joint.dot(cov_std).dot(&a_joint.t());
cov_raw.iter().copied().collect::<Vec<f64>>()
})
.ok_or_else(|| {
EstimationError::InvalidInput(format!(
"multinomial REML converged without the required {expected_joint}x{expected_joint} joint posterior covariance"
))
})?;
let coefficient_influence_flat = match joint_recon.as_ref() {
Some((f, _, _, _, _)) => Some(f.iter().copied().collect::<Vec<f64>>()),
None => fit
.covariance_conditional
.as_ref()
.filter(|c| c.nrows() == expected_joint && c.ncols() == expected_joint)
.and_then(|hinv| {
if fit.blocks.len() != m {
return None;
}
let mut s_lambda = Array2::<f64>::zeros((expected_joint, expected_joint));
for (a, block) in fit.blocks.iter().enumerate() {
if block.lambdas.len() != penalties_arc.len() {
return None;
}
let base = a * p_per_class;
for (t, pen) in penalties_arc.iter().enumerate() {
let lam = block.lambdas[t];
if lam == 0.0 {
continue;
}
let dense = pen.to_dense();
if dense.nrows() != p_per_class || dense.ncols() != p_per_class {
return None;
}
for i in 0..p_per_class {
for j in 0..p_per_class {
s_lambda[[base + i, base + j]] += lam * dense[[i, j]];
}
}
}
}
let hinv_s = hinv.dot(&s_lambda);
let mut f = Array2::<f64>::eye(expected_joint);
f -= &hinv_s;
Some(f.iter().copied().collect::<Vec<f64>>())
}),
};
let mut smooth_term_spans: Vec<MultinomialSmoothTermSpan> = Vec::new();
for (pen_idx, bp) in design.penalties.iter().enumerate() {
let col_start = bp.col_range.start;
let col_end = bp.col_range.end;
if col_start >= col_end || col_end > p_per_class {
continue;
}
if smooth_term_spans
.iter()
.any(|s| s.col_start == col_start && s.col_end == col_end)
{
continue;
}
let label = design
.penaltyinfo
.get(pen_idx)
.and_then(|info| info.termname.clone())
.unwrap_or_else(|| format!("s{pen_idx}"));
let nullspace_dim = design
.nullspace_dims
.get(pen_idx)
.copied()
.unwrap_or(0)
.min(col_end - col_start);
smooth_term_spans.push(MultinomialSmoothTermSpan {
label,
col_start,
col_end,
nullspace_dim,
});
}
let lambda_labels: Vec<String> = design
.penalties
.iter()
.enumerate()
.map(|(pen_idx, _)| penalty_component_label(design.penaltyinfo.get(pen_idx), pen_idx))
.collect();
let deviance = -2.0 * fit.log_likelihood;
Ok(MultinomialSavedModel {
formula: formula.to_string(),
class_levels: class_levels.clone(),
reference_class_index: class_levels.len() - 1,
resolved_termspec: spec,
coefficients_flat,
p_per_class,
n_active_classes: m,
training_headers: data.headers.clone(),
training_table_kind: config.training_table_kind.clone(),
lambdas: lambdas_flat,
lambdas_per_block,
iterations: fit.inner_cycles,
penalized_neg_log_likelihood: -fit.log_likelihood + 0.5 * fit.stable_penalty_term,
deviance,
edf_per_class,
edf_per_penalty,
coefficient_covariance_flat,
coefficient_influence_flat,
smooth_term_spans,
lambda_labels,
})
}
fn build_multinomial_predict_design(
model: &MultinomialSavedModel,
data: &EncodedDataset,
) -> Result<Array2<f64>, EstimationError> {
let predict_columns = data.column_map();
let realigned = model.resolved_termspec.remap_feature_columns(
|index| -> Result<usize, EstimationError> {
let name = model.training_headers.get(index).ok_or_else(|| {
EstimationError::InvalidInput(format!(
"multinomial predict: saved training column index {index} is out of bounds \
for {} training headers",
model.training_headers.len()
))
})?;
resolve_role_col(&predict_columns, name, "feature")
.map_err(|err| EstimationError::InvalidInput(err.to_string()))
},
)?;
let design = build_term_collection_design(data.values.view(), &realigned).map_err(|err| {
EstimationError::InvalidInput(format!(
"multinomial predict: rebuild design from saved termspec: {err}"
))
})?;
if design.affine_offset.iter().any(|value| *value != 0.0) {
crate::bail_invalid_estim!(
"multinomial predict does not support non-zero smooth anchors: the saved \
reference-coded softmax has no per-class affine offset channel"
);
}
let x_dense = design
.design
.try_to_dense_by_chunks("multinomial predict design")
.map_err(EstimationError::InvalidInput)?;
if x_dense.ncols() != model.p_per_class {
crate::bail_invalid_estim!(
"multinomial predict: predict design has {} cols, saved model expects {}",
x_dense.ncols(),
model.p_per_class
);
}
Ok(x_dense)
}
pub fn predict_multinomial_formula(
model: &MultinomialSavedModel,
data: &EncodedDataset,
) -> Result<Array2<f64>, EstimationError> {
model.validate()?;
let x_dense = build_multinomial_predict_design(model, data)?;
model.predict_probabilities(x_dense.view())
}
pub fn posterior_predict_multinomial_formula(
model: &MultinomialSavedModel,
data: &EncodedDataset,
n_draws: usize,
seed: u64,
) -> Result<Array2<u32>, EstimationError> {
if n_draws == 0 {
crate::bail_invalid_estim!("multinomial posterior_predict: n_draws must be >= 1");
}
model.validate()?;
let x_dense = build_multinomial_predict_design(model, data)?;
model.sample_replicate_classes(x_dense.view(), n_draws, seed)
}
pub fn predict_multinomial_formula_with_se(
model: &MultinomialSavedModel,
data: &EncodedDataset,
) -> Result<(Array2<f64>, Array2<f64>), EstimationError> {
model.validate()?;
let x_dense = build_multinomial_predict_design(model, data)?;
model.predict_probabilities_with_se(x_dense.view())
}
#[derive(Debug, Clone)]
pub struct MultinomialPredictionIntervals {
pub mean: Array2<f64>,
pub standard_error: Array2<f64>,
pub mean_lower: Array2<f64>,
pub mean_upper: Array2<f64>,
pub level: f64,
}
pub fn predict_multinomial_formula_with_intervals(
model: &MultinomialSavedModel,
data: &EncodedDataset,
level: f64,
) -> Result<MultinomialPredictionIntervals, EstimationError> {
if !(level.is_finite() && level > 0.0 && level < 1.0) {
crate::bail_invalid_estim!(
"multinomial prediction interval level must be finite and in (0, 1), got {level}"
);
}
let (mean, standard_error) = predict_multinomial_formula_with_se(model, data)?;
let z = gam_math::probability::standard_normal_quantile(0.5 + 0.5 * level)
.map_err(EstimationError::InvalidInput)?;
let mut mean_lower = mean.clone();
let mut mean_upper = mean.clone();
for ((row, class), &se) in standard_error.indexed_iter() {
mean_lower[[row, class]] = (mean[[row, class]] - z * se).clamp(0.0, 1.0);
mean_upper[[row, class]] = (mean[[row, class]] + z * se).clamp(0.0, 1.0);
}
Ok(MultinomialPredictionIntervals {
mean,
standard_error,
mean_lower,
mean_upper,
level,
})
}
#[cfg(test)]
mod fisher_override_tests {
use super::*;
fn multinomial_formula_unresolved_probe_separation_evidence(
block_states: &[ParameterBlockState],
) -> Option<String> {
if let Some(evidence) = multinomial_formula_separation_evidence(block_states) {
return Some(evidence);
}
let mut best = (0.0_f64, 0usize, 0usize);
for (active_class, state) in block_states.iter().enumerate() {
for (row, &value) in state.eta.iter().enumerate() {
let abs = value.abs();
if abs > best.0 {
best = (abs, row, active_class);
}
}
}
if best.0 >= MULTINOMIAL_SEPARATION_ETA_THRESHOLD {
Some(format!(
"separation-scale finite logit |eta[row {}, active class {}]| = {:.3e} \
after capped unbiased probe",
best.1, best.2, best.0
))
} else {
None
}
}
use ndarray::Array3;
fn toy() -> (Array2<f64>, Array2<f64>, Array2<f64>, Array1<f64>) {
let n = 15;
let p = 2;
let k = 3;
let design =
Array2::<f64>::from_shape_fn(
(n, p),
|(i, j)| {
if j == 0 { 1.0 } else { ((i + 2) as f64).cos() }
},
);
let mut y = Array2::<f64>::zeros((n, k));
for i in 0..n {
y[[i, i % k]] = 1.0;
}
let penalty = Array2::<f64>::eye(p);
let lambdas = Array1::<f64>::from_elem(k, 0.5);
(design, y, penalty, lambdas)
}
#[test]
fn fisher_override_none_reproduces_analytic() {
let (design, y, penalty, lambdas) = toy();
let mk = |over: Option<ndarray::ArrayView3<'_, f64>>| {
fit_penalized_multinomial(MultinomialFitInputs {
design: design.view(),
y_one_hot: y.view(),
penalty: penalty.view(),
lambdas: lambdas.view(),
row_weights: None,
fisher_w_override: over,
max_iter: 50,
tol: 1.0e-9,
resume_from: None,
})
.expect("fit must succeed")
};
let a = mk(None);
let b = mk(None);
for (x, z) in a
.coefficients_active
.iter()
.zip(b.coefficients_active.iter())
{
assert_eq!(x, z);
}
}
#[test]
fn exhausted_fixed_lambda_budget_is_typed_error_not_fit() {
let (design, y, penalty, lambdas) = toy();
let error = 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: 0,
tol: 1.0e-9,
resume_from: None,
})
.expect_err("a zero-budget Newton solve must not mint a multinomial fit");
assert!(matches!(
error,
EstimationError::FixedLambdaNewtonDidNotConverge {
objective_value,
checkpoint,
..
} if objective_value.is_finite()
&& checkpoint.stage() == FixedLambdaSolverStage::MultinomialNewton
&& checkpoint.completed_iterations() == 0
));
}
#[test]
fn fixed_lambda_checkpoint_resume_matches_uninterrupted_solve() {
let (design, y, penalty, lambdas) = toy();
let interrupted = 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: 1,
tol: 1.0e-9,
resume_from: None,
})
.expect_err("one Newton step must leave this coupled fit uncertified");
let checkpoint = match interrupted {
EstimationError::FixedLambdaNewtonDidNotConverge { checkpoint, .. } => checkpoint,
other => panic!("unexpected interruption error: {other}"),
};
assert_eq!(
checkpoint.stage(),
FixedLambdaSolverStage::MultinomialNewton
);
assert_eq!(checkpoint.completed_iterations(), 1);
let resumed = 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: 49,
tol: 1.0e-9,
resume_from: Some(&checkpoint),
})
.expect("resumed multinomial solve must converge");
let uninterrupted = 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: 50,
tol: 1.0e-9,
resume_from: None,
})
.expect("uninterrupted multinomial solve must converge");
assert_eq!(resumed.iterations, uninterrupted.iterations);
assert_eq!(
resumed.coefficients_active,
uninterrupted.coefficients_active
);
assert_eq!(
resumed.penalized_neg_log_likelihood,
uninterrupted.penalized_neg_log_likelihood,
);
assert_eq!(
resumed.coefficient_covariance,
uninterrupted.coefficient_covariance,
);
}
#[test]
fn fisher_override_wrong_shape_is_rejected() {
let (design, y, penalty, lambdas) = toy();
let n = design.nrows();
let m = y.ncols(); let bad = Array3::<f64>::zeros((n, m, m));
let err = fit_penalized_multinomial(MultinomialFitInputs {
design: design.view(),
y_one_hot: y.view(),
penalty: penalty.view(),
lambdas: lambdas.view(),
row_weights: None,
fisher_w_override: Some(bad.view()),
max_iter: 50,
tol: 1.0e-9,
resume_from: None,
})
.expect_err("wrong active-block shape must error");
assert!(format!("{err}").contains("fisher_w_override shape"));
}
#[test]
fn covariance_and_delta_method_se_are_finite_and_wellformed_1101() {
let (design, y, penalty, lambdas) = toy();
let p = design.ncols();
let k = y.ncols();
let m = k - 1;
let d = p * m;
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: 50,
tol: 1.0e-9,
resume_from: None,
})
.expect("fit must succeed");
let cov = &fit.coefficient_covariance;
assert_eq!(
cov.dim(),
(d, d),
"covariance must be (P·(K−1))² = ({d},{d})"
);
for &v in cov.iter() {
assert!(v.is_finite(), "covariance entry must be finite (got {v})");
}
for i in 0..d {
for j in 0..d {
let asym = (cov[[i, j]] - cov[[j, i]]).abs();
assert!(
asym <= 1e-9 * (1.0 + cov[[i, j]].abs()),
"covariance must be symmetric at ({i},{j}): |Σ_ij − Σ_ji| = {asym:.3e}"
);
}
}
for i in 0..d {
assert!(
cov[[i, i]] >= 0.0,
"covariance diagonal[{i}] must be ≥ 0 (got {})",
cov[[i, i]]
);
}
let mut probes: Vec<Vec<f64>> = Vec::new();
for i in 0..d {
let mut e = vec![0.0_f64; d];
e[i] = 1.0;
probes.push(e);
}
probes.push(vec![1.0_f64; d]);
for v in &probes {
let mut q = 0.0_f64;
for i in 0..d {
for j in 0..d {
q += v[i] * cov[[i, j]] * v[j];
}
}
assert!(q >= -1e-9, "covariance must be PSD: vᵀΣv = {q:.3e} < 0");
}
let (probs, prob_se) = fit
.predict_probabilities_with_se(design.view())
.expect("delta-method SE must succeed");
let n = design.nrows();
assert_eq!(probs.dim(), (n, k));
assert_eq!(prob_se.dim(), (n, k));
for row in 0..n {
let mut rowsum = 0.0_f64;
for c in 0..k {
let pc = probs[[row, c]];
assert!(
pc.is_finite() && (0.0..=1.0).contains(&pc),
"prob[{row},{c}]={pc}"
);
rowsum += pc;
let se = prob_se[[row, c]];
assert!(
se.is_finite(),
"prob_se[{row},{c}] must be finite (got {se})"
);
assert!(
(0.0..=1.0).contains(&se),
"prob_se[{row},{c}] must be in [0,1] (got {se})"
);
}
assert!(
(rowsum - 1.0).abs() < 1e-9,
"row {row} probabilities must sum to 1 (got {rowsum})"
);
}
}
#[test]
fn formula_outer_route_uses_exact_curvature_for_medium_d() {
assert!(
multinomial_formula_use_outer_hessian(8),
"D=8 loaded multinomial fits need exact curvature to avoid over-smoothed lambda caps"
);
assert!(
multinomial_formula_use_outer_hessian(12),
"D=12 (3 double-penalty smooth terms, K=3) stays on exact curvature"
);
}
#[test]
fn formula_outer_route_uses_exact_curvature_for_d16_penguin_fixture() {
assert!(
multinomial_formula_use_outer_hessian(16),
"D=16 multinomial fits need exact ARC curvature for the #1082 stall halt"
);
}
#[test]
fn formula_min_lambda_floor_is_continuous_and_information_scaled() {
fn floor_for_min_count(count: usize) -> f64 {
let n = 1000 + count;
let mut y = Array2::<f64>::zeros((n, 2));
for r in 0..1000 {
y[[r, 0]] = 1.0;
}
for r in 1000..n {
y[[r, 1]] = 1.0;
}
multinomial_formula_min_lambda(y.view())
}
let base = MULTINOMIAL_FORMULA_PRIOR_PSEUDO_OBS * MULTINOMIAL_FORMULA_FISHER_INFO_PER_OBS;
let sparse = MULTINOMIAL_FORMULA_SPARSE_PRIOR_PSEUDO_OBS_MAX
* MULTINOMIAL_FORMULA_FISHER_INFO_PER_OBS;
assert!(
(base - 2.0e-4).abs() < 1e-18,
"derived base floor must equal the calibrated 2e-4"
);
assert!(
(sparse - 1.0e-3).abs() < 1e-18,
"derived sparse floor must equal the calibrated 1e-3"
);
assert!((floor_for_min_count(50) - base).abs() < 1e-18);
assert!((floor_for_min_count(200) - base).abs() < 1e-18);
assert!((floor_for_min_count(10) - sparse).abs() < 1e-18);
assert!((floor_for_min_count(5) - sparse).abs() < 1e-18);
let f49 = floor_for_min_count(49);
let f50 = floor_for_min_count(50);
assert!(
f49 >= f50 && f49 <= f50 * 1.05,
"floor must be continuous across c0, got {f49} vs {f50}"
);
let f25 = floor_for_min_count(25);
assert!(
f25 > f50 && f25 < floor_for_min_count(10),
"mid-support floor must interpolate strictly between the two endpoints"
);
for &n_c in &[12usize, 16, 20, 30, 40] {
let expected = base * (MULTINOMIAL_FORMULA_SPARSE_REFERENCE_SUPPORT / n_c as f64);
assert!(
(floor_for_min_count(n_c) - expected).abs() < 1e-15,
"floor at n_c={n_c} must be τ·I₁·n_ref/n_c = {expected}, got {}",
floor_for_min_count(n_c)
);
}
assert!(
(floor_for_min_count(20) - 2.0 * floor_for_min_count(40)).abs() < 1e-15,
"floor must scale like 1/n_c (effective Fisher information) in the interior band"
);
}
#[test]
fn formula_penalty_scale_tracks_softmax_fisher_curvature() {
assert!(
(multinomial_formula_penalty_scale(2) - 0.5).abs() < 1.0e-12,
"binary-logit neutral-simplex curvature scale should remain at 1/2"
);
assert!(
(multinomial_formula_penalty_scale(3) - 4.0 / 9.0).abs() < 1.0e-12,
"three-class softmax penalties should be calibrated to 2*(K-1)/K^2"
);
assert!(
multinomial_formula_penalty_scale(5) < multinomial_formula_penalty_scale(3),
"active-class Fisher curvature decreases as the simplex gains classes"
);
}
#[test]
fn fixed_lambda_multinomial_firth_keeps_complete_separation_finite() {
let n = 90;
let design = Array2::<f64>::from_shape_fn((n, 2), |(row, col)| match col {
0 => 1.0,
_ => -3.0 + 6.0 * (row as f64) / ((n - 1) as f64),
});
let mut y = Array2::<f64>::zeros((n, 3));
for row in 0..n {
let x = design[[row, 1]];
let class = if x < -1.0 {
0
} else if x > 1.0 {
1
} else {
2
};
y[[row, class]] = 1.0;
}
let penalty = Array2::<f64>::zeros((2, 2));
let lambdas = Array1::<f64>::zeros(3);
let out = 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: 80,
tol: 1.0e-12,
resume_from: None,
})
.expect("Firth/Jeffreys prior keeps the separated multinomial fit finite (#1854)");
for &b in out.coefficients_active.iter() {
assert!(
b.is_finite(),
"Firth-penalized coefficients must be finite, got {b}"
);
}
for row in 0..n {
let mut mass = 0.0_f64;
for c in 0..3 {
let p = out.fitted_probabilities[[row, c]];
assert!(
p.is_finite() && (0.0..=1.0 + 1e-9).contains(&p),
"row {row} class {c} probability {p} out of [0,1]"
);
mass += p;
}
assert!(
(mass - 1.0).abs() < 1e-6,
"row {row} probabilities must sum to 1, got {mass}"
);
}
let predict = |x: f64| -> usize {
let mut eta = [0.0_f64; 3];
for a in 0..2 {
eta[a] = out.coefficients_active[[0, a]] + out.coefficients_active[[1, a]] * x;
}
let mut best = 0usize;
for c in 1..3 {
if eta[c] > eta[best] {
best = c;
}
}
best
};
assert_eq!(predict(-2.5), 0, "deep-left region should predict class 0");
assert_eq!(predict(2.5), 1, "deep-right region should predict class 1");
assert_eq!(predict(0.0), 2, "central region should predict class 2");
}
#[test]
fn formula_multinomial_accepts_finite_saturated_logits() {
let saturated_states = vec![
ParameterBlockState {
beta: Array1::from_vec(vec![1.0, 2.0]),
eta: Array1::from_vec(vec![0.2, 4.0, -7.0]),
},
ParameterBlockState {
beta: Array1::from_vec(vec![-1.0, 3.0]),
eta: Array1::from_vec(vec![1.0, 25.5, -0.1]),
},
];
assert!(
multinomial_formula_separation_diagnostic(17, 9, &saturated_states).is_none(),
"a finite (even saturated, |eta|>25) formula optimum is a valid fit, \
not a separation diagnostic"
);
let blown_up = vec![
ParameterBlockState {
beta: Array1::from_vec(vec![1.0, 2.0]),
eta: Array1::from_vec(vec![0.2, 4.0, -7.0]),
},
ParameterBlockState {
beta: Array1::from_vec(vec![-1.0, 3.0]),
eta: Array1::from_vec(vec![1.0, f64::INFINITY, -0.1]),
},
];
let err = multinomial_formula_separation_diagnostic(17, 9, &blown_up)
.expect("a non-finite formula logit must raise the separation diagnostic");
assert!(
matches!(
err,
EstimationError::MultinomialSeparationDetected {
iteration: 17,
max_abs_eta,
active_class_index: 1,
row_index: 1,
} if !max_abs_eta.is_finite()
),
"expected typed multinomial separation diagnostic at the non-finite channel, got {err:?}"
);
}
#[test]
fn separation_evidence_gate_arms_firth_only_on_blowup() {
let interior = vec![
ParameterBlockState {
beta: Array1::from_vec(vec![1.0, 2.0]),
eta: Array1::from_vec(vec![0.2, 4.0, -7.0]),
},
ParameterBlockState {
beta: Array1::from_vec(vec![-1.0, 3.0]),
eta: Array1::from_vec(vec![1.0, -3.5, -0.1]),
},
];
assert!(
multinomial_formula_separation_evidence(&interior).is_none(),
"an interior finite mode must not arm the Firth refit"
);
let saturated = vec![
ParameterBlockState {
beta: Array1::from_vec(vec![1.0, 2.0]),
eta: Array1::from_vec(vec![0.2, 4.0, -7.0]),
},
ParameterBlockState {
beta: Array1::from_vec(vec![-1.0, 3.0]),
eta: Array1::from_vec(vec![1.0, 25.5, -0.1]),
},
];
assert!(
multinomial_formula_separation_evidence(&saturated).is_none(),
"a finite saturated formula-mode logit must not arm the Firth refit"
);
let blown_up = vec![ParameterBlockState {
beta: Array1::from_vec(vec![1.0, 2.0]),
eta: Array1::from_vec(vec![0.2, f64::NAN, -7.0]),
}];
let evidence = multinomial_formula_separation_evidence(&blown_up)
.expect("a non-finite logit is separation evidence");
assert!(
evidence.contains("non-finite logit") && evidence.contains("row 1"),
"evidence must name the non-finite logit, got {evidence}"
);
let near = vec![ParameterBlockState {
beta: Array1::from_vec(vec![1.0, 2.0]),
eta: Array1::from_vec(vec![0.2, 24.9, -24.9]),
}];
assert!(
multinomial_formula_separation_evidence(&near).is_none(),
"logits below the saturation threshold must not arm the Firth refit"
);
}
#[test]
fn unresolved_probe_evidence_arms_firth_on_saturated_finite_logits() {
let saturated = vec![
ParameterBlockState {
beta: Array1::from_vec(vec![1.0, 2.0]),
eta: Array1::from_vec(vec![0.2, 4.0, -7.0]),
},
ParameterBlockState {
beta: Array1::from_vec(vec![-1.0, 3.0]),
eta: Array1::from_vec(vec![1.0, 25.5, -0.1]),
},
];
assert!(
multinomial_formula_separation_evidence(&saturated).is_none(),
"a converged finite saturated formula optimum remains unbiased"
);
let evidence = multinomial_formula_unresolved_probe_separation_evidence(&saturated)
.expect("a non-converged saturated probe should arm the Firth refit");
assert!(
evidence.contains("separation-scale finite logit")
&& evidence.contains("row 1")
&& evidence.contains("active class 1"),
"unresolved-probe evidence should name the saturated channel, got {evidence}"
);
let near = vec![ParameterBlockState {
beta: Array1::from_vec(vec![1.0, 2.0]),
eta: Array1::from_vec(vec![0.2, 24.9, -24.9]),
}];
assert!(
multinomial_formula_unresolved_probe_separation_evidence(&near).is_none(),
"finite logits below the separation threshold still get the full unbiased retry"
);
}
#[test]
fn scaled_fisher_override_changes_first_step() {
let (design, y, penalty, lambdas) = toy();
let n = design.nrows();
let m = y.ncols() - 1;
let engine_lambdas = Array1::<f64>::from_elem(m, lambdas[0]);
let pk = 1.0 / (y.ncols() as f64);
let mut over = Array3::<f64>::zeros((n, m, m));
for row in 0..n {
for a in 0..m {
for b in 0..m {
let analytic = if a == b { pk * (1.0 - pk) } else { -pk * pk };
over[[row, a, b]] = 4.0 * analytic;
}
}
}
let likelihood =
MultinomialLogitLikelihood::with_classes(y.ncols()).expect("test class count is valid");
let scaled = fit_penalized_vector_glm(
PenalizedVectorGlmInputs {
design: design.view(),
y: y.view(),
penalty: penalty.view(),
lambdas: engine_lambdas.view(),
fisher_w_override: Some(over.view()),
max_iter: 1,
tol: 1.0e-9,
class_penalty_metric: crate::penalized_vector_glm::ClassPenaltyMetric::Centered,
resume_from: None,
},
&likelihood,
"multinomial scaled-curvature first-step test",
)
.expect("scaled-curvature engine step must be finite");
let analytic = fit_penalized_vector_glm(
PenalizedVectorGlmInputs {
design: design.view(),
y: y.view(),
penalty: penalty.view(),
lambdas: engine_lambdas.view(),
fisher_w_override: None,
max_iter: 1,
tol: 1.0e-9,
class_penalty_metric: crate::penalized_vector_glm::ClassPenaltyMetric::Centered,
resume_from: None,
},
&likelihood,
"multinomial analytic-curvature first-step test",
)
.expect("analytic-curvature engine step must be finite");
let checkpoint_coefficients = |solve| match solve {
VectorGlmSolve::Converged(fit) => fit.coefficients,
VectorGlmSolve::Stalled(stall) => stall.coefficients,
};
let scaled = checkpoint_coefficients(scaled);
let analytic = checkpoint_coefficients(analytic);
let differs = scaled
.iter()
.zip(analytic.iter())
.any(|(a, b)| (a - b).abs() > 1.0e-6);
assert!(differs, "scaled curvature must change the first step");
}
}
#[cfg(test)]
mod separation_firth_tests {
use super::*;
fn separated_three_class() -> (Array2<f64>, Array2<f64>, Array2<f64>, Array1<f64>) {
let n = 21;
let p = 2; let k = 3;
let mut design = Array2::<f64>::zeros((n, p));
let mut y = Array2::<f64>::zeros((n, k));
for i in 0..n {
let x = -3.0 + 6.0 * (i as f64) / ((n - 1) as f64);
design[[i, 0]] = 1.0;
design[[i, 1]] = x;
let cls = if x < -1.0 {
0
} else if x < 1.0 {
1
} else {
2
};
y[[i, cls]] = 1.0;
}
let penalty = Array2::<f64>::zeros((p, p));
let lambdas = Array1::<f64>::from_elem(k, 1.0);
(design, y, penalty, lambdas)
}
#[test]
fn separation_engages_firth_finite_converged_fit() {
let (design, y, penalty, lambdas) = separated_three_class();
let out = 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: 300,
tol: 1e-10,
resume_from: None,
})
.expect("separated multinomial must engage Firth and return a fit, not error");
assert!(
out.coefficients_active.iter().all(|v| v.is_finite()),
"all coefficients must be finite under the Firth prior"
);
assert!(out.deviance.is_finite(), "deviance must be finite");
for v in out.fitted_probabilities.iter() {
assert!(
*v > 0.0 && *v < 1.0,
"Firth fit must stay interior, got p={v}"
);
}
let n = design.nrows();
let k = y.ncols();
for i in 0..n {
let mut best = 0usize;
for c in 1..k {
if out.fitted_probabilities[[i, c]] > out.fitted_probabilities[[i, best]] {
best = c;
}
}
let truth = (0..k)
.find(|&c| y[[i, c]] == 1.0)
.expect("one-hot truth class");
assert_eq!(best, truth, "row {i} misclassified under separation");
}
}
#[test]
fn separation_firth_returns_finite_wellshaped_covariance() {
let (design, y, penalty, lambdas) = separated_three_class();
let p = design.ncols();
let k = y.ncols();
let m = k - 1;
let out = 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: 300,
tol: 1e-10,
resume_from: None,
})
.expect("separated multinomial must return a Firth fit");
assert_eq!(
out.coefficient_covariance.dim(),
(p * m, p * m),
"covariance must be P·M square"
);
assert!(
out.coefficient_covariance.iter().all(|v| v.is_finite()),
"Firth covariance entries must be finite"
);
for i in 0..(p * m) {
assert!(
out.coefficient_covariance[[i, i]] >= -1e-9,
"covariance diagonal must be non-negative, got {}",
out.coefficient_covariance[[i, i]]
);
}
}
#[test]
fn firth_solver_rejects_a_truncated_iterate() {
let (design, y, penalty, lambdas) = separated_three_class();
let truncated = fit_penalized_multinomial_firth_fallback(
design.view(),
y.view(),
penalty.view(),
lambdas.view(),
None,
1, 1e-12,
None,
)
.expect_err("a one-iteration Firth solve must not mint a fit");
let checkpoint = match truncated {
EstimationError::FixedLambdaNewtonDidNotConverge {
objective_value,
stationarity,
checkpoint,
..
} => {
assert!(objective_value.is_finite());
assert_eq!(stationarity.kind, FixedLambdaResidualKind::NewtonDecrement);
assert_eq!(checkpoint.stage(), FixedLambdaSolverStage::MultinomialFirth);
assert_eq!(checkpoint.completed_iterations(), 1);
checkpoint
}
other => panic!("unexpected Firth interruption error: {other}"),
};
let resumed = 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: 299,
tol: 1e-10,
resume_from: Some(&checkpoint),
})
.expect("Firth checkpoint must resume to the certified mode");
let uninterrupted = fit_penalized_multinomial_firth_fallback(
design.view(),
y.view(),
penalty.view(),
lambdas.view(),
None,
300,
1e-10,
None,
)
.expect("Firth fallback must converge under a full budget");
assert_eq!(resumed.iterations, uninterrupted.iterations);
assert_eq!(
resumed.coefficients_active,
uninterrupted.coefficients_active
);
assert_eq!(
resumed.penalized_neg_log_likelihood,
uninterrupted.penalized_neg_log_likelihood,
);
assert_eq!(
resumed.coefficient_covariance,
uninterrupted.coefficient_covariance,
);
}
}
#[cfg(test)]
mod reference_class_invariance_tests {
use super::*;
use gam_data::load_dataset_projected;
use std::fmt::Write as _;
use std::fs;
use tempfile::tempdir;
struct SplitMix64(u64);
impl SplitMix64 {
fn next_u64(&mut self) -> u64 {
self.0 = self.0.wrapping_add(0x9E37_79B9_7F4A_7C15);
let mut z = self.0;
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
z ^ (z >> 31)
}
fn unit(&mut self) -> f64 {
(self.next_u64() >> 11) as f64 / (1u64 << 53) as f64
}
}
fn sample_classes(seed: u64, n: usize) -> (Vec<f64>, Vec<usize>) {
let mut rng = SplitMix64(seed.wrapping_add(0x1234_5678));
let mut x = Vec::with_capacity(n);
let mut cls = Vec::with_capacity(n);
for _ in 0..n {
let xi = -2.0 + 4.0 * rng.unit();
let eta = [0.5 + 0.8 * xi, -0.3 - 0.5 * xi, 0.0];
let mut p = [eta[0].exp(), eta[1].exp(), eta[2].exp()];
let s: f64 = p.iter().sum();
for v in &mut p {
*v /= s;
}
let u = rng.unit();
let c = if u < p[0] {
0
} else if u < p[0] + p[1] {
1
} else {
2
};
x.push(xi);
cls.push(c);
}
(x, cls)
}
fn dataset_xy(
dir: &std::path::Path,
tag: &str,
x: &[f64],
y: &[String],
) -> gam_data::EncodedDataset {
let path = dir.join(format!("data_{tag}.csv"));
let mut csv = String::from("x,y\n");
for (xi, yi) in x.iter().zip(y.iter()) {
writeln!(csv, "{xi},{yi}").unwrap();
}
fs::write(&path, csv).expect("write training csv");
load_dataset_projected(&path, &["x".to_string(), "y".to_string()])
.expect("load training dataset")
}
fn fit_predict_aligned(
dir: &std::path::Path,
tag: &str,
x: &[f64],
cls: &[usize],
name_map: [&str; 3],
grid: &[f64],
) -> Array2<f64> {
let labels: Vec<String> = cls.iter().map(|&c| name_map[c].to_string()).collect();
let train = dataset_xy(dir, tag, x, &labels);
let config = FitConfig::default();
let model = fit_penalized_multinomial_formula(&MultinomialFitRequest {
init_lambda: 1.0,
max_iter: 60,
tol: 1e-6,
..MultinomialFitRequest::new(&train, "y ~ s(x)", &config)
})
.expect("multinomial formula fit must succeed");
let grid_y: Vec<String> = grid.iter().map(|_| name_map[0].to_string()).collect();
let grid_ds = dataset_xy(dir, &format!("{tag}_grid"), grid, &grid_y);
let probs = predict_multinomial_formula(&model, &grid_ds)
.expect("multinomial predict must succeed");
let mut sorted: Vec<&str> = name_map.to_vec();
sorted.sort_unstable();
let col_of_orig: Vec<usize> = (0..3)
.map(|c| sorted.iter().position(|l| *l == name_map[c]).unwrap())
.collect();
assert_eq!(
model.class_levels,
sorted.iter().map(|s| s.to_string()).collect::<Vec<_>>(),
"class_levels must be the sorted label order"
);
let n = grid.len();
let mut aligned = Array2::<f64>::zeros((n, 3));
for r in 0..n {
for c in 0..3 {
aligned[[r, c]] = probs[[r, col_of_orig[c]]];
}
}
aligned
}
fn max_abs_diff(a: &Array2<f64>, b: &Array2<f64>) -> f64 {
a.iter()
.zip(b.iter())
.map(|(p, q)| (p - q).abs())
.fold(0.0_f64, f64::max)
}
#[test]
fn multinomial_fit_is_invariant_to_reference_class_1587() {
let td = tempdir().expect("tempdir");
let dir = td.path();
let (x, cls) = sample_classes(0, 300);
let grid: Vec<f64> = (0..7).map(|i| -1.5 + 3.0 * (i as f64) / 6.0).collect();
let a = fit_predict_aligned(dir, "abc", &x, &cls, ["A", "B", "C"], &grid);
let b = fit_predict_aligned(dir, "bca", &x, &cls, ["B", "C", "A"], &grid);
let c = fit_predict_aligned(dir, "cab", &x, &cls, ["C", "A", "B"], &grid);
let a2 = fit_predict_aligned(dir, "abc2", &x, &cls, ["A", "B", "C"], &grid);
let refit_noise = max_abs_diff(&a, &a2);
assert!(
refit_noise < 1e-6,
"refitting the same labeling must be deterministic (got {refit_noise:.3e})"
);
let drift = max_abs_diff(&a, &b)
.max(max_abs_diff(&a, &c))
.max(max_abs_diff(&b, &c));
assert!(
drift < 1e-3,
"predicted probabilities must be invariant to the reference class; \
cross-labeling drift = {drift:.3e} (refit noise = {refit_noise:.3e})"
);
}
#[test]
fn zz_measure_2349_outer_gradient_fd_at_refusal_checkpoint() {
let td = tempdir().expect("tempdir");
let dir = td.path();
let (x, cls) = sample_classes(0, 300);
let labels: Vec<String> = cls
.iter()
.map(|&c| ["A", "B", "C"][c].to_string())
.collect();
let train = dataset_xy(dir, "fd2349", &x, &labels);
let config = FitConfig::default();
let request = MultinomialFitRequest {
init_lambda: 1.0,
max_iter: 60,
tol: 1e-6,
..MultinomialFitRequest::new(&train, "y ~ s(x)", &config)
};
let parts = penalized_multinomial_formula_parts(&request)
.expect("production formula parts must build");
let rho_star = [
6.50584039279757,
-1.6183906983083074,
5.922109861708934,
-0.5810545109816936,
-0.4894709703255621,
1.299144316808675,
];
let mut probe_options = parts.options.clone();
probe_options.compute_covariance = false;
eprintln!(
"#2349 gate state: use_remlobjective={} (RidgedQuadraticReml default => \
logdet_h/logdet_s included in the fixed-lambda score iff this is true)",
probe_options.use_remlobjective
);
let v_at_with = |rho: &[f64], use_reml: bool| -> f64 {
let fam = parts
.family
.clone()
.with_joint_initial_log_lambdas(rho.to_vec());
let mut opts = probe_options.clone();
opts.use_remlobjective = use_reml;
let fit = crate::custom_family::fit_custom_family_fixed_log_lambdas(
&fam,
&parts.blocks,
&opts,
None,
)
.expect("fixed-lambda inner solve at the checkpoint must converge");
fit.reml_score
};
let v_plain = v_at_with(&rho_star, false);
let v_laml = v_at_with(&rho_star, true);
eprintln!(
"#2349 V(rho*): plain(penalized NLL)={v_plain:.9e} \
laml(+0.5logdetH-0.5logdetS)={v_laml:.9e} logdet_pair={:.9e} \
(the refusal reported final objective 2.687403e2 at this checkpoint — \
whichever variant matches IS the outer criterion)",
v_laml - v_plain
);
let outer_uses_laml = (v_laml - 2.687403e2).abs() < (v_plain - 2.687403e2).abs();
let v_at = |rho: &[f64]| -> f64 { v_at_with(rho, outer_uses_laml) };
{
let fam = parts
.family
.clone()
.with_joint_initial_log_lambdas(rho_star.to_vec());
let fit = crate::custom_family::fit_custom_family_fixed_log_lambdas(
&fam,
&parts.blocks,
&probe_options,
None,
)
.expect("fixed-lambda decomposition fit at the checkpoint");
eprintln!(
"#2349 decompose: reml_score={:.9e} penalized_objective={:.9e} \
log_likelihood={:.9e} deviance={:.9e}",
fit.reml_score, fit.penalized_objective, fit.log_likelihood, fit.deviance
);
}
let h = 1.0e-3;
let mut grad_fd = [0.0_f64; 6];
for s in 0..6 {
let mut plus = rho_star;
plus[s] += h;
let mut minus = rho_star;
minus[s] -= h;
grad_fd[s] = (v_at(&plus) - v_at(&minus)) / (2.0 * h);
eprintln!("#2349 FD dV/drho[{s}] = {:+.6e}", grad_fd[s]);
}
let norm = grad_fd.iter().map(|g| g * g).sum::<f64>().sqrt();
eprintln!(
"#2349 |FD grad| = {norm:.6e} on the {} criterion \
(certificate claimed |Pg|=2.047e0, bound 2.697e-3)",
if outer_uses_laml { "LAML" } else { "plain penalized-NLL" }
);
for delta in [2.0_f64, -2.0] {
let rho_far: Vec<f64> = rho_star.iter().map(|r| r + delta).collect();
let fam_far = parts
.family
.clone()
.with_joint_initial_log_lambdas(rho_far.clone());
let far_fit = crate::custom_family::fit_custom_family_fixed_log_lambdas(
&fam_far,
&parts.blocks,
&probe_options,
None,
)
.expect("cold fixed-lambda solve at the far point");
let far_beta: Vec<f64> = far_fit
.block_states
.iter()
.flat_map(|bs| bs.beta.iter().copied())
.collect();
let block_cols: Vec<usize> =
parts.blocks.iter().map(|s| s.design.ncols()).collect();
let warm = crate::custom_family::CustomFamilyWarmStart::from_cached_beta(
&block_cols,
&ndarray::Array1::from(far_beta),
)
.expect("warm start from far-point mode");
let fam_star = parts
.family
.clone()
.with_joint_initial_log_lambdas(rho_star.to_vec());
match crate::custom_family::fit_custom_family_fixed_log_lambdas(
&fam_star,
&parts.blocks,
&probe_options,
Some(&warm),
) {
Ok(fit) => eprintln!(
"#2349 warm-from(delta={delta:+.1}): V={:.9e} (cold {:.9e}, refusal 2.687403e2) \
gap_to_cold={:+.3e}",
fit.reml_score,
v_laml,
fit.reml_score - v_laml
),
Err(e) => eprintln!(
"#2349 warm-from(delta={delta:+.1}): inner REFUSED honestly: {}",
format!("{e}").chars().take(220).collect::<String>()
),
}
}
{
let fam = parts
.family
.clone()
.with_joint_initial_log_lambdas(rho_star.to_vec());
let eval_at = |rho_vec: &[f64]| -> (f64, ndarray::Array1<f64>, bool) {
let diagnostics =
crate::custom_family::evaluate_labeled_outer_criterion_for_diagnostics(
&fam,
&parts.blocks,
&probe_options,
&ndarray::Array1::from(rho_vec.to_vec()),
gam_problem::EvalMode::ValueAndGradient,
)
.expect("labeled outer evaluation at the checkpoint");
(
diagnostics.objective,
diagnostics.gradient,
diagnostics.inner_converged,
)
};
let (v0, g0, conv0) = eval_at(&rho_star);
eprintln!(
"#2349 labeled-evaluator at rho*: V={v0:.9e} (refusal 2.687403e2, \
fixed-lambda LAML 2.561663540e2) inner_converged={conv0} |analytic g|={:.6e}",
g0.iter().map(|g| g * g).sum::<f64>().sqrt()
);
let h = 1.0e-3;
for s in 0..6 {
let mut plus = rho_star;
plus[s] += h;
let mut minus = rho_star;
minus[s] -= h;
let (vp, _, _) = eval_at(&plus);
let (vm, _, _) = eval_at(&minus);
let fd = (vp - vm) / (2.0 * h);
eprintln!(
"#2349 labeled grad[{s}]: analytic={:+.6e} fd={fd:+.6e} diff={:+.3e}",
g0[s],
g0[s] - fd
);
}
}
}
}