Skip to main content

gam_models/survival/
lognormal_kernel.rs

1//! Shared analytic kernel for latent-variable families with lognormal structure.
2//!
3//! The kernel object `K_{k,m}(μ, σ) := E[exp(k·U − m·exp(U))]`, where
4//! `U ~ N(μ, σ²)`, is the only special function required by all latent families.
5//!
6//! It satisfies exact μ-recurrences (see [`kernel_ratio_jet`]) and the
7//! corresponding heat-equation σ-identities, so fixed-σ latent families reduce
8//! to evaluating kernel bundles at shifted arguments.
9//!
10//! Row likelihoods for binary and survival models are small signed sums of
11//! kernel terms; [`LogKernelSumJet`] evaluates their log-derivatives from
12//! log-space kernel bundles and treats non-positive signed sums as invalid rows.
13
14use 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// ─── Typed errors ────────────────────────────────────────────────────────────
23
24/// Errors produced by the lognormal-kernel frailty/marginal-slope validators.
25///
26/// Public boundaries that historically returned `Result<_, String>` continue to
27/// do so via `.map_err(|e| e.to_string())`; the `Display` impl reproduces the
28/// original error strings byte-for-byte.
29#[derive(Debug, Clone)]
30pub enum LognormalKernelError {
31    /// The chosen frailty modifier is not finite-state exact with the
32    /// requested marginal-slope family.
33    InvalidSpec { reason: String },
34}
35
36impl_reason_error_boilerplate! {
37    LognormalKernelError {
38        InvalidSpec,
39    }
40}
41
42// ─── Frailty specification ───────────────────────────────────────────────────
43
44/// How the hazard multiplier frailty loads onto the hazard components.
45#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
46#[serde(rename_all = "kebab-case")]
47pub enum HazardLoading {
48    /// Frailty multiplies the entire hazard: h(t|U) = exp(U) · h_0(t).
49    Full,
50    /// Frailty multiplies only the disease-like component; an exogenous
51    /// background ("Makeham") component is unloaded:
52    ///   h(t|U) = exp(U) · h_loaded(t) + h_unloaded(t).
53    /// This is the faithful model for Gompertz-Makeham.
54    LoadedVsUnloaded,
55}
56
57/// Frailty modifier specification at the family level.
58///
59/// Two structurally different exact modifiers exist:
60///
61/// 1. **GaussianShift**: additive Gaussian on the final transformation index.
62///    Exact for probit families — the existing sextic microcell kernel survives
63///    unchanged (just scale denested cell coefficients by 1/√(1+σ²)).
64///
65/// 2. **HazardMultiplier**: lognormal multiplier on the loaded cumulative hazard.
66///    Exact for PH/cloglog families — row likelihoods are finite sums of
67///    K_{k,m}(μ, σ) kernel terms.
68///
69/// These are mathematically distinct families.  Do not mix them.
70#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
71#[serde(tag = "frailty_kind", rename_all = "kebab-case")]
72pub enum FrailtySpec {
73    /// No frailty modifier.
74    #[default]
75    None,
76    /// Gaussian shift on the final scalar index: U ~ N(0, σ²) added to η.
77    /// Exact for probit: E[Φ(η + U)] = Φ(η / √(1+σ²)).
78    /// The existing sextic microcell kernel is preserved.
79    GaussianShift {
80        /// Fixed σ, or None if learnable.
81        sigma_fixed: Option<f64>,
82    },
83    /// Lognormal hazard multiplier: conditional hazard h(t|U) involves exp(U).
84    /// Exact for PH/cloglog/survival via K_{k,m} kernel.
85    HazardMultiplier {
86        /// Fixed σ, or None if learnable.
87        sigma_fixed: Option<f64>,
88        /// How the multiplier loads onto hazard components.
89        loading: HazardLoading,
90    },
91}
92
93impl FrailtySpec {
94    /// Whether this spec requests an actual frailty modifier.
95    ///
96    /// [`FrailtySpec::None`] is the sole "no frailty" value. Family/mode guards
97    /// that only support the no-frailty case must reject on this predicate, or they
98    /// misclassify every ordinary CLI fit as a frailty request.
99    #[inline]
100    pub fn is_active(&self) -> bool {
101        !matches!(self, Self::None)
102    }
103
104    /// Validate the frailty scale domain independently of a model family.
105    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    /// Resolve the exact frailty subset supported by Gaussian-shift
122    /// marginal-slope models.
123    ///
124    /// These models admit either no frailty or a Gaussian shift with a fixed
125    /// standard deviation. A learnable Gaussian-shift scale and the
126    /// structurally different hazard-multiplier frailty require other fitting
127    /// machinery and are rejected here at the type that owns those states.
128    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    /// Validate that this frailty spec is compatible with score_warp/linkwiggle
154    /// cubic marginal-slope families.
155    ///
156    /// - `GaussianShift` is exact: the sextic microcell kernel is preserved
157    ///   (probit scaling by 1/τ, τ = √(1+σ²)).
158    /// - `HazardMultiplier` is exact only for PH/cloglog rowwise families.
159    ///   It is NOT finite-state exact with score_warp/linkwiggle cubic
160    ///   marginal-slope, because the multiplicative frailty breaks the
161    ///   polynomial kernel closure that the cubic cell derivatives require.
162    ///
163    /// Returns an error if the combination is not exactly integrable.
164    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    /// Typed variant of [`Self::validate_for_marginal_slope`] used internally;
170    /// the `String`-returning entry point above is preserved as a one-line
171    /// shim for external callers.
172    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// ─── Probit frailty scaling ──────────────────────────────────────────────────
187
188#[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/// Probit frailty scaling factor **with** t-derivatives (t = log σ).
203///
204/// Provides exact closed-form derivatives of s = 1/√(1+σ²) with respect to
205/// t = log(σ) for learnable Gaussian-shift frailty in the marginal-slope
206/// families.  For Gaussian frailty on the final probit index
207/// E[Φ(η + U)] = Φ(η · s) with s = 1/√(1+σ²); writing α = σ²/(1+σ²) the
208/// derivatives are ∂_t s = −α·s and ∂_{tt} s = α(3α−2)·s.
209#[derive(Clone, Copy, Debug)]
210pub struct ProbitFrailtyScaleJet {
211    /// s = 1/√(1+σ²)
212    pub s: f64,
213    /// α = σ²/(1+σ²)  — shared auxiliary for all derivative levels.
214    pub alpha: f64,
215    /// ∂_t s = -α·s
216    pub ds: f64,
217    /// ∂_{tt} s = α(3α−2)·s
218    pub d2s: f64,
219}
220
221impl ProbitFrailtyScaleJet {
222    /// Build the jet from σ (not from t = log σ).
223    ///
224    /// At σ = 0 the jet degenerates to (s=1, α=0, ds=0, d2s=0), which is
225    /// correct: zero frailty means s ≡ 1 independent of t.
226    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    /// Build the jet from t = log(σ) directly.
237    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// ─── Log-space kernel infrastructure ──────────────────────────────────────────
251//
252// The runtime kernel path stays in log-space until the final ratios are formed,
253// avoiding the overflow/underflow and cancellation problems that come from
254// exponentiating individual terms too early.
255
256/// Returns `log K_{k,m}(μ,σ)` directly, without exponentiation.
257///
258/// The value is always finite (or `NEG_INFINITY` when the kernel is zero), so
259/// it cannot overflow or underflow.
260#[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    // Survival carried in log space: prefix + ln S(shifted_mu, σ). This keeps the
308    // kernel's true magnitude when S underflows in value space at large σ — the
309    // old `laplace <= 0.0 → −∞` collapse discarded a large-but-finite log-value
310    // (#798) and the value-space asymptotic was biased low at σ ≥ 8 (#799).
311    let (log_laplace, mode) = lognormal_laplace_unit_log_term_shared(quadctx, shifted_mu, sigma);
312    Ok((prefix + log_laplace, mode))
313}
314
315/// Kernel bundle storing `log K_{k,m}` values instead of `K_{k,m}`.
316#[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
334/// Builds a log-space kernel bundle for `k = 0, 1, …, max_k` at fixed
335/// `(m, μ, σ)`.
336pub 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
395/// Computes the value-space derivative ratios `∂ⁿ_μ K_{k,m} / K_{k,m}`
396/// from a log-space bundle.
397///
398/// Returns `[1, K'/K, K''/K, K'''/K, K''''/K]` where only the first
399/// `order + 1` entries are valid.
400///
401/// The recurrences are applied in ratio form, with each `K_{k+r}/K_k`
402/// computed as `exp(log K_{k+r} − log K_k)`, which remains finite even when
403/// the individual kernel values would overflow or underflow.
404pub 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    // Precompute ratios K_{k+r}/K_k for r = 1..=order, each from a single
414    // log-difference.  This avoids redundant exp() calls when the same ratio
415    // appears in multiple derivative orders.
416    let mut rk = [0.0f64; 5]; // rk[0] unused; rk[r] = K_{k+r}/K_k
417    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
458// `LatentCLogLogJet5` + `latent_cloglog_jet5` / `latent_cloglog_inverse_link_jet`
459// moved DOWN to `crate::quadrature` (#1135), co-located with their analytic
460// backend, so the `solver` link layer names them without importing up into
461// `families::survival`. Re-exported here so the in-family callers (e.g.
462// `family_runtime`) keep resolving.
463pub use crate::quadrature::{
464    LatentCLogLogJet5, latent_cloglog_inverse_link_jet, latent_cloglog_jet5,
465};
466
467// ─── LogKernelSumJet: log-sum derivatives from log-space bundles ─────────────
468
469/// A single signed term in a kernel sum: coefficient × K_{k,m}.
470#[derive(Clone, Copy, Debug)]
471pub struct KernelSumTerm {
472    /// Multiplicative coefficient (can be negative for difference terms).
473    pub coeff: f64,
474    /// Kernel order parameter k.
475    pub k: usize,
476    /// Kernel mass parameter m (≥ 0).
477    pub m: f64,
478}
479
480/// Derivatives of `log(Σ_j a_j · K_{k_j, m_j}(μ, σ))` with respect to μ.
481///
482/// This is the workhorse for row-level log-likelihood derivatives in all
483/// latent families.  The numerator and denominator of a row likelihood are
484/// each a small signed sum of kernel terms.
485///
486/// The value path is assembled from log-space kernel bundles and ratio jets,
487/// so individual kernel terms are never exponentiated before the final signed
488/// sum. That avoids the old overflow/underflow problems from value-space
489/// kernels. When the signed sum is zero or negative, this returns an invalid
490/// row (`value = -∞`) instead of trying to continue with a floored surrogate.
491/// Signed two-term differences (e.g. interval censoring `K_{0,M_L} − K_{0,M_R}`)
492/// are still combined through the shared sign-aware log-sum path.
493#[derive(Clone, Copy, Debug)]
494pub struct LogKernelSumJet {
495    /// log(Σ a_j K_j)
496    pub value: f64,
497    /// d/dμ log(Σ a_j K_j)
498    pub d1: f64,
499    /// d²/dμ² log(Σ a_j K_j)
500    pub d2: f64,
501    /// d³/dμ³ log(Σ a_j K_j)
502    pub d3: f64,
503    /// d⁴/dμ⁴ log(Σ a_j K_j)
504    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            // d4 is used by the exact log-sigma curvature, so this must carry
549            // ratios through order 4 rather than truncating at order 3.
550            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    /// Evaluate for a single positive kernel term (fast path).
601    ///
602    /// Computes `log(K_{k,m})` and its μ-derivatives from exact recurrences,
603    /// entirely in log-space.
604    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    /// Evaluate `log(Σ a_j K_j)` and its μ-derivatives for a small signed sum.
621    ///
622    /// All terms share the same `(μ, σ)`.  Both the value and derivative
623    /// ratios are computed entirely in log-space.  The runtime latent-survival
624    /// rows in this repo are almost always one-term or two-term sums, so those
625    /// cases stay on dedicated stack paths; the heap-backed logic below is only
626    /// for genuinely longer symbolic sums:
627    ///
628    /// 1. Per-term log-magnitudes `log|a_j| + log K_{k_j,m_j}` and signs.
629    /// 2. Sign-aware log-sum-exp to get `log|S|` and `sign(S)`.
630    /// 3. Importance weights `w_j = a_j K_j / S` formed in log-space.
631    /// 4. Weighted ratio sums `R_n = Σ w_j · (∂ⁿK_j / K_j)` for the
632    ///    final log-derivatives.
633    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            // Empty sums are a caller-contract violation, not a degenerate row.
641            // Return an input error so callers can report the malformed kernel sum.
642            crate::bail_invalid_estim!("KernelSumJet requires at least one term");
643        }
644
645        // Fast path for single term.
646        if terms.len() == 1 {
647            let t = &terms[0];
648            if t.coeff <= 0.0 {
649                // Negative or zero coefficient: the sum is non-positive, so
650                // log(sum) is undefined.  Return −∞ (impossible observation),
651                // matching the general path's sign_s ≤ 0 branch.
652                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        // Build log-bundles for each unique mass.
673        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        // Per-term: log magnitude, sign, and ratio jet.
695        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        // Sign-aware log-sum-exp: compute log|S| and sign(S).
706        let (log_s, sign_s) = signed_log_sum_exp(&log_mags, &signs);
707
708        if !log_s.is_finite() || sign_s <= 0.0 {
709            // Sum is zero or negative — degenerate row.
710            return Ok(Self::non_positive(overall_mode));
711        }
712
713        // Importance weights w_j = sign(S) · sign(a_j) · exp(log|a_j K_j| − log|S|).
714        // When S > 0 and all terms have well-defined kernels, Σ w_j = 1.
715        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// ─── Latent survival sufficient statistics ───────────────────────────────────
740
741/// Event type for compiled survival sufficient statistics.
742#[derive(Clone, Copy, Debug, PartialEq, Eq)]
743pub enum LatentSurvivalEventType {
744    /// Right-censored: observed alive in the observation window.
745    RightCensored,
746    /// Exact event: event observed at a known time.
747    ExactEvent,
748    /// Interval-censored: event known to occur in (t_left, t_right].
749    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/// Row-level sufficient statistics for one latent survival observation.
763///
764/// This is the canonical row representation used by both fitted-family
765/// evaluation and saved-model prediction.
766///
767/// For the full-loading model (frailty multiplies entire hazard):
768///   mass_loaded = total cumulative hazard mass
769///   mass_unloaded = 0
770///
771/// For the loaded-vs-unloaded model (Gompertz-Makeham):
772///   mass_loaded = integrated disease hazard component
773///   mass_unloaded = integrated background hazard component (not frailty-modified)
774///
775/// The unloaded mass contributes a simple exp(-M_U) prefactor.
776#[derive(Clone, Copy, Debug)]
777pub struct LatentSurvivalRow {
778    pub event_type: LatentSurvivalEventType,
779    /// Cumulative nuisance mass at entry: B(a_in).
780    /// Zero if there is no left truncation.
781    pub mass_entry: f64,
782    /// Cumulative nuisance mass at exit/event: B(a_out) or B(a_event).
783    pub mass_exit: f64,
784    /// For interval censoring: mass at left boundary B(a_L).
785    pub mass_left: f64,
786    /// For interval censoring: mass at right boundary B(a_R).
787    pub mass_right: f64,
788    /// For interval censoring: unloaded mass at left boundary.
789    pub mass_unloaded_left: f64,
790    /// For interval censoring: unloaded mass at right boundary.
791    pub mass_unloaded_right: f64,
792    /// Unloaded (background) cumulative mass at entry (0 for full loading).
793    pub mass_unloaded_entry: f64,
794    /// Unloaded (background) cumulative mass at exit.
795    pub mass_unloaded_exit: f64,
796    /// Loaded instantaneous hazard at event time (for exact events).
797    pub hazard_loaded: f64,
798    /// Unloaded instantaneous hazard at event time (for exact events).
799    pub hazard_unloaded: f64,
800}
801
802impl LatentSurvivalRow {
803    /// Delayed-entry right-censored row with explicit loaded/unloaded masses.
804    ///
805    /// `mass_entry` and `mass_exit` are cumulative loaded masses `B_L(a_in)`
806    /// and `B_L(a_out)` for this row object. They are not an increment over
807    /// `(a_in, a_out]`.
808    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    /// Delayed-entry exact-event row with explicit loaded/unloaded hazard parts.
830    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    /// Delayed-entry interval-censored row with explicit loaded/unloaded masses.
854    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/// Row-level log-likelihood and μ-derivatives for the latent survival model.
1049///
1050/// The conditional model is:
1051///   `Λ(a | U) = B(a) · exp(U)`,  `U ~ N(μ, σ²)`
1052///
1053/// All likelihoods reduce to algebra on `K_{k,m}(μ, σ)`.
1054#[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    // For S = Σ a_j K_j, D = σ ∂_σ, and D S = σ² S_μμ:
1080    // D² log S = 2σ² (S''/S) + σ⁴ (S''''/S - (S''/S)²).
1081    // Express the final parenthesized term directly in log-derivatives to
1082    // avoid the larger cancellation in `r4 - r2²`.
1083    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    /// Right-censoring with loaded/unloaded mass decomposition.
1105    ///
1106    /// Full formula:
1107    ///   `ℓ = -M_U_exit + log K_{0,M_L_exit} + M_U_entry - log K_{0,M_L_entry}`
1108    ///
1109    /// When `mass_unloaded_exit == 0` and `mass_unloaded_entry == 0`, this
1110    /// falls back to the original formula using `mass_exit` / `mass_entry`.
1111    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        // Loaded mass for the kernel terms: when unloaded mass is present,
1121        // mass_exit contains only the loaded component; otherwise it is the
1122        // total mass.
1123        let mass_exit_loaded = row.mass_exit;
1124        let mass_entry_loaded = row.mass_entry;
1125
1126        // Unloaded mass contributes a simple additive constant to log-lik
1127        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    /// Exact event with loaded/unloaded hazard decomposition.
1159    ///
1160    /// `ℓ = log(h_U · K_{0,M_L} + h_L · K_{1,M_L}) - M_U_event + M_U_entry - log K_{0,M_L_entry}`
1161    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    /// Interval event: `ℓ = log(K_{0,M_L} − K_{0,M_R}) − log K_{0,M_in}`.
1200    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        // Second μ-derivative of ℓ = log[S(L) − S(R)] for the interval kernel,
1453        // FD-checked. `neg_hessian` stores −d²ℓ/dμ², so compare against the
1454        // negated central second difference.
1455        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        // σ-recovery for interval data is driven by `score_log_sigma`, the
1478        // derivative of ℓ = log[S(L) − S(R)] w.r.t. log σ. FD-check it directly
1479        // against the row log-likelihood (this is the channel the interval fit's
1480        // latent_sd estimate moves along, the test's primary metric).
1481        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        // d/d(log σ) = σ · d/dσ, so FD over log σ directly.
1492        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}