gam_models/fit_orchestration/materialize/
survival_time.rs1use 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 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}