solow_duration/
survfunc.rs1use ndarray::Array1;
11use solow_core::error::{Error, Result};
12
13#[derive(Clone, Debug)]
18pub struct SurvfuncRight {
19 pub surv_times: Array1<f64>,
21 pub surv_prob: Array1<f64>,
23 pub surv_prob_se: Array1<f64>,
26 pub n_risk: Array1<f64>,
28 pub n_events: Array1<f64>,
30}
31
32impl SurvfuncRight {
33 pub fn new(time: &[f64], status: &[f64]) -> Result<Self> {
38 if time.len() != status.len() {
39 return Err(Error::Shape("time and status length differ".into()));
40 }
41 if time.is_empty() {
42 return Err(Error::Shape("empty time vector".into()));
43 }
44
45 let mut order: Vec<usize> = (0..time.len()).collect();
48 order.sort_by(|&a, &b| time[a].total_cmp(&time[b]));
49 let mut utime: Vec<f64> = Vec::new();
50 let mut rtime: Vec<usize> = vec![0; time.len()];
51 for &i in &order {
52 if utime.is_empty() || time[i] != *utime.last().unwrap() {
53 utime.push(time[i]);
54 }
55 rtime[i] = utime.len() - 1;
56 }
57 let ml = utime.len();
58
59 let mut d = vec![0.0_f64; ml];
62 let mut raw_n = vec![0.0_f64; ml];
63 for i in 0..time.len() {
64 let k = rtime[i];
65 raw_n[k] += 1.0;
66 if status[i].round() as i64 == 1 {
67 d[k] += 1.0;
68 }
69 }
70
71 let mut n = vec![0.0_f64; ml];
74 let mut acc = 0.0;
75 for k in (0..ml).rev() {
76 acc += raw_n[k];
77 n[k] = acc;
78 }
79
80 let keep: Vec<usize> = (0..ml).filter(|&k| d[k] > 0.0).collect();
82 let nk = keep.len();
83 let dk: Vec<f64> = keep.iter().map(|&k| d[k]).collect();
84 let nrisk: Vec<f64> = keep.iter().map(|&k| n[k]).collect();
85 let times: Vec<f64> = keep.iter().map(|&k| utime[k]).collect();
86
87 let mut sp = vec![0.0_f64; nk];
89 let mut zero_flag = vec![false; nk];
90 let mut log_cumsum = 0.0;
91 for j in 0..nk {
92 let mut frac = 1.0 - dk[j] / nrisk[j];
93 if frac < 1e-16 {
94 frac = 1e-16;
95 zero_flag[j] = true;
96 }
97 log_cumsum += frac.ln();
98 sp[j] = log_cumsum.exp();
99 if zero_flag[j] {
100 sp[j] = 0.0;
101 }
102 }
103
104 let mut se = vec![0.0_f64; nk];
108 let mut cum = 0.0;
109 for j in 0..nk {
110 let denom = (nrisk[j] * (nrisk[j] - dk[j])).max(1e-12);
111 let mut term = dk[j] / denom;
112 if nrisk[j] == dk[j] || nrisk[j] == 0.0 {
113 term = f64::NAN;
114 }
115 cum += term;
116 let s = cum.sqrt();
117 if s.is_finite() || sp[j] != 0.0 {
119 se[j] = s * sp[j];
120 } else {
121 se[j] = f64::NAN;
122 }
123 }
124
125 Ok(SurvfuncRight {
126 surv_times: Array1::from_vec(times),
127 surv_prob: Array1::from_vec(sp),
128 surv_prob_se: Array1::from_vec(se),
129 n_risk: Array1::from_vec(nrisk),
130 n_events: Array1::from_vec(dk),
131 })
132 }
133}
134
135#[cfg(test)]
136mod tests {
137 use super::*;
138
139 #[test]
140 fn all_events_no_censoring() {
141 let time = [1.0, 2.0, 3.0, 4.0];
144 let status = [1.0, 1.0, 1.0, 1.0];
145 let s = SurvfuncRight::new(&time, &status).unwrap();
146 assert_eq!(s.surv_times.to_vec(), vec![1.0, 2.0, 3.0, 4.0]);
147 let exp = [0.75, 0.5, 0.25, 0.0];
148 for (i, &e) in exp.iter().enumerate() {
149 assert!((s.surv_prob[i] - e).abs() < 1e-12);
150 }
151 assert!(s.surv_prob_se[3].is_nan());
153 }
154
155 #[test]
156 fn censored_times_excluded_from_surv_times() {
157 let time = [1.0, 2.0, 3.0];
159 let status = [1.0, 0.0, 1.0];
160 let s = SurvfuncRight::new(&time, &status).unwrap();
161 assert_eq!(s.surv_times.to_vec(), vec![1.0, 3.0]);
162 assert!((s.surv_prob[0] - (2.0 / 3.0)).abs() < 1e-12);
164 assert!((s.surv_prob[1]).abs() < 1e-12);
165 }
166
167 #[test]
168 fn mismatched_lengths_error() {
169 assert!(SurvfuncRight::new(&[1.0, 2.0], &[1.0]).is_err());
170 }
171}