qualia_core_db/solvers/learning/survival/
kaplan_meier.rs1use crate::solvers::learning::LearningError;
9
10#[derive(Debug, Clone)]
12pub struct KaplanMeier {
13 pub event_times: Vec<f64>,
15 pub survival: Vec<f64>,
17 pub at_risk: Vec<usize>,
19 pub events: Vec<usize>,
21}
22
23impl KaplanMeier {
24 pub fn fit(times: &[f64], event: &[bool]) -> Result<Self, LearningError> {
27 let n = times.len();
28 if n == 0 || n != event.len() {
29 return Err(LearningError::InvalidDimension);
30 }
31 let mut order: Vec<usize> = (0..n).collect();
34 order.sort_by(|&a, &b| {
35 times[a]
36 .partial_cmp(×[b])
37 .unwrap_or(core::cmp::Ordering::Equal)
38 });
39
40 let mut event_times = Vec::new();
41 let mut survival = Vec::new();
42 let mut at_risk_v = Vec::new();
43 let mut events_v = Vec::new();
44
45 let mut s = 1.0;
46 let mut i = 0;
47 while i < n {
48 let t = times[order[i]];
49 let mut d = 0usize; let mut tied = 0usize; let mut j = i;
53 while j < n && times[order[j]] == t {
54 if event[order[j]] {
55 d += 1;
56 }
57 tied += 1;
58 j += 1;
59 }
60 let n_at_risk = n - i; if d > 0 {
62 s *= 1.0 - d as f64 / n_at_risk as f64;
63 event_times.push(t);
64 survival.push(s);
65 at_risk_v.push(n_at_risk);
66 events_v.push(d);
67 }
68 let _ = tied;
69 i = j;
70 }
71
72 Ok(Self {
73 event_times,
74 survival,
75 at_risk: at_risk_v,
76 events: events_v,
77 })
78 }
79
80 pub fn survival_at(&self, t: f64) -> f64 {
83 let mut s = 1.0;
84 for (k, &et) in self.event_times.iter().enumerate() {
85 if et <= t {
86 s = self.survival[k];
87 } else {
88 break;
89 }
90 }
91 s
92 }
93
94 pub fn median_survival(&self) -> Option<f64> {
97 self.event_times
98 .iter()
99 .zip(self.survival.iter())
100 .find(|(_, &s)| s <= 0.5)
101 .map(|(&t, _)| t)
102 }
103}
104
105#[cfg(test)]
106mod tests {
107 use super::*;
108
109 #[test]
110 fn no_censoring_matches_empirical() {
111 let times = [1.0, 2.0, 3.0, 4.0];
113 let event = [true, true, true, true];
114 let km = KaplanMeier::fit(×, &event).unwrap();
115 assert!((km.survival_at(1.0) - 0.75).abs() < 1e-12);
116 assert!((km.survival_at(2.0) - 0.5).abs() < 1e-12);
117 assert!((km.survival_at(3.0) - 0.25).abs() < 1e-12);
118 assert!((km.survival_at(0.5) - 1.0).abs() < 1e-12); assert_eq!(km.median_survival(), Some(2.0));
120 }
121
122 #[test]
123 fn censoring_keeps_survival_higher() {
124 let times = [1.0, 2.0, 3.0, 4.0];
126 let event = [true, false, true, false];
127 let km = KaplanMeier::fit(×, &event).unwrap();
128 assert!((km.survival_at(1.0) - 0.75).abs() < 1e-12);
130 assert!((km.survival_at(3.0) - 0.375).abs() < 1e-12);
132 }
133
134 #[test]
135 fn tied_events_drop_together() {
136 let times = [1.0, 2.0, 2.0, 4.0];
138 let event = [true, true, true, true];
139 let km = KaplanMeier::fit(×, &event).unwrap();
140 assert!((km.survival_at(2.0) - 0.25).abs() < 1e-12);
142 }
143
144 #[test]
145 fn guards() {
146 assert_eq!(
147 KaplanMeier::fit(&[], &[]).unwrap_err(),
148 LearningError::InvalidDimension
149 );
150 assert_eq!(
151 KaplanMeier::fit(&[1.0], &[true, false]).unwrap_err(),
152 LearningError::InvalidDimension
153 );
154 }
155}