Skip to main content

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}