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