use crate::bms::deviation_runtime::AnchorComponentTag;
use crate::bms::{
DeviationRuntime, LatentMeasureKind, LatentZConditionalCalibration, LatentZRankIntCalibration,
};
use crate::cubic_cell_kernel::ANCHORED_DEVIATION_KERNEL;
use crate::fit_orchestration::drivers::freeze_term_collection_from_design;
use crate::fit_orchestration::{FitConfig, StandardFitResult, expectile_tau_for_config};
use crate::inference::model::{
FittedEstimator, FittedFamily, FittedModelPayload, MODEL_PAYLOAD_VERSION, ModelKind,
SavedAnchorComponent, SavedAnchorKind, SavedCompiledFlexBlock, SavedLatentZNormalization,
SavedResidualCascade, SavedSplineScan, SavedSurvivalLocationScaleStructure,
SavedTransformationNormalGeometry, TransformationNormalParameterization,
TransformationScoreCalibration,
};
use crate::scale_design::ScaleDeviationTransform;
use crate::survival::construction::{
SavedSurvivalTimeBasis, SurvivalBaselineConfig, survival_baseline_targetname,
};
use crate::survival::location_scale::{
ResidualDistribution, SurvivalCovariateTimeBasis, SurvivalLocationScaleTimeParameterization,
residual_distribution_from_inverse_link,
};
use crate::transformation_normal::TransformationNormalFamily;
use faer::Side;
use gam_data::{DataSchema, EncodedDataset};
use gam_linalg::faer_ndarray::{FaerCholesky, array2_to_nested_vec};
use gam_problem::types::{
InverseLink, LikelihoodSpec, ResponseFamily, StandardLink, inverse_link_to_binomial_spec,
};
use gam_solve::estimate::{
FittedLinkState, UnifiedFitResult, saved_latent_cloglog_state_from_fit,
saved_mixture_state_from_fit, saved_sas_state_from_fit,
};
use gam_terms::smooth::{TermCollectionDesign, TermCollectionSpec};
use ndarray::{Array1, Array2, s};
const FAMILY_BERNOULLI_MARGINAL_SLOPE: &str = "bernoulli-marginal-slope";
const FAMILY_TRANSFORMATION_NORMAL: &str = "transformation-normal";
pub fn serialize_anchored_deviation_runtime(runtime: &DeviationRuntime) -> SavedCompiledFlexBlock {
let mut anchor_correction: Option<Vec<Vec<f64>>> = None;
let mut anchor_components: Vec<SavedAnchorComponent> = Vec::new();
if let Some(installed) = runtime.installed_flex_block() {
anchor_correction = Some(
installed
.anchor_correction
.rows()
.into_iter()
.map(|row| row.to_vec())
.collect::<Vec<Vec<f64>>>(),
);
for component in &installed.anchor_components {
anchor_components.push(SavedAnchorComponent {
kind: match component {
AnchorComponentTag::Parametric { block, ncols } => {
SavedAnchorKind::Parametric {
block: *block,
ncols: *ncols,
}
}
AnchorComponentTag::FlexEvaluation { ncols } => {
SavedAnchorKind::FlexEvaluation { ncols: *ncols }
}
},
});
}
}
SavedCompiledFlexBlock {
kernel: ANCHORED_DEVIATION_KERNEL.to_string(),
breakpoints: runtime.breakpoints().to_vec(),
basis_dim: runtime.basis_dim(),
span_c0: runtime
.span_c0()
.rows()
.into_iter()
.map(|row| row.to_vec())
.collect(),
span_c1: runtime
.span_c1()
.rows()
.into_iter()
.map(|row| row.to_vec())
.collect(),
span_c2: runtime
.span_c2()
.rows()
.into_iter()
.map(|row| row.to_vec())
.collect(),
span_c3: runtime
.span_c3()
.rows()
.into_iter()
.map(|row| row.to_vec())
.collect(),
anchor_correction,
anchor_components,
}
}
pub struct SavedModelSourceMetadata {
pub training_headers: Vec<String>,
pub training_feature_ranges: Option<Vec<(f64, f64)>>,
pub offset_column: Option<String>,
pub noise_offset_column: Option<String>,
}
impl SavedModelSourceMetadata {
fn apply_to(self, payload: &mut FittedModelPayload) {
match self.training_feature_ranges {
Some(ranges) => payload.set_training_feature_metadata(self.training_headers, ranges),
None => payload.training_headers = Some(self.training_headers),
}
payload.offset_column = self.offset_column;
payload.noise_offset_column = self.noise_offset_column;
}
}
pub struct StandardPayloadInputs<'a> {
pub formula: String,
pub dataset: &'a EncodedDataset,
pub fit_config: &'a FitConfig,
pub result: StandardFitResult,
}
fn fitted_inverse_link(state: &FittedLinkState) -> Option<InverseLink> {
match state {
FittedLinkState::Standard(Some(link)) => Some(InverseLink::Standard(*link)),
FittedLinkState::Standard(None) => None,
FittedLinkState::LatentCLogLog { state } => Some(InverseLink::LatentCLogLog(*state)),
FittedLinkState::Sas { state, .. } => Some(InverseLink::Sas(*state)),
FittedLinkState::BetaLogistic { state, .. } => Some(InverseLink::BetaLogistic(*state)),
FittedLinkState::Mixture { state, .. } => Some(InverseLink::Mixture(state.clone())),
}
}
fn standard_null_space_metadata(
design: &TermCollectionDesign,
fit: &UnifiedFitResult,
) -> Result<(usize, f64), String> {
let hessian = fit
.penalized_hessian()
.ok_or_else(|| "null-space Hessian logdet requires fitted penalized Hessian".to_string())?;
let hessian_dim = hessian.nrows();
if hessian.ncols() != hessian_dim {
return Err(format!(
"null-space Hessian logdet requires a square Hessian, got {}x{}",
hessian.nrows(),
hessian.ncols()
));
}
let p = design.design.ncols();
if hessian_dim < p {
return Err(format!(
"null-space Hessian logdet design/Hessian mismatch: design has {p} columns but \
Hessian is only {hessian_dim}x{hessian_dim}"
));
}
if design.penalties.is_empty() {
return Ok((0, 0.0));
}
let hessian = if hessian_dim > p {
hessian.slice(s![0..p, 0..p]).to_owned()
} else {
hessian.clone()
};
let mut penalty = Array2::<f64>::zeros((p, p));
for (idx, block) in design.penalties.iter().enumerate() {
let range = block.col_range.clone();
if range.start > range.end
|| range.end > p
|| block.local.nrows() != range.len()
|| block.local.ncols() != range.len()
{
return Err(format!(
"null-space Hessian logdet penalty {idx} shape mismatch: range {}..{}, local {}x{}, p={p}",
range.start,
range.end,
block.local.nrows(),
block.local.ncols()
));
}
penalty
.slice_mut(s![range.clone(), range])
.scaled_add(1.0, &block.local);
}
let (null_basis, _) = gam_linalg::faer_ndarray::rrqr_nullspace_basis(
&penalty,
gam_linalg::faer_ndarray::default_rrqr_rank_alpha(),
)
.map_err(|err| format!("failed to compute penalty null-space basis: {err}"))?;
let q = null_basis.ncols();
if q == 0 {
return Ok((0, 0.0));
}
let projected = hessian.dot(&null_basis);
let mut restricted = null_basis.t().dot(&projected);
restricted = (&restricted + &restricted.t()) * 0.5;
let chol = restricted
.cholesky(Side::Lower)
.map_err(|err| format!("null-space Hessian is not positive definite: {err}"))?;
let logdet = 2.0 * chol.diag().iter().map(|value| value.ln()).sum::<f64>();
if logdet.is_finite() {
Ok((q, logdet))
} else {
Err(format!("null-space Hessian logdet is not finite: {logdet}"))
}
}
fn response_for_standard_payload(formula: &str, dataset: &EncodedDataset) -> Option<Array1<f64>> {
let response = gam_terms::inference::formula_dsl::parse_formula(formula)
.ok()?
.response;
let column = *dataset.column_map().get(&response)?;
Some(dataset.values.column(column).to_owned())
}
fn standard_conformal_substrates(
formula: &str,
dataset: &EncodedDataset,
fit_config: &FitConfig,
family: &LikelihoodSpec,
fit: &UnifiedFitResult,
design: &TermCollectionDesign,
) -> (
Option<crate::inference::full_conformal::GaussianJackknifePlusStats>,
Option<crate::inference::full_conformal::ExactFullConformalSubstrate>,
) {
let expectile = fit_config.family.as_deref().is_some_and(|family| {
let family = family.trim().to_ascii_lowercase();
family == "expectile" || family.starts_with("expectile(")
});
if expectile
|| !family.is_gaussian_identity()
|| fit_config.weight_column.is_some()
|| fit_config.offset_column.is_some()
|| fit_config.flexible_link
|| design.affine_offset.iter().any(|value| *value != 0.0)
{
return (None, None);
}
let Some(y) = response_for_standard_payload(formula, dataset) else {
return (None, None);
};
let Ok(x) = design.design.try_to_dense_arc("standard conformal design") else {
return (None, None);
};
let Some(normal_matrix) = fit.penalized_hessian() else {
return (None, None);
};
if x.nrows() != y.len()
|| normal_matrix.nrows() != x.ncols()
|| normal_matrix.ncols() != x.ncols()
{
return (None, None);
}
let weights = Array1::<f64>::ones(y.len());
let jackknife = crate::inference::full_conformal::GaussianJackknifePlusStats::from_design_unit_weight_normal_matrix(
x.as_ref(),
&y,
&weights,
normal_matrix,
)
.ok();
let full = crate::inference::full_conformal::ExactFullConformalSubstrate::from_design_unit_weight_normal_matrix(
x.as_ref(),
&y,
&weights,
normal_matrix,
)
.ok();
(jackknife, full)
}
pub fn assemble_standard_payload(
inputs: StandardPayloadInputs<'_>,
) -> Result<FittedModelPayload, String> {
let StandardPayloadInputs {
formula,
dataset,
fit_config,
result,
} = inputs;
let StandardFitResult {
mut fit,
design,
resolvedspec,
adaptive_diagnostics,
saved_link_state,
wiggle_knots,
wiggle_degree,
wiggle_penalty_metadata,
wiggle_saved_warp_beta,
wiggle_saved_index_shift,
..
} = result;
fit.fitted_link = saved_link_state;
let resolved_termspec = freeze_term_collection_from_design(&resolvedspec, &design)
.map_err(|err| format!("failed to freeze standard term specification: {err}"))?;
let (null_space_dim, null_space_logdet) = standard_null_space_metadata(&design, &fit)?;
fit.artifacts.null_space_dim = Some(null_space_dim);
fit.artifacts.null_space_logdet = Some(null_space_logdet);
let family = fit
.likelihood_family
.clone()
.unwrap_or_else(LikelihoodSpec::gaussian_identity);
let estimator = expectile_tau_for_config(fit_config)
.map_err(|error| format!("failed to persist estimator metadata: {error}"))?
.map_or(FittedEstimator::Likelihood, |tau| {
FittedEstimator::Expectile { tau }
});
let family_label = match estimator {
FittedEstimator::Likelihood => family.name().to_string(),
FittedEstimator::Expectile { tau } => format!("expectile({tau})"),
};
let (gaussian_jackknife_plus, full_conformal) =
standard_conformal_substrates(&formula, dataset, fit_config, &family, &fit, &design);
let latent_cloglog_state = if family.is_latent_cloglog() {
Some(saved_latent_cloglog_state_from_fit(&fit).ok_or_else(|| {
"latent-cloglog-binomial fit did not produce a fitted latent-cloglog state".to_string()
})?)
} else {
saved_latent_cloglog_state_from_fit(&fit)
};
let mut payload = FittedModelPayload::new(
MODEL_PAYLOAD_VERSION,
formula,
ModelKind::Standard,
FittedFamily::Standard {
likelihood: family.clone(),
link: StandardLink::try_from(family.link_function()).ok(),
latent_cloglog_state,
mixture_state: saved_mixture_state_from_fit(&fit),
sas_state: saved_sas_state_from_fit(&fit),
},
family_label,
);
payload.estimator = estimator;
payload.unified = Some(fit.clone());
payload.fit_result = Some(fit.clone());
payload.data_schema = Some(dataset.schema.clone());
payload.link = fitted_inverse_link(&fit.fitted_link).or_else(|| Some(family.link.clone()));
payload.linkwiggle_knots = wiggle_knots.map(|knots| knots.to_vec());
payload.linkwiggle_degree = wiggle_degree;
payload.linkwiggle_penalty_metadata = wiggle_penalty_metadata;
payload.beta_link_wiggle = wiggle_saved_warp_beta;
payload.link_wiggle_index_shift = wiggle_saved_index_shift;
match &fit.fitted_link {
FittedLinkState::Mixture { covariance, .. } => {
payload.mixture_link_param_covariance = covariance.as_ref().map(array2_to_nested_vec);
}
FittedLinkState::Sas { covariance, .. }
| FittedLinkState::BetaLogistic { covariance, .. } => {
payload.sas_param_covariance = covariance.as_ref().map(array2_to_nested_vec);
}
FittedLinkState::Standard(_) | FittedLinkState::LatentCLogLog { .. } => {}
}
payload.set_training_feature_metadata(dataset.headers.clone(), dataset.feature_ranges());
payload.resolved_termspec = Some(resolved_termspec);
payload.adaptive_regularization_diagnostics = adaptive_diagnostics;
payload.offset_column = fit_config.offset_column.clone();
payload.noise_offset_column = fit_config.noise_offset_column.clone();
payload.weight_column = fit_config.weight_column.clone();
payload.gaussian_jackknife_plus = gaussian_jackknife_plus;
payload.full_conformal = full_conformal;
Ok(payload)
}
pub struct BernoulliMarginalSlopeInputs<'a> {
pub formula: String,
pub data_schema: DataSchema,
pub logslope_formula: String,
pub z_column: String,
pub resolved_marginalspec: TermCollectionSpec,
pub resolved_logslopespec: TermCollectionSpec,
pub fit_result: UnifiedFitResult,
pub p_marginal: usize,
pub baseline_marginal: f64,
pub baseline_logslope: f64,
pub latent_z_normalization: SavedLatentZNormalization,
pub latent_measure: LatentMeasureKind,
pub latent_z_rank_int_calibration: Option<LatentZRankIntCalibration>,
pub latent_z_conditional_calibration: Option<LatentZConditionalCalibration>,
pub score_warp_runtime: Option<&'a DeviationRuntime>,
pub link_dev_runtime: Option<&'a DeviationRuntime>,
pub base_link: InverseLink,
pub frailty: crate::survival::lognormal_kernel::FrailtySpec,
}
fn truncate_marginal_slope_influence_absorber(
fit_result: UnifiedFitResult,
p_marginal: usize,
) -> Result<UnifiedFitResult, String> {
let Some(block0) = fit_result.blocks.first() else {
return Err("marginal-slope fit result has no coefficient blocks".to_string());
};
let widened_len = block0.beta.len();
if widened_len <= p_marginal {
return Ok(fit_result);
}
let p_influence = widened_len - p_marginal;
let pirls_status = fit_result.convergence_evidence().inner_status();
let UnifiedFitResult {
mut blocks,
log_lambdas,
lambdas,
likelihood_family,
likelihood_scale,
log_likelihood_normalization,
log_likelihood,
deviance,
reml_score,
stable_penalty_term,
penalized_objective,
used_device,
outer_iterations,
outer_gradient_norm,
standard_deviation,
covariance_conditional,
covariance_corrected,
inference,
fitted_link,
geometry: _,
mut block_states,
beta: _,
max_abs_eta,
constraint_kkt,
artifacts,
inner_cycles,
outer_cost_evals: _,
inner_pirls_solves: _,
..
} = fit_result;
blocks[0].beta = blocks[0].beta.slice(ndarray::s![..p_marginal]).to_owned();
if let Some(state0) = block_states.first_mut() {
state0.beta = state0.beta.slice(ndarray::s![..p_marginal]).to_owned();
}
let drop_gamma_block = |cov: Option<Array2<f64>>| -> Option<Array2<f64>> {
cov.map(|cov| {
let total = cov.nrows();
let kept: Vec<usize> = (0..p_marginal)
.chain((p_marginal + p_influence)..total)
.collect();
let mut out = Array2::<f64>::zeros((kept.len(), kept.len()));
for (ri, &r) in kept.iter().enumerate() {
for (ci, &c) in kept.iter().enumerate() {
out[[ri, ci]] = cov[[r, c]];
}
}
out
})
};
let covariance_conditional = drop_gamma_block(covariance_conditional);
let covariance_corrected = drop_gamma_block(covariance_corrected);
UnifiedFitResult::try_from_parts(gam_solve::estimate::UnifiedFitResultParts {
blocks,
log_lambdas,
lambdas,
likelihood_family,
likelihood_scale,
log_likelihood_normalization,
log_likelihood,
deviance,
reml_score,
stable_penalty_term,
penalized_objective,
used_device,
outer_iterations,
outer_converged: true,
outer_gradient_norm,
standard_deviation,
covariance_conditional,
covariance_corrected,
inference,
fitted_link,
geometry: None,
block_states,
pirls_status,
max_abs_eta,
constraint_kkt,
artifacts,
inner_cycles,
})
.map_err(|e| {
format!("marginal-slope influence-absorber truncation produced an invalid fit result: {e}")
})
}
pub fn assemble_spline_scan_payload(
formula: String,
feature_column: String,
fit: &gam_solve::spline_scan::SplineScanFit,
data_schema: DataSchema,
training_headers: Vec<String>,
training_feature_ranges: Vec<(f64, f64)>,
) -> FittedModelPayload {
let mut payload = FittedModelPayload::new(
MODEL_PAYLOAD_VERSION,
formula,
ModelKind::Standard,
FittedFamily::Standard {
likelihood: LikelihoodSpec::gaussian_identity(),
link: None,
latent_cloglog_state: None,
mixture_state: None,
sas_state: None,
},
"gaussian".to_string(),
);
payload.spline_scan = Some(SavedSplineScan {
feature_column,
state: fit.to_state(),
});
payload.data_schema = Some(data_schema);
payload.set_training_feature_metadata(training_headers, training_feature_ranges);
payload
}
pub fn assemble_residual_cascade_payload(
formula: String,
feature_columns: Vec<String>,
fit: &gam_solve::residual_cascade::ResidualCascadeFit,
data_schema: DataSchema,
training_headers: Vec<String>,
training_feature_ranges: Vec<(f64, f64)>,
) -> Result<FittedModelPayload, String> {
let mut payload = FittedModelPayload::new(
MODEL_PAYLOAD_VERSION,
formula,
ModelKind::Standard,
FittedFamily::Standard {
likelihood: gam_problem::types::LikelihoodSpec::gaussian_identity(),
link: None,
latent_cloglog_state: None,
mixture_state: None,
sas_state: None,
},
"gaussian".to_string(),
);
payload.residual_cascade = Some(SavedResidualCascade {
feature_columns,
state: fit.to_state().map_err(|e| {
format!("residual-cascade to_state failed during payload assembly: {e}")
})?,
});
payload.data_schema = Some(data_schema);
payload.set_training_feature_metadata(training_headers, training_feature_ranges);
Ok(payload)
}
pub fn assemble_bernoulli_marginal_slope_payload(
inputs: BernoulliMarginalSlopeInputs<'_>,
source: SavedModelSourceMetadata,
) -> Result<FittedModelPayload, String> {
let BernoulliMarginalSlopeInputs {
formula,
data_schema,
logslope_formula,
z_column,
resolved_marginalspec,
resolved_logslopespec,
fit_result,
p_marginal,
baseline_marginal,
baseline_logslope,
latent_z_normalization,
latent_measure,
latent_z_rank_int_calibration,
latent_z_conditional_calibration,
score_warp_runtime,
link_dev_runtime,
base_link,
frailty,
} = inputs;
let fit_result = truncate_marginal_slope_influence_absorber(fit_result, p_marginal)?;
let marginal_likelihood_spec =
inverse_link_to_binomial_spec(&base_link).map_err(|e| e.to_string())?;
let mut payload = FittedModelPayload::new(
MODEL_PAYLOAD_VERSION,
formula,
ModelKind::MarginalSlope,
FittedFamily::MarginalSlope {
likelihood: marginal_likelihood_spec,
base_link: base_link.clone(),
frailty,
},
FAMILY_BERNOULLI_MARGINAL_SLOPE.to_string(),
);
payload.unified = Some(fit_result.clone());
payload.fit_result = Some(fit_result);
payload.data_schema = Some(data_schema);
payload.formula_logslope = Some(logslope_formula.clone());
payload.z_column = Some(z_column.clone());
payload.formula_logslopes = Some(vec![logslope_formula]);
payload.z_columns = Some(vec![z_column]);
payload.latent_z_normalization = Some(latent_z_normalization);
payload.latent_measure = Some(latent_measure);
payload.latent_z_rank_int_calibration = latent_z_rank_int_calibration;
payload.latent_z_conditional_calibration = latent_z_conditional_calibration;
payload.marginal_baseline = Some(baseline_marginal);
payload.logslope_baseline = Some(baseline_logslope);
payload.logslope_baselines = Some(vec![baseline_logslope]);
payload.link = Some(base_link);
payload.resolved_termspec = Some(resolved_marginalspec);
payload.resolved_termspec_logslopes = Some(vec![resolved_logslopespec.clone()]);
payload.resolved_termspec_logslope = Some(resolved_logslopespec);
payload.score_warp_runtime = score_warp_runtime.map(serialize_anchored_deviation_runtime);
payload.link_deviation_runtime = link_dev_runtime.map(serialize_anchored_deviation_runtime);
source.apply_to(&mut payload);
Ok(payload)
}
pub struct TransformationNormalInputs<'a> {
pub formula: String,
pub data_schema: DataSchema,
pub resolved_covariate_spec: TermCollectionSpec,
pub fit_result: UnifiedFitResult,
pub family: &'a TransformationNormalFamily,
pub score_calibration: TransformationScoreCalibration,
}
pub fn assemble_transformation_normal_payload(
inputs: TransformationNormalInputs<'_>,
source: SavedModelSourceMetadata,
) -> FittedModelPayload {
let TransformationNormalInputs {
formula,
data_schema,
resolved_covariate_spec,
fit_result,
family,
score_calibration,
} = inputs;
let mut payload = FittedModelPayload::new(
MODEL_PAYLOAD_VERSION,
formula,
ModelKind::TransformationNormal,
FittedFamily::TransformationNormal {
likelihood: LikelihoodSpec::new(
ResponseFamily::Gaussian,
InverseLink::Standard(StandardLink::Identity),
),
},
FAMILY_TRANSFORMATION_NORMAL.to_string(),
);
payload.unified = Some(fit_result.clone());
payload.fit_result = Some(fit_result);
payload.data_schema = Some(data_schema);
payload.resolved_termspec = Some(resolved_covariate_spec);
payload.transformation_response_knots = Some(family.response_knots().to_vec());
payload.transformation_response_transform = Some(
family
.response_transform()
.rows()
.into_iter()
.map(|row| row.to_vec())
.collect(),
);
payload.transformation_response_degree = Some(family.response_degree());
payload.transformation_response_median = Some(family.response_median());
payload.transformation_geometry = Some(transformation_normal_geometry(family));
let cone_carrier = family
.covariate_dense_arc()
.expect("CTN covariate design must materialize for the persisted cone carrier");
payload.transformation_cone_carrier = Some(cone_carrier.iter().copied().collect());
payload.transformation_score_calibration = Some(score_calibration);
source.apply_to(&mut payload);
payload
}
fn transformation_normal_geometry(
family: &TransformationNormalFamily,
) -> SavedTransformationNormalGeometry {
let knots = family.response_knots();
let lo = knots.iter().copied().fold(f64::INFINITY, f64::min);
let hi = knots.iter().copied().fold(f64::NEG_INFINITY, f64::max);
SavedTransformationNormalGeometry {
parameterization: TransformationNormalParameterization::DirectAlpha,
response_degree: family.response_degree(),
response_knot_count: knots.len(),
shape_coordinate_count: family.p_resp().saturating_sub(1),
cone_carrier_covariate_width: family.p_cov(),
cone_carrier_row_count: family.n_obs(),
certified_response_support: (lo, hi),
response_median: family.response_median(),
}
}
pub enum LocationScaleResponse<'a> {
Gaussian {
response_scale: f64,
base_link: Option<InverseLink>,
},
Binomial {
link: InverseLink,
noise_transform: &'a ScaleDeviationTransform,
},
Dispersion {
likelihood: LikelihoodSpec,
base_link: InverseLink,
family_tag: &'static str,
},
}
pub struct LocationScaleWiggle {
pub knots: Vec<f64>,
pub degree: usize,
pub beta_link_wiggle: Vec<f64>,
}
pub struct LocationScaleInputs {
pub formula: String,
pub data_schema: DataSchema,
pub noise_formula: String,
pub resolved_termspec: TermCollectionSpec,
pub resolved_termspec_noise: TermCollectionSpec,
pub fit_result: UnifiedFitResult,
pub beta_noise: Option<Vec<f64>>,
pub wiggle: Option<LocationScaleWiggle>,
}
pub fn assemble_location_scale_payload(
inputs: LocationScaleInputs,
response: LocationScaleResponse<'_>,
source: SavedModelSourceMetadata,
) -> Result<FittedModelPayload, String> {
let (family_tag, likelihood, base_link, link, response_scale, noise_transform) = match response
{
LocationScaleResponse::Gaussian {
response_scale,
base_link,
} => (
"gaussian-location-scale".to_string(),
LikelihoodSpec::gaussian_identity(),
None,
Some(base_link.unwrap_or(InverseLink::Standard(StandardLink::Identity))),
Some(response_scale),
None,
),
LocationScaleResponse::Binomial {
link,
noise_transform,
} => {
let likelihood = inverse_link_to_binomial_spec(&link).map_err(|e| {
format!("failed to resolve LikelihoodSpec for binomial location-scale link {link:?}: {e}")
})?;
(
"binomial-location-scale".to_string(),
likelihood,
Some(link.clone()),
Some(link),
None,
Some(noise_transform),
)
}
LocationScaleResponse::Dispersion {
likelihood,
base_link,
family_tag,
} => (
family_tag.to_string(),
likelihood,
Some(base_link.clone()),
Some(base_link),
None,
None,
),
};
let mut payload = FittedModelPayload::new(
MODEL_PAYLOAD_VERSION,
inputs.formula,
ModelKind::LocationScale,
FittedFamily::LocationScale {
likelihood,
base_link,
},
family_tag,
);
payload.unified = Some(inputs.fit_result.clone());
payload.fit_result = Some(inputs.fit_result);
payload.data_schema = Some(inputs.data_schema);
payload.link = link;
payload.formula_noise = Some(inputs.noise_formula);
payload.beta_noise = inputs.beta_noise;
payload.gaussian_response_scale = response_scale;
if let Some(transform) = noise_transform {
payload.noise_projection = Some(
transform
.projection_coef
.rows()
.into_iter()
.map(|row| row.to_vec())
.collect(),
);
payload.noise_center = Some(transform.weighted_column_mean.to_vec());
payload.noise_scale = Some(transform.rescale.to_vec());
payload.noise_non_intercept_start = Some(transform.non_intercept_start);
payload.noise_projection_ridge_alpha = Some(transform.projection_ridge_alpha);
}
payload.resolved_termspec = Some(inputs.resolved_termspec);
payload.resolved_termspec_noise = Some(inputs.resolved_termspec_noise);
if let Some(wiggle) = inputs.wiggle {
payload.linkwiggle_knots = Some(wiggle.knots);
payload.linkwiggle_degree = Some(wiggle.degree);
payload.beta_link_wiggle = Some(wiggle.beta_link_wiggle);
}
source.apply_to(&mut payload);
Ok(payload)
}
pub struct SurvivalMarginalSlopeInputs<'a> {
pub formula: String,
pub data_schema: DataSchema,
pub fit_result: UnifiedFitResult,
pub frailty: crate::survival::lognormal_kernel::FrailtySpec,
pub survival_entry: Option<String>,
pub survival_exit: String,
pub survival_event: String,
pub survivalspec: String,
pub baseline_cfg: SurvivalBaselineConfig,
pub time_basis: SavedSurvivalTimeBasis,
pub ridge_lambda: f64,
pub survival_likelihood_label: String,
pub resolved_marginalspec: TermCollectionSpec,
pub resolved_logslopespec: TermCollectionSpec,
pub logslope_formula: String,
pub z_column: String,
pub latent_z_normalization: SavedLatentZNormalization,
pub baseline_logslope: f64,
pub timewiggle: Option<SurvivalTimewiggle>,
pub score_warp_runtime: Option<&'a DeviationRuntime>,
pub link_dev_runtime: Option<&'a DeviationRuntime>,
pub influence_absorber_width: Option<usize>,
pub influence_absorber_design: Option<&'a Array2<f64>>,
pub score_covariance: &'a Array2<f64>,
}
fn new_royston_parmar_survival_payload(
formula: String,
fit_result: UnifiedFitResult,
data_schema: DataSchema,
survival_likelihood_label: &str,
survival_distribution: Option<ResidualDistribution>,
frailty: crate::survival::lognormal_kernel::FrailtySpec,
) -> FittedModelPayload {
let mut payload = FittedModelPayload::new(
MODEL_PAYLOAD_VERSION,
formula,
ModelKind::Survival,
FittedFamily::Survival {
likelihood: LikelihoodSpec::new(
ResponseFamily::RoystonParmar,
InverseLink::Standard(StandardLink::Identity),
),
survival_likelihood: Some(survival_likelihood_label.to_string()),
survival_distribution,
frailty,
},
ResponseFamily::RoystonParmar.name().to_string(),
);
payload.unified = Some(fit_result.clone());
payload.fit_result = Some(fit_result);
payload.data_schema = Some(data_schema);
payload
}
pub fn assemble_survival_marginal_slope_payload(
inputs: SurvivalMarginalSlopeInputs<'_>,
source: SavedModelSourceMetadata,
) -> FittedModelPayload {
let mut payload = new_royston_parmar_survival_payload(
inputs.formula,
inputs.fit_result,
inputs.data_schema,
&inputs.survival_likelihood_label,
Some(ResidualDistribution::Gaussian),
inputs.frailty,
);
payload.survival_entry = inputs.survival_entry;
payload.survival_exit = Some(inputs.survival_exit);
payload.survival_event = Some(inputs.survival_event);
payload.survivalspec = Some(inputs.survivalspec);
payload.survival_baseline_target =
Some(survival_baseline_targetname(inputs.baseline_cfg.target).to_string());
payload.survival_baseline_scale = inputs.baseline_cfg.scale;
payload.survival_baseline_shape = inputs.baseline_cfg.shape;
payload.survival_baseline_rate = inputs.baseline_cfg.rate;
payload.survival_baseline_makeham = inputs.baseline_cfg.makeham;
payload.apply_survival_time_basis(&inputs.time_basis);
payload.survivalridge_lambda = Some(inputs.ridge_lambda);
payload.survival_likelihood = Some(inputs.survival_likelihood_label);
payload.survival_distribution = Some(ResidualDistribution::Gaussian);
payload.link = Some(InverseLink::Standard(StandardLink::Probit));
payload.resolved_termspec = Some(inputs.resolved_marginalspec);
payload.resolved_termspec_logslopes = Some(vec![inputs.resolved_logslopespec.clone()]);
payload.resolved_termspec_logslope = Some(inputs.resolved_logslopespec);
payload.formula_logslope = Some(inputs.logslope_formula.clone());
payload.formula_logslopes = Some(vec![inputs.logslope_formula]);
payload.z_column = Some(inputs.z_column.clone());
payload.z_columns = Some(vec![inputs.z_column]);
payload.latent_z_normalization = Some(inputs.latent_z_normalization);
payload.latent_measure = Some(LatentMeasureKind::StandardNormal);
payload.logslope_baseline = Some(inputs.baseline_logslope);
payload.logslope_baselines = Some(vec![inputs.baseline_logslope]);
if let Some(timewiggle) = inputs.timewiggle {
payload.baseline_timewiggle_degree = Some(timewiggle.degree);
payload.baseline_timewiggle_knots = Some(timewiggle.knots);
payload.baseline_timewiggle_penalty_orders = timewiggle.penalty_orders;
payload.baseline_timewiggle_double_penalty = timewiggle.double_penalty;
apply_timewiggle_beta(&mut payload, timewiggle.beta);
}
payload.score_warp_runtime = inputs
.score_warp_runtime
.map(serialize_anchored_deviation_runtime);
payload.link_deviation_runtime = inputs
.link_dev_runtime
.map(serialize_anchored_deviation_runtime);
payload.influence_absorber_width = inputs.influence_absorber_width;
payload.influence_absorber_design = inputs
.influence_absorber_design
.map(|design| design.rows().into_iter().map(|row| row.to_vec()).collect());
payload.survival_marginal_slope_score_covariance = Some(
inputs
.score_covariance
.rows()
.into_iter()
.map(|row| row.to_vec())
.collect(),
);
source.apply_to(&mut payload);
payload
}
pub enum SurvivalTimewiggleBeta {
Single(Vec<f64>),
ByCause(Vec<Vec<f64>>),
}
fn apply_timewiggle_beta(payload: &mut FittedModelPayload, beta: SurvivalTimewiggleBeta) {
match beta {
SurvivalTimewiggleBeta::Single(beta) => {
payload.beta_baseline_timewiggle = Some(beta);
}
SurvivalTimewiggleBeta::ByCause(by_cause) => {
payload.beta_baseline_timewiggle_by_cause = Some(by_cause);
}
}
}
pub struct SurvivalTimewiggle {
pub degree: usize,
pub knots: Vec<f64>,
pub penalty_orders: Option<Vec<usize>>,
pub double_penalty: Option<bool>,
pub beta: SurvivalTimewiggleBeta,
}
pub struct SurvivalTransformationInputs {
pub formula: String,
pub data_schema: DataSchema,
pub fit_result: UnifiedFitResult,
pub survival_entry: Option<String>,
pub survival_exit: String,
pub survival_event: String,
pub survivalspec: String,
pub cause_count: Option<usize>,
pub baseline_cfg: SurvivalBaselineConfig,
pub time_basis: SavedSurvivalTimeBasis,
pub ridge_lambda: f64,
pub survival_likelihood_label: String,
pub resolved_termspec: TermCollectionSpec,
pub survival_beta_time: Option<Vec<f64>>,
pub timewiggle: Option<SurvivalTimewiggle>,
}
pub fn assemble_survival_transformation_payload(
inputs: SurvivalTransformationInputs,
source: SavedModelSourceMetadata,
) -> FittedModelPayload {
let mut payload = new_royston_parmar_survival_payload(
inputs.formula,
inputs.fit_result,
inputs.data_schema,
&inputs.survival_likelihood_label,
None,
crate::survival::lognormal_kernel::FrailtySpec::None,
);
payload.survival_entry = inputs.survival_entry;
payload.survival_exit = Some(inputs.survival_exit);
payload.survival_event = Some(inputs.survival_event);
payload.survivalspec = Some(inputs.survivalspec);
if let Some(cause_count) = inputs.cause_count {
payload.survival_cause_count = Some(cause_count);
payload.survival_endpoint_names = Some(
(1..=cause_count)
.map(|idx| format!("cause_{idx}"))
.collect(),
);
}
payload.survival_baseline_target =
Some(survival_baseline_targetname(inputs.baseline_cfg.target).to_string());
payload.survival_baseline_scale = inputs.baseline_cfg.scale;
payload.survival_baseline_shape = inputs.baseline_cfg.shape;
payload.survival_baseline_rate = inputs.baseline_cfg.rate;
payload.survival_baseline_makeham = inputs.baseline_cfg.makeham;
payload.apply_survival_time_basis(&inputs.time_basis);
if let Some(timewiggle) = inputs.timewiggle {
payload.baseline_timewiggle_degree = Some(timewiggle.degree);
payload.baseline_timewiggle_knots = Some(timewiggle.knots);
payload.baseline_timewiggle_penalty_orders = timewiggle.penalty_orders;
payload.baseline_timewiggle_double_penalty = timewiggle.double_penalty;
apply_timewiggle_beta(&mut payload, timewiggle.beta);
}
payload.survivalridge_lambda = Some(inputs.ridge_lambda);
payload.survival_likelihood = Some(inputs.survival_likelihood_label);
payload.survival_beta_time = inputs.survival_beta_time;
payload.resolved_termspec = Some(inputs.resolved_termspec);
source.apply_to(&mut payload);
payload
}
pub struct SurvivalLocationScaleInputs {
pub formula: String,
pub data_schema: DataSchema,
pub fit_result: UnifiedFitResult,
pub fitted_inverse_link: InverseLink,
pub linkwiggle_degree: Option<usize>,
pub linkwiggle_knots: Option<Vec<f64>>,
pub beta_link_wiggle: Option<Vec<f64>>,
pub baseline_timewiggle: Option<SurvivalTimewiggle>,
pub survival_entry: Option<String>,
pub survival_exit: String,
pub survival_event: String,
pub survivalspec: String,
pub baseline_cfg: SurvivalBaselineConfig,
pub time_basis: SavedSurvivalTimeBasis,
pub ridge_lambda: f64,
pub survival_likelihood_label: String,
pub time_parameterization: SurvivalLocationScaleTimeParameterization,
pub threshold_time_basis: Option<SurvivalCovariateTimeBasis>,
pub log_sigma_time_basis: Option<SurvivalCovariateTimeBasis>,
pub formula_noise: Option<String>,
pub survival_beta_time: Vec<f64>,
pub survival_beta_threshold: Vec<f64>,
pub survival_beta_log_sigma: Vec<f64>,
pub resolved_thresholdspec: TermCollectionSpec,
pub resolved_log_sigmaspec: TermCollectionSpec,
}
pub fn assemble_survival_location_scale_payload(
inputs: SurvivalLocationScaleInputs,
source: SavedModelSourceMetadata,
) -> FittedModelPayload {
let survival_distribution =
residual_distribution_from_inverse_link(&inputs.fitted_inverse_link);
let mut payload = new_royston_parmar_survival_payload(
inputs.formula,
inputs.fit_result,
inputs.data_schema,
&inputs.survival_likelihood_label,
survival_distribution,
crate::survival::lognormal_kernel::FrailtySpec::None,
);
payload.link = Some(inputs.fitted_inverse_link);
payload.linkwiggle_degree = inputs.linkwiggle_degree;
payload.linkwiggle_knots = inputs.linkwiggle_knots;
payload.beta_link_wiggle = inputs.beta_link_wiggle;
if let Some(timewiggle) = inputs.baseline_timewiggle {
payload.baseline_timewiggle_degree = Some(timewiggle.degree);
payload.baseline_timewiggle_knots = Some(timewiggle.knots);
payload.baseline_timewiggle_penalty_orders = timewiggle.penalty_orders;
payload.baseline_timewiggle_double_penalty = timewiggle.double_penalty;
apply_timewiggle_beta(&mut payload, timewiggle.beta);
}
payload.survival_entry = inputs.survival_entry;
payload.survival_exit = Some(inputs.survival_exit);
payload.survival_event = Some(inputs.survival_event);
payload.survivalspec = Some(inputs.survivalspec);
payload.survival_baseline_target =
Some(survival_baseline_targetname(inputs.baseline_cfg.target).to_string());
payload.survival_baseline_scale = inputs.baseline_cfg.scale;
payload.survival_baseline_shape = inputs.baseline_cfg.shape;
payload.survival_baseline_rate = inputs.baseline_cfg.rate;
payload.survival_baseline_makeham = inputs.baseline_cfg.makeham;
payload.apply_survival_time_basis(&inputs.time_basis);
payload.survivalridge_lambda = Some(inputs.ridge_lambda);
payload.survival_likelihood = Some(inputs.survival_likelihood_label);
payload.survival_location_scale_structure = Some(SavedSurvivalLocationScaleStructure {
time_parameterization: inputs.time_parameterization,
threshold_time_basis: inputs.threshold_time_basis,
log_sigma_time_basis: inputs.log_sigma_time_basis,
});
payload.formula_noise = inputs.formula_noise;
payload.survival_beta_time = Some(inputs.survival_beta_time);
payload.survival_beta_threshold = Some(inputs.survival_beta_threshold);
payload.survival_beta_log_sigma = Some(inputs.survival_beta_log_sigma);
payload.survival_distribution = survival_distribution;
payload.resolved_termspec = Some(inputs.resolved_thresholdspec);
payload.resolved_termspec_noise = Some(inputs.resolved_log_sigmaspec);
source.apply_to(&mut payload);
payload
}
pub struct LatentWindowInputs {
pub formula: String,
pub data_schema: DataSchema,
pub fit_result: UnifiedFitResult,
pub family: FittedFamily,
pub model_class_label: String,
pub likelihood_label: String,
pub survival_entry: Option<String>,
pub survival_exit: String,
pub survival_event: String,
pub baseline_cfg: SurvivalBaselineConfig,
pub time_basis: SavedSurvivalTimeBasis,
pub ridge_lambda: f64,
pub beta_time: Vec<f64>,
pub resolved_termspec: TermCollectionSpec,
}
pub fn assemble_latent_window_payload(
inputs: LatentWindowInputs,
source: SavedModelSourceMetadata,
) -> FittedModelPayload {
let mut payload = FittedModelPayload::new(
MODEL_PAYLOAD_VERSION,
inputs.formula,
ModelKind::Survival,
inputs.family,
inputs.model_class_label,
);
payload.unified = Some(inputs.fit_result.clone());
payload.fit_result = Some(inputs.fit_result);
payload.data_schema = Some(inputs.data_schema);
payload.survival_entry = inputs.survival_entry;
payload.survival_exit = Some(inputs.survival_exit);
payload.survival_event = Some(inputs.survival_event);
payload.survivalspec = Some("net".to_string());
payload.survival_baseline_target =
Some(survival_baseline_targetname(inputs.baseline_cfg.target).to_string());
payload.survival_baseline_scale = inputs.baseline_cfg.scale;
payload.survival_baseline_shape = inputs.baseline_cfg.shape;
payload.survival_baseline_rate = inputs.baseline_cfg.rate;
payload.survival_baseline_makeham = inputs.baseline_cfg.makeham;
payload.apply_survival_time_basis(&inputs.time_basis);
payload.survival_likelihood = Some(inputs.likelihood_label);
payload.survival_beta_time = Some(inputs.beta_time);
payload.survivalridge_lambda = Some(inputs.ridge_lambda);
payload.resolved_termspec = Some(inputs.resolved_termspec);
source.apply_to(&mut payload);
payload
}
#[cfg(test)]
mod apply_timewiggle_beta_tests {
use super::*;
fn empty_payload() -> FittedModelPayload {
FittedModelPayload::new(
MODEL_PAYLOAD_VERSION,
"y ~ 1".to_string(),
ModelKind::Survival,
FittedFamily::LatentBinary {
frailty: crate::survival::lognormal_kernel::FrailtySpec::None,
},
"test".to_string(),
)
}
#[test]
fn by_cause_beta_populates_only_the_by_cause_slot() {
let mut payload = empty_payload();
apply_timewiggle_beta(
&mut payload,
SurvivalTimewiggleBeta::ByCause(vec![vec![1.0, 2.0], vec![3.0]]),
);
assert_eq!(
payload.beta_baseline_timewiggle_by_cause,
Some(vec![vec![1.0, 2.0], vec![3.0]]),
"ByCause coefficients must land in the by-cause slot (regression: the \
location-scale assembler used to silently drop them)"
);
assert!(
payload.beta_baseline_timewiggle.is_none(),
"ByCause must not populate the single-block slot"
);
}
#[test]
fn single_beta_populates_only_the_flat_slot() {
let mut payload = empty_payload();
apply_timewiggle_beta(&mut payload, SurvivalTimewiggleBeta::Single(vec![4.0, 5.0]));
assert_eq!(payload.beta_baseline_timewiggle, Some(vec![4.0, 5.0]));
assert!(payload.beta_baseline_timewiggle_by_cause.is_none());
}
}