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, 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    /// Exact log-sigma coordinate and declared finite chart domain for a
100    /// learned scale. Fixed scales have no optimizer coordinate.
101    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    /// No frailty modifier.
113    #[default]
114    None,
115    /// Gaussian shift on the final scalar index: U ~ N(0, σ²) added to η.
116    /// Exact for probit: E[Φ(η + U)] = Φ(η / √(1+σ²)).
117    /// The existing sextic microcell kernel is preserved.
118    GaussianShift {
119        scale: FrailtyScale,
120    },
121    /// Lognormal hazard multiplier: conditional hazard h(t|U) involves exp(U).
122    /// Exact for PH/cloglog/survival via K_{k,m} kernel.
123    HazardMultiplier {
124        scale: FrailtyScale,
125        /// How the multiplier loads onto hazard components.
126        loading: HazardLoading,
127    },
128}
129
130impl FrailtySpec {
131    /// Whether this spec requests an actual frailty modifier.
132    ///
133    /// [`FrailtySpec::None`] is the sole "no frailty" value. Family/mode guards
134    /// that only support the no-frailty case must reject on this predicate, or they
135    /// misclassify every ordinary CLI fit as a frailty request.
136    #[inline]
137    pub fn is_active(&self) -> bool {
138        !matches!(self, Self::None)
139    }
140
141    /// Validate the frailty scale domain independently of a model family.
142    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    /// Validate that this frailty spec is compatible with score_warp/linkwiggle
152    /// cubic marginal-slope families.
153    ///
154    /// - `GaussianShift` is exact: the sextic microcell kernel is preserved
155    ///   (probit scaling by 1/τ, τ = √(1+σ²)).
156    /// - `HazardMultiplier` is exact only for PH/cloglog rowwise families.
157    ///   It is NOT finite-state exact with score_warp/linkwiggle cubic
158    ///   marginal-slope, because the multiplicative frailty breaks the
159    ///   polynomial kernel closure that the cubic cell derivatives require.
160    ///
161    /// Returns an error if the combination is not exactly integrable.
162    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    /// Typed variant of [`Self::validate_for_marginal_slope`] used internally;
168    /// the `String`-returning entry point above is preserved as a one-line
169    /// shim for external callers.
170    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// ─── Probit frailty scaling ──────────────────────────────────────────────────
186
187#[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/// Probit frailty scaling factor **with** t-derivatives (t = log σ).
202///
203/// Provides exact closed-form derivatives of s = 1/√(1+σ²) with respect to
204/// t = log(σ) for learnable Gaussian-shift frailty in the marginal-slope
205/// families.  For Gaussian frailty on the final probit index
206/// E[Φ(η + U)] = Φ(η · s) with s = 1/√(1+σ²); writing α = σ²/(1+σ²) the
207/// derivatives are ∂_t s = −α·s and ∂_{tt} s = α(3α−2)·s.
208#[derive(Clone, Copy, Debug)]
209pub struct ProbitFrailtyScaleJet {
210    /// s = 1/√(1+σ²)
211    pub s: f64,
212    /// α = σ²/(1+σ²)  — shared auxiliary for all derivative levels.
213    pub alpha: f64,
214    /// ∂_t s = -α·s
215    pub ds: f64,
216    /// ∂_{tt} s = α(3α−2)·s
217    pub d2s: f64,
218}
219
220impl ProbitFrailtyScaleJet {
221    /// Build the jet from σ (not from t = log σ).
222    ///
223    /// At σ = 0 the jet degenerates to (s=1, α=0, ds=0, d2s=0), which is
224    /// correct: zero frailty means s ≡ 1 independent of t.
225    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    /// Build the jet from t = log(σ) directly.
236    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// ─── Log-space kernel infrastructure ──────────────────────────────────────────
250//
251// The runtime kernel path stays in log-space until the final ratios are formed,
252// avoiding the overflow/underflow and cancellation problems that come from
253// exponentiating individual terms too early.
254
255/// Returns `log K_{k,m}(μ,σ)` directly, without exponentiation.
256///
257/// The value is always finite (or `NEG_INFINITY` when the kernel is zero), so
258/// it cannot overflow or underflow.
259#[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    // Survival carried in log space: prefix + ln S(shifted_mu, σ). This keeps the
307    // kernel's true magnitude when S underflows in value space at large σ — the
308    // old `laplace <= 0.0 → −∞` collapse discarded a large-but-finite log-value
309    // (#798) and the value-space asymptotic was biased low at σ ≥ 8 (#799).
310    let (log_laplace, mode) = lognormal_laplace_unit_log_term_shared(quadctx, shifted_mu, sigma);
311    Ok((prefix + log_laplace, mode))
312}
313
314/// Kernel bundle storing `log K_{k,m}` values instead of `K_{k,m}`.
315#[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
333/// Builds a log-space kernel bundle for `k = 0, 1, …, max_k` at fixed
334/// `(m, μ, σ)`.
335pub 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
394/// Computes the value-space derivative ratios `∂ⁿ_μ K_{k,m} / K_{k,m}`
395/// from a log-space bundle.
396///
397/// Returns `[1, K'/K, K''/K, K'''/K, K''''/K]` where only the first
398/// `order + 1` entries are valid.
399///
400/// The recurrences are applied in ratio form, with each `K_{k+r}/K_k`
401/// computed as `exp(log K_{k+r} − log K_k)`, which remains finite even when
402/// the individual kernel values would overflow or underflow.
403pub 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    // Precompute ratios K_{k+r}/K_k for r = 1..=order, each from a single
413    // log-difference.  This avoids redundant exp() calls when the same ratio
414    // appears in multiple derivative orders.
415    let mut rk = [0.0f64; 5]; // rk[0] unused; rk[r] = K_{k+r}/K_k
416    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
457// `LatentCLogLogJet5` + `latent_cloglog_jet5` / `latent_cloglog_inverse_link_jet`
458// moved DOWN to `crate::quadrature` (#1135), co-located with their analytic
459// backend, so the `solver` link layer names them without importing up into
460// `families::survival`. Re-exported here so the in-family callers (e.g.
461// `family_runtime`) keep resolving.
462pub use crate::quadrature::{
463    LatentCLogLogJet5, latent_cloglog_inverse_link_jet, latent_cloglog_jet5,
464};
465
466// ─── LogKernelSumJet: log-sum derivatives from log-space bundles ─────────────
467
468/// A single signed term in a kernel sum: coefficient × K_{k,m}.
469#[derive(Clone, Copy, Debug)]
470pub struct KernelSumTerm {
471    /// Multiplicative coefficient (can be negative for difference terms).
472    pub coeff: f64,
473    /// Kernel order parameter k.
474    pub k: usize,
475    /// Kernel mass parameter m (≥ 0).
476    pub m: f64,
477}
478
479/// Derivatives of `log(Σ_j a_j · K_{k_j, m_j}(μ, σ))` with respect to μ.
480///
481/// This is the workhorse for row-level log-likelihood derivatives in all
482/// latent families.  The numerator and denominator of a row likelihood are
483/// each a small signed sum of kernel terms.
484///
485/// The value path is assembled from log-space kernel bundles and ratio jets,
486/// so individual kernel terms are never exponentiated before the final signed
487/// sum. That avoids the old overflow/underflow problems from value-space
488/// kernels. When the signed sum is zero or negative, this returns an invalid
489/// row (`value = -∞`) instead of trying to continue with a floored surrogate.
490/// Signed two-term differences (e.g. interval censoring `K_{0,M_L} − K_{0,M_R}`)
491/// are still combined through the shared sign-aware log-sum path.
492#[derive(Clone, Copy, Debug)]
493pub struct LogKernelSumJet {
494    /// log(Σ a_j K_j)
495    pub value: f64,
496    /// d/dμ log(Σ a_j K_j)
497    pub d1: f64,
498    /// d²/dμ² log(Σ a_j K_j)
499    pub d2: f64,
500    /// d³/dμ³ log(Σ a_j K_j)
501    pub d3: f64,
502    /// d⁴/dμ⁴ log(Σ a_j K_j)
503    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            // d4 is used by the exact log-sigma curvature, so this must carry
548            // ratios through order 4 rather than truncating at order 3.
549            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    /// Evaluate for a single positive kernel term (fast path).
600    ///
601    /// Computes `log(K_{k,m})` and its μ-derivatives from exact recurrences,
602    /// entirely in log-space.
603    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    /// Evaluate `log(Σ a_j K_j)` and its μ-derivatives for a small signed sum.
620    ///
621    /// All terms share the same `(μ, σ)`.  Both the value and derivative
622    /// ratios are computed entirely in log-space.  The runtime latent-survival
623    /// rows in this repo are almost always one-term or two-term sums, so those
624    /// cases stay on dedicated stack paths; the heap-backed logic below is only
625    /// for genuinely longer symbolic sums:
626    ///
627    /// 1. Per-term log-magnitudes `log|a_j| + log K_{k_j,m_j}` and signs.
628    /// 2. Sign-aware log-sum-exp to get `log|S|` and `sign(S)`.
629    /// 3. Importance weights `w_j = a_j K_j / S` formed in log-space.
630    /// 4. Weighted ratio sums `R_n = Σ w_j · (∂ⁿK_j / K_j)` for the
631    ///    final log-derivatives.
632    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            // Empty sums are a caller-contract violation, not a degenerate row.
640            // Return an input error so callers can report the malformed kernel sum.
641            crate::bail_invalid_estim!("KernelSumJet requires at least one term");
642        }
643
644        // Fast path for single term.
645        if terms.len() == 1 {
646            let t = &terms[0];
647            if t.coeff <= 0.0 {
648                // Negative or zero coefficient: the sum is non-positive, so
649                // log(sum) is undefined.  Return −∞ (impossible observation),
650                // matching the general path's sign_s ≤ 0 branch.
651                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        // Build log-bundles for each unique mass.
672        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        // Per-term: log magnitude, sign, and ratio jet.
687        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        // Sign-aware log-sum-exp: compute log|S| and sign(S).
698        let (log_s, sign_s) = signed_log_sum_exp(&log_mags, &signs);
699
700        if !log_s.is_finite() || sign_s <= 0.0 {
701            // Sum is zero or negative — degenerate row.
702            return Ok(Self::non_positive(overall_mode));
703        }
704
705        // Importance weights w_j = sign(S) · sign(a_j) · exp(log|a_j K_j| − log|S|).
706        // When S > 0 and all terms have well-defined kernels, Σ w_j = 1.
707        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// ─── Latent survival sufficient statistics ───────────────────────────────────
732
733/// Event type for compiled survival sufficient statistics.
734#[derive(Clone, Copy, Debug, PartialEq, Eq)]
735pub enum LatentSurvivalEventType {
736    /// Right-censored: observed alive in the observation window.
737    RightCensored,
738    /// Exact event: event observed at a known time.
739    ExactEvent,
740    /// Interval-censored: event known to occur in (t_left, t_right].
741    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/// Row-level sufficient statistics for one latent survival observation.
755///
756/// This is the canonical row representation used by both fitted-family
757/// evaluation and saved-model prediction.
758///
759/// For the full-loading model (frailty multiplies entire hazard):
760///   mass_loaded = total cumulative hazard mass
761///   mass_unloaded = 0
762///
763/// For the loaded-vs-unloaded model (Gompertz-Makeham):
764///   mass_loaded = integrated disease hazard component
765///   mass_unloaded = integrated background hazard component (not frailty-modified)
766///
767/// The unloaded mass contributes a simple exp(-M_U) prefactor.
768#[derive(Clone, Copy, Debug)]
769pub struct LatentSurvivalRow {
770    pub event_type: LatentSurvivalEventType,
771    /// Cumulative nuisance mass at entry: B(a_in).
772    /// Zero if there is no left truncation.
773    pub mass_entry: f64,
774    /// Cumulative nuisance mass at exit/event: B(a_out) or B(a_event).
775    pub mass_exit: f64,
776    /// For interval censoring: mass at left boundary B(a_L).
777    pub mass_left: f64,
778    /// For interval censoring: mass at right boundary B(a_R).
779    pub mass_right: f64,
780    /// For interval censoring: unloaded mass at left boundary.
781    pub mass_unloaded_left: f64,
782    /// For interval censoring: unloaded mass at right boundary.
783    pub mass_unloaded_right: f64,
784    /// Unloaded (background) cumulative mass at entry (0 for full loading).
785    pub mass_unloaded_entry: f64,
786    /// Unloaded (background) cumulative mass at exit.
787    pub mass_unloaded_exit: f64,
788    /// Loaded instantaneous hazard at event time (for exact events).
789    pub hazard_loaded: f64,
790    /// Unloaded instantaneous hazard at event time (for exact events).
791    pub hazard_unloaded: f64,
792}
793
794impl LatentSurvivalRow {
795    /// Delayed-entry right-censored row with explicit loaded/unloaded masses.
796    ///
797    /// `mass_entry` and `mass_exit` are cumulative loaded masses `B_L(a_in)`
798    /// and `B_L(a_out)` for this row object. They are not an increment over
799    /// `(a_in, a_out]`.
800    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    /// Delayed-entry exact-event row with explicit loaded/unloaded hazard parts.
822    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    /// Delayed-entry interval-censored row with explicit loaded/unloaded masses.
846    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/// Row-level log-likelihood and μ-derivatives for the latent survival model.
1041///
1042/// The conditional model is:
1043///   `Λ(a | U) = B(a) · exp(U)`,  `U ~ N(μ, σ²)`
1044///
1045/// All likelihoods reduce to algebra on `K_{k,m}(μ, σ)`.
1046#[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    // For S = Σ a_j K_j, D = σ ∂_σ, and D S = σ² S_μμ:
1072    // D² log S = 2σ² (S''/S) + σ⁴ (S''''/S - (S''/S)²).
1073    // Express the final parenthesized term directly in log-derivatives to
1074    // avoid the larger cancellation in `r4 - r2²`.
1075    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    /// Right-censoring with loaded/unloaded mass decomposition.
1097    ///
1098    /// Full formula:
1099    ///   `ℓ = -M_U_exit + log K_{0,M_L_exit} + M_U_entry - log K_{0,M_L_entry}`
1100    ///
1101    /// When `mass_unloaded_exit == 0` and `mass_unloaded_entry == 0`, this
1102    /// falls back to the original formula using `mass_exit` / `mass_entry`.
1103    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        // Loaded mass for the kernel terms: when unloaded mass is present,
1112        // mass_exit contains only the loaded component; otherwise it is the
1113        // total mass.
1114        let mass_exit_loaded = row.mass_exit;
1115        let mass_entry_loaded = row.mass_entry;
1116
1117        // Unloaded mass contributes a simple additive constant to log-lik
1118        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    /// Exact event with loaded/unloaded hazard decomposition.
1150    ///
1151    /// `ℓ = 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}`
1152    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    /// Interval event: `ℓ = log(K_{0,M_L} − K_{0,M_R}) − log K_{0,M_in}`.
1190    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    /// #2277 hardening: a NARROW interval-censored window (S(L) ≈ S(R)) must
1445    /// stay numerically stable. The interval contribution `log[S(L) − S(R)]` is
1446    /// evaluated in the log domain (sign-aware log-sum-exp + `log1mexp`), never
1447    /// as a probability-space `S(L) − S(R)` subtraction, so for a small gap
1448    /// `Δ = M_R − M_L` the interval mass is `≈ |S'(M_L)|·Δ` and the
1449    /// log-likelihood behaves as `const + log Δ` — finite and accurate down to
1450    /// gaps where subtracting two nearly-equal survival probabilities would
1451    /// catastrophically cancel. Pin the log-linear-in-Δ law: the shared,
1452    /// cancellation-prone `const` drops out of the difference, leaving exactly
1453    /// `log(Δ₁/Δ₂)`.
1454    #[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        // const + log Δ ⇒ ll(g1) − ll(g2) = log(g1/g2), independent of the
1472        // shared const a probability-space subtraction would destroy here.
1473        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        // Second μ-derivative of ℓ = log[S(L) − S(R)] for the interval kernel,
1485        // FD-checked. `neg_hessian` stores −d²ℓ/dμ², so compare against the
1486        // negated central second difference.
1487        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        // σ-recovery for interval data is driven by `score_log_sigma`, the
1510        // derivative of ℓ = log[S(L) − S(R)] w.r.t. log σ. FD-check it directly
1511        // against the row log-likelihood (this is the channel the interval fit's
1512        // latent_sd estimate moves along, the test's primary metric).
1513        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        // d/d(log σ) = σ · d/dσ, so FD over log σ directly.
1524        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}