use num_traits::{One, Zero};
use crate::api::context::Context;
use crate::api::expr::Ex;
use crate::base::errors::SymplexError;
use crate::base::interval::Interval;
use crate::domains::stats::data::Q;
use crate::domains::stats::family::Distribution;
use crate::domains::stats::hypothesis::{Alternative, TestResult};
fn invalid(reason: impl Into<String>) -> SymplexError {
SymplexError::invalid_argument("stats::survival", reason)
}
fn qi(n: i64) -> Q {
Q::from_integer(n.into())
}
fn qu(n: usize) -> Q {
Q::from_integer(n.into())
}
#[derive(Clone, Debug, PartialEq)]
pub struct Observation {
pub time: Q,
pub event: bool,
}
impl Observation {
pub fn from_i64(times: &[i64], events: &[bool]) -> Vec<Observation> {
times
.iter()
.zip(events)
.map(|(&t, &e)| Observation {
time: qi(t),
event: e,
})
.collect()
}
pub fn from_q(times: &[Q], events: &[bool]) -> Vec<Observation> {
times
.iter()
.zip(events)
.map(|(t, &e)| Observation {
time: t.clone(),
event: e,
})
.collect()
}
}
#[derive(Clone, Debug, PartialEq)]
pub struct LifeTableRow {
pub time: Q,
pub at_risk: usize,
pub events: usize,
pub censored: usize,
pub survival: Q,
pub variance: Q,
pub cumulative_hazard: Q,
}
#[derive(Clone, Debug, PartialEq)]
pub struct KaplanMeier {
rows: Vec<LifeTableRow>,
n: usize,
}
fn validate(obs: &[Observation]) -> Result<(), SymplexError> {
if obs.is_empty() {
return Err(invalid("at least one observation is required"));
}
if obs.iter().any(|o| o.time < Q::zero()) {
return Err(invalid("times must be non-negative"));
}
Ok(())
}
fn tally(obs: &[Observation]) -> Vec<(Q, usize, usize)> {
let mut times: Vec<Q> = obs.iter().map(|o| o.time.clone()).collect();
times.sort();
times.dedup();
times
.into_iter()
.map(|t| {
let events = obs.iter().filter(|o| o.time == t && o.event).count();
let censored = obs.iter().filter(|o| o.time == t && !o.event).count();
(t, events, censored)
})
.collect()
}
impl KaplanMeier {
pub fn fit(obs: &[Observation]) -> Result<Self, SymplexError> {
validate(obs)?;
let n = obs.len();
let mut at_risk = n;
let mut survival = Q::one();
let mut greenwood = Q::zero();
let mut hazard = Q::zero();
let mut rows = Vec::new();
for (time, events, censored) in tally(obs) {
if events > 0 {
let d = qu(events);
let nn = qu(at_risk);
survival *= (&nn - &d) / &nn;
if at_risk > events {
greenwood += &d / (&nn * (&nn - &d));
}
hazard += &d / &nn;
rows.push(LifeTableRow {
time,
at_risk,
events,
censored,
survival: survival.clone(),
variance: &survival * &survival * &greenwood,
cumulative_hazard: hazard.clone(),
});
}
at_risk -= events + censored;
}
Ok(KaplanMeier { rows, n })
}
pub fn n(&self) -> usize {
self.n
}
pub fn table(&self) -> &[LifeTableRow] {
&self.rows
}
pub fn event_times(&self) -> Vec<Q> {
self.rows.iter().map(|r| r.time.clone()).collect()
}
pub fn survival_at(&self, t: &Q) -> Q {
self.rows
.iter()
.rev()
.find(|r| r.time <= *t)
.map_or_else(Q::one, |r| r.survival.clone())
}
pub fn variance_at(&self, t: &Q) -> Q {
self.rows
.iter()
.rev()
.find(|r| r.time <= *t)
.map_or_else(Q::zero, |r| r.variance.clone())
}
pub fn cumulative_hazard_at(&self, t: &Q) -> Q {
self.rows
.iter()
.rev()
.find(|r| r.time <= *t)
.map_or_else(Q::zero, |r| r.cumulative_hazard.clone())
}
pub fn quantile(&self, p: &Q) -> Option<Q> {
let target = Q::one() - p;
self.rows
.iter()
.find(|r| r.survival <= target)
.map(|r| r.time.clone())
}
pub fn median(&self) -> Option<Q> {
self.quantile(&Q::new(1.into(), 2.into()))
}
pub fn confidence_interval(
&self,
t: &Q,
confidence: f64,
method: CiMethod,
) -> Result<Interval<f64>, SymplexError> {
if !(confidence > 0.0 && confidence < 1.0) {
return Err(invalid("the confidence level must lie in (0, 1)"));
}
let ctx = Context::new();
let z =
Distribution::normal(ctx.int(0), ctx.int(1)).quantile_f64(0.5 + confidence / 2.0)?;
let s = ratio_f64(&self.survival_at(t));
let se = ratio_f64(&self.variance_at(t)).sqrt();
Ok(match method {
CiMethod::Linear => Interval::closed((s - z * se).max(0.0), (s + z * se).min(1.0)),
CiMethod::LogLog => {
if s <= 0.0 || s >= 1.0 {
Interval::closed(s, s)
} else {
let theta = z * se / (s * s.ln());
let lo = s.powf(theta.exp());
let hi = s.powf((-theta).exp());
Interval::closed(lo.min(hi), lo.max(hi))
}
}
})
}
pub fn restricted_mean(&self, tau: &Q) -> Q {
let mut area = Q::zero();
let mut prev_t = Q::zero();
let mut prev_s = Q::one();
for r in &self.rows {
if r.time >= *tau {
break;
}
area += &prev_s * (&r.time - &prev_t);
prev_t = r.time.clone();
prev_s = r.survival.clone();
}
if *tau > prev_t {
area += &prev_s * (tau - &prev_t);
}
area
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum CiMethod {
Linear,
LogLog,
}
fn ratio_f64(q: &Q) -> f64 {
use num_traits::ToPrimitive;
q.numer().to_f64().unwrap_or(f64::NAN) / q.denom().to_f64().unwrap_or(f64::NAN)
}
pub fn log_rank_test(
ctx: &Context,
obs: &[Observation],
groups: &[usize],
) -> Result<TestResult, SymplexError> {
validate(obs)?;
if obs.len() != groups.len() {
return Err(invalid("one group label per observation is required"));
}
let k = groups.iter().copied().max().map_or(0, |m| m + 1);
if k < 2 {
return Err(invalid("the log-rank test needs at least two groups"));
}
for g in 0..k {
if !groups.contains(&g) {
return Err(invalid(format!("group {g} has no observations")));
}
}
let m = k - 1;
let mut diff = vec![Q::zero(); m];
let mut v = vec![vec![Q::zero(); m]; m];
let mut any_event = false;
for (time, events, _) in tally(obs) {
if events == 0 {
continue;
}
any_event = true;
let at_risk: Vec<usize> = (0..k)
.map(|g| {
obs.iter()
.zip(groups)
.filter(|(o, gg)| **gg == g && o.time >= time)
.count()
})
.collect();
let n: usize = at_risk.iter().sum();
let d = qu(events);
let nn = qu(n);
for g in 0..m {
let observed = obs
.iter()
.zip(groups)
.filter(|(o, gg)| **gg == g && o.time == time && o.event)
.count();
let expected = &d * qu(at_risk[g]) / &nn;
diff[g] += qu(observed) - expected;
}
if n > 1 {
let scale = &d * (&nn - &d) / (&nn - Q::one());
for g in 0..m {
for h in 0..m {
let pg = qu(at_risk[g]) / &nn;
let ph = qu(at_risk[h]) / &nn;
let delta = if g == h { Q::one() } else { Q::zero() };
v[g][h] += &scale * &pg * (delta - ph);
}
}
}
}
if !any_event {
return Err(invalid("no events were observed"));
}
let vm = crate::prelude::QMatrix::new(v).map_err(|e| invalid(e.to_string()))?;
let rhs = crate::prelude::QMatrix::new(diff.iter().map(|d| vec![d.clone()]).collect())
.map_err(|e| invalid(e.to_string()))?;
let x = vm.solve(&rhs).map_err(|_| {
SymplexError::computation_failed(
"stats::survival",
"the log-rank covariance matrix is singular (a group has no risk set at every event time)",
)
})?;
let mut stat = Q::zero();
for (g, d) in diff.iter().enumerate() {
stat += d * x.get(g, 0);
}
let statistic = ctx.from_ratio(stat);
let chi = Distribution::chi_squared(ctx.int(m as i64));
let p_value = (ctx.one() - chi.cdf(&statistic)).simplify();
Ok(TestResult {
statistic,
p_value,
df: Some(ctx.int(m as i64)),
alternative: Alternative::TwoSided,
})
}
pub fn exponential_rate(obs: &[Observation]) -> Result<Q, SymplexError> {
validate(obs)?;
let total: Q = obs.iter().fold(Q::zero(), |acc, o| acc + &o.time);
if total.is_zero() {
return Err(invalid("the total time at risk is zero"));
}
let events = obs.iter().filter(|o| o.event).count();
Ok(qu(events) / total)
}
pub fn mean_event_time(obs: &[Observation]) -> Result<Q, SymplexError> {
validate(obs)?;
let events: Vec<Q> = obs
.iter()
.filter(|o| o.event)
.map(|o| o.time.clone())
.collect();
if events.is_empty() {
return Err(invalid("no events were observed"));
}
let n = events.len();
Ok(events.into_iter().fold(Q::zero(), |a, t| a + t) / qu(n))
}
pub fn survival_function(dist: &Distribution, t: &Ex) -> Ex {
let cdf = dist.family().cdf(t).unwrap_or_else(|| dist.cdf(t));
(dist.context().one() - cdf).simplify()
}
pub fn hazard_function(dist: &Distribution, t: &Ex) -> Ex {
(dist.density(t) / survival_function(dist, t)).simplify()
}