use super::*;
use gam_linalg::matrix::LinearOperator;
use gam_solve::estimate::reml::reml_outer_engine::penalty_matrix_root;
#[derive(Default)]
pub struct StandardFitOptionsInputs {
pub latent_cloglog: Option<LatentCLogLogState>,
pub mixture_link: Option<MixtureLinkSpec>,
pub optimize_mixture: bool,
pub sas_link: Option<SasLinkSpec>,
pub optimize_sas: bool,
pub linear_constraints: Option<gam_solve::pirls::LinearInequalityConstraints>,
pub firth_bias_reduction: bool,
pub adaptive_regularization: Option<AdaptiveRegularizationOptions>,
pub penalty_shrinkage_floor_override: Option<Option<f64>>,
}
pub fn canonical_standard_fit_options(
config: &FitConfig,
inputs: StandardFitOptionsInputs,
) -> FitOptions {
FitOptions {
resource_policy: resolved_resource_policy(
config,
gam_runtime::resource::ProblemHints::default(),
),
latent_cloglog: inputs.latent_cloglog,
mixture_link: inputs.mixture_link,
optimize_mixture: inputs.optimize_mixture,
sas_link: inputs.sas_link,
optimize_sas: inputs.optimize_sas,
compute_inference: true,
skip_rho_posterior_inference: true,
max_iter: config.outer_max_iter.unwrap_or(200),
tol: 1e-10,
nullspace_dims: vec![],
linear_constraints: inputs.linear_constraints,
firth_bias_reduction: inputs.firth_bias_reduction,
adaptive_regularization: inputs.adaptive_regularization,
penalty_shrinkage_floor: inputs
.penalty_shrinkage_floor_override
.unwrap_or(Some(1e-6)),
rho_prior: Default::default(),
kronecker_penalty_system: None,
kronecker_factored: None,
persist_warm_start_disk: config.persist_warm_start_disk,
}
}
pub fn fit_model(request: FitRequest<'_>) -> Result<FitResult, WorkflowError> {
let request = request;
let wrap_solver_err =
|reason: String| -> WorkflowError { WorkflowError::IntegrationFailed { reason } };
match request {
FitRequest::Standard(request) => fit_standard_model(request)
.map(FitResult::Standard)
.map_err(wrap_solver_err),
FitRequest::GaussianLocationScale(request) => fit_gaussian_location_scale_model(request)
.map(FitResult::GaussianLocationScale)
.map_err(wrap_solver_err),
FitRequest::BinomialLocationScale(request) => fit_binomial_location_scale_model(request)
.map(FitResult::BinomialLocationScale)
.map_err(wrap_solver_err),
FitRequest::DispersionLocationScale(request) => {
fit_dispersion_location_scale_model(request)
.map(FitResult::DispersionLocationScale)
.map_err(wrap_solver_err)
}
FitRequest::SurvivalLocationScale(request) => fit_survival_location_scale_model(request)
.map(FitResult::SurvivalLocationScale)
.map_err(wrap_solver_err),
FitRequest::SurvivalTransformation(request) => fit_survival_transformation_model(request)
.map(FitResult::SurvivalTransformation)
.map_err(wrap_solver_err),
FitRequest::BernoulliMarginalSlope(request) => fit_bernoulli_marginal_slope_model(request)
.map(FitResult::BernoulliMarginalSlope)
.map_err(wrap_solver_err),
FitRequest::SurvivalMarginalSlope(request) => fit_survival_marginal_slope_model(request)
.map(FitResult::SurvivalMarginalSlope)
.map_err(wrap_solver_err),
FitRequest::LatentSurvival(request) => fit_latent_survival_model(request)
.map(FitResult::LatentSurvival)
.map_err(wrap_solver_err),
FitRequest::LatentBinary(request) => fit_latent_binary_model(request)
.map(FitResult::LatentBinary)
.map_err(wrap_solver_err),
FitRequest::TransformationNormal(request) => fit_transformation_normal_model(request)
.map(FitResult::TransformationNormal)
.map_err(wrap_solver_err),
}
}
pub(crate) fn resolved_resource_policy(
config: &FitConfig,
hints: gam_runtime::resource::ProblemHints,
) -> gam_runtime::resource::ResourcePolicy {
if let Some(p) = config.resource_policy.clone() {
return p;
}
gam_runtime::resource::ResourcePolicy::for_problem(hints)
}
pub(crate) fn marginal_slope_hints(config: &FitConfig) -> gam_runtime::resource::ProblemHints {
gam_runtime::resource::ProblemHints {
marginal_slope_large_scale_active: requests_bernoulli_marginal_slope(config),
}
}
pub fn expectile_tau_for_config(config: &FitConfig) -> Result<Option<f64>, WorkflowError> {
let Some(raw) = config.family.as_deref() else {
return Ok(None);
};
let trimmed = raw.trim();
let lower = trimmed.to_ascii_lowercase();
if !(lower == "expectile" || lower.starts_with("expectile(")) {
return Ok(None);
}
let invalid = |reason: String| WorkflowError::InvalidConfig { reason };
let inline_tau = if let Some(rest) = lower.strip_prefix("expectile(") {
let inner = rest.strip_suffix(')').ok_or_else(|| {
invalid(format!(
"expectile family asymmetry must be written as `expectile(τ)`; got `{trimmed}`"
))
})?;
let value: f64 = inner.trim().parse().map_err(|_| {
invalid(format!(
"expectile asymmetry `{}` is not a finite number",
inner.trim()
))
})?;
Some(value)
} else {
None
};
let tau = match (inline_tau, config.expectile_tau) {
(Some(a), Some(b)) if (a - b).abs() > 0.0 => {
return Err(invalid(format!(
"expectile asymmetry given both inline (`expectile({a})`) and via expectile_tau \
({b}); supply exactly one"
)));
}
(Some(a), _) => a,
(None, Some(b)) => b,
(None, None) => 0.5,
};
if !(tau.is_finite() && tau > 0.0 && tau < 1.0) {
return Err(invalid(format!(
"expectile asymmetry τ must be finite and strictly in (0, 1); got {tau}"
)));
}
Ok(Some(tau))
}
fn expectile_row_weights(
y: ArrayView1<f64>,
mu: ArrayView1<f64>,
base: ArrayView1<f64>,
tau: f64,
) -> Array1<f64> {
Array1::from_shape_fn(y.len(), |i| {
let asym = if y[i] > mu[i] { tau } else { 1.0 - tau };
base[i] * asym
})
}
#[derive(Debug, Default)]
struct ExpectileSignCycle {
anchor: Option<Vec<bool>>,
power: usize,
span: usize,
}
impl ExpectileSignCycle {
fn observe(&mut self, sign: &[bool]) -> Option<usize> {
let Some(anchor) = self.anchor.as_deref() else {
self.anchor = Some(sign.to_vec());
self.power = 1;
return None;
};
self.span += 1;
if anchor == sign {
return Some(self.span);
}
if self.span == self.power {
self.anchor = Some(sign.to_vec());
self.power = self.power.saturating_mul(2);
self.span = 0;
}
None
}
}
fn expectile_kkt_residual(
design: &gam_linalg::matrix::DesignMatrix,
residual: ArrayView1<'_, f64>,
frozen_weights: ArrayView1<'_, f64>,
target_weights: ArrayView1<'_, f64>,
) -> Result<f64, String> {
use gam_linalg::matrix::LinearOperator;
let n = design.nrows();
if residual.len() != n || frozen_weights.len() != n || target_weights.len() != n {
return Err(format!(
"expectile KKT dimension mismatch: design rows={n}, residual={}, frozen weights={}, \
target weights={}",
residual.len(),
frozen_weights.len(),
target_weights.len(),
));
}
if residual.iter().any(|v| !v.is_finite())
|| frozen_weights
.iter()
.chain(target_weights.iter())
.any(|v| !v.is_finite() || *v < 0.0)
{
return Err(
"expectile KKT audit requires finite residuals and finite non-negative weights"
.to_string(),
);
}
let mut row_scratch =
Array1::from_shape_fn(n, |i| (frozen_weights[i] - target_weights[i]) * residual[i]);
let defect = design.apply_transpose(&row_scratch);
for i in 0..n {
row_scratch[i] = frozen_weights[i].max(target_weights[i]);
}
let energy = (0..n)
.map(|i| row_scratch[i] * residual[i] * residual[i])
.sum::<f64>();
if !energy.is_finite() || energy < 0.0 {
return Err(format!(
"expectile KKT audit produced invalid residual energy {energy:?}"
));
}
let gram_diag = design.diag_gram(&row_scratch)?;
if defect.len() != gram_diag.len()
|| defect.iter().any(|v| !v.is_finite())
|| gram_diag.iter().any(|v| !v.is_finite() || *v < 0.0)
{
return Err("expectile KKT audit produced invalid score/Gram evidence".to_string());
}
let mut max_scaled = 0.0_f64;
for (&d, &q) in defect.iter().zip(gram_diag.iter()) {
let denominator_squared = q * energy;
let scaled = if denominator_squared > 0.0 {
d.abs() / denominator_squared.sqrt()
} else if d == 0.0 {
0.0
} else {
f64::INFINITY
};
max_scaled = max_scaled.max(scaled);
}
Ok(max_scaled)
}
#[cfg(test)]
mod expectile_convergence_tests {
use super::{ExpectileSignCycle, expectile_kkt_residual};
use gam_linalg::matrix::{DenseDesignMatrix, DesignMatrix};
use ndarray::array;
#[test]
fn brent_detector_finds_fixed_sign_state() {
let mut detector = ExpectileSignCycle::default();
let sign = vec![true, false, true, true];
assert_eq!(detector.observe(&sign), None);
assert_eq!(detector.observe(&sign), Some(1));
}
#[test]
fn brent_detector_finds_longer_cycle_without_storing_history() {
let mut detector = ExpectileSignCycle::default();
let cycle = [
vec![true, false, false],
vec![false, true, false],
vec![false, false, true],
];
let mut detected = None;
for sign in cycle.iter().cycle().take(9) {
detected = detector.observe(sign);
if detected.is_some() {
break;
}
}
assert_eq!(detected, Some(3));
assert_eq!(detector.anchor.as_ref().map(Vec::len), Some(3));
}
#[test]
fn normalized_kkt_residual_handles_a_cancelling_frozen_score() {
let design = DesignMatrix::Dense(DenseDesignMatrix::from(array![[1.0], [1.0]]));
let residual = array![-1.0, 1.0];
let frozen = array![1.0, 1.0];
let target = array![1.0, 1.0 + 1.0e-12];
let kkt = expectile_kkt_residual(&design, residual.view(), frozen.view(), target.view())
.expect("finite KKT audit");
assert!(kkt < 1.0e-10, "normalized residual was {kkt:.3e}");
}
#[test]
fn normalized_kkt_residual_is_column_and_weight_scale_invariant() {
let residual = array![-2.0, 1.0, 1.0];
let frozen = array![1.0, 1.0, 1.0];
let target = array![1.0, 1.25, 0.75];
let x = array![[1.0], [2.0], [-1.0]];
let base = DesignMatrix::Dense(DenseDesignMatrix::from(x.clone()));
let scaled = DesignMatrix::Dense(DenseDesignMatrix::from(x * 1.0e6));
let base_kkt = expectile_kkt_residual(&base, residual.view(), frozen.view(), target.view())
.expect("base KKT audit");
let scaled_kkt = expectile_kkt_residual(
&scaled,
residual.view(),
(frozen.clone() * 1.0e4).view(),
(target.clone() * 1.0e4).view(),
)
.expect("scaled KKT audit");
assert!((base_kkt - scaled_kkt).abs() <= f64::EPSILON.sqrt());
}
}
fn deterministic_gaussian_standard_fit(
request: &StandardFitRequest<'_>,
exact_unpenalized_beta: Option<Array1<f64>>,
) -> Result<StandardFitResult, WorkflowError> {
if !request.family.is_gaussian_identity() || request.y.is_empty() {
return Err(WorkflowError::InvalidConfig {
reason: "deterministic Gaussian shortcut requires a non-empty Gaussian identity request"
.to_string(),
});
}
if request.y.iter().any(|value| !value.is_finite())
|| request.offset.iter().any(|value| !value.is_finite())
|| request
.weights
.iter()
.any(|value| !value.is_finite() || *value < 0.0)
{
return Err(WorkflowError::InvalidConfig {
reason: "deterministic Gaussian shortcut requires finite response, offset, and non-negative weights"
.to_string(),
});
}
let weight_sum = request.weights.sum();
if !(weight_sum.is_finite() && weight_sum > 0.0) {
return Err(WorkflowError::InvalidConfig {
reason: "deterministic Gaussian shortcut requires positive total weight".to_string(),
});
}
let design =
build_term_collection_design(request.data.view(), &request.spec).map_err(|err| {
WorkflowError::InvalidConfig {
reason: format!("deterministic Gaussian shortcut could not rebuild design: {err}"),
}
})?;
let p = design.design.ncols();
let beta = match exact_unpenalized_beta {
Some(beta) => {
if beta.len() != p {
return Err(WorkflowError::IntegrationFailed {
reason: format!(
"deterministic Gaussian coefficient width {} does not match rebuilt design width {p}",
beta.len()
),
});
}
beta
}
None => {
let intercept = request.y[0] - request.offset[0];
let mut beta = Array1::<f64>::zeros(p);
for col in design.intercept_range.clone() {
if col < p {
beta[col] = intercept;
}
}
beta
}
};
let fitted_eta = design.design.apply(&beta) + request.offset.as_ref();
let max_abs_eta = fitted_eta
.iter()
.copied()
.map(f64::abs)
.fold(0.0_f64, f64::max);
let x_dense = design.design.to_dense();
let weights = request.weights.as_ref().clone();
let xtwx = gam_linalg::faer_ndarray::fast_xt_diag_x(&x_dense, &weights);
let n_penalties = design.penalties.len();
let mut unit_penalty = Array2::<f64>::zeros((p, p));
for (penalty_index, block) in design.penalties.iter().enumerate() {
let r = block.col_range.clone();
if r.is_empty()
|| r.end > p
|| block.local.nrows() != r.len()
|| block.local.ncols() != r.len()
{
return Err(WorkflowError::IntegrationFailed {
reason: format!(
"deterministic Gaussian shortcut received malformed penalty {penalty_index}: \
range={r:?}, local={}x{}, design width={p}",
block.local.nrows(),
block.local.ncols()
),
});
}
if block.local.iter().any(|value| !value.is_finite()) {
return Err(WorkflowError::IntegrationFailed {
reason: format!(
"deterministic Gaussian shortcut received non-finite penalty {penalty_index}"
),
});
}
unit_penalty
.slice_mut(ndarray::s![r.clone(), r])
.scaled_add(1.0, &block.local);
}
let lambda_full = if n_penalties == 0 {
0.0
} else {
use gam_linalg::faer_ndarray::FaerEigh;
let symmetric_penalty = (&unit_penalty + &unit_penalty.t().to_owned()) * 0.5;
let (penalty_eigenvalues, _) =
symmetric_penalty.eigh(faer::Side::Lower).map_err(|error| {
WorkflowError::IntegrationFailed {
reason: format!(
"deterministic Gaussian shortcut could not resolve the penalty spectrum: {error}"
),
}
})?;
let largest_penalty = penalty_eigenvalues
.iter()
.fold(0.0_f64, |largest, &value| largest.max(value.abs()));
if !(largest_penalty.is_finite() && largest_penalty > 0.0) {
return Err(WorkflowError::IntegrationFailed {
reason: "deterministic Gaussian shortcut received penalties with zero numerical rank"
.to_string(),
});
}
let rank_floor = f64::EPSILON * (p.max(1) as f64) * largest_penalty;
if let Some(&negative) = penalty_eigenvalues
.iter()
.filter(|&&value| value < -rank_floor)
.min_by(|left, right| left.total_cmp(right))
{
return Err(WorkflowError::IntegrationFailed {
reason: format!(
"deterministic Gaussian shortcut received a non-PSD penalty \
(minimum eigenvalue {negative:.6e}, numerical floor {rank_floor:.6e})"
),
});
}
let weakest_penalty = penalty_eigenvalues
.iter()
.copied()
.filter(|&value| value > rank_floor)
.min_by(|left, right| left.total_cmp(right))
.ok_or_else(|| WorkflowError::IntegrationFailed {
reason: "deterministic Gaussian shortcut could not identify a penalized direction"
.to_string(),
})?;
let information_scale = xtwx
.rows()
.into_iter()
.map(|row| row.iter().map(|value| value.abs()).sum::<f64>())
.fold(0.0_f64, f64::max)
.max(f64::MIN_POSITIVE);
let lambda = information_scale / (f64::EPSILON.sqrt() * weakest_penalty);
if !(lambda.is_finite() && lambda > 0.0) {
return Err(WorkflowError::IntegrationFailed {
reason: format!(
"deterministic Gaussian shortcut produced invalid boundary precision {lambda}"
),
});
}
lambda
};
let log_lambda_full = lambda_full.max(f64::MIN_POSITIVE).ln();
let lambda_full = if n_penalties == 0 {
lambda_full
} else {
gam_problem::checked_exp_log_strength(log_lambda_full).map_err(|error| {
WorkflowError::IntegrationFailed {
reason: format!(
"deterministic Gaussian shortcut produced a boundary precision outside the \
log-strength domain: {error}"
),
}
})?
};
let mut penalized_hessian = xtwx.clone();
penalized_hessian.scaled_add(lambda_full, &unit_penalty);
penalized_hessian = (&penalized_hessian + &penalized_hessian.t()) * 0.5;
let (edf_total, edf_by_block, penalty_block_trace, coefficient_influence) = {
use gam_linalg::faer_ndarray::FaerCholesky;
let chol = penalized_hessian
.cholesky(faer::Side::Lower)
.map_err(|error| WorkflowError::IntegrationFailed {
reason: format!(
"deterministic Gaussian boundary precision is not positive definite: {error}"
),
})?;
{
let influence = chol.solve_mat(&xtwx);
let mut raw_traces = vec![0.0_f64; n_penalties];
let mut block_ranks = vec![0_usize; n_penalties];
for (kk, block) in design.penalties.iter().enumerate() {
let r = block.col_range.clone();
let block_cols = r.len();
block_ranks[kk] = penalty_matrix_root(&block.local)
.map_err(|reason| WorkflowError::IntegrationFailed {
reason: format!(
"deterministic Gaussian shortcut penalty {kk} rank factorization failed: {reason}"
),
})?
.nrows();
let mut rhs = Array2::<f64>::zeros((p, block_cols));
for c in 0..block_cols {
for rr in 0..block_cols {
rhs[[r.start + rr, c]] = block.local[[rr, c]];
}
}
let sol = chol.solve_mat(&rhs);
let mut trace = 0.0_f64;
for j in 0..block_cols {
trace += sol[[r.start + j, j]];
}
raw_traces[kk] = lambda_full * trace;
}
let joint_penalty_rank = penalty_matrix_root(&unit_penalty)
.map_err(|reason| WorkflowError::IntegrationFailed {
reason: format!(
"deterministic Gaussian shortcut joint penalty rank factorization failed: {reason}"
),
})?
.nrows();
let bundle = gam_solve::estimate::penalized_edf_bundle(
&raw_traces,
&block_ranks,
p,
(p - joint_penalty_rank.min(p)) as f64,
);
(
bundle.edf_total,
bundle.edf_by_block,
bundle.penalty_block_trace,
Some(influence),
)
}
};
let working_response = request.y.as_ref().clone();
let lambdas = Array1::<f64>::from_elem(n_penalties, lambda_full);
let log_lambdas = Array1::<f64>::from_elem(n_penalties, log_lambda_full);
let penalized_hessian_precision =
gam_problem::dispersion_cov::UnscaledPrecision::wrap(penalized_hessian.clone());
let inference = gam_solve::estimate::FitInference {
edf_by_block,
penalty_block_trace,
edf_total,
smoothing_correction: None,
smoothing_correction_method: None,
smoothing_correction_first_order: None,
smoothing_correction_method_first_order: None,
penalized_hessian: penalized_hessian_precision.clone(),
reparam_qs: None,
dispersion: gam_solve::estimate::Dispersion::ZERO_ESTIMATE,
beta_covariance: Some(gam_problem::dispersion_cov::PhiScaledCovariance::wrap(
ndarray::Array2::<f64>::zeros((p, p)),
)),
beta_standard_errors: Some(Array1::<f64>::zeros(p)),
beta_covariance_corrected: None,
beta_standard_errors_corrected: None,
beta_covariance_frequentist: None,
coefficient_influence,
weighted_gram: Some(xtwx),
bias_correction_beta: None,
bias_correction_jacobian: None,
};
let geometry = Some(gam_solve::estimate::FitGeometry {
coefficient_gauge: gam_problem::gauge::Gauge::identity(&[beta.len()]),
penalized_hessian: penalized_hessian_precision,
constrained_posterior: None,
working: Some(gam_solve::estimate::WorkingGeometry {
weights,
response: working_response,
}),
});
let fit = gam_solve::estimate::UnifiedFitResult::try_from_parts(
gam_solve::estimate::UnifiedFitResultParts {
blocks: vec![gam_solve::estimate::FittedBlock {
beta: beta.clone(),
role: gam_problem::BlockRole::Mean,
edf: edf_total,
lambdas: lambdas.clone(),
}],
log_lambdas,
lambdas,
likelihood_family: Some(request.family.clone()),
likelihood_scale: gam_problem::LikelihoodScaleMetadata::ProfiledGaussian,
log_likelihood_normalization: gam_problem::LogLikelihoodNormalization::UserProvided,
log_likelihood: 0.0,
deviance: 0.0,
reml_score: 0.0,
stable_penalty_term: 0.0,
penalized_objective: 0.0,
used_device: false,
outer_iterations: 0,
outer_converged: true,
outer_gradient_norm: Some(0.0),
standard_deviation: 0.0,
covariance_conditional: Some(ndarray::Array2::<f64>::zeros((p, p))),
covariance_corrected: None,
inference: Some(inference),
fitted_link: gam_solve::estimate::FittedLinkState::Standard(None),
geometry,
block_states: Vec::new(),
pirls_status: gam_solve::pirls::PirlsStatus::Converged,
max_abs_eta,
constraint_kkt: None,
artifacts: gam_solve::estimate::FitArtifacts {
pirls: None,
..Default::default()
},
inner_cycles: 0,
},
)
.map_err(|err| WorkflowError::IntegrationFailed {
reason: format!("deterministic Gaussian shortcut produced invalid fit: {err}"),
})?;
let resolvedspec =
freeze_term_collection_from_design(&request.spec, &design).map_err(|err| {
WorkflowError::InvalidConfig {
reason: format!("deterministic Gaussian shortcut could not freeze design: {err}"),
}
})?;
Ok(StandardFitResult {
fit,
design,
resolvedspec,
adaptive_spatial_terms: adaptive_spatial_term_mask(&request.spec),
adaptive_spatial_center_counts: adaptive_spatial_center_counts(&request.spec),
adaptive_diagnostics: None,
kappa_timing: None,
saved_link_state: gam_solve::estimate::FittedLinkState::Standard(None),
wiggle_knots: None,
wiggle_degree: None,
wiggle_penalty_metadata: None,
wiggle_saved_warp_beta: None,
wiggle_saved_index_shift: None,
})
}
fn gaussian_response_is_constant(request: &StandardFitRequest<'_>) -> bool {
if !request.family.is_gaussian_identity() || request.y.is_empty() {
return false;
}
if gam_terms::smooth::term_collection_has_nonzero_anchor(&request.spec) {
return false;
}
if request.y.len() != request.offset.len() {
return false;
}
let mut adjusted = request.y.iter().zip(request.offset.iter());
let Some((&first_y, &first_offset)) = adjusted.next() else {
return false;
};
let first = first_y - first_offset;
if !first.is_finite() {
return false;
}
for (&yi, &oi) in adjusted {
let value = yi - oi;
if !value.is_finite() || value != first {
return false;
}
}
true
}
fn exact_unpenalized_gaussian_beta(
request: &StandardFitRequest<'_>,
) -> Result<Option<Array1<f64>>, WorkflowError> {
if !request.family.is_gaussian_identity()
|| request.y.is_empty()
|| !request.spec.smooth_terms.is_empty()
|| !request.spec.random_effect_terms.is_empty()
|| request.options.linear_constraints.is_some()
|| request.spec.linear_terms.iter().any(|term| {
!matches!(
&term.coefficient_geometry,
gam_terms::smooth::LinearCoefficientGeometry::Unconstrained
) || term.coefficient_min.is_some()
|| term.coefficient_max.is_some()
})
|| request.y.len() != request.offset.len()
|| request.y.len() != request.weights.len()
{
return Ok(None);
}
let design =
build_term_collection_design(request.data.view(), &request.spec).map_err(|err| {
WorkflowError::InvalidConfig {
reason: format!(
"deterministic Gaussian candidate could not build its parametric design: {err}"
),
}
})?;
if !design.penalties.is_empty() || design.design.ncols() == 0 {
return Ok(None);
}
let adjusted_response = request.y.as_ref() - request.offset.as_ref();
if adjusted_response.iter().any(|value| !value.is_finite())
|| request
.weights
.iter()
.any(|weight| !weight.is_finite() || *weight < 0.0)
{
return Ok(None);
}
let x = design.design.to_dense();
let gram = gam_linalg::faer_ndarray::fast_xt_diag_x(&x, request.weights.as_ref());
let rhs_matrix = gam_linalg::faer_ndarray::fast_xt_diag_y(
&x,
request.weights.as_ref(),
&adjusted_response.view().insert_axis(ndarray::Axis(1)),
);
let rhs = rhs_matrix.column(0).to_owned();
let beta = match gam_linalg::utils::certified_symmetric_solve(
&gram,
&rhs,
"deterministic Gaussian normal equations",
) {
Ok(solution) => solution.into_solution(),
Err(_) => return Ok(None),
};
let fitted = design.design.apply(&beta);
let operations = (x.ncols() + 1) as f64;
let roundoff = operations * f64::EPSILON;
if !(roundoff < 1.0) {
return Ok(None);
}
let gamma = roundoff / (1.0 - roundoff);
for row in 0..x.nrows() {
if request.weights[row] == 0.0 {
continue;
}
let operand_scale = adjusted_response[row].abs()
+ x.row(row)
.iter()
.zip(beta.iter())
.map(|(&value, &coefficient)| (value * coefficient).abs())
.sum::<f64>();
let residual = (adjusted_response[row] - fitted[row]).abs();
if !residual.is_finite() || residual > gamma * operand_scale {
return Ok(None);
}
}
Ok(Some(beta))
}
pub fn fit_from_formula(
formula: &str,
data: &Dataset,
config: &FitConfig,
) -> Result<FitResult, WorkflowError> {
fit_from_formula_with_notes(formula, data, config).map(|outcome| outcome.result)
}
pub struct FormulaFitResult {
pub result: FitResult,
pub inference_notes: Vec<String>,
}
pub fn fit_from_formula_with_notes(
formula: &str,
data: &Dataset,
config: &FitConfig,
) -> Result<FormulaFitResult, WorkflowError> {
let mut config = config
.clone()
.resolve()
.map_err(|reason| WorkflowError::InvalidConfig { reason })?;
config.spatial_center_counts = Some(Vec::new());
let current = fit_from_formula_once_with_notes(formula, data, &config)?;
finish_adaptive_spatial_fit(formula, data, config, current)
}
pub fn fit_materialized_standard_with_notes(
formula: &str,
data: &Dataset,
config: &FitConfig,
request: StandardFitRequest<'_>,
inference_notes: Vec<String>,
) -> Result<FormulaFitResult, WorkflowError> {
let mut config = config
.clone()
.resolve()
.map_err(|reason| WorkflowError::InvalidConfig { reason })?;
config.spatial_center_counts = Some(Vec::new());
let current = fit_materialized_once_with_notes(MaterializedModel {
request: FitRequest::Standard(request),
inference_notes,
})?;
finish_adaptive_spatial_fit(formula, data, config, current)
}
fn finish_adaptive_spatial_fit(
formula: &str,
data: &Dataset,
mut config: FitConfig,
mut current: FormulaFitResult,
) -> Result<FormulaFitResult, WorkflowError> {
loop {
let Some(current_standard) = standard_result(¤t) else {
return Ok(current);
};
let standard_options =
canonical_standard_fit_options(&config, StandardFitOptionsInputs::default());
let resolution_tol = standard_options
.tol
.max(standard_options.penalty_shrinkage_floor.unwrap_or(0.0));
let candidates =
adaptive_spatial_candidates(current_standard, data.values.nrows(), resolution_tol)?;
if candidates.is_empty() {
return Ok(current);
}
let term_count = candidates.term_count;
let candidate = candidates
.terms
.into_iter()
.next()
.expect("non-empty adaptive candidate set");
drop(current);
let mut candidate_config = config.clone();
let center_counts = candidate_config
.spatial_center_counts
.get_or_insert_with(Vec::new);
if center_counts.len() < term_count {
center_counts.resize(term_count, None);
}
center_counts[candidate.term_index] = Some(candidate.proposed_centers);
let candidate_outcome = fit_from_formula_once_with_notes(formula, data, &candidate_config)
.map_err(|error| WorkflowError::SpatialUnderresolved {
term: candidate.term_name.clone(),
current_centers: candidate.current_centers,
attempted_centers: candidate.proposed_centers,
reason: error.to_string(),
})?;
if standard_result(&candidate_outcome).is_none() {
return Err(WorkflowError::SpatialUnderresolved {
term: candidate.term_name.clone(),
current_centers: candidate.current_centers,
attempted_centers: candidate.proposed_centers,
reason: "the certification refit changed estimator representation".to_string(),
});
}
config = candidate_config;
current = candidate_outcome;
}
}
struct AdaptiveSpatialCandidates {
term_count: usize,
terms: Vec<AdaptiveSpatialCandidate>,
}
impl AdaptiveSpatialCandidates {
fn is_empty(&self) -> bool {
self.terms.is_empty()
}
}
struct AdaptiveSpatialCandidate {
term_index: usize,
term_name: String,
current_centers: usize,
proposed_centers: usize,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum AdaptiveCenterDecision {
Certified,
Expand(usize),
Exhausted,
}
fn adaptive_center_decision(
current_centers: usize,
ceiling_centers: usize,
edf: f64,
realized_width: usize,
nullspace_dim: usize,
resolution_tol: f64,
) -> AdaptiveCenterDecision {
if !gam_terms::basis::basis_is_saturated(edf, realized_width, nullspace_dim, resolution_tol) {
return AdaptiveCenterDecision::Certified;
}
match gam_terms::basis::expanded_num_centers(current_centers, ceiling_centers) {
Some(proposed) => AdaptiveCenterDecision::Expand(proposed),
None => AdaptiveCenterDecision::Exhausted,
}
}
fn standard_result(outcome: &FormulaFitResult) -> Option<&StandardFitResult> {
match &outcome.result {
FitResult::Standard(result) => Some(result),
_ => None,
}
}
fn adaptive_spatial_candidates(
result: &StandardFitResult,
n_rows: usize,
resolution_tol: f64,
) -> Result<AdaptiveSpatialCandidates, WorkflowError> {
let term_count = result.resolvedspec.smooth_terms.len();
if result.adaptive_spatial_terms.len() != term_count
|| result.adaptive_spatial_center_counts.len() != term_count
|| result.design.smooth.terms.len() != term_count
{
return Err(WorkflowError::IntegrationFailed {
reason: format!(
"adaptive spatial provenance mismatch: resolved terms={term_count}, mask={}, \
requested counts={}, realized terms={}",
result.adaptive_spatial_terms.len(),
result.adaptive_spatial_center_counts.len(),
result.design.smooth.terms.len(),
),
});
}
let smooth_offset = result
.design
.design
.ncols()
.saturating_sub(result.design.smooth.total_smooth_cols());
let mut candidates = Vec::new();
for term_index in 0..term_count {
let realized = &result.design.smooth.terms[term_index];
if result.adaptive_spatial_terms[term_index]
&& let Some(current_centers) = result.adaptive_spatial_center_counts[term_index]
{
let penalty_range = result
.design
.smooth_term_penalty_range(term_index)
.map_err(|reason| WorkflowError::IntegrationFailed { reason })?
.ok_or_else(|| WorkflowError::IntegrationFailed {
reason: format!(
"adaptive spatial term '{}' emitted no penalty block",
result.resolvedspec.smooth_terms[term_index].name,
),
})?;
let spatial_dimension = result.resolvedspec.smooth_terms[term_index]
.basis
.structural_feature_cols()
.len();
if spatial_dimension == 0 {
return Err(WorkflowError::IntegrationFailed {
reason: format!(
"adaptive spatial term '{}' has no structural feature columns",
result.resolvedspec.smooth_terms[term_index].name,
),
});
}
let ceiling_centers = gam_terms::basis::default_num_centers(n_rows, spatial_dimension)
.max(current_centers);
let global_range = (smooth_offset + realized.coeff_range.start)
..(smooth_offset + realized.coeff_range.end);
let edf =
result
.fit
.per_term_edf(global_range, penalty_range.start, penalty_range.len());
let nullspace_dim = realized.wald_unpenalized_dim();
match adaptive_center_decision(
current_centers,
ceiling_centers,
edf,
realized.coeff_range.len(),
nullspace_dim,
resolution_tol,
) {
AdaptiveCenterDecision::Certified => {}
AdaptiveCenterDecision::Expand(proposed_centers) => {
candidates.push(AdaptiveSpatialCandidate {
term_index,
term_name: result.resolvedspec.smooth_terms[term_index].name.clone(),
current_centers,
proposed_centers,
});
}
AdaptiveCenterDecision::Exhausted => {
return Err(WorkflowError::SpatialUnderresolved {
term: result.resolvedspec.smooth_terms[term_index].name.clone(),
current_centers,
attempted_centers: ceiling_centers,
reason: format!(
"term EDF {edf:.6} remains at its realized basis ceiling with all \
{ceiling_centers} validated default centers already requested"
),
});
}
}
}
}
Ok(AdaptiveSpatialCandidates {
term_count,
terms: candidates,
})
}
#[cfg(test)]
mod adaptive_spatial_resolution_tests {
use super::{AdaptiveCenterDecision, adaptive_center_decision};
#[test]
fn unsaturated_basis_is_certified_without_a_probe_refit() {
assert_eq!(
adaptive_center_decision(8, 100, 5.0, 10, 2, 1.0e-6),
AdaptiveCenterDecision::Certified
);
}
#[test]
fn saturated_basis_expands_geometrically_and_respects_validated_ceiling() {
assert_eq!(
adaptive_center_decision(8, 100, 10.0, 10, 2, 1.0e-6),
AdaptiveCenterDecision::Expand(16)
);
assert_eq!(
adaptive_center_decision(64, 100, 10.0, 10, 2, 1.0e-6),
AdaptiveCenterDecision::Expand(100)
);
}
#[test]
fn saturated_basis_at_validated_ceiling_is_typed_exhaustion() {
assert_eq!(
adaptive_center_decision(100, 100, 10.0, 10, 2, 1.0e-6),
AdaptiveCenterDecision::Exhausted
);
}
}
fn fit_from_formula_once_with_notes(
formula: &str,
data: &Dataset,
config: &FitConfig,
) -> Result<FormulaFitResult, WorkflowError> {
if let Some(result) = fit_expectile_if_requested(formula, data, &config)? {
return Ok(FormulaFitResult {
result: FitResult::Standard(result),
inference_notes: Vec::new(),
});
}
let mat = materialize(formula, data, &config)?;
fit_materialized_once_with_notes(mat)
}
fn fit_materialized_once_with_notes(
mat: MaterializedModel<'_>,
) -> Result<FormulaFitResult, WorkflowError> {
let inference_notes = mat.inference_notes;
if let FitRequest::Standard(request) = &mat.request {
if gaussian_response_is_constant(request) {
return deterministic_gaussian_standard_fit(request, None).map(|result| {
FormulaFitResult {
result: FitResult::Standard(result),
inference_notes,
}
});
}
if let Some(beta) = exact_unpenalized_gaussian_beta(request)? {
return deterministic_gaussian_standard_fit(request, Some(beta)).map(|result| {
FormulaFitResult {
result: FitResult::Standard(result),
inference_notes,
}
});
}
if let Some(inputs) = spline_scan_fast_path(request) {
let scan = gam_solve::spline_scan::fit_spline_scan(
&inputs.x,
&inputs.y,
&inputs.w,
inputs.order,
)
.map_err(|reason| WorkflowError::IntegrationFailed { reason })?;
return Ok(FormulaFitResult {
result: FitResult::SplineScan(scan),
inference_notes,
});
}
if let Some(inputs) = residual_cascade_fast_path(request) {
let coord_refs: Vec<&[f64]> = inputs.coords.iter().map(Vec::as_slice).collect();
if let Ok(fit) = gam_solve::residual_cascade::fit_residual_cascade(
&coord_refs,
&inputs.y,
&inputs.w,
&inputs.metric,
inputs.sobolev_s,
) {
return Ok(FormulaFitResult {
result: FitResult::ResidualCascade(fit),
inference_notes,
});
}
}
}
fit_model(mat.request).map(|result| FormulaFitResult {
result,
inference_notes,
})
}
pub fn fit_expectile_if_requested(
formula: &str,
data: &Dataset,
config: &FitConfig,
) -> Result<Option<StandardFitResult>, WorkflowError> {
match expectile_tau_for_config(config)? {
Some(tau) => Ok(Some(fit_expectile_laws(formula, data, config, tau)?)),
None => Ok(None),
}
}
fn fit_expectile_laws(
formula: &str,
data: &Dataset,
config: &FitConfig,
tau: f64,
) -> Result<StandardFitResult, WorkflowError> {
if config.frailty.is_active() {
return Err(WorkflowError::InvalidConfig {
reason: "expectile regression does not support frailty; use a survival/frailty-aware family instead"
.to_string(),
});
}
let gaussian_config = FitConfig {
family: Some("gaussian".to_string()),
link: Some("identity".to_string()),
expectile_tau: None,
frailty: FrailtySpec::None,
..config.clone()
};
let base_mat = materialize(formula, data, &gaussian_config)?;
let FitRequest::Standard(base_request) = base_mat.request else {
return Err(WorkflowError::InvalidConfig {
reason: "expectile regression is only defined for standard (non-survival, \
non-location-scale) responses"
.to_string(),
});
};
let StandardFitRequest {
data: design_data,
y,
weights: base_weights,
offset,
spec,
family: materialized_family,
estimate_tweedie_p: _,
options,
kappa_options,
wiggle,
coefficient_groups,
penalty_block_gamma_priors,
latent_coord,
} = base_request;
if !materialized_family.is_gaussian_identity() {
return Err(WorkflowError::InvalidConfig {
reason: format!(
"expectile LAWS requires a Gaussian-identity inner family; materializer produced {}",
materialized_family.name()
),
});
}
if wiggle.is_some() || latent_coord.is_some() {
return Err(WorkflowError::InvalidConfig {
reason: "expectile regression does not support flexible-link wiggle or latent \
coordinates"
.to_string(),
});
}
let n = y.len();
let gaussian_family = LikelihoodSpec::gaussian_identity();
let mut weights = Arc::clone(&base_weights);
let mut sign_cycle = ExpectileSignCycle::default();
let mut last_kkt = (f64::NAN, f64::NAN);
let mut last_rho_checkpoint = Vec::new();
let max_laws_iters = options.max_iter;
if max_laws_iters == 0 || !(options.tol.is_finite() && options.tol > 0.0) {
return Err(WorkflowError::InvalidConfig {
reason: format!(
"expectile LAWS requires a positive iteration budget and finite positive KKT \
tolerance; got max_iter={max_laws_iters}, tol={}",
options.tol,
),
});
}
for iteration in 1..=max_laws_iters {
let request = StandardFitRequest {
data: design_data.clone(),
y: Arc::clone(&y),
weights: Arc::clone(&weights),
offset: Arc::clone(&offset),
spec: spec.clone(),
family: gaussian_family.clone(),
estimate_tweedie_p: false,
options: options.clone(),
kappa_options: kappa_options.clone(),
wiggle: None,
coefficient_groups: coefficient_groups.clone(),
penalty_block_gamma_priors: penalty_block_gamma_priors.clone(),
latent_coord: None,
};
let result = fit_standard_model(request)
.map_err(|reason| WorkflowError::IntegrationFailed { reason })?;
let mu = result
.design
.apply(result.fit.beta.view())
.map_err(|error| WorkflowError::IntegrationFailed {
reason: format!("expectile LAWS could not evaluate fitted design: {error}"),
})?;
if mu.len() != n {
return Err(WorkflowError::IntegrationFailed {
reason: format!(
"expectile LAWS: fitted mean length {} disagrees with response length {n}",
mu.len()
),
});
}
let mut mu_off = mu;
mu_off += offset.as_ref();
let sign: Vec<bool> = (0..n).map(|i| y[i] > mu_off[i]).collect();
let next_weights = expectile_row_weights(y.view(), mu_off.view(), base_weights.view(), tau);
let residual = y.as_ref() - &mu_off;
let kkt = expectile_kkt_residual(
&result.design.design,
residual.view(),
weights.view(),
next_weights.view(),
)
.map_err(|reason| WorkflowError::IntegrationFailed {
reason: format!(
"expectile LAWS KKT audit failed at iteration {iteration} \
(rho_checkpoint={:?}): {reason}",
result.fit.log_lambdas.to_vec(),
),
})?;
let kkt_bound = options.tol;
if kkt <= kkt_bound {
return Ok(result);
}
last_kkt = (kkt, kkt_bound);
last_rho_checkpoint = result.fit.log_lambdas.to_vec();
if let Some(cycle_length) = sign_cycle.observe(&sign) {
return Err(WorkflowError::IntegrationFailed {
reason: format!(
"expectile LAWS entered a deterministic sign-pattern cycle without \
reaching the KKT fixed point of the convex asymmetric least-squares \
problem (tau={tau}, iterations={iteration}, cycle_length={cycle_length}, \
KKT residual={:.3e} vs scaled tolerance {:.3e}, \
rho_checkpoint={:?}); non-convergence is a typed error, never a \
best-effort fit",
kkt,
kkt_bound,
result.fit.log_lambdas.to_vec(),
),
});
}
weights = Arc::new(next_weights);
}
Err(WorkflowError::IntegrationFailed {
reason: format!(
"expectile LAWS exhausted its {max_laws_iters}-iteration safety cap without a \
KKT certificate for the convex asymmetric least-squares problem (tau={tau}, \
final KKT residual={:.3e} vs scaled tolerance {:.3e}, \
rho_checkpoint={last_rho_checkpoint:?}); the iteration cap \
never selects the estimator — non-convergence is a typed error",
last_kkt.0, last_kkt.1,
),
})
}
pub fn spline_scan_fast_path(request: &StandardFitRequest<'_>) -> Option<SplineScanInputs> {
if !request.family.is_gaussian_identity() {
return None;
}
if request.wiggle.is_some()
|| request.latent_coord.is_some()
|| !request.coefficient_groups.is_empty()
|| !request.penalty_block_gamma_priors.is_empty()
{
return None;
}
let options = &request.options;
if options.latent_cloglog.is_some()
|| options.mixture_link.is_some()
|| options.sas_link.is_some()
|| options.linear_constraints.is_some()
|| options.adaptive_regularization.is_some()
|| options.kronecker_penalty_system.is_some()
|| options.kronecker_factored.is_some()
|| options.firth_bias_reduction
|| !options.nullspace_dims.is_empty()
{
return None;
}
let spec = &request.spec;
if !spec.linear_terms.is_empty()
|| !spec.random_effect_terms.is_empty()
|| spec.smooth_terms.len() != 1
{
return None;
}
let term = &spec.smooth_terms[0];
if !matches!(term.shape, gam_terms::smooth::ShapeConstraint::None)
|| term.joint_null_rotation.is_some()
{
return None;
}
let gam_terms::smooth::SmoothBasisSpec::BSpline1D {
feature_col,
spec: bspec,
} = &term.basis
else {
return None;
};
let order = bspec.penalty_order;
if !(1..=3).contains(&order)
|| bspec.degree != 2 * order - 1
|| bspec.double_penalty
|| !bspec.boundary_conditions.is_free()
|| !matches!(bspec.boundary, gam_terms::basis::OneDimensionalBoundary::Open)
|| matches!(
bspec.knotspec,
gam_terms::basis::BSplineKnotSpec::PeriodicUniform { .. }
| gam_terms::basis::BSplineKnotSpec::NaturalCubicRegression { .. }
)
|| matches!(
bspec.knotspec,
gam_terms::basis::BSplineKnotSpec::NaturalCubicRegression { .. }
)
{
return None;
}
if request.offset.iter().any(|&v| v != 0.0) {
return None;
}
if request.weights.iter().any(|&v| !(v.is_finite() && v > 0.0)) {
return None;
}
if *feature_col >= request.data.ncols() || request.y.len() != request.data.nrows() {
return None;
}
let x: Vec<f64> = request.data.column(*feature_col).iter().copied().collect();
let y: Vec<f64> = request.y.iter().copied().collect();
let w: Vec<f64> = request.weights.iter().copied().collect();
if x.iter().any(|v| !v.is_finite()) || y.iter().any(|v| !v.is_finite()) {
return None;
}
let mut sorted = x.clone();
sorted.sort_by(f64::total_cmp);
sorted.dedup();
if sorted.len() < order + 1 {
return None;
}
Some(SplineScanInputs { x, y, w, order })
}
pub fn fit_spline_scan_from_formula(
formula: &str,
data: &Dataset,
config: &FitConfig,
) -> Result<Option<gam_solve::spline_scan::SplineScanFit>, WorkflowError> {
let mat = materialize(formula, data, config)?;
let FitRequest::Standard(request) = mat.request else {
return Ok(None);
};
let Some(inputs) = spline_scan_fast_path(&request) else {
return Ok(None);
};
gam_solve::spline_scan::fit_spline_scan(&inputs.x, &inputs.y, &inputs.w, inputs.order)
.map(Some)
.map_err(|reason| WorkflowError::IntegrationFailed { reason })
}
pub fn constant_curvature_profiled_reml_scores(
formula: &str,
data: &Dataset,
config: &FitConfig,
kappas: &[f64],
) -> Result<Vec<(f64, f64)>, WorkflowError> {
let mat = materialize(formula, data, config)?;
let FitRequest::Standard(request) = mat.request else {
return Err(WorkflowError::IntegrationFailed {
reason: "constant_curvature_profiled_reml_scores: formula did not materialise to a \
standard fit request"
.to_string(),
});
};
let term_idx =
*crate::fit_orchestration::drivers::constant_curvature_term_indices(&request.spec)
.first()
.ok_or_else(|| WorkflowError::IntegrationFailed {
reason:
"constant_curvature_profiled_reml_scores: formula has no constant-curvature \
curv() term"
.to_string(),
})?;
let mut out = Vec::with_capacity(kappas.len());
for &kappa in kappas {
let score = crate::fit_orchestration::drivers::fixed_kappa_profiled_reml_score(
request.data.view(),
request.y.view(),
request.weights.view(),
request.offset.view(),
&request.spec,
term_idx,
kappa,
request.family.clone(),
&request.options,
)
.map_err(|e| WorkflowError::IntegrationFailed {
reason: format!(
"constant_curvature_profiled_reml_scores: fixed-κ fit at κ={kappa} failed: {e}"
),
})?;
out.push((kappa, score));
}
Ok(out)
}
fn past_dense_kernel_cliff(n: usize, d: usize) -> bool {
const DENSE_CENTER_CAP: usize = 2000;
gam_terms::basis::default_num_centers(n, d) >= DENSE_CENTER_CAP
}
fn cascade_sobolev_order(requested: f64, d: usize) -> f64 {
let lo = d as f64 / 2.0;
let hi = (d as f64 + 3.0) / 2.0;
let eps = 1e-6 * (hi - lo);
requested.clamp(lo + eps, hi)
}
pub fn residual_cascade_fast_path(
request: &StandardFitRequest<'_>,
) -> Option<ResidualCascadeInputs> {
if !request.family.is_gaussian_identity() {
return None;
}
if request.wiggle.is_some()
|| request.latent_coord.is_some()
|| !request.coefficient_groups.is_empty()
|| !request.penalty_block_gamma_priors.is_empty()
{
return None;
}
let options = &request.options;
if options.latent_cloglog.is_some()
|| options.mixture_link.is_some()
|| options.sas_link.is_some()
|| options.linear_constraints.is_some()
|| options.adaptive_regularization.is_some()
|| options.kronecker_penalty_system.is_some()
|| options.kronecker_factored.is_some()
|| options.firth_bias_reduction
|| !options.nullspace_dims.is_empty()
{
return None;
}
let spec = &request.spec;
if !spec.linear_terms.is_empty()
|| !spec.random_effect_terms.is_empty()
|| spec.smooth_terms.len() != 1
{
return None;
}
let term = &spec.smooth_terms[0];
if !matches!(term.shape, gam_terms::smooth::ShapeConstraint::None)
|| term.joint_null_rotation.is_some()
{
return None;
}
let (feature_cols, requested_s) = match &term.basis {
gam_terms::smooth::SmoothBasisSpec::Duchon {
feature_cols, spec, ..
} => {
let p = match spec.nullspace_order {
gam_terms::basis::DuchonNullspaceOrder::Zero => 0.0,
gam_terms::basis::DuchonNullspaceOrder::Linear => 1.0,
gam_terms::basis::DuchonNullspaceOrder::Degree(k) => k as f64,
};
(feature_cols, spec.power + p)
}
gam_terms::smooth::SmoothBasisSpec::Matern {
feature_cols, spec, ..
} => {
let nu = spec.nu.half_integer_value();
(feature_cols, nu + feature_cols.len() as f64 / 2.0)
}
_ => return None,
};
let d = feature_cols.len();
if !(2..=3).contains(&d) {
return None;
}
if request.offset.iter().any(|&v| v != 0.0) {
return None;
}
if request.weights.iter().any(|&v| !(v.is_finite() && v > 0.0)) {
return None;
}
let n = request.y.len();
if n != request.data.nrows() || feature_cols.iter().any(|&c| c >= request.data.ncols()) {
return None;
}
if !past_dense_kernel_cliff(n, d) {
return None;
}
let coords: Vec<Vec<f64>> = feature_cols
.iter()
.map(|&c| request.data.column(c).iter().copied().collect())
.collect();
let y: Vec<f64> = request.y.iter().copied().collect();
let w: Vec<f64> = request.weights.iter().copied().collect();
if coords
.iter()
.any(|axis| axis.iter().any(|v| !v.is_finite()))
|| y.iter().any(|v| !v.is_finite())
{
return None;
}
let metric = vec![1.0_f64; d];
let sobolev_s = cascade_sobolev_order(requested_s, d);
Some(ResidualCascadeInputs {
coords,
y,
w,
metric,
sobolev_s,
})
}
pub fn fit_residual_cascade_from_formula(
formula: &str,
data: &Dataset,
config: &FitConfig,
) -> Result<Option<gam_solve::residual_cascade::ResidualCascadeFit>, WorkflowError> {
let mat = materialize(formula, data, config)?;
let FitRequest::Standard(request) = mat.request else {
return Ok(None);
};
let Some(inputs) = residual_cascade_fast_path(&request) else {
return Ok(None);
};
let coord_refs: Vec<&[f64]> = inputs.coords.iter().map(Vec::as_slice).collect();
match gam_solve::residual_cascade::fit_residual_cascade(
&coord_refs,
&inputs.y,
&inputs.w,
&inputs.metric,
inputs.sobolev_s,
) {
Ok(fit) => Ok(Some(fit)),
Err(_) => Ok(None),
}
}
fn family_requests_transformation_normal(family: Option<&str>) -> bool {
family
.map(|name| name.trim().to_ascii_lowercase().replace('_', "-"))
.as_deref()
== Some("transformation-normal")
}
pub fn materialize<'a>(
formula: &str,
data: &'a Dataset,
config: &FitConfig,
) -> Result<MaterializedModel<'a>, WorkflowError> {
materialize_impl(formula, data, config, false)
}
pub fn materialize_structural<'a>(
formula: &str,
data: &'a Dataset,
config: &FitConfig,
) -> Result<MaterializedModel<'a>, WorkflowError> {
materialize_impl(formula, data, config, true)
}
fn materialize_impl<'a>(
formula: &str,
data: &'a Dataset,
config: &FitConfig,
structural_only: bool,
) -> Result<MaterializedModel<'a>, WorkflowError> {
let config = config
.clone()
.resolve()
.map_err(|reason| WorkflowError::InvalidConfig { reason })?;
let config = &config;
gam_gpu::configure_global_policy(config.gpu_policy);
let parsed = parse_formula(formula)?;
let col_map = data.column_map();
let family_transformation_normal =
family_requests_transformation_normal(config.family.as_deref());
let transformation_normal_config;
let effective_config = if family_transformation_normal && !config.transformation_normal {
transformation_normal_config = FitConfig {
transformation_normal: true,
..config.clone()
};
&transformation_normal_config
} else {
config
};
if let Some((left_col, right_col, event_col)) = parse_surv_interval_response(&parsed.response)?
{
if effective_config.transformation_normal {
return Err(WorkflowError::InvalidConfig {
reason:
"transformation_normal cannot be combined with a SurvInterval(...) response"
.to_string(),
});
}
materialize_survival(
&parsed,
data,
&col_map,
effective_config,
None,
&left_col,
&event_col,
Some(&right_col),
structural_only,
)
} else if let Some((entry_col, exit_col, event_col)) = parse_surv_response(&parsed.response)? {
if effective_config.transformation_normal {
return Err(WorkflowError::InvalidConfig {
reason: "transformation_normal cannot be combined with a Surv(...) response"
.to_string(),
});
}
materialize_survival(
&parsed,
data,
&col_map,
effective_config,
entry_col.as_deref(),
&exit_col,
&event_col,
None,
structural_only,
)
} else {
reject_survival_only_terms_for_nonsurvival(&parsed)?;
reject_survival_likelihood_for_nonsurvival(effective_config)?;
if effective_config.transformation_normal {
reject_marginal_slope_controls_for_transformation_normal(effective_config)?;
if effective_config.noise_formula.is_some() {
return Err(WorkflowError::InvalidConfig {
reason: "transformation_normal cannot be combined with noise_formula"
.to_string(),
});
}
materialize_transformation_normal(&parsed, data, &col_map, effective_config)
} else if requests_bernoulli_marginal_slope(effective_config) {
materialize_bernoulli_marginal_slope(&parsed, data, &col_map, effective_config)
} else if effective_config.noise_formula.is_some() {
materialize_location_scale(&parsed, data, &col_map, effective_config)
} else {
materialize_standard(&parsed, data, &col_map, effective_config)
}
}
}
#[cfg(test)]
mod sz_factor_smooth_recovery_tests {
use super::*;
const NOISE_SD: f64 = 0.20;
const N: usize = 4000;
const N_GROUPS: usize = 4;
struct Lcg(u64);
impl Lcg {
fn next_u64(&mut self) -> u64 {
self.0 = self
.0
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
self.0
}
fn unif(&mut self) -> f64 {
(self.next_u64() >> 11) as f64 / (1u64 << 53) as f64
}
fn normal(&mut self) -> f64 {
let u1 = (self.unif()).max(1e-12);
let u2 = self.unif();
(-2.0 * u1.ln()).sqrt() * (std::f64::consts::TAU * u2).cos()
}
}
fn sz_class_dataset() -> (Dataset, tempfile::TempDir) {
let mut rng = Lcg(0x5326_2026_0628_1605);
let phases: Vec<f64> = (0..N_GROUPS)
.map(|k| 1.2 * k as f64 / (N_GROUPS as f64 - 1.0))
.collect();
let deviations = |xi: f64| -> Vec<f64> {
let vals: Vec<f64> = phases
.iter()
.map(|p| 0.6 * (std::f64::consts::TAU * xi + std::f64::consts::TAU * p).sin())
.collect();
let mean = vals.iter().sum::<f64>() / vals.len() as f64;
vals.iter().map(|v| v - mean).collect()
};
let mut csv = String::from("y,x,g\n");
for _ in 0..N {
let x = rng.unif();
let g = ((rng.unif() * N_GROUPS as f64) as usize).min(N_GROUPS - 1);
let f0 = (std::f64::consts::TAU * x).sin();
let mu = f0 + deviations(x)[g];
let y = mu + NOISE_SD * rng.normal();
csv.push_str(&format!("{y},{x},g{g}\n"));
}
let td = tempfile::tempdir().expect("tempdir");
let path = td.path().join("sz_class.csv");
std::fs::write(&path, csv).expect("write sz-class csv");
let mut roles = std::collections::HashSet::new();
roles.insert("g");
let data = gam_data::load_dataset_projected_with_categorical_roles(
&path,
&["y".to_string(), "x".to_string(), "g".to_string()],
&roles,
)
.expect("load sz-class dataset");
(data, td)
}
fn gaussian_config() -> FitConfig {
FitConfig {
family: Some("gaussian".to_string()),
..FitConfig::default()
}
}
fn residual_sd(fit: &StandardFitResult, data: &Dataset) -> f64 {
let beta = &fit.fit.beta;
let design = &fit.design.design;
let n = design.nrows();
assert_eq!(design.ncols(), beta.len(), "design/beta width mismatch");
let mut fitted = vec![0.0f64; n];
const CHUNK: usize = 512;
let mut start = 0usize;
while start < n {
let end = (start + CHUNK).min(n);
let block = design
.try_row_chunk(start..end)
.expect("materialize design row chunk");
for (r, row) in block.rows().into_iter().enumerate() {
let mut acc = 0.0;
for (c, &xv) in row.iter().enumerate() {
acc += xv * beta[c];
}
fitted[start + r] = acc;
}
start = end;
}
let y = data.values.column(0);
let resid: Vec<f64> = y
.iter()
.zip(fitted.iter())
.map(|(&yi, &fi)| yi - fi)
.collect();
let mean = resid.iter().sum::<f64>() / resid.len() as f64;
let var = resid.iter().map(|r| (r - mean).powi(2)).sum::<f64>() / resid.len() as f64;
var.sqrt()
}
fn fit_standard(formula: &str, data: &Dataset) -> StandardFitResult {
match fit_from_formula(formula, data, &gaussian_config())
.unwrap_or_else(|e| panic!("fit `{formula}` failed: {e:?}"))
{
FitResult::Standard(r) => r,
other => panic!(
"expected Standard fit for `{formula}`, got a different variant: {}",
std::any::type_name_of_val(&other)
),
}
}
#[test]
fn sz_factor_smooth_recovers_its_own_model_class_end_to_end() {
let (data, _td) = sz_class_dataset();
let fs_fit = fit_standard("y ~ s(x, g, bs='fs')", &data);
let fs_resid = residual_sd(&fs_fit, &data);
assert!(
fs_resid < 1.2 * NOISE_SD,
"control bs='fs' did not reach the noise floor: resid_sd={fs_resid:.4} \
vs noise_sd={NOISE_SD} (data/floor sanity check)",
);
let sz_fit = fit_standard("y ~ s(x) + s(g, x, bs='sz')", &data);
let sz_resid = residual_sd(&sz_fit, &data);
assert!(
sz_resid < 1.4 * NOISE_SD,
"bs='sz' under-fits its own model class: resid_sd={sz_resid:.4} \
({:.2}x the noise floor {NOISE_SD}); the bs='fs' superset reached \
{fs_resid:.4}. The sz fit leaves systematic signal in the residual.",
sz_resid / NOISE_SD,
);
assert!(
sz_resid < 1.5 * fs_resid,
"bs='sz' residual {sz_resid:.4} is {:.2}x the bs='fs' residual \
{fs_resid:.4} on identical sz-class data",
sz_resid / fs_resid,
);
}
}