1use crate::model_types::EstimationError;
15use crate::probability::signed_log_sum_exp;
16use crate::quadrature::{
17 IntegratedExpectationMode, QuadratureContext, lognormal_laplace_unit_log_term_shared,
18};
19use serde::{Deserialize, Serialize};
20use std::fmt;
21
22#[derive(Debug, Clone)]
30pub enum LognormalKernelError {
31 InvalidSpec { reason: String },
34}
35
36impl_reason_error_boilerplate! {
37 LognormalKernelError {
38 InvalidSpec,
39 }
40}
41
42#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
46#[serde(rename_all = "kebab-case")]
47pub enum HazardLoading {
48 Full,
50 LoadedVsUnloaded,
55}
56
57#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
71#[serde(tag = "frailty_kind", rename_all = "kebab-case")]
72pub enum FrailtySpec {
73 #[default]
75 None,
76 GaussianShift {
80 sigma_fixed: Option<f64>,
82 },
83 HazardMultiplier {
86 sigma_fixed: Option<f64>,
88 loading: HazardLoading,
90 },
91}
92
93impl FrailtySpec {
94 #[inline]
100 pub fn is_active(&self) -> bool {
101 !matches!(self, Self::None)
102 }
103
104 pub fn validate(&self) -> Result<(), LognormalKernelError> {
106 let (kind, sigma) = match self {
107 Self::None => return Ok(()),
108 Self::GaussianShift { sigma_fixed } => ("GaussianShift", sigma_fixed),
109 Self::HazardMultiplier { sigma_fixed, .. } => ("HazardMultiplier", sigma_fixed),
110 };
111 if let Some(sigma) = sigma
112 && (!sigma.is_finite() || *sigma < 0.0)
113 {
114 return Err(LognormalKernelError::InvalidSpec {
115 reason: format!("{kind} frailty requires a finite fixed sigma >= 0, got {sigma}"),
116 });
117 }
118 Ok(())
119 }
120
121 pub fn resolve_fixed_gaussian_shift(
129 &self,
130 context: &str,
131 ) -> Result<Self, LognormalKernelError> {
132 self.validate()?;
133 match self {
134 Self::None => Ok(Self::None),
135 Self::GaussianShift {
136 sigma_fixed: Some(sigma),
137 } => Ok(Self::GaussianShift {
138 sigma_fixed: Some(*sigma),
139 }),
140 Self::GaussianShift { sigma_fixed: None } => Err(LognormalKernelError::InvalidSpec {
141 reason: format!(
142 "{context} requires a fixed GaussianShift sigma; learnable GaussianShift sigma is not supported"
143 ),
144 }),
145 Self::HazardMultiplier { .. } => Err(LognormalKernelError::InvalidSpec {
146 reason: format!(
147 "{context} requires GaussianShift frailty or no frailty; HazardMultiplier is a distinct hazard-scale model"
148 ),
149 }),
150 }
151 }
152
153 pub fn validate_for_marginal_slope(&self) -> Result<(), String> {
165 self.validate_for_marginal_slope_typed()
166 .map_err(|e| e.to_string())
167 }
168
169 pub fn validate_for_marginal_slope_typed(&self) -> Result<(), LognormalKernelError> {
173 match self {
174 Self::None | Self::GaussianShift { .. } => Ok(()),
175 Self::HazardMultiplier { .. } => Err(LognormalKernelError::InvalidSpec {
176 reason:
177 "HazardMultiplier frailty is not finite-state exact with score_warp/linkwiggle \
178 cubic marginal-slope families. Use GaussianShift frailty (exact probit scaling) \
179 or use the standalone latent-cloglog/latent-survival families instead."
180 .to_string(),
181 }),
182 }
183 }
184}
185
186#[inline]
189fn probit_frailty_scale_components(sigma: f64) -> (f64, f64) {
190 let abs_sigma = sigma.abs();
191 if abs_sigma > 1.0 {
192 let inv = 1.0 / abs_sigma;
193 let denom = 1.0 + inv * inv;
194 (inv / denom.sqrt(), 1.0 / denom)
195 } else {
196 let sigma2 = sigma * sigma;
197 let denom = 1.0 + sigma2;
198 (1.0 / denom.sqrt(), sigma2 / denom)
199 }
200}
201
202#[derive(Clone, Copy, Debug)]
210pub struct ProbitFrailtyScaleJet {
211 pub s: f64,
213 pub alpha: f64,
215 pub ds: f64,
217 pub d2s: f64,
219}
220
221impl ProbitFrailtyScaleJet {
222 pub fn new(sigma: f64) -> Self {
227 let (s, alpha) = probit_frailty_scale_components(sigma);
228 Self {
229 s,
230 alpha,
231 ds: -alpha * s,
232 d2s: alpha * (3.0 * alpha - 2.0) * s,
233 }
234 }
235
236 pub fn from_log_sigma(log_sigma: f64) -> Self {
238 Self::new(log_sigma.exp())
239 }
240}
241
242#[inline]
243fn worst_mode(
244 a: IntegratedExpectationMode,
245 b: IntegratedExpectationMode,
246) -> IntegratedExpectationMode {
247 if a.rank() >= b.rank() { a } else { b }
248}
249
250#[inline]
261fn validate_kernel_inputs(m: f64, mu: f64, sigma: f64) -> Result<(), EstimationError> {
262 if !m.is_finite() || m < 0.0 {
263 crate::bail_invalid_estim!("lognormal kernel requires finite m >= 0, got {m}");
264 }
265 if !mu.is_finite() || !sigma.is_finite() || sigma < 0.0 {
266 crate::bail_invalid_estim!(
267 "lognormal kernel requires finite mu and sigma >= 0, got mu={mu}, sigma={sigma}"
268 );
269 }
270 Ok::<(), _>(())
271}
272
273#[inline]
274pub fn log_kernel_term(
275 quadctx: &QuadratureContext,
276 k: usize,
277 m: f64,
278 mu: f64,
279 sigma: f64,
280) -> Result<(f64, IntegratedExpectationMode), EstimationError> {
281 validate_kernel_inputs(m, mu, sigma)?;
282 let kf = k as f64;
283 let sigma2 = sigma * sigma;
284 if !sigma2.is_finite() {
285 crate::bail_invalid_estim!(
286 "lognormal kernel sigma is outside the finite exact-derivative range: sigma={sigma}"
287 );
288 }
289 let prefix_bound = kf * mu.abs() + 0.5 * kf * kf * sigma2;
290 if !prefix_bound.is_finite() {
291 crate::bail_invalid_estim!(
292 "lognormal kernel prefix is outside the finite exact-derivative range: k={k}, mu={mu}, sigma={sigma}"
293 );
294 }
295 let prefix = kf * mu + 0.5 * kf * kf * sigma2;
296 if m == 0.0 {
297 return Ok((prefix, IntegratedExpectationMode::ExactClosedForm));
298 }
299 let log_m = m.ln();
300 let shifted_bound = mu.abs() + kf * sigma2 + log_m.abs();
301 if !shifted_bound.is_finite() {
302 crate::bail_invalid_estim!(
303 "lognormal kernel shifted location is outside the finite exact-derivative range: k={k}, m={m}, mu={mu}, sigma={sigma}"
304 );
305 }
306 let shifted_mu = mu + kf * sigma2 + log_m;
307 let (log_laplace, mode) = lognormal_laplace_unit_log_term_shared(quadctx, shifted_mu, sigma);
312 Ok((prefix + log_laplace, mode))
313}
314
315#[derive(Clone, Debug)]
317pub struct LogLognormalKernelBundle {
318 pub log_values: Vec<f64>,
319 pub mode: IntegratedExpectationMode,
320}
321
322impl LogLognormalKernelBundle {
323 #[inline]
324 pub fn get(&self, k: usize) -> f64 {
325 self.log_values[k]
326 }
327
328 #[inline]
329 pub fn len(&self) -> usize {
330 self.log_values.len()
331 }
332}
333
334pub fn log_kernel_bundle(
337 quadctx: &QuadratureContext,
338 m: f64,
339 mu: f64,
340 sigma: f64,
341 max_k: usize,
342) -> Result<LogLognormalKernelBundle, EstimationError> {
343 validate_kernel_inputs(m, mu, sigma)?;
344 let mut log_values = Vec::with_capacity(max_k + 1);
345 let sigma2 = sigma * sigma;
346 if !sigma2.is_finite() {
347 crate::bail_invalid_estim!(
348 "lognormal kernel sigma is outside the finite exact-derivative range: sigma={sigma}"
349 );
350 }
351 let max_kf = max_k as f64;
352 let prefix_bound = max_kf * mu.abs() + 0.5 * max_kf * max_kf * sigma2;
353 if !prefix_bound.is_finite() {
354 crate::bail_invalid_estim!(
355 "lognormal kernel bundle prefix is outside the finite exact-derivative range: max_k={max_k}, mu={mu}, sigma={sigma}"
356 );
357 }
358 if m == 0.0 {
359 let mut prefix = 0.0;
360 for k in 0..=max_k {
361 log_values.push(prefix);
362 prefix += mu + (k as f64 + 0.5) * sigma2;
363 }
364 return Ok(LogLognormalKernelBundle {
365 log_values,
366 mode: IntegratedExpectationMode::ExactClosedForm,
367 });
368 }
369
370 let log_m = m.ln();
371 let shifted_bound = mu.abs() + max_kf * sigma2 + log_m.abs();
372 if !shifted_bound.is_finite() {
373 crate::bail_invalid_estim!(
374 "lognormal kernel bundle shifted location is outside the finite exact-derivative range: max_k={max_k}, m={m}, mu={mu}, sigma={sigma}"
375 );
376 }
377 let mut shifted_mu = mu + log_m;
378 let mut prefix = 0.0;
379 let mut mode = IntegratedExpectationMode::ExactClosedForm;
380 for k in 0..=max_k {
381 let (log_laplace, val_mode) =
382 lognormal_laplace_unit_log_term_shared(quadctx, shifted_mu, sigma);
383 log_values.push(if log_laplace.is_finite() {
384 prefix + log_laplace
385 } else {
386 f64::NEG_INFINITY
387 });
388 mode = worst_mode(mode, val_mode);
389 prefix += mu + (k as f64 + 0.5) * sigma2;
390 shifted_mu += sigma2;
391 }
392 Ok(LogLognormalKernelBundle { log_values, mode })
393}
394
395pub fn kernel_ratio_jet(
405 log_bundle: &LogLognormalKernelBundle,
406 k: usize,
407 m: f64,
408 order: usize,
409) -> [f64; 5] {
410 let kf = k as f64;
411 let log_k0 = log_bundle.get(k);
412
413 let mut rk = [0.0f64; 5]; for r in 1..=order.min(4) {
418 let delta = log_bundle.get(k + r) - log_k0;
419 rk[r] = if delta.is_finite() {
420 delta.exp()
421 } else if delta > 0.0 {
422 f64::INFINITY
423 } else {
424 0.0
425 };
426 }
427
428 let mut jet = [0.0; 5];
429 jet[0] = 1.0;
430
431 if order >= 1 {
432 jet[1] = kf - m * rk[1];
433 }
434 if order >= 2 {
435 jet[2] = kf * kf - (2.0 * kf + 1.0) * m * rk[1] + m * m * rk[2];
436 }
437 if order >= 3 {
438 jet[3] = kf * kf * kf - (3.0 * kf * kf + 3.0 * kf + 1.0) * m * rk[1]
439 + 3.0 * (kf + 1.0) * m * m * rk[2]
440 - m * m * m * rk[3];
441 }
442 if order >= 4 {
443 let k2 = kf * kf;
444 let k3 = k2 * kf;
445 let k4 = k3 * kf;
446 let m2 = m * m;
447 let m3 = m2 * m;
448 let m4 = m3 * m;
449 jet[4] = k4 - (4.0 * k3 + 6.0 * k2 + 4.0 * kf + 1.0) * m * rk[1]
450 + (6.0 * k2 + 12.0 * kf + 7.0) * m2 * rk[2]
451 - (4.0 * kf + 6.0) * m3 * rk[3]
452 + m4 * rk[4];
453 }
454
455 jet
456}
457
458pub use crate::quadrature::{
464 LatentCLogLogJet5, latent_cloglog_inverse_link_jet, latent_cloglog_jet5,
465};
466
467#[derive(Clone, Copy, Debug)]
471pub struct KernelSumTerm {
472 pub coeff: f64,
474 pub k: usize,
476 pub m: f64,
478}
479
480#[derive(Clone, Copy, Debug)]
494pub struct LogKernelSumJet {
495 pub value: f64,
497 pub d1: f64,
499 pub d2: f64,
501 pub d3: f64,
503 pub d4: f64,
505 pub mode: IntegratedExpectationMode,
506}
507
508impl LogKernelSumJet {
509 #[inline]
510 fn non_positive(mode: IntegratedExpectationMode) -> Self {
511 Self {
512 value: f64::NEG_INFINITY,
513 d1: 0.0,
514 d2: 0.0,
515 d3: 0.0,
516 d4: 0.0,
517 mode,
518 }
519 }
520
521 #[inline]
522 fn from_log_value_and_ratios(
523 value: f64,
524 ratio: [f64; 5],
525 mode: IntegratedExpectationMode,
526 ) -> Self {
527 let r1 = ratio[1];
528 let r2 = ratio[2];
529 let r3 = ratio[3];
530 let r4 = ratio[4];
531 Self {
532 value,
533 d1: r1,
534 d2: r2 - r1 * r1,
535 d3: r3 - 3.0 * r1 * r2 + 2.0 * r1 * r1 * r1,
536 d4: r4 - 4.0 * r1 * r3 - 3.0 * r2 * r2 + 12.0 * r1 * r1 * r2 - 6.0 * r1.powi(4),
537 mode,
538 }
539 }
540
541 #[inline]
542 fn term_log_mag_and_ratio(
543 bundle: &LogLognormalKernelBundle,
544 term: KernelSumTerm,
545 ) -> (f64, [f64; 5]) {
546 (
547 term.coeff.abs().ln() + bundle.get(term.k),
548 kernel_ratio_jet(bundle, term.k, term.m, 4),
551 )
552 }
553
554 fn evaluate_two_terms(
555 quadctx: &QuadratureContext,
556 t0: KernelSumTerm,
557 t1: KernelSumTerm,
558 mu: f64,
559 sigma: f64,
560 ) -> Result<Self, EstimationError> {
561 let max_k_needed = t0.k.max(t1.k) + 4;
562 let bundle0 = log_kernel_bundle(quadctx, t0.m, mu, sigma, max_k_needed)?;
563 let mut overall_mode = bundle0.mode;
564 let bundle1_owned = if (t0.m - t1.m).abs() < 1e-300 {
565 None
566 } else {
567 let bundle1 = log_kernel_bundle(quadctx, t1.m, mu, sigma, max_k_needed)?;
568 overall_mode = worst_mode(overall_mode, bundle1.mode);
569 Some(bundle1)
570 };
571 let bundle1 = bundle1_owned.as_ref().unwrap_or(&bundle0);
572
573 let (log_mag0, ratio0) = Self::term_log_mag_and_ratio(&bundle0, t0);
574 let (log_mag1, ratio1) = Self::term_log_mag_and_ratio(bundle1, t1);
575 let log_mags = [log_mag0, log_mag1];
576 let signs = [t0.coeff.signum(), t1.coeff.signum()];
577 let (log_s, sign_s) = signed_log_sum_exp(&log_mags, &signs);
578 if !log_s.is_finite() || sign_s <= 0.0 {
579 return Ok(Self::non_positive(overall_mode));
580 }
581
582 let w0 = sign_s * signs[0] * (log_mag0 - log_s).exp();
583 let w1 = sign_s * signs[1] * (log_mag1 - log_s).exp();
584 let wr1 = w0 * ratio0[1] + w1 * ratio1[1];
585 let wr2 = w0 * ratio0[2] + w1 * ratio1[2];
586 let wr3 = w0 * ratio0[3] + w1 * ratio1[3];
587 let wr4 = w0 * ratio0[4] + w1 * ratio1[4];
588
589 Ok(Self {
590 value: log_s,
591 d1: wr1,
592 d2: wr2 - wr1 * wr1,
593 d3: wr3 - 3.0 * wr1 * wr2 + 2.0 * wr1 * wr1 * wr1,
594 d4: wr4 - 4.0 * wr1 * wr3 - 3.0 * wr2 * wr2 + 12.0 * wr1 * wr1 * wr2
595 - 6.0 * wr1.powi(4),
596 mode: overall_mode,
597 })
598 }
599
600 pub fn single_term(
605 quadctx: &QuadratureContext,
606 k: usize,
607 m: f64,
608 mu: f64,
609 sigma: f64,
610 ) -> Result<Self, EstimationError> {
611 let max_k_needed = k + 4;
612 let lb = log_kernel_bundle(quadctx, m, mu, sigma, max_k_needed)?;
613 Ok(Self::from_log_value_and_ratios(
614 lb.get(k),
615 kernel_ratio_jet(&lb, k, m, 4),
616 lb.mode,
617 ))
618 }
619
620 pub fn evaluate(
634 quadctx: &QuadratureContext,
635 terms: &[KernelSumTerm],
636 mu: f64,
637 sigma: f64,
638 ) -> Result<Self, EstimationError> {
639 if terms.is_empty() {
640 crate::bail_invalid_estim!("KernelSumJet requires at least one term");
643 }
644
645 if terms.len() == 1 {
647 let t = &terms[0];
648 if t.coeff <= 0.0 {
649 return Ok(Self::non_positive(
653 IntegratedExpectationMode::ExactClosedForm,
654 ));
655 }
656 let jet = Self::single_term(quadctx, t.k, t.m, mu, sigma)?;
657 return Ok(Self {
658 value: t.coeff.ln() + jet.value,
659 d1: jet.d1,
660 d2: jet.d2,
661 d3: jet.d3,
662 d4: jet.d4,
663 mode: jet.mode,
664 });
665 }
666 if terms.len() == 2 {
667 return Self::evaluate_two_terms(quadctx, terms[0], terms[1], mu, sigma);
668 }
669
670 let max_k_needed = terms.iter().map(|t| t.k).max().unwrap_or(0) + 4;
671
672 let mut log_bundles: Vec<(f64, LogLognormalKernelBundle)> = Vec::with_capacity(2);
674 let mut overall_mode = IntegratedExpectationMode::ExactClosedForm;
675 for term in terms {
676 if !log_bundles
677 .iter()
678 .any(|(m, _)| (*m - term.m).abs() < 1e-300)
679 {
680 let b = log_kernel_bundle(quadctx, term.m, mu, sigma, max_k_needed)?;
681 overall_mode = worst_mode(overall_mode, b.mode);
682 log_bundles.push((term.m, b));
683 }
684 }
685
686 let get_lb = |m: f64| -> &LogLognormalKernelBundle {
687 &log_bundles
688 .iter()
689 .find(|(bm, _)| (*bm - m).abs() < 1e-300)
690 .unwrap()
691 .1
692 };
693
694 let mut log_mags: Vec<f64> = Vec::with_capacity(terms.len());
696 let mut signs: Vec<f64> = Vec::with_capacity(terms.len());
697 let mut ratios: Vec<[f64; 5]> = Vec::with_capacity(terms.len());
698 for term in terms {
699 let lb = get_lb(term.m);
700 log_mags.push(term.coeff.abs().ln() + lb.get(term.k));
701 signs.push(term.coeff.signum());
702 ratios.push(kernel_ratio_jet(lb, term.k, term.m, 4));
703 }
704
705 let (log_s, sign_s) = signed_log_sum_exp(&log_mags, &signs);
707
708 if !log_s.is_finite() || sign_s <= 0.0 {
709 return Ok(Self::non_positive(overall_mode));
711 }
712
713 let mut wr1 = 0.0;
716 let mut wr2 = 0.0;
717 let mut wr3 = 0.0;
718 let mut wr4 = 0.0;
719 for i in 0..terms.len() {
720 let w = sign_s * signs[i] * (log_mags[i] - log_s).exp();
721 wr1 += w * ratios[i][1];
722 wr2 += w * ratios[i][2];
723 wr3 += w * ratios[i][3];
724 wr4 += w * ratios[i][4];
725 }
726
727 Ok(Self {
728 value: log_s,
729 d1: wr1,
730 d2: wr2 - wr1 * wr1,
731 d3: wr3 - 3.0 * wr1 * wr2 + 2.0 * wr1 * wr1 * wr1,
732 d4: wr4 - 4.0 * wr1 * wr3 - 3.0 * wr2 * wr2 + 12.0 * wr1 * wr1 * wr2
733 - 6.0 * wr1.powi(4),
734 mode: overall_mode,
735 })
736 }
737}
738
739#[derive(Clone, Copy, Debug, PartialEq, Eq)]
743pub enum LatentSurvivalEventType {
744 RightCensored,
746 ExactEvent,
748 IntervalCensored,
750}
751
752impl fmt::Display for LatentSurvivalEventType {
753 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
754 match self {
755 Self::RightCensored => write!(f, "right_censored"),
756 Self::ExactEvent => write!(f, "exact_event"),
757 Self::IntervalCensored => write!(f, "interval_censored"),
758 }
759 }
760}
761
762#[derive(Clone, Copy, Debug)]
777pub struct LatentSurvivalRow {
778 pub event_type: LatentSurvivalEventType,
779 pub mass_entry: f64,
782 pub mass_exit: f64,
784 pub mass_left: f64,
786 pub mass_right: f64,
788 pub mass_unloaded_left: f64,
790 pub mass_unloaded_right: f64,
792 pub mass_unloaded_entry: f64,
794 pub mass_unloaded_exit: f64,
796 pub hazard_loaded: f64,
798 pub hazard_unloaded: f64,
800}
801
802impl LatentSurvivalRow {
803 pub fn right_censored(
809 mass_entry: f64,
810 mass_exit: f64,
811 mass_unloaded_entry: f64,
812 mass_unloaded_exit: f64,
813 ) -> Self {
814 Self {
815 event_type: LatentSurvivalEventType::RightCensored,
816 mass_entry,
817 mass_exit,
818 mass_left: 0.0,
819 mass_right: 0.0,
820 mass_unloaded_left: 0.0,
821 mass_unloaded_right: 0.0,
822 mass_unloaded_entry,
823 mass_unloaded_exit,
824 hazard_loaded: 0.0,
825 hazard_unloaded: 0.0,
826 }
827 }
828
829 pub fn exact_event(
831 mass_entry: f64,
832 mass_exit: f64,
833 mass_unloaded_entry: f64,
834 mass_unloaded_exit: f64,
835 hazard_loaded: f64,
836 hazard_unloaded: f64,
837 ) -> Self {
838 Self {
839 event_type: LatentSurvivalEventType::ExactEvent,
840 mass_entry,
841 mass_exit,
842 mass_left: 0.0,
843 mass_right: 0.0,
844 mass_unloaded_left: 0.0,
845 mass_unloaded_right: 0.0,
846 mass_unloaded_entry,
847 mass_unloaded_exit,
848 hazard_loaded,
849 hazard_unloaded,
850 }
851 }
852
853 pub fn interval_censored(
855 mass_entry: f64,
856 mass_left: f64,
857 mass_right: f64,
858 mass_unloaded_entry: f64,
859 mass_unloaded_left: f64,
860 mass_unloaded_right: f64,
861 ) -> Self {
862 Self {
863 event_type: LatentSurvivalEventType::IntervalCensored,
864 mass_entry,
865 mass_exit: 0.0,
866 mass_left,
867 mass_right,
868 mass_unloaded_left,
869 mass_unloaded_right,
870 mass_unloaded_entry,
871 mass_unloaded_exit: 0.0,
872 hazard_loaded: 0.0,
873 hazard_unloaded: 0.0,
874 }
875 }
876
877 pub fn validate(&self) -> Result<(), EstimationError> {
878 let fields = [
879 ("mass_entry", self.mass_entry),
880 ("mass_exit", self.mass_exit),
881 ("mass_left", self.mass_left),
882 ("mass_right", self.mass_right),
883 ("mass_unloaded_left", self.mass_unloaded_left),
884 ("mass_unloaded_right", self.mass_unloaded_right),
885 ("mass_unloaded_entry", self.mass_unloaded_entry),
886 ("mass_unloaded_exit", self.mass_unloaded_exit),
887 ("hazard_loaded", self.hazard_loaded),
888 ("hazard_unloaded", self.hazard_unloaded),
889 ];
890 for (name, value) in fields {
891 if !value.is_finite() || value < 0.0 {
892 crate::bail_invalid_estim!(
893 "latent survival row has invalid {name}={value}; expected a finite non-negative value"
894 );
895 }
896 }
897
898 match self.event_type {
899 LatentSurvivalEventType::RightCensored => {
900 if self.mass_exit < self.mass_entry {
901 crate::bail_invalid_estim!(
902 "latent survival right-censored row requires mass_exit >= mass_entry, got {} < {}",
903 self.mass_exit,
904 self.mass_entry
905 );
906 }
907 if self.mass_unloaded_exit < self.mass_unloaded_entry {
908 crate::bail_invalid_estim!(
909 "latent survival right-censored row requires unloaded exit mass >= unloaded entry mass, got {} < {}",
910 self.mass_unloaded_exit,
911 self.mass_unloaded_entry
912 );
913 }
914 if self.mass_left > 0.0
915 || self.mass_right > 0.0
916 || self.mass_unloaded_left > 0.0
917 || self.mass_unloaded_right > 0.0
918 || self.hazard_loaded > 0.0
919 || self.hazard_unloaded > 0.0
920 {
921 crate::bail_invalid_estim!("latent survival right-censored row cannot carry interval masses or event hazards"
922 .to_string(),);
923 }
924 }
925 LatentSurvivalEventType::ExactEvent => {
926 if self.mass_exit < self.mass_entry {
927 crate::bail_invalid_estim!(
928 "latent survival exact-event row requires mass_exit >= mass_entry, got {} < {}",
929 self.mass_exit,
930 self.mass_entry
931 );
932 }
933 if self.mass_unloaded_exit < self.mass_unloaded_entry {
934 crate::bail_invalid_estim!(
935 "latent survival exact-event row requires unloaded exit mass >= unloaded entry mass, got {} < {}",
936 self.mass_unloaded_exit,
937 self.mass_unloaded_entry
938 );
939 }
940 if self.mass_left > 0.0
941 || self.mass_right > 0.0
942 || self.mass_unloaded_left > 0.0
943 || self.mass_unloaded_right > 0.0
944 {
945 crate::bail_invalid_estim!(
946 "latent survival exact-event row cannot carry interval masses"
947 );
948 }
949 if self.hazard_loaded == 0.0 && self.hazard_unloaded == 0.0 {
950 crate::bail_invalid_estim!("latent survival exact-event row requires a positive loaded or unloaded hazard"
951 .to_string(),);
952 }
953 }
954 LatentSurvivalEventType::IntervalCensored => {
955 if self.mass_left < self.mass_entry || self.mass_right < self.mass_left {
956 crate::bail_invalid_estim!(
957 "latent survival interval row requires mass_entry <= mass_left <= mass_right, got entry={}, left={}, right={}",
958 self.mass_entry,
959 self.mass_left,
960 self.mass_right
961 );
962 }
963 if self.mass_unloaded_left < self.mass_unloaded_entry
964 || self.mass_unloaded_right < self.mass_unloaded_left
965 {
966 crate::bail_invalid_estim!(
967 "latent survival interval row requires unloaded_entry <= unloaded_left <= unloaded_right, got entry={}, left={}, right={}",
968 self.mass_unloaded_entry,
969 self.mass_unloaded_left,
970 self.mass_unloaded_right
971 );
972 }
973 if self.mass_exit > 0.0
974 || self.mass_unloaded_exit > 0.0
975 || self.hazard_loaded > 0.0
976 || self.hazard_unloaded > 0.0
977 {
978 crate::bail_invalid_estim!(
979 "latent survival interval row cannot carry exit masses or event hazards"
980 .to_string(),
981 );
982 }
983 }
984 }
985
986 Ok(())
987 }
988}
989
990fn exact_event_kernel_jet(
991 quadctx: &QuadratureContext,
992 row: &LatentSurvivalRow,
993 mu: f64,
994 sigma: f64,
995) -> Result<LogKernelSumJet, EstimationError> {
996 if row.hazard_loaded < 0.0 || row.hazard_unloaded < 0.0 {
997 crate::bail_invalid_estim!(
998 "latent survival exact-event hazards must be non-negative, got loaded={} unloaded={}",
999 row.hazard_loaded,
1000 row.hazard_unloaded
1001 );
1002 }
1003 match (row.hazard_unloaded > 0.0, row.hazard_loaded > 0.0) {
1004 (true, true) => {
1005 let terms = [
1006 KernelSumTerm {
1007 coeff: row.hazard_unloaded,
1008 k: 0,
1009 m: row.mass_exit,
1010 },
1011 KernelSumTerm {
1012 coeff: row.hazard_loaded,
1013 k: 1,
1014 m: row.mass_exit,
1015 },
1016 ];
1017 LogKernelSumJet::evaluate(quadctx, &terms, mu, sigma)
1018 }
1019 (true, false) => {
1020 let jet = LogKernelSumJet::single_term(quadctx, 0, row.mass_exit, mu, sigma)?;
1021 Ok(LogKernelSumJet {
1022 value: row.hazard_unloaded.ln() + jet.value,
1023 d1: jet.d1,
1024 d2: jet.d2,
1025 d3: jet.d3,
1026 d4: jet.d4,
1027 mode: jet.mode,
1028 })
1029 }
1030 (false, true) => {
1031 let jet = LogKernelSumJet::single_term(quadctx, 1, row.mass_exit, mu, sigma)?;
1032 Ok(LogKernelSumJet {
1033 value: row.hazard_loaded.ln() + jet.value,
1034 d1: jet.d1,
1035 d2: jet.d2,
1036 d3: jet.d3,
1037 d4: jet.d4,
1038 mode: jet.mode,
1039 })
1040 }
1041 (false, false) => Err(EstimationError::InvalidInput(
1042 "latent survival exact-event row requires a positive loaded or unloaded hazard"
1043 .to_string(),
1044 )),
1045 }
1046}
1047
1048#[derive(Clone, Copy, Debug)]
1055pub struct LatentSurvivalRowJet {
1056 pub log_lik: f64,
1057 pub score: f64,
1058 pub neg_hessian: f64,
1059 pub d3: f64,
1060 pub score_log_sigma: f64,
1061 pub neg_hessian_log_sigma: f64,
1062}
1063
1064#[inline]
1065fn log_sigma_score_from_log_sum(jet: &LogKernelSumJet, sigma: f64) -> f64 {
1066 let sigma2 = sigma * sigma;
1067 sigma2 * (jet.d2 + jet.d1 * jet.d1)
1068}
1069
1070#[inline]
1071fn log_sigma_neg_hessian_from_log_sum(jet: &LogKernelSumJet, sigma: f64) -> f64 {
1072 let sigma2 = sigma * sigma;
1073 let sigma4 = sigma2 * sigma2;
1074 let d1 = jet.d1;
1075 let d2 = jet.d2;
1076 let d3 = jet.d3;
1077 let d4 = jet.d4;
1078 let s2_over_s = d2 + d1 * d1;
1079 let s4_over_s_minus_s2_sq = d4 + 4.0 * d1 * d3 + 2.0 * d2 * d2 + 4.0 * d1 * d1 * d2;
1084 -(2.0 * sigma2 * s2_over_s + sigma4 * s4_over_s_minus_s2_sq)
1085}
1086
1087impl LatentSurvivalRowJet {
1088 pub fn evaluate(
1089 quadctx: &QuadratureContext,
1090 row: &LatentSurvivalRow,
1091 mu: f64,
1092 sigma: f64,
1093 ) -> Result<Self, EstimationError> {
1094 row.validate()?;
1095 match row.event_type {
1096 LatentSurvivalEventType::RightCensored => Self::right_censored(quadctx, mu, sigma, row),
1097 LatentSurvivalEventType::ExactEvent => Self::exact_event(quadctx, mu, sigma, row),
1098 LatentSurvivalEventType::IntervalCensored => {
1099 Self::interval_censored(quadctx, mu, sigma, row)
1100 }
1101 }
1102 }
1103
1104 fn right_censored(
1112 quadctx: &QuadratureContext,
1113 mu: f64,
1114 sigma: f64,
1115 row: &LatentSurvivalRow,
1116 ) -> Result<Self, EstimationError> {
1117 let has_unloaded =
1118 row.mass_unloaded_exit.abs() > 1e-300 || row.mass_unloaded_entry.abs() > 1e-300;
1119
1120 let mass_exit_loaded = row.mass_exit;
1124 let mass_entry_loaded = row.mass_entry;
1125
1126 let unloaded_offset = if has_unloaded {
1128 -row.mass_unloaded_exit + row.mass_unloaded_entry
1129 } else {
1130 0.0
1131 };
1132
1133 let num = LogKernelSumJet::single_term(quadctx, 0, mass_exit_loaded, mu, sigma)?;
1134 if mass_entry_loaded > 1e-300 {
1135 let den = LogKernelSumJet::single_term(quadctx, 0, mass_entry_loaded, mu, sigma)?;
1136 Ok(Self {
1137 log_lik: unloaded_offset + num.value - den.value,
1138 score: num.d1 - den.d1,
1139 neg_hessian: -(num.d2 - den.d2),
1140 d3: num.d3 - den.d3,
1141 score_log_sigma: log_sigma_score_from_log_sum(&num, sigma)
1142 - log_sigma_score_from_log_sum(&den, sigma),
1143 neg_hessian_log_sigma: log_sigma_neg_hessian_from_log_sum(&num, sigma)
1144 - log_sigma_neg_hessian_from_log_sum(&den, sigma),
1145 })
1146 } else {
1147 Ok(Self {
1148 log_lik: unloaded_offset + num.value,
1149 score: num.d1,
1150 neg_hessian: -num.d2,
1151 d3: num.d3,
1152 score_log_sigma: log_sigma_score_from_log_sum(&num, sigma),
1153 neg_hessian_log_sigma: log_sigma_neg_hessian_from_log_sum(&num, sigma),
1154 })
1155 }
1156 }
1157
1158 fn exact_event(
1162 quadctx: &QuadratureContext,
1163 mu: f64,
1164 sigma: f64,
1165 row: &LatentSurvivalRow,
1166 ) -> Result<Self, EstimationError> {
1167 let unloaded_offset =
1168 if row.mass_unloaded_exit.abs() > 1e-300 || row.mass_unloaded_entry.abs() > 1e-300 {
1169 -row.mass_unloaded_exit + row.mass_unloaded_entry
1170 } else {
1171 0.0
1172 };
1173 let num = exact_event_kernel_jet(quadctx, row, mu, sigma)?;
1174
1175 if row.mass_entry > 1e-300 {
1176 let den = LogKernelSumJet::single_term(quadctx, 0, row.mass_entry, mu, sigma)?;
1177 Ok(Self {
1178 log_lik: unloaded_offset + num.value - den.value,
1179 score: num.d1 - den.d1,
1180 neg_hessian: -(num.d2 - den.d2),
1181 d3: num.d3 - den.d3,
1182 score_log_sigma: log_sigma_score_from_log_sum(&num, sigma)
1183 - log_sigma_score_from_log_sum(&den, sigma),
1184 neg_hessian_log_sigma: log_sigma_neg_hessian_from_log_sum(&num, sigma)
1185 - log_sigma_neg_hessian_from_log_sum(&den, sigma),
1186 })
1187 } else {
1188 Ok(Self {
1189 log_lik: unloaded_offset + num.value,
1190 score: num.d1,
1191 neg_hessian: -num.d2,
1192 d3: num.d3,
1193 score_log_sigma: log_sigma_score_from_log_sum(&num, sigma),
1194 neg_hessian_log_sigma: log_sigma_neg_hessian_from_log_sum(&num, sigma),
1195 })
1196 }
1197 }
1198
1199 fn interval_censored(
1201 quadctx: &QuadratureContext,
1202 mu: f64,
1203 sigma: f64,
1204 row: &LatentSurvivalRow,
1205 ) -> Result<Self, EstimationError> {
1206 let num_terms = [
1207 KernelSumTerm {
1208 coeff: (-row.mass_unloaded_left).exp(),
1209 k: 0,
1210 m: row.mass_left,
1211 },
1212 KernelSumTerm {
1213 coeff: -(-row.mass_unloaded_right).exp(),
1214 k: 0,
1215 m: row.mass_right,
1216 },
1217 ];
1218 let num = LogKernelSumJet::evaluate(quadctx, &num_terms, mu, sigma)?;
1219
1220 if row.mass_entry > 1e-300 {
1221 let den = LogKernelSumJet::single_term(quadctx, 0, row.mass_entry, mu, sigma)?;
1222 Ok(Self {
1223 log_lik: num.value + row.mass_unloaded_entry - den.value,
1224 score: num.d1 - den.d1,
1225 neg_hessian: -(num.d2 - den.d2),
1226 d3: num.d3 - den.d3,
1227 score_log_sigma: log_sigma_score_from_log_sum(&num, sigma)
1228 - log_sigma_score_from_log_sum(&den, sigma),
1229 neg_hessian_log_sigma: log_sigma_neg_hessian_from_log_sum(&num, sigma)
1230 - log_sigma_neg_hessian_from_log_sum(&den, sigma),
1231 })
1232 } else {
1233 Ok(Self {
1234 log_lik: num.value + row.mass_unloaded_entry,
1235 score: num.d1,
1236 neg_hessian: -num.d2,
1237 d3: num.d3,
1238 score_log_sigma: log_sigma_score_from_log_sum(&num, sigma),
1239 neg_hessian_log_sigma: log_sigma_neg_hessian_from_log_sum(&num, sigma),
1240 })
1241 }
1242 }
1243}
1244
1245#[cfg(test)]
1246mod tests {
1247 use super::*;
1248
1249 #[test]
1250 fn fixed_gaussian_shift_resolution_accepts_only_exact_fixed_states() {
1251 assert_eq!(
1252 FrailtySpec::None
1253 .resolve_fixed_gaussian_shift("marginal-slope")
1254 .unwrap(),
1255 FrailtySpec::None
1256 );
1257 let fixed = FrailtySpec::GaussianShift {
1258 sigma_fixed: Some(0.75),
1259 };
1260 assert_eq!(
1261 fixed
1262 .resolve_fixed_gaussian_shift("marginal-slope")
1263 .unwrap(),
1264 fixed
1265 );
1266 assert!(
1267 FrailtySpec::GaussianShift { sigma_fixed: None }
1268 .resolve_fixed_gaussian_shift("marginal-slope")
1269 .is_err()
1270 );
1271 assert!(
1272 FrailtySpec::HazardMultiplier {
1273 sigma_fixed: Some(0.75),
1274 loading: HazardLoading::Full,
1275 }
1276 .resolve_fixed_gaussian_shift("marginal-slope")
1277 .is_err()
1278 );
1279 assert!(
1280 FrailtySpec::GaussianShift {
1281 sigma_fixed: Some(-0.1),
1282 }
1283 .resolve_fixed_gaussian_shift("marginal-slope")
1284 .is_err()
1285 );
1286 assert!(
1287 FrailtySpec::GaussianShift {
1288 sigma_fixed: Some(f64::NAN),
1289 }
1290 .resolve_fixed_gaussian_shift("marginal-slope")
1291 .is_err()
1292 );
1293 }
1294
1295 fn latent_binomial_row_log_lik(
1296 ctx: &QuadratureContext,
1297 eta: f64,
1298 sigma: f64,
1299 y: f64,
1300 weight: f64,
1301 ) -> f64 {
1302 let mu = latent_cloglog_jet5(ctx, eta, sigma)
1303 .expect("latent jet")
1304 .mean;
1305 let mu = mu.clamp(1e-12, 1.0 - 1e-12);
1306 weight * (y * mu.ln() + (1.0 - y) * (1.0 - mu).ln())
1307 }
1308
1309 #[test]
1310 fn kernel_ratio_jet_d1_fd_check() {
1311 let ctx = QuadratureContext::new();
1312 let mu = 0.3;
1313 let sigma = 0.5;
1314 let m = 1.0;
1315 let k = 0usize;
1316 let h = 1e-5;
1317
1318 let bundle = log_kernel_bundle(&ctx, m, mu, sigma, k + 4).unwrap();
1319 let log_k = bundle.get(k);
1320 let ratios = kernel_ratio_jet(&bundle, k, m, 2);
1321 let kc = log_k.exp();
1322 let d1 = kc * ratios[1];
1323 let d2 = kc * ratios[2];
1324
1325 let kp = log_kernel_term(&ctx, k, m, mu + h, sigma).unwrap().0.exp();
1326 let km = log_kernel_term(&ctx, k, m, mu - h, sigma).unwrap().0.exp();
1327 let fd_d1 = (kp - km) / (2.0 * h);
1328 assert!(
1329 (d1 - fd_d1).abs() / fd_d1.abs().max(1e-15) < 1e-4,
1330 "d1: jet={d1}, fd={fd_d1}",
1331 );
1332
1333 let fd_d2 = (kp - 2.0 * kc + km) / (h * h);
1334 assert!(
1335 (d2 - fd_d2).abs() / fd_d2.abs().max(1e-15) < 1e-3,
1336 "d2: jet={d2}, fd={fd_d2}",
1337 );
1338 }
1339
1340 #[test]
1341 fn survival_right_censored_score_fd() {
1342 let ctx = QuadratureContext::new();
1343 let mu = -0.5;
1344 let sigma = 0.3;
1345 let h = 1e-6;
1346 let row = LatentSurvivalRow::right_censored(0.0, 2.0, 0.0, 0.0);
1347 let ll_p = LatentSurvivalRowJet::evaluate(&ctx, &row, mu + h, sigma)
1348 .unwrap()
1349 .log_lik;
1350 let ll_m = LatentSurvivalRowJet::evaluate(&ctx, &row, mu - h, sigma)
1351 .unwrap()
1352 .log_lik;
1353 let fd_score = (ll_p - ll_m) / (2.0 * h);
1354 let jet = LatentSurvivalRowJet::evaluate(&ctx, &row, mu, sigma).unwrap();
1355 assert!(
1356 (jet.score - fd_score).abs() / fd_score.abs().max(1e-15) < 1e-3,
1357 "score={}, fd={fd_score}",
1358 jet.score
1359 );
1360 }
1361
1362 #[test]
1363 fn survival_exact_event_score_fd() {
1364 let ctx = QuadratureContext::new();
1365 let mu = 0.2;
1366 let sigma = 0.5;
1367 let h = 1e-6;
1368 let row = LatentSurvivalRow::exact_event(0.0, 1.5, 0.0, 0.0, (-0.3f64).exp(), 0.0);
1369 let ll_p = LatentSurvivalRowJet::evaluate(&ctx, &row, mu + h, sigma)
1370 .unwrap()
1371 .log_lik;
1372 let ll_m = LatentSurvivalRowJet::evaluate(&ctx, &row, mu - h, sigma)
1373 .unwrap()
1374 .log_lik;
1375 let fd_score = (ll_p - ll_m) / (2.0 * h);
1376 let jet = LatentSurvivalRowJet::evaluate(&ctx, &row, mu, sigma).unwrap();
1377 assert!(
1378 (jet.score - fd_score).abs() / fd_score.abs().max(1e-15) < 1e-3,
1379 "score={}, fd={fd_score}",
1380 jet.score
1381 );
1382 }
1383
1384 #[test]
1385 fn survival_exact_event_loaded_vs_unloaded_score_fd() {
1386 let ctx = QuadratureContext::new();
1387 let mu = -0.1;
1388 let sigma = 0.4;
1389 let h = 1e-6;
1390 let row = LatentSurvivalRow::exact_event(0.3, 1.2, 0.2, 0.6, 0.9, 0.15);
1391 let ll_p = LatentSurvivalRowJet::evaluate(&ctx, &row, mu + h, sigma)
1392 .unwrap()
1393 .log_lik;
1394 let ll_m = LatentSurvivalRowJet::evaluate(&ctx, &row, mu - h, sigma)
1395 .unwrap()
1396 .log_lik;
1397 let fd_score = (ll_p - ll_m) / (2.0 * h);
1398 let jet = LatentSurvivalRowJet::evaluate(&ctx, &row, mu, sigma).unwrap();
1399 assert!(
1400 (jet.score - fd_score).abs() / fd_score.abs().max(1e-15) < 1e-3,
1401 "score={}, fd={fd_score}",
1402 jet.score
1403 );
1404 }
1405
1406 #[test]
1407 fn survival_right_censored_loaded_vs_unloaded_score_fd() {
1408 let ctx = QuadratureContext::new();
1409 let mu = 0.15;
1410 let sigma: f64 = 0.35;
1411 let h = 1e-6;
1412 let row = LatentSurvivalRow::right_censored(0.4, 1.7, 0.1, 0.5);
1413 let ll_p = LatentSurvivalRowJet::evaluate(&ctx, &row, mu + h, sigma)
1414 .unwrap()
1415 .log_lik;
1416 let ll_m = LatentSurvivalRowJet::evaluate(&ctx, &row, mu - h, sigma)
1417 .unwrap()
1418 .log_lik;
1419 let fd_score = (ll_p - ll_m) / (2.0 * h);
1420 let jet = LatentSurvivalRowJet::evaluate(&ctx, &row, mu, sigma).unwrap();
1421 assert!(
1422 (jet.score - fd_score).abs() / fd_score.abs().max(1e-15) < 1e-3,
1423 "score={}, fd={fd_score}",
1424 jet.score
1425 );
1426 }
1427
1428 #[test]
1429 fn survival_interval_censored_score_fd() {
1430 let ctx = QuadratureContext::new();
1431 let mu = 0.0;
1432 let sigma = 0.6;
1433 let h = 1e-6;
1434 let row = LatentSurvivalRow::interval_censored(0.0, 1.0, 2.0, 0.0, 0.0, 0.0);
1435 let ll_p = LatentSurvivalRowJet::evaluate(&ctx, &row, mu + h, sigma)
1436 .unwrap()
1437 .log_lik;
1438 let ll_m = LatentSurvivalRowJet::evaluate(&ctx, &row, mu - h, sigma)
1439 .unwrap()
1440 .log_lik;
1441 let fd_score = (ll_p - ll_m) / (2.0 * h);
1442 let jet = LatentSurvivalRowJet::evaluate(&ctx, &row, mu, sigma).unwrap();
1443 assert!(
1444 (jet.score - fd_score).abs() / fd_score.abs().max(1e-15) < 1e-3,
1445 "score={}, fd={fd_score}",
1446 jet.score
1447 );
1448 }
1449
1450 #[test]
1451 fn survival_interval_censored_neg_hessian_fd() {
1452 let ctx = QuadratureContext::new();
1456 let mu = -0.2;
1457 let sigma = 0.55;
1458 let h = 2e-4;
1459 let row = LatentSurvivalRow::interval_censored(0.0, 0.7, 1.9, 0.0, 0.0, 0.0);
1460 let ll = |m: f64| {
1461 LatentSurvivalRowJet::evaluate(&ctx, &row, m, sigma)
1462 .unwrap()
1463 .log_lik
1464 };
1465 let fd_d2 = (ll(mu + h) - 2.0 * ll(mu) + ll(mu - h)) / (h * h);
1466 let jet = LatentSurvivalRowJet::evaluate(&ctx, &row, mu, sigma).unwrap();
1467 assert!(
1468 (jet.neg_hessian - (-fd_d2)).abs() / fd_d2.abs().max(1e-12) < 1e-2,
1469 "interval neg_hessian={}, fd(-d2)={}",
1470 jet.neg_hessian,
1471 -fd_d2
1472 );
1473 }
1474
1475 #[test]
1476 fn survival_interval_censored_log_sigma_score_fd() {
1477 let ctx = QuadratureContext::new();
1482 let mu = 0.1;
1483 let sigma: f64 = 0.6;
1484 let h = 1e-5;
1485 let row = LatentSurvivalRow::interval_censored(0.0, 0.8, 2.1, 0.0, 0.0, 0.0);
1486 let ll_at = |s: f64| {
1487 LatentSurvivalRowJet::evaluate(&ctx, &row, mu, s)
1488 .unwrap()
1489 .log_lik
1490 };
1491 let fd_dlogsigma =
1493 (ll_at((sigma.ln() + h).exp()) - ll_at((sigma.ln() - h).exp())) / (2.0 * h);
1494 let jet = LatentSurvivalRowJet::evaluate(&ctx, &row, mu, sigma).unwrap();
1495 assert!(
1496 (jet.score_log_sigma - fd_dlogsigma).abs() / fd_dlogsigma.abs().max(1e-12) < 1e-3,
1497 "interval score_log_sigma={}, fd={fd_dlogsigma}",
1498 jet.score_log_sigma
1499 );
1500 }
1501
1502 #[test]
1503 fn log_kernel_single_term_log_sigma_derivatives_match_ghq_reference() {
1504 let ctx = QuadratureContext::new();
1505 let mu = 0.2;
1506 let sigma = 1.0;
1507 let jet = LogKernelSumJet::single_term(&ctx, 0, 1.0, mu, sigma).unwrap();
1508 let ghq = crate::inference::quadrature::cloglog_ghq_derivatives_adaptive(&ctx, mu, sigma);
1509 let survival = (1.0 - ghq.l).max(1e-300);
1510 let survival_sigma_over_survival = -ghq.l_sigma / survival;
1511 let ref_score = sigma * survival_sigma_over_survival;
1512 let ref_neg_hessian = -(ref_score
1513 + sigma
1514 * sigma
1515 * (-ghq.l_sigmasigma / survival - survival_sigma_over_survival.powi(2)));
1516
1517 assert!(
1518 (log_sigma_score_from_log_sum(&jet, sigma) - ref_score).abs()
1519 / ref_score.abs().max(1e-12)
1520 < 1e-4,
1521 "log-sigma score={}, ref={ref_score}",
1522 log_sigma_score_from_log_sum(&jet, sigma)
1523 );
1524 assert!(
1525 (log_sigma_neg_hessian_from_log_sum(&jet, sigma) - ref_neg_hessian).abs()
1526 / ref_neg_hessian.abs().max(1e-12)
1527 < 1e-3,
1528 "log-sigma neg_hessian={}, ref={ref_neg_hessian}",
1529 log_sigma_neg_hessian_from_log_sum(&jet, sigma)
1530 );
1531 }
1532
1533 #[test]
1534 fn log_kernel_sum_jet_single_term_d1_fd() {
1535 let ctx = QuadratureContext::new();
1536 let mu = 0.5;
1537 let sigma = 0.4;
1538 let m = 1.0;
1539 let k = 0usize;
1540 let h = 1e-6;
1541
1542 let jet = LogKernelSumJet::single_term(&ctx, k, m, mu, sigma).unwrap();
1543 let val_p = log_kernel_term(&ctx, k, m, mu + h, sigma).unwrap().0;
1544 let val_m = log_kernel_term(&ctx, k, m, mu - h, sigma).unwrap().0;
1545 let fd_d1 = (val_p - val_m) / (2.0 * h);
1546 assert!(
1547 (jet.d1 - fd_d1).abs() / fd_d1.abs().max(1e-15) < 1e-3,
1548 "d1={}, fd={fd_d1}",
1549 jet.d1
1550 );
1551 }
1552
1553 #[test]
1554 fn log_kernel_sum_jet_single_term_d4_fd() {
1555 let ctx = QuadratureContext::new();
1556 let mu = 0.35;
1557 let sigma = 0.45;
1558 let m = 1.2;
1559 let k = 1usize;
1560 let h = 2e-3;
1561
1562 let jet = LogKernelSumJet::single_term(&ctx, k, m, mu, sigma).unwrap();
1563 let v_pp = log_kernel_term(&ctx, k, m, mu + 2.0 * h, sigma).unwrap().0;
1564 let v_p = log_kernel_term(&ctx, k, m, mu + h, sigma).unwrap().0;
1565 let v_0 = log_kernel_term(&ctx, k, m, mu, sigma).unwrap().0;
1566 let v_m = log_kernel_term(&ctx, k, m, mu - h, sigma).unwrap().0;
1567 let v_mm = log_kernel_term(&ctx, k, m, mu - 2.0 * h, sigma).unwrap().0;
1568 let fd_d4 = (v_mm - 4.0 * v_m + 6.0 * v_0 - 4.0 * v_p + v_pp) / h.powi(4);
1569 assert!(
1570 (jet.d4 - fd_d4).abs() / jet.d4.abs().max(fd_d4.abs()).max(1e-8) < 2e-2,
1571 "d4={}, fd={fd_d4}",
1572 jet.d4
1573 );
1574 }
1575
1576 #[test]
1577 fn latent_cloglog_jet_matches_point_limit_at_zero_sigma() {
1578 let ctx = QuadratureContext::new();
1579 let eta = -0.4;
1580 let jet = latent_cloglog_jet5(&ctx, eta, 0.0).expect("latent jet");
1581 let t = eta.exp();
1582 let d1 = (eta - t).exp();
1583 let d2 = (1.0 - t) * d1;
1584 let d3 = (t * t - 3.0 * t + 1.0) * d1;
1585 let d4 = (-t * t * t + 6.0 * t * t - 7.0 * t + 1.0) * d1;
1586 let d5 = (t.powi(4) - 10.0 * t.powi(3) + 25.0 * t * t - 15.0 * t + 1.0) * d1;
1587 assert!((jet.mean - (1.0 - (-t).exp())).abs() < 1e-12);
1588 assert!((jet.d1 - d1).abs() < 1e-12);
1589 assert!((jet.d2 - d2).abs() < 1e-12);
1590 assert!((jet.d3 - d3).abs() < 1e-12);
1591 assert!((jet.d4 - d4).abs() < 1e-12);
1592 assert!((jet.d5 - d5).abs() < 1e-12);
1593 }
1594
1595 #[test]
1596 fn latent_cloglog_jet_matches_exact_kernel_recurrence() {
1597 let ctx = QuadratureContext::new();
1598 let cases = [(-4.0, 0.15), (-1.2, 0.35), (0.4, 0.6), (1.3, 0.9)];
1599
1600 for (eta, sigma) in cases {
1601 let jet = latent_cloglog_jet5(&ctx, eta, sigma).expect("latent jet");
1602 let bundle = log_kernel_bundle(&ctx, 1.0, eta, sigma, 5).expect("kernel bundle");
1603 let k0 = bundle.get(0);
1604 let k1 = bundle.get(1).exp();
1605 let k2 = bundle.get(2).exp();
1606 let k3 = bundle.get(3).exp();
1607 let k4 = bundle.get(4).exp();
1608 let k5 = bundle.get(5).exp();
1609
1610 let mean = if k0.is_finite() { -k0.exp_m1() } else { 1.0 };
1611 let d1 = k1;
1612 let d2 = k1 - k2;
1613 let d3 = k1 - 3.0 * k2 + k3;
1614 let d4 = k1 - 7.0 * k2 + 6.0 * k3 - k4;
1615 let d5 = k1 - 15.0 * k2 + 25.0 * k3 - 10.0 * k4 + k5;
1616
1617 assert!((jet.mean - mean).abs() < 1e-12);
1618 assert!((jet.d1 - d1).abs() < 1e-12);
1619 assert!((jet.d2 - d2).abs() < 1e-12);
1620 assert!((jet.d3 - d3).abs() < 1e-12);
1621 assert!((jet.d4 - d4).abs() < 1e-12);
1622 assert!((jet.d5 - d5).abs() < 1e-12);
1623 }
1624 }
1625
1626 #[test]
1627 fn latent_cloglog_binomial_row_neg_hessian_matches_fd() {
1628 let ctx = QuadratureContext::new();
1629 let eta = 0.4;
1630 let sigma = 0.6;
1631 let y = 0.35;
1632 let weight = 2.0;
1633 let h = 1e-4;
1634
1635 let jet = latent_cloglog_jet5(&ctx, eta, sigma).expect("latent jet");
1636 let mu = jet.mean.clamp(1e-12, 1.0 - 1e-12);
1637 let ellmu = y / mu - (1.0 - y) / (1.0 - mu);
1638 let ellmumu = -y / (mu * mu) - (1.0 - y) / ((1.0 - mu) * (1.0 - mu));
1639 let neg_hessian = -weight * (ellmumu * jet.d1 * jet.d1 + ellmu * jet.d2);
1640
1641 let ll_minus = latent_binomial_row_log_lik(&ctx, eta - h, sigma, y, weight);
1642 let ll0 = latent_binomial_row_log_lik(&ctx, eta, sigma, y, weight);
1643 let ll_plus = latent_binomial_row_log_lik(&ctx, eta + h, sigma, y, weight);
1644 let neg_hessian_fd = -(ll_plus - 2.0 * ll0 + ll_minus) / (h * h);
1645
1646 let err = (neg_hessian - neg_hessian_fd).abs();
1647 let tol = 2e-5_f64.max(3e-3 * neg_hessian_fd.abs());
1648 assert!(
1649 err <= tol,
1650 "latent cloglog Bernoulli row curvature mismatch: analytic={} fd={}",
1651 neg_hessian,
1652 neg_hessian_fd
1653 );
1654 }
1655}