Clustering Individualized Survival Curves with unsurv

knitr::opts_chunk$set(
  collapse = TRUE,
  comment = "#>",
  fig.width = 7,
  fig.height = 4.5
)
has_survival <- requireNamespace("survival", quietly = TRUE)

unsurv clusters full predicted survival trajectories rather than baseline covariates or a single risk score. Each observation is one survival probability curve on a shared time grid, and the package groups patients whose whole predicted trajectory has the same shape.

This vignette has two parts:

  1. a minimal simulated example that shows the mechanics of unsurv(), plot(), predict(), and unsurv_stability();
  2. the worked example from the accompanying paper --- clustering deep-learning survival predictions for the METABRIC breast-cancer cohort --- run from a small set of predicted curves shipped with the package.

1. Mechanics on a simulated example

Three prognosis groups with different exponential hazards, plus a little noise so the curves are not perfectly smooth.

library(unsurv)

set.seed(2026)
n <- 150
Q <- 60
sim_times <- seq(0, 5, length.out = Q)

group <- sample(1:3, n, TRUE, prob = c(0.35, 0.4, 0.25))
haz   <- c(0.18, 0.45, 0.8)[group]

S <- sapply(sim_times, function(t) exp(-haz * t))
S <- S + matrix(rnorm(n * Q, 0, 0.02), nrow = n)
S[S < 0] <- 0
S[S > 1] <- 1

Leaving K = NULL lets unsurv choose the number of clusters by the mean silhouette width over 2:K_max.

sim_fit <- unsurv(S, sim_times, K = NULL, K_max = 6, distance = "L2",
                  enforce_monotone = TRUE, smooth_median_width = 5,
                  standardize_cols = TRUE, eps_jitter = 0.0005, seed = 1)
sim_fit

Key slots: K (chosen cluster count), clusters (assignment per curve), medoids (representative curves), silhouette_mean.

plot(sim_fit)

New curves must be on the same time grid; the fitted preprocessing (clamping, monotonicity, smoothing, standardization) is reused automatically.

predict(sim_fit, S[1:5, ])

2. Worked example: METABRIC predicted curves (from the paper)

Where the data come from

The file extdata/metabric_survdnn_curves.rds reproduces the worked example in

El Badisy I (2026). "unsurv: clustering individualized survival curves." Bioinformatics Advances, 6(1), vbag218.

It was produced by data-raw/metabric_survdnn_curves.R with the same pipeline as the paper's scripts:

The shipped object contains only model-predicted probabilities and the (time, status) pair for each patient --- no METABRIC covariates or patient identifiers --- so the vignette reproduces the paper figure without importing survdnn, biostatlab, or torch at build time. Regenerate it by running data-raw/metabric_survdnn_curves.R from the package root.

f <- system.file("extdata", "metabric_survdnn_curves.rds", package = "unsurv")
mb <- readRDS(f)

str(mb, max.level = 1)
cat(mb$provenance)

mb$S_partition is a 367 x 100 matrix of predicted survival probabilities (one row per patient), mb$times the shared time grid in months, and mb$os_time / mb$os_event the observed outcomes for each set.

Cluster the predicted curves

Following the paper, we fix K = 3 and cluster the partition-set curves.

fit <- unsurv(mb$S_partition, mb$times, K = 3, distance = "L2",
              enforce_monotone = TRUE, smooth_median_width = 3,
              standardize_cols = FALSE, eps_jitter = 0, seed = 20260615)
fit
library(ggplot2)

grid_n  <- length(mb$times)
curf <- data.frame(
  id       = rep(seq_len(nrow(mb$S_partition)), each = grid_n),
  time     = rep(mb$times, times = nrow(mb$S_partition)),
  survival = as.vector(t(mb$S_partition)),
  cluster  = factor(rep(fit$clusters, each = grid_n))
)
medf <- data.frame(
  time     = rep(mb$times, times = nrow(fit$medoids)),
  survival = as.vector(t(fit$medoids)),
  cluster  = factor(rep(seq_len(nrow(fit$medoids)), each = grid_n))
)

ggplot(curf, aes(time, survival, group = id, colour = cluster)) +
  geom_line(linewidth = 0.25, alpha = 0.20) +
  geom_line(data = medf, aes(group = cluster), linewidth = 1.1) +
  labs(x = "Months", y = "Predicted survival probability", colour = "Cluster") +
  theme_minimal(base_size = 11)

Out-of-sample assignment

The clusters are defined on the partition set; the validation patients are assigned out-of-sample by nearest medoid.

val_clusters <- predict(fit, mb$S_validation)
table(validation = val_clusters)

Compare against a scalar-risk baseline

A common alternative is to cluster a single risk summary. Here we take the predicted risk 1 - S(t) at the median grid time and run PAM on it, then assign the validation set by nearest cluster mean.

h <- which.min(abs(mb$times - stats::median(mb$times)))
risk_p <- 1 - mb$S_partition[, h]
risk_v <- 1 - mb$S_validation[, h]

pam_scalar <- cluster::pam(as.matrix(risk_p), k = 3)
centres <- tapply(risk_p, pam_scalar$clustering, mean)
scalar_v <- vapply(risk_v, function(v) which.min(abs(v - centres)), integer(1))

unsurv_compare() summarizes each partition against the observed outcomes --- per-cluster Kaplan-Meier medians --- and reports the Adjusted Rand Index of each partition against a reference. (A log-rank test is deliberately not reported: the partitions are fit to separate the curves, so a log-rank test against those same labels is circular.)

cmp_part <- unsurv_compare(
  list(`unsurv curve` = fit$clusters, `scalar risk` = pam_scalar$clustering),
  mb$os_time$partition, mb$os_event$partition,
  reference = "unsurv curve"
)
cmp_part$summary
cmp_part$cluster_summary

Repeating on the held-out validation set checks that the survival ordering of the clusters generalizes:

cmp_val <- unsurv_compare(
  list(`unsurv curve` = val_clusters, `scalar risk` = scalar_v),
  mb$os_time$validation, mb$os_event$validation,
  reference = "unsurv curve"
)
cmp_val$cluster_summary
autoplot(unsurv_compare(
  list(`unsurv curve` = val_clusters),
  mb$os_time$validation, mb$os_event$validation
))

Stability

Resampling the partition set gives a sense of how reproducible the clustering is under perturbation.

stab <- unsurv_stability(
  mb$S_partition, mb$times, fit,
  B = 30, frac = 0.7, mode = "subsample",
  jitter_sd = 0.005, weight_perturb = 0.1, eps_jitter = 0,
  return_distribution = TRUE
)
stab$mean

Higher mean ARI indicates more reproducible clusters.

Tips and troubleshooting



Try the unsurv package in your browser

Any scripts or data that you put into this service are public.

unsurv documentation built on Sept. 1, 2026, 1:06 a.m.