use ndarray::Array1;
use crate::error::{RegressionError, Result};
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct KmStep {
pub time: f64,
pub at_risk: usize,
pub events: usize,
pub survival: f64,
pub std_error: f64,
}
#[derive(Debug, Clone)]
pub struct KaplanMeier {
steps: Vec<KmStep>,
n: usize,
}
impl KaplanMeier {
pub fn new(time: Array1<f64>, event: Array1<f64>) -> Result<Self> {
let n = time.len();
if n == 0 {
return Err(RegressionError::EmptyInput { what: "time" });
}
if event.len() != n {
return Err(RegressionError::ShapeMismatch {
what: "event length vs time",
expected: n,
got: event.len(),
});
}
for &t in time.iter() {
if !t.is_finite() || t <= 0.0 {
return Err(RegressionError::InvalidResponse {
msg: format!("survival times must be positive, found {t}"),
});
}
}
for &e in event.iter() {
if e != 0.0 && e != 1.0 {
return Err(RegressionError::InvalidResponse {
msg: format!("event indicator must be 0 or 1, found {e}"),
});
}
}
let mut ev_times: Vec<f64> = time
.iter()
.zip(event.iter())
.filter(|&(_, &e)| e == 1.0)
.map(|(&t, _)| t)
.collect();
ev_times.sort_by(|a, b| a.partial_cmp(b).unwrap());
ev_times.dedup();
let mut steps = Vec::with_capacity(ev_times.len());
let mut survival = 1.0;
let mut greenwood_sum = 0.0; for &t in &ev_times {
let at_risk = time.iter().filter(|&&tj| tj >= t).count();
let events = time
.iter()
.zip(event.iter())
.filter(|&(&tj, &e)| tj == t && e == 1.0)
.count();
let ni = at_risk as f64;
let di = events as f64;
survival *= 1.0 - di / ni;
if ni > di {
greenwood_sum += di / (ni * (ni - di));
}
let std_error = survival * greenwood_sum.sqrt();
steps.push(KmStep {
time: t,
at_risk,
events,
survival,
std_error,
});
}
Ok(Self { steps, n })
}
pub fn n_observations(&self) -> usize {
self.n
}
pub fn steps(&self) -> &[KmStep] {
&self.steps
}
pub fn survival_at(&self, t: f64) -> f64 {
let mut s = 1.0;
for step in &self.steps {
if step.time <= t {
s = step.survival;
} else {
break;
}
}
s
}
pub fn median_survival(&self) -> Option<f64> {
self.steps
.iter()
.find(|s| s.survival <= 0.5)
.map(|s| s.time)
}
}