use super::*;
const PROBIT_FAR_TAIL_ENTRY: f64 = 1.0e3;
const PROBIT_LEFT_TAIL_CUTOFF: f64 = -38.0;
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct PairedNeglogStacks {
pub(crate) event: [f64; 4],
pub(crate) censored: [f64; 4],
}
pub fn paired_neglog_stacks(
link: &InverseLink,
u0: f64,
u1: f64,
delta_u: f64,
) -> Option<PairedNeglogStacks> {
match link {
InverseLink::Standard(StandardLink::Probit) => probit_paired(link, u0, u1, delta_u),
InverseLink::Standard(StandardLink::Logit) => logit_paired(link, u0, u1),
InverseLink::Standard(StandardLink::CLogLog) => Some(cloglog_paired(u0, delta_u)),
InverseLink::Standard(StandardLink::Identity) => identity_paired(link, u0, u1),
_ => None,
}
}
fn surv_derivs(link: &InverseLink, u: f64) -> Option<[f64; 4]> {
let (_, r, dr, ddr, dddr) =
SurvivalLocationScaleFamily::exact_survival_neglog_derivatives_fourth_rescaled(link, u, 0.0)
.ok()?;
Some([r, dr, ddr, dddr])
}
fn pdf_derivs(link: &InverseLink, u: f64) -> Option<[f64; 4]> {
let (_, d1, d2, d3, d4) =
SurvivalLocationScaleFamily::exact_log_pdf_derivatives_rescaled(link, u, 0.0).ok()?;
Some([d1, d2, d3, d4])
}
fn sigmoid_of_neg(u: f64) -> f64 {
if u >= 0.0 {
let z = (-u).exp();
z / (1.0 + z)
} else {
1.0 / (1.0 + u.exp())
}
}
fn mills_residual_derivs(u: f64) -> (f64, f64, f64, f64) {
let inv = 1.0 / u;
let p = inv * inv; let rho = inv * (1.0 + p * (-2.0 + p * (10.0 + p * (-74.0))));
let d_rho = p * (-1.0 + p * (6.0 + p * (-50.0 + p * 518.0)));
let dd_rho = inv * p * (2.0 + p * (-24.0 + p * (300.0 + p * (-4144.0))));
let ddd_rho = p * p * (-6.0 + p * (120.0 + p * (-2100.0 + p * 37296.0)));
(rho, d_rho, dd_rho, ddd_rho)
}
fn probit_paired(
link: &InverseLink,
u0: f64,
u1: f64,
delta_u: f64,
) -> Option<PairedNeglogStacks> {
if u0 <= PROBIT_LEFT_TAIL_CUTOFF {
return None;
}
if u0 <= PROBIT_FAR_TAIL_ENTRY {
let a0 = surv_derivs(link, u0)?;
let a1 = surv_derivs(link, u1)?;
let b1 = pdf_derivs(link, u1)?;
return Some(PairedNeglogStacks {
event: [
a0[0] + b1[0],
a0[1] + b1[1],
a0[2] + b1[2],
a0[3] + b1[3],
],
censored: [
a0[0] - a1[0],
a0[1] - a1[1],
a0[2] - a1[2],
a0[3] - a1[3],
],
});
}
let (rho0, d_rho0, dd_rho0, ddd_rho0) = mills_residual_derivs(u0);
let (rho1, d_rho1, dd_rho1, ddd_rho1) = mills_residual_derivs(u1);
Some(PairedNeglogStacks {
event: [rho0 - delta_u, d_rho0, dd_rho0, ddd_rho0],
censored: [
(-delta_u) + (rho0 - rho1),
d_rho0 - d_rho1,
dd_rho0 - dd_rho1,
ddd_rho0 - ddd_rho1,
],
})
}
fn logit_paired(link: &InverseLink, u0: f64, u1: f64) -> Option<PairedNeglogStacks> {
let a0 = surv_derivs(link, u0)?;
let a1 = surv_derivs(link, u1)?;
let b1 = pdf_derivs(link, u1)?;
let s1_event = 2.0 * sigmoid_of_neg(u1) - sigmoid_of_neg(u0);
let s1_censored = sigmoid_of_neg(u1) - sigmoid_of_neg(u0);
Some(PairedNeglogStacks {
event: [s1_event, a0[1] + b1[1], a0[2] + b1[2], a0[3] + b1[3]],
censored: [s1_censored, a0[1] - a1[1], a0[2] - a1[2], a0[3] - a1[3]],
})
}
fn cloglog_paired(u0: f64, delta_u: f64) -> PairedNeglogStacks {
let growth = if delta_u == 0.0 {
0.0
} else {
u0.exp() * delta_u.exp_m1()
};
PairedNeglogStacks {
event: [1.0 - growth, -growth, -growth, -growth],
censored: [-growth, -growth, -growth, -growth],
}
}
fn identity_paired(link: &InverseLink, u0: f64, u1: f64) -> Option<PairedNeglogStacks> {
let a0 = surv_derivs(link, u0)?;
let a1 = surv_derivs(link, u1)?;
let b1 = pdf_derivs(link, u1)?;
Some(PairedNeglogStacks {
event: [
a0[0] + b1[0],
a0[1] + b1[1],
a0[2] + b1[2],
a0[3] + b1[3],
],
censored: [
a0[0] - a1[0],
a0[1] - a1[1],
a0[2] - a1[2],
a0[3] - a1[3],
],
})
}
pub(crate) fn paired_contraction_needs_regroup(link: &InverseLink, u0: f64) -> bool {
matches!(link, InverseLink::Standard(StandardLink::Probit))
&& u0.is_finite()
&& u0 > PROBIT_FAR_TAIL_ENTRY
}
pub(crate) fn weighted_paired_index_sums(
link: &InverseLink,
u0: f64,
u1: f64,
delta_u: f64,
d: f64,
w: f64,
) -> Option<[f64; 3]> {
let p = paired_neglog_stacks(link, u0, u1, delta_u)?;
Some([
w * event_mix(d, p.event[0], p.censored[0]),
w * event_mix(d, p.event[1], p.censored[1]),
w * event_mix(d, p.event[2], p.censored[2]),
])
}
#[cfg(test)]
mod paired_stack_tests {
use super::*;
fn probit() -> InverseLink {
InverseLink::Standard(StandardLink::Probit)
}
fn logit() -> InverseLink {
InverseLink::Standard(StandardLink::Logit)
}
fn cloglog() -> InverseLink {
InverseLink::Standard(StandardLink::CLogLog)
}
fn identity() -> InverseLink {
InverseLink::Standard(StandardLink::Identity)
}
fn assert_close(a: f64, b: f64, rel: f64, abs: f64, ctx: &str) {
let diff = (a - b).abs();
let scale = a.abs().max(b.abs());
assert!(
diff <= rel * scale + abs,
"{ctx}: a={a:e} b={b:e} diff={diff:e} rel_budget={:e}",
rel * scale + abs
);
}
#[test]
fn paired_matches_naive_stacks_in_moderate_regime() {
let probit_pts = [
(-3.0, -2.5),
(-1.0, 0.0),
(0.0, 0.0),
(0.5, 2.0),
(2.0, 2.0),
(-5.0, 3.0),
];
let logit_pts = [
(-4.0, -3.0),
(-1.0, 1.0),
(0.0, 0.0),
(2.0, 2.0),
(3.0, 5.0),
(6.0, 6.0),
];
let cloglog_pts = [
(-2.0, -1.0),
(-1.0, 0.0),
(0.0, 0.0),
(0.5, 1.5),
(1.0, 1.0),
(-3.0, 2.0),
];
let identity_pts = [
(-3.0, -2.0),
(-1.0, 0.0),
(0.0, 0.0),
(0.3, 0.6),
(0.8, 0.9),
(-10.0, 0.5),
];
let bitwise_all = |link: &InverseLink, pts: &[(f64, f64)]| {
for &(u0, u1) in pts {
let du = u1 - u0;
let p = paired_neglog_stacks(link, u0, u1, du).unwrap();
let a0 = surv_derivs(link, u0).unwrap();
let a1 = surv_derivs(link, u1).unwrap();
let b1 = pdf_derivs(link, u1).unwrap();
for k in 0..4 {
assert_eq!(
p.event[k].to_bits(),
(a0[k] + b1[k]).to_bits(),
"event[{k}] u0={u0} u1={u1}"
);
assert_eq!(
p.censored[k].to_bits(),
(a0[k] - a1[k]).to_bits(),
"censored[{k}] u0={u0} u1={u1}"
);
}
}
};
bitwise_all(&probit(), &probit_pts);
bitwise_all(&identity(), &identity_pts);
for &(u0, u1) in &logit_pts {
let du = u1 - u0;
let p = paired_neglog_stacks(&logit(), u0, u1, du).unwrap();
let a0 = surv_derivs(&logit(), u0).unwrap();
let a1 = surv_derivs(&logit(), u1).unwrap();
let b1 = pdf_derivs(&logit(), u1).unwrap();
for k in 1..4 {
assert_eq!(p.event[k].to_bits(), (a0[k] + b1[k]).to_bits());
assert_eq!(p.censored[k].to_bits(), (a0[k] - a1[k]).to_bits());
}
assert_close(p.event[0], a0[0] + b1[0], 1e-12, 1e-300, "logit s1 event");
assert_close(p.censored[0], a0[0] - a1[0], 1e-12, 1e-300, "logit s1 censored");
}
for &(u0, u1) in &cloglog_pts {
let du = u1 - u0;
let p = paired_neglog_stacks(&cloglog(), u0, u1, du).unwrap();
let a0 = surv_derivs(&cloglog(), u0).unwrap();
let a1 = surv_derivs(&cloglog(), u1).unwrap();
let b1 = pdf_derivs(&cloglog(), u1).unwrap();
for k in 0..4 {
assert_close(p.event[k], a0[k] + b1[k], 1e-12, 1e-300, "cloglog event");
assert_close(p.censored[k], a0[k] - a1[k], 1e-12, 1e-300, "cloglog censored");
}
}
}
#[test]
fn probit_far_tail_matches_mills_series() {
let entries = [1.0e4, 1.0e8, 1.0e50, 1.0e150, 3.66e150];
let gaps = [0.0, 1.0e-3, 0.48];
for &u0 in &entries {
for &du in &gaps {
let u1 = u0 + du;
let p = paired_neglog_stacks(&probit(), u0, u1, du).unwrap();
let ref2 = -1.0 / (u0 * u0);
assert_close(p.event[1], ref2, 10.0 / (u0 * u0), 0.0, "s2 event");
let ref3 = 2.0 / (u0 * u0 * u0);
if ref3 != 0.0 && ref3.is_finite() {
assert_close(p.event[2], ref3, 16.0 / (u0 * u0), 0.0, "s3 event");
}
let ref4 = -6.0 / (u0 * u0 * u0 * u0);
if ref4 != 0.0 && ref4.is_finite() {
assert_close(p.event[3], ref4, 28.0 / (u0 * u0), 0.0, "s4 event");
}
if du == 0.0 {
assert_close(p.event[0], 1.0 / u0, 3.0 / (u0 * u0), 0.0, "s1 event du=0");
} else {
assert!(
(p.event[0] + du).abs() <= 2.0 / u0,
"s1 event residual too large: u0={u0} du={du} val={:e}",
p.event[0] + du
);
}
if du != 0.0 {
let (_, d_rho0, _, _) = mills_residual_derivs(u0);
let ref_cens = -du * (1.0 + d_rho0);
assert_close(p.censored[0], ref_cens, 1e-10, 0.0, "s1 censored");
}
}
}
}
#[test]
fn probit_crossover_is_continuous() {
let thr = PROBIT_FAR_TAIL_ENTRY;
let off = 1.0e-11; let du = 0.1;
let lo = thr * (1.0 - off);
let hi = thr * (1.0 + off);
let p_lo = paired_neglog_stacks(&probit(), lo, lo + du, du).unwrap();
let p_hi = paired_neglog_stacks(&probit(), hi, hi + du, du).unwrap();
for k in 0..2 {
assert_close(p_lo.event[k], p_hi.event[k], 1e-6, 0.0, "crossover event lo-order");
assert_close(p_lo.censored[k], p_hi.censored[k], 1e-6, 0.0, "crossover censored lo-order");
}
for k in 2..4 {
assert_close(p_lo.event[k], p_hi.event[k], 1e-1, 1e-300, "crossover event hi-order");
assert_close(
p_lo.censored[k],
p_hi.censored[k],
1e-1,
1e-300,
"crossover censored hi-order",
);
}
}
#[test]
fn logistic_s1_regrouping_beats_naive_at_saturation() {
for &(u0, u1) in &[(-2.0, 1.0), (0.0, 0.0), (1.5, 3.0)] {
let p = paired_neglog_stacks(&logit(), u0, u1, u1 - u0).unwrap();
let want = 2.0 * sigmoid_of_neg(u1) - sigmoid_of_neg(u0);
assert_eq!(p.event[0].to_bits(), want.to_bits(), "logit s1 exact form");
}
let p = paired_neglog_stacks(&logit(), 40.0, 40.0, 0.0).unwrap();
let sig = sigmoid_of_neg(40.0);
assert!(sig > 0.0, "sigma(-40) must be positive: {sig:e}");
assert_eq!(p.event[0].to_bits(), sig.to_bits());
let a0 = surv_derivs(&logit(), 40.0).unwrap();
let b1 = pdf_derivs(&logit(), 40.0).unwrap();
let naive = a0[0] + b1[0];
assert_eq!(naive, 0.0, "naive s1 should saturate to 0, got {naive:e}");
assert!(p.event[0] > naive, "regrouping must recover a positive s1");
}
#[test]
fn cloglog_zero_gap_guard_stays_finite() {
assert!((800.0_f64).exp().is_infinite(), "fixture must overflow e^u0");
let p = paired_neglog_stacks(&cloglog(), 800.0, 800.0, 0.0).unwrap();
assert_eq!(p.event[0], 1.0);
for k in 1..4 {
assert_eq!(p.event[k], 0.0, "event[{k}]");
}
for k in 0..4 {
assert_eq!(p.censored[k], 0.0, "censored[{k}]");
}
assert!(
p.event.iter().chain(p.censored.iter()).all(|v| v.is_finite()),
"all entries finite"
);
let q = paired_neglog_stacks(&cloglog(), 1.0, 1.5, 0.5).unwrap();
let growth = 1.0_f64.exp() * 0.5_f64.exp_m1();
assert_close(q.event[0], 1.0 - growth, 1e-14, 0.0, "cloglog s1 growth");
assert_close(q.censored[3], -growth, 1e-14, 0.0, "cloglog cens growth");
}
#[test]
fn regroup_gate_fires_only_in_probit_far_tail() {
assert!(paired_contraction_needs_regroup(&probit(), 3.6e150));
assert!(paired_contraction_needs_regroup(&probit(), 1.0e4));
assert!(!paired_contraction_needs_regroup(&probit(), 5.0));
assert!(!paired_contraction_needs_regroup(&probit(), 1.0e3));
assert!(!paired_contraction_needs_regroup(&logit(), 3.6e150));
assert!(!paired_contraction_needs_regroup(&cloglog(), 3.6e150));
assert!(!paired_contraction_needs_regroup(&identity(), 0.9));
}
#[test]
fn weighted_paired_index_sums_are_stable() {
let (u0, u1, w) = (0.3_f64, 0.5_f64, 1.5_f64);
let s = weighted_paired_index_sums(&probit(), u0, u1, u1 - u0, 1.0, w).unwrap();
let a = surv_derivs(&probit(), u0).unwrap();
let b = pdf_derivs(&probit(), u1).unwrap();
assert_close(s[0], w * (a[0] + b[0]), 1e-12, 0.0, "moderate S1 event");
let (u0, delta_u, w) = (3.6e150_f64, 0.08_f64, 1.2_f64);
let u1 = u0 + delta_u;
let s = weighted_paired_index_sums(&probit(), u0, u1, delta_u, 1.0, w).unwrap();
assert_close(s[0], w * (-delta_u), 1e-6, 0.0, "far-tail S1 event");
let naive = w * (surv_derivs(&probit(), u0).unwrap()[0] + pdf_derivs(&probit(), u1).unwrap()[0]);
assert!(
naive.abs() < 1e-3 && (s[0] - w * (-delta_u)).abs() < 1e-6,
"naive S1 collapses ({naive:e}) while paired recovers {:e}",
s[0]
);
}
}