use std::collections::HashMap;
use std::panic::{catch_unwind, AssertUnwindSafe};
use crate::formula::{label_ranef, lower, Column, Table};
use crate::{fit_warm, Boundary, Family, Note, StartValues, WaldSe};
use crate::{BinomialLink, GammaLink, NegBinomialLink, PoissonLink};
pub fn family_from_str(family: &str, link: &str) -> Result<Family, String> {
match family {
"gaussian" => Ok(Family::Gaussian),
"binomial" => match link {
"logit" => Ok(Family::Binomial {
link: BinomialLink::Logit,
}),
"probit" => Ok(Family::Binomial {
link: BinomialLink::Probit,
}),
"cloglog" => Err(
"link 'cloglog' requires GLMM 0.1.1; not yet implemented in the kernel".to_string(),
),
other => Err(format!("unsupported binomial link {other:?}")),
},
"poisson" => match link {
"log" => Ok(Family::Poisson {
link: PoissonLink::Log,
}),
other => Err(format!("unsupported poisson link {other:?}")),
},
"gamma" => match link {
"log" => Ok(Family::Gamma {
link: GammaLink::Log,
}),
"inverse" => Ok(Family::Gamma {
link: GammaLink::Inverse,
}),
other => Err(format!("unsupported gamma link {other:?}")),
},
"negativebinomial" => match link {
"log" => Ok(Family::NegativeBinomial {
link: NegBinomialLink::Log,
}),
other => Err(format!("unsupported negativebinomial link {other:?}")),
},
"inversegaussian" => Err(
"family 'inversegaussian' requires GLMM 0.1.1; not yet implemented \
in the kernel"
.to_string(),
),
other => Err(format!("unknown family {other:?}")),
}
}
#[derive(Debug)]
pub struct NoteInfo {
pub kind: &'static str,
pub columns: Vec<u32>,
pub pivot: f64,
pub evals: u32,
pub final_eval: bool,
pub detail: String,
pub ratio: f64,
}
pub type RanefBlockTuple = (String, Vec<String>, Vec<String>, Vec<f64>);
fn boundary_name(boundary: Boundary) -> &'static str {
match boundary {
Boundary::Interior => "interior",
Boundary::AtBoundary => "at_boundary",
Boundary::NoOptimum => "no_optimum",
}
}
fn note_infos(notes: Vec<Note>) -> Vec<NoteInfo> {
notes
.into_iter()
.map(|note| match note {
Note::IllConditioned { columns, pivot } => NoteInfo {
kind: "ill_conditioned",
columns,
pivot,
evals: 0,
final_eval: false,
detail: String::new(),
ratio: f64::NAN,
},
Note::PirlsExhausted { evals, final_eval } => NoteInfo {
kind: "pirls_exhausted",
columns: Vec::new(),
pivot: f64::NAN,
evals,
final_eval,
detail: String::new(),
ratio: f64::NAN,
},
Note::UnusedGroupingLevels { grouping, levels } => NoteInfo {
kind: "unused_grouping_levels",
columns: Vec::new(),
pivot: f64::NAN,
evals: 0,
final_eval: false,
detail: format!("{grouping}: {}", levels.join(", ")),
ratio: f64::NAN,
},
Note::ReDesignScaleSpread { grouping, ratio } => NoteInfo {
kind: "re_design_scale_spread",
columns: Vec::new(),
pivot: f64::NAN,
evals: 0,
final_eval: false,
detail: grouping,
ratio,
},
Note::HessianSeFallback => NoteInfo {
kind: "hessian_se_fallback",
columns: Vec::new(),
pivot: f64::NAN,
evals: 0,
final_eval: false,
detail: String::new(),
ratio: f64::NAN,
},
})
.collect()
}
#[derive(Debug)]
pub struct FitResult {
pub beta: Vec<f64>,
pub se: Vec<f64>,
pub vcov: Vec<Vec<f64>>,
pub tau2: Vec<f64>,
pub varcorr: Vec<Vec<f64>>,
pub stddev_se: Vec<f64>,
pub aliased: Vec<bool>,
pub dispersion: f64,
pub converged: bool,
pub n_eval: usize,
pub deviance: f64,
pub singular: bool,
pub names: Vec<String>,
pub re_groups: Vec<(String, Vec<String>)>,
pub agq_warning: Option<String>,
pub loglik: f64,
pub df: usize,
pub reml: bool,
pub fitted: Vec<f64>,
pub ranef: Vec<f64>,
pub ranef_levels: Vec<usize>,
pub ranef_blocks: Vec<RanefBlockTuple>,
pub boundary: &'static str,
pub pinned: Vec<Vec<bool>>,
pub notes: Vec<NoteInfo>,
}
#[allow(clippy::too_many_arguments)]
pub fn run_fit(
formula: &str,
numeric_columns: HashMap<String, Vec<f64>>,
factor_columns: HashMap<String, (Vec<String>, Vec<u32>)>,
family: &str,
link: &str,
wald_se: &str,
nagq: u8,
dispersion: Option<f64>,
weights: Option<Vec<f64>>,
offset: Option<Vec<f64>>,
warm_start: Option<(Vec<f64>, Vec<f64>)>,
) -> Result<FitResult, String> {
let fam = family_from_str(family, link)?;
let mut columns: Vec<(String, Column)> = Vec::new();
let mut lengths: Vec<(String, usize)> = Vec::new();
for (name, values) in numeric_columns {
lengths.push((name.clone(), values.len()));
columns.push((name, Column::Numeric(values)));
}
for (name, (levels, codes)) in factor_columns {
if let Some(&bad) = codes.iter().find(|&&c| c as usize >= levels.len()) {
return Err(format!(
"factor column {name:?}: code {bad} is out of range for {} levels",
levels.len()
));
}
lengths.push((name.clone(), codes.len()));
columns.push((name, Column::Factor { levels, codes }));
}
let n = match lengths.first() {
Some((_, first_len)) => {
let n = *first_len;
if lengths.iter().any(|(_, len)| *len != n) {
let mut detail: Vec<String> = lengths
.iter()
.map(|(name, len)| format!("{name:?}: {len}"))
.collect();
detail.sort();
return Err(format!(
"columns have mismatched lengths ({}); all columns must have the same length",
detail.join(", ")
));
}
n
}
None => 0,
};
let table = Table { columns, n };
let mut lowered = catch_unwind(AssertUnwindSafe(|| lower(formula, &table, fam)))
.map_err(panic_message)?
.map_err(|e| e.to_string())?;
lowered.opts.wald_se = match wald_se {
"hessian" => WaldSe::Hessian,
"rx" => WaldSe::Rx,
other => return Err(format!("unsupported wald_se {other:?}")),
};
let mut nagq = nagq;
let mut agq_warning: Option<String> = None;
if nagq > 1 {
let agq_family = matches!(fam, Family::Binomial { .. } | Family::Poisson { .. });
let eligible = match lowered.model.re.as_ref() {
Some(re) => {
let q_p = 1 + re.slopes.len(); agq_family && re.extra_groupings.is_empty() && q_p <= 3
}
None => false,
};
if !eligible {
agq_warning = Some(format!(
"nagq={nagq} (adaptive quadrature) applies only to binomial/Poisson \
mixed models with a single grouping factor and at most 3 random \
effects per group; fitting with Laplace (nagq=1)"
));
nagq = 1;
}
}
lowered.opts.nagq = nagq;
lowered.opts.dispersion = dispersion;
lowered.opts.weights = weights;
lowered.opts.offset = offset;
let lowered_notes = std::mem::take(&mut lowered.notes);
let start = warm_start.map(|(beta, theta)| StartValues { beta, theta });
let fit = catch_unwind(AssertUnwindSafe(|| {
fit_warm(
&lowered.x,
&lowered.y,
lowered.n,
lowered.p,
&lowered.model,
&lowered.ids,
start.as_ref(),
&lowered.opts,
)
}))
.map_err(panic_message)?;
let ranef_blocks: Vec<RanefBlockTuple> = label_ranef(&fit, &lowered.re_groups)
.map_err(|e| e.to_string())?
.into_iter()
.map(|b| (b.group, b.terms, b.levels, b.values))
.collect();
let re_groups: Vec<(String, Vec<String>)> = lowered
.re_groups
.into_iter()
.map(|g| (g.name, g.terms))
.collect();
if !fit.varcorr.is_empty() && fit.varcorr.len() != re_groups.len() {
return Err(format!(
"re_groups and varcorr must agree in length and order: \
{} grouping(s) lowered, {} varcorr block(s) returned",
re_groups.len(),
fit.varcorr.len()
));
}
let diagnostics = fit.diagnostics;
Ok(FitResult {
beta: fit.beta,
se: fit.se,
vcov: fit.vcov,
tau2: fit.tau2,
varcorr: fit.varcorr,
stddev_se: fit.stddev_se,
aliased: diagnostics.aliased,
dispersion: fit.dispersion,
converged: diagnostics.converged,
n_eval: fit.n_eval,
deviance: fit.deviance,
singular: diagnostics.singular,
names: lowered.col_names,
re_groups,
agq_warning,
loglik: fit.loglik,
df: fit.df,
reml: fit.reml,
fitted: fit.fitted,
ranef: fit.ranef,
ranef_levels: fit.ranef_levels,
ranef_blocks,
boundary: boundary_name(diagnostics.boundary),
pinned: diagnostics.pinned,
notes: note_infos(lowered_notes.into_iter().chain(diagnostics.notes).collect()),
})
}
fn panic_message(payload: Box<dyn std::any::Any + Send>) -> String {
if let Some(s) = payload.downcast_ref::<&str>() {
s.to_string()
} else if let Some(s) = payload.downcast_ref::<String>() {
s.clone()
} else {
"glmm kernel panicked with a non-string payload".to_string()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::collections::HashMap;
#[test]
fn gaussian_maps() {
assert_eq!(
family_from_str("gaussian", "identity"),
Ok(Family::Gaussian)
);
}
#[test]
fn binomial_logit_and_probit_map() {
assert_eq!(
family_from_str("binomial", "logit"),
Ok(Family::Binomial {
link: BinomialLink::Logit
})
);
assert_eq!(
family_from_str("binomial", "probit"),
Ok(Family::Binomial {
link: BinomialLink::Probit
})
);
}
#[test]
fn binomial_cloglog_is_a_kernel_gap() {
let err = family_from_str("binomial", "cloglog").unwrap_err();
assert!(err.contains("not yet implemented in the kernel"), "{err}");
}
#[test]
fn poisson_maps() {
assert_eq!(
family_from_str("poisson", "log"),
Ok(Family::Poisson {
link: PoissonLink::Log
})
);
}
#[test]
fn gamma_log_and_inverse_map() {
assert_eq!(
family_from_str("gamma", "log"),
Ok(Family::Gamma {
link: GammaLink::Log
})
);
assert_eq!(
family_from_str("gamma", "inverse"),
Ok(Family::Gamma {
link: GammaLink::Inverse
})
);
}
#[test]
fn negativebinomial_maps() {
assert_eq!(
family_from_str("negativebinomial", "log"),
Ok(Family::NegativeBinomial {
link: NegBinomialLink::Log
})
);
}
#[test]
fn inversegaussian_is_a_kernel_gap() {
let err = family_from_str("inversegaussian", "log").unwrap_err();
assert!(err.contains("not yet implemented in the kernel"), "{err}");
}
#[test]
fn pirls_exhausted_payload_survives_flattening() {
let notes = note_infos(vec![
Note::PirlsExhausted {
evals: 3,
final_eval: false,
},
Note::PirlsExhausted {
evals: 0,
final_eval: true,
},
]);
assert_eq!(notes[0].kind, "pirls_exhausted");
assert_eq!(notes[0].evals, 3);
assert!(!notes[0].final_eval);
assert_eq!(notes[1].kind, "pirls_exhausted");
assert_eq!(notes[1].evals, 0);
assert!(notes[1].final_eval);
}
#[test]
fn re_design_scale_spread_and_hessian_fallback_payloads_survive_flattening() {
let notes = note_infos(vec![
Note::ReDesignScaleSpread {
grouping: "Subject".to_string(),
ratio: 4200.0,
},
Note::HessianSeFallback,
]);
assert_eq!(notes[0].kind, "re_design_scale_spread");
assert_eq!(notes[0].detail, "Subject");
assert_eq!(notes[0].ratio, 4200.0);
assert_eq!(notes[1].kind, "hessian_se_fallback");
assert_eq!(notes[1].detail, "");
assert!(notes[1].ratio.is_nan());
}
fn factor_col(labels: &[&str]) -> (Vec<String>, Vec<u32>) {
let mut levels: Vec<String> = labels.iter().map(|s| s.to_string()).collect();
levels.sort();
levels.dedup();
let codes = labels
.iter()
.map(|l| levels.iter().position(|v| v == l).unwrap() as u32)
.collect();
(levels, codes)
}
#[allow(clippy::type_complexity)] fn toy_ols() -> (
HashMap<String, Vec<f64>>,
HashMap<String, (Vec<String>, Vec<u32>)>,
) {
let y = vec![1.0, 2.0, 2.9, 4.1, 5.0, 6.2, 6.8, 8.1, 9.0, 10.2];
let x = vec![0.0, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0];
let mut numeric = HashMap::new();
numeric.insert("y".to_string(), y);
numeric.insert("x".to_string(), x);
(numeric, HashMap::new())
}
#[test]
fn gaussian_ols_end_to_end() {
let (numeric, factor) = toy_ols();
let result = run_fit(
"y ~ x", numeric, factor, "gaussian", "identity", "hessian", 1, None, None, None, None,
)
.expect("fit should succeed");
assert_eq!(
result.names,
vec!["(Intercept)".to_string(), "x".to_string()]
);
assert_eq!(result.beta.len(), 2);
assert!(result.converged);
assert!(
(result.beta[1] - 1.0).abs() < 0.1,
"slope = {}",
result.beta[1]
);
}
#[test]
fn unknown_column_is_a_clean_error() {
let (numeric, factor) = toy_ols();
let err = run_fit(
"y ~ z", numeric, factor, "gaussian", "identity", "hessian", 1, None, None, None, None,
)
.unwrap_err();
assert!(err.contains("z"), "{err}");
}
#[test]
fn ineligible_nagq_is_stripped_with_a_warning_not_an_error() {
let (numeric, mut factor) = toy_ols();
let g: Vec<&str> = ["a", "b", "c", "d", "e"]
.iter()
.copied()
.cycle()
.take(10)
.collect();
factor.insert("g".to_string(), factor_col(&g));
let result = run_fit(
"y ~ x + (1 | g)",
numeric,
factor,
"gaussian",
"identity",
"hessian",
3,
None,
None,
None,
None,
)
.expect("ineligible nagq must be stripped, not an error");
let msg = result.agq_warning.as_deref().expect("warning expected");
assert!(msg.contains("nagq=3"), "{msg}");
}
#[test]
fn malformed_formula_panic_becomes_a_clean_error_not_a_process_abort() {
let (numeric, factor) = toy_ols();
let err = run_fit(
"y ~ :", numeric, factor, "gaussian", "identity", "hessian", 1, None, None, None, None,
)
.unwrap_err();
assert!(!err.is_empty());
}
#[test]
fn mismatched_column_lengths_is_a_clean_error() {
let mut numeric = std::collections::HashMap::new();
numeric.insert("y".to_string(), vec![1.0, 2.0, 3.0, 4.0, 5.0]);
numeric.insert("x".to_string(), vec![0.0, 1.0]); let err = run_fit(
"y ~ x",
numeric,
std::collections::HashMap::new(),
"gaussian",
"identity",
"hessian",
1,
None,
None,
None,
None,
)
.unwrap_err();
assert!(err.contains("x"), "{err}");
assert!(err.contains('5') && err.contains('2'), "{err}");
}
#[test]
fn unused_longer_column_does_not_silently_inflate_n() {
let (mut numeric, factor) = toy_ols();
numeric.insert("junk".to_string(), vec![0.0; 1000]); let err = run_fit(
"y ~ x", numeric, factor, "gaussian", "identity", "hessian", 1, None, None, None, None,
)
.unwrap_err();
assert!(err.contains("junk"), "{err}");
}
}