use super::*;
use gam_math::probability::{normal_cdf, normal_logcdf_derivatives, normal_pdf};
pub(crate) use crate::sigma_link::survival_q0_from_eta;
#[inline]
pub(crate) fn probit_survival_value(eta: f64) -> f64 {
if eta.is_nan() {
f64::NAN
} else if eta == f64::INFINITY {
0.0
} else if eta == f64::NEG_INFINITY {
1.0
} else {
normal_cdf(-eta)
}
}
#[inline]
pub(crate) fn probit_log_survival_and_ratio_derivatives(eta: f64) -> (f64, f64, f64, f64, f64) {
let d = normal_logcdf_derivatives(-eta);
(d[0], d[1], -d[2], d[3], -d[4])
}
#[cfg(test)]
mod probit_tail_tests {
use super::probit_log_survival_and_ratio_derivatives;
#[test]
fn probit_survival_derivatives_have_exact_infinite_limits() {
assert_eq!(
probit_log_survival_and_ratio_derivatives(f64::NEG_INFINITY),
(0.0, 0.0, -0.0, 0.0, -0.0)
);
assert_eq!(
probit_log_survival_and_ratio_derivatives(f64::INFINITY),
(f64::NEG_INFINITY, f64::INFINITY, 1.0, 0.0, -0.0)
);
let nan_derivatives = probit_log_survival_and_ratio_derivatives(f64::NAN);
assert!(
[
nan_derivatives.0,
nan_derivatives.1,
nan_derivatives.2,
nan_derivatives.3,
nan_derivatives.4,
]
.into_iter()
.all(f64::is_nan)
);
}
#[test]
fn probit_survival_derivatives_preserve_both_extreme_tails() {
let (_, ratio, dr, ddr, dddr) = probit_log_survival_and_ratio_derivatives(1.0e100);
assert_eq!(ratio, 1.0e100);
assert_eq!(dr, 1.0);
assert!(ddr > 0.0 && ddr.is_finite());
assert_eq!(dddr, -0.0);
let (_, ratio, dr, ddr, dddr) = probit_log_survival_and_ratio_derivatives(-38.6);
assert_eq!(ratio, 0.0);
assert!(dr > 0.0 && dr.is_subnormal());
assert!(ddr > 0.0 && ddr.is_subnormal());
assert!(dddr > 0.0 && dddr.is_subnormal());
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub enum ResidualDistribution {
Gaussian,
Gumbel,
Logistic,
}
pub trait ResidualDistributionOps {
fn cdf(&self, z: f64) -> f64;
fn pdf(&self, z: f64) -> f64;
fn pdf_derivative(&self, z: f64) -> f64;
fn pdfsecond_derivative(&self, z: f64) -> f64;
fn pdfthird_derivative(&self, z: f64) -> f64;
fn pdffourth_derivative(&self, z: f64) -> f64;
}
impl ResidualDistributionOps for ResidualDistribution {
fn cdf(&self, z: f64) -> f64 {
match self {
ResidualDistribution::Gaussian => normal_cdf(z),
ResidualDistribution::Gumbel => {
component_inverse_link_jet(gam_problem::LinkComponent::CLogLog, z).mu
}
ResidualDistribution::Logistic => {
component_inverse_link_jet(gam_problem::LinkComponent::Logit, z).mu
}
}
}
fn pdf(&self, z: f64) -> f64 {
match self {
ResidualDistribution::Gaussian => normal_pdf(z),
ResidualDistribution::Gumbel => {
component_inverse_link_jet(gam_problem::LinkComponent::CLogLog, z).d1
}
ResidualDistribution::Logistic => {
component_inverse_link_jet(gam_problem::LinkComponent::Logit, z).d1
}
}
}
fn pdf_derivative(&self, z: f64) -> f64 {
match self {
ResidualDistribution::Gaussian => -z * normal_pdf(z),
ResidualDistribution::Gumbel => {
component_inverse_link_jet(gam_problem::LinkComponent::CLogLog, z).d2
}
ResidualDistribution::Logistic => {
component_inverse_link_jet(gam_problem::LinkComponent::Logit, z).d2
}
}
}
fn pdfsecond_derivative(&self, z: f64) -> f64 {
match self {
ResidualDistribution::Gaussian => {
let f = normal_pdf(z);
(z * z - 1.0) * f
}
ResidualDistribution::Gumbel => {
component_inverse_link_jet(gam_problem::LinkComponent::CLogLog, z).d3
}
ResidualDistribution::Logistic => {
component_inverse_link_jet(gam_problem::LinkComponent::Logit, z).d3
}
}
}
fn pdfthird_derivative(&self, z: f64) -> f64 {
match self {
ResidualDistribution::Gaussian => {
let f = normal_pdf(z);
-(z * z * z - 3.0 * z) * f
}
ResidualDistribution::Gumbel => inverse_link_pdfthird_derivative_for_inverse_link(
&InverseLink::Standard(StandardLink::CLogLog),
z,
)
.expect("standard cloglog inverse-link third derivative should evaluate"),
ResidualDistribution::Logistic => inverse_link_pdfthird_derivative_for_inverse_link(
&InverseLink::Standard(StandardLink::Logit),
z,
)
.expect("standard logit inverse-link third derivative should evaluate"),
}
}
fn pdffourth_derivative(&self, z: f64) -> f64 {
match self {
ResidualDistribution::Gaussian => {
let f = normal_pdf(z);
let z2 = z * z;
(z2 * z2 - 6.0 * z2 + 3.0) * f
}
ResidualDistribution::Gumbel => inverse_link_pdffourth_derivative_for_inverse_link(
&InverseLink::Standard(StandardLink::CLogLog),
z,
)
.expect("standard cloglog inverse-link fourth derivative should evaluate"),
ResidualDistribution::Logistic => inverse_link_pdffourth_derivative_for_inverse_link(
&InverseLink::Standard(StandardLink::Logit),
z,
)
.expect("standard logit inverse-link fourth derivative should evaluate"),
}
}
}
#[inline]
pub(crate) fn residual_distribution_link(distribution: ResidualDistribution) -> StandardLink {
match distribution {
ResidualDistribution::Gaussian => StandardLink::Probit,
ResidualDistribution::Gumbel => StandardLink::CLogLog,
ResidualDistribution::Logistic => StandardLink::Logit,
}
}
#[inline]
pub fn residual_distribution_inverse_link(distribution: ResidualDistribution) -> InverseLink {
InverseLink::Standard(residual_distribution_link(distribution))
}
#[inline]
pub fn residual_distribution_from_inverse_link(link: &InverseLink) -> Option<ResidualDistribution> {
match link {
InverseLink::Standard(StandardLink::Probit) => Some(ResidualDistribution::Gaussian),
InverseLink::Standard(StandardLink::CLogLog) => Some(ResidualDistribution::Gumbel),
InverseLink::Standard(StandardLink::Logit) => Some(ResidualDistribution::Logistic),
_ => None,
}
}
pub(crate) fn inverse_link_pdffourth_derivative(
inverse_link: &InverseLink,
eta: f64,
) -> Result<f64, SurvivalLocationScaleError> {
match inverse_link {
InverseLink::Standard(StandardLink::Probit) => {
Ok(ResidualDistribution::Gaussian.pdffourth_derivative(eta))
}
InverseLink::Standard(StandardLink::Logit) => {
Ok(ResidualDistribution::Logistic.pdffourth_derivative(eta))
}
InverseLink::Standard(StandardLink::CLogLog) => {
Ok(ResidualDistribution::Gumbel.pdffourth_derivative(eta))
}
_ => gam_solve::mixture_link::inverse_link_pdffourth_derivative_for_inverse_link(
inverse_link,
eta,
)
.map_err(|e| SurvivalLocationScaleError::NumericalFailure {
reason: format!("inverse link fourth-derivative evaluation failed at eta={eta}: {e}"),
}),
}
}