use ndarray::{Array1, Array2};
use solow_core::error::{Error, Result};
use solow_distributions::chi2_sf;
use solow_linalg::solve;
#[derive(Clone, Copy, Debug, PartialEq)]
pub enum WeightType {
LogRank,
GehanBreslow,
TaroneWare,
FlemingHarrington(f64),
}
#[derive(Clone, Debug)]
pub struct SurvDiffResult {
pub chisq: f64,
pub pvalue: f64,
pub df: usize,
}
pub fn survdiff(
time: &[f64],
status: &[f64],
group: &[f64],
weight_type: WeightType,
) -> Result<SurvDiffResult> {
let n = time.len();
if status.len() != n || group.len() != n {
return Err(Error::Shape("time/status/group length mismatch".into()));
}
if n == 0 {
return Err(Error::Shape("empty sample".into()));
}
let mut gr: Vec<f64> = group.to_vec();
gr.sort_by(|a, b| a.total_cmp(b));
gr.dedup();
let ng = gr.len();
if ng < 2 {
return Err(Error::Shape("survdiff requires at least two groups".into()));
}
let mut utimes: Vec<f64> = time.to_vec();
utimes.sort_by(|a, b| a.total_cmp(b));
utimes.dedup();
let ml = utimes.len();
let time_rank = |t: f64| -> usize { utimes.partition_point(|&u| u < t) };
let mut obsv: Vec<Array1<f64>> = vec![Array1::zeros(ml); ng];
let mut nrisk: Vec<Array1<f64>> = vec![Array1::zeros(ml); ng];
let group_idx = |g: f64| -> usize { gr.iter().position(|&x| x == g).unwrap() };
let mut nbin: Vec<Array1<f64>> = vec![Array1::zeros(ml); ng];
for i in 0..n {
let gi = group_idx(group[i]);
let k = time_rank(time[i]);
nbin[gi][k] += 1.0;
if status[i].round() as i64 == 1 {
obsv[gi][k] += 1.0;
}
}
for g in 0..ng {
let mut acc = 0.0;
for k in (0..ml).rev() {
acc += nbin[g][k];
nrisk[g][k] = acc;
}
}
let mut obs = Array1::<f64>::zeros(ml);
let mut nrisk_tot = Array1::<f64>::zeros(ml);
for g in 0..ng {
for k in 0..ml {
obs[k] += obsv[g][k];
nrisk_tot[k] += nrisk[g][k];
}
}
let ix: Vec<usize> = (0..ml).filter(|&k| nrisk_tot[k] > 1.0).collect();
let weights: Option<Array1<f64>> = match weight_type {
WeightType::LogRank => None,
WeightType::GehanBreslow => Some(nrisk_tot.clone()),
WeightType::TaroneWare => Some(nrisk_tot.mapv(f64::sqrt)),
WeightType::FlemingHarrington(p) => {
let mut sp = Array1::<f64>::zeros(ml);
let mut logcum = 0.0;
for k in 0..ml {
let frac = 1.0 - obs[k] / nrisk_tot[k];
logcum += frac.ln();
sp[k] = logcum.exp();
}
let mut w = sp.mapv(|v| v.powf(p));
let mut rolled = Array1::<f64>::zeros(ml);
for k in 0..ml {
rolled[k] = w[(k + ml - 1) % ml];
}
rolled[0] = 1.0;
w = rolled;
Some(w)
}
};
let dfs = ng - 1;
let mut r: Vec<Array1<f64>> = vec![Array1::zeros(ml); ng];
for g in 0..ng {
for k in 0..ml {
let denom = nrisk_tot[k].max(1e-10);
r[g][k] = nrisk[g][k] / denom;
}
}
let var_denom: Array1<f64> = nrisk_tot.mapv(|v| (v - 1.0).max(1e-10));
let var_scalar: Array1<f64> =
Array1::from_iter((0..ml).map(|k| obs[k] * (nrisk_tot[k] - obs[k]) / var_denom[k]));
let mut obs_vec = Array1::<f64>::zeros(dfs);
let mut var_mat = Array2::<f64>::zeros((dfs, dfs));
for g in 1..=dfs {
let mut oe = Array1::<f64>::zeros(ml);
for k in 0..ml {
oe[k] = obsv[g][k] - r[g][k] * obs[k];
}
let mut var_row = Array1::<f64>::zeros(dfs);
let (oe_w, w2): (Array1<f64>, Option<Array1<f64>>) = match &weights {
None => (oe, None),
Some(w) => {
let mut oew = Array1::<f64>::zeros(ml);
for k in 0..ml {
oew[k] = w[k] * oe[k];
}
(oew, Some(w.mapv(|v| v * v)))
}
};
for &k in &ix {
obs_vec[g - 1] += oe_w[k];
for (ci, c) in (1..=dfs).enumerate() {
let ind = if c == g { 1.0 } else { 0.0 };
let mut v = r[c][k] * (ind - r[g][k]) * var_scalar[k];
if let Some(w2v) = &w2 {
v *= w2v[k];
}
var_row[ci] += v;
}
}
for ci in 0..dfs {
var_mat[[g - 1, ci]] = var_row[ci];
}
}
let sol = solve(&var_mat, &obs_vec)?;
let chisq = obs_vec.dot(&sol);
let pvalue = chi2_sf(chisq, dfs as f64);
Ok(SurvDiffResult {
chisq,
pvalue,
df: dfs,
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn identical_groups_zero_statistic() {
let time = [1.0, 2.0, 3.0, 4.0, 1.0, 2.0, 3.0, 4.0];
let status = [1.0, 1.0, 1.0, 0.0, 1.0, 1.0, 1.0, 0.0];
let group = [0.0, 0.0, 0.0, 0.0, 1.0, 1.0, 1.0, 1.0];
let res = survdiff(&time, &status, &group, WeightType::LogRank).unwrap();
assert!(res.chisq.abs() < 1e-12, "chisq={}", res.chisq);
assert!((res.pvalue - 1.0).abs() < 1e-12);
assert_eq!(res.df, 1);
}
#[test]
fn requires_two_groups() {
let time = [1.0, 2.0, 3.0];
let status = [1.0, 1.0, 1.0];
let group = [0.0, 0.0, 0.0];
assert!(survdiff(&time, &status, &group, WeightType::LogRank).is_err());
}
#[test]
fn statistic_nonnegative() {
let time = [1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
let status = [1.0, 1.0, 0.0, 1.0, 1.0, 1.0];
let group = [0.0, 1.0, 0.0, 1.0, 0.0, 1.0];
for wt in [
WeightType::LogRank,
WeightType::GehanBreslow,
WeightType::TaroneWare,
WeightType::FlemingHarrington(1.0),
] {
let res = survdiff(&time, &status, &group, wt).unwrap();
assert!(res.chisq >= 0.0);
assert!((0.0..=1.0).contains(&res.pvalue));
}
}
}