use ndarray::{Array1, Array2, ArrayView1, ArrayView2};
use statrs::distribution::{ContinuousCDF, Normal};
use crate::error::{RegressionError, Result};
use crate::linalg::dmatrix_from_rows;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Ties {
Efron,
Breslow,
}
#[derive(Debug, Clone)]
pub struct CoxFit {
start: Array1<f64>,
time: Array1<f64>,
event: Array1<f64>,
strata: Vec<usize>,
x: Array2<f64>,
coefficients: Array1<f64>,
cov: Array2<f64>,
log_partial_likelihood: f64,
baseline: Vec<(usize, f64, f64)>,
ties: Ties,
n_strata: usize,
iterations: usize,
n: usize,
p: usize,
}
impl CoxFit {
pub fn new(time: Array1<f64>, event: Array1<f64>, x: Array2<f64>) -> Result<Self> {
Self::with_options(time, event, x, Ties::Efron, 100, 1e-9)
}
pub fn with_options(
time: Array1<f64>,
event: Array1<f64>,
x: Array2<f64>,
ties: Ties,
max_iter: usize,
tol: f64,
) -> Result<Self> {
let n = x.nrows();
let start = Array1::<f64>::zeros(n);
let strata = vec![0usize; n];
Self::fit(start, time, event, x, strata, ties, max_iter, tol)
}
pub fn stratified(
time: Array1<f64>,
event: Array1<f64>,
x: Array2<f64>,
strata: &[usize],
ties: Ties,
) -> Result<Self> {
let n = x.nrows();
if strata.len() != n {
return Err(RegressionError::ShapeMismatch {
what: "strata length vs X rows",
expected: n,
got: strata.len(),
});
}
let start = Array1::<f64>::zeros(n);
Self::fit(start, time, event, x, densify(strata), ties, 100, 1e-9)
}
pub fn counting_process(
start: Array1<f64>,
stop: Array1<f64>,
event: Array1<f64>,
x: Array2<f64>,
ties: Ties,
) -> Result<Self> {
let n = x.nrows();
if start.len() != n {
return Err(RegressionError::ShapeMismatch {
what: "start length vs X rows",
expected: n,
got: start.len(),
});
}
for i in 0..n {
if !matches!(
start[i].partial_cmp(&stop[i]),
Some(std::cmp::Ordering::Less)
) {
return Err(RegressionError::InvalidResponse {
msg: format!("interval {i} has start {} >= stop {}", start[i], stop[i]),
});
}
}
let strata = vec![0usize; n];
Self::fit(start, stop, event, x, strata, ties, 100, 1e-9)
}
#[allow(clippy::too_many_arguments)]
fn fit(
start: Array1<f64>,
time: Array1<f64>,
event: Array1<f64>,
x: Array2<f64>,
strata: Vec<usize>,
ties: Ties,
max_iter: usize,
tol: f64,
) -> Result<Self> {
let n = x.nrows();
let p = x.ncols();
if n == 0 || p == 0 {
return Err(RegressionError::EmptyInput { what: "X" });
}
if time.len() != n || event.len() != n {
return Err(RegressionError::ShapeMismatch {
what: "time/event length vs X rows",
expected: n,
got: time.len().min(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}"),
});
}
}
let mut n_events = 0usize;
for &e in event.iter() {
if e == 1.0 {
n_events += 1;
} else if e != 0.0 {
return Err(RegressionError::InvalidResponse {
msg: format!("event indicator must be 0 or 1, found {e}"),
});
}
}
if n_events == 0 {
return Err(RegressionError::InvalidResponse {
msg: "no events observed; the partial likelihood is empty".into(),
});
}
if detect_constant_column(&x).is_some() {
return Err(RegressionError::InvalidResponse {
msg: "Cox design must not include an intercept/constant column".into(),
});
}
let n_strata = strata.iter().copied().max().map_or(0, |m| m + 1);
let mut ev_points: Vec<(usize, f64)> = (0..n)
.filter(|&i| event[i] == 1.0)
.map(|i| (strata[i], time[i]))
.collect();
ev_points.sort_by(|a, b| a.0.cmp(&b.0).then(a.1.partial_cmp(&b.1).unwrap()));
ev_points.dedup();
let mut beta = Array1::<f64>::zeros(p);
let mut cov = Array2::<f64>::zeros((p, p));
let mut iterations = 0usize;
let mut converged = false;
while iterations < max_iter {
iterations += 1;
let (_ll, score, info) =
partial_likelihood(&start, &time, &event, &x, &strata, &beta, &ev_points, ties);
let info_dm = dmatrix_from_rows(p, p, info.as_standard_layout().as_slice().unwrap());
let inv = info_dm.try_inverse().ok_or(RegressionError::RankDeficient)?;
let inv_arr = Array2::from_shape_fn((p, p), |(i, j)| inv[(i, j)]);
let delta = inv_arr.dot(&score);
beta = &beta + δ
cov = inv_arr;
let step = delta.iter().fold(0.0_f64, |m, v| m.max(v.abs()));
if !beta.iter().all(|v| v.is_finite()) || beta.iter().any(|v| v.abs() > 1e8) {
return Err(RegressionError::NotConverged {
iterations,
msg: "coefficients diverging".into(),
});
}
if step < tol {
converged = true;
break;
}
}
if !converged {
return Err(RegressionError::NotConverged {
iterations,
msg: "Newton iteration did not reach tolerance".into(),
});
}
let (log_partial_likelihood, _, _) =
partial_likelihood(&start, &time, &event, &x, &strata, &beta, &ev_points, ties);
let baseline = breslow_baseline(&start, &time, &event, &x, &strata, &beta, &ev_points);
Ok(Self {
start,
time,
event,
strata,
x,
coefficients: beta,
cov,
log_partial_likelihood,
baseline,
ties,
n_strata,
iterations,
n,
p,
})
}
pub fn n_observations(&self) -> usize {
self.n
}
pub fn n_parameters(&self) -> usize {
self.p
}
pub fn n_events(&self) -> usize {
self.event.iter().filter(|&&e| e == 1.0).count()
}
pub fn n_strata(&self) -> usize {
self.n_strata.max(1)
}
pub fn ties(&self) -> Ties {
self.ties
}
pub fn iterations(&self) -> usize {
self.iterations
}
pub fn design_matrix(&self) -> ArrayView2<'_, f64> {
self.x.view()
}
pub fn time(&self) -> ArrayView1<'_, f64> {
self.time.view()
}
pub fn event(&self) -> ArrayView1<'_, f64> {
self.event.view()
}
pub fn coefficients(&self) -> ArrayView1<'_, f64> {
self.coefficients.view()
}
pub fn hazard_ratios(&self) -> Array1<f64> {
self.coefficients.mapv(f64::exp)
}
pub fn covariance(&self) -> ArrayView2<'_, f64> {
self.cov.view()
}
pub fn log_partial_likelihood(&self) -> f64 {
self.log_partial_likelihood
}
pub fn linear_predictors(&self) -> Array1<f64> {
self.x.dot(&self.coefficients)
}
pub fn coefficient_standard_errors(&self) -> Array1<f64> {
Array1::from_shape_fn(self.p, |j| self.cov[(j, j)].max(0.0).sqrt())
}
pub fn z_values(&self) -> Array1<f64> {
let se = self.coefficient_standard_errors();
Array1::from_shape_fn(self.p, |j| {
if se[j] > 0.0 {
self.coefficients[j] / se[j]
} else {
f64::NAN
}
})
}
pub fn p_values(&self) -> Array1<f64> {
let z = self.z_values();
let normal = Normal::new(0.0, 1.0).expect("standard normal");
Array1::from_shape_fn(self.p, |j| {
if z[j].is_finite() {
2.0 * (1.0 - normal.cdf(z[j].abs()))
} else {
f64::NAN
}
})
}
pub fn aic(&self) -> f64 {
-2.0 * self.log_partial_likelihood + 2.0 * self.p as f64
}
pub fn baseline_cumulative_hazard_stratum(&self, s: usize) -> Vec<(f64, f64)> {
let mut cum = 0.0;
self.baseline
.iter()
.filter(|(st, _, _)| *st == s)
.map(|&(_, t, dh)| {
cum += dh;
(t, cum)
})
.collect()
}
pub fn baseline_cumulative_hazard(&self) -> Vec<(f64, f64)> {
self.baseline_cumulative_hazard_stratum(0)
}
pub(crate) fn cumulative_hazard_at(&self, i: usize) -> f64 {
let eta_i: f64 = (0..self.p).map(|j| self.x[(i, j)] * self.coefficients[j]).sum();
let s = self.strata[i];
let (a, b) = (self.start[i], self.time[i]);
let h0: f64 = self
.baseline
.iter()
.filter(|(st, t, _)| *st == s && *t > a && *t <= b)
.map(|(_, _, dh)| *dh)
.sum();
eta_i.exp() * h0
}
pub fn concordance(&self) -> f64 {
let eta = self.linear_predictors();
let mut concordant = 0.0;
let mut comparable = 0.0;
for i in 0..self.n {
if self.event[i] != 1.0 {
continue;
}
for j in 0..self.n {
if i == j || self.strata[i] != self.strata[j] {
continue;
}
if self.time[j] > self.time[i]
|| (self.time[j] == self.time[i] && self.event[j] == 0.0)
{
comparable += 1.0;
if eta[i] > eta[j] {
concordant += 1.0;
} else if (eta[i] - eta[j]).abs() < 1e-12 {
concordant += 0.5;
}
}
}
}
if comparable > 0.0 {
concordant / comparable
} else {
f64::NAN
}
}
pub(crate) fn coef_slice(&self) -> ArrayView1<'_, f64> {
self.coefficients.view()
}
pub(crate) fn at_risk(&self, j: usize, t: f64, s: usize) -> bool {
self.strata[j] == s && self.start[j] < t && self.time[j] >= t
}
pub(crate) fn stratum_of(&self, i: usize) -> usize {
self.strata[i]
}
}
#[allow(clippy::too_many_arguments)]
fn partial_likelihood(
start: &Array1<f64>,
time: &Array1<f64>,
event: &Array1<f64>,
x: &Array2<f64>,
strata: &[usize],
beta: &Array1<f64>,
ev_points: &[(usize, f64)],
ties: Ties,
) -> (f64, Array1<f64>, Array2<f64>) {
let n = x.nrows();
let p = x.ncols();
let eta: Vec<f64> = (0..n)
.map(|i| (0..p).map(|j| x[(i, j)] * beta[j]).sum::<f64>())
.collect();
let w: Vec<f64> = eta.iter().map(|e| e.exp()).collect();
let mut ll = 0.0;
let mut score = Array1::<f64>::zeros(p);
let mut info = Array2::<f64>::zeros((p, p));
for &(stratum, t) in ev_points {
let mut sr0 = 0.0;
let mut sr1 = vec![0.0; p];
let mut sr2 = vec![0.0; p * p];
let mut sd0 = 0.0;
let mut sd1 = vec![0.0; p];
let mut sd2 = vec![0.0; p * p];
let mut m = 0usize;
for i in 0..n {
if strata[i] != stratum {
continue;
}
if start[i] < t && time[i] >= t {
sr0 += w[i];
for a in 0..p {
sr1[a] += w[i] * x[(i, a)];
for b in 0..p {
sr2[a * p + b] += w[i] * x[(i, a)] * x[(i, b)];
}
}
if time[i] == t && event[i] == 1.0 {
m += 1;
ll += eta[i];
for a in 0..p {
score[a] += x[(i, a)];
sd1[a] += w[i] * x[(i, a)];
for b in 0..p {
sd2[a * p + b] += w[i] * x[(i, a)] * x[(i, b)];
}
}
sd0 += w[i];
}
}
}
if m == 0 {
continue;
}
let steps = match ties {
Ties::Breslow => 1,
Ties::Efron => m,
};
for l in 0..steps {
let frac = match ties {
Ties::Breslow => 0.0,
Ties::Efron => l as f64 / m as f64,
};
let mult = match ties {
Ties::Breslow => m as f64,
Ties::Efron => 1.0,
};
let d0 = sr0 - frac * sd0;
ll -= mult * d0.ln();
for a in 0..p {
let d1a = sr1[a] - frac * sd1[a];
score[a] -= mult * d1a / d0;
for b in 0..p {
let d1b = sr1[b] - frac * sd1[b];
let d2ab = sr2[a * p + b] - frac * sd2[a * p + b];
info[(a, b)] += mult * (d2ab / d0 - (d1a * d1b) / (d0 * d0));
}
}
}
}
(ll, score, info)
}
#[allow(clippy::too_many_arguments)]
fn breslow_baseline(
start: &Array1<f64>,
time: &Array1<f64>,
event: &Array1<f64>,
x: &Array2<f64>,
strata: &[usize],
beta: &Array1<f64>,
ev_points: &[(usize, f64)],
) -> Vec<(usize, f64, f64)> {
let n = x.nrows();
let p = x.ncols();
let w: Vec<f64> = (0..n)
.map(|i| (0..p).map(|j| x[(i, j)] * beta[j]).sum::<f64>().exp())
.collect();
ev_points
.iter()
.map(|&(stratum, t)| {
let mut risk = 0.0;
let mut d = 0.0;
for i in 0..n {
if strata[i] == stratum && start[i] < t && time[i] >= t {
risk += w[i];
if time[i] == t && event[i] == 1.0 {
d += 1.0;
}
}
}
(stratum, t, if risk > 0.0 { d / risk } else { 0.0 })
})
.collect()
}
fn densify(labels: &[usize]) -> Vec<usize> {
let mut map = std::collections::BTreeMap::new();
labels
.iter()
.map(|&l| {
let next = map.len();
*map.entry(l).or_insert(next)
})
.collect()
}
fn detect_constant_column(x: &Array2<f64>) -> Option<usize> {
for (j, col) in x.columns().into_iter().enumerate() {
let first = col[0];
let scale = first.abs().max(1.0);
if col.iter().all(|&v| (v - first).abs() <= 1e-12 * scale) {
return Some(j);
}
}
None
}