use ndarray::Array1;
use solow_core::error::{Error, Result};
#[derive(Clone, Debug)]
pub struct CumIncidenceRight {
pub times: Array1<f64>,
pub cinc: Vec<Array1<f64>>,
pub cinc_se: Vec<Array1<f64>>,
}
impl CumIncidenceRight {
pub fn new(time: &[f64], status: &[f64]) -> Result<Self> {
let nobs = time.len();
if status.len() != nobs {
return Err(Error::Shape("time and status length differ".into()));
}
if nobs == 0 {
return Err(Error::Shape("empty sample".into()));
}
let mut order: Vec<usize> = (0..nobs).collect();
order.sort_by(|&a, &b| time[a].total_cmp(&time[b]));
let mut utime: Vec<f64> = Vec::new();
let mut rtime: Vec<usize> = vec![0; nobs];
for &i in &order {
if utime.is_empty() || time[i] != *utime.last().unwrap() {
utime.push(time[i]);
}
rtime[i] = utime.len() - 1;
}
let ml = utime.len();
let mut d_all = vec![0.0_f64; ml];
let mut nbin = vec![0.0_f64; ml];
for i in 0..nobs {
let k = rtime[i];
nbin[k] += 1.0;
if status[i] >= 1.0 {
d_all[k] += 1.0;
}
}
let mut n = vec![0.0_f64; ml];
let mut acc = 0.0;
for k in (0..ml).rev() {
acc += nbin[k];
n[k] = acc;
}
let mut sp = vec![0.0_f64; ml];
let mut logcum = 0.0;
for k in 0..ml {
let mut frac = 1.0 - d_all[k] / n[k];
if frac < 1e-16 {
frac = 1e-16;
}
logcum += frac.ln();
sp[k] = logcum.exp();
}
let ngrp = status.iter().cloned().fold(0.0_f64, f64::max).round() as usize;
if ngrp == 0 {
return Err(Error::Shape(
"no events: status has no positive cause labels".into(),
));
}
let mut d: Vec<Vec<f64>> = vec![vec![0.0_f64; ml]; ngrp];
for i in 0..nobs {
let lab = status[i].round() as i64;
if lab >= 1 {
let j = (lab - 1) as usize;
if j < ngrp {
d[j][rtime[i]] += 1.0;
}
}
}
let mut sp0 = vec![0.0_f64; ml];
for k in 0..ml {
let s_prev = if k == 0 { 1.0 } else { sp[k - 1] };
sp0[k] = s_prev / n[k];
}
let mut ip: Vec<Vec<f64>> = vec![vec![0.0_f64; ml]; ngrp];
for j in 0..ngrp {
let mut c = 0.0;
for k in 0..ml {
c += sp0[k] * d[j][k];
ip[j][k] = c;
}
}
let mut da = vec![0.0_f64; ml];
for k in 0..ml {
for dj in d.iter() {
da[k] += dj[k];
}
}
let mut se: Vec<Array1<f64>> = Vec::with_capacity(ngrp);
for j in 0..ngrp {
let mut v = vec![0.0_f64; ml];
let mut c_ra1 = 0.0;
let mut c_ip_ra1 = 0.0;
let mut c_ip2_ra1 = 0.0;
for k in 0..ml {
let denom = n[k] * (n[k] - da[k]);
let ra1 = da[k] / denom; c_ra1 += ra1;
c_ip_ra1 += ip[j][k] * ra1;
c_ip2_ra1 += ip[j][k] * ip[j][k] * ra1;
v[k] = ip[j][k] * ip[j][k] * c_ra1 - 2.0 * ip[j][k] * c_ip_ra1 + c_ip2_ra1;
}
let mut c_sp0_ra2 = 0.0;
for k in 0..ml {
let ra2 = (n[k] - d[j][k]) * d[j][k] / n[k];
c_sp0_ra2 += sp0[k] * sp0[k] * ra2;
v[k] += c_sp0_ra2;
}
let mut c_ra3 = 0.0;
let mut c_ip_ra3 = 0.0;
for k in 0..ml {
let ra3 = sp0[k] * d[j][k] / n[k];
c_ra3 += ra3;
c_ip_ra3 += ip[j][k] * ra3;
v[k] += -2.0 * ip[j][k] * c_ra3 + 2.0 * c_ip_ra3;
}
let se_j: Vec<f64> = v.iter().map(|&x| x.sqrt()).collect();
se.push(Array1::from_vec(se_j));
}
Ok(CumIncidenceRight {
times: Array1::from_vec(utime),
cinc: ip.into_iter().map(Array1::from_vec).collect(),
cinc_se: se,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn single_cause_matches_one_minus_km() {
let time = [1.0, 2.0, 3.0, 4.0];
let status = [1.0, 1.0, 1.0, 1.0];
let ci = CumIncidenceRight::new(&time, &status).unwrap();
assert_eq!(ci.cinc.len(), 1);
let exp = [0.25, 0.5, 0.75, 1.0];
for (k, &e) in exp.iter().enumerate() {
assert!((ci.cinc[0][k] - e).abs() < 1e-12, "k={k}");
}
}
#[test]
fn two_causes_sum_to_one_minus_survival() {
let time = [1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
let status = [1.0, 2.0, 1.0, 2.0, 1.0, 2.0];
let ci = CumIncidenceRight::new(&time, &status).unwrap();
assert_eq!(ci.cinc.len(), 2);
let last = ci.times.len() - 1;
let total = ci.cinc[0][last] + ci.cinc[1][last];
assert!((total - 1.0).abs() < 1e-12, "total={total}");
}
#[test]
fn length_mismatch_errors() {
assert!(CumIncidenceRight::new(&[1.0, 2.0], &[1.0]).is_err());
}
}