Skip to main content

gam_models/fit_orchestration/materialize/
survival_time.rs

1use super::*;
2
3pub struct PreparedSurvivalTimeStack {
4    pub eta_offset_entry: Array1<f64>,
5    pub eta_offset_exit: Array1<f64>,
6    pub derivative_offset_exit: Array1<f64>,
7    pub unloaded_mass_entry: Array1<f64>,
8    pub unloaded_mass_exit: Array1<f64>,
9    pub unloaded_hazard_exit: Array1<f64>,
10    pub time_design_entry: gam_linalg::matrix::DesignMatrix,
11    pub time_design_exit: gam_linalg::matrix::DesignMatrix,
12    pub time_design_derivative_exit: gam_linalg::matrix::DesignMatrix,
13    pub time_penalties: Vec<Array2<f64>>,
14    pub time_nullspace_dims: Vec<usize>,
15    pub timewiggle_build: Option<crate::survival::construction::SurvivalTimeWiggleBuild>,
16    pub timewiggle_block: Option<TimeWiggleBlockInput>,
17}
18
19pub fn prepare_survival_time_stack(
20    age_entry: &Array1<f64>,
21    age_exit: &Array1<f64>,
22    baseline_cfg: &crate::survival::construction::SurvivalBaselineConfig,
23    likelihood_mode: SurvivalLikelihoodMode,
24    inverse_link: Option<&InverseLink>,
25    time_anchor: f64,
26    derivative_guard: f64,
27    time_build: &crate::survival::construction::SurvivalTimeBuildOutput,
28    effective_timewiggle: Option<&LinkWiggleFormulaSpec>,
29    latent_loading: Option<crate::survival::lognormal_kernel::HazardLoading>,
30) -> Result<PreparedSurvivalTimeStack, String> {
31    let (
32        mut eta_offset_entry,
33        mut eta_offset_exit,
34        mut derivative_offset_exit,
35        unloaded_mass_entry,
36        unloaded_mass_exit,
37        unloaded_hazard_exit,
38    ) = if let Some(loading) = latent_loading {
39        let offsets =
40            build_latent_survival_baseline_offsets(age_entry, age_exit, baseline_cfg, loading)?;
41        (
42            offsets.loaded_eta_entry,
43            offsets.loaded_eta_exit,
44            offsets.loaded_derivative_exit,
45            offsets.unloaded_mass_entry,
46            offsets.unloaded_mass_exit,
47            offsets.unloaded_hazard_exit,
48        )
49    } else {
50        // Baseline-hazard barrier conditioning for the marginal-slope likelihood
51        // (gam#797). That likelihood carries `-d·log(qd1)`, a log-barrier on the
52        // baseline-hazard time derivative `qd1 = X_d·β_time + derivative_offset`.
53        // The default `baseline-target=linear` is DEGENERATE for this barrier:
54        // `evaluate_survival_baseline` returns `(0, 0)` for Linear, so the offset
55        // collapses to `derivative_guard` (1e-6) and the I-spline time seed starts
56        // at `qd1 ≈ 1e-6` — exactly ON the barrier boundary, where the
57        // self-concordant Newton step is `∝ qd1` (intrinsically ~1e-4), the
58        // barrier gradient/Hessian are ~1e6 / ~1e12, and the inner joint-Newton
59        // crawls and never reaches the data-scale baseline within the cycle
60        // budget — every outer seed is rejected and the fit hard-fails.
61        //
62        // Condition the COLD START by building the baseline OFFSET from a fixed,
63        // data-seeded Weibull (scale = mean positive exit time, shape = 1) instead
64        // of the zero-derivative Linear baseline, but ONLY for the offset: the
65        // outer `baseline_cfg.target` stays `Linear`, so the
66        // `baseline_cfg.target != Linear` optimize gate
67        // (the gradient baseline optimizers) never fires and no baseline-shape
68        // search is introduced. With shape = 1 the Weibull baseline-hazard
69        // derivative is `1/age_exit` (the natural data hazard scale), so the seed
70        // starts with `qd1` at O(1/T) interior — barrier gradient O(10-10²),
71        // comparable to the marginal/logslope blocks — and `β_time ≈ 0`. This
72        // changes only the STARTING point / offset split: the I-spline still learns
73        // the data-driven deviation from this parametric baseline (the converged
74        // fitted hazard is the same flexible family), so the fix is a pure
75        // preconditioning of the cold start. Gated to MarginalSlope with a Linear
76        // target so every other Linear-baseline survival path is byte-unchanged.
77        let marginal_slope_offset_cfg;
78        let offset_cfg = if likelihood_mode == SurvivalLikelihoodMode::MarginalSlope {
79            marginal_slope_offset_cfg =
80                crate::survival::construction::survival_marginal_slope_offset_baseline_config(
81                    age_exit,
82                    baseline_cfg,
83                );
84            &marginal_slope_offset_cfg
85        } else {
86            baseline_cfg
87        };
88        let (eta_offset_entry, eta_offset_exit, derivative_offset_exit) =
89            build_survival_time_offsets_for_likelihood(
90                age_entry,
91                age_exit,
92                offset_cfg,
93                likelihood_mode,
94                inverse_link,
95            )?;
96        let n = age_entry.len();
97        (
98            eta_offset_entry,
99            eta_offset_exit,
100            derivative_offset_exit,
101            Array1::zeros(n),
102            Array1::zeros(n),
103            Array1::zeros(n),
104        )
105    };
106    add_survival_time_derivative_guard_offset(
107        age_entry,
108        age_exit,
109        time_anchor,
110        derivative_guard,
111        &mut eta_offset_entry,
112        &mut eta_offset_exit,
113        &mut derivative_offset_exit,
114    )?;
115    let timewiggle_build = if let Some(cfg) = effective_timewiggle {
116        Some(build_survival_timewiggle_from_baseline(
117            &eta_offset_entry,
118            &eta_offset_exit,
119            &derivative_offset_exit,
120            cfg,
121        )?)
122    } else {
123        None
124    };
125    let mut time_design_entry = time_build.x_entry_time.clone();
126    let mut time_design_exit = time_build.x_exit_time.clone();
127    let mut time_design_derivative_exit = time_build.x_derivative_time.clone();
128    let mut time_penalties = time_build.penalties.clone();
129    let mut time_nullspace_dims = time_build.nullspace_dims.clone();
130    let mut timewiggle_block = None;
131    if let Some(wiggle) = timewiggle_build.as_ref() {
132        let p_base = time_design_exit.ncols();
133        append_zero_tail_columns(
134            &mut time_design_entry,
135            &mut time_design_exit,
136            &mut time_design_derivative_exit,
137            wiggle.ncols,
138        );
139        for (idx, penalty) in wiggle.penalties.iter().enumerate() {
140            let mut embedded = Array2::<f64>::zeros((p_base + wiggle.ncols, p_base + wiggle.ncols));
141            embedded
142                .slice_mut(s![
143                    p_base..p_base + wiggle.ncols,
144                    p_base..p_base + wiggle.ncols
145                ])
146                .assign(penalty);
147            time_penalties.push(embedded);
148            time_nullspace_dims.push(wiggle.nullspace_dims.get(idx).copied().unwrap_or(0));
149        }
150        timewiggle_block = Some(TimeWiggleBlockInput {
151            knots: wiggle.knots.clone(),
152            degree: wiggle.degree,
153            ncols: wiggle.ncols,
154        });
155    }
156    Ok(PreparedSurvivalTimeStack {
157        eta_offset_entry,
158        eta_offset_exit,
159        derivative_offset_exit,
160        unloaded_mass_entry,
161        unloaded_mass_exit,
162        unloaded_hazard_exit,
163        time_design_entry,
164        time_design_exit,
165        time_design_derivative_exit,
166        time_penalties,
167        time_nullspace_dims,
168        timewiggle_build,
169        timewiggle_block,
170    })
171}