regression_diagnostics/survival/kaplan_meier.rs
1use ndarray::Array1;
2
3use crate::error::{RegressionError, Result};
4
5/// One step of a **Kaplan–Meier** survival curve, recorded at a distinct event
6/// time.
7#[derive(Debug, Clone, Copy, PartialEq)]
8pub struct KmStep {
9 /// The event time.
10 pub time: f64,
11 /// Number at risk just before this time, `nᵢ`.
12 pub at_risk: usize,
13 /// Number of events at this time, `dᵢ`.
14 pub events: usize,
15 /// Survivor function `Ŝ(t)` at and after this time.
16 pub survival: f64,
17 /// Greenwood standard error of `Ŝ(t)`.
18 pub std_error: f64,
19}
20
21/// The **Kaplan–Meier** product-limit estimator of the survivor function
22/// `S(t) = P(T > t)` from right-censored data.
23///
24/// `Ŝ(t) = Π_{tᵢ ≤ t} (1 − dᵢ/nᵢ)` over distinct event times, with `dᵢ` events
25/// and `nᵢ` subjects at risk. Censored observations leave the risk set without
26/// producing a step. Standard errors use **Greenwood's formula**. This is the
27/// non-parametric, covariate-free counterpart to the Cox model's baseline.
28#[derive(Debug, Clone)]
29pub struct KaplanMeier {
30 steps: Vec<KmStep>,
31 n: usize,
32}
33
34impl KaplanMeier {
35 /// Estimate the survival curve from event/censoring `time` and `event`
36 /// indicators (`1.0` event, `0.0` right-censored).
37 ///
38 /// # Errors
39 ///
40 /// * [`RegressionError::EmptyInput`] if there are no observations.
41 /// * [`RegressionError::ShapeMismatch`] if `time` and `event` differ in
42 /// length.
43 /// * [`RegressionError::InvalidResponse`] for non-positive times or event
44 /// flags outside `{0, 1}`.
45 pub fn new(time: Array1<f64>, event: Array1<f64>) -> Result<Self> {
46 let n = time.len();
47 if n == 0 {
48 return Err(RegressionError::EmptyInput { what: "time" });
49 }
50 if event.len() != n {
51 return Err(RegressionError::ShapeMismatch {
52 what: "event length vs time",
53 expected: n,
54 got: event.len(),
55 });
56 }
57 for &t in time.iter() {
58 if !t.is_finite() || t <= 0.0 {
59 return Err(RegressionError::InvalidResponse {
60 msg: format!("survival times must be positive, found {t}"),
61 });
62 }
63 }
64 for &e in event.iter() {
65 if e != 0.0 && e != 1.0 {
66 return Err(RegressionError::InvalidResponse {
67 msg: format!("event indicator must be 0 or 1, found {e}"),
68 });
69 }
70 }
71
72 // Distinct event times, ascending.
73 let mut ev_times: Vec<f64> = time
74 .iter()
75 .zip(event.iter())
76 .filter(|&(_, &e)| e == 1.0)
77 .map(|(&t, _)| t)
78 .collect();
79 ev_times.sort_by(|a, b| a.partial_cmp(b).unwrap());
80 ev_times.dedup();
81
82 let mut steps = Vec::with_capacity(ev_times.len());
83 let mut survival = 1.0;
84 let mut greenwood_sum = 0.0; // Σ dᵢ / (nᵢ (nᵢ − dᵢ))
85 for &t in &ev_times {
86 let at_risk = time.iter().filter(|&&tj| tj >= t).count();
87 let events = time
88 .iter()
89 .zip(event.iter())
90 .filter(|&(&tj, &e)| tj == t && e == 1.0)
91 .count();
92 let ni = at_risk as f64;
93 let di = events as f64;
94 survival *= 1.0 - di / ni;
95 if ni > di {
96 greenwood_sum += di / (ni * (ni - di));
97 }
98 let std_error = survival * greenwood_sum.sqrt();
99 steps.push(KmStep {
100 time: t,
101 at_risk,
102 events,
103 survival,
104 std_error,
105 });
106 }
107
108 Ok(Self { steps, n })
109 }
110
111 /// Number of observations the estimator was built from.
112 pub fn n_observations(&self) -> usize {
113 self.n
114 }
115
116 /// The estimated steps, one per distinct event time (ascending).
117 pub fn steps(&self) -> &[KmStep] {
118 &self.steps
119 }
120
121 /// Survivor function `Ŝ(t)` at an arbitrary time (right-continuous step
122 /// function; `1.0` before the first event).
123 pub fn survival_at(&self, t: f64) -> f64 {
124 let mut s = 1.0;
125 for step in &self.steps {
126 if step.time <= t {
127 s = step.survival;
128 } else {
129 break;
130 }
131 }
132 s
133 }
134
135 /// Median survival time — the earliest event time at which `Ŝ(t) ≤ 0.5`, or
136 /// `None` if the curve never drops to `0.5` (heavy censoring).
137 pub fn median_survival(&self) -> Option<f64> {
138 self.steps
139 .iter()
140 .find(|s| s.survival <= 0.5)
141 .map(|s| s.time)
142 }
143}