use std::collections::HashMap;
use ndarray::{Array1, Array2, ArrayView2, s};
use crate::fit_orchestration::prepare_survival_time_stack;
use crate::inference::model::{
FittedFamily, FittedModel as SavedModel, SavedBaselineTimeWiggleRuntime,
load_survival_time_basis_config_from_model, survival_baseline_config_from_model,
};
use crate::inference::predict_io::{BernoulliMarginalSlopePredictor, PredictInput};
use crate::model_types::{BlockRole, FittedBlock, FittedLinkState, UnifiedFitResult};
use crate::probability::signed_probit_logcdf_and_mills_ratio;
use crate::survival::construction::{
SurvivalBaselineConfig, SurvivalBaselineTarget, SurvivalLikelihoodMode,
SurvivalTimeBuildOutput, add_survival_time_derivative_guard_offset, build_survival_time_basis,
build_survival_time_offsets_for_likelihood, build_survival_timewiggle_derivative_design,
center_survival_time_designs_at_anchor, evaluate_survival_time_basis_row,
normalize_survival_time_pair, parse_survival_likelihood_mode,
require_structural_survival_time_basis, resolved_survival_time_basis_config_from_build,
survival_derivative_guard_for_likelihood, survival_likelihood_modename,
};
use crate::survival::latent::fixed_latent_hazard_frailty;
use crate::survival::lognormal_kernel::FrailtySpec;
use crate::survival::{CompetingRisksCifResult, assemble_competing_risks_cif_from_endpoints};
use crate::wiggle::buildwiggle_block_input_from_knots;
use gam_linalg::matrix::DesignMatrix;
use gam_problem::{InverseLink, LikelihoodSpec, ResponseFamily, StandardLink};
use gam_solve::mixture_link::inverse_link_jet_for_inverse_link;
use gam_terms::smooth::TermCollectionSpec;
use gam_terms::smooth::build_term_collection_design;
use gam_terms::term_builder::resolve_role_col;
pub struct SurvivalTimeColumns {
pub entry_col: Option<usize>,
pub exit_col: usize,
}
impl SurvivalTimeColumns {
#[inline]
pub fn row_entry_time(&self, data: ArrayView2<'_, f64>, i: usize) -> f64 {
self.entry_col.map_or(0.0, |idx| data[[i, idx]])
}
}
pub fn resolve_saved_survival_time_columns(
model: &SavedModel,
col_map: &HashMap<String, usize>,
) -> Result<SurvivalTimeColumns, String> {
let entry_col: Option<usize> = model
.survival_entry
.as_deref()
.map(|name| resolve_role_col(col_map, name, "entry"))
.transpose()?;
let exitname = model
.survival_exit
.as_ref()
.ok_or_else(|| "survival model missing exit column metadata".to_string())?;
let exit_col = resolve_role_col(col_map, exitname, "exit")?;
Ok(SurvivalTimeColumns {
entry_col,
exit_col,
})
}
const SURVIVAL_PROB_MIN_FOR_LOG: f64 = 1e-300;
#[derive(Debug, Clone)]
pub enum SurvivalPredictError {
InvalidInput { reason: String },
MissingFitMetadata { reason: String },
IncompatibleSchema { reason: String },
UnsupportedConfiguration { reason: String },
PosteriorCovariance { reason: String },
NumericalFailure { reason: String },
ModelPayload {
context: &'static str,
source: crate::inference::model::FittedModelError,
},
}
impl std::fmt::Display for SurvivalPredictError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
SurvivalPredictError::InvalidInput { reason }
| SurvivalPredictError::MissingFitMetadata { reason }
| SurvivalPredictError::IncompatibleSchema { reason }
| SurvivalPredictError::UnsupportedConfiguration { reason }
| SurvivalPredictError::PosteriorCovariance { reason }
| SurvivalPredictError::NumericalFailure { reason } => f.write_str(reason),
SurvivalPredictError::ModelPayload { context, source } => {
write!(f, "{context}: {source}")
}
}
}
}
impl std::error::Error for SurvivalPredictError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
SurvivalPredictError::ModelPayload { source, .. } => Some(source),
SurvivalPredictError::InvalidInput { .. }
| SurvivalPredictError::MissingFitMetadata { .. }
| SurvivalPredictError::IncompatibleSchema { .. }
| SurvivalPredictError::UnsupportedConfiguration { .. }
| SurvivalPredictError::PosteriorCovariance { .. }
| SurvivalPredictError::NumericalFailure { .. } => None,
}
}
}
impl From<SurvivalPredictError> for String {
fn from(err: SurvivalPredictError) -> String {
err.to_string()
}
}
impl From<String> for SurvivalPredictError {
fn from(reason: String) -> SurvivalPredictError {
SurvivalPredictError::InvalidInput { reason }
}
}
impl From<gam_data::DataError> for SurvivalPredictError {
fn from(err: gam_data::DataError) -> SurvivalPredictError {
SurvivalPredictError::InvalidInput {
reason: err.to_string(),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum SurvivalPredictEstimand {
#[default]
PosteriorMean,
Plugin,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SurvivalPredictionCovarianceMode {
Conditional,
SmoothingCorrected,
}
impl SurvivalPredictionCovarianceMode {
pub const fn as_str(self) -> &'static str {
match self {
Self::Conditional => "conditional",
Self::SmoothingCorrected => "smoothing-corrected",
}
}
}
pub struct SurvivalPredictRequest<'a> {
pub model: &'a SavedModel,
pub data: ArrayView2<'a, f64>,
pub col_map: &'a HashMap<String, usize>,
pub training_headers: Option<&'a Vec<String>>,
pub primary_offset: &'a Array1<f64>,
pub noise_offset: &'a Array1<f64>,
pub time_grid: Option<&'a [f64]>,
pub with_uncertainty: bool,
pub estimand: SurvivalPredictEstimand,
}
pub struct SurvivalPredictResult {
pub times: Vec<f64>,
pub hazard: Array2<f64>,
pub survival: Array2<f64>,
pub cumulative_hazard: Array2<f64>,
pub linear_predictor: Array1<f64>,
pub likelihood_mode: SurvivalLikelihoodMode,
pub survival_se: Option<Array2<f64>>,
pub eta_se: Option<Array1<f64>>,
pub covariance_source: Option<SurvivalPredictionCovarianceMode>,
}
pub struct LatentWindowSurvivalResult {
pub window_survival: Array1<f64>,
pub likelihood_mode: SurvivalLikelihoodMode,
}
pub fn predict_latent_window_survival(
req: SurvivalPredictRequest<'_>,
) -> Result<LatentWindowSurvivalResult, SurvivalPredictError> {
let SurvivalPredictRequest {
model,
data,
col_map,
training_headers,
primary_offset,
noise_offset,
time_grid,
with_uncertainty,
estimand,
} = req;
if time_grid.is_some() {
return Err(SurvivalPredictError::InvalidInput {
reason: "latent-window prediction consumes each row's saved entry/exit columns; an independent time_grid is not a window law".to_string(),
});
}
if with_uncertainty || estimand != SurvivalPredictEstimand::Plugin {
return Err(SurvivalPredictError::UnsupportedConfiguration {
reason: "latent-window observation generation requires the fitted plug-in hazard law; posterior coefficient integration is a different sampling target".to_string(),
});
}
let likelihood_mode = require_saved_survival_likelihood_mode(model)?;
if !matches!(
likelihood_mode,
SurvivalLikelihoodMode::Latent | SurvivalLikelihoodMode::LatentBinary
) {
return Err(SurvivalPredictError::UnsupportedConfiguration {
reason: format!(
"latent-window prediction requires latent or latent-binary likelihood mode, got {}",
survival_likelihood_modename(likelihood_mode)
),
});
}
if model.has_baseline_time_wiggle() {
return Err(SurvivalPredictError::IncompatibleSchema {
reason:
"saved latent survival/binary model contains forbidden baseline timewiggle metadata"
.to_string(),
});
}
let n = data.nrows();
if primary_offset.len() != n || noise_offset.len() != n {
return Err(SurvivalPredictError::InvalidInput {
reason: format!(
"latent-window offset length mismatch: rows={n}, primary={}, noise={}",
primary_offset.len(),
noise_offset.len()
),
});
}
if noise_offset.iter().any(|value| *value != 0.0) {
return Err(SurvivalPredictError::InvalidInput {
reason: "latent-window survival has no secondary offset coordinate".to_string(),
});
}
let termspec = resolve_termspec_for_prediction(
&model.resolved_termspec,
training_headers,
col_map,
"resolved_termspec",
)?;
let clipped = model.axis_clip_to_training_ranges(data, col_map);
let covariate_input = clipped.as_ref().map_or(data, |array| array.view());
let covariate_design = build_term_collection_design(covariate_input, &termspec)
.map_err(|error| format!("failed to build latent-window covariate design: {error}"))?;
let effective_primary_offset = covariate_design
.compose_offset(primary_offset.view(), "latent-window covariate block")
.map_err(|error| error.to_string())?;
let time_columns = resolve_saved_survival_time_columns(model, col_map)?;
let mut age_entry = Array1::<f64>::zeros(n);
let mut age_exit = Array1::<f64>::zeros(n);
for row in 0..n {
let (entry, exit) = normalize_survival_time_pair(
time_columns.row_entry_time(data, row),
data[[row, time_columns.exit_col]],
row,
)?;
age_entry[row] = entry;
age_exit[row] = exit;
}
let time_config = load_survival_time_basis_config_from_model(model)?;
let mut time_build = build_survival_time_basis(&age_entry, &age_exit, time_config, None)?;
let resolved_time_config = resolved_survival_time_basis_config_from_build(
&time_build.basisname,
time_build.degree,
time_build.knots.as_ref(),
time_build.keep_cols.as_ref(),
time_build.smooth_lambda,
)?;
let time_anchor =
model
.survival_time_anchor
.ok_or_else(|| SurvivalPredictError::MissingFitMetadata {
reason: "saved latent-window model is missing survival_time_anchor".to_string(),
})?;
let anchor_row = evaluate_survival_time_basis_row(time_anchor, &resolved_time_config)?;
center_survival_time_designs_at_anchor(
&mut time_build.x_entry_time,
&mut time_build.x_exit_time,
&anchor_row,
)?;
require_structural_survival_time_basis(
&time_build.basisname,
"saved latent-window prediction",
)?;
let frailty =
model
.family_state
.frailty()
.ok_or_else(|| SurvivalPredictError::MissingFitMetadata {
reason: "saved latent-window model is missing its hazard-multiplier frailty"
.to_string(),
})?;
let (sigma, loading) = fixed_latent_hazard_frailty(frailty, "saved latent-window prediction")
.map_err(|reason| SurvivalPredictError::MissingFitMetadata { reason })?;
let baseline_config = saved_survival_runtime_baseline_config(model)?;
let prepared = prepare_survival_time_stack(
&age_entry,
&age_exit,
&baseline_config,
likelihood_mode,
None,
time_anchor,
survival_derivative_guard_for_likelihood(likelihood_mode),
&time_build,
None,
Some(loading),
)?;
let fit = fit_result_from_saved_model_for_prediction(model)?;
let mean_block = fit.block_by_role(BlockRole::Mean).ok_or_else(|| {
SurvivalPredictError::MissingFitMetadata {
reason: "saved latent-window model is missing its mean coefficient block".to_string(),
}
})?;
let time_block = fit.block_by_role(BlockRole::Time).ok_or_else(|| {
SurvivalPredictError::MissingFitMetadata {
reason: "saved latent-window model is missing its time coefficient block".to_string(),
}
})?;
if mean_block.beta.len() != covariate_design.design.ncols() {
return Err(SurvivalPredictError::IncompatibleSchema {
reason: format!(
"latent-window mean/design mismatch: beta has {} coefficients but design has {} columns",
mean_block.beta.len(),
covariate_design.design.ncols()
),
});
}
if time_block.beta.len() != prepared.time_design_exit.ncols() {
let hint = stale_weibull_time_basis_hint(
&time_build.basisname,
time_block.beta.len() == prepared.time_design_exit.ncols() + 1,
);
return Err(SurvivalPredictError::IncompatibleSchema {
reason: format!(
"latent-window time/design mismatch: beta has {} coefficients but design has {} columns{hint}",
time_block.beta.len(),
prepared.time_design_exit.ncols()
),
});
}
let eta = covariate_design.design.dot(&mean_block.beta) + &effective_primary_offset;
let q_entry = prepared.time_design_entry.dot(&time_block.beta) + &prepared.eta_offset_entry;
let q_exit = prepared.time_design_exit.dot(&time_block.beta) + &prepared.eta_offset_exit;
let quadrature = gam_solve::quadrature::QuadratureContext::new();
let mut window_survival = Array1::<f64>::zeros(n);
for row in 0..n {
let latent_row = crate::survival::lognormal_kernel::LatentSurvivalRow::right_censored(
q_entry[row].exp(),
q_exit[row].exp(),
prepared.unloaded_mass_entry[row],
prepared.unloaded_mass_exit[row],
);
let jet = crate::survival::lognormal_kernel::LatentSurvivalRowJet::evaluate(
&quadrature,
&latent_row,
eta[row],
sigma,
)
.map_err(|error| SurvivalPredictError::NumericalFailure {
reason: format!("latent-window row {row} evaluation failed: {error}"),
})?;
let survival = jet.log_lik.exp();
if !(survival.is_finite() && (0.0..=1.0).contains(&survival)) {
return Err(SurvivalPredictError::NumericalFailure {
reason: format!(
"latent-window row {row} produced invalid conditional survival {survival}"
),
});
}
window_survival[row] = survival;
}
Ok(LatentWindowSurvivalResult {
window_survival,
likelihood_mode,
})
}
fn select_survival_prediction_covariance<'a>(
conditional: Option<&'a Array2<f64>>,
smoothing_corrected: Option<&'a Array2<f64>>,
mode: SurvivalPredictionCovarianceMode,
) -> Result<&'a Array2<f64>, SurvivalPredictError> {
match mode {
SurvivalPredictionCovarianceMode::Conditional => {
conditional.ok_or_else(|| SurvivalPredictError::PosteriorCovariance {
reason: "fit result does not contain conditional covariance".to_string(),
})
}
SurvivalPredictionCovarianceMode::SmoothingCorrected => {
smoothing_corrected.ok_or_else(|| SurvivalPredictError::PosteriorCovariance {
reason: "fit result does not contain smoothing-corrected covariance".to_string(),
})
}
}
}
fn survival_prediction_posterior_factor(
model: &SavedModel,
covariance_mode: SurvivalPredictionCovarianceMode,
) -> Result<(Array1<f64>, Array2<f64>, Vec<usize>), SurvivalPredictError> {
let fit = fit_result_from_saved_model_for_prediction(model)?;
let inactive_tail = if require_saved_survival_likelihood_mode(model)?
== SurvivalLikelihoodMode::MarginalSlope
{
model
.saved_prediction_runtime()?
.influence_absorber_width
.unwrap_or(0)
} else {
0
};
let active_len = fit.beta.len().checked_sub(inactive_tail).ok_or_else(|| {
SurvivalPredictError::IncompatibleSchema {
reason: format!(
"saved survival influence-absorber width {inactive_tail} exceeds the {} fitted coefficients",
fit.beta.len()
),
}
})?;
let covariance = select_survival_prediction_covariance(
fit.beta_covariance(),
fit.beta_covariance_corrected(),
covariance_mode,
)?;
if covariance.nrows() != fit.beta.len() || covariance.ncols() != fit.beta.len() {
return Err(SurvivalPredictError::PosteriorCovariance {
reason: format!(
"saved survival {} covariance has shape {}x{}, expected {}x{} in fitted block order",
covariance_mode.as_str(),
covariance.nrows(),
covariance.ncols(),
fit.beta.len(),
fit.beta.len(),
),
});
}
let cone_coords = survival_posterior_cone_coordinates(model, active_len)?;
Ok((
fit.beta.clone(),
covariance.slice(s![..active_len, ..active_len]).to_owned(),
cone_coords,
))
}
fn saved_model_with_survival_coefficients(
model: &SavedModel,
coefficients: &Array1<f64>,
) -> Result<SavedModel, SurvivalPredictError> {
let mut draw_model = model.clone();
let payload = match &mut draw_model {
SavedModel::Standard { payload }
| SavedModel::LocationScale { payload }
| SavedModel::MarginalSlope { payload }
| SavedModel::Survival { payload }
| SavedModel::TransformationNormal { payload } => payload,
};
let (beta_time, beta_threshold, beta_log_sigma, beta_link_wiggle, beta_time_blocks) = {
let fit = payload.fit_result.as_mut().ok_or_else(|| {
SurvivalPredictError::MissingFitMetadata {
reason: "saved survival model is missing canonical fit_result".to_string(),
}
})?;
if coefficients.len() != fit.beta.len() {
return Err(SurvivalPredictError::IncompatibleSchema {
reason: format!(
"posterior survival coefficient draw has length {}, expected {}",
coefficients.len(),
fit.beta.len()
),
});
}
fit.beta.assign(coefficients);
let mut cursor = 0usize;
for block in &mut fit.blocks {
let end = cursor + block.beta.len();
block.beta.assign(&coefficients.slice(s![cursor..end]));
cursor = end;
}
if cursor != coefficients.len() {
return Err(SurvivalPredictError::IncompatibleSchema {
reason: format!(
"saved survival coefficient blocks total {cursor} entries, but the joint vector has {}",
coefficients.len()
),
});
}
(
fit.block_by_role(BlockRole::Time)
.map(|block| block.beta.to_vec()),
fit.block_by_role(BlockRole::Threshold)
.map(|block| block.beta.to_vec()),
fit.block_by_role(BlockRole::Scale)
.map(|block| block.beta.to_vec()),
fit.block_by_role(BlockRole::LinkWiggle)
.map(|block| block.beta.to_vec()),
fit.blocks
.iter()
.map(|block| block.beta.to_vec())
.collect::<Vec<_>>(),
)
};
if payload.survival_beta_time.is_some() {
payload.survival_beta_time = beta_time.clone();
}
if payload.survival_beta_threshold.is_some() {
payload.survival_beta_threshold = beta_threshold;
}
if payload.survival_beta_log_sigma.is_some() {
payload.survival_beta_log_sigma = beta_log_sigma;
}
if payload.beta_link_wiggle.is_some() {
payload.beta_link_wiggle = beta_link_wiggle;
}
if let (Some(saved), Some(time_beta)) = (
payload.beta_baseline_timewiggle.as_mut(),
beta_time.as_ref(),
) {
if saved.len() > time_beta.len() {
return Err(SurvivalPredictError::IncompatibleSchema {
reason: format!(
"saved baseline-timewiggle has {} coefficients, but the time block has {}",
saved.len(),
time_beta.len()
),
});
}
*saved = time_beta[time_beta.len() - saved.len()..].to_vec();
}
if let Some(saved_by_cause) = payload.beta_baseline_timewiggle_by_cause.as_mut() {
if saved_by_cause.len() != beta_time_blocks.len() {
return Err(SurvivalPredictError::IncompatibleSchema {
reason: format!(
"saved cause-specific timewiggles have {} blocks, but the fit has {} cause blocks",
saved_by_cause.len(),
beta_time_blocks.len()
),
});
}
for (saved, block) in saved_by_cause.iter_mut().zip(&beta_time_blocks) {
if saved.len() > block.len() {
return Err(SurvivalPredictError::IncompatibleSchema {
reason: format!(
"saved cause-specific timewiggle has {} coefficients, but its endpoint block has {}",
saved.len(),
block.len()
),
});
}
*saved = block[block.len() - saved.len()..].to_vec();
}
}
Ok(draw_model)
}
fn conditional_event_density(
survival: f64,
cumulative_hazard: f64,
hazard: f64,
) -> Result<f64, SurvivalPredictError> {
if hazard == 0.0 {
return Ok(0.0);
}
if survival > 0.0 && hazard.is_finite() {
return Ok(survival * hazard);
}
if cumulative_hazard.is_finite() && hazard > 0.0 {
return Ok((hazard.ln() - cumulative_hazard).exp());
}
if cumulative_hazard == f64::INFINITY && hazard.is_finite() && hazard >= 0.0 {
return Ok(0.0);
}
Err(SurvivalPredictError::NumericalFailure {
reason: format!(
"posterior survival quadrature could not resolve conditional density from S={survival}, H={cumulative_hazard}, h={hazard}"
),
})
}
fn for_each_survival_posterior_node<F>(
posterior_mean: &Array1<f64>,
active_covariance: &Array2<f64>,
cone_coords: &[usize],
mut consume: F,
) -> Result<(), SurvivalPredictError>
where
F: FnMut(&Array1<f64>, f64) -> Result<(), SurvivalPredictError>,
{
let active_len = active_covariance.nrows();
if active_covariance.ncols() != active_len || active_len > posterior_mean.len() {
return Err(SurvivalPredictError::PosteriorCovariance {
reason: format!(
"survival posterior quadrature received mean length {} and active covariance {}x{}",
posterior_mean.len(),
active_covariance.nrows(),
active_covariance.ncols(),
),
});
}
let factorization = crate::survival::location_scale::factorize_psd_covariance(
active_covariance,
"survival posterior coefficient covariance",
)
.map_err(|reason| SurvivalPredictError::PosteriorCovariance { reason })?;
let rank = factorization.factor.ncols();
if rank == 0 {
return consume(posterior_mean, 1.0);
}
let nominal_scale = (rank as f64).sqrt();
let weight = 1.0 / (2 * rank) as f64;
for column in 0..rank {
let mut scale = nominal_scale;
for &j in cone_coords {
if j >= active_len {
continue;
}
let load = factorization.factor[[j, column]].abs();
if load == 0.0 {
continue;
}
let limit = posterior_mean[j].max(0.0) / load;
if limit < scale {
scale = limit;
}
}
for sign in [-1.0_f64, 1.0_f64] {
let mut node = posterior_mean.clone();
for row in 0..active_len {
node[row] += sign * scale * factorization.factor[[row, column]];
}
consume(&node, weight)?;
}
}
Ok(())
}
fn survival_posterior_cone_coordinates(
model: &SavedModel,
active_len: usize,
) -> Result<Vec<usize>, SurvivalPredictError> {
if require_saved_survival_likelihood_mode(model)? != SurvivalLikelihoodMode::Transformation {
return Ok(Vec::new());
}
let time_cfg = load_survival_time_basis_config_from_model(model)
.map_err(|err| SurvivalPredictError::MissingFitMetadata {
reason: err.to_string(),
})?;
let p_time_base = if matches!(
time_cfg,
crate::survival::construction::SurvivalTimeBasisConfig::None
) {
0
} else {
let dummy = Array1::from_elem(1, 1.0_f64);
build_survival_time_basis(&dummy, &dummy, time_cfg, None)
.map_err(|reason| SurvivalPredictError::MissingFitMetadata { reason })?
.x_exit_time
.ncols()
};
let fit = fit_result_from_saved_model_for_prediction(model)?;
let cause_count = model
.survival_cause_count
.unwrap_or(fit.blocks.len())
.max(1);
let per_cause_wiggle: Vec<usize> = if cause_count > 1 {
saved_cause_specific_timewiggles(model, &fit, cause_count)?
.iter()
.map(|w| w.as_ref().map_or(0, |runtime| runtime.beta.len()))
.collect()
} else {
vec![
model
.saved_baseline_time_wiggle()
.map_err(|err| SurvivalPredictError::MissingFitMetadata {
reason: err.to_string(),
})?
.map_or(0, |runtime| runtime.beta.len()),
]
};
let mut cone = Vec::new();
let mut cursor = 0usize;
for (cause, block) in fit.blocks.iter().enumerate() {
let block_len = block.beta.len();
let width = (p_time_base + per_cause_wiggle.get(cause).copied().unwrap_or(0)).min(block_len);
for j in cursor..cursor + width {
if j < active_len {
cone.push(j);
}
}
cursor += block_len;
}
Ok(cone)
}
fn posterior_standard_error_matrix(
mean: &Array2<f64>,
second_moment: &Array2<f64>,
label: &str,
) -> Result<Array2<f64>, SurvivalPredictError> {
if second_moment.dim() != mean.dim() {
return Err(SurvivalPredictError::IncompatibleSchema {
reason: format!(
"posterior {label} moment shape mismatch: mean={:?}, second={:?}",
mean.dim(),
second_moment.dim(),
),
});
}
let mut standard_error = Array2::<f64>::zeros(mean.raw_dim());
for ((row, column), slot) in standard_error.indexed_iter_mut() {
let first = mean[[row, column]];
let second = second_moment[[row, column]];
if !(first.is_finite() && second.is_finite()) {
return Err(SurvivalPredictError::NumericalFailure {
reason: format!(
"posterior {label} moments must be finite at row {row}, time column {column}: mean={first}, second={second}"
),
});
}
let variance = second - first * first;
let roundoff_tolerance =
128.0 * f64::EPSILON * second.abs().max((first * first).abs()).max(1.0);
if variance < -roundoff_tolerance {
return Err(SurvivalPredictError::NumericalFailure {
reason: format!(
"posterior {label} variance is negative beyond roundoff at row {row}, time column {column}: {variance}"
),
});
}
*slot = variance.max(0.0).sqrt();
}
Ok(standard_error)
}
fn posterior_standard_error_vector(
mean: &Array1<f64>,
second_moment: &Array1<f64>,
label: &str,
) -> Result<Array1<f64>, SurvivalPredictError> {
if second_moment.len() != mean.len() {
return Err(SurvivalPredictError::IncompatibleSchema {
reason: format!(
"posterior {label} moment length mismatch: mean={}, second={}",
mean.len(),
second_moment.len(),
),
});
}
let mut standard_error = Array1::<f64>::zeros(mean.len());
for row in 0..mean.len() {
let first = mean[row];
let second = second_moment[row];
if !(first.is_finite() && second.is_finite()) {
return Err(SurvivalPredictError::NumericalFailure {
reason: format!(
"posterior {label} moments must be finite at row {row}: mean={first}, second={second}"
),
});
}
let variance = second - first * first;
let roundoff_tolerance =
128.0 * f64::EPSILON * second.abs().max((first * first).abs()).max(1.0);
if variance < -roundoff_tolerance {
return Err(SurvivalPredictError::NumericalFailure {
reason: format!(
"posterior {label} variance is negative beyond roundoff at row {row}: {variance}"
),
});
}
standard_error[row] = variance.max(0.0).sqrt();
}
Ok(standard_error)
}
fn posterior_standard_error_surfaces(
mean: &[Array2<f64>],
second_moment: &[Array2<f64>],
label: &str,
) -> Result<Vec<Array2<f64>>, SurvivalPredictError> {
if second_moment.len() != mean.len() {
return Err(SurvivalPredictError::IncompatibleSchema {
reason: format!(
"posterior {label} cause count mismatch: mean={}, second={}",
mean.len(),
second_moment.len(),
),
});
}
mean.iter()
.zip(second_moment)
.enumerate()
.map(|(cause, (first, second))| {
posterior_standard_error_matrix(first, second, &format!("{label} cause {}", cause + 1))
})
.collect()
}
fn posterior_standard_error_vectors(
mean: &[Array1<f64>],
second_moment: &[Array1<f64>],
label: &str,
) -> Result<Vec<Array1<f64>>, SurvivalPredictError> {
if second_moment.len() != mean.len() {
return Err(SurvivalPredictError::IncompatibleSchema {
reason: format!(
"posterior {label} cause count mismatch: mean={}, second={}",
mean.len(),
second_moment.len(),
),
});
}
mean.iter()
.zip(second_moment)
.enumerate()
.map(|(cause, (first, second))| {
posterior_standard_error_vector(first, second, &format!("{label} cause {}", cause + 1))
})
.collect()
}
fn predict_survival_posterior_mean(
req: SurvivalPredictRequest<'_>,
covariance_mode: SurvivalPredictionCovarianceMode,
) -> Result<SurvivalPredictResult, SurvivalPredictError> {
let (posterior_mean, active_covariance, cone_coords) =
survival_prediction_posterior_factor(req.model, covariance_mode)?;
let mut result = predict_survival(
SurvivalPredictRequest {
model: req.model,
data: req.data,
col_map: req.col_map,
training_headers: req.training_headers,
primary_offset: req.primary_offset,
noise_offset: req.noise_offset,
time_grid: req.time_grid,
with_uncertainty: false,
estimand: SurvivalPredictEstimand::Plugin,
},
covariance_mode,
)?;
let (n_rows, n_times) = result.survival.dim();
let mut survival_mean = Array2::<f64>::zeros((n_rows, n_times));
let mut survival_second = Array2::<f64>::zeros((n_rows, n_times));
let mut density_mean = Array2::<f64>::zeros((n_rows, n_times));
let mut hazard_mean = Array2::<f64>::zeros((n_rows, n_times));
let mut eta_mean = Array1::<f64>::zeros(n_rows);
let mut eta_second = Array1::<f64>::zeros(n_rows);
for_each_survival_posterior_node(&posterior_mean, &active_covariance, &cone_coords, |node, weight| {
let draw_model = saved_model_with_survival_coefficients(req.model, node)?;
let draw = predict_survival(
SurvivalPredictRequest {
model: &draw_model,
data: req.data,
col_map: req.col_map,
training_headers: req.training_headers,
primary_offset: req.primary_offset,
noise_offset: req.noise_offset,
time_grid: req.time_grid,
with_uncertainty: false,
estimand: SurvivalPredictEstimand::Plugin,
},
covariance_mode,
)?;
if draw.survival.dim() != (n_rows, n_times)
|| draw.hazard.dim() != (n_rows, n_times)
|| draw.cumulative_hazard.dim() != (n_rows, n_times)
|| draw.linear_predictor.len() != n_rows
|| draw.times != result.times
|| draw.likelihood_mode != result.likelihood_mode
{
return Err(SurvivalPredictError::IncompatibleSchema {
reason: "posterior survival quadrature node changed the prediction schema"
.to_string(),
});
}
for row in 0..n_rows {
let eta = draw.linear_predictor[row];
eta_mean[row] += weight * eta;
eta_second[row] += weight * eta * eta;
for time in 0..n_times {
let survival = draw.survival[[row, time]];
let hazard = draw.hazard[[row, time]];
let density = conditional_event_density(
survival,
draw.cumulative_hazard[[row, time]],
hazard,
)?;
survival_mean[[row, time]] += weight * survival;
survival_second[[row, time]] += weight * survival * survival;
density_mean[[row, time]] += weight * density;
hazard_mean[[row, time]] += weight * hazard;
}
}
Ok(())
})?;
for row in 0..n_rows {
for time in 0..n_times {
let survival = survival_mean[[row, time]].clamp(0.0, 1.0);
let density = density_mean[[row, time]];
if !(density.is_finite() && density >= 0.0) {
return Err(SurvivalPredictError::NumericalFailure {
reason: format!(
"posterior survival density is invalid at row {row}, time column {time}: {density}"
),
});
}
result.survival[[row, time]] = survival;
result.cumulative_hazard[[row, time]] = -survival.ln();
result.hazard[[row, time]] = if survival > 0.0 {
density / survival
} else if hazard_mean[[row, time]] == 0.0 {
0.0
} else {
f64::INFINITY
};
}
}
result.survival_se = req.with_uncertainty.then(|| {
Array2::from_shape_fn((n_rows, n_times), |(row, time)| {
(survival_second[[row, time]] - survival_mean[[row, time]] * survival_mean[[row, time]])
.max(0.0)
.sqrt()
})
});
result.eta_se = req.with_uncertainty.then(|| {
Array1::from_shape_fn(n_rows, |row| {
(eta_second[row] - eta_mean[row] * eta_mean[row])
.max(0.0)
.sqrt()
})
});
result.covariance_source = req.with_uncertainty.then_some(covariance_mode);
Ok(result)
}
fn predict_competing_risks_with_posterior(
req: SurvivalPredictRequest<'_>,
covariance_mode: SurvivalPredictionCovarianceMode,
) -> Result<CompetingRisksPredictResult, SurvivalPredictError> {
let posterior_mean_estimand = req.estimand == SurvivalPredictEstimand::PosteriorMean;
let (posterior_mean, active_covariance, cone_coords) =
survival_prediction_posterior_factor(req.model, covariance_mode)?;
let separate_conditional_point = posterior_mean_estimand
&& req.with_uncertainty
&& covariance_mode == SurvivalPredictionCovarianceMode::SmoothingCorrected;
let mut result = if separate_conditional_point {
predict_competing_risks_with_posterior(
SurvivalPredictRequest {
model: req.model,
data: req.data,
col_map: req.col_map,
training_headers: req.training_headers,
primary_offset: req.primary_offset,
noise_offset: req.noise_offset,
time_grid: req.time_grid,
with_uncertainty: false,
estimand: SurvivalPredictEstimand::PosteriorMean,
},
SurvivalPredictionCovarianceMode::Conditional,
)?
} else {
predict_competing_risks_survival(
SurvivalPredictRequest {
model: req.model,
data: req.data,
col_map: req.col_map,
training_headers: req.training_headers,
primary_offset: req.primary_offset,
noise_offset: req.noise_offset,
time_grid: req.time_grid,
with_uncertainty: false,
estimand: SurvivalPredictEstimand::Plugin,
},
SurvivalPredictionCovarianceMode::Conditional,
)?
};
let cause_count = result.cif.len();
let (n_rows, n_times) = result.overall_survival.dim();
let mut survival_mean = (0..cause_count)
.map(|_| Array2::<f64>::zeros((n_rows, n_times)))
.collect::<Vec<_>>();
let mut survival_second = (0..cause_count)
.map(|_| Array2::<f64>::zeros((n_rows, n_times)))
.collect::<Vec<_>>();
let mut hazard_mean = (0..cause_count)
.map(|_| Array2::<f64>::zeros((n_rows, n_times)))
.collect::<Vec<_>>();
let mut hazard_second = (0..cause_count)
.map(|_| Array2::<f64>::zeros((n_rows, n_times)))
.collect::<Vec<_>>();
let mut cumulative_hazard_mean = (0..cause_count)
.map(|_| Array2::<f64>::zeros((n_rows, n_times)))
.collect::<Vec<_>>();
let mut cumulative_hazard_second = (0..cause_count)
.map(|_| Array2::<f64>::zeros((n_rows, n_times)))
.collect::<Vec<_>>();
let mut cif_mean = (0..cause_count)
.map(|_| Array2::<f64>::zeros((n_rows, n_times)))
.collect::<Vec<_>>();
let mut cif_second = (0..cause_count)
.map(|_| Array2::<f64>::zeros((n_rows, n_times)))
.collect::<Vec<_>>();
let mut overall_mean = Array2::<f64>::zeros((n_rows, n_times));
let mut overall_second = Array2::<f64>::zeros((n_rows, n_times));
let mut eta_mean = (0..cause_count)
.map(|_| Array1::<f64>::zeros(n_rows))
.collect::<Vec<_>>();
let mut eta_second = (0..cause_count)
.map(|_| Array1::<f64>::zeros(n_rows))
.collect::<Vec<_>>();
for_each_survival_posterior_node(&posterior_mean, &active_covariance, &cone_coords, |node, weight| {
let draw_model = saved_model_with_survival_coefficients(req.model, node)?;
let draw = predict_competing_risks_survival(
SurvivalPredictRequest {
model: &draw_model,
data: req.data,
col_map: req.col_map,
training_headers: req.training_headers,
primary_offset: req.primary_offset,
noise_offset: req.noise_offset,
time_grid: req.time_grid,
with_uncertainty: false,
estimand: SurvivalPredictEstimand::Plugin,
},
SurvivalPredictionCovarianceMode::Conditional,
)?;
if draw.cif.len() != cause_count
|| draw.survival.len() != cause_count
|| draw.hazard.len() != cause_count
|| draw.cumulative_hazard.len() != cause_count
|| draw.linear_predictor.len() != cause_count
|| draw.overall_survival.dim() != (n_rows, n_times)
|| draw.times != result.times
|| draw.endpoint_names != result.endpoint_names
|| draw.likelihood_mode != result.likelihood_mode
{
return Err(SurvivalPredictError::IncompatibleSchema {
reason: "posterior competing-risks quadrature node changed the prediction schema"
.to_string(),
});
}
for cause in 0..cause_count {
if draw.survival[cause].dim() != (n_rows, n_times)
|| draw.hazard[cause].dim() != (n_rows, n_times)
|| draw.cumulative_hazard[cause].dim() != (n_rows, n_times)
|| draw.cif[cause].dim() != (n_rows, n_times)
|| draw.linear_predictor[cause].len() != n_rows
{
return Err(SurvivalPredictError::IncompatibleSchema {
reason: format!(
"posterior competing-risks quadrature node changed cause {} surface dimensions",
cause + 1
),
});
}
for row in 0..n_rows {
let eta = draw.linear_predictor[cause][row];
eta_mean[cause][row] += weight * eta;
eta_second[cause][row] += weight * eta * eta;
for time in 0..n_times {
let survival = draw.survival[cause][[row, time]];
let hazard = draw.hazard[cause][[row, time]];
let cumulative_hazard = draw.cumulative_hazard[cause][[row, time]];
let cif = draw.cif[cause][[row, time]];
survival_mean[cause][[row, time]] += weight * survival;
survival_second[cause][[row, time]] += weight * survival * survival;
hazard_mean[cause][[row, time]] += weight * hazard;
hazard_second[cause][[row, time]] += weight * hazard * hazard;
cumulative_hazard_mean[cause][[row, time]] += weight * cumulative_hazard;
cumulative_hazard_second[cause][[row, time]] +=
weight * cumulative_hazard * cumulative_hazard;
cif_mean[cause][[row, time]] += weight * cif;
cif_second[cause][[row, time]] += weight * cif * cif;
}
}
}
for row in 0..n_rows {
for time in 0..n_times {
let overall_survival = draw.overall_survival[[row, time]];
overall_mean[[row, time]] += weight * overall_survival;
overall_second[[row, time]] += weight * overall_survival * overall_survival;
}
}
Ok(())
})?;
let (hazard_se, survival_se, cumulative_hazard_se, cif_se, overall_survival_se, eta_se) =
if req.with_uncertainty {
(
Some(posterior_standard_error_surfaces(
&hazard_mean,
&hazard_second,
"competing-risks hazard",
)?),
Some(posterior_standard_error_surfaces(
&survival_mean,
&survival_second,
"competing-risks survival",
)?),
Some(posterior_standard_error_surfaces(
&cumulative_hazard_mean,
&cumulative_hazard_second,
"competing-risks cumulative hazard",
)?),
Some(posterior_standard_error_surfaces(
&cif_mean,
&cif_second,
"competing-risks cumulative incidence",
)?),
Some(posterior_standard_error_matrix(
&overall_mean,
&overall_second,
"competing-risks overall survival",
)?),
Some(posterior_standard_error_vectors(
&eta_mean,
&eta_second,
"competing-risks linear predictor",
)?),
)
} else {
(None, None, None, None, None, None)
};
if posterior_mean_estimand && !separate_conditional_point {
result.hazard = hazard_mean;
result.survival = survival_mean
.into_iter()
.map(|surface| surface.mapv(|value| value.clamp(0.0, 1.0)))
.collect();
result.cumulative_hazard = cumulative_hazard_mean;
result.cif = cif_mean
.into_iter()
.map(|surface| surface.mapv(|value| value.clamp(0.0, 1.0)))
.collect();
result.overall_survival = overall_mean.mapv(|value| value.clamp(0.0, 1.0));
result.linear_predictor = eta_mean;
}
result.hazard_se = hazard_se;
result.survival_se = survival_se;
result.cumulative_hazard_se = cumulative_hazard_se;
result.cif_se = cif_se;
result.overall_survival_se = overall_survival_se;
result.eta_se = eta_se;
result.covariance_source = req.with_uncertainty.then_some(covariance_mode);
Ok(result)
}
fn restricted_mean_survival_time_from_curve(
times: &[f64],
survival_row: ndarray::ArrayView1<'_, f64>,
tau: f64,
) -> Option<f64> {
if times.is_empty() || !(tau > 0.0) || !tau.is_finite() {
return None;
}
if times.len() != survival_row.len() {
return None;
}
let mut prev_t = 0.0_f64;
let mut prev_s = 1.0_f64;
let mut area = 0.0_f64;
for (idx, &t) in times.iter().enumerate() {
if !t.is_finite() || t < prev_t {
return None;
}
let s = survival_row[idx];
if !s.is_finite() {
return None;
}
if t >= tau {
let span = t - prev_t;
let s_tau = if span > 0.0 {
let w = (tau - prev_t) / span;
prev_s + w * (s - prev_s)
} else {
prev_s
};
area += 0.5 * (prev_s + s_tau) * (tau - prev_t);
return Some(area);
}
area += 0.5 * (prev_s + s) * (t - prev_t);
prev_t = t;
prev_s = s;
}
area += prev_s * (tau - prev_t);
Some(area)
}
impl SurvivalPredictResult {
pub fn restricted_mean_survival_time(&self, tau: f64) -> Option<Array1<f64>> {
let n = self.survival.nrows();
let mut out = Array1::<f64>::zeros(n);
for i in 0..n {
let rmst =
restricted_mean_survival_time_from_curve(&self.times, self.survival.row(i), tau)?;
out[i] = rmst;
}
Some(out)
}
}
impl CompetingRisksPredictResult {
pub fn restricted_mean_overall_survival_time(&self, tau: f64) -> Option<Array1<f64>> {
let n = self.overall_survival.nrows();
let mut out = Array1::<f64>::zeros(n);
for i in 0..n {
let rmst = restricted_mean_survival_time_from_curve(
&self.times,
self.overall_survival.row(i),
tau,
)?;
out[i] = rmst;
}
Some(out)
}
}
pub fn harrell_concordance(time: &[f64], event: &[f64], risk: &[f64]) -> Option<f64> {
let n = time.len();
if n != event.len() || n != risk.len() {
return None;
}
let mut comparable = 0.0_f64;
let mut concordant = 0.0_f64;
for i in 0..n {
for j in (i + 1)..n {
let (early, late) = if time[i] < time[j] {
(i, j)
} else if time[j] < time[i] {
(j, i)
} else {
if event[i] > 0.5 && event[j] > 0.5 {
comparable += 1.0;
concordant += 0.5;
}
continue;
};
if event[early] < 0.5 {
continue;
}
comparable += 1.0;
if risk[early] > risk[late] {
concordant += 1.0;
} else if risk[early] == risk[late] {
concordant += 0.5;
}
}
}
if comparable == 0.0 {
return None;
}
Some(concordant / comparable)
}
pub fn ipcw_brier_score(
s_pred: &[f64],
time: &[f64],
event: &[f64],
tau: f64,
g_cens: impl Fn(f64) -> f64,
) -> Option<f64> {
let n = s_pred.len();
if n != time.len() || n != event.len() {
return None;
}
let mut n_valid = 0.0_f64;
let mut acc = 0.0_f64;
for i in 0..n {
if !time[i].is_finite() || !event[i].is_finite() || time[i] <= 0.0 {
continue;
}
n_valid += 1.0;
let (target, weight) = if time[i] <= tau && event[i] > 0.5 {
let g = g_cens(time[i]);
if !(g > 0.0) {
continue;
}
(0.0, 1.0 / g)
} else if time[i] > tau {
let g = g_cens(tau);
if !(g > 0.0) {
continue;
}
(1.0, 1.0 / g)
} else {
continue;
};
let resid = target - s_pred[i];
acc += weight * resid * resid;
}
if n_valid == 0.0 {
return None;
}
Some(acc / n_valid)
}
pub fn integrated_ipcw_brier_score(
s_pred: ArrayView2<f64>,
time: &[f64],
event: &[f64],
grid: &[f64],
horizon: f64,
g_cens: impl Fn(f64) -> f64,
) -> Option<f64> {
let m = grid.len();
if m < 2 || s_pred.ncols() != m || s_pred.nrows() != time.len() {
return None;
}
if grid.windows(2).any(|pair| !(pair[1] > pair[0])) {
return None;
}
let mut pts: Vec<(f64, f64)> = Vec::with_capacity(m);
for k in 0..m {
if grid[k] > horizon {
break;
}
let col = s_pred.column(k);
let col_slice: Vec<f64> = col.to_vec();
if let Some(bs) = ipcw_brier_score(&col_slice, time, event, grid[k], &g_cens) {
pts.push((grid[k], bs));
}
}
if pts.len() < 2 {
return None;
}
let span = pts[pts.len() - 1].0 - pts[0].0;
if !(span > 0.0) {
return None;
}
let mut integral = 0.0_f64;
for w in pts.windows(2) {
integral += 0.5 * (w[1].1 + w[0].1) * (w[1].0 - w[0].0);
}
Some(integral / span)
}
#[derive(Clone, Debug, Default)]
pub struct KaplanMeier {
steps: Vec<(f64, f64)>,
}
impl KaplanMeier {
pub fn fit(time: &[f64], event: &[f64]) -> Self {
let mut rows: Vec<(f64, bool)> = time
.iter()
.zip(event.iter())
.filter_map(|(&t, &e)| {
(t.is_finite() && e.is_finite() && t > 0.0).then_some((t, e > 0.5))
})
.collect();
rows.sort_by(|a, b| a.0.total_cmp(&b.0));
let mut steps = Vec::new();
let mut at_risk = rows.len() as f64;
let mut survival = 1.0_f64;
let mut i = 0usize;
while i < rows.len() {
let t = rows[i].0;
let mut j = i;
let mut deaths = 0usize;
while j < rows.len() && rows[j].0 == t {
deaths += usize::from(rows[j].1);
j += 1;
}
if deaths > 0 && at_risk > 0.0 {
survival *= ((at_risk - deaths as f64) / at_risk).max(0.0);
steps.push((t, survival));
}
at_risk -= (j - i) as f64;
i = j;
}
Self { steps }
}
pub fn fit_censoring(time: &[f64], event: &[f64]) -> Self {
let flipped: Vec<f64> = event
.iter()
.map(|&e| if e > 0.5 { 0.0 } else { 1.0 })
.collect();
Self::fit(time, &flipped)
}
pub fn at(&self, t: f64) -> f64 {
let mut s = 1.0_f64;
for &(time, surv) in &self.steps {
if time <= t {
s = surv;
} else {
break;
}
}
s
}
}
pub struct CompetingRisksPredictResult {
pub times: Vec<f64>,
pub endpoint_names: Vec<String>,
pub hazard: Vec<Array2<f64>>,
pub survival: Vec<Array2<f64>>,
pub cumulative_hazard: Vec<Array2<f64>>,
pub cif: Vec<Array2<f64>>,
pub overall_survival: Array2<f64>,
pub linear_predictor: Vec<Array1<f64>>,
pub likelihood_mode: SurvivalLikelihoodMode,
pub covariance_source: Option<SurvivalPredictionCovarianceMode>,
pub hazard_se: Option<Vec<Array2<f64>>>,
pub survival_se: Option<Vec<Array2<f64>>>,
pub cumulative_hazard_se: Option<Vec<Array2<f64>>>,
pub cif_se: Option<Vec<Array2<f64>>>,
pub overall_survival_se: Option<Array2<f64>>,
pub eta_se: Option<Vec<Array1<f64>>>,
}
pub fn predict_survival(
req: SurvivalPredictRequest<'_>,
covariance_mode: SurvivalPredictionCovarianceMode,
) -> Result<SurvivalPredictResult, SurvivalPredictError> {
if req.estimand == SurvivalPredictEstimand::PosteriorMean {
return predict_survival_posterior_mean(req, covariance_mode);
}
let SurvivalPredictRequest {
model,
data,
col_map,
training_headers,
primary_offset,
noise_offset,
time_grid,
with_uncertainty,
estimand: _,
} = req;
let time_cols = resolve_saved_survival_time_columns(model, col_map)?;
let exit_col = time_cols.exit_col;
let termspec = resolve_termspec_for_prediction(
&model.resolved_termspec,
training_headers,
col_map,
"resolved_termspec",
)?;
let cov_clipped = model.axis_clip_to_training_ranges(data, col_map);
let cov_input = cov_clipped.as_ref().map_or(data, |arr| arr.view());
let cov_design = build_term_collection_design(cov_input, &termspec)
.map_err(|e| format!("failed to build survival prediction design: {e}"))?;
let n = data.nrows();
if primary_offset.len() != n || noise_offset.len() != n {
return Err(SurvivalPredictError::InvalidInput {
reason: format!(
"survival prediction offset length mismatch: rows={n}, offset={}, noise_offset={}",
primary_offset.len(),
noise_offset.len()
),
});
}
let effective_primary_offset = cov_design
.compose_offset(primary_offset.view(), "survival prediction covariate block")
.map_err(|error| error.to_string())?;
use rayon::iter::{IntoParallelIterator, ParallelIterator};
let pairs: Result<Vec<(f64, f64)>, String> = (0..n)
.into_par_iter()
.map(|i| {
normalize_survival_time_pair(time_cols.row_entry_time(data, i), data[[i, exit_col]], i)
})
.collect();
let pairs = pairs?;
let mut age_entry = Array1::<f64>::zeros(n);
let mut age_exit = Array1::<f64>::zeros(n);
for (i, (t0, t1)) in pairs.into_iter().enumerate() {
age_entry[i] = t0;
age_exit[i] = t1;
}
let saved_likelihood_mode = require_saved_survival_likelihood_mode(model)?;
if matches!(
saved_likelihood_mode,
SurvivalLikelihoodMode::Latent | SurvivalLikelihoodMode::LatentBinary
) {
return Err(SurvivalPredictError::UnsupportedConfiguration {
reason: format!(
"survival prediction via predict_survival does not support likelihood_mode={} yet; \
latent window prediction lives in the CLI's run_predict_saved_latent_window_impl \
pipeline and has not yet been ported to the library. Use the CLI predict command.",
survival_likelihood_modename(saved_likelihood_mode)
),
});
}
if saved_likelihood_mode == SurvivalLikelihoodMode::LocationScale {
return predict_survival_location_scale_batch(
model,
&age_entry,
&age_exit,
&cov_design,
&effective_primary_offset,
noise_offset,
training_headers,
col_map,
data,
time_grid,
with_uncertainty,
covariance_mode,
)
.map_err(SurvivalPredictError::from);
}
if with_uncertainty {
return Err(SurvivalPredictError::from(format!(
"predict_survival: with_uncertainty is currently supported only for the \
location-scale likelihood mode; got {}",
survival_likelihood_modename(saved_likelihood_mode)
)));
}
let time_cfg = load_survival_time_basis_config_from_model(model)?;
let mut time_build = build_survival_time_basis(&age_entry, &age_exit, time_cfg.clone(), None)?;
let resolved_time_cfg = resolved_survival_time_basis_config_from_build(
&time_build.basisname,
time_build.degree,
time_build.knots.as_ref(),
time_build.keep_cols.as_ref(),
time_build.smooth_lambda,
)?;
let weibull_baseline_in_beta = saved_likelihood_mode == SurvivalLikelihoodMode::Weibull
&& !model.has_baseline_time_wiggle();
let mut time_anchor: Option<f64> = None;
let mut time_anchor_row_cached: Option<Array1<f64>> = None;
if matches!(
saved_likelihood_mode,
SurvivalLikelihoodMode::LocationScale | SurvivalLikelihoodMode::MarginalSlope
) || weibull_baseline_in_beta
{
let anchor = model
.survival_time_anchor
.ok_or_else(|| "saved survival model missing survival_time_anchor".to_string())?;
let time_anchor_row = evaluate_survival_time_basis_row(anchor, &resolved_time_cfg)?;
center_survival_time_designs_at_anchor(
&mut time_build.x_entry_time,
&mut time_build.x_exit_time,
&time_anchor_row,
)?;
time_anchor = Some(anchor);
time_anchor_row_cached = Some(time_anchor_row);
}
if saved_likelihood_mode != SurvivalLikelihoodMode::Weibull && !model.has_baseline_time_wiggle()
{
require_structural_survival_time_basis(&time_build.basisname, "saved survival sampling")?;
}
let mut baseline_cfg = saved_survival_runtime_baseline_config(model)?;
if weibull_baseline_in_beta {
baseline_cfg = SurvivalBaselineConfig {
target: SurvivalBaselineTarget::Linear,
scale: None,
shape: None,
rate: None,
makeham: None,
};
}
let per_row_eval = time_grid.is_none();
let eval_times: Vec<f64> = match time_grid {
Some(grid) => {
if grid.is_empty() {
return Err(SurvivalPredictError::InvalidInput {
reason: "survival time_grid must contain at least one time".to_string(),
});
}
for (idx, &t) in grid.iter().enumerate() {
if !t.is_finite() || t < 0.0 {
return Err(SurvivalPredictError::InvalidInput {
reason: format!(
"survival time_grid requires finite non-negative times (index {idx})",
),
});
}
}
grid.to_vec()
}
None => Vec::new(),
};
let t_cols = if per_row_eval { 1 } else { eval_times.len() };
let mut hazard = Array2::<f64>::zeros((n, t_cols));
let mut survival = Array2::<f64>::zeros((n, t_cols));
let mut cumulative_hazard = Array2::<f64>::zeros((n, t_cols));
let mut linear_predictor = Array1::<f64>::zeros(n);
let marginal_slope_ctx = if saved_likelihood_mode == SurvivalLikelihoodMode::MarginalSlope {
let (mut eta_offset_entry, mut eta_offset_exit, mut derivative_offset_exit) =
build_survival_time_offsets_for_likelihood(
&age_entry,
&age_exit,
&baseline_cfg,
saved_likelihood_mode,
None,
)?;
add_survival_time_derivative_guard_offset(
&age_entry,
&age_exit,
time_anchor.ok_or_else(|| {
"saved survival marginal-slope model missing survival_time_anchor".to_string()
})?,
survival_derivative_guard_for_likelihood(saved_likelihood_mode),
&mut eta_offset_entry,
&mut eta_offset_exit,
&mut derivative_offset_exit,
)?;
Some(build_marginal_slope_predict_context(
model,
data,
col_map,
training_headers,
&cov_design.design,
&effective_primary_offset,
noise_offset,
&time_build,
&eta_offset_entry,
&eta_offset_exit,
&derivative_offset_exit,
)?)
} else {
None
};
struct SurvivalPredictionRow {
hazard: Vec<f64>,
survival: Vec<f64>,
cumulative_hazard: Vec<f64>,
linear_predictor: f64,
}
let row_results: Result<Vec<SurvivalPredictionRow>, SurvivalPredictError> = (0..n)
.into_par_iter()
.map(|i| {
let cov_row = if matches!(
saved_likelihood_mode,
SurvivalLikelihoodMode::Transformation | SurvivalLikelihoodMode::Weibull
) {
Some(design_row_owned(
&cov_design.design,
i,
"survival predict covariate row",
)?)
} else {
None
};
let evaluate_at = |t_query: f64| -> Result<(f64, f64, f64), SurvivalPredictError> {
let t_entry = age_entry[i].min(t_query);
let single_entry = Array1::from_elem(1, t_entry);
let single_exit = Array1::from_elem(1, t_query);
let mut row_time =
build_survival_time_basis(&single_entry, &single_exit, time_cfg.clone(), None)?;
if let Some(anchor_row) = time_anchor_row_cached.as_ref() {
center_survival_time_designs_at_anchor(
&mut row_time.x_entry_time,
&mut row_time.x_exit_time,
anchor_row,
)?;
}
let (mut r_eta_entry, mut r_eta_exit, mut r_deriv_exit) =
build_survival_time_offsets_for_likelihood(
&single_entry,
&single_exit,
&baseline_cfg,
saved_likelihood_mode,
None,
)?;
if saved_likelihood_mode == SurvivalLikelihoodMode::MarginalSlope {
add_survival_time_derivative_guard_offset(
&single_entry,
&single_exit,
time_anchor.ok_or_else(|| {
"saved survival marginal-slope model missing survival_time_anchor"
.to_string()
})?,
survival_derivative_guard_for_likelihood(saved_likelihood_mode),
&mut r_eta_entry,
&mut r_eta_exit,
&mut r_deriv_exit,
)?;
}
match saved_likelihood_mode {
SurvivalLikelihoodMode::MarginalSlope => {
let ctx = marginal_slope_ctx.as_ref().ok_or_else(|| {
"internal error: marginal-slope context missing for marginal-slope mode"
.to_string()
})?;
evaluate_marginal_slope_row(
i,
ctx,
&row_time,
&r_eta_exit,
&r_deriv_exit,
effective_primary_offset[i],
)
}
SurvivalLikelihoodMode::Transformation | SurvivalLikelihoodMode::Weibull => {
let cov_row = cov_row.as_ref().ok_or_else(|| {
"internal error: covariate row missing for Royston-Parmar prediction"
.to_string()
})?;
evaluate_rp_row(
model,
&row_time,
cov_row,
r_eta_exit[0],
r_deriv_exit[0],
effective_primary_offset[i],
)
}
SurvivalLikelihoodMode::Latent
| SurvivalLikelihoodMode::LatentBinary
| SurvivalLikelihoodMode::LocationScale => {
Err(SurvivalPredictError::NumericalFailure {
reason: "unreachable: unsupported likelihood_mode filtered earlier"
.to_string(),
})
}
}
};
let mut row = SurvivalPredictionRow {
hazard: vec![0.0; t_cols],
survival: vec![0.0; t_cols],
cumulative_hazard: vec![0.0; t_cols],
linear_predictor: 0.0,
};
if per_row_eval {
let (eta_t, cum_t, haz_t) = evaluate_at(age_exit[i])?;
row.linear_predictor = eta_t;
row.hazard[0] = haz_t;
row.cumulative_hazard[0] = cum_t;
row.survival[0] = (-cum_t).exp().clamp(0.0, 1.0);
} else {
for (j, &t_query) in eval_times.iter().enumerate() {
if t_query <= 0.0 {
row.hazard[j] = 0.0;
row.cumulative_hazard[j] = 0.0;
row.survival[j] = 1.0;
} else {
let (_eta_t, cum_t, haz_t) = evaluate_at(t_query)?;
row.hazard[j] = haz_t;
row.cumulative_hazard[j] = cum_t;
row.survival[j] = (-cum_t).exp().clamp(0.0, 1.0);
}
}
let (eta_t, _, _) = evaluate_at(age_exit[i])?;
row.linear_predictor = eta_t;
}
Ok(row)
})
.collect();
for (i, row) in row_results?.into_iter().enumerate() {
linear_predictor[i] = row.linear_predictor;
for j in 0..t_cols {
hazard[[i, j]] = row.hazard[j];
cumulative_hazard[[i, j]] = row.cumulative_hazard[j];
survival[[i, j]] = row.survival[j];
}
}
let times_out: Vec<f64> = if per_row_eval {
age_exit.to_vec()
} else {
eval_times
};
Ok(SurvivalPredictResult {
times: times_out,
hazard,
survival,
cumulative_hazard,
linear_predictor,
likelihood_mode: saved_likelihood_mode,
survival_se: None,
eta_se: None,
covariance_source: None,
})
}
pub fn predict_competing_risks_survival(
req: SurvivalPredictRequest<'_>,
covariance_mode: SurvivalPredictionCovarianceMode,
) -> Result<CompetingRisksPredictResult, SurvivalPredictError> {
if req.estimand == SurvivalPredictEstimand::PosteriorMean || req.with_uncertainty {
return predict_competing_risks_with_posterior(req, covariance_mode);
}
let SurvivalPredictRequest {
model,
data,
col_map,
training_headers,
primary_offset,
noise_offset,
time_grid,
with_uncertainty: _,
estimand: _,
} = req;
let saved_likelihood_mode = require_saved_survival_likelihood_mode(model)?;
if !matches!(
saved_likelihood_mode,
SurvivalLikelihoodMode::Transformation | SurvivalLikelihoodMode::Weibull
) {
return Err(SurvivalPredictError::UnsupportedConfiguration {
reason: format!(
"joint cause-specific prediction supports transformation/weibull survival only; got {}",
survival_likelihood_modename(saved_likelihood_mode)
),
});
}
let fit = fit_result_from_saved_model_for_prediction(model)?;
let cause_count = model
.survival_cause_count
.unwrap_or(fit.blocks.len())
.max(1);
if cause_count <= 1 {
return Err(SurvivalPredictError::MissingFitMetadata {
reason: "competing-risks survival prediction requires a saved model with at least two causes"
.to_string(),
});
}
if fit.blocks.len() != cause_count {
return Err(SurvivalPredictError::IncompatibleSchema {
reason: format!(
"saved competing-risks survival fit has {} coefficient blocks but metadata says {cause_count} causes",
fit.blocks.len()
),
});
}
let endpoint_names = model.survival_endpoint_names.clone().unwrap_or_else(|| {
(1..=cause_count)
.map(|idx| format!("cause_{idx}"))
.collect()
});
if endpoint_names.len() != cause_count {
return Err(SurvivalPredictError::IncompatibleSchema {
reason: format!(
"saved competing-risks survival endpoint_names has length {}, expected {cause_count}",
endpoint_names.len()
),
});
}
let time_cols = resolve_saved_survival_time_columns(model, col_map)?;
let exit_col = time_cols.exit_col;
let termspec = resolve_termspec_for_prediction(
&model.resolved_termspec,
training_headers,
col_map,
"resolved_termspec",
)?;
let cov_clipped = model.axis_clip_to_training_ranges(data, col_map);
let cov_input = cov_clipped.as_ref().map_or(data, |arr| arr.view());
let cov_design = build_term_collection_design(cov_input, &termspec)
.map_err(|e| format!("failed to build competing-risks prediction design: {e}"))?;
let n = data.nrows();
if primary_offset.len() != n || noise_offset.len() != n {
return Err(SurvivalPredictError::InvalidInput {
reason: format!(
"competing-risks prediction offset length mismatch: rows={n}, offset={}, noise_offset={}",
primary_offset.len(),
noise_offset.len()
),
});
}
let effective_primary_offset = cov_design
.compose_offset(
primary_offset.view(),
"competing-risks prediction covariate block",
)
.map_err(|error| error.to_string())?;
use rayon::iter::{IntoParallelIterator, ParallelIterator};
let pairs: Result<Vec<(f64, f64)>, String> = (0..n)
.into_par_iter()
.map(|i| {
normalize_survival_time_pair(time_cols.row_entry_time(data, i), data[[i, exit_col]], i)
})
.collect();
let pairs = pairs?;
let mut age_entry = Array1::<f64>::zeros(n);
let mut age_exit = Array1::<f64>::zeros(n);
for (i, (t0, t1)) in pairs.into_iter().enumerate() {
age_entry[i] = t0;
age_exit[i] = t1;
}
let time_cfg = load_survival_time_basis_config_from_model(model)?;
let time_build = build_survival_time_basis(&age_entry, &age_exit, time_cfg.clone(), None)?;
let resolved_time_cfg = resolved_survival_time_basis_config_from_build(
&time_build.basisname,
time_build.degree,
time_build.knots.as_ref(),
time_build.keep_cols.as_ref(),
time_build.smooth_lambda,
)?;
let weibull_baseline_in_beta = saved_likelihood_mode == SurvivalLikelihoodMode::Weibull
&& !model.has_baseline_time_wiggle();
let cr_time_anchor_row: Option<Array1<f64>> = if weibull_baseline_in_beta {
let anchor = model
.survival_time_anchor
.ok_or_else(|| "saved survival model missing survival_time_anchor".to_string())?;
Some(evaluate_survival_time_basis_row(
anchor,
&resolved_time_cfg,
)?)
} else {
None
};
if saved_likelihood_mode != SurvivalLikelihoodMode::Weibull && !model.has_baseline_time_wiggle()
{
require_structural_survival_time_basis(
&time_build.basisname,
"saved competing-risks survival prediction",
)?;
}
let baseline_cfg = saved_survival_runtime_baseline_config(model)?;
let per_row_eval = time_grid.is_none();
let eval_times: Vec<f64> = match time_grid {
Some(grid) => {
if grid.is_empty() {
return Err(SurvivalPredictError::InvalidInput {
reason: "survival time_grid must contain at least one time".to_string(),
});
}
for (idx, &t) in grid.iter().enumerate() {
if !t.is_finite() || t < 0.0 {
return Err(SurvivalPredictError::InvalidInput {
reason: format!(
"survival time_grid requires finite non-negative times (index {idx})",
),
});
}
}
grid.to_vec()
}
None => Vec::new(),
};
let t_cols = if per_row_eval { 1 } else { eval_times.len() };
const CIF_REFINE_SUBINTERVALS: usize = 32;
let (refined_times, user_time_to_refined_index): (Vec<f64>, Vec<usize>) = if per_row_eval {
(Vec::new(), Vec::new())
} else {
let mut order: Vec<usize> = (0..eval_times.len()).collect();
order.sort_by(|&a, &b| {
eval_times[a]
.partial_cmp(&eval_times[b])
.expect("survival time_grid entries are validated finite above")
});
let mut refined: Vec<f64> = Vec::new();
let mut user_index: Vec<usize> = vec![0; eval_times.len()];
let mut prev = 0.0_f64;
for &j_user in &order {
let t_user = eval_times[j_user];
let gap = t_user - prev;
if gap > 0.0 {
for s in 1..CIF_REFINE_SUBINTERVALS {
let t_mid = prev + gap * (s as f64) / (CIF_REFINE_SUBINTERVALS as f64);
if refined.last().is_none_or(|&last| t_mid > last) {
refined.push(t_mid);
}
}
}
if refined.last().is_none_or(|&last| t_user > last) {
refined.push(t_user);
}
user_index[j_user] = refined.len() - 1;
prev = t_user;
}
(refined, user_index)
};
let refined_cols = if per_row_eval {
CIF_REFINE_SUBINTERVALS
} else {
refined_times.len()
};
let saved_timewiggle_by_cause = saved_cause_specific_timewiggles(model, &fit, cause_count)?;
let cov_rows = (0..n)
.map(|i| design_row_owned(&cov_design.design, i, "competing-risks covariate row"))
.collect::<Result<Vec<_>, _>>()?;
let mut hazard = (0..cause_count)
.map(|_| Array2::<f64>::zeros((n, t_cols)))
.collect::<Vec<_>>();
let mut survival = (0..cause_count)
.map(|_| Array2::<f64>::zeros((n, t_cols)))
.collect::<Vec<_>>();
let mut cumulative_hazard = (0..cause_count)
.map(|_| Array2::<f64>::zeros((n, t_cols)))
.collect::<Vec<_>>();
let mut cumulative_hazard_refined = (0..cause_count)
.map(|_| Array2::<f64>::zeros((n, refined_cols)))
.collect::<Vec<_>>();
let mut linear_predictor = (0..cause_count)
.map(|_| Array1::<f64>::zeros(n))
.collect::<Vec<_>>();
struct CauseRow {
cause: usize,
row: usize,
hazard: Vec<f64>,
survival: Vec<f64>,
cumulative: Vec<f64>,
cumulative_refined: Vec<f64>,
eta_exit: f64,
}
let rows: Result<Vec<CauseRow>, SurvivalPredictError> = (0..cause_count * n)
.into_par_iter()
.map(|flat| {
let cause = flat / n;
let i = flat % n;
let block = &fit.blocks[cause];
let timewiggle = saved_timewiggle_by_cause[cause].as_ref();
let evaluate_at = |t_query: f64| -> Result<(f64, f64, f64), SurvivalPredictError> {
let t_entry = age_entry[i].min(t_query);
let single_entry = Array1::from_elem(1, t_entry);
let single_exit = Array1::from_elem(1, t_query);
let mut row_time =
build_survival_time_basis(&single_entry, &single_exit, time_cfg.clone(), None)?;
if let Some(anchor_row) = cr_time_anchor_row.as_ref() {
center_survival_time_designs_at_anchor(
&mut row_time.x_entry_time,
&mut row_time.x_exit_time,
anchor_row,
)?;
}
let (r_eta_exit, r_deriv_exit) = if weibull_baseline_in_beta {
(0.0, 0.0)
} else {
let (_, eta_exit, deriv_exit) = build_survival_time_offsets_for_likelihood(
&single_entry,
&single_exit,
&baseline_cfg,
saved_likelihood_mode,
None,
)?;
(eta_exit[0], deriv_exit[0])
};
evaluate_rp_row_with_beta(
&block.beta,
timewiggle,
&row_time,
&cov_rows[i],
r_eta_exit,
r_deriv_exit,
effective_primary_offset[i],
)
};
let mut out = CauseRow {
cause,
row: i,
hazard: vec![0.0; t_cols],
survival: vec![0.0; t_cols],
cumulative: vec![0.0; t_cols],
cumulative_refined: vec![0.0; refined_cols],
eta_exit: 0.0,
};
if per_row_eval {
let (eta_t, cum_t, haz_t) = evaluate_at(age_exit[i])?;
out.eta_exit = eta_t;
out.hazard[0] = haz_t;
out.cumulative[0] = cum_t;
out.survival[0] = (-cum_t).exp().clamp(0.0, 1.0);
for s in 1..=CIF_REFINE_SUBINTERVALS {
let frac = (s as f64) / (CIF_REFINE_SUBINTERVALS as f64);
let t_query = age_exit[i] * frac;
out.cumulative_refined[s - 1] = if t_query <= 0.0 {
0.0
} else if s == CIF_REFINE_SUBINTERVALS {
cum_t
} else {
evaluate_at(t_query)?.1
};
}
} else {
for (j, &t_query) in eval_times.iter().enumerate() {
if t_query <= 0.0 {
out.hazard[j] = 0.0;
out.cumulative[j] = 0.0;
out.survival[j] = 1.0;
} else {
let (_eta_t, cum_t, haz_t) = evaluate_at(t_query)?;
out.hazard[j] = haz_t;
out.cumulative[j] = cum_t;
out.survival[j] = (-cum_t).exp().clamp(0.0, 1.0);
}
}
for (jr, &t_query) in refined_times.iter().enumerate() {
out.cumulative_refined[jr] = if t_query <= 0.0 {
0.0
} else {
evaluate_at(t_query)?.1
};
}
let (eta_t, _, _) = evaluate_at(age_exit[i])?;
out.eta_exit = eta_t;
}
Ok(out)
})
.collect();
for row in rows? {
linear_predictor[row.cause][row.row] = row.eta_exit;
for j in 0..t_cols {
hazard[row.cause][[row.row, j]] = row.hazard[j];
survival[row.cause][[row.row, j]] = row.survival[j];
cumulative_hazard[row.cause][[row.row, j]] = row.cumulative[j];
}
for jr in 0..refined_cols {
cumulative_hazard_refined[row.cause][[row.row, jr]] = row.cumulative_refined[jr];
}
}
let assembled = if per_row_eval {
let assembly_times = Array1::from_shape_fn(CIF_REFINE_SUBINTERVALS, |s| {
((s + 1) as f64) / (CIF_REFINE_SUBINTERVALS as f64)
});
let refined_assembled = assemble_competing_risks_cif_from_endpoints(
assembly_times.view(),
&cumulative_hazard_refined,
)
.map_err(|err| err.to_string())?;
let last = CIF_REFINE_SUBINTERVALS - 1;
let mut cif_user = (0..cause_count)
.map(|_| Array2::<f64>::zeros((n, 1)))
.collect::<Vec<_>>();
let mut overall_user = Array2::<f64>::zeros((n, 1));
for cause in 0..cause_count {
for row in 0..n {
cif_user[cause][[row, 0]] = refined_assembled.cif[cause][[row, last]];
}
}
for row in 0..n {
overall_user[[row, 0]] = refined_assembled.overall_survival[[row, last]];
}
CompetingRisksCifResult {
cif: cif_user,
overall_survival: overall_user,
}
} else {
let assembly_times = Array1::from_vec(refined_times.clone());
let refined_assembled = assemble_competing_risks_cif_from_endpoints(
assembly_times.view(),
&cumulative_hazard_refined,
)
.map_err(|err| err.to_string())?;
let mut cif_user = (0..cause_count)
.map(|_| Array2::<f64>::zeros((n, t_cols)))
.collect::<Vec<_>>();
let mut overall_user = Array2::<f64>::zeros((n, t_cols));
for (j_user, &jr) in user_time_to_refined_index.iter().enumerate() {
for cause in 0..cause_count {
for row in 0..n {
cif_user[cause][[row, j_user]] = refined_assembled.cif[cause][[row, jr]];
}
}
for row in 0..n {
overall_user[[row, j_user]] = refined_assembled.overall_survival[[row, jr]];
}
}
CompetingRisksCifResult {
cif: cif_user,
overall_survival: overall_user,
}
};
if assembled.cif.len() != cause_count {
return Err(format!(
"competing-risks CIF assembly produced {} endpoint matrices, expected {cause_count}",
assembled.cif.len()
)
.into());
}
let cif = assembled.cif;
let overall_survival = assembled.overall_survival;
let times_out = if per_row_eval {
age_exit.to_vec()
} else {
eval_times
};
Ok(CompetingRisksPredictResult {
times: times_out,
endpoint_names,
hazard,
survival,
cumulative_hazard,
cif,
overall_survival,
linear_predictor,
likelihood_mode: saved_likelihood_mode,
covariance_source: None,
hazard_se: None,
survival_se: None,
cumulative_hazard_se: None,
cif_se: None,
overall_survival_se: None,
eta_se: None,
})
}
fn saved_cause_specific_timewiggles(
model: &SavedModel,
fit: &UnifiedFitResult,
cause_count: usize,
) -> Result<Vec<Option<SavedBaselineTimeWiggleRuntime>>, SurvivalPredictError> {
let has_metadata = model.baseline_timewiggle_knots.is_some()
|| model.baseline_timewiggle_degree.is_some()
|| model.baseline_timewiggle_penalty_orders.is_some()
|| model.baseline_timewiggle_double_penalty.is_some()
|| model.beta_baseline_timewiggle_by_cause.is_some();
if !has_metadata {
return Ok(vec![None; cause_count]);
}
let knots = model.baseline_timewiggle_knots.clone().ok_or_else(|| {
"joint cause-specific survival missing baseline_timewiggle_knots".to_string()
})?;
let degree = model.baseline_timewiggle_degree.ok_or_else(|| {
"joint cause-specific survival missing baseline_timewiggle_degree".to_string()
})?;
let penalty_orders = model
.baseline_timewiggle_penalty_orders
.clone()
.ok_or_else(|| {
"joint cause-specific survival missing baseline_timewiggle_penalty_orders".to_string()
})?;
let double_penalty = model.baseline_timewiggle_double_penalty.ok_or_else(|| {
"joint cause-specific survival missing baseline_timewiggle_double_penalty".to_string()
})?;
let by_cause = model
.beta_baseline_timewiggle_by_cause
.as_ref()
.ok_or_else(|| {
"joint cause-specific survival missing beta_baseline_timewiggle_by_cause".to_string()
})?;
if by_cause.len() != cause_count {
return Err(SurvivalPredictError::IncompatibleSchema {
reason: format!(
"joint cause-specific survival has {} timewiggle coefficient blocks, expected {cause_count}",
by_cause.len()
),
});
}
for (cause, (block, beta_w)) in fit.blocks.iter().zip(by_cause).enumerate() {
if beta_w.len() > block.beta.len() {
return Err(SurvivalPredictError::IncompatibleSchema {
reason: format!(
"joint cause-specific survival cause {} timewiggle beta has length {}, but endpoint beta has {} coefficients",
cause + 1,
beta_w.len(),
block.beta.len()
),
});
}
}
Ok(by_cause
.iter()
.map(|beta| {
Some(SavedBaselineTimeWiggleRuntime {
knots: knots.clone(),
degree,
penalty_orders: penalty_orders.clone(),
double_penalty,
beta: beta.clone(),
})
})
.collect())
}
struct MarginalSlopePredictContext {
predictor: BernoulliMarginalSlopePredictor,
beta_time: Array1<f64>,
beta_marginal: Array1<f64>,
saved_timewiggle: Option<SavedBaselineTimeWiggleRuntime>,
cov_design: DesignMatrix,
logslope_design: DesignMatrix,
cov_eta: Array1<f64>,
z_raw: Array1<f64>,
noise_offset: Array1<f64>,
}
fn design_row_owned(
design: &DesignMatrix,
row: usize,
context: &str,
) -> Result<Array1<f64>, SurvivalPredictError> {
let chunk = design
.try_row_chunk(row..row + 1)
.map_err(|e| format!("{context}: {e}"))?;
Ok(chunk.row(0).to_owned())
}
fn build_marginal_slope_predict_context(
model: &SavedModel,
data: ArrayView2<'_, f64>,
col_map: &HashMap<String, usize>,
training_headers: Option<&Vec<String>>,
cov_design: &DesignMatrix,
primary_offset: &Array1<f64>,
noise_offset: &Array1<f64>,
time_build: &SurvivalTimeBuildOutput,
eta_offset_entry: &Array1<f64>,
eta_offset_exit: &Array1<f64>,
derivative_offset_exit: &Array1<f64>,
) -> Result<MarginalSlopePredictContext, SurvivalPredictError> {
let z_name = model
.z_column
.as_ref()
.ok_or_else(|| "saved survival marginal-slope model missing z_column".to_string())?;
let z_col = resolve_role_col(col_map, z_name, "z")?;
let z_raw = data.column(z_col).to_owned();
let logslopespec = resolve_termspec_for_prediction(
&model.resolved_termspec_logslope.as_ref().cloned(),
training_headers,
col_map,
"resolved_termspec_logslope",
)?;
let logslope_clipped = model.axis_clip_to_training_ranges(data, col_map);
let logslope_input = logslope_clipped.as_ref().map_or(data, |arr| arr.view());
let logslope_design = build_term_collection_design(logslope_input, &logslopespec)
.map_err(|e| format!("failed to build survival marginal-slope logslope design: {e}"))?;
let effective_noise_offset = logslope_design
.compose_offset(
noise_offset.view(),
"survival marginal-slope logslope block",
)
.map_err(|error| error.to_string())?;
let fit_saved = fit_result_from_saved_model_for_prediction(model)?;
let (predictor, _pred_input, _predictor_fit) = build_saved_survival_marginal_slope_predictor(
model,
&fit_saved,
z_name,
&z_raw,
cov_design,
&logslope_design.design,
time_build,
eta_offset_entry,
eta_offset_exit,
derivative_offset_exit,
primary_offset,
&effective_noise_offset,
)?;
let blocks = &fit_saved.blocks;
if blocks.len() < 3 {
return Err(SurvivalPredictError::IncompatibleSchema {
reason: format!(
"saved survival marginal-slope model requires at least 3 blocks [time, marginal, slope], got {}",
blocks.len()
),
});
}
let beta_time = blocks[0].beta.clone();
let beta_marginal = blocks[1].beta.clone();
let saved_runtime = model.saved_prediction_runtime()?;
let saved_timewiggle = saved_runtime.baseline_time_wiggle.clone();
let cov_eta = cov_design.dot(&beta_marginal);
Ok(MarginalSlopePredictContext {
predictor,
beta_time,
beta_marginal,
saved_timewiggle,
cov_design: cov_design.clone(),
logslope_design: logslope_design.design.clone(),
cov_eta,
z_raw,
noise_offset: effective_noise_offset,
})
}
fn evaluate_marginal_slope_row(
row_index: usize,
ctx: &MarginalSlopePredictContext,
row_time: &SurvivalTimeBuildOutput,
r_eta_exit: &Array1<f64>,
r_deriv_exit: &Array1<f64>,
primary_offset_row: f64,
) -> Result<(f64, f64, f64), SurvivalPredictError> {
let beta_time = &ctx.beta_time;
let p_time_base = row_time.x_exit_time.ncols();
let p_timewiggle = ctx
.saved_timewiggle
.as_ref()
.map_or(0, |runtime| runtime.beta.len());
if beta_time.len() != p_time_base + p_timewiggle {
let hint = stale_weibull_time_basis_hint(
&row_time.basisname,
beta_time.len() == p_time_base + p_timewiggle + 1,
);
return Err(SurvivalPredictError::IncompatibleSchema {
reason: format!(
"saved survival marginal-slope time coefficient mismatch: beta has {} entries but expected base={} plus timewiggle={}{hint}",
beta_time.len(),
p_time_base,
p_timewiggle
),
});
}
let beta_time_base = beta_time.slice(s![..p_time_base]).to_owned();
let q_exit_base = row_time.x_exit_time.dot(&beta_time_base)[0]
+ ctx.cov_eta[row_index]
+ r_eta_exit[0]
+ primary_offset_row;
let qd_exit_base = row_time.x_derivative_time.dot(&beta_time_base)[0] + r_deriv_exit[0];
let (qd_with_wiggle, exit_wiggle_design) = if let Some(runtime) = ctx.saved_timewiggle.as_ref()
{
let knots = Array1::from_vec(runtime.knots.clone());
let beta_w = beta_time.slice(s![p_time_base..]).to_owned();
let eta_exit_row = Array1::from_elem(1, q_exit_base);
let deriv_row = Array1::from_elem(1, qd_exit_base);
let exit_design = match buildwiggle_block_input_from_knots(
eta_exit_row.view(),
&knots,
runtime.degree,
2,
false,
)?
.design
{
DesignMatrix::Dense(m) => m.to_dense_arc().as_ref().clone(),
_ => {
return Err(SurvivalPredictError::IncompatibleSchema {
reason: "saved baseline-timewiggle exit design must be dense".to_string(),
});
}
};
let derivative_design = build_survival_timewiggle_derivative_design(
&eta_exit_row,
&deriv_row,
&knots,
runtime.degree,
)?;
(
qd_exit_base + derivative_design.dot(&beta_w)[0],
Some(exit_design),
)
} else {
(qd_exit_base, None)
};
let cov_dim = ctx.beta_marginal.len();
let q_design_ncols = p_time_base + p_timewiggle + cov_dim;
let mut q_design_full = Array2::<f64>::zeros((1, q_design_ncols));
q_design_full
.slice_mut(s![.., ..p_time_base])
.assign(&row_time.x_exit_time.to_dense());
if let Some(exit_w) = exit_wiggle_design.as_ref() {
q_design_full
.slice_mut(s![.., p_time_base..p_time_base + p_timewiggle])
.assign(exit_w);
}
if cov_dim > 0 {
let cov_row = design_row_owned(
&ctx.cov_design,
row_index,
"survival marginal covariate row",
)?;
q_design_full
.slice_mut(s![.., p_time_base + p_timewiggle..])
.row_mut(0)
.assign(&cov_row);
}
let logslope_row = design_row_owned(
&ctx.logslope_design,
row_index,
"survival marginal logslope row",
)?;
let mut logslope_design_2d = Array2::<f64>::zeros((1, logslope_row.len()));
logslope_design_2d.row_mut(0).assign(&logslope_row);
let pred_input = PredictInput {
design: DesignMatrix::Dense(gam_linalg::matrix::DenseDesignMatrix::from(q_design_full)),
offset: Array1::from_elem(1, r_eta_exit[0] + primary_offset_row),
design_noise: Some(DesignMatrix::Dense(
gam_linalg::matrix::DenseDesignMatrix::from(logslope_design_2d),
)),
offset_noise: Some(Array1::from_elem(1, ctx.noise_offset[row_index])),
auxiliary_scalar: Some(Array1::from_elem(1, ctx.z_raw[row_index])),
auxiliary_matrix: None,
};
let (eta_arr, deta_dq_arr) = ctx
.predictor
.predict_eta_and_q_chain(&pred_input)
.map_err(|e| format!("saved survival marginal-slope predictor eta failed: {e}"))?;
let eta = eta_arr[0];
let eta_derivative = marginal_slope_index_derivative_at_horizon(deta_dq_arr[0], qd_with_wiggle);
let (cum, haz) = probit_survival_hazard_components(eta, eta_derivative)?;
Ok((eta, cum, haz))
}
#[inline]
fn marginal_slope_index_derivative_at_horizon(deta_dq: f64, qd_with_wiggle: f64) -> f64 {
let eta_derivative = deta_dq * qd_with_wiggle;
if eta_derivative.is_finite() {
eta_derivative.max(0.0)
} else {
eta_derivative
}
}
#[inline]
fn probit_survival_hazard_components(
eta: f64,
eta_derivative: f64,
) -> Result<(f64, f64), SurvivalPredictError> {
if !(eta.is_finite() && eta_derivative.is_finite() && eta_derivative >= 0.0) {
return Err(SurvivalPredictError::NumericalFailure {
reason: format!(
"saved survival marginal-slope prediction produced invalid survival index derivative: eta={eta}, eta_t={eta_derivative}"
),
});
}
let (log_survival, mills_ratio) = signed_probit_logcdf_and_mills_ratio(-eta);
let cumulative_hazard = -log_survival;
let hazard = if eta_derivative == 0.0 {
0.0
} else {
mills_ratio * eta_derivative
};
if !(cumulative_hazard >= 0.0 && hazard >= 0.0) {
return Err(SurvivalPredictError::NumericalFailure {
reason: format!(
"saved survival marginal-slope prediction produced invalid survival components: eta={eta}, eta_t={eta_derivative}, log_survival={log_survival}, hazard={hazard}"
),
});
}
Ok((cumulative_hazard, hazard))
}
fn evaluate_rp_row(
model: &SavedModel,
row_time: &SurvivalTimeBuildOutput,
cov_row: &Array1<f64>,
eta_time_offset_row: f64,
derivative_time_offset_row: f64,
primary_offset_row: f64,
) -> Result<(f64, f64, f64), SurvivalPredictError> {
let fit_saved = fit_result_from_saved_model_for_prediction(model)?;
let saved_runtime = model.saved_prediction_runtime()?;
evaluate_rp_row_with_beta(
&fit_saved.beta,
saved_runtime.baseline_time_wiggle.as_ref(),
row_time,
cov_row,
eta_time_offset_row,
derivative_time_offset_row,
primary_offset_row,
)
}
fn evaluate_rp_row_with_beta(
beta: &Array1<f64>,
saved_timewiggle: Option<&SavedBaselineTimeWiggleRuntime>,
row_time: &SurvivalTimeBuildOutput,
cov_row: &Array1<f64>,
eta_time_offset_row: f64,
derivative_time_offset_row: f64,
primary_offset_row: f64,
) -> Result<(f64, f64, f64), SurvivalPredictError> {
let p_time = row_time.x_exit_time.ncols();
let p_timewiggle = saved_timewiggle.map_or(0, |runtime| runtime.beta.len());
let p_cov = cov_row.len();
let p = p_time + p_timewiggle + p_cov;
if beta.len() != p {
let hint = stale_weibull_time_basis_hint(&row_time.basisname, beta.len() == p + 1);
return Err(SurvivalPredictError::IncompatibleSchema {
reason: format!(
"survival RP coefficient mismatch: beta has {} entries but design has {} columns{hint}",
beta.len(),
p
),
});
}
let mut x_exit = Array2::<f64>::zeros((1, p));
if p_time > 0 {
x_exit
.slice_mut(s![.., ..p_time])
.assign(&row_time.x_exit_time.to_dense());
}
let offset_derivative_component = derivative_time_offset_row;
let mut eta_derivative = offset_derivative_component;
let mut time_derivative_component = 0.0_f64;
if p_time > 0 {
time_derivative_component = row_time
.x_derivative_time
.dot(&beta.slice(s![..p_time]).to_owned())[0];
eta_derivative += time_derivative_component;
}
let mut wiggle_derivative_component = 0.0_f64;
if let Some(runtime) = saved_timewiggle {
let knots = Array1::from_vec(runtime.knots.clone());
let beta_w = beta.slice(s![p_time..p_time + p_timewiggle]).to_owned();
let eta_exit_row = Array1::from_elem(1, eta_time_offset_row);
let derivative_exit_row = Array1::from_elem(1, derivative_time_offset_row);
let exit_design = match buildwiggle_block_input_from_knots(
eta_exit_row.view(),
&knots,
runtime.degree,
2,
false,
)?
.design
{
DesignMatrix::Dense(m) => m.to_dense_arc().as_ref().clone(),
_ => {
return Err(SurvivalPredictError::IncompatibleSchema {
reason: "saved baseline-timewiggle exit design must be dense".to_string(),
});
}
};
if exit_design.ncols() != p_timewiggle {
return Err(SurvivalPredictError::IncompatibleSchema {
reason: format!(
"survival RP timewiggle design mismatch: rebuilt {} columns but runtime expects {}",
exit_design.ncols(),
p_timewiggle
),
});
}
x_exit
.slice_mut(s![.., p_time..p_time + p_timewiggle])
.assign(&exit_design);
let derivative_design = build_survival_timewiggle_derivative_design(
&eta_exit_row,
&derivative_exit_row,
&knots,
runtime.degree,
)?;
wiggle_derivative_component = derivative_design.dot(&beta_w)[0];
eta_derivative += wiggle_derivative_component;
}
if !(eta_derivative.is_finite() && eta_derivative >= 0.0) {
let time_beta = beta.slice(s![..p_time]);
let beta_min = time_beta.iter().copied().fold(f64::INFINITY, f64::min);
let beta_max = time_beta.iter().copied().fold(f64::NEG_INFINITY, f64::max);
let dtime = row_time.x_derivative_time.to_dense();
let dmin = dtime.iter().copied().fold(f64::INFINITY, f64::min);
let dmax = dtime.iter().copied().fold(f64::NEG_INFINITY, f64::max);
log::info!(
"[rp-predict/eta_t-refusal] eta_t={eta_derivative:.12e} = offset({offset_derivative_component:.12e}) + time({time_derivative_component:.12e}) + wiggle({wiggle_derivative_component:.12e}); p_time={p_time} p_timewiggle={p_timewiggle} p_cov={p_cov} time_beta=[{beta_min:.6e},{beta_max:.6e}] x_derivative_time=[{dmin:.6e},{dmax:.6e}] has_wiggle={}",
saved_timewiggle.is_some(),
);
}
if p_cov > 0 {
x_exit
.slice_mut(s![
..,
(p_time + p_timewiggle)..(p_time + p_timewiggle + p_cov)
])
.row_mut(0)
.assign(cov_row);
}
let offset_view = Array1::from_elem(1, eta_time_offset_row + primary_offset_row);
let likelihood = LikelihoodSpec::new(
ResponseFamily::RoystonParmar,
InverseLink::Standard(StandardLink::Identity),
);
let eta =
predict_royston_parmar_eta(x_exit.view(), beta.view(), offset_view.view(), &likelihood)?[0];
let (cum, haz) = royston_parmar_survival_hazard_components(eta, eta_derivative)?;
Ok((eta, cum, haz))
}
fn predict_royston_parmar_eta<X>(
x: X,
beta: ndarray::ArrayView1<'_, f64>,
offset: ndarray::ArrayView1<'_, f64>,
likelihood: &LikelihoodSpec,
) -> Result<Array1<f64>, SurvivalPredictError>
where
X: Into<DesignMatrix>,
{
if !matches!(likelihood.response, ResponseFamily::RoystonParmar)
|| !matches!(
likelihood.link,
InverseLink::Standard(StandardLink::Identity)
)
{
return Err(SurvivalPredictError::UnsupportedConfiguration {
reason: "survival prediction requires RoystonParmar with identity link".to_string(),
});
}
let x = x.into();
if x.nrows() != offset.len() || x.ncols() != beta.len() {
return Err(SurvivalPredictError::IncompatibleSchema {
reason: format!(
"survival prediction design dimensions disagree: design is {}x{}, beta has length {}, offset has length {}",
x.nrows(),
x.ncols(),
beta.len(),
offset.len()
),
});
}
let mut eta = x.matrixvectormultiply(&beta.to_owned());
eta += &offset;
Ok(eta)
}
#[inline]
fn royston_parmar_survival_hazard_components(
eta: f64,
eta_derivative: f64,
) -> Result<(f64, f64), SurvivalPredictError> {
if !(eta.is_finite() && eta_derivative.is_finite() && eta_derivative >= 0.0) {
return Err(SurvivalPredictError::NumericalFailure {
reason: format!(
"saved Royston-Parmar survival prediction produced invalid log-cumulative-hazard derivative: eta={eta}, eta_t={eta_derivative}"
),
});
}
let cumulative_hazard = eta.exp();
let hazard = if eta_derivative == 0.0 {
0.0
} else {
cumulative_hazard * eta_derivative
};
if !(cumulative_hazard >= 0.0 && hazard >= 0.0) {
return Err(SurvivalPredictError::NumericalFailure {
reason: format!(
"saved Royston-Parmar survival prediction produced invalid survival components: eta={eta}, eta_t={eta_derivative}, cumulative_hazard={cumulative_hazard}, hazard={hazard}"
),
});
}
Ok((cumulative_hazard, hazard))
}
fn predict_survival_location_scale_batch(
model: &SavedModel,
age_entry: &Array1<f64>,
age_exit: &Array1<f64>,
cov_design: &gam_terms::smooth::TermCollectionDesign,
primary_offset: &Array1<f64>,
noise_offset: &Array1<f64>,
training_headers: Option<&Vec<String>>,
col_map: &HashMap<String, usize>,
data: ArrayView2<'_, f64>,
time_grid: Option<&[f64]>,
with_uncertainty: bool,
covariance_mode: SurvivalPredictionCovarianceMode,
) -> Result<SurvivalPredictResult, String> {
use crate::survival::construction::evaluate_survival_time_basis_row;
use crate::survival::location_scale::{
SurvivalLocationScalePredictInput, predict_survival_location_scale,
predict_survival_location_scalewith_uncertainty, replay_survival_covariate_channels,
};
use gam_linalg::matrix::DesignMatrix;
let n = age_entry.len();
let per_row_eval = time_grid.is_none();
let eval_times: Vec<f64> = match time_grid {
Some(grid) => {
if grid.is_empty() {
return Err("survival time_grid must contain at least one time".to_string());
}
for (idx, &t) in grid.iter().enumerate() {
if !t.is_finite() || t < 0.0 {
return Err(format!(
"survival time_grid requires finite non-negative times (index {idx})",
));
}
}
grid.to_vec()
}
None => Vec::new(),
};
let t_cols = if per_row_eval { 1 } else { eval_times.len() };
let eval_width = if per_row_eval { 1 } else { t_cols + 1 };
let saved_likelihood_mode = SurvivalLikelihoodMode::LocationScale;
let baseline_cfg = saved_survival_runtime_baseline_config(model)?;
let saved_fit = saved_survival_location_scale_fit_result(model)?;
let saved_structure = model
.survival_location_scale_structure
.as_ref()
.ok_or_else(|| {
"saved location-scale survival model is missing exact replay structure".to_string()
})?;
let reduced_parametric_aft = matches!(
saved_structure.time_parameterization,
crate::survival::location_scale::SurvivalLocationScaleTimeParameterization::ReducedParametricAft
);
let time_cfg = load_survival_time_basis_config_from_model(model)?;
let mut time_build = build_survival_time_basis(age_entry, age_exit, time_cfg.clone(), None)?;
let resolved_time_cfg = resolved_survival_time_basis_config_from_build(
&time_build.basisname,
time_build.degree,
time_build.knots.as_ref(),
time_build.keep_cols.as_ref(),
time_build.smooth_lambda,
)?;
let time_anchor = model
.survival_time_anchor
.ok_or_else(|| "saved survival model missing survival_time_anchor".to_string())?;
let time_anchor_row = evaluate_survival_time_basis_row(time_anchor, &resolved_time_cfg)?;
center_survival_time_designs_at_anchor(
&mut time_build.x_entry_time,
&mut time_build.x_exit_time,
&time_anchor_row,
)?;
if !model.has_baseline_time_wiggle() && !reduced_parametric_aft {
require_structural_survival_time_basis(&time_build.basisname, "saved survival sampling")?;
}
let saved_inverse_link = resolve_survival_inverse_link_from_saved(model)?;
let (eval_entry, eval_exit) = if per_row_eval {
(age_entry.clone(), age_exit.clone())
} else {
let total = n * eval_width;
let mut entry = Array1::<f64>::zeros(total);
let mut exit = Array1::<f64>::zeros(total);
{
use rayon::iter::{IntoParallelIterator, ParallelIterator};
let pairs: Vec<(f64, f64)> = (0..total)
.into_par_iter()
.map(|k| {
let i = k / eval_width;
let col = k % eval_width;
let t = if col < t_cols {
eval_times[col]
} else {
age_exit[i]
};
(age_entry[i].min(t), t)
})
.collect();
for (k, (t0, t1)) in pairs.into_iter().enumerate() {
entry[k] = t0;
exit[k] = t1;
}
}
(entry, exit)
};
let mut time_build =
build_survival_time_basis(&eval_entry, &eval_exit, time_cfg.clone(), None)?;
center_survival_time_designs_at_anchor(
&mut time_build.x_entry_time,
&mut time_build.x_exit_time,
&time_anchor_row,
)?;
let (mut eta_offset_entry, mut eta_offset_exit, mut derivative_offset_exit) =
build_survival_time_offsets_for_likelihood(
&eval_entry,
&eval_exit,
&baseline_cfg,
saved_likelihood_mode,
Some(&saved_inverse_link),
)?;
add_survival_time_derivative_guard_offset(
&eval_entry,
&eval_exit,
time_anchor,
survival_derivative_guard_for_likelihood(saved_likelihood_mode),
&mut eta_offset_entry,
&mut eta_offset_exit,
&mut derivative_offset_exit,
)?;
if reduced_parametric_aft {
eta_offset_exit = Array1::<f64>::zeros(eval_exit.len());
}
let saved_timewiggle_runtime = model.saved_baseline_time_wiggle()?;
let threshold_design = cov_design;
let log_sigmaspec = resolve_termspec_for_prediction(
&model.resolved_termspec_noise,
training_headers,
col_map,
"resolved_termspec_noise",
)?;
let sigma_clipped = model.axis_clip_to_training_ranges(data, col_map);
let sigma_input = sigma_clipped.as_ref().map_or(data, |arr| arr.view());
let raw_sigma_design =
gam_terms::smooth::build_term_collection_design(sigma_input, &log_sigmaspec)
.map_err(|err| format!("failed to build survival log-sigma design: {err}"))?;
let effective_noise_offset = raw_sigma_design
.compose_offset(
noise_offset.view(),
"survival location-scale log-sigma block",
)
.map_err(|error| error.to_string())?;
let x_time_exit_dense = time_build
.x_exit_time
.try_to_dense_by_chunks("survival location-scale prediction time-exit design")?;
let total_rows = eval_exit.len();
let x_time_exit = if let Some(runtime) = saved_timewiggle_runtime.as_ref() {
let mut full =
Array2::<f64>::zeros((total_rows, x_time_exit_dense.ncols() + runtime.beta.len()));
full.slice_mut(s![.., 0..x_time_exit_dense.ncols()])
.assign(&x_time_exit_dense);
full
} else {
x_time_exit_dense
};
let repeat_rows =
|matrix: &DesignMatrix, label: &str| -> Result<DesignMatrix, SurvivalPredictError> {
if per_row_eval {
return Ok(matrix.clone());
}
let dense = matrix.try_to_dense_by_chunks(label)?;
let mut repeated = Array2::<f64>::zeros((total_rows, dense.ncols()));
use rayon::iter::{IntoParallelIterator, ParallelIterator};
let rows: Vec<Vec<f64>> = (0..total_rows)
.into_par_iter()
.map(|k| dense.row(k / eval_width).to_vec())
.collect();
for (k, row) in rows.into_iter().enumerate() {
for (j, value) in row.into_iter().enumerate() {
repeated[[k, j]] = value;
}
}
Ok(DesignMatrix::from(repeated))
};
let expand_vector = |values: &Array1<f64>| -> Array1<f64> {
if per_row_eval {
values.clone()
} else {
Array1::from_shape_fn(total_rows, |k| values[k / eval_width])
}
};
if saved_structure.threshold_time_basis.is_some()
&& threshold_design
.affine_offset
.iter()
.any(|value| *value != 0.0)
{
return Err(
"saved time-varying survival threshold cannot carry a non-zero smooth anchor"
.to_string(),
);
}
if saved_structure.log_sigma_time_basis.is_some()
&& raw_sigma_design
.affine_offset
.iter()
.any(|value| *value != 0.0)
{
return Err(
"saved time-varying survival log-sigma cannot carry a non-zero smooth anchor"
.to_string(),
);
}
let threshold_base_matrix = repeat_rows(
&threshold_design.design,
"survival location-scale prediction threshold design",
)?;
let raw_sigma_base_matrix = repeat_rows(
&raw_sigma_design.design,
"survival location-scale prediction log-sigma design",
)?;
let mut threshold_replay = replay_survival_covariate_channels(
&threshold_base_matrix,
&expand_vector(primary_offset),
&eval_entry,
&eval_exit,
saved_structure.threshold_time_basis.as_ref(),
"survival location-scale threshold",
)?;
let sigma_replay = replay_survival_covariate_channels(
&raw_sigma_base_matrix,
&expand_vector(&effective_noise_offset),
&eval_entry,
&eval_exit,
saved_structure.log_sigma_time_basis.as_ref(),
"survival location-scale log-sigma",
)?;
let link_wiggle_knots = model
.linkwiggle_knots
.as_ref()
.map(|k| Array1::from_vec(k.clone()));
let link_wiggle_degree = model.linkwiggle_degree;
let time_wiggle_knots = saved_timewiggle_runtime
.as_ref()
.map(|w| Array1::from_vec(w.knots.clone()));
let time_wiggle_degree = saved_timewiggle_runtime.as_ref().map(|w| w.degree);
let time_wiggle_ncols = saved_timewiggle_runtime
.as_ref()
.map_or(0, |w| w.beta.len());
if reduced_parametric_aft {
for (slot, &t) in threshold_replay.offset.iter_mut().zip(eval_exit.iter()) {
*slot -= t
.max(crate::survival::construction::SURVIVAL_TIME_FLOOR)
.ln();
}
}
let pred_input = SurvivalLocationScalePredictInput {
x_time_exit,
eta_time_offset_exit: eta_offset_exit.clone(),
time_wiggle_knots: time_wiggle_knots.clone(),
time_wiggle_degree,
time_wiggle_ncols,
x_threshold: threshold_replay.design_exit.clone(),
eta_threshold_offset: threshold_replay.offset.clone(),
x_log_sigma: sigma_replay.design_exit.clone(),
eta_log_sigma_offset: sigma_replay.offset.clone(),
x_link_wiggle: None,
link_wiggle_knots: link_wiggle_knots.clone(),
link_wiggle_degree,
inverse_link: saved_inverse_link.clone(),
};
let (eta_full, survival_prob_full, response_se_full, eta_se_full): (
Array1<f64>,
Array1<f64>,
Option<Array1<f64>>,
Option<Array1<f64>>,
) = if with_uncertainty {
let cov = match select_survival_prediction_covariance(
saved_fit.beta_covariance(),
saved_fit.beta_covariance_corrected(),
covariance_mode,
) {
Ok(cov) => cov,
Err(SurvivalPredictError::PosteriorCovariance { reason })
if covariance_mode == SurvivalPredictionCovarianceMode::Conditional =>
{
return Err(format!(
"survival location-scale uncertainty: {reason}; refit with the \
current CLI / library to populate beta_covariance"
));
}
Err(err) => return Err(String::from(err)),
};
let unc = predict_survival_location_scalewith_uncertainty(
&pred_input,
&saved_fit,
cov,
false,
true,
)
.map_err(|err| format!("survival location-scale uncertainty predict failed: {err}"))?;
let response_se = unc.response_standard_error.ok_or_else(|| {
"survival location-scale uncertainty: response_standard_error \
missing despite include_response_sd=true"
.to_string()
})?;
(
unc.eta,
unc.survival_prob,
Some(response_se),
Some(unc.eta_standard_error),
)
} else {
let pred = predict_survival_location_scale(&pred_input, &saved_fit)
.map_err(|err| format!("survival location-scale predict failed: {err}"))?;
(pred.eta, pred.survival_prob, None, None)
};
let beta_threshold = saved_fit.beta_threshold();
let beta_log_sigma = saved_fit.beta_log_sigma();
let eta_threshold = threshold_replay
.design_exit
.matrixvectormultiply(&beta_threshold)
+ &threshold_replay.offset;
let mut eta_threshold_derivative = threshold_replay
.design_derivative_exit
.as_ref()
.map(|design| design.matrixvectormultiply(&beta_threshold))
.unwrap_or_else(|| Array1::zeros(total_rows));
if reduced_parametric_aft {
for (slot, &time) in eta_threshold_derivative.iter_mut().zip(eval_exit.iter()) {
*slot -= 1.0 / time.max(crate::survival::construction::SURVIVAL_TIME_FLOOR);
}
}
let eta_log_sigma = sigma_replay
.design_exit
.matrixvectormultiply(&beta_log_sigma)
+ &sigma_replay.offset;
let eta_log_sigma_derivative = sigma_replay
.design_derivative_exit
.as_ref()
.map(|design| design.matrixvectormultiply(&beta_log_sigma))
.unwrap_or_else(|| Array1::zeros(total_rows));
let hdot = if reduced_parametric_aft {
Array1::zeros(total_rows)
} else {
let x_time_derivative = time_build
.x_derivative_time
.try_to_dense_by_chunks("survival location-scale prediction time-derivative design")?;
location_scale_eta_derivative_components(
&x_time_derivative,
&derivative_offset_exit,
&pred_input.x_time_exit,
&pred_input.eta_time_offset_exit,
time_wiggle_knots.as_ref(),
time_wiggle_degree,
time_wiggle_ncols,
&saved_fit,
)?
};
let inv_sigma = eta_log_sigma.mapv(crate::sigma_link::exp_sigma_inverse_from_eta_scalar);
let q_base = -&eta_threshold * &inv_sigma;
let mut qdot =
&inv_sigma * &(&eta_threshold * &eta_log_sigma_derivative - &eta_threshold_derivative);
if let Some(beta_wiggle) = saved_fit.beta_link_wiggle() {
let knots = link_wiggle_knots.as_ref().ok_or_else(|| {
"saved location-scale link-wiggle coefficients are missing knots".to_string()
})?;
let degree = link_wiggle_degree.ok_or_else(|| {
"saved location-scale link-wiggle coefficients are missing degree".to_string()
})?;
let derivative_basis = crate::wiggle::monotone_wiggle_basis_with_derivative_order(
q_base.view(),
knots,
degree,
1,
)?;
if derivative_basis.ncols() != beta_wiggle.len() {
return Err(format!(
"saved location-scale link-wiggle derivative width mismatch: design={}, beta={}",
derivative_basis.ncols(),
beta_wiggle.len()
));
}
qdot *= &(derivative_basis.dot(&beta_wiggle) + 1.0);
}
let eta_derivative_full = hdot + qdot;
if eta_derivative_full
.iter()
.any(|value| !(value.is_finite() && *value > 0.0))
{
return Err(
"saved location-scale survival event-rate derivative must be finite and positive"
.to_string(),
);
}
let hazard_full = location_scale_hazard_from_eta_derivative(
&eta_full,
&eta_derivative_full,
&saved_inverse_link,
)?;
let mut survival = Array2::<f64>::zeros((n, t_cols));
let mut cumulative_hazard = Array2::<f64>::zeros((n, t_cols));
let mut hazard = Array2::<f64>::zeros((n, t_cols));
ndarray::Zip::indexed(&mut survival)
.and(&mut cumulative_hazard)
.and(&mut hazard)
.par_for_each(|(i, j), s, ch, h| {
let query_time = if per_row_eval {
age_exit[i]
} else {
eval_times[j]
};
if query_time <= 0.0 {
*s = 1.0;
*ch = 0.0;
*h = 0.0;
return;
}
let k = if per_row_eval { i } else { i * eval_width + j };
let surv = survival_prob_full[k].clamp(SURVIVAL_PROB_MIN_FOR_LOG, 1.0);
*s = surv;
*ch = -surv.ln();
*h = hazard_full[k];
});
let linear_predictor = if per_row_eval {
eta_full.clone()
} else {
Array1::from_shape_fn(n, |i| eta_full[i * eval_width + t_cols])
};
let times = if per_row_eval {
age_exit.to_vec()
} else {
eval_times.clone()
};
let survival_se = response_se_full.as_ref().map(|response_se| {
let mut out = Array2::<f64>::zeros((n, t_cols));
ndarray::Zip::indexed(&mut out).par_for_each(|(i, j), slot| {
let query_time = if per_row_eval {
age_exit[i]
} else {
eval_times[j]
};
if query_time <= 0.0 {
*slot = 0.0;
return;
}
let k = if per_row_eval { i } else { i * eval_width + j };
*slot = response_se[k].max(0.0);
});
out
});
let eta_se_per_row = eta_se_full.as_ref().map(|eta_se| {
if per_row_eval {
eta_se.clone()
} else {
Array1::from_shape_fn(n, |i| eta_se[i * eval_width + t_cols])
}
});
Ok(SurvivalPredictResult {
times,
hazard,
survival,
cumulative_hazard,
linear_predictor,
likelihood_mode: saved_likelihood_mode,
survival_se,
eta_se: eta_se_per_row,
covariance_source: with_uncertainty.then_some(covariance_mode),
})
}
pub(crate) struct LocationScaleEtaComponents {
pub h: Array1<f64>,
pub time_jac: Array2<f64>,
pub eta_t: Array1<f64>,
pub eta_ls: Array1<f64>,
pub inv_sigma: Array1<f64>,
}
pub(crate) struct LocationScaleTimeWarpComponents {
pub(crate) h: Array1<f64>,
pub(crate) time_jac: Array2<f64>,
pub(crate) time_wiggle_dq: Option<Array1<f64>>,
}
pub(crate) fn location_scale_time_warp_components(
x_time_exit: &Array2<f64>,
eta_time_offset_exit: &Array1<f64>,
time_wiggle_knots: Option<&Array1<f64>>,
time_wiggle_degree: Option<usize>,
time_wiggle_ncols: usize,
fit: &UnifiedFitResult,
) -> Result<LocationScaleTimeWarpComponents, String> {
let n = x_time_exit.nrows();
if eta_time_offset_exit.len() != n {
return Err("survival location-scale time-warp row mismatch across inputs".to_string());
}
let beta_time = fit.beta_time();
if x_time_exit.ncols() != beta_time.len() {
return Err(format!(
"survival location-scale time-warp design mismatch: x_exit={} beta_time={}",
x_time_exit.ncols(),
beta_time.len()
));
}
let p_time_total = beta_time.len();
let p_wiggle = time_wiggle_ncols.min(p_time_total);
let p_base = p_time_total - p_wiggle;
let beta_base = beta_time.slice(s![..p_base]).to_owned();
let h_base = if p_base > 0 {
x_time_exit.slice(s![.., ..p_base]).dot(&beta_base) + eta_time_offset_exit
} else {
eta_time_offset_exit.clone()
};
let mut h = h_base.clone();
let mut time_jac = x_time_exit.clone();
let mut time_wiggle_dq = None;
if p_wiggle > 0 {
if x_time_exit
.slice(s![.., p_base..p_time_total])
.iter()
.any(|&value| value != 0.0)
{
return Err(
"survival location-scale timewiggle prediction requires zero placeholder tail columns"
.to_string(),
);
}
let knots = time_wiggle_knots.ok_or_else(|| {
"survival location-scale time-warp: timewiggle coefficients are missing knot metadata"
.to_string()
})?;
let degree = time_wiggle_degree.ok_or_else(|| {
"survival location-scale time-warp: timewiggle coefficients are missing degree metadata"
.to_string()
})?;
let beta_w = beta_time.slice(s![p_base..p_time_total]).to_owned();
let time_basis = crate::wiggle::monotone_wiggle_basis_with_derivative_order(
h_base.view(),
knots,
degree,
0,
)?;
let time_basis_d1 = crate::wiggle::monotone_wiggle_basis_with_derivative_order(
h_base.view(),
knots,
degree,
1,
)?;
if time_basis.ncols() != p_wiggle || time_basis_d1.ncols() != p_wiggle {
return Err(format!(
"survival location-scale time-warp timewiggle mismatch: value basis has {} columns, derivative basis has {}, beta has {}",
time_basis.ncols(),
time_basis_d1.ncols(),
p_wiggle
));
}
let dq = time_basis_d1.dot(&beta_w) + 1.0;
h = &h_base + &time_basis.dot(&beta_w);
time_jac = Array2::<f64>::zeros((n, p_time_total));
if p_base > 0 {
let scaled_base = crate::survival::location_scale::scale_dense_rows(
&x_time_exit.slice(s![.., ..p_base]).to_owned(),
&dq,
)?;
time_jac.slice_mut(s![.., ..p_base]).assign(&scaled_base);
}
time_jac
.slice_mut(s![.., p_base..p_time_total])
.assign(&time_basis);
time_wiggle_dq = Some(dq);
}
Ok(LocationScaleTimeWarpComponents {
h,
time_jac,
time_wiggle_dq,
})
}
pub(crate) fn location_scale_eta_components(
x_time_exit: &Array2<f64>,
eta_time_offset_exit: &Array1<f64>,
time_wiggle_knots: Option<&Array1<f64>>,
time_wiggle_degree: Option<usize>,
time_wiggle_ncols: usize,
x_threshold: &gam_linalg::matrix::DesignMatrix,
eta_threshold_offset: &Array1<f64>,
x_log_sigma: &gam_linalg::matrix::DesignMatrix,
eta_log_sigma_offset: &Array1<f64>,
fit: &UnifiedFitResult,
) -> Result<LocationScaleEtaComponents, String> {
let n = x_time_exit.nrows();
if x_threshold.nrows() != n
|| eta_threshold_offset.len() != n
|| x_log_sigma.nrows() != n
|| eta_log_sigma_offset.len() != n
{
return Err("survival location-scale eta component row mismatch across inputs".to_string());
}
let time_components = location_scale_time_warp_components(
x_time_exit,
eta_time_offset_exit,
time_wiggle_knots,
time_wiggle_degree,
time_wiggle_ncols,
fit,
)?;
let beta_threshold = fit.beta_threshold();
let beta_log_sigma = fit.beta_log_sigma();
let eta_t = x_threshold.matrixvectormultiply(&beta_threshold) + eta_threshold_offset;
let eta_ls = x_log_sigma.matrixvectormultiply(&beta_log_sigma) + eta_log_sigma_offset;
let inv_sigma = eta_ls.mapv(crate::sigma_link::exp_sigma_inverse_from_eta_scalar);
Ok(LocationScaleEtaComponents {
h: time_components.h,
time_jac: time_components.time_jac,
eta_t,
eta_ls,
inv_sigma,
})
}
fn location_scale_eta_derivative_components(
x_time_derivative: &Array2<f64>,
derivative_offset_exit: &Array1<f64>,
x_time_exit: &Array2<f64>,
eta_time_offset_exit: &Array1<f64>,
time_wiggle_knots: Option<&Array1<f64>>,
time_wiggle_degree: Option<usize>,
time_wiggle_ncols: usize,
fit: &UnifiedFitResult,
) -> Result<Array1<f64>, String> {
let n = x_time_exit.nrows();
if x_time_derivative.nrows() != n
|| derivative_offset_exit.len() != n
|| eta_time_offset_exit.len() != n
{
return Err(
"survival location-scale hazard derivative row mismatch across inputs".to_string(),
);
}
let beta_time = fit.beta_time();
let p_time_total = beta_time.len();
let p_wiggle = time_wiggle_ncols.min(p_time_total);
let p_base = p_time_total - p_wiggle;
if x_time_exit.ncols() != p_time_total || x_time_derivative.ncols() != p_base {
return Err(format!(
"survival location-scale hazard derivative design mismatch: x_exit={} beta_time={} x_derivative={} base={}",
x_time_exit.ncols(),
p_time_total,
x_time_derivative.ncols(),
p_base
));
}
let time_components = location_scale_time_warp_components(
x_time_exit,
eta_time_offset_exit,
time_wiggle_knots,
time_wiggle_degree,
time_wiggle_ncols,
fit,
)?;
let beta_base = beta_time.slice(s![..p_base]).to_owned();
let mut eta_derivative = if p_base > 0 {
x_time_derivative.dot(&beta_base) + derivative_offset_exit
} else {
derivative_offset_exit.clone()
};
if let Some(dq) = time_components.time_wiggle_dq.as_ref() {
eta_derivative *= dq;
}
if eta_derivative
.iter()
.any(|value| !(value.is_finite() && *value > 0.0))
{
return Err(
"survival location-scale hazard derivative must be finite and positive".to_string(),
);
}
Ok(eta_derivative)
}
fn location_scale_hazard_from_eta_derivative(
eta: &Array1<f64>,
eta_derivative: &Array1<f64>,
inverse_link: &InverseLink,
) -> Result<Array1<f64>, String> {
if eta.len() != eta_derivative.len() {
return Err(format!(
"survival location-scale hazard row mismatch: eta={} eta_derivative={}",
eta.len(),
eta_derivative.len()
));
}
let values = eta
.iter()
.zip(eta_derivative.iter())
.map(|(&q, &q_t)| location_scale_hazard_component(q, q_t, inverse_link))
.collect::<Result<Vec<_>, _>>()?;
Ok(Array1::from_vec(values))
}
fn location_scale_hazard_component(
eta: f64,
eta_derivative: f64,
inverse_link: &InverseLink,
) -> Result<f64, String> {
if !(eta.is_finite() && eta_derivative.is_finite() && eta_derivative > 0.0) {
return Err(format!(
"survival location-scale hazard requires finite eta and positive eta_t, got eta={eta}, eta_t={eta_derivative}"
));
}
match inverse_link {
InverseLink::Standard(StandardLink::Probit) => {
let (_, hazard) = probit_survival_hazard_components(eta, eta_derivative)?;
Ok(hazard)
}
InverseLink::Standard(StandardLink::CLogLog) => {
let (_, hazard) = royston_parmar_survival_hazard_components(eta, eta_derivative)?;
Ok(hazard)
}
InverseLink::Standard(StandardLink::Logit) => {
let failure = if eta >= 0.0 {
1.0 / (1.0 + (-eta).exp())
} else {
let exp_eta = eta.exp();
exp_eta / (1.0 + exp_eta)
};
Ok(failure * eta_derivative)
}
InverseLink::Standard(StandardLink::Identity) => {
let survival = 1.0 - eta;
if !(survival.is_finite() && survival > 0.0) {
return Err(format!(
"survival location-scale identity link produced invalid survival={survival} at eta={eta}"
));
}
Ok(eta_derivative / survival)
}
_ => {
let jet = inverse_link_jet_for_inverse_link(inverse_link, eta)
.map_err(|err| format!("survival location-scale inverse-link jet failed: {err}"))?;
let survival = 1.0 - jet.mu;
let hazard = jet.d1 * eta_derivative / survival;
if !(survival.is_finite() && survival > 0.0 && hazard.is_finite() && hazard >= 0.0) {
return Err(format!(
"survival location-scale inverse link produced invalid hazard components: eta={eta}, eta_t={eta_derivative}, failure={}, d_failure={}, survival={survival}, hazard={hazard}",
jet.mu, jet.d1
));
}
Ok(hazard)
}
}
}
pub fn require_saved_survival_likelihood_mode(
model: &SavedModel,
) -> Result<SurvivalLikelihoodMode, SurvivalPredictError> {
if matches!(&model.family_state, FittedFamily::LatentSurvival { .. }) {
return match model.survival_likelihood.as_deref() {
Some("latent") => Ok(SurvivalLikelihoodMode::Latent),
Some(other) => Err(SurvivalPredictError::MissingFitMetadata { reason: format!(
"saved latent survival model has contradictory survival_likelihood metadata: expected 'latent', got '{other}'"
) }),
None => Err(SurvivalPredictError::MissingFitMetadata {
reason:
"saved latent survival model is missing survival_likelihood=latent metadata; refit"
.to_string(),
}),
};
}
if matches!(&model.family_state, FittedFamily::LatentBinary { .. }) {
return match model.survival_likelihood.as_deref() {
Some("latent-binary") => Ok(SurvivalLikelihoodMode::LatentBinary),
Some(other) => Err(SurvivalPredictError::MissingFitMetadata { reason: format!(
"saved latent binary model has contradictory survival_likelihood metadata: expected 'latent-binary', got '{other}'"
) }),
None => Err(SurvivalPredictError::MissingFitMetadata {
reason:
"saved latent binary model is missing survival_likelihood=latent-binary metadata; refit"
.to_string(),
}),
};
}
let raw = model.survival_likelihood.as_deref().ok_or_else(|| {
"saved survival model is missing survival_likelihood metadata; refit".to_string()
})?;
parse_survival_likelihood_mode(raw).map_err(SurvivalPredictError::from)
}
pub fn saved_survival_runtime_baseline_config(
model: &SavedModel,
) -> Result<SurvivalBaselineConfig, SurvivalPredictError> {
survival_baseline_config_from_model(model).map_err(SurvivalPredictError::from)
}
pub fn resolve_termspec_for_prediction(
modelspec: &Option<TermCollectionSpec>,
training_headers: Option<&Vec<String>>,
col_map: &HashMap<String, usize>,
spec_label: &str,
) -> Result<TermCollectionSpec, SurvivalPredictError> {
let saved = modelspec.as_ref().ok_or_else(|| {
format!(
"model is missing {spec_label}; refit to guarantee train/predict design consistency"
)
})?;
saved.validate_frozen(spec_label)?;
let headers = training_headers.ok_or_else(|| {
"model is missing training_headers; refit to guarantee stable feature mapping at prediction time"
.to_string()
})?;
let remapped = remap_term_collectionspec_columns(saved, headers, col_map)?;
remapped.validate_frozen(spec_label)?;
Ok(remapped)
}
fn remap_term_collectionspec_columns(
spec: &TermCollectionSpec,
training_headers: &[String],
prediction_column_map: &HashMap<String, usize>,
) -> Result<TermCollectionSpec, SurvivalPredictError> {
spec.remap_feature_columns(|index| -> Result<usize, SurvivalPredictError> {
let name = training_headers
.get(index)
.ok_or_else(|| format!("saved training column index {index} is out of bounds"))?;
resolve_role_col(prediction_column_map, name, "prediction")
.map_err(SurvivalPredictError::from)
})
}
pub fn fit_result_from_saved_model_for_prediction(
model: &SavedModel,
) -> Result<UnifiedFitResult, String> {
model
.fit_result
.clone()
.ok_or_else(|| "model is missing canonical fit_result payload; refit".to_string())
}
pub fn saved_survival_location_scale_fit_result(
model: &SavedModel,
) -> Result<UnifiedFitResult, SurvivalPredictError> {
model.saved_prediction_runtime()?;
let mut fit = model.fit_result.clone().ok_or_else(|| {
"saved location-scale survival model missing canonical fit_result; refit".to_string()
})?;
let inverse_link = resolve_survival_inverse_link_from_saved(model)?;
apply_inverse_link_state_to_fit_result(&mut fit, &inverse_link);
Ok(fit)
}
pub fn apply_inverse_link_state_to_fit_result(
fit_result: &mut UnifiedFitResult,
inverse_link: &InverseLink,
) {
fit_result.fitted_link = match inverse_link {
InverseLink::LatentCLogLog(state) => FittedLinkState::LatentCLogLog { state: *state },
InverseLink::Sas(state) => FittedLinkState::Sas {
state: *state,
covariance: None,
},
InverseLink::BetaLogistic(state) => FittedLinkState::BetaLogistic {
state: *state,
covariance: None,
},
InverseLink::Mixture(state) => FittedLinkState::Mixture {
state: state.clone(),
covariance: None,
},
InverseLink::Standard(_) => FittedLinkState::Standard(None),
};
}
pub fn resolve_survival_inverse_link_from_saved(
model: &SavedModel,
) -> Result<InverseLink, SurvivalPredictError> {
if let Some(link) = model.link.as_ref() {
return Ok(link.clone());
}
Err(SurvivalPredictError::MissingFitMetadata {
reason: "saved survival model is missing link metadata; refit".to_string(),
})
}
pub fn concat_array1_refs(parts: &[&Array1<f64>]) -> Array1<f64> {
let total: usize = parts.iter().map(|part| part.len()).sum();
let mut out = Array1::<f64>::zeros(total);
let mut offset = 0usize;
for part in parts {
let width = part.len();
out.slice_mut(s![offset..offset + width]).assign(part);
offset += width;
}
out
}
pub fn saved_baseline_timewiggle_components(
eta_entry: &Array1<f64>,
eta_exit: &Array1<f64>,
derivative_exit: &Array1<f64>,
model: &SavedModel,
) -> Result<Option<(Array2<f64>, Array2<f64>, Array2<f64>)>, SurvivalPredictError> {
match model.saved_baseline_time_wiggle()? {
None => Ok(None),
Some(runtime) => {
runtime.validate_global_monotonicity()?;
let SavedBaselineTimeWiggleRuntime {
knots,
degree,
beta,
..
} = runtime;
let knots = Array1::from_vec(knots);
let entry = match buildwiggle_block_input_from_knots(
eta_entry.view(),
&knots,
degree,
2,
false,
)?
.design
{
DesignMatrix::Dense(m) => m.to_dense_arc().as_ref().clone(),
_ => {
return Err(SurvivalPredictError::IncompatibleSchema {
reason: "saved baseline-timewiggle entry design must be dense".to_string(),
});
}
};
let exit = match buildwiggle_block_input_from_knots(
eta_exit.view(),
&knots,
degree,
2,
false,
)?
.design
{
DesignMatrix::Dense(m) => m.to_dense_arc().as_ref().clone(),
_ => {
return Err(SurvivalPredictError::IncompatibleSchema {
reason: "saved baseline-timewiggle exit design must be dense".to_string(),
});
}
};
let betaw = beta;
if entry.ncols() != betaw.len() || exit.ncols() != betaw.len() {
return Err(SurvivalPredictError::IncompatibleSchema {
reason: format!(
"saved baseline-timewiggle dimension mismatch: coefficients have {} entries but basis has entry={} exit={}",
betaw.len(),
entry.ncols(),
exit.ncols()
),
});
}
let derivative = build_survival_timewiggle_derivative_design(
eta_exit,
derivative_exit,
&knots,
degree,
)
.map_err(|e| {
e.replace(
"build baseline-timewiggle",
"evaluate saved baseline-timewiggle",
)
})?;
if derivative.ncols() != betaw.len() {
return Err(SurvivalPredictError::IncompatibleSchema {
reason: format!(
"saved baseline-timewiggle derivative dimension mismatch: coefficients have {} entries but derivative basis has {} columns",
betaw.len(),
derivative.ncols()
),
});
}
Ok(Some((entry, exit, derivative)))
}
}
}
pub fn build_saved_survival_marginal_slope_predictor(
model: &SavedModel,
fit_saved: &UnifiedFitResult,
z_name: &str,
z: &Array1<f64>,
cov_design: &DesignMatrix,
logslope_design: &DesignMatrix,
time_build: &SurvivalTimeBuildOutput,
eta_offset_entry: &Array1<f64>,
eta_offset_exit: &Array1<f64>,
derivative_offset_exit: &Array1<f64>,
primary_offset: &Array1<f64>,
noise_offset: &Array1<f64>,
) -> Result<
(
BernoulliMarginalSlopePredictor,
PredictInput,
UnifiedFitResult,
),
SurvivalPredictError,
> {
let saved_runtime = model.saved_prediction_runtime()?;
if saved_runtime.link_wiggle.is_some() {
return Err(SurvivalPredictError::MissingFitMetadata {
reason:
"saved survival marginal-slope model contains legacy linkwiggle metadata; refit with the anchored link-deviation runtime"
.to_string(),
});
}
let saved_score_runtime = saved_runtime.score_warp;
let saved_link_runtime = saved_runtime.link_deviation;
let influence_absorber_width = saved_runtime.influence_absorber_width;
let blocks = &fit_saved.blocks;
let expected_blocks = 3
+ usize::from(saved_score_runtime.is_some())
+ usize::from(saved_link_runtime.is_some())
+ usize::from(influence_absorber_width.is_some());
if blocks.len() != expected_blocks {
return Err(SurvivalPredictError::IncompatibleSchema {
reason: format!(
"saved survival marginal-slope model requires {} blocks [time, marginal, slope{}{}{}], got {}",
expected_blocks,
if saved_score_runtime.is_some() {
", score-warp"
} else {
""
},
if saved_link_runtime.is_some() {
", link-deviation"
} else {
""
},
if influence_absorber_width.is_some() {
", influence-absorber(dropped)"
} else {
""
},
blocks.len(),
),
});
}
let beta_time = &blocks[0].beta;
let beta_marginal = &blocks[1].beta;
let beta_logslope = &blocks[2].beta;
if let Some(runtime) = saved_score_runtime.as_ref() {
let beta = &blocks[3].beta;
if beta.len() != runtime.basis_dim {
return Err(SurvivalPredictError::IncompatibleSchema {
reason: format!(
"saved survival marginal-slope score-warp coefficient mismatch: beta has {} entries but runtime expects {}",
beta.len(),
runtime.basis_dim
),
});
}
}
if let Some(runtime) = saved_link_runtime.as_ref() {
let idx = 3 + usize::from(saved_score_runtime.is_some());
let beta = &blocks[idx].beta;
if beta.len() != runtime.basis_dim {
return Err(SurvivalPredictError::IncompatibleSchema {
reason: format!(
"saved survival marginal-slope link-deviation coefficient mismatch: beta has {} entries but runtime expects {}",
beta.len(),
runtime.basis_dim
),
});
}
}
if beta_marginal.len() != cov_design.ncols() {
return Err(SurvivalPredictError::IncompatibleSchema {
reason: format!(
"saved survival marginal-slope marginal coefficient mismatch: beta has {} entries but baseline design has {} columns",
beta_marginal.len(),
cov_design.ncols()
),
});
}
if beta_logslope.len() != logslope_design.ncols() {
return Err(SurvivalPredictError::IncompatibleSchema {
reason: format!(
"saved survival marginal-slope slope coefficient mismatch: beta has {} entries but slope design has {} columns",
beta_logslope.len(),
logslope_design.ncols()
),
});
}
let p_time_base = time_build.x_exit_time.ncols();
let saved_timewiggle = saved_runtime.baseline_time_wiggle;
let p_timewiggle = saved_timewiggle
.as_ref()
.map_or(0, |runtime| runtime.beta.len());
if beta_time.len() != p_time_base + p_timewiggle {
let hint = stale_weibull_time_basis_hint(
&time_build.basisname,
beta_time.len() == p_time_base + p_timewiggle + 1,
);
return Err(SurvivalPredictError::IncompatibleSchema {
reason: format!(
"saved survival marginal-slope time coefficient mismatch: beta has {} entries but expected base={} plus timewiggle={}{hint}",
beta_time.len(),
p_time_base,
p_timewiggle
),
});
}
let beta_time_base = beta_time.slice(s![..p_time_base]).to_owned();
let cov_eta_marginal = cov_design.dot(beta_marginal);
let q_entry_base = time_build.x_entry_time.dot(&beta_time_base)
+ &cov_eta_marginal
+ eta_offset_entry
+ primary_offset;
let q_exit_base = time_build.x_exit_time.dot(&beta_time_base)
+ &cov_eta_marginal
+ eta_offset_exit
+ primary_offset;
let qd_exit_base = time_build.x_derivative_time.dot(&beta_time_base) + derivative_offset_exit;
let mut q_design_parts = vec![time_build.x_exit_time.clone()];
if saved_timewiggle.is_some() {
let (_, exit_w, _) = saved_baseline_timewiggle_components(
&q_entry_base,
&q_exit_base,
&qd_exit_base,
model,
)?
.ok_or_else(|| {
"saved survival marginal-slope model is missing baseline-timewiggle runtime metadata"
.to_string()
})?;
if exit_w.ncols() != p_timewiggle {
return Err(SurvivalPredictError::IncompatibleSchema {
reason: format!(
"saved survival marginal-slope timewiggle design mismatch: rebuilt {} columns but runtime expects {}",
exit_w.ncols(),
p_timewiggle
),
});
}
q_design_parts.push(DesignMatrix::from(exit_w));
}
q_design_parts.push(cov_design.clone());
let q_design = DesignMatrix::hstack(q_design_parts)?;
let combined_q_beta = concat_array1_refs(&[beta_time, beta_marginal]);
let combined_q_lambdas = concat_array1_refs(&[&blocks[0].lambdas, &blocks[1].lambdas]);
let mut predictor_blocks = Vec::with_capacity(
2 + usize::from(saved_score_runtime.is_some()) + usize::from(saved_link_runtime.is_some()),
);
predictor_blocks.push(FittedBlock {
beta: combined_q_beta.clone(),
role: BlockRole::Mean,
edf: blocks[0].edf + blocks[1].edf,
lambdas: combined_q_lambdas,
});
predictor_blocks.push(FittedBlock {
beta: beta_logslope.clone(),
role: BlockRole::Scale,
edf: blocks[2].edf,
lambdas: blocks[2].lambdas.clone(),
});
if saved_score_runtime.is_some() {
let mut block = blocks[3].clone();
block.role = BlockRole::Mean;
predictor_blocks.push(block);
}
if saved_link_runtime.is_some() {
let idx = 3 + usize::from(saved_score_runtime.is_some());
let mut block = blocks[idx].clone();
block.role = BlockRole::LinkWiggle;
predictor_blocks.push(block);
}
let mut predictor_fit = fit_saved.clone();
predictor_fit.blocks = predictor_blocks;
predictor_fit.beta = concat_array1_refs(
&predictor_fit
.blocks
.iter()
.map(|block| &block.beta)
.collect::<Vec<_>>(),
);
predictor_fit.block_states.clear();
let predictor = BernoulliMarginalSlopePredictor::from_unified(
&predictor_fit,
z_name.to_string(),
model.latent_z_normalization.ok_or_else(|| {
"saved survival marginal-slope model missing latent_z_normalization".to_string()
})?,
model.latent_measure.clone().ok_or_else(|| {
"saved survival marginal-slope model missing latent_measure".to_string()
})?,
0.0,
model.logslope_baseline.ok_or_else(|| {
"saved survival marginal-slope model missing logslope_baseline".to_string()
})?,
model
.resolved_inverse_link()?
.unwrap_or(InverseLink::Standard(StandardLink::Probit)),
model
.family_state
.frailty()
.cloned()
.unwrap_or(FrailtySpec::None),
saved_score_runtime,
saved_link_runtime,
model.latent_z_rank_int_calibration.clone(),
model.latent_z_conditional_calibration.clone(),
)?;
let pred_input = PredictInput {
design: q_design,
offset: eta_offset_exit + primary_offset,
design_noise: Some(logslope_design.clone()),
offset_noise: Some(noise_offset.clone()),
auxiliary_scalar: Some(z.clone()),
auxiliary_matrix: None,
};
Ok((predictor, pred_input, predictor_fit))
}
fn stale_weibull_time_basis_hint(basisname: &str, extra_time_coefficient: bool) -> &'static str {
if basisname == "linear" && extra_time_coefficient {
" (this looks like a model saved before the #2301 Weibull time-basis \
change, which removed the redundant constant column; refit the model)"
} else {
""
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::probability::{normal_cdf, normal_pdf};
#[test]
fn competing_risks_covariance_mode_selects_exact_requested_matrix() {
let conditional = ndarray::array![[1.0, 0.2], [0.2, 2.0]];
let corrected = ndarray::array![[1.5, 0.4], [0.4, 3.0]];
let selected_conditional = select_survival_prediction_covariance(
Some(&conditional),
Some(&corrected),
SurvivalPredictionCovarianceMode::Conditional,
)
.expect("conditional covariance");
let selected_corrected = select_survival_prediction_covariance(
Some(&conditional),
Some(&corrected),
SurvivalPredictionCovarianceMode::SmoothingCorrected,
)
.expect("smoothing-corrected covariance");
assert_eq!(selected_conditional, &conditional);
assert_eq!(selected_corrected, &corrected);
assert_eq!(
SurvivalPredictionCovarianceMode::Conditional.as_str(),
"conditional"
);
assert_eq!(
SurvivalPredictionCovarianceMode::SmoothingCorrected.as_str(),
"smoothing-corrected"
);
}
#[test]
fn competing_risks_smoothing_covariance_never_falls_back() {
let conditional = ndarray::array![[1.0]];
let error = select_survival_prediction_covariance(
Some(&conditional),
None,
SurvivalPredictionCovarianceMode::SmoothingCorrected,
)
.expect_err("a corrected request must not substitute conditional covariance");
assert_eq!(
error.to_string(),
"fit result does not contain smoothing-corrected covariance"
);
}
#[test]
fn posterior_quadrature_second_moment_honors_cross_coordinate_covariance() {
let posterior_mean = ndarray::array![0.4, -0.2];
let covariance = ndarray::array![[0.9, 0.35], [0.35, 0.6]];
let mut functional_mean = 0.0_f64;
let mut functional_second = 0.0_f64;
let mut recovered_cross_covariance = 0.0_f64;
for_each_survival_posterior_node(&posterior_mean, &covariance, &[], |node, weight| {
let functional = node[0] + 2.0 * node[1];
functional_mean += weight * functional;
functional_second += weight * functional * functional;
recovered_cross_covariance +=
weight * (node[0] - posterior_mean[0]) * (node[1] - posterior_mean[1]);
Ok(())
})
.expect("joint posterior quadrature");
let expected_mean = posterior_mean[0] + 2.0 * posterior_mean[1];
let expected_variance =
covariance[[0, 0]] + 4.0 * covariance[[1, 1]] + 4.0 * covariance[[0, 1]];
assert!((functional_mean - expected_mean).abs() <= 1e-12);
assert!((recovered_cross_covariance - covariance[[0, 1]]).abs() <= 1e-12);
let mean_surface = Array2::from_elem((1, 1), functional_mean);
let second_surface = Array2::from_elem((1, 1), functional_second);
let standard_error = posterior_standard_error_matrix(
&mean_surface,
&second_surface,
"joint-covariance witness",
)
.expect("posterior standard error");
assert!((standard_error[[0, 0]].powi(2) - expected_variance).abs() <= 1e-11);
}
#[test]
fn posterior_quadrature_zero_covariance_has_zero_standard_error() {
let posterior_mean = ndarray::array![0.25, -0.75];
let covariance = Array2::<f64>::zeros((2, 2));
let mut functional_mean = 0.0_f64;
let mut functional_second = 0.0_f64;
let mut node_count = 0usize;
for_each_survival_posterior_node(&posterior_mean, &covariance, &[], |node, weight| {
let functional = node[0].exp() + node[1].sin();
functional_mean += weight * functional;
functional_second += weight * functional * functional;
node_count += 1;
Ok(())
})
.expect("rank-zero posterior quadrature");
assert_eq!(node_count, 1, "rank-zero covariance has one exact node");
let standard_error = posterior_standard_error_matrix(
&Array2::from_elem((1, 1), functional_mean),
&Array2::from_elem((1, 1), functional_second),
"rank-zero witness",
)
.expect("rank-zero posterior standard error");
assert_eq!(standard_error[[0, 0]], 0.0);
}
#[test]
fn posterior_quadrature_keeps_cone_coordinates_feasible_and_unbiased() {
let posterior_mean = ndarray::array![0.354, -8.30];
let covariance = ndarray::array![[0.2304, 0.30], [0.30, 0.9604]];
let mut min_cone0_unconstrained = f64::INFINITY;
for_each_survival_posterior_node(&posterior_mean, &covariance, &[], |node, _weight| {
min_cone0_unconstrained = min_cone0_unconstrained.min(node[0]);
Ok(())
})
.expect("unconstrained quadrature");
assert!(
min_cone0_unconstrained < 0.0,
"fixture must reproduce the infeasible-node bug (min β_0 = {min_cone0_unconstrained})"
);
let mut mean0 = 0.0_f64;
let mut mean1 = 0.0_f64;
let mut weight_sum = 0.0_f64;
let mut min_cone0 = f64::INFINITY;
for_each_survival_posterior_node(&posterior_mean, &covariance, &[0], |node, weight| {
assert!(
node[0] >= -1e-12,
"cone coordinate stepped below its β_0 ≥ 0 wall: {}",
node[0]
);
min_cone0 = min_cone0.min(node[0]);
mean0 += weight * node[0];
mean1 += weight * node[1];
weight_sum += weight;
Ok(())
})
.expect("cone-truncated quadrature");
assert!((weight_sum - 1.0).abs() <= 1e-12, "weights must sum to one");
assert!(
(mean0 - posterior_mean[0]).abs() <= 1e-12
&& (mean1 - posterior_mean[1]).abs() <= 1e-12,
"cone truncation must leave the posterior mean unbiased (got [{mean0}, {mean1}])"
);
let mut var0_unconstrained = 0.0_f64;
for_each_survival_posterior_node(&posterior_mean, &covariance, &[], |node, weight| {
var0_unconstrained += weight * (node[0] - posterior_mean[0]).powi(2);
Ok(())
})
.expect("unconstrained spread");
let mut var0_cone = 0.0_f64;
for_each_survival_posterior_node(&posterior_mean, &covariance, &[0], |node, weight| {
var0_cone += weight * (node[0] - posterior_mean[0]).powi(2);
Ok(())
})
.expect("cone spread");
assert!(
var0_cone <= var0_unconstrained + 1e-12 && var0_cone < var0_unconstrained,
"cone spread {var0_cone} must be strictly smaller than the untruncated {var0_unconstrained}"
);
}
#[test]
fn posterior_quadrature_cone_is_a_noop_far_from_the_wall() {
let posterior_mean = ndarray::array![40.0, -0.2];
let covariance = ndarray::array![[0.9, 0.35], [0.35, 0.6]];
let mut recovered_var0 = 0.0_f64;
let mut recovered_cross = 0.0_f64;
for_each_survival_posterior_node(&posterior_mean, &covariance, &[0], |node, weight| {
recovered_var0 += weight * (node[0] - posterior_mean[0]).powi(2);
recovered_cross +=
weight * (node[0] - posterior_mean[0]) * (node[1] - posterior_mean[1]);
Ok(())
})
.expect("cone quadrature far from the wall");
assert!((recovered_var0 - covariance[[0, 0]]).abs() <= 1e-11);
assert!((recovered_cross - covariance[[0, 1]]).abs() <= 1e-11);
}
#[test]
fn posterior_quadrature_radius_collapses_on_an_active_bound() {
let posterior_mean = ndarray::array![0.0, 0.75];
let covariance = ndarray::array![[0.5, 0.0], [0.0, 0.2]];
let mut min_pinned = f64::INFINITY;
let mut max_pinned = f64::NEG_INFINITY;
let mut spread_unpinned = 0.0_f64;
for_each_survival_posterior_node(&posterior_mean, &covariance, &[0], |node, weight| {
min_pinned = min_pinned.min(node[0]);
max_pinned = max_pinned.max(node[0]);
spread_unpinned += weight * (node[1] - posterior_mean[1]).powi(2);
Ok(())
})
.expect("active-bound quadrature");
assert!(
min_pinned >= 0.0,
"an active bound must never be crossed, got {min_pinned}"
);
assert!(
max_pinned.abs() <= 1e-12,
"a direction loading an active-bound coordinate carries zero symmetric spread, \
but the coordinate reached {max_pinned}"
);
assert!(
(spread_unpinned - covariance[[1, 1]]).abs() <= 1e-11,
"a coordinate outside the cone keeps its full spread, got {spread_unpinned} want {}",
covariance[[1, 1]]
);
}
#[test]
fn posterior_quadrature_clamps_a_roundoff_negative_cone_coordinate() {
let roundoff_below_wall = -1e-15_f64;
let posterior_mean = ndarray::array![roundoff_below_wall, 0.75];
let covariance = ndarray::array![[0.5, 0.0], [0.0, 0.2]];
let mut nodes = Vec::new();
for_each_survival_posterior_node(&posterior_mean, &covariance, &[0], |node, _weight| {
nodes.push(node[0]);
Ok(())
})
.expect("round-off-negative cone quadrature");
for value in &nodes {
assert!(
*value >= roundoff_below_wall,
"truncation must never push a cone coordinate further below the wall than the \
fit left it: node {value} < β̂ {roundoff_below_wall}"
);
assert!(
(*value - roundoff_below_wall).abs() <= 1e-12,
"a coordinate at the wall carries no spread, got {value}"
);
}
}
#[test]
fn probit_survival_hazard_uses_density_over_survival() {
let eta = 2.0;
let eta_t = 0.3;
let (cum, hazard) =
probit_survival_hazard_components(eta, eta_t).expect("valid components");
let survival = normal_cdf(-eta);
let expected_cum = -survival.ln();
let expected_hazard = normal_pdf(eta) * eta_t / survival;
assert!((cum - expected_cum).abs() <= 1e-14);
assert!((hazard - expected_hazard).abs() <= 1e-14);
}
#[test]
fn probit_survival_hazard_stays_finite_in_right_tail() {
let eta = 40.0;
let eta_t = 9.694_340_360_912_401e-5;
let event_density =
(-0.5_f64 * eta * eta).exp() / (2.0 * std::f64::consts::PI).sqrt() * eta_t;
assert_eq!(event_density, 0.0);
let (cum, hazard) =
probit_survival_hazard_components(eta, eta_t).expect("valid tail components");
assert!(cum > 800.0, "right-tail cumulative hazard was {cum}");
assert!(
(3.87e-3..3.89e-3).contains(&hazard),
"right-tail hazard was {hazard}"
);
}
#[test]
fn probit_survival_hazard_accepts_zero_time_derivative_as_flat_hazard() {
let (cum, hazard) =
probit_survival_hazard_components(1.0, 0.0).expect("zero derivative is flat hazard");
assert!(cum > 0.0);
assert_eq!(hazard, 0.0);
}
#[test]
fn marginal_slope_index_derivative_clamps_extrapolation_negative_to_flat_hazard() {
let deta_dq = (1.0_f64 + 0.4 * 0.4).sqrt(); let qd_with_wiggle = -1.35e-3;
let eta_t = marginal_slope_index_derivative_at_horizon(deta_dq, qd_with_wiggle);
assert_eq!(
eta_t, 0.0,
"negative extrapolation derivative must clamp to 0"
);
let (cum, hazard) = probit_survival_hazard_components(-0.563, eta_t)
.expect("clamped flat-hazard prediction must validate");
assert!(
cum >= 0.0,
"cumulative hazard must be well-posed, got {cum}"
);
assert_eq!(
hazard, 0.0,
"clamped derivative gives zero instantaneous hazard"
);
}
#[test]
fn marginal_slope_index_derivative_preserves_positive_and_nonfinite() {
let positive = marginal_slope_index_derivative_at_horizon(1.25, 0.8);
assert!(
(positive - 1.0).abs() <= 1e-15,
"positive derivative scaled by chain factor"
);
let nonfinite = marginal_slope_index_derivative_at_horizon(1.25, f64::NAN);
assert!(
nonfinite.is_nan(),
"non-finite derivative passes through unclamped"
);
assert!(
probit_survival_hazard_components(0.5, nonfinite).is_err(),
"non-finite derivative must still be rejected by the validator"
);
}
#[test]
fn probit_survival_hazard_rejects_infinite_time_derivative() {
let err = probit_survival_hazard_components(1.0, f64::INFINITY)
.expect_err("infinite derivative should be invalid");
assert!(
err.to_string()
.contains("invalid survival index derivative")
);
}
#[test]
fn probit_survival_hazard_rejects_nan_inputs() {
let err_eta =
probit_survival_hazard_components(f64::NAN, 0.5).expect_err("NaN eta must be rejected");
assert!(
err_eta
.to_string()
.contains("invalid survival index derivative")
);
let err_dt = probit_survival_hazard_components(1.0, f64::NAN)
.expect_err("NaN eta_derivative must be rejected");
assert!(
err_dt
.to_string()
.contains("invalid survival index derivative")
);
}
#[test]
fn probit_survival_hazard_rejects_negative_time_derivative() {
let err = probit_survival_hazard_components(1.0, -0.5)
.expect_err("negative derivative should be invalid");
assert!(
err.to_string()
.contains("invalid survival index derivative")
);
}
#[test]
fn royston_parmar_hazard_is_cumulative_hazard_derivative() {
let eta = 2.0_f64.ln();
let eta_t = 0.25;
let (cum, hazard) =
royston_parmar_survival_hazard_components(eta, eta_t).expect("valid components");
assert!((cum - 2.0).abs() <= 1e-14);
assert!((hazard - 0.5).abs() <= 1e-14);
assert_ne!(hazard, cum);
}
#[test]
fn royston_parmar_hazard_rejects_negative_log_hazard_derivative() {
let err = royston_parmar_survival_hazard_components(0.0, -0.5)
.expect_err("negative derivative should be invalid");
assert!(
err.to_string()
.contains("invalid log-cumulative-hazard derivative")
);
}
#[test]
fn royston_parmar_hazard_accepts_zero_derivative_as_flat_boundary() {
let eta = 1.9909019457445971_f64; let (cum, hazard) = royston_parmar_survival_hazard_components(eta, 0.0)
.expect("zero derivative is a valid flat boundary, not an error");
assert!((cum - eta.exp()).abs() <= 1e-12, "cum = Λ(t) = exp(η)");
assert_eq!(
hazard, 0.0,
"flat cumulative hazard ⇒ zero instantaneous hazard"
);
let survival = (-cum).exp().clamp(0.0, 1.0);
assert!(survival.is_finite() && (0.0..=1.0).contains(&survival));
}
#[test]
fn royston_parmar_hazard_zero_derivative_in_saturated_tail_is_zero_not_nan() {
let eta = 1000.0_f64;
assert!(
eta.exp().is_infinite(),
"test premise: exp(1000) overflows to +∞"
);
assert!(
(f64::INFINITY * 0.0).is_nan(),
"test premise: the naive product is NaN"
);
let (cum, hazard) = royston_parmar_survival_hazard_components(eta, 0.0)
.expect("saturated + flat boundary must be valid");
assert!(cum.is_infinite() && cum > 0.0, "cum saturates to +∞");
assert_eq!(hazard, 0.0, "hazard at a flat boundary is 0, never NaN");
}
#[test]
fn royston_parmar_hazard_propagates_saturation_as_infinity() {
let eta = 1000.0_f64;
let eta_t = 0.5_f64;
assert!(eta.exp().is_infinite(), "test premise: exp(1000) overflows");
let (cum, hazard) = royston_parmar_survival_hazard_components(eta, eta_t)
.expect("saturated RP fit must yield a result, not an error");
assert!(cum.is_infinite() && cum > 0.0, "expected +∞ cum, got {cum}");
assert!(
hazard.is_infinite() && hazard > 0.0,
"expected +∞ hazard, got {hazard}"
);
let survival = (-cum).exp().clamp(0.0, 1.0);
assert_eq!(survival, 0.0, "saturated cum_hazard must give survival 0");
}
#[test]
fn royston_parmar_hazard_rejects_nan_eta() {
let err = royston_parmar_survival_hazard_components(f64::NAN, 0.5)
.expect_err("NaN eta should be invalid");
assert!(
err.to_string()
.contains("invalid log-cumulative-hazard derivative")
);
}
#[test]
fn royston_parmar_hazard_left_tail_collapses_to_zero() {
let eta = -1000.0_f64;
let eta_t = 2.0_f64;
assert_eq!(eta.exp(), 0.0, "test premise: exp(-1000) underflows to 0");
let (cum, hazard) = royston_parmar_survival_hazard_components(eta, eta_t)
.expect("RP left tail must remain valid");
assert_eq!(
cum, 0.0,
"left-tail cum_hazard should underflow to 0, got {cum}"
);
assert_eq!(
hazard, 0.0,
"left-tail hazard should underflow to 0, got {hazard}"
);
let survival = (-cum).exp().clamp(0.0, 1.0);
assert_eq!(survival, 1.0);
}
#[test]
fn probit_survival_hazard_left_tail_collapses_to_zero() {
let eta = -40.0_f64;
let eta_t = 1.5_f64;
let (cum, hazard) =
probit_survival_hazard_components(eta, eta_t).expect("left tail must remain valid");
assert!(
(0.0..1e-300).contains(&cum),
"left-tail cum should be ~0, got {cum}"
);
assert_eq!(
hazard, 0.0,
"left-tail hazard should underflow to 0, got {hazard}"
);
}
#[test]
fn location_scale_logit_hazard_is_failure_slope_over_survival() {
let eta = 0.7;
let eta_t = 0.4;
let hazard = location_scale_hazard_component(
eta,
eta_t,
&InverseLink::Standard(StandardLink::Logit),
)
.expect("valid logit hazard");
let failure = 1.0 / (1.0 + (-eta).exp());
assert!((hazard - failure * eta_t).abs() <= 1e-14);
}
#[test]
fn location_scale_cloglog_hazard_matches_log_cumulative_hazard_derivative() {
let eta = 1.5;
let eta_t = 0.2;
let hazard = location_scale_hazard_component(
eta,
eta_t,
&InverseLink::Standard(StandardLink::CLogLog),
)
.expect("valid cloglog hazard");
assert!((hazard - eta.exp() * eta_t).abs() <= 1e-14);
}
#[test]
fn kaplan_meier_censoring_is_right_continuous_step() {
let time = [2.0, 4.0, 6.0, 8.0];
let event = [1.0, 0.0, 1.0, 0.0];
let g = KaplanMeier::fit_censoring(&time, &event);
assert!((g.at(0.0) - 1.0).abs() <= 1e-15);
assert!((g.at(2.0) - 1.0).abs() <= 1e-15);
assert!((g.at(3.999) - 1.0).abs() <= 1e-15);
assert!((g.at(4.0) - 2.0 / 3.0).abs() <= 1e-12);
assert!((g.at(5.0) - 2.0 / 3.0).abs() <= 1e-12);
assert!((g.at(6.0) - 2.0 / 3.0).abs() <= 1e-12);
assert!(g.at(8.0).abs() <= 1e-15);
}
#[test]
fn ipcw_brier_no_censoring_reduces_to_plain_brier() {
let s_pred = [0.3, 0.7, 0.6, 0.2];
let time = [2.0, 8.0, 10.0, 3.0];
let event = [1.0, 1.0, 0.0, 1.0];
let tau = 5.0;
let g = KaplanMeier::fit_censoring(&time, &event);
let bs = ipcw_brier_score(&s_pred, &time, &event, tau, |t| g.at(t)).unwrap();
let expected =
(0.3f64.powi(2) + (1.0 - 0.7f64).powi(2) + (1.0 - 0.6f64).powi(2) + 0.2f64.powi(2))
/ 4.0;
assert!(
(bs - expected).abs() <= 1e-12,
"bs={bs} expected={expected}"
);
}
#[test]
fn ipcw_brier_reweights_by_inverse_censoring_probability() {
let s_pred = [0.4, 0.5, 0.7, 0.8];
let time = [2.0, 4.0, 6.0, 8.0];
let event = [1.0, 0.0, 1.0, 0.0];
let tau = 5.0;
let g = KaplanMeier::fit_censoring(&time, &event);
let bs = ipcw_brier_score(&s_pred, &time, &event, tau, |t| g.at(t)).unwrap();
let expected = (0.16 + 0.0 + 0.135 + 0.06) / 4.0;
assert!(
(bs - expected).abs() <= 1e-12,
"bs={bs} expected={expected}"
);
}
#[test]
fn ipcw_brier_drops_invalid_rows_from_both_numerator_and_denominator() {
let s_pred = [0.3, 0.7, 0.5, 0.5];
let time = [2.0, 8.0, f64::NAN, -1.0];
let event = [1.0, 1.0, 1.0, 0.0];
let g = KaplanMeier::fit_censoring(&time, &event);
let bs = ipcw_brier_score(&s_pred, &time, &event, 5.0, |t| g.at(t)).unwrap();
let expected = (0.3f64.powi(2) + (1.0 - 0.7f64).powi(2)) / 2.0;
assert!(
(bs - expected).abs() <= 1e-12,
"bs={bs} expected={expected}"
);
}
#[test]
fn integrated_ipcw_brier_of_constant_brier_is_that_constant() {
let time = [2.0, 8.0, 10.0, 3.0];
let event = [1.0, 1.0, 0.0, 1.0];
let grid = [0.0, 1.0, 2.5, 4.0, 6.0];
let col = [0.3, 0.7, 0.6, 0.2];
let mut surv = Array2::<f64>::zeros((4, grid.len()));
for k in 0..grid.len() {
for i in 0..4 {
surv[[i, k]] = col[i];
}
}
let g = KaplanMeier::fit_censoring(&time, &event);
let per_time = ipcw_brier_score(&col, &time, &event, grid[2], |t| g.at(t)).unwrap();
let mut oracle_pts = Vec::new();
for k in 0..grid.len() {
oracle_pts.push((
grid[k],
ipcw_brier_score(&col, &time, &event, grid[k], |t| g.at(t)).unwrap(),
));
}
let mut integral = 0.0;
for w in oracle_pts.windows(2) {
integral += 0.5 * (w[0].1 + w[1].1) * (w[1].0 - w[0].0);
}
let oracle = integral / (grid[grid.len() - 1] - grid[0]);
let ibs =
integrated_ipcw_brier_score(surv.view(), &time, &event, &grid, f64::INFINITY, |t| {
g.at(t)
})
.unwrap();
assert!((ibs - oracle).abs() <= 1e-12, "ibs={ibs} oracle={oracle}");
assert!(per_time >= 0.0);
}
#[test]
fn integrated_ipcw_brier_respects_the_horizon_cutoff() {
let time = [2.0, 8.0, 10.0, 3.0];
let event = [1.0, 1.0, 0.0, 1.0];
let grid = [0.0, 2.0, 4.0, 100.0];
let col = [0.3, 0.7, 0.6, 0.2];
let mut surv = Array2::<f64>::zeros((4, grid.len()));
for k in 0..grid.len() {
for i in 0..4 {
surv[[i, k]] = col[i];
}
}
let g = KaplanMeier::fit_censoring(&time, &event);
let restricted =
integrated_ipcw_brier_score(surv.view(), &time, &event, &grid, 5.0, |t| g.at(t))
.unwrap();
let full =
integrated_ipcw_brier_score(surv.view(), &time, &event, &grid, f64::INFINITY, |t| {
g.at(t)
})
.unwrap();
assert!(
(restricted - full).abs() > 1e-3,
"horizon cutoff had no effect: restricted={restricted} full={full}"
);
}
#[test]
fn integrated_ipcw_brier_rejects_malformed_grids() {
let time = [2.0, 8.0];
let event = [1.0, 0.0];
let surv = Array2::<f64>::from_elem((2, 3), 0.5);
let g = KaplanMeier::fit_censoring(&time, &event);
let bad = [0.0, 2.0, 1.0];
assert!(
integrated_ipcw_brier_score(surv.view(), &time, &event, &bad, f64::INFINITY, |t| g
.at(t))
.is_none()
);
let short = [0.0, 1.0];
assert!(
integrated_ipcw_brier_score(surv.view(), &time, &event, &short, f64::INFINITY, |t| g
.at(t))
.is_none()
);
}
}