use super::*;
use gam_solve::estimate::reml::reml_outer_engine::penalty_matrix_root;
pub(crate) fn survival_inverse_link_has_free_parameters(link: &InverseLink) -> bool {
match link {
InverseLink::Sas(_) | InverseLink::BetaLogistic(_) => true,
InverseLink::Mixture(state) => !state.rho.is_empty(),
InverseLink::LatentCLogLog(_) | InverseLink::Standard(_) => false,
}
}
#[derive(Debug)]
struct ProfiledOuterPayload<T> {
theta: Array1<f64>,
objective: f64,
gradient: Array1<f64>,
value: T,
}
fn consume_certified_profiled_outer_payload<T>(
selected: Option<ProfiledOuterPayload<T>>,
outer: &gam_solve::rho_optimizer::CertifiedOuterResult,
context: &str,
) -> Result<ProfiledOuterPayload<T>, String> {
let selected = selected
.ok_or_else(|| format!("{context} retained no optimizer-installed terminal profile"))?;
if selected.theta.len() != outer.rho().len()
|| selected
.theta
.iter()
.zip(outer.rho().iter())
.any(|(selected, certified)| selected.to_bits() != certified.to_bits())
{
return Err(format!(
"{context} terminal profile hyperparameters do not bitwise match the certified optimum"
));
}
if selected.objective.to_bits() != outer.final_value().to_bits() {
return Err(format!(
"{context} terminal profile objective does not bitwise match the certified optimum: selected={:.17e}, certified={:.17e}",
selected.objective,
outer.final_value(),
));
}
let certified_gradient = outer.final_gradient().ok_or_else(|| {
format!("{context} certified result retained no analytic terminal gradient")
})?;
if selected.gradient.len() != certified_gradient.len()
|| selected
.gradient
.iter()
.zip(certified_gradient.iter())
.any(|(selected, certified)| selected.to_bits() != certified.to_bits())
{
return Err(format!(
"{context} terminal profile gradient does not bitwise match the certified optimum"
));
}
Ok(selected)
}
#[cfg(test)]
mod profiled_outer_payload_tests {
use super::*;
use gam_problem::{DeclaredHessianForm, Derivative, HessianValue, OuterEval};
use gam_solve::rho_optimizer::OuterProblem;
fn certified_quadratic() -> gam_solve::rho_optimizer::CertifiedOuterResult {
let problem = OuterProblem::new(1)
.with_gradient(Derivative::Analytic)
.with_hessian(DeclaredHessianForm::Unavailable)
.with_tolerance(1.0e-8)
.with_max_iter(40)
.with_initial_rho(Array1::from_vec(vec![0.5]))
.with_seed_config(gam_problem::SeedConfig {
max_seeds: 1,
seed_budget: 1,
..Default::default()
});
let mut objective = problem.build_objective(
(),
|_: &mut (), theta: &Array1<f64>| Ok(0.5 * (theta[0] - 0.25).powi(2)),
|_: &mut (), theta: &Array1<f64>| {
Ok(OuterEval {
cost: 0.5 * (theta[0] - 0.25).powi(2),
gradient: Array1::from_vec(vec![theta[0] - 0.25]),
hessian: HessianValue::Unavailable,
inner_beta_hint: None,
})
},
None::<fn(&mut ())>,
None::<
fn(
&mut (),
&Array1<f64>,
)
-> Result<gam_problem::EfsEval, gam_solve::estimate::EstimationError>,
>,
);
problem
.run_certified(&mut objective, "profiled-payload unit")
.expect("quadratic outer problem must certify")
}
fn matching_payload(
outer: &gam_solve::rho_optimizer::CertifiedOuterResult,
) -> ProfiledOuterPayload<&'static str> {
ProfiledOuterPayload {
theta: outer.rho().clone(),
objective: outer.final_value(),
gradient: outer
.final_gradient()
.expect("analytic fixture must retain its terminal gradient")
.clone(),
value: "terminal profile",
}
}
#[test]
fn selected_profile_requires_theta_objective_and_gradient_identity() {
let outer = certified_quadratic();
let mut wrong_theta = matching_payload(&outer);
wrong_theta.theta[0] += 1.0;
assert!(
consume_certified_profiled_outer_payload(
Some(wrong_theta),
&outer,
"theta substitution",
)
.expect_err("theta substitution must be rejected")
.contains("hyperparameters")
);
let mut wrong_objective = matching_payload(&outer);
wrong_objective.objective = f64::from_bits(wrong_objective.objective.to_bits() + 1);
assert!(
consume_certified_profiled_outer_payload(
Some(wrong_objective),
&outer,
"objective substitution",
)
.expect_err("objective substitution must be rejected")
.contains("objective")
);
let mut wrong_gradient = matching_payload(&outer);
wrong_gradient.gradient[0] += 1.0;
assert!(
consume_certified_profiled_outer_payload(
Some(wrong_gradient),
&outer,
"gradient substitution",
)
.expect_err("gradient substitution must be rejected")
.contains("gradient")
);
let selected = consume_certified_profiled_outer_payload(
Some(matching_payload(&outer)),
&outer,
"valid terminal profile",
)
.expect("the exact runner-installed terminal payload must be consumable");
assert_eq!(selected.value, "terminal profile");
assert!(outer.criterion_certificate().certifies());
}
}
const SURVIVAL_TRANSFORMATION_PIRLS_MAX_ITERATIONS: usize = 400;
const SURVIVAL_TRANSFORMATION_PIRLS_CONVERGENCE_TOL: f64 =
crate::survival::SURVIVAL_LAML_STATIONARITY_RELATIVE_TOL;
const SURVIVAL_TRANSFORMATION_PIRLS_MAX_STEP_HALVING: usize = 40;
const SURVIVAL_TRANSFORMATION_PIRLS_MIN_STEP_SIZE: f64 = 1e-12;
const SURVIVAL_TRANSFORMATION_OUTER_STALL_RESTARTS: usize = 10;
struct SurvivalLocationScaleProfile {
fit: SurvivalLocationScaleTermFitResult,
inverse_link: InverseLink,
wiggle_knots: Option<Array1<f64>>,
wiggle_degree: Option<usize>,
inverse_link_outer: Option<gam_solve::rho_optimizer::CertifiedOuterResult>,
}
fn survival_inverse_link_profile_objective(
profile: &SurvivalLocationScaleProfile,
context: &str,
) -> Result<f64, String> {
let objective = -profile.fit.fit.log_likelihood + 0.5 * profile.fit.fit.stable_penalty_term;
if objective.is_finite() {
Ok(objective)
} else {
Err(format!(
"{context}: non-finite profile objective (log_likelihood={}, stable_penalty_term={})",
profile.fit.fit.log_likelihood, profile.fit.fit.stable_penalty_term,
))
}
}
fn survival_pirls_status_is_certified(status: gam_solve::pirls::PirlsStatus) -> bool {
status.is_converged()
}
fn require_certified_survival_pirls(
summary: &gam_solve::pirls::WorkingModelPirlsResult,
context: &str,
parameter_checkpoint: &[f64],
durable_checkpoint_key: Option<&str>,
) -> Result<(), String> {
if survival_pirls_status_is_certified(summary.status) {
return Ok(());
}
Err(format!(
"{context} did not produce a strict PIRLS convergence certificate \
(status={:?}, iterations={}, projected_gradient_norm={:.6e}, \
deviance={:.6e}, min_penalized_deviance={:.6e}, last_step_size={:.6e}, \
last_step_halving={}, parameter_checkpoint={parameter_checkpoint:?}{}). The accepted \
iterate is checkpoint evidence only; no fit was minted.",
summary.status,
summary.iterations,
summary.lastgradient_norm,
summary.state.deviance,
summary.min_penalized_deviance,
summary.last_step_size,
summary.last_step_halving,
durable_checkpoint_key
.map(|key| format!(", durable_checkpoint_key={key}"))
.unwrap_or_default(),
))
}
fn survival_baseline_parameter_checkpoint(
config: &crate::survival::construction::SurvivalBaselineConfig,
) -> Result<Vec<f64>, String> {
let required = |name: &str, value: Option<f64>| {
value
.filter(|candidate| candidate.is_finite())
.ok_or_else(|| format!("survival baseline checkpoint is missing finite {name}"))
};
let positive_log = |name: &str, value: Option<f64>| {
let value = required(name, value)?;
if value > 0.0 {
Ok(value.ln())
} else {
Err(format!(
"survival baseline checkpoint requires positive {name}, got {value}"
))
}
};
use crate::survival::construction::SurvivalBaselineTarget;
match config.target {
SurvivalBaselineTarget::Linear => Ok(Vec::new()),
SurvivalBaselineTarget::Weibull => Ok(vec![
positive_log("Weibull scale", config.scale)?,
positive_log("Weibull shape", config.shape)?,
]),
SurvivalBaselineTarget::Gompertz => Ok(vec![
positive_log("Gompertz rate", config.rate)?,
required("Gompertz shape", config.shape)?,
]),
SurvivalBaselineTarget::GompertzMakeham => Ok(vec![
positive_log("Gompertz-Makeham rate", config.rate)?,
required("Gompertz-Makeham shape", config.shape)?,
positive_log("Gompertz-Makeham makeham", config.makeham)?,
]),
}
}
impl SurvivalLocationScaleProfile {
fn into_result(self) -> SurvivalLocationScaleFitResult {
SurvivalLocationScaleFitResult {
fit: self.fit,
inverse_link: self.inverse_link,
wiggle_knots: self.wiggle_knots,
wiggle_degree: self.wiggle_degree,
inverse_link_outer: self.inverse_link_outer,
}
}
}
fn resolved_wiggle_inverse_link(
spec: &LikelihoodSpec,
fit: &UnifiedFitResult,
fallback: &InverseLink,
) -> Result<InverseLink, String> {
let resolved = match fit.fitted_link_state(spec).map_err(|e| e.to_string())? {
FittedLinkState::Standard(Some(link)) => InverseLink::Standard(link),
FittedLinkState::Standard(None) => fallback.clone(),
FittedLinkState::LatentCLogLog { state } => InverseLink::LatentCLogLog(state),
FittedLinkState::Sas { state, .. } => InverseLink::Sas(state),
FittedLinkState::BetaLogistic { state, .. } => InverseLink::BetaLogistic(state),
FittedLinkState::Mixture { state, .. } => InverseLink::Mixture(state),
};
require_inverse_link_supports_joint_wiggle(&resolved, "standard link wiggle")?;
Ok(resolved)
}
type StandardBaseFit = crate::fit_orchestration::drivers::FittedTermCollectionWithSpec;
fn fit_standard_base(
request: &StandardFitRequest<'_>,
family: &LikelihoodSpec,
options: &FitOptions,
) -> Result<StandardBaseFit, gam_solve::estimate::EstimationError> {
if let Some(latent_coord) = request.latent_coord.as_ref() {
if !request.coefficient_groups.is_empty() || !request.penalty_block_gamma_priors.is_empty()
{
return Err(gam_solve::estimate::EstimationError::InvalidInput(
"latent-coordinate standard fits do not support coefficient_groups or \
penalty_block_gamma_priors in the same request"
.to_string(),
));
}
fit_term_collectionwith_latent_coord_optimization(
request.data.view(),
request.y.as_ref().clone(),
request.weights.as_ref().clone(),
request.offset.as_ref().clone(),
&request.spec,
latent_coord,
family.clone(),
options,
)
} else if !request.coefficient_groups.is_empty()
|| !request.penalty_block_gamma_priors.is_empty()
{
let fitted = fit_term_collection_with_coefficient_groups_and_penalty_block_gamma_priors(
request.data.view(),
request.y.view(),
request.weights.view(),
request.offset.view(),
&request.spec,
&request.coefficient_groups,
&request.penalty_block_gamma_priors,
family.clone(),
options,
)?;
let resolvedspec = crate::fit_orchestration::drivers::freeze_term_collection_from_design(
&request.spec,
&fitted.design,
)?;
Ok(
crate::fit_orchestration::drivers::FittedTermCollectionWithSpec {
fit: fitted.fit,
design: fitted.design,
resolvedspec,
adaptive_diagnostics: fitted.adaptive_diagnostics,
kappa_timing: None,
},
)
} else {
fit_term_collectionwith_spatial_length_scale_optimization(
request.data.view(),
request.y.as_ref().clone(),
request.weights.as_ref().clone(),
request.offset.as_ref().clone(),
&request.spec,
family.clone(),
options,
&request.kappa_options,
)
}
}
fn firth_can_rescue(error: &gam_solve::estimate::EstimationError) -> bool {
use gam_solve::estimate::EstimationError;
error.is_inner_solve_retreat()
|| matches!(
error,
EstimationError::PrefitPerfectSeparationDetected { .. }
| EstimationError::PrefitLinearSeparationDetected { .. }
| EstimationError::RemlDidNotConverge { .. }
)
}
fn certified_retry_or_original<T, E>(original: E, retry: Result<T, E>) -> Result<T, E> {
match retry {
Ok(value) => Ok(value),
Err(_) => Err(original),
}
}
fn rescale_covariance_coordinates(covariance: &mut Array2<f64>, factors: &[f64]) {
let dimension = factors.len();
assert_eq!(
covariance.dim(),
(dimension, dimension),
"covariance must align with the remapped coefficient vector"
);
for i in 0..dimension {
for j in 0..dimension {
covariance[[i, j]] *= factors[i] * factors[j];
}
}
}
fn rescale_precision_coordinates(precision: &mut Array2<f64>, factors: &[f64]) {
let dimension = factors.len();
assert_eq!(
precision.dim(),
(dimension, dimension),
"precision must align with the remapped coefficient vector"
);
for i in 0..dimension {
for j in 0..dimension {
precision[[i, j]] /= factors[i] * factors[j];
}
}
}
fn rescale_influence_coordinates(matrix: &mut Array2<f64>, factors: &[f64]) {
let dimension = factors.len();
assert_eq!(
matrix.dim(),
(dimension, dimension),
"influence map must align with the remapped coefficient vector"
);
for i in 0..dimension {
for j in 0..dimension {
matrix[[i, j]] *= factors[i] / factors[j];
}
}
}
#[cfg(test)]
mod standard_convergence_gate_tests {
use super::{
certified_retry_or_original, firth_can_rescue, rescale_covariance_coordinates,
rescale_precision_coordinates, survival_baseline_parameter_checkpoint,
survival_pirls_status_is_certified,
};
use crate::survival::construction::{SurvivalBaselineConfig, SurvivalBaselineTarget};
use gam_solve::estimate::EstimationError;
use gam_solve::pirls::PirlsStatus;
use ndarray::array;
#[test]
fn raw_coordinate_precision_is_the_inverse_congruence_of_covariance() {
let mut covariance = array![[0.30, -0.10], [-0.10, 0.70]];
let mut precision = array![[3.5, 0.5], [0.5, 1.5]];
let factors = [4.0, 1.0];
rescale_covariance_coordinates(&mut covariance, &factors);
rescale_precision_coordinates(&mut precision, &factors);
assert_eq!(covariance, array![[4.8, -0.4], [-0.4, 0.7]]);
assert_eq!(precision, array![[0.21875, 0.125], [0.125, 1.5]]);
let identity = precision.dot(&covariance);
for i in 0..2 {
for j in 0..2 {
let target = if i == j { 1.0 } else { 0.0 };
assert!((identity[[i, j]] - target).abs() <= 2e-15);
}
}
}
#[test]
fn failed_retry_returns_original_evidence() {
let result = certified_retry_or_original::<(), _>("base evidence", Err("retry evidence"));
assert_eq!(result, Err("base evidence"));
assert_eq!(
certified_retry_or_original("base evidence", Ok::<_, &str>(7)),
Ok(7)
);
}
#[test]
fn survival_gate_rejects_every_exhausted_or_stalled_status() {
assert!(survival_pirls_status_is_certified(PirlsStatus::Converged));
for status in [
PirlsStatus::StalledAtValidMinimum,
PirlsStatus::MaxIterationsReached,
PirlsStatus::LmStepSearchExhausted,
PirlsStatus::Unstable,
] {
assert!(!survival_pirls_status_is_certified(status));
}
}
#[test]
fn survival_baseline_checkpoint_matches_outer_coordinates() {
let checkpoint = survival_baseline_parameter_checkpoint(&SurvivalBaselineConfig {
target: SurvivalBaselineTarget::GompertzMakeham,
scale: None,
shape: Some(-0.25),
rate: Some(2.0),
makeham: Some(4.0),
})
.expect("valid baseline checkpoint");
assert_eq!(checkpoint, vec![2.0_f64.ln(), -0.25, 4.0_f64.ln()]);
}
#[test]
fn firth_retry_is_limited_to_separation_and_nonconvergence() {
assert!(firth_can_rescue(&EstimationError::PirlsDidNotConverge {
max_iterations: 20,
last_change: 1.0,
}));
assert!(firth_can_rescue(
&EstimationError::PrefitPerfectSeparationDetected {
column_index: 0,
threshold: 0.0,
positive_above_threshold: true,
}
));
assert!(!firth_can_rescue(&EstimationError::InvalidInput(
"structural mismatch".to_string()
)));
}
}
fn tweedie_profile_loglik(request: &StandardFitRequest<'_>, p: f64) -> Option<f64> {
if !gam_spec::is_valid_tweedie_power(p) {
return None;
}
let family = LikelihoodSpec::new(ResponseFamily::Tweedie { p }, request.family.link.clone());
let fitted = fit_standard_base(request, &family, &request.options).ok()?;
let mut eta = fitted.design.apply(fitted.fit.beta.view()).ok()?;
if eta.len() != request.y.len() {
return None;
}
eta += request.offset.as_ref();
let mu = eta.mapv(f64::exp);
const PHI_MIN: f64 = 1e-6;
const PHI_MAX: f64 = 1e12;
let mut weighted_pearson = 0.0_f64;
let mut total_weight = 0.0_f64;
for ((&yi, &mui), &wi) in request.y.iter().zip(mu.iter()).zip(request.weights.iter()) {
let wi = wi.max(0.0);
if wi == 0.0 {
continue;
}
let resid = yi - mui;
let var_unit = mui.powf(p).max(f64::MIN_POSITIVE);
weighted_pearson += wi * resid * resid / var_unit;
total_weight += wi;
}
if total_weight <= 0.0 || !weighted_pearson.is_finite() || weighted_pearson <= 0.0 {
return None;
}
let phi = (weighted_pearson / total_weight).clamp(PHI_MIN, PHI_MAX);
let ll = gam_solve::pirls::tweedie_exact_loglik_total_from_eta(
request.y.view(),
eta.view(),
request.weights.view(),
p,
phi,
)
.ok()?;
Some(ll)
}
fn estimate_tweedie_power(request: &StandardFitRequest<'_>) -> Result<f64, String> {
const EPS: f64 = 1e-3;
const TOL: f64 = 1e-3;
let mut a = 1.0 + EPS;
let mut b = 2.0 - EPS;
let inv_phi_gr = (5.0_f64.sqrt() - 1.0) / 2.0; let eval = |p: f64| tweedie_profile_loglik(request, p).unwrap_or(f64::NEG_INFINITY);
let n_iter = (((TOL / (b - a)).ln() / inv_phi_gr.ln()).ceil() as i64).max(1) as usize;
let mut c = b - inv_phi_gr * (b - a);
let mut d = a + inv_phi_gr * (b - a);
let mut fc = eval(c);
let mut fd = eval(d);
for _ in 0..n_iter {
if fc >= fd {
b = d;
d = c;
fd = fc;
c = b - inv_phi_gr * (b - a);
fc = eval(c);
} else {
a = c;
c = d;
fc = fd;
d = a + inv_phi_gr * (b - a);
fd = eval(d);
}
}
let p_hat = (0.5 * (a + b)).clamp(1.0 + EPS, 2.0 - EPS);
if !tweedie_profile_loglik(request, p_hat).is_some_and(f64::is_finite) {
return Err(
"tweedie power profiling failed: the profile likelihood is non-finite across \
(1, 2); set an explicit power via family=\"tweedie(p)\""
.to_string(),
);
}
log::info!(
"[tweedie#2026] estimated variance power p={p_hat:.4} by golden-section profile \
likelihood (bare family=\"tweedie\"); set family=\"tweedie(p)\" to pin it."
);
Ok(p_hat)
}
pub(crate) fn fit_standard_model(
mut request: StandardFitRequest<'_>,
) -> Result<StandardFitResult, String> {
if request.estimate_tweedie_p
&& matches!(request.family.response, ResponseFamily::Tweedie { .. })
{
let p_hat = estimate_tweedie_power(&request)?;
request.family = LikelihoodSpec::new(
ResponseFamily::Tweedie { p: p_hat },
request.family.link.clone(),
);
request.estimate_tweedie_p = false;
}
let is_firth_capable_binomial = request.family.supports_firth();
let base = fit_standard_base(&request, &request.family, &request.options);
let fitted = match base {
Ok(fitted) => fitted,
Err(original_error)
if is_firth_capable_binomial
&& !request.options.firth_bias_reduction
&& firth_can_rescue(&original_error) =>
{
let original_report = original_error.to_string();
let mut firth_options = request.options.clone();
firth_options.firth_bias_reduction = true;
let firth = fit_standard_base(&request, &request.family, &firth_options);
let firth_failure = firth.as_ref().err().map(ToString::to_string);
match certified_retry_or_original(original_error, firth) {
Ok(firth_fitted) => {
log::info!(
"[#1762/#2273] Firth-capable binomial base fit ({}) failed with \
retryable separation/non-convergence evidence ({original_report}); Firth \
bias-reduction retry certified — adopting it (Firth edf {:.2}).",
request.family.pretty_name(),
firth_fitted.fit.edf_total().unwrap_or(f64::NAN),
);
firth_fitted
}
Err(original_error) => {
log::warn!(
"[#1762/#2273] Firth-capable binomial base fit ({}) failed \
({original_report}); Firth retry also failed to certify ({}) — returning \
the original typed base evidence, not either abandoned iterate.",
request.family.pretty_name(),
firth_failure.unwrap_or_else(|| "unknown retry failure".to_string()),
);
return Err(original_error.to_string());
}
}
}
Err(error) => return Err(error.to_string()),
};
let adaptive_spatial_terms = adaptive_spatial_term_mask(&request.spec);
let adaptive_spatial_center_counts = adaptive_spatial_center_counts(&request.spec);
let result = StandardFitResult {
saved_link_state: fitted.fit.fitted_link.clone(),
fit: fitted.fit,
design: fitted.design,
resolvedspec: fitted.resolvedspec,
adaptive_spatial_terms: adaptive_spatial_terms.clone(),
adaptive_spatial_center_counts: adaptive_spatial_center_counts.clone(),
adaptive_diagnostics: fitted.adaptive_diagnostics,
kappa_timing: fitted.kappa_timing,
wiggle_knots: None,
wiggle_degree: None,
wiggle_penalty_metadata: None,
wiggle_saved_warp_beta: None,
wiggle_saved_index_shift: None,
};
let Some(wiggle) = request.wiggle else {
return Ok(result);
};
let mut wiggle_options = wiggle.refit_options.clone();
wiggle_options.compute_covariance = true;
let wiggle_link_kind =
resolved_wiggle_inverse_link(&request.family, &result.fit, &wiggle.link_kind)?;
let selected_wiggle_basis = select_binomial_mean_link_wiggle_basis_from_pilot(
&result.design,
&result.fit,
&WiggleBlockConfig {
degree: wiggle.wiggle.degree,
num_internal_knots: wiggle.wiggle.num_internal_knots,
penalty_order: 2,
double_penalty: wiggle.wiggle.double_penalty,
},
&wiggle.wiggle.penalty_orders,
)?;
let wiggle_penalty_metadata = selected_wiggle_basis.penalty_metadata.clone();
let solved = match fit_binomial_mean_wiggle_terms_with_selected_basis(
request.data.view(),
&result.resolvedspec,
&result.design,
&result.fit,
request.y.as_ref(),
request.weights.as_ref(),
wiggle_link_kind,
selected_wiggle_basis,
&wiggle_options,
&request.kappa_options,
) {
Ok(solved) => solved,
Err(e) => {
log::warn!("[linkwiggle] binomial mean link-wiggle joint solve did not converge ({e})");
return Err(format!(
"flexible/learnable link requested via link(type=flexible(...)) / \
linkwiggle(...), but the binomial mean link-wiggle joint solve did not \
converge ({e}). The fit was NOT silently downgraded to the fixed base \
link. Refit with a fixed link (e.g. logit/probit/cloglog) or adjust the \
wiggle spec (linkwiggle(internal_knots=...)). See gam#1596."
));
}
};
if solved.fit.beta_covariance().is_none() {
return Err(
"link-wiggle fit reached assembly without its joint [Mean, LinkWiggle] posterior covariance; no model was minted"
.to_string(),
);
}
Ok(StandardFitResult {
saved_link_state: result.saved_link_state,
fit: solved.fit,
design: solved.design,
resolvedspec: solved.resolvedspec,
adaptive_spatial_terms,
adaptive_spatial_center_counts,
adaptive_diagnostics: result.adaptive_diagnostics,
kappa_timing: result.kappa_timing,
wiggle_knots: Some(solved.wiggle_knots),
wiggle_degree: Some(solved.wiggle_degree),
wiggle_penalty_metadata: Some(wiggle_penalty_metadata),
wiggle_saved_warp_beta: solved.saved_warp_beta,
wiggle_saved_index_shift: solved.saved_index_shift,
})
}
struct LocationScaleWorkflowParts<'a, S> {
data: ArrayView2<'a, f64>,
spec: S,
wiggle: Option<LinkWiggleConfig>,
options: BlockwiseFitOptions,
kappa_options: SpatialLengthScaleOptimizationOptions,
}
trait LocationScaleWorkflowAdapter {
type Spec;
type Request<'a>;
type Result;
fn into_parts<'a>(request: Self::Request<'a>) -> LocationScaleWorkflowParts<'a, Self::Spec>;
fn fit_pilot(
data: ArrayView2<'_, f64>,
spec: &Self::Spec,
options: &BlockwiseFitOptions,
kappa_options: &SpatialLengthScaleOptimizationOptions,
) -> Result<BlockwiseTermFitResult, String>;
fn refit_with_selected_wiggle(
data: ArrayView2<'_, f64>,
spec: Self::Spec,
pilot: &BlockwiseTermFitResult,
wiggle_cfg: &LinkWiggleConfig,
options: &BlockwiseFitOptions,
kappa_options: &SpatialLengthScaleOptimizationOptions,
) -> Result<BlockwiseTermWiggleFitResult, String>;
fn fit_plain(
data: ArrayView2<'_, f64>,
spec: Self::Spec,
options: &BlockwiseFitOptions,
kappa_options: &SpatialLengthScaleOptimizationOptions,
) -> Result<BlockwiseTermFitResult, String>;
fn assemble_plain(fit: BlockwiseTermFitResult) -> Self::Result;
fn assemble_with_wiggle(
fit: BlockwiseTermFitResult,
wiggle_knots: Array1<f64>,
wiggle_degree: usize,
beta_link_wiggle: Option<Vec<f64>>,
) -> Self::Result;
}
fn fit_location_scale_with_optional_wiggle<A: LocationScaleWorkflowAdapter>(
request: A::Request<'_>,
) -> Result<A::Result, String> {
let LocationScaleWorkflowParts {
data,
spec,
wiggle,
options,
kappa_options,
} = A::into_parts(request);
let Some(wiggle_cfg) = wiggle else {
let mut fit_options = options.clone();
fit_options.compute_covariance = true;
let fit = A::fit_plain(data, spec, &fit_options, &kappa_options)?;
if fit.fit.beta_covariance().is_none() {
return Err(
"plain location-scale fit reached assembly without its joint posterior covariance; no model was minted"
.to_string(),
);
}
return Ok(A::assemble_plain(fit));
};
let pilot = A::fit_pilot(data, &spec, &options, &kappa_options)?;
let mut refit_options = options.clone();
refit_options.compute_covariance = true;
let solved = A::refit_with_selected_wiggle(
data,
spec,
&pilot,
&wiggle_cfg,
&refit_options,
&kappa_options,
)?;
let fit = solved.fit.fit;
if fit.beta_covariance().is_none() {
return Err(
"location-scale link-wiggle fit reached assembly without its joint posterior covariance; no model was minted"
.to_string(),
);
}
let beta_link_wiggle = fit.block_states.get(2).map(|b| b.beta.to_vec());
let assembled_fit = BlockwiseTermFitResult::try_from_parts(BlockwiseTermFitResultParts {
fit,
meanspec_resolved: solved.fit.meanspec_resolved,
noisespec_resolved: solved.fit.noisespec_resolved,
mean_design: solved.fit.mean_design,
noise_design: solved.fit.noise_design,
})?;
Ok(A::assemble_with_wiggle(
assembled_fit,
solved.wiggle_knots,
solved.wiggle_degree,
beta_link_wiggle,
))
}
struct GaussianLocationScaleWorkflow;
impl LocationScaleWorkflowAdapter for GaussianLocationScaleWorkflow {
type Spec = GaussianLocationScaleTermSpec;
type Request<'a> = GaussianLocationScaleFitRequest<'a>;
type Result = GaussianLocationScaleFitResult;
fn into_parts<'a>(request: Self::Request<'a>) -> LocationScaleWorkflowParts<'a, Self::Spec> {
LocationScaleWorkflowParts {
data: request.data,
spec: request.spec,
wiggle: request.wiggle,
options: request.options,
kappa_options: request.kappa_options,
}
}
fn fit_pilot(
data: ArrayView2<'_, f64>,
spec: &Self::Spec,
options: &BlockwiseFitOptions,
kappa_options: &SpatialLengthScaleOptimizationOptions,
) -> Result<BlockwiseTermFitResult, String> {
fit_gaussian_location_scale_terms(
data,
GaussianLocationScaleTermSpec {
y: spec.y.clone(),
weights: spec.weights.clone(),
meanspec: spec.meanspec.clone(),
log_sigmaspec: spec.log_sigmaspec.clone(),
mean_offset: spec.mean_offset.clone(),
log_sigma_offset: spec.log_sigma_offset.clone(),
},
options,
kappa_options,
)
}
fn refit_with_selected_wiggle(
data: ArrayView2<'_, f64>,
spec: Self::Spec,
pilot: &BlockwiseTermFitResult,
wiggle_cfg: &LinkWiggleConfig,
options: &BlockwiseFitOptions,
kappa_options: &SpatialLengthScaleOptimizationOptions,
) -> Result<BlockwiseTermWiggleFitResult, String> {
let selected_wiggle_basis = select_gaussian_location_scale_link_wiggle_basis_from_pilot(
pilot,
&WiggleBlockConfig {
degree: wiggle_cfg.degree,
num_internal_knots: wiggle_cfg.num_internal_knots,
penalty_order: 2,
double_penalty: wiggle_cfg.double_penalty,
},
&wiggle_cfg.penalty_orders,
)?;
fit_gaussian_location_scale_terms_with_selected_wiggle(
data,
spec,
selected_wiggle_basis,
options,
kappa_options,
)
}
fn fit_plain(
data: ArrayView2<'_, f64>,
spec: Self::Spec,
options: &BlockwiseFitOptions,
kappa_options: &SpatialLengthScaleOptimizationOptions,
) -> Result<BlockwiseTermFitResult, String> {
fit_gaussian_location_scale_terms(data, spec, options, kappa_options)
}
fn assemble_plain(fit: BlockwiseTermFitResult) -> Self::Result {
GaussianLocationScaleFitResult {
fit,
wiggle_knots: None,
wiggle_degree: None,
beta_link_wiggle: None,
response_scale: 1.0,
}
}
fn assemble_with_wiggle(
fit: BlockwiseTermFitResult,
wiggle_knots: Array1<f64>,
wiggle_degree: usize,
beta_link_wiggle: Option<Vec<f64>>,
) -> Self::Result {
GaussianLocationScaleFitResult {
fit,
wiggle_knots: Some(wiggle_knots),
wiggle_degree: Some(wiggle_degree),
beta_link_wiggle,
response_scale: 1.0,
}
}
}
struct BinomialLocationScaleWorkflow;
impl LocationScaleWorkflowAdapter for BinomialLocationScaleWorkflow {
type Spec = BinomialLocationScaleTermSpec;
type Request<'a> = BinomialLocationScaleFitRequest<'a>;
type Result = BinomialLocationScaleFitResult;
fn into_parts<'a>(request: Self::Request<'a>) -> LocationScaleWorkflowParts<'a, Self::Spec> {
LocationScaleWorkflowParts {
data: request.data,
spec: request.spec,
wiggle: request.wiggle,
options: request.options,
kappa_options: request.kappa_options,
}
}
fn fit_pilot(
data: ArrayView2<'_, f64>,
spec: &Self::Spec,
options: &BlockwiseFitOptions,
kappa_options: &SpatialLengthScaleOptimizationOptions,
) -> Result<BlockwiseTermFitResult, String> {
require_inverse_link_supports_joint_wiggle(
&spec.link_kind,
"binomial location-scale link wiggle",
)?;
fit_binomial_location_scale_terms(
data,
BinomialLocationScaleTermSpec {
y: spec.y.clone(),
weights: spec.weights.clone(),
link_kind: spec.link_kind.clone(),
thresholdspec: spec.thresholdspec.clone(),
log_sigmaspec: spec.log_sigmaspec.clone(),
threshold_offset: spec.threshold_offset.clone(),
log_sigma_offset: spec.log_sigma_offset.clone(),
},
options,
kappa_options,
)
}
fn refit_with_selected_wiggle(
data: ArrayView2<'_, f64>,
spec: Self::Spec,
pilot: &BlockwiseTermFitResult,
wiggle_cfg: &LinkWiggleConfig,
options: &BlockwiseFitOptions,
kappa_options: &SpatialLengthScaleOptimizationOptions,
) -> Result<BlockwiseTermWiggleFitResult, String> {
let selected_wiggle_basis = select_binomial_location_scale_link_wiggle_basis_from_pilot(
pilot,
&WiggleBlockConfig {
degree: wiggle_cfg.degree,
num_internal_knots: wiggle_cfg.num_internal_knots,
penalty_order: 2,
double_penalty: wiggle_cfg.double_penalty,
},
&wiggle_cfg.penalty_orders,
)?;
fit_binomial_location_scale_terms_with_selected_wiggle(
data,
spec,
selected_wiggle_basis,
options,
kappa_options,
)
}
fn fit_plain(
data: ArrayView2<'_, f64>,
spec: Self::Spec,
options: &BlockwiseFitOptions,
kappa_options: &SpatialLengthScaleOptimizationOptions,
) -> Result<BlockwiseTermFitResult, String> {
fit_binomial_location_scale_terms(data, spec, options, kappa_options)
}
fn assemble_plain(fit: BlockwiseTermFitResult) -> Self::Result {
BinomialLocationScaleFitResult {
fit,
wiggle_knots: None,
wiggle_degree: None,
beta_link_wiggle: None,
}
}
fn assemble_with_wiggle(
fit: BlockwiseTermFitResult,
wiggle_knots: Array1<f64>,
wiggle_degree: usize,
beta_link_wiggle: Option<Vec<f64>>,
) -> Self::Result {
BinomialLocationScaleFitResult {
fit,
wiggle_knots: Some(wiggle_knots),
wiggle_degree: Some(wiggle_degree),
beta_link_wiggle,
}
}
}
pub(crate) fn gaussian_response_sample_std(v: ArrayView1<'_, f64>) -> f64 {
if v.is_empty() {
return 0.0;
}
let n = v.len() as f64;
let mean = v.iter().copied().sum::<f64>() / n;
let var = v
.iter()
.copied()
.map(|x| {
let d = x - mean;
d * d
})
.sum::<f64>()
/ n.max(1.0);
var.max(0.0).sqrt()
}
pub(crate) fn rescale_gaussian_location_scale_to_raw(
result: &mut GaussianLocationScaleFitResult,
response_scale: f64,
) {
use gam_problem::BlockRole;
let s = response_scale;
assert!(
s.is_finite() && s > 0.0,
"Gaussian location-scale response rescale must be finite and positive, got {s}"
);
let ln_s = s.ln();
let scale_intercept_range = result.fit.noise_design.intercept_range.clone();
let mut joint_offset = 0usize;
for (block_idx, block) in result.fit.fit.blocks.iter_mut().enumerate() {
let block_len = block.beta.len();
match block.role {
BlockRole::Mean | BlockRole::Location | BlockRole::LinkWiggle => {
block.beta.mapv_inplace(|v| v * s);
if result.fit.fit.beta.len() >= joint_offset + block_len {
for i in 0..block_len {
result.fit.fit.beta[joint_offset + i] *= s;
}
}
if let Some(state) = result.fit.fit.block_states.get_mut(block_idx) {
state.beta.mapv_inplace(|v| v * s);
state.eta.mapv_inplace(|v| v * s);
}
}
BlockRole::Scale => {
for col in scale_intercept_range.clone() {
if col < block.beta.len() {
block.beta[col] += ln_s;
}
let joint_col = joint_offset + col;
if joint_col < result.fit.fit.beta.len() {
result.fit.fit.beta[joint_col] += ln_s;
}
if let Some(state) = result.fit.fit.block_states.get_mut(block_idx)
&& col < state.beta.len()
{
state.beta[col] += ln_s;
}
}
if let Some(state) = result.fit.fit.block_states.get_mut(block_idx) {
state.eta.mapv_inplace(|v| v + ln_s);
}
}
BlockRole::Time | BlockRole::Threshold => {
}
}
joint_offset += block_len;
}
if let Some(knots) = result.wiggle_knots.as_mut() {
knots.mapv_inplace(|v| v * s);
}
if let Some(beta_w) = result.beta_link_wiggle.as_mut() {
for coef in beta_w.iter_mut() {
*coef *= s;
}
}
let mut row_factors: Vec<f64> = Vec::new();
for block in &result.fit.fit.blocks {
let f = match block.role {
BlockRole::Mean | BlockRole::Location | BlockRole::LinkWiggle => s,
BlockRole::Scale | BlockRole::Time | BlockRole::Threshold => 1.0,
};
row_factors.extend(std::iter::repeat_n(f, block.beta.len()));
}
if let Some(cov) = result.fit.fit.covariance_conditional.as_mut() {
rescale_covariance_coordinates(cov, &row_factors);
}
if let Some(cov) = result.fit.fit.covariance_corrected.as_mut() {
rescale_covariance_coordinates(cov, &row_factors);
}
if let Some(geometry) = result.fit.fit.geometry.as_mut() {
rescale_precision_coordinates(&mut geometry.penalized_hessian.0, &row_factors);
}
if let Some(inference) = result.fit.fit.inference.as_mut() {
rescale_precision_coordinates(&mut inference.penalized_hessian.0, &row_factors);
if let Some(cov) = inference.beta_covariance.as_mut() {
rescale_covariance_coordinates(&mut cov.0, &row_factors);
}
if let Some(cov) = inference.beta_covariance_corrected.as_mut() {
rescale_covariance_coordinates(cov, &row_factors);
}
if let Some(cov) = inference.beta_covariance_frequentist.as_mut() {
rescale_covariance_coordinates(cov, &row_factors);
}
if let Some(correction) = inference.smoothing_correction.as_mut() {
rescale_covariance_coordinates(correction, &row_factors);
}
for se in [
inference.beta_standard_errors.as_mut(),
inference.beta_standard_errors_corrected.as_mut(),
]
.into_iter()
.flatten()
{
for (value, &factor) in se.iter_mut().zip(row_factors.iter()) {
*value *= factor;
}
}
if let Some(gram) = inference.weighted_gram.as_mut() {
rescale_precision_coordinates(gram, &row_factors);
}
if let Some(influence) = inference.coefficient_influence.as_mut() {
rescale_influence_coordinates(influence, &row_factors);
}
if let Some(jacobian) = inference.bias_correction_jacobian.as_mut() {
rescale_influence_coordinates(jacobian, &row_factors);
}
if let Some(bias) = inference.bias_correction_beta.as_mut() {
for (value, &factor) in bias.iter_mut().zip(row_factors.iter()) {
*value *= factor;
}
}
if let Some(qs) = inference.reparam_qs.as_mut() {
for (mut row, &factor) in qs.rows_mut().into_iter().zip(row_factors.iter()) {
row.mapv_inplace(|v| v * factor);
}
}
}
result.fit.fit.standard_deviation *= s;
result.fit.fit.max_abs_eta *= s;
if let Some(n_obs) = result
.fit
.fit
.block_states
.first()
.map(|state| state.eta.len() as f64)
.filter(|&n| n > 0.0)
{
let ln_s = s.ln();
result.fit.fit.log_likelihood -= n_obs * ln_s;
result.fit.fit.deviance += 2.0 * n_obs * ln_s;
result.fit.fit.reml_score += n_obs * ln_s;
result.fit.fit.penalized_objective += n_obs * ln_s;
}
result.response_scale = s;
}
pub(crate) fn fit_gaussian_location_scale_model(
mut request: GaussianLocationScaleFitRequest<'_>,
) -> Result<GaussianLocationScaleFitResult, String> {
let response_scale = gaussian_response_sample_std(request.spec.y.view()).max(1e-6);
if response_scale != 1.0 {
request.spec.y.mapv_inplace(|v| v / response_scale);
request
.spec
.mean_offset
.mapv_inplace(|v| v / response_scale);
}
let mut result =
fit_location_scale_with_optional_wiggle::<GaussianLocationScaleWorkflow>(request)?;
rescale_gaussian_location_scale_to_raw(&mut result, response_scale);
Ok(result)
}
pub(crate) fn fit_dispersion_location_scale_model(
request: DispersionLocationScaleFitRequest<'_>,
) -> Result<DispersionLocationScaleFitResult, String> {
let kind = request.spec.kind;
let fit = fit_dispersion_glm_location_scale_terms(
request.data,
request.spec,
&request.options,
&request.kappa_options,
)?;
Ok(DispersionLocationScaleFitResult { fit, kind })
}
pub(crate) fn fit_binomial_location_scale_model(
request: BinomialLocationScaleFitRequest<'_>,
) -> Result<BinomialLocationScaleFitResult, String> {
fit_location_scale_with_optional_wiggle::<BinomialLocationScaleWorkflow>(request)
}
fn survival_transformation_edf(
state: &gam_solve::pirls::WorkingState,
penalty_blocks: &[PenaltyBlock],
) -> Result<(f64, Vec<f64>, Vec<f64>, Array2<f64>), String> {
let h_dense = state.hessian.to_dense();
let (edf_total, edf_by_block, penalty_block_trace) =
survival_edf_from_dense_hessian(&h_dense, penalty_blocks)?;
Ok((edf_total, edf_by_block, penalty_block_trace, h_dense))
}
fn survival_edf_from_dense_hessian(
h_dense: &Array2<f64>,
penalty_blocks: &[PenaltyBlock],
) -> Result<(f64, Vec<f64>, Vec<f64>), String> {
let p = h_dense.nrows();
let h_sym = gam_linalg::matrix::SymmetricMatrix::Dense(h_dense.clone());
let factor = h_sym.factorize().map_err(|error| {
format!("survival edf: exact penalized-Hessian factorization failed: {error}")
})?;
let mut raw_traces = vec![0.0_f64; penalty_blocks.len()];
let mut block_ranks = vec![0_usize; penalty_blocks.len()];
let mut joint_penalty = Array2::<f64>::zeros((p, p));
for (kk, block) in penalty_blocks.iter().enumerate() {
let block_cols = block.range.end - block.range.start;
let penalty_rank = if block_cols == 0 {
0
} else {
penalty_matrix_root(&block.matrix)
.map_err(|error| {
format!("survival edf: penalty {kk} rank factorization failed: {error}")
})?
.nrows()
};
block_ranks[kk] = penalty_rank;
if block_cols > 0 {
let r = block.range.start..block.range.end;
let mut target = joint_penalty.slice_mut(ndarray::s![r.clone(), r]);
target += &block.matrix;
}
if block.lambda <= 0.0 || block_cols == 0 {
raw_traces[kk] = 0.0;
continue;
}
let mut rhs = Array2::<f64>::zeros((p, block_cols));
for c in 0..block_cols {
for r in 0..block_cols {
rhs[[block.range.start + r, c]] = block.matrix[[r, c]];
}
}
let sol = factor.solvemulti(&rhs).map_err(|e| {
let spectrum_note =
match gam_linalg::faer_ndarray::FaerEigh::eigh(h_dense, faer::Side::Lower) {
Ok((eigenvalues, eigenvectors)) => {
let mut min_idx = 0usize;
for (idx, value) in eigenvalues.iter().enumerate() {
if value.abs() < eigenvalues[min_idx].abs() {
min_idx = idx;
}
}
let max_abs = eigenvalues
.iter()
.fold(0.0_f64, |acc, &value| acc.max(value.abs()));
let flat_direction: Vec<f64> =
eigenvectors.column(min_idx).iter().copied().collect();
format!(
"penalized-Hessian spectrum: min_abs_eig={:.6e}, max_abs_eig={:.6e}, \
eigenvalues={:?}, flattest direction (coefficient loadings)={:?}",
eigenvalues[min_idx], max_abs, eigenvalues, flat_direction
)
}
Err(error) => {
format!("penalized-Hessian eigendecomposition also failed: {error:?}")
}
};
format!(
"survival edf trace solve failed for penalty block {kk} \
(lambda={:.6e}, block_cols={block_cols}): {e}; {spectrum_note}",
block.lambda
)
})?;
let mut trace = 0.0_f64;
for j in 0..block_cols {
trace += sol[[block.range.start + j, j]];
}
raw_traces[kk] = block.lambda * trace;
}
let joint_penalty_rank = penalty_matrix_root(&joint_penalty)
.map_err(|error| format!("survival edf: joint penalty rank failed: {error}"))?
.nrows();
let bundle = gam_solve::estimate::penalized_edf_bundle(
&raw_traces,
&block_ranks,
p,
(p - joint_penalty_rank.min(p)) as f64,
);
let edf_by_block = bundle.edf_by_block;
let penalty_block_trace = bundle.penalty_block_trace;
let edf_total = bundle.edf_total;
if !edf_total.is_finite()
|| edf_by_block.iter().any(|v| !v.is_finite())
|| penalty_block_trace.iter().any(|v| !v.is_finite())
{
return Err("survival edf: non-finite effective degrees of freedom".to_string());
}
Ok((edf_total, edf_by_block, penalty_block_trace))
}
struct SurvivalSmoothingSelection {
lambdas: Vec<f64>,
outer_iterations: usize,
criterion_certificate: Option<gam_solve::estimate::OuterCriterionCertificate>,
}
fn optimize_survival_transformation_smoothing(
model: &crate::survival::WorkingModelSurvival,
penalty_blocks: &[PenaltyBlock],
num_smoothing: usize,
beta0: &Array1<f64>,
structural_lower_bounds: Option<&Array1<f64>>,
time_block_cols: usize,
left_truncated: bool,
) -> Result<Option<SurvivalSmoothingSelection>, String> {
use gam_problem::{Derivative, HessianValue, OuterEval};
use gam_solve::rho_optimizer::OuterProblem;
if num_smoothing == 0 {
return Ok(None);
}
if num_smoothing > penalty_blocks.len() {
return Err(format!(
"survival transformation smoothing count {num_smoothing} exceeds penalty count {}",
penalty_blocks.len()
));
}
let seed_lambdas: Vec<f64> = penalty_blocks.iter().map(|b| b.lambda).collect();
let seed_log_lambdas = seed_lambdas
.iter()
.copied()
.enumerate()
.map(|(coordinate, value)| {
gam_problem::checked_log_strength(value).map_err(|error| {
format!("survival transformation seed lambda {coordinate}: {error}")
})
})
.collect::<Result<Vec<_>, _>>()?;
let seed_rho = Array1::from_vec(seed_log_lambdas[..num_smoothing].to_vec());
let eval_cache: std::cell::RefCell<Option<(Array1<f64>, f64, Array1<f64>)>> =
std::cell::RefCell::new(None);
let warm_beta: std::cell::RefCell<Array1<f64>> = std::cell::RefCell::new(beta0.clone());
let eval_at = |rho_smooth: &Array1<f64>| -> Result<
(f64, Array1<f64>),
gam_solve::estimate::EstimationError,
> {
let physical_smoothing =
gam_problem::checked_exp_log_strengths(rho_smooth.iter().copied())?;
if let Some((cached_rho, cached_cost, cached_grad)) = eval_cache.borrow().as_ref()
&& cached_rho == rho_smooth
{
return Ok((*cached_cost, cached_grad.clone()));
}
let mut candidate = model.clone();
let mut lambdas = seed_lambdas.clone();
for k in 0..num_smoothing {
lambdas[k] = physical_smoothing[k];
}
candidate
.set_penalty_lambdas(&lambdas)
.map_err(|error| {
gam_solve::estimate::EstimationError::InvalidInput(error.to_string())
})?;
let opts = gam_solve::pirls::WorkingModelPirlsOptions {
max_iterations: SURVIVAL_TRANSFORMATION_PIRLS_MAX_ITERATIONS,
convergence_tolerance: SURVIVAL_TRANSFORMATION_PIRLS_CONVERGENCE_TOL,
adaptive_kkt_tolerance: None,
max_step_halving: SURVIVAL_TRANSFORMATION_PIRLS_MAX_STEP_HALVING,
min_step_size: SURVIVAL_TRANSFORMATION_PIRLS_MIN_STEP_SIZE,
firth_bias_reduction: false,
coefficient_lower_bounds: structural_lower_bounds.cloned(),
linear_constraints: None,
initial_lm_lambda: None,
arrow_schur: None,
};
let summary = gam_solve::pirls::runworking_model_pirls(
&mut candidate,
gam_problem::Coefficients::new(warm_beta.borrow().clone()),
&opts,
|_| {},
)?;
if !survival_pirls_status_is_certified(summary.status) {
return Err(gam_solve::estimate::EstimationError::PirlsDidNotConverge {
max_iterations: opts.max_iterations,
last_change: summary.lastgradient_norm,
});
}
let beta = summary.beta.as_ref().to_owned();
*warm_beta.borrow_mut() = beta.clone();
let state = candidate.update_state(&beta).map_err(|error| {
gam_solve::estimate::EstimationError::InvalidInput(format!(
"survival smoothing inner state evaluation failed: {error}"
))
})?;
let full_rho = Array1::from_vec(
lambdas
.iter()
.copied()
.enumerate()
.map(|(coordinate, value)| {
gam_problem::checked_log_strength(value).map_err(|error| {
gam_solve::estimate::EstimationError::InvalidInput(format!(
"survival smoothing candidate lambda {coordinate}: {error}"
))
})
})
.collect::<Result<Vec<_>, _>>()?,
);
let (cost, grad_full) = candidate
.unified_lamlobjective_and_rhogradient(&beta, &state, &full_rho)
.map_err(|error| {
gam_solve::estimate::EstimationError::InvalidInput(format!(
"survival smoothing LAML evaluation failed: {error}"
))
})?;
if grad_full.len() < num_smoothing || !cost.is_finite() {
return Err(gam_solve::estimate::EstimationError::InvalidInput(
"survival smoothing LAML cost was non-finite or gradient was too short"
.to_string(),
));
}
let grad = grad_full.slice(s![..num_smoothing]).to_owned();
if grad.iter().any(|g| !g.is_finite()) {
return Err(gam_solve::estimate::EstimationError::InvalidInput(
"survival smoothing LAML gradient was non-finite".to_string(),
));
}
*eval_cache.borrow_mut() = Some((rho_smooth.to_owned(), cost, grad.clone()));
Ok((cost, grad))
};
let mut lower = seed_rho.mapv(|v| v - 12.0);
let upper = seed_rho.mapv(|v| v + 12.0);
if left_truncated {
for k in 0..num_smoothing {
let is_time_block = penalty_blocks
.get(k)
.is_some_and(|block| block.range.start < time_block_cols);
if is_time_block {
lower[k] = seed_rho[k];
}
}
}
let context =
format!("survival transformation smoothing-parameter selection (dim={num_smoothing})");
let mut current_seed = seed_rho.clone();
let mut carried_iterations = 0usize;
let mut resumes_remaining = SURVIVAL_TRANSFORMATION_OUTER_STALL_RESTARTS;
let (outer_iterations, criterion_certificate, selected_rho) = loop {
let problem = OuterProblem::new(num_smoothing)
.with_gradient(Derivative::Analytic)
.with_hessian(gam_problem::DeclaredHessianForm::Unavailable)
.with_tolerance(1e-4)
.with_max_iter(120)
.with_bounds(lower.clone(), upper.clone())
.with_initial_rho(current_seed.clone())
.with_seed_config(gam_problem::SeedConfig {
max_seeds: 1,
seed_budget: 1,
..Default::default()
});
let mut obj = problem.build_objective(
(),
|_: &mut (), rho: &Array1<f64>| eval_at(rho).map(|(c, _)| c),
|_: &mut (), rho: &Array1<f64>| {
let (cost, gradient) = eval_at(rho)?;
Ok(OuterEval {
cost,
gradient,
hessian: HessianValue::Unavailable,
inner_beta_hint: None,
})
},
None::<fn(&mut ())>,
None::<
fn(
&mut (),
&Array1<f64>,
)
-> Result<gam_problem::EfsEval, gam_solve::estimate::EstimationError>,
>,
);
match problem.run(&mut obj, &context) {
Ok(result) => {
break (
result.iterations.saturating_add(carried_iterations),
result.criterion_certificate,
result.rho,
);
}
Err(error) => {
if resumes_remaining > 0
&& let gam_solve::estimate::EstimationError::RemlDidNotConverge {
rho_checkpoint,
iterations,
..
} = &error
&& rho_checkpoint.len() == num_smoothing
{
log::info!(
"[OUTER] {context}: resuming from refused checkpoint {rho_checkpoint:?} \
with a fresh BFGS metric ({resumes_remaining} resume(s) left)"
);
carried_iterations = carried_iterations.saturating_add(*iterations);
resumes_remaining -= 1;
current_seed = Array1::from_vec(rho_checkpoint.clone());
continue;
}
return Err(error.to_string());
}
}
};
if selected_rho.len() != num_smoothing {
return Err(format!(
"survival transformation smoothing selector returned {} coordinates for \
{num_smoothing} smoothing parameters; selected-rho checkpoint={:?}",
selected_rho.len(),
selected_rho.to_vec(),
));
}
let selected_lambdas = gam_problem::checked_exp_log_strengths(selected_rho.iter().copied())
.map_err(|error| format!("survival transformation selected rho: {error}"))?;
let mut lambdas = seed_lambdas;
for (slot, lambda) in lambdas.iter_mut().zip(selected_lambdas) {
*slot = lambda;
}
Ok(Some(SurvivalSmoothingSelection {
lambdas,
outer_iterations,
criterion_certificate,
}))
}
fn survival_conditional_covariance_from_penalized_hessian(
penalized_hessian: &Array2<f64>,
) -> Option<Array2<f64>> {
use gam_linalg::faer_ndarray::FaerCholesky;
let p = penalized_hessian.nrows();
let identity = Array2::<f64>::eye(p);
let cov = match penalized_hessian.cholesky(faer::Side::Lower) {
Ok(chol) => chol.solve_mat(&identity),
Err(_) => return None,
};
if !cov.iter().all(|v| v.is_finite()) {
return None;
}
let mut symm = cov.clone();
for i in 0..p {
for j in (i + 1)..p {
let avg = 0.5 * (cov[[i, j]] + cov[[j, i]]);
symm[[i, j]] = avg;
symm[[j, i]] = avg;
}
}
Some(symm)
}
fn survival_unified_fit_result(
beta: Array1<f64>,
lambdas: Array1<f64>,
summary: &gam_solve::pirls::WorkingModelPirlsResult,
state: &gam_solve::pirls::WorkingState,
penalty_blocks: &[PenaltyBlock],
outer_iterations: usize,
criterion_certificate: Option<gam_solve::estimate::OuterCriterionCertificate>,
) -> Result<UnifiedFitResult, String> {
let log_lambdas = Array1::from_vec(
lambdas
.iter()
.copied()
.enumerate()
.map(|(coordinate, value)| {
gam_problem::checked_log_strength(value).map_err(|error| {
format!("survival fit lambda coordinate {coordinate}: {error}")
})
})
.collect::<Result<Vec<_>, _>>()?,
);
let lambdas = Array1::from_vec(
log_lambdas
.iter()
.copied()
.enumerate()
.map(|(coordinate, log_value)| {
gam_problem::checked_exp_log_strength(log_value).map_err(|error| {
format!("survival fit log-lambda coordinate {coordinate}: {error}")
})
})
.collect::<Result<Vec<_>, _>>()?,
);
require_certified_survival_pirls(
summary,
"survival transformation fit assembly",
log_lambdas.as_slice().unwrap_or(&[]),
None,
)?;
let reml_score = state.penalized_objective();
gam_solve::estimate::validate_all_finite("survival fit beta", beta.iter().copied())?;
gam_solve::estimate::validate_all_finite("survival fit lambdas", lambdas.iter().copied())?;
gam_solve::estimate::ensure_finite_scalar("survival fit log_likelihood", state.log_likelihood)?;
gam_solve::estimate::ensure_finite_scalar("survival fit deviance", state.deviance)?;
gam_solve::estimate::ensure_finite_scalar("survival fit penalty", state.penalty_term)?;
gam_solve::estimate::ensure_finite_scalar("survival fit reml_score", reml_score)?;
gam_solve::estimate::ensure_finite_scalar(
"survival fit gradient_norm",
summary.lastgradient_norm,
)?;
gam_solve::estimate::ensure_finite_scalar("survival fit max_abs_eta", summary.max_abs_eta)?;
let (edf_total, edf_by_block, penalty_block_trace, penalized_hessian) =
survival_transformation_edf(state, penalty_blocks)?;
assert_eq!(edf_by_block.len(), lambdas.len());
assert_eq!(penalty_block_trace.len(), lambdas.len());
let covariance_conditional =
survival_conditional_covariance_from_penalized_hessian(&penalized_hessian);
let beta_standard_errors = covariance_conditional
.as_ref()
.map(gam_problem::se_from_covariance)
.transpose()
.map_err(|reason| {
format!("survival transformation conditional standard errors are invalid: {reason}")
})?;
let beta_covariance = covariance_conditional
.clone()
.map(gam_problem::dispersion_cov::PhiScaledCovariance::wrap);
let penalized_hessian = gam_problem::dispersion_cov::UnscaledPrecision::wrap(penalized_hessian);
let inference = gam_solve::estimate::FitInference {
edf_by_block: edf_by_block.clone(),
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.clone(),
reparam_qs: None,
dispersion: gam_solve::estimate::Dispersion::UNIT,
beta_covariance,
beta_standard_errors,
beta_covariance_corrected: None,
beta_standard_errors_corrected: None,
beta_covariance_frequentist: None,
coefficient_influence: None,
weighted_gram: None,
bias_correction_beta: None,
bias_correction_jacobian: None,
};
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(LikelihoodSpec::royston_parmar()),
likelihood_scale: gam_problem::LikelihoodScaleMetadata::Unspecified,
log_likelihood_normalization: gam_problem::LogLikelihoodNormalization::UserProvided,
log_likelihood: state.log_likelihood,
deviance: state.deviance,
reml_score,
stable_penalty_term: state.penalty_term,
penalized_objective: reml_score,
used_device: false,
outer_iterations,
outer_converged: true,
outer_gradient_norm: criterion_certificate
.as_ref()
.map(|certificate| certificate.stationarity.projected_norm())
.or(Some(summary.lastgradient_norm)),
standard_deviation: 1.0,
covariance_conditional,
covariance_corrected: None,
inference: Some(inference),
fitted_link: FittedLinkState::Standard(None),
geometry: Some(gam_solve::estimate::FitGeometry {
coefficient_gauge: gam_problem::gauge::Gauge::identity(&[beta.len()]),
penalized_hessian,
constrained_posterior: None,
working: None,
}),
block_states: Vec::new(),
pirls_status: summary.status,
max_abs_eta: summary.max_abs_eta,
constraint_kkt: None,
artifacts: gam_solve::estimate::FitArtifacts {
pirls: None,
criterion_certificate,
..Default::default()
},
inner_cycles: 0,
})
.map_err(|err| err.to_string())
}
pub(crate) fn replicate_pooled_baseline_seed_per_cause(
pooled_seed: ArrayView1<'_, f64>,
cause_count: usize,
) -> Array1<f64> {
let p = pooled_seed.len();
let mut beta0_flat = Array1::<f64>::zeros(p * cause_count);
for cause in 0..cause_count {
beta0_flat
.slice_mut(s![cause * p..(cause + 1) * p])
.assign(&pooled_seed);
}
beta0_flat
}
fn fit_cause_specific_survival_transformation_custom(
spec: &SurvivalTransformationTermSpec,
resolvedspec: TermCollectionSpec,
baseline_cfg: crate::survival::construction::SurvivalBaselineConfig,
prepared: PreparedSurvivalTimeStack,
dense_cov_design: &Array2<f64>,
penalty_blocks: Vec<PenaltyBlock>,
beta0_flat: Array1<f64>,
derivative_floor: f64,
penalty_block_gamma_priors: &[(String, f64, f64)],
) -> Result<SurvivalTransformationFitResult, String> {
let cause_count = crate::survival::cause_count_from_event_codes(spec.event_target.view())
.into_workflow_result()?;
if cause_count == 0 {
return Err(WorkflowError::MissingDependency {
reason: "cause-specific custom survival fit requires at least one cause".to_string(),
}
.into());
}
let n = spec.event_target.len();
let p_time_total = prepared.time_design_exit.ncols();
let p_cov = dense_cov_design.ncols();
let p = p_time_total + p_cov;
if beta0_flat.len() != p * cause_count {
return Err(WorkflowError::SchemaMismatch {
reason: format!(
"cause-specific survival initial beta length mismatch: got {}, expected {}",
beta0_flat.len(),
p * cause_count
),
}
.into());
}
let dense_time_entry = prepared.time_design_entry.to_dense();
let dense_time_exit = prepared.time_design_exit.to_dense();
let dense_time_derivative = prepared.time_design_derivative_exit.to_dense();
let mut x_entry = Array2::<f64>::zeros((n, p));
let mut x_exit = Array2::<f64>::zeros((n, p));
let mut x_derivative = Array2::<f64>::zeros((n, p));
if p_time_total > 0 {
x_entry
.slice_mut(s![.., ..p_time_total])
.assign(&dense_time_entry);
x_exit
.slice_mut(s![.., ..p_time_total])
.assign(&dense_time_exit);
x_derivative
.slice_mut(s![.., ..p_time_total])
.assign(&dense_time_derivative);
}
if p_cov > 0 {
x_entry
.slice_mut(s![.., p_time_total..])
.assign(dense_cov_design);
x_exit
.slice_mut(s![.., p_time_total..])
.assign(dense_cov_design);
}
let mut family_blocks = Vec::with_capacity(cause_count);
let mut block_specs = Vec::with_capacity(cause_count);
for cause in 0..cause_count {
let cause_code = (cause + 1) as u8;
let event_target = spec
.event_target
.mapv(|observed| u8::from(observed == cause_code));
family_blocks.push(crate::survival::CauseSpecificRoystonParmarBlock {
age_entry: spec.age_entry.clone(),
age_exit: spec.age_exit.clone(),
event_target,
sampleweight: spec.weights.clone(),
x_entry: x_entry.clone(),
x_exit: x_exit.clone(),
x_derivative: x_derivative.clone(),
offset_eta_entry: prepared.eta_offset_entry.clone() + &spec.covariate_offset,
offset_eta_exit: prepared.eta_offset_exit.clone() + &spec.covariate_offset,
offset_derivative_exit: prepared.derivative_offset_exit.clone(),
derivative_floor,
structural_time_columns: if spec.likelihood_mode == SurvivalLikelihoodMode::Weibull {
0
} else {
p_time_total
},
});
let mut penalties = Vec::with_capacity(penalty_blocks.len());
let mut nullspace_dims = Vec::with_capacity(penalty_blocks.len());
let mut initial_log_lambdas = Array1::<f64>::zeros(penalty_blocks.len());
for (penalty_idx, block) in penalty_blocks.iter().enumerate() {
if block.range.end > p || block.range.start > block.range.end {
return Err(WorkflowError::SchemaMismatch {
reason: "cause-specific survival penalty range is out of bounds".to_string(),
}
.into());
}
let block_dim = block.range.end - block.range.start;
if block.matrix.nrows() != block_dim || block.matrix.ncols() != block_dim {
return Err(WorkflowError::SchemaMismatch {
reason: format!(
"cause-specific survival penalty {penalty_idx} has shape {}x{} but range has width {block_dim}",
block.matrix.nrows(),
block.matrix.ncols()
),
}
.into());
}
penalties.push(
PenaltyMatrix::Blockwise {
local: block.matrix.clone(),
col_range: block.range.clone(),
total_dim: p,
}
.with_precision_label(format!(
"cause_specific_survival_cause_{}_penalty_{penalty_idx}",
cause + 1
)),
);
nullspace_dims.push(block.nullspace_dim);
initial_log_lambdas[penalty_idx] = gam_problem::checked_log_strength(block.lambda)
.map_err(|error| {
format!("cause-specific survival penalty {penalty_idx} strength: {error}")
})?;
}
let beta_start = beta0_flat.slice(s![cause * p..(cause + 1) * p]).to_owned();
let cause_priority =
100u8.saturating_add(u8::try_from(cause_count - cause).unwrap_or(u8::MAX));
let cause_jacobian = std::sync::Arc::new(AdditiveBlockJacobian {
design: x_exit.clone(),
own_output: cause,
n_family_outputs: cause_count,
});
block_specs.push(ParameterBlockSpec {
name: format!("time_cause_{}", cause + 1),
design: gam_linalg::matrix::DesignMatrix::from(x_exit.clone()),
offset: prepared.eta_offset_exit.clone() + &spec.covariate_offset,
penalties,
nullspace_dims,
initial_log_lambdas,
initial_beta: Some(beta_start),
gauge_priority: cause_priority,
jacobian_callback: Some(cause_jacobian),
stacked_design: None,
stacked_offset: None,
});
}
let family = crate::survival::CauseSpecificRoystonParmarFamily::new(family_blocks)?;
let fit_options = BlockwiseFitOptions {
compute_covariance: true,
..Default::default()
};
let rho_prior = cause_specific_survival_rho_prior(
cause_count,
penalty_blocks.len(),
penalty_block_gamma_priors,
)?;
let mut fit = fit_custom_family_with_rho_prior(&family, &block_specs, &fit_options, rho_prior)
.map_err(|err| format!("cause-specific survival custom-family fit failed: {err}"))?;
fit.likelihood_family = Some(LikelihoodSpec::royston_parmar());
let time_basis = crate::survival::construction::SavedSurvivalTimeBasis::from_build(
&spec.time_build,
spec.time_anchor,
);
let fitted_baseline_cfg = if spec.likelihood_mode == SurvivalLikelihoodMode::Weibull
&& spec.timewiggle.is_none()
{
let first_block = fit.blocks.first().ok_or_else(|| {
"cause-specific survival fit produced no coefficient blocks".to_string()
})?;
let time_beta = first_block
.beta
.slice(s![..spec.time_build.x_exit_time.ncols()])
.to_owned();
fitted_weibull_baseline_from_linear_time_beta(&time_beta, spec.time_anchor).ok_or_else(|| {
"failed to recover fitted Weibull scale/shape from the cause-specific linear time coefficients"
.to_string()
})?
} else {
baseline_cfg
};
Ok(SurvivalTransformationFitResult {
fit,
resolvedspec,
baseline_cfg: fitted_baseline_cfg,
likelihood_mode: spec.likelihood_mode,
time_basis,
time_base_ncols: spec.time_build.x_exit_time.ncols(),
baseline_timewiggle: prepared.timewiggle_block,
})
}
fn cause_specific_survival_rho_prior(
cause_count: usize,
penalty_count: usize,
penalty_block_gamma_priors: &[(String, f64, f64)],
) -> Result<gam_problem::RhoPrior, String> {
if penalty_block_gamma_priors.is_empty() {
return Ok(gam_problem::RhoPrior::Flat);
}
let mut keyed = BTreeMap::<String, (f64, f64)>::new();
for (label, shape, rate) in penalty_block_gamma_priors {
if keyed.insert(label.clone(), (*shape, *rate)).is_some() {
return Err(WorkflowError::InvalidConfig {
reason: format!(
"duplicate Gamma precision hyperprior for penalty block label '{label}'"
),
}
.into());
}
if !shape.is_finite() || *shape <= 0.0 {
return Err(WorkflowError::InvalidConfig {
reason: format!(
"Gamma precision hyperprior for penalty block '{label}' requires shape > 0, got {shape}"
),
}
.into());
}
if !rate.is_finite() || *rate < 0.0 {
return Err(WorkflowError::InvalidConfig {
reason: format!(
"Gamma precision hyperprior for penalty block '{label}' requires rate >= 0, got {rate}"
),
}
.into());
}
}
let mut consumed = Vec::<String>::new();
let mut priors = Vec::<gam_problem::RhoPrior>::with_capacity(cause_count * penalty_count);
for cause in 0..cause_count {
for penalty_idx in 0..penalty_count {
let label = format!(
"cause_specific_survival_cause_{}_penalty_{penalty_idx}",
cause + 1
);
if let Some((shape, rate)) = keyed.get(&label) {
consumed.push(label);
priors.push(gam_problem::RhoPrior::GammaPrecision {
shape: *shape,
rate: *rate,
});
} else {
priors.push(gam_problem::RhoPrior::Flat);
}
}
}
let unknown = keyed
.keys()
.filter(|label| !consumed.iter().any(|known| known == *label))
.cloned()
.collect::<Vec<_>>();
if !unknown.is_empty() {
let available = (0..cause_count)
.flat_map(|cause| {
(0..penalty_count).map(move |idx| {
format!("cause_specific_survival_cause_{}_penalty_{idx}", cause + 1)
})
})
.collect::<Vec<_>>()
.join(", ");
return Err(WorkflowError::InvalidConfig {
reason: format!(
"unknown Gamma precision hyperprior penalty block label(s): {}; available labels: {available}",
unknown.join(", ")
),
}
.into());
}
Ok(gam_problem::RhoPrior::Independent(priors))
}
fn hash_workflow_array_view(
hasher: &mut gam_runtime::warm_start::Fingerprinter,
array: ArrayView1<'_, f64>,
) {
hasher.write_usize(array.len());
for &value in array {
hasher.write_f64(value);
}
}
fn hash_workflow_u8_array(
hasher: &mut gam_runtime::warm_start::Fingerprinter,
array: ArrayView1<'_, u8>,
) {
hasher.write_usize(array.len());
for &value in array {
hasher.write_usize(usize::from(value));
}
}
fn hash_workflow_array2(
hasher: &mut gam_runtime::warm_start::Fingerprinter,
array: ArrayView2<'_, f64>,
) {
hasher.write_usize(array.nrows());
hasher.write_usize(array.ncols());
for row in array.rows() {
for &value in row {
hasher.write_f64(value);
}
}
}
fn hash_workflow_design_matrix(
hasher: &mut gam_runtime::warm_start::Fingerprinter,
matrix: &gam_linalg::matrix::DesignMatrix,
) {
let dense = matrix.to_dense();
hash_workflow_array2(hasher, dense.view());
}
fn survival_transformation_log_lambdas(
penalty_blocks: &[crate::survival::PenaltyBlock],
) -> Result<Vec<f64>, String> {
penalty_blocks
.iter()
.enumerate()
.map(|(coordinate, block)| {
gam_problem::checked_log_strength(block.lambda)
.map_err(|error| format!("survival transformation penalty {coordinate}: {error}"))
})
.collect()
}
fn persistent_survival_transformation_key(
spec: &SurvivalTransformationTermSpec,
baseline_cfg: &crate::survival::construction::SurvivalBaselineConfig,
dense_cov_design: ArrayView2<'_, f64>,
prepared: &PreparedSurvivalTimeStack,
penalty_blocks: &[crate::survival::PenaltyBlock],
opts: &gam_solve::pirls::WorkingModelPirlsOptions,
n_cols: usize,
) -> String {
let mut hasher = gam_runtime::warm_start::Fingerprinter::new();
hasher.write_str("gamfit-persistent-survival-transformation-working-pirls");
hasher.write_str(&gam_solve::persistent_warm_start::cache_schema_tag());
hasher.write_str(&format!("{:?}", spec.likelihood_mode));
hasher.write_f64(spec.time_anchor);
hasher.write_f64(spec.ridge_lambda);
hasher.write_str(&format!("{:?}", baseline_cfg.target));
for value in [
baseline_cfg.scale,
baseline_cfg.shape,
baseline_cfg.rate,
baseline_cfg.makeham,
] {
hasher.write_bool(value.is_some());
if let Some(value) = value {
hasher.write_f64(value);
}
}
hasher.write_str(&spec.time_build.basisname);
hasher.write_usize(spec.time_build.x_entry_time.nrows());
hasher.write_usize(spec.time_build.x_entry_time.ncols());
hasher.write_usize(spec.time_build.x_exit_time.nrows());
hasher.write_usize(spec.time_build.x_exit_time.ncols());
hasher.write_usize(spec.time_build.x_derivative_time.nrows());
hasher.write_usize(spec.time_build.x_derivative_time.ncols());
hasher.write_bool(spec.time_build.degree.is_some());
if let Some(degree) = spec.time_build.degree {
hasher.write_usize(degree);
}
match spec.time_build.knots.as_ref() {
Some(knots) => {
hasher.write_bool(true);
hasher.write_usize(knots.len());
for &knot in knots {
hasher.write_f64(knot);
}
}
None => hasher.write_bool(false),
}
match spec.time_build.keep_cols.as_ref() {
Some(cols) => {
hasher.write_bool(true);
hasher.write_usize(cols.len());
for &col in cols {
hasher.write_usize(col);
}
}
None => hasher.write_bool(false),
}
hasher.write_bool(spec.time_build.smooth_lambda.is_some());
if let Some(lambda) = spec.time_build.smooth_lambda {
hasher.write_f64(lambda);
}
hasher.write_usize(n_cols);
hash_workflow_array_view(&mut hasher, spec.age_entry.view());
hash_workflow_array_view(&mut hasher, spec.age_exit.view());
hash_workflow_u8_array(&mut hasher, spec.event_target.view());
hash_workflow_array_view(&mut hasher, spec.weights.view());
hash_workflow_array_view(&mut hasher, spec.covariate_offset.view());
hash_workflow_array2(&mut hasher, dense_cov_design);
hash_workflow_array_view(&mut hasher, prepared.eta_offset_entry.view());
hash_workflow_array_view(&mut hasher, prepared.eta_offset_exit.view());
hash_workflow_array_view(&mut hasher, prepared.derivative_offset_exit.view());
hash_workflow_design_matrix(&mut hasher, &prepared.time_design_entry);
hash_workflow_design_matrix(&mut hasher, &prepared.time_design_exit);
hash_workflow_design_matrix(&mut hasher, &prepared.time_design_derivative_exit);
hasher.write_usize(penalty_blocks.len());
for block in penalty_blocks {
hasher.write_f64(block.lambda);
hasher.write_usize(block.range.start);
hasher.write_usize(block.range.end);
hasher.write_usize(block.nullspace_dim);
hash_workflow_array2(&mut hasher, block.matrix.view());
}
hasher.write_usize(opts.max_iterations);
hasher.write_f64(opts.convergence_tolerance);
hasher.write_usize(opts.max_step_halving);
hasher.write_f64(opts.min_step_size);
hasher.write_bool(opts.firth_bias_reduction);
hasher.write_bool(opts.coefficient_lower_bounds.is_some());
if let Some(bounds) = opts.coefficient_lower_bounds.as_ref() {
hash_workflow_array_view(&mut hasher, bounds.view());
}
hasher.write_bool(opts.linear_constraints.is_some());
format!("surv-transform-{}", hasher.finish_hex())
}
fn load_survival_transformation_persistent_warm_start(
key: &str,
spec: &SurvivalTransformationTermSpec,
n_cols: usize,
rho: &[f64],
) -> Option<(Array1<f64>, Option<f64>)> {
let record = gam_solve::persistent_warm_start::load_record(key)?;
if !record.is_compatible(key, spec.age_entry.len(), n_cols)
|| record.rho.len() != rho.len()
|| !record
.rho
.iter()
.zip(rho.iter())
.all(|(cached, expected)| (*cached - *expected).abs() <= 1e-10)
{
return None;
}
log::info!("[warm-start-cache] restored survival transformation warm start key={key}");
let lm_lambda = record
.last_pirls_lm_lambda
.filter(|value| value.is_finite() && *value > 0.0);
Some((Array1::from_vec(record.beta), lm_lambda))
}
fn store_survival_transformation_persistent_warm_start(
key: &str,
spec: &SurvivalTransformationTermSpec,
n_cols: usize,
rho: Vec<f64>,
beta: &Array1<f64>,
summary: &gam_solve::pirls::WorkingModelPirlsResult,
) -> bool {
if beta.len() != n_cols
|| beta.iter().any(|value| !value.is_finite())
|| rho.iter().any(|value| !value.is_finite())
{
return false;
}
let mut record = gam_solve::persistent_warm_start::PersistentWarmStartRecord::new(
key.to_string(),
spec.age_entry.len(),
n_cols,
);
record.rho = rho;
record.beta = beta.to_vec();
record.last_inner_iters = summary.iterations;
record.last_inner_converged = summary.status.is_converged();
record.last_pirls_lm_lambda = (summary.final_lm_lambda.is_finite()
&& summary.final_lm_lambda > 0.0)
.then_some(summary.final_lm_lambda);
record.last_pirls_accept_rho = summary
.final_accept_rho
.filter(|value| value.is_finite() && *value >= 0.0);
match gam_solve::persistent_warm_start::store_record(&record) {
Ok(()) => {
gam_solve::persistent_warm_start::load_record(&record.key).is_some_and(|stored| {
stored.rho == record.rho
&& stored.beta == record.beta
&& stored.last_inner_iters == record.last_inner_iters
&& stored.last_inner_converged == record.last_inner_converged
})
}
Err(err) => {
log::warn!(
"[warm-start-cache] failed to persist survival transformation warm start: {err}"
);
false
}
}
}
pub(crate) fn fit_survival_transformation_model(
request: SurvivalTransformationFitRequest<'_>,
) -> Result<SurvivalTransformationFitResult, String> {
use crate::survival::{PenaltyBlock, PenaltyBlocks, SurvivalMonotonicityPenalty, SurvivalSpec};
let SurvivalTransformationFitRequest {
data,
spec,
cache_session: _cache_session,
} = request;
let mut baseline_cfg = spec.baseline_cfg.clone();
let covariate_design =
build_term_collection_design(data, &spec.covariate_spec).map_err(|err| err.to_string())?;
let resolvedspec = crate::fit_orchestration::drivers::freeze_term_collection_from_design(
&spec.covariate_spec,
&covariate_design,
)
.map_err(|err| err.to_string())?;
let dense_cov_design = covariate_design.design.to_dense();
let p_cov = dense_cov_design.ncols();
let cause_count = crate::survival::cause_count_from_event_codes(spec.event_target.view())
.into_workflow_result()?;
let exact_derivative_guard = survival_derivative_guard_for_likelihood(spec.likelihood_mode);
let build_working_model =
|candidate: &crate::survival::construction::SurvivalBaselineConfig| {
let prepared = prepare_survival_time_stack(
&spec.age_entry,
&spec.age_exit,
candidate,
spec.likelihood_mode,
None,
spec.time_anchor,
exact_derivative_guard,
&spec.time_build,
spec.timewiggle.as_ref(),
None,
)?;
let mut eta_offset_entry = prepared.eta_offset_entry.clone();
let mut eta_offset_exit = prepared.eta_offset_exit.clone();
eta_offset_entry += &spec.covariate_offset;
eta_offset_exit += &spec.covariate_offset;
eta_offset_entry += &covariate_design.affine_offset;
eta_offset_exit += &covariate_design.affine_offset;
let p_time_total = prepared.time_design_exit.ncols();
let p = p_time_total + p_cov;
let mut penalty_blocks = Vec::<PenaltyBlock>::new();
for (idx, penalty) in prepared.time_penalties.iter().enumerate() {
if penalty.nrows() == p_time_total && penalty.ncols() == p_time_total {
penalty_blocks.push(PenaltyBlock {
matrix: penalty.clone(),
lambda: spec.time_build.smooth_lambda.unwrap_or(1e-2),
range: 0..p_time_total,
nullspace_dim: prepared.time_nullspace_dims.get(idx).copied().unwrap_or(0),
});
}
}
for (penalty_idx, cov_penalty) in covariate_design.penalties.iter().enumerate() {
let cr = &cov_penalty.col_range;
let block_dim = cr.end - cr.start;
let matches_dims = cov_penalty.local.nrows() == block_dim
&& cov_penalty.local.ncols() == block_dim;
let zero_prior = matches!(
cov_penalty.prior_mean,
gam_problem::CoefficientPriorMean::Zero
);
if block_dim > 0 && matches_dims && zero_prior && cr.end <= p_cov {
penalty_blocks.push(PenaltyBlock {
matrix: cov_penalty.local.clone(),
lambda: 1e-2,
range: (p_time_total + cr.start)..(p_time_total + cr.end),
nullspace_dim: covariate_design
.nullspace_dims
.get(penalty_idx)
.copied()
.unwrap_or(0),
});
}
}
let num_smoothing_blocks = penalty_blocks.len();
let ridge_range_start = 0usize;
if spec.ridge_lambda > 0.0 && p > ridge_range_start {
let dim = p - ridge_range_start;
let mut ridge = Array2::<f64>::zeros((dim, dim));
for d in 0..dim {
ridge[[d, d]] = 1.0;
}
penalty_blocks.push(PenaltyBlock {
matrix: ridge,
lambda: spec.ridge_lambda,
range: ridge_range_start..p,
nullspace_dim: 0,
});
}
let dense_time_entry = prepared.time_design_entry.to_dense();
let dense_time_exit = prepared.time_design_exit.to_dense();
let dense_time_derivative = prepared.time_design_derivative_exit.to_dense();
let event_competing = Array1::<u8>::zeros(spec.event_target.len());
let baseline_event_indicator = spec.event_target.mapv(|label| u8::from(label > 0));
let mut model =
crate::survival::royston_parmar::working_model_from_time_covariateshared(
PenaltyBlocks::new(penalty_blocks.clone()),
SurvivalMonotonicityPenalty { tolerance: 0.0 },
SurvivalSpec::Net,
crate::survival::royston_parmar::RoystonParmarSharedTimeCovariateInputs {
age_entry: spec.age_entry.view(),
age_exit: spec.age_exit.view(),
event_target: baseline_event_indicator.view(),
event_competing: event_competing.view(),
weights: spec.weights.view(),
time_entry: dense_time_entry.view(),
time_exit: dense_time_exit.view(),
time_derivative: dense_time_derivative.view(),
covariates: dense_cov_design.view(),
monotonicity_constraint_rows: None,
monotonicity_constraint_offsets: None,
eta_offset_entry: Some(eta_offset_entry.view()),
eta_offset_exit: Some(eta_offset_exit.view()),
derivative_offset_exit: Some(prepared.derivative_offset_exit.view()),
},
)
.map_err(|err| format!("failed to construct survival model: {err}"))?;
if spec.likelihood_mode != SurvivalLikelihoodMode::Weibull {
model
.set_structural_monotonicity(true, p_time_total)
.map_err(|err| format!("failed to enable structural monotonicity: {err}"))?;
}
let mut beta0 = Array1::<f64>::zeros(p);
if spec.likelihood_mode == SurvivalLikelihoodMode::Weibull && spec.timewiggle.is_none()
{
let (scale, shape) = spec
.weibull_seed
.ok_or_else(|| "weibull survival fit missing scale/shape seed".to_string())?;
if p_time_total < 1 {
return Err(format!(
"weibull built-in time basis has {p_time_total} columns but needs 1 for the shape"
));
}
if covariate_design.intercept_range.is_empty() {
return Err(
"weibull survival fit requires a mean intercept to carry the baseline \
location, but the covariate design has none (intercept suppression such \
as `~ x - 1` is unsupported; see formula_dsl.rs:2456)"
.to_string(),
);
}
beta0[0] = shape;
let intercept_col = p_time_total + covariate_design.intercept_range.start;
beta0[intercept_col] = -shape * scale.ln();
}
let structural_lower_bounds =
if spec.likelihood_mode != SurvivalLikelihoodMode::Weibull && p_time_total > 0 {
let mut lb = Array1::from_elem(p, f64::NEG_INFINITY);
for j in 0..p_time_total {
lb[j] = 0.0;
beta0[j] = 1e-4;
}
Some(lb)
} else {
None
};
Ok::<_, String>((
prepared,
penalty_blocks,
beta0,
structural_lower_bounds,
model,
num_smoothing_blocks,
))
};
if baseline_cfg.target != SurvivalBaselineTarget::Linear {
baseline_cfg = optimize_survival_baseline_config_with_gradient_only(
&baseline_cfg,
"workflow survival transformation baseline",
|candidate| {
let (_, _, beta0, structural_lower_bounds, mut model, _) =
build_working_model(candidate)?;
let opts = gam_solve::pirls::WorkingModelPirlsOptions {
max_iterations: SURVIVAL_TRANSFORMATION_PIRLS_MAX_ITERATIONS,
convergence_tolerance: SURVIVAL_TRANSFORMATION_PIRLS_CONVERGENCE_TOL,
adaptive_kkt_tolerance: None,
max_step_halving: SURVIVAL_TRANSFORMATION_PIRLS_MAX_STEP_HALVING,
min_step_size: SURVIVAL_TRANSFORMATION_PIRLS_MIN_STEP_SIZE,
firth_bias_reduction: false,
coefficient_lower_bounds: structural_lower_bounds,
linear_constraints: None,
initial_lm_lambda: None,
arrow_schur: None,
};
let parameter_checkpoint = survival_baseline_parameter_checkpoint(candidate)?;
let summary = gam_solve::pirls::runworking_model_pirls(
&mut model,
gam_problem::Coefficients::new(beta0),
&opts,
|_| {},
)
.map_err(|error| {
format!(
"survival baseline PIRLS failed at parameter_checkpoint=\
{parameter_checkpoint:?}: {error}; no fit was minted"
)
})?;
require_certified_survival_pirls(
&summary,
"survival transformation baseline profile",
¶meter_checkpoint,
None,
)?;
let beta = summary.beta.as_ref().to_owned();
let state = model.update_state(&beta).map_err(|err| {
format!("failed to evaluate survival baseline candidate: {err}")
})?;
let cost = state.penalized_objective();
let residuals = model.offset_channel_residuals(&beta).map_err(|err| {
format!("failed to form survival baseline offset residuals: {err}")
})?;
let gradient = baseline_chain_rule_gradient(
spec.age_entry.view(),
spec.age_exit.view(),
spec.age_exit.view(),
candidate,
&residuals,
)?
.ok_or_else(|| {
"workflow survival transformation baseline unexpectedly has no theta gradient"
.to_string()
})?;
Ok((cost, gradient))
},
)?;
}
let (
prepared,
mut penalty_blocks,
beta0,
structural_lower_bounds,
mut model,
num_smoothing_blocks,
) = build_working_model(&baseline_cfg)?;
if cause_count > 1 || !spec.penalty_block_gamma_priors.is_empty() {
let beta0_flat = replicate_pooled_baseline_seed_per_cause(beta0.view(), cause_count);
return fit_cause_specific_survival_transformation_custom(
&spec,
resolvedspec,
baseline_cfg,
prepared,
&dense_cov_design,
penalty_blocks,
beta0_flat,
exact_derivative_guard,
&spec.penalty_block_gamma_priors,
);
}
let is_left_truncated = spec
.age_entry
.iter()
.any(|&t| t > crate::survival::ENTRY_AT_ORIGIN_THRESHOLD);
let p_time_total = prepared.time_design_exit.ncols();
let (survival_outer_iterations, survival_outer_certificate) =
if let Some(selection) = optimize_survival_transformation_smoothing(
&model,
&penalty_blocks,
num_smoothing_blocks,
&beta0,
structural_lower_bounds.as_ref(),
p_time_total,
is_left_truncated,
)? {
model
.set_penalty_lambdas(&selection.lambdas)
.map_err(|e| e.to_string())?;
for (block, &lam) in penalty_blocks.iter_mut().zip(selection.lambdas.iter()) {
block.lambda = lam;
}
(selection.outer_iterations, selection.criterion_certificate)
} else {
(0, None)
};
let opts = gam_solve::pirls::WorkingModelPirlsOptions {
max_iterations: SURVIVAL_TRANSFORMATION_PIRLS_MAX_ITERATIONS,
convergence_tolerance: SURVIVAL_TRANSFORMATION_PIRLS_CONVERGENCE_TOL,
adaptive_kkt_tolerance: None,
max_step_halving: SURVIVAL_TRANSFORMATION_PIRLS_MAX_STEP_HALVING,
min_step_size: SURVIVAL_TRANSFORMATION_PIRLS_MIN_STEP_SIZE,
firth_bias_reduction: false,
coefficient_lower_bounds: structural_lower_bounds,
linear_constraints: None,
initial_lm_lambda: None,
arrow_schur: None,
};
let rho_for_cache = survival_transformation_log_lambdas(&penalty_blocks)?;
let expected_beta_len = beta0.len();
let persistent_warm_start_key = persistent_survival_transformation_key(
&spec,
&baseline_cfg,
dense_cov_design.view(),
&prepared,
&penalty_blocks,
&opts,
expected_beta_len,
);
let mut opts = opts;
let beta_start = match load_survival_transformation_persistent_warm_start(
&persistent_warm_start_key,
&spec,
expected_beta_len,
&rho_for_cache,
) {
Some((beta, lm_lambda)) => {
opts.initial_lm_lambda = lm_lambda;
beta
}
None => beta0,
};
let summary = gam_solve::pirls::runworking_model_pirls(
&mut model,
gam_problem::Coefficients::new(beta_start),
&opts,
|_| {},
)
.map_err(|error| {
format!(
"survival transformation final fixed-lambda PIRLS failed at \
parameter_checkpoint={rho_for_cache:?} (warm_start_key=\
{persistent_warm_start_key}): {error}; no fit was minted"
)
})?;
let beta = summary.beta.as_ref().to_owned();
let checkpoint_persisted = store_survival_transformation_persistent_warm_start(
&persistent_warm_start_key,
&spec,
expected_beta_len,
rho_for_cache.clone(),
&beta,
&summary,
);
require_certified_survival_pirls(
&summary,
"survival transformation final fixed-lambda PIRLS",
&rho_for_cache,
checkpoint_persisted.then_some(persistent_warm_start_key.as_str()),
)?;
let state = model
.update_state(&beta)
.map_err(|err| format!("failed to evaluate survival optimum: {err}"))?;
let lambdas = Array1::from_iter(penalty_blocks.iter().map(|block| block.lambda));
let fitted_baseline_cfg =
if spec.likelihood_mode == SurvivalLikelihoodMode::Weibull && spec.timewiggle.is_none() {
let time_beta = beta
.slice(s![..spec.time_build.x_exit_time.ncols()])
.to_owned();
fitted_weibull_baseline_from_linear_time_beta(&time_beta, spec.time_anchor).ok_or_else(
|| {
"failed to recover fitted Weibull scale/shape from the linear time coefficients"
.to_string()
},
)?
} else {
baseline_cfg
};
let fit = survival_unified_fit_result(
beta,
lambdas,
&summary,
&state,
&penalty_blocks,
survival_outer_iterations,
survival_outer_certificate,
)?;
let time_base_ncols = spec.time_build.x_exit_time.ncols();
let time_basis = crate::survival::construction::SavedSurvivalTimeBasis::from_build(
&spec.time_build,
spec.time_anchor,
);
Ok(SurvivalTransformationFitResult {
fit,
resolvedspec,
baseline_cfg: fitted_baseline_cfg,
likelihood_mode: spec.likelihood_mode,
time_basis,
time_base_ncols,
baseline_timewiggle: prepared.timewiggle_block,
})
}
pub(crate) fn fit_survival_location_scale_model(
request: SurvivalLocationScaleFitRequest<'_>,
) -> Result<SurvivalLocationScaleFitResult, String> {
fn profile_survival_location_scale(
data: ArrayView2<'_, f64>,
spec: SurvivalLocationScaleTermSpec,
wiggle: Option<LinkWiggleConfig>,
kappa_options: &SpatialLengthScaleOptimizationOptions,
) -> Result<SurvivalLocationScaleProfile, String> {
let mut wiggle_knots = None;
let mut wiggle_degree = None;
let inverse_link = spec.inverse_link.clone();
let fit = if let Some(wiggle) = wiggle {
require_inverse_link_supports_joint_wiggle(&inverse_link, "survival link wiggle")?;
let mut pilot_spec = spec.clone();
pilot_spec.linkwiggle_block = None;
let pilot = fit_survival_location_scale_terms(data, pilot_spec, kappa_options)?;
let selected_wiggle_basis = select_survival_link_wiggle_basis_from_pilot(
&pilot,
&WiggleBlockConfig {
degree: wiggle.degree,
num_internal_knots: wiggle.num_internal_knots,
penalty_order: 2,
double_penalty: wiggle.double_penalty,
},
&wiggle.penalty_orders,
)?;
wiggle_knots = Some(selected_wiggle_basis.knots.clone());
wiggle_degree = Some(selected_wiggle_basis.degree);
fit_survival_location_scale_terms_with_selected_wiggle(
data,
spec,
selected_wiggle_basis,
kappa_options,
)?
} else {
fit_survival_location_scale_terms(data, spec, kappa_options)?
};
Ok(SurvivalLocationScaleProfile {
fit,
inverse_link,
wiggle_knots,
wiggle_degree,
inverse_link_outer: None,
})
}
fn profile_survival_location_scale_with_inverse_link(
data: ArrayView2<'_, f64>,
spec: &SurvivalLocationScaleTermSpec,
inverse_link: InverseLink,
wiggle: Option<LinkWiggleConfig>,
kappa_options: &SpatialLengthScaleOptimizationOptions,
) -> Result<SurvivalLocationScaleProfile, String> {
let mut spec_at_link = spec.clone();
spec_at_link.inverse_link = inverse_link;
profile_survival_location_scale(data, spec_at_link, wiggle, kappa_options)
}
fn optimize_survival_inverse_link_profile(
data: ArrayView2<'_, f64>,
spec: &SurvivalLocationScaleTermSpec,
wiggle: Option<LinkWiggleConfig>,
kappa_options: &SpatialLengthScaleOptimizationOptions,
) -> Result<SurvivalLocationScaleProfile, String> {
fn optimize_link_parameters(
data: ArrayView2<'_, f64>,
spec: &SurvivalLocationScaleTermSpec,
kappa_options: &SpatialLengthScaleOptimizationOptions,
init: Array1<f64>,
name: &str,
wiggle_cfg: Option<LinkWiggleConfig>,
make_link: impl Fn(&Array1<f64>) -> Result<InverseLink, String> + Clone,
) -> Result<SurvivalLocationScaleProfile, String> {
use gam_problem::{DeclaredHessianForm, Derivative, HessianValue, OuterEval};
use gam_solve::rho_optimizer::OuterProblem;
let dim = init.len();
let lower = init.mapv(|v| v - 6.0);
let upper = init.mapv(|v| v + 6.0);
let problem = OuterProblem::new(dim)
.with_gradient(Derivative::Analytic)
.with_hessian(DeclaredHessianForm::Unavailable)
.with_tolerance(1e-4)
.with_max_iter(240)
.with_bounds(lower, upper)
.with_initial_rho(init.clone())
.with_seed_config(gam_problem::SeedConfig {
max_seeds: 1,
seed_budget: 1,
num_auxiliary_trailing: dim,
..Default::default()
});
let context = format!("survival inverse-link optimization ({name}, dim={dim})");
let eval_link = move |theta: &Array1<f64>| -> Result<
ProfiledOuterPayload<SurvivalLocationScaleProfile>,
String,
> {
let link = make_link(theta)?;
let profile = profile_survival_location_scale_with_inverse_link(
data,
spec,
link,
wiggle_cfg.clone(),
kappa_options,
)?;
let cost = survival_inverse_link_profile_objective(
&profile,
&format!("survival inverse-link ({name})"),
)?;
let gradient = profile
.fit
.link_param_data_fit_gradient
.clone()
.ok_or_else(|| {
format!(
"survival inverse-link ({name}): fit reported no link-parameter \
data-fit gradient"
)
})?;
if gradient.len() != theta.len() {
return Err(format!(
"survival inverse-link ({name}): gradient dim {} != theta dim {}",
gradient.len(),
theta.len()
));
}
Ok(ProfiledOuterPayload {
theta: theta.clone(),
objective: cost,
gradient,
value: profile,
})
};
let cost_eval = eval_link.clone();
let cost_fn =
move |selected: &mut Option<ProfiledOuterPayload<SurvivalLocationScaleProfile>>,
theta: &Array1<f64>| {
let payload = cost_eval(theta)
.map_err(gam_solve::estimate::EstimationError::InvalidInput)?;
let cost = payload.objective;
*selected = Some(payload);
Ok(cost)
};
let eval_fn =
move |selected: &mut Option<ProfiledOuterPayload<SurvivalLocationScaleProfile>>,
theta: &Array1<f64>| {
let payload = eval_link(theta)
.map_err(gam_solve::estimate::EstimationError::InvalidInput)?;
let evaluation = OuterEval {
cost: payload.objective,
gradient: payload.gradient.clone(),
hessian: HessianValue::Unavailable,
inner_beta_hint: None,
};
*selected = Some(payload);
Ok(evaluation)
};
let mut obj = problem.build_objective(
None::<ProfiledOuterPayload<SurvivalLocationScaleProfile>>,
cost_fn,
eval_fn,
None::<fn(&mut Option<ProfiledOuterPayload<SurvivalLocationScaleProfile>>)>,
None::<
fn(
&mut Option<ProfiledOuterPayload<SurvivalLocationScaleProfile>>,
&Array1<f64>,
)
-> Result<gam_problem::EfsEval, gam_solve::estimate::EstimationError>,
>,
);
let certified_outer = problem
.run_certified(&mut obj, &context)
.map_err(|err| format!("{context} failed: {err}"))?;
let selected = consume_certified_profiled_outer_payload(
obj.state.take(),
&certified_outer,
&context,
)?;
let replayed_objective =
survival_inverse_link_profile_objective(&selected.value, &context)?;
if replayed_objective.to_bits() != certified_outer.final_value().to_bits() {
return Err(format!(
"{context} retained profile no longer reproduces its certified objective: replayed={replayed_objective:.17e}, certified={:.17e}",
certified_outer.final_value(),
));
}
let mut profile = selected.value;
profile.inverse_link_outer = Some(certified_outer);
Ok(profile)
}
match spec.inverse_link.clone() {
InverseLink::Sas(state0) => optimize_link_parameters(
data,
spec,
kappa_options,
Array1::from_vec(vec![state0.epsilon, state0.log_delta]),
"SAS",
wiggle.clone(),
|theta| {
state_from_sasspec(SasLinkSpec {
initial_epsilon: theta[0],
initial_log_delta: theta[1],
})
.map(InverseLink::Sas)
},
),
InverseLink::BetaLogistic(state0) => optimize_link_parameters(
data,
spec,
kappa_options,
Array1::from_vec(vec![state0.epsilon, state0.log_delta]),
"BetaLogistic",
wiggle.clone(),
|theta| {
state_from_beta_logisticspec(SasLinkSpec {
initial_epsilon: theta[0],
initial_log_delta: theta[1],
})
.map(InverseLink::BetaLogistic)
},
),
InverseLink::Mixture(state0) if !state0.rho.is_empty() => {
let components = state0.components.clone();
optimize_link_parameters(
data,
spec,
kappa_options,
state0.rho.clone(),
"mixture",
wiggle.clone(),
move |rho| {
state_fromspec(&MixtureLinkSpec {
components: components.clone(),
initial_rho: rho.clone(),
})
.map(InverseLink::Mixture)
},
)
}
_ => profile_survival_location_scale(data, spec.clone(), wiggle, kappa_options),
}
}
let profile = if request.optimize_inverse_link {
optimize_survival_inverse_link_profile(
request.data,
&request.spec,
request.wiggle.clone(),
&request.kappa_options,
)?
} else {
profile_survival_location_scale(
request.data,
request.spec.clone(),
request.wiggle.clone(),
&request.kappa_options,
)?
};
Ok(profile.into_result())
}
pub(crate) fn fit_bernoulli_marginal_slope_model(
request: BernoulliMarginalSlopeFitRequest<'_>,
) -> Result<BernoulliMarginalSlopeFitResult, String> {
fit_bernoulli_marginal_slope_terms(
request.data,
request.spec,
&request.options,
&request.kappa_options,
&request.policy,
)
}
pub(crate) fn fit_survival_marginal_slope_model(
request: SurvivalMarginalSlopeFitRequest<'_>,
) -> Result<SurvivalMarginalSlopeFitResult, String> {
fit_survival_marginal_slope_terms(
request.data,
request.spec,
&request.options,
&request.kappa_options,
)
}
pub(crate) fn fit_latent_survival_model(
request: LatentSurvivalFitRequest<'_>,
) -> Result<LatentSurvivalTermFitResult, String> {
fit_latent_survival_terms(
request.data,
request.spec,
request.frailty,
&request.options,
)
}
pub(crate) fn fit_latent_binary_model(
request: LatentBinaryFitRequest<'_>,
) -> Result<LatentBinaryTermFitResult, String> {
fit_latent_binary_terms(
request.data,
request.spec,
request.frailty,
&request.options,
)
}
pub(crate) fn fit_transformation_normal_model(
request: TransformationNormalFitRequest<'_>,
) -> Result<TransformationNormalFitResult, String> {
fit_transformation_normal(
&request.response,
&request.weights,
&request.offset,
request.data,
&request.covariate_spec,
&request.config,
&request.options,
&request.kappa_options,
request.warm_start.as_ref(),
)
}
fn crossfit_fold_count(n: usize) -> usize {
if n < 250 {
n.min(3).max(2)
} else if n < 200_000 {
5
} else if n < 2_000_000 {
3
} else {
2
}
}
fn crossfit_partition(n: usize, k: usize) -> Vec<Vec<usize>> {
let mut folds: Vec<Vec<usize>> = Vec::with_capacity(k);
let base = n / k;
let remainder = n % k;
let mut start = 0usize;
for f in 0..k {
let len = base + usize::from(f < remainder);
let end = start + len;
folds.push((start..end).collect());
start = end;
}
folds
}
fn crossfit_select_rows_1d(source: &Array1<f64>, indices: &[usize]) -> Array1<f64> {
Array1::from_iter(indices.iter().map(|&i| source[i]))
}
pub(crate) fn crossfit_score_calibration(
data: &Dataset,
col_map: &HashMap<String, usize>,
recipe: Option<&CtnStage1Recipe>,
policy: &gam_runtime::resource::ResourcePolicy,
) -> Result<Option<CrossFitScoreCalibration>, String> {
let Some(recipe) = recipe else {
return Ok(None);
};
let n = data.values.nrows();
if n == 0 {
return Err("cross-fit score calibration requires a non-empty dataset".to_string());
}
let y_col = resolve_role_col(col_map, &recipe.response_column, "response")
.map_err(|e| e.to_string())?;
let response_full = data.values.column(y_col).to_owned();
let weights_full = resolve_weight_column(data, col_map, recipe.weight_column.as_deref())
.map_err(|e| e.to_string())?;
let offset_full = resolve_offset_column(data, col_map, recipe.offset_column.as_deref())
.map_err(|e| e.to_string())?;
let parsed_cov = parse_formula(&format!(
"{} ~ {}",
recipe.response_column, recipe.covariate_formula_rhs
))
.map_err(|e| e.to_string())?;
let mut frozen_notes = Vec::new();
let covariate_spec_raw = build_termspec_with_geometry_and_overrides(
&parsed_cov.terms,
data,
col_map,
&mut frozen_notes,
false,
policy,
None,
None,
)
.map_err(|e| e.to_string())?;
let full_cov_design = build_term_collection_design(data.values.view(), &covariate_spec_raw)
.map_err(|e| e.to_string())?;
let frozen_cov_spec = crate::fit_orchestration::drivers::freeze_term_collection_from_design(
&covariate_spec_raw,
&full_cov_design,
)
.map_err(|e| e.to_string())?;
let p_cov = full_cov_design.design.ncols();
let k = crossfit_fold_count(n);
let folds = crossfit_partition(n, k);
let min_complement = folds.iter().map(|held| n - held.len()).min().unwrap_or(n);
let mut fold_config = recipe.config.clone();
fold_config.response_num_internal_knots =
crate::transformation_normal::effective_response_num_internal_knots(
&recipe.config,
min_complement,
p_cov,
response_full.view(),
);
fold_config.response_num_internal_knots_pinned = true;
let mut z_oof = Array1::<f64>::zeros(n);
let mut jac_oof: Option<Array2<f64>> = None;
for held in &folds {
if held.is_empty() {
continue;
}
let held_set: std::collections::HashSet<usize> = held.iter().copied().collect();
let complement: Vec<usize> = (0..n).filter(|i| !held_set.contains(i)).collect();
if complement.is_empty() {
return Err(
"cross-fit fold left an empty training complement; too few rows for K folds"
.to_string(),
);
}
let train_cov = data.values.select(Axis(0), &complement);
let train_resp = crossfit_select_rows_1d(&response_full, &complement);
let train_weights = crossfit_select_rows_1d(&weights_full, &complement);
let train_offset = crossfit_select_rows_1d(&offset_full, &complement);
let fold_fit = fit_transformation_normal(
&train_resp,
&train_weights,
&train_offset,
train_cov.view(),
&frozen_cov_spec,
&fold_config,
&BlockwiseFitOptions::default(),
&SpatialLengthScaleOptimizationOptions::default(),
None,
)?;
let held_cov = data.values.select(Axis(0), held);
let held_resp = crossfit_select_rows_1d(&response_full, held);
let held_offset = crossfit_select_rows_1d(&offset_full, held);
let jac = crate::marginal_slope_orthogonal::score_influence_jacobian(
&fold_fit,
&held_resp,
held_cov.view(),
&held_offset,
)?;
if jac.columns.nrows() != held.len() {
return Err(format!(
"cross-fit fold Jacobian row count {} != held-out fold size {}",
jac.columns.nrows(),
held.len()
));
}
if jac.z.len() != held.len() {
return Err(format!(
"cross-fit fold OOF z length {} != held-out fold size {}",
jac.z.len(),
held.len()
));
}
let p1 = jac.columns.ncols();
let jac_full = jac_oof.get_or_insert_with(|| Array2::<f64>::zeros((n, p1)));
if jac_full.ncols() != p1 {
return Err(format!(
"cross-fit fold p₁ mismatch: this fold has {p1} columns but a prior fold had {}; \
the frozen response/covariate basis failed to align across folds",
jac_full.ncols()
));
}
for (local, &global) in held.iter().enumerate() {
z_oof[global] = jac.z[local];
for c in 0..p1 {
jac_full[[global, c]] = jac.columns[[local, c]];
}
}
}
let jac_oof = jac_oof.ok_or_else(|| {
"cross-fit produced no folds with held-out rows; cannot assemble OOF Jacobian".to_string()
})?;
Ok(Some(CrossFitScoreCalibration { z_oof, jac_oof }))
}
#[cfg(test)]
mod survival_edf_tests {
use super::*;
use crate::survival::PenaltyBlock;
use ndarray::array;
fn penalty_block(matrix: Array2<f64>, lambda: f64, start: usize) -> PenaltyBlock {
let cols = matrix.ncols();
PenaltyBlock {
matrix,
lambda,
range: start..start + cols,
nullspace_dim: 0,
}
}
#[test]
fn survival_edf_exact_trace_on_well_conditioned_hessian() {
let h = array![[4.0, 1.0, 0.0], [1.0, 3.0, 0.0], [0.0, 0.0, 2.0]];
let blocks = vec![penalty_block(array![[1.0, 0.0], [0.0, 1.0]], 1.0, 0)];
let (edf_total, edf_by_block, penalty_block_trace) =
survival_edf_from_dense_hessian(&h, &blocks).expect("PD Hessian must compute EDF");
let expected_trace = 7.0 / 11.0;
assert!(
(penalty_block_trace[0] - expected_trace).abs() < 1e-9,
"penalty trace {:.9} != analytic 7/11",
penalty_block_trace[0]
);
assert!(
(edf_by_block[0] - (2.0 - expected_trace)).abs() < 1e-9,
"per-block EDF {:.9} != 15/11",
edf_by_block[0]
);
assert!(
(edf_total - (3.0 - expected_trace)).abs() < 1e-9,
"total EDF {:.9} != 26/11",
edf_total
);
}
}