use super::common::{information_cholesky, invalid, qu, wald_summary};
use super::data::Q;
use super::hypothesis::{Alternative, TestResult};
use super::regression::{LikelihoodFit, WaldFit, chi_squared_test_result};
use super::survival::Observation;
use crate::api::context::Context;
use crate::base::dense_f64::{self, dot};
use crate::base::errors::SymplexError;
use crate::base::interval::Interval;
const DIVERGENCE_BOUND: f64 = 25.0;
const MAX_HALVINGS: u32 = 30;
fn failed(op: &'static str, reason: impl Into<String>) -> SymplexError {
SymplexError::computation_failed(op, reason)
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub enum Ties {
Breslow,
#[default]
Efron,
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct CoxOpts {
pub ties: Ties,
pub max_iter: usize,
pub tol: f64,
}
impl Default for CoxOpts {
fn default() -> Self {
Self {
ties: Ties::Efron,
max_iter: 50,
tol: 1e-9,
}
}
}
#[derive(Clone, Debug, PartialEq)]
pub struct BaselineHazardRow {
pub stratum: usize,
pub time: Q,
pub hazard: f64,
pub cumulative: f64,
}
#[derive(Clone, Debug, PartialEq)]
pub struct CoxModel {
pub coefficients: Vec<f64>,
pub standard_errors: Vec<f64>,
pub z_values: Vec<f64>,
pub p_values: Vec<f64>,
pub log_likelihood: f64,
pub null_log_likelihood: f64,
pub cov_params: Vec<Vec<f64>>,
pub nobs: usize,
pub n_events: usize,
pub ties: Ties,
pub iterations: usize,
pub converged: bool,
obs: Vec<Observation>,
x: Vec<Vec<f64>>,
strata: Vec<usize>,
table: Vec<EventTime>,
score_statistic: f64,
wald_statistic: f64,
}
#[derive(Clone, Debug, PartialEq)]
struct EventTime {
stratum: usize,
time: Q,
enter: Vec<usize>,
events: Vec<usize>,
}
struct RiskSums {
s0: f64,
s1: Vec<f64>,
s2: Vec<Vec<f64>>,
d0: f64,
d1: Vec<f64>,
d2: Vec<Vec<f64>>,
}
struct Evaluation {
ll: f64,
score: Vec<f64>,
info: Vec<Vec<f64>>,
}
fn build_table(obs: &[Observation], strata: &[usize]) -> Vec<EventTime> {
let mut labels: Vec<usize> = strata.to_vec();
labels.sort_unstable();
labels.dedup();
let mut table = Vec::new();
for s in labels {
let members: Vec<usize> = (0..obs.len()).filter(|&i| strata[i] == s).collect();
let mut times: Vec<Q> = members
.iter()
.filter(|&&i| obs[i].event)
.map(|&i| obs[i].time.clone())
.collect();
times.sort();
times.dedup();
for (k, t) in times.iter().enumerate().rev() {
let upper = times.get(k + 1);
let enter = members
.iter()
.copied()
.filter(|&i| obs[i].time >= *t && upper.is_none_or(|u| obs[i].time < *u))
.collect();
let events = members
.iter()
.copied()
.filter(|&i| obs[i].event && obs[i].time == *t)
.collect();
table.push(EventTime {
stratum: s,
time: t.clone(),
enter,
events,
});
}
}
table
}
fn sweep(table: &[EventTime], x: &[Vec<f64>], w: &[f64], p: usize) -> Vec<RiskSums> {
let mut out = Vec::with_capacity(table.len());
let mut s0 = 0.0;
let mut s1 = vec![0.0; p];
let mut s2 = vec![vec![0.0; p]; p];
let mut stratum = None;
for entry in table {
if stratum != Some(entry.stratum) {
stratum = Some(entry.stratum);
s0 = 0.0;
s1.iter_mut().for_each(|v| *v = 0.0);
s2.iter_mut().flatten().for_each(|v| *v = 0.0);
}
for &i in &entry.enter {
let wi = w[i];
s0 += wi;
for a in 0..p {
s1[a] += wi * x[i][a];
for b in 0..p {
s2[a][b] += wi * x[i][a] * x[i][b];
}
}
}
let mut d0 = 0.0;
let mut d1 = vec![0.0; p];
let mut d2 = vec![vec![0.0; p]; p];
for &i in &entry.events {
let wi = w[i];
d0 += wi;
for a in 0..p {
d1[a] += wi * x[i][a];
for b in 0..p {
d2[a][b] += wi * x[i][a] * x[i][b];
}
}
}
out.push(RiskSums {
s0,
s1: s1.clone(),
s2: s2.clone(),
d0,
d1,
d2,
});
}
out
}
struct LinearPredictor {
eta: Vec<f64>,
shift: f64,
w: Vec<f64>,
}
fn linear_predictor(x: &[Vec<f64>], beta: &[f64]) -> LinearPredictor {
let eta: Vec<f64> = x.iter().map(|row| dot(row, beta)).collect();
let shift = eta.iter().copied().fold(f64::NEG_INFINITY, f64::max);
let w = eta.iter().map(|e| (e - shift).exp()).collect();
LinearPredictor { eta, shift, w }
}
fn evaluate(table: &[EventTime], x: &[Vec<f64>], beta: &[f64], ties: Ties) -> Evaluation {
let p = beta.len();
let lp = linear_predictor(x, beta);
let sums = sweep(table, x, &lp.w, p);
let mut ll = 0.0;
let mut score = vec![0.0; p];
let mut info = vec![vec![0.0; p]; p];
for (entry, s) in table.iter().zip(&sums) {
let d = entry.events.len();
for &i in &entry.events {
ll += lp.eta[i] - lp.shift;
for a in 0..p {
score[a] += x[i][a];
}
}
for l in 0..d {
let f = match ties {
Ties::Efron => l as f64 / d as f64,
Ties::Breslow => 0.0,
};
let a0 = s.s0 - f * s.d0;
ll -= a0.ln();
let mean: Vec<f64> = (0..p).map(|a| (s.s1[a] - f * s.d1[a]) / a0).collect();
for a in 0..p {
score[a] -= mean[a];
for b in 0..p {
info[a][b] += (s.s2[a][b] - f * s.d2[a][b]) / a0 - mean[a] * mean[b];
}
}
}
}
Evaluation { ll, score, info }
}
fn validate(
op: &'static str,
obs: &[Observation],
x: &[Vec<f64>],
strata: &[usize],
opts: &CoxOpts,
) -> Result<usize, SymplexError> {
let n = obs.len();
if n == 0 {
return Err(invalid(op, "at least one observation is required"));
}
if x.len() != n {
return Err(invalid(
op,
format!("{n} observations but x has {} rows", x.len()),
));
}
if strata.len() != n {
return Err(invalid(
op,
format!("{n} observations but {} stratum labels", strata.len()),
));
}
let p = x.first().map_or(0, Vec::len);
if p == 0 {
return Err(invalid(op, "at least one covariate is required"));
}
if let Some((i, r)) = x.iter().enumerate().find(|(_, r)| r.len() != p) {
return Err(invalid(
op,
format!("row {i} of x has {} entries, expected {p}", r.len()),
));
}
if let Some((i, j)) = x
.iter()
.enumerate()
.find_map(|(i, r)| r.iter().position(|v| !v.is_finite()).map(|j| (i, j)))
{
return Err(invalid(op, format!("x[{i}][{j}] is not finite")));
}
if let Some(j) = (0..p).find(|&j| x.iter().all(|r| r[j] == x[0][j])) {
return Err(invalid(
op,
format!("covariate {j} is constant: its coefficient is not identified"),
));
}
if !obs.iter().any(|o| o.event) {
return Err(invalid(op, "no events were observed"));
}
if opts.max_iter == 0 {
return Err(invalid(op, "max_iter must be positive"));
}
if opts.tol.is_nan() || opts.tol <= 0.0 {
return Err(invalid(
op,
format!("tol must be positive, got {}", opts.tol),
));
}
Ok(p)
}
pub fn cox_ph(
obs: &[Observation],
x: &[Vec<f64>],
opts: &CoxOpts,
) -> Result<CoxModel, SymplexError> {
let strata = vec![0; obs.len()];
fit("cox_ph", obs, x, &strata, opts)
}
pub fn cox_ph_stratified(
obs: &[Observation],
x: &[Vec<f64>],
strata: &[usize],
opts: &CoxOpts,
) -> Result<CoxModel, SymplexError> {
fit("cox_ph_stratified", obs, x, strata, opts)
}
fn fit(
op: &'static str,
obs: &[Observation],
x: &[Vec<f64>],
strata: &[usize],
opts: &CoxOpts,
) -> Result<CoxModel, SymplexError> {
let p = validate(op, obs, x, strata, opts)?;
let table = build_table(obs, strata);
let ties = opts.ties;
let mut beta = vec![0.0; p];
let mut current = evaluate(&table, x, &beta, ties);
let null_log_likelihood = current.ll;
let singular_at_start = || {
failed(
op,
"the information matrix at β = 0 is singular: a covariate is collinear with the others or constant within every risk set",
)
};
let l0 = information_cholesky(¤t.info).ok_or_else(singular_at_start)?;
let score_statistic = dense_f64::quadratic_form(&l0, p, ¤t.score);
let mut iterations = 0;
let mut converged = false;
for iter in 1..=opts.max_iter {
iterations = iter;
let l = if iter == 1 {
l0.clone()
} else {
information_cholesky(¤t.info).ok_or_else(|| {
failed(
op,
"the information matrix became singular: the partial likelihood has no finite maximiser",
)
})?
};
let direction = dense_f64::cholesky_solve(&l, p, ¤t.score);
let mut lambda = 1.0;
let mut halvings = 0;
let (next_beta, next) = loop {
let candidate: Vec<f64> = beta
.iter()
.zip(&direction)
.map(|(b, d)| b + lambda * d)
.collect();
let eval = evaluate(&table, x, &candidate, ties);
let acceptable =
eval.ll.is_finite() && eval.ll >= current.ll - 1e-12 * current.ll.abs();
if acceptable || halvings >= MAX_HALVINGS {
break (candidate, eval);
}
lambda *= 0.5;
halvings += 1;
};
let max_step = direction
.iter()
.fold(0.0_f64, |m, d| m.max((lambda * d).abs()));
let delta_ll = (next.ll - current.ll).abs();
beta = next_beta;
current = next;
if beta.iter().any(|b| !b.is_finite()) || !current.ll.is_finite() {
return Err(failed(
op,
"the coefficients diverged: the partial likelihood has no finite maximiser",
));
}
let scale = beta.iter().fold(1.0_f64, |m, b| m.max(b.abs()));
let step_small = max_step <= opts.tol * scale;
let runaway = (0..p)
.filter(|&j| beta[j].abs() > DIVERGENCE_BOUND)
.max_by(|&a, &b| beta[a].abs().total_cmp(&beta[b].abs()));
if let (false, Some(j)) = (step_small, runaway) {
return Err(failed(
op,
format!(
"monotone likelihood: the coefficient of covariate {j} passed {} (β = {:.3}) with the Newton step still large, so the partial likelihood has no finite maximiser (every subject with a larger value of covariate {j} fails before every subject with a smaller one, or the reverse; if the covariate is merely on a tiny scale, rescale it)",
DIVERGENCE_BOUND, beta[j]
),
));
}
if step_small && delta_ll <= opts.tol * current.ll.abs().max(1.0) {
converged = true;
break;
}
}
if !converged {
return Err(failed(
op,
format!(
"no convergence in {} Newton steps (last |Δβ| criterion not met); raise max_iter or rescale the covariates",
opts.max_iter
),
));
}
let wald = wald_summary(¤t.info, &beta).ok_or_else(|| {
failed(
op,
"the information matrix at the estimate is singular: the standard errors are undefined",
)
})?;
let wald_statistic = dot(
&beta,
&dense_f64::matvec(&dense_f64::flatten(¤t.info), p, p, &beta),
);
Ok(CoxModel {
coefficients: beta,
standard_errors: wald.se,
z_values: wald.z,
p_values: wald.p,
log_likelihood: current.ll,
null_log_likelihood,
cov_params: wald.cov,
nobs: obs.len(),
n_events: obs.iter().filter(|o| o.event).count(),
ties,
iterations,
converged,
obs: obs.to_vec(),
x: x.to_vec(),
strata: strata.to_vec(),
table,
score_statistic,
wald_statistic,
})
}
impl LikelihoodFit for CoxModel {
fn log_likelihood(&self) -> f64 {
self.log_likelihood
}
fn null_log_likelihood(&self) -> f64 {
self.null_log_likelihood
}
fn n_params(&self) -> usize {
self.coefficients.len()
}
fn nobs(&self) -> usize {
self.nobs
}
fn df_model(&self) -> usize {
self.coefficients.len()
}
fn bic(&self) -> f64 {
-2.0 * self.log_likelihood + self.coefficients.len() as f64 * (self.n_events as f64).ln()
}
fn llr_test(&self, ctx: &Context) -> Result<TestResult, SymplexError> {
self.chi_squared_test(ctx, LikelihoodFit::llr(self))
}
}
impl WaldFit for CoxModel {
fn coefficients(&self) -> &[f64] {
&self.coefficients
}
fn standard_errors(&self) -> &[f64] {
&self.standard_errors
}
}
impl CoxModel {
#[must_use]
pub fn n_params(&self) -> usize {
<Self as LikelihoodFit>::n_params(self)
}
#[must_use]
pub fn strata(&self) -> &[usize] {
&self.strata
}
#[must_use]
pub fn hazard_ratios(&self) -> Vec<f64> {
self.coefficients.iter().map(|b| b.exp()).collect()
}
pub fn conf_int(&self, confidence: f64) -> Result<Vec<Interval<f64>>, SymplexError> {
<Self as WaldFit>::conf_int(self, confidence)
}
pub fn hazard_ratio_conf_int(
&self,
confidence: f64,
) -> Result<Vec<Interval<f64>>, SymplexError> {
Ok(self
.conf_int(confidence)?
.into_iter()
.map(|iv| Interval::closed(iv.lower.exp(), iv.upper.exp()))
.collect())
}
#[must_use]
pub fn llr(&self) -> f64 {
<Self as LikelihoodFit>::llr(self)
}
#[must_use]
pub fn wald_statistic(&self) -> f64 {
self.wald_statistic
}
#[must_use]
pub fn score_statistic(&self) -> f64 {
self.score_statistic
}
fn chi_squared_test(&self, ctx: &Context, statistic: f64) -> Result<TestResult, SymplexError> {
chi_squared_test_result(ctx, statistic, self.n_params(), Alternative::TwoSided)
}
pub fn llr_test(&self, ctx: &Context) -> Result<TestResult, SymplexError> {
<Self as LikelihoodFit>::llr_test(self, ctx)
}
pub fn wald_test(&self, ctx: &Context) -> Result<TestResult, SymplexError> {
self.chi_squared_test(ctx, self.wald_statistic())
}
pub fn score_test(&self, ctx: &Context) -> Result<TestResult, SymplexError> {
self.chi_squared_test(ctx, self.score_statistic)
}
#[must_use]
pub fn aic(&self) -> f64 {
<Self as LikelihoodFit>::aic(self)
}
#[must_use]
pub fn bic(&self) -> f64 {
<Self as LikelihoodFit>::bic(self)
}
fn check_row(&self, op: &'static str, x_row: &[f64]) -> Result<(), SymplexError> {
let p = self.n_params();
if x_row.len() != p {
return Err(invalid(
op,
format!("x_row has {} entries, expected {p}", x_row.len()),
));
}
if let Some(j) = x_row.iter().position(|v| !v.is_finite()) {
return Err(invalid(op, format!("x_row[{j}] is not finite")));
}
Ok(())
}
pub fn predict_log_partial_hazard(&self, x_row: &[f64]) -> Result<f64, SymplexError> {
self.check_row("predict_log_partial_hazard", x_row)?;
Ok(dot(x_row, &self.coefficients))
}
pub fn predict_partial_hazard(&self, x_row: &[f64]) -> Result<f64, SymplexError> {
self.check_row("predict_partial_hazard", x_row)?;
Ok(dot(x_row, &self.coefficients).exp())
}
#[must_use]
pub fn linear_predictors(&self) -> Vec<f64> {
linear_predictor(&self.x, &self.coefficients).eta
}
pub fn concordance(&self) -> Result<Q, SymplexError> {
let eta = self.linear_predictors();
let n = self.obs.len();
let mut usable = 0usize;
let mut twice_concordant = 0usize;
for i in 0..n {
if !self.obs[i].event {
continue;
}
for j in 0..n {
if self.strata[i] != self.strata[j] || self.obs[i].time >= self.obs[j].time {
continue;
}
usable += 1;
if eta[i] > eta[j] {
twice_concordant += 2;
} else if eta[i] == eta[j] {
twice_concordant += 1;
}
}
}
if usable == 0 {
return Err(failed(
"concordance",
"no usable pairs: every event occurs at the last observed time of its stratum",
));
}
Ok(qu(twice_concordant) / qu(2 * usable))
}
#[must_use]
pub fn baseline_hazard(&self) -> Vec<BaselineHazardRow> {
let p = self.n_params();
let lp = linear_predictor(&self.x, &self.coefficients);
let unshift = (-lp.shift).exp();
let sums = sweep(&self.table, &self.x, &lp.w, p);
let mut rows: Vec<BaselineHazardRow> = Vec::with_capacity(self.table.len());
let mut start = 0;
while start < self.table.len() {
let stratum = self.table[start].stratum;
let end = (start..self.table.len())
.find(|&k| self.table[k].stratum != stratum)
.unwrap_or(self.table.len());
let mut cumulative = 0.0;
for k in (start..end).rev() {
let hazard = self.table[k].events.len() as f64 * unshift / sums[k].s0;
cumulative += hazard;
rows.push(BaselineHazardRow {
stratum,
time: self.table[k].time.clone(),
hazard,
cumulative,
});
}
start = end;
}
rows
}
#[must_use]
pub fn schoenfeld_residuals(&self) -> Vec<Vec<f64>> {
let p = self.n_params();
let lp = linear_predictor(&self.x, &self.coefficients);
let sums = sweep(&self.table, &self.x, &lp.w, p);
let mut entry_of = vec![usize::MAX; self.obs.len()];
for (k, entry) in self.table.iter().enumerate() {
for &i in &entry.events {
entry_of[i] = k;
}
}
(0..self.obs.len())
.filter(|&i| self.obs[i].event)
.map(|i| {
let s = &sums[entry_of[i]];
(0..p).map(|a| self.x[i][a] - s.s1[a] / s.s0).collect()
})
.collect()
}
#[must_use]
pub fn martingale_residuals(&self) -> Vec<f64> {
let p = self.n_params();
let lp = linear_predictor(&self.x, &self.coefficients);
let w = &lp.w;
let sums = sweep(&self.table, &self.x, w, p);
let mut resid: Vec<f64> = self
.obs
.iter()
.map(|o| if o.event { 1.0 } else { 0.0 })
.collect();
for (k, entry) in self.table.iter().enumerate() {
let increment = entry.events.len() as f64 / sums[k].s0;
for i in 0..self.obs.len() {
if self.strata[i] == entry.stratum && self.obs[i].time >= entry.time {
resid[i] -= w[i] * increment;
}
}
}
resid
}
}
#[cfg(test)]
mod tests {
use super::*;
fn d1() -> (Vec<Observation>, Vec<Vec<f64>>) {
let obs = Observation::from_i64(
&[4, 7, 2, 9, 12, 5, 15, 3, 11, 8],
&[
true, true, true, false, true, true, false, true, true, false,
],
);
let x = [3.0, 1.0, 5.0, 2.0, 0.0, 4.0, 1.0, 6.0, 2.0, 3.0]
.iter()
.map(|&v| vec![v])
.collect();
(obs, x)
}
#[test]
fn table_risk_sets_include_censored_at_event_time() {
let obs = Observation::from_i64(&[2, 2, 3, 5], &[true, false, true, false]);
let table = build_table(&obs, &[0; 4]);
assert_eq!(table.len(), 2);
assert_eq!(table[0].time, qu(3));
assert_eq!(table[0].enter, vec![2, 3]);
assert_eq!(table[0].events, vec![2]);
assert_eq!(table[1].time, qu(2));
assert_eq!(table[1].enter, vec![0, 1]);
assert_eq!(table[1].events, vec![0]);
}
#[test]
fn efron_and_breslow_agree_without_ties() {
let (obs, x) = d1();
let table = build_table(&obs, &[0; 10]);
let b = [0.3];
let e = evaluate(&table, &x, &b, Ties::Efron);
let br = evaluate(&table, &x, &b, Ties::Breslow);
assert!((e.ll - br.ll).abs() < 1e-12);
assert!((e.score[0] - br.score[0]).abs() < 1e-12);
assert!((e.info[0][0] - br.info[0][0]).abs() < 1e-12);
}
#[test]
fn score_is_the_derivative_of_the_log_likelihood() {
let (obs, x) = d1();
let table = build_table(&obs, &[0; 10]);
let h = 1e-6;
for ties in [Ties::Efron, Ties::Breslow] {
let at = |b: f64| evaluate(&table, &x, &[b], ties);
let e = at(0.4);
let numeric = (at(0.4 + h).ll - at(0.4 - h).ll) / (2.0 * h);
assert!((e.score[0] - numeric).abs() < 1e-6, "{ties:?}");
let numeric_info = -(at(0.4 + h).score[0] - at(0.4 - h).score[0]) / (2.0 * h);
assert!((e.info[0][0] - numeric_info).abs() < 1e-5, "{ties:?}");
}
}
}