Skip to content

Cross-validate and tune a model

Three functions cover the workflow, each answering a different question:

FunctionQuestion it answers
cross_validateHow well does this model, with these settings, do on unseen patients?
tuneWhich settings work best?
nested_cvHow well does the tuned model do on unseen patients?

All three take the same two arguments: a build function that returns a fresh, unfitted model, and the data. They use stratified folds, so every fold gets its share of events, and the IPCW scorers estimate the censoring distribution on the training fold only. tune and nested_cv need Optuna: pip install 'tausurv[tune]'.

import numpy as np
import polars as pl
import tausurv as ts
from tausurv.model_selection import cross_validate, nested_cv, scoring, tune

We use the PBC cohort with the five covariates from Tutorial 2, times in years.

pbc = ts.datasets.load_pbc()
df = (
pbc.X.with_columns(
sex_male=(pl.col("sex") == "m").cast(pl.Int8),
event_time=pl.Series(pbc.event_time / 365.25),
event_indicator=pl.Series(pbc.event_indicator),
)
.select(
"age", "stage", "bili", "albumin", "sex_male", "event_time", "event_indicator"
)
.drop_nulls()
)
X = df.drop("event_time", "event_indicator").to_numpy().astype(np.float64)
Y = df["event_time"].to_numpy()
E = df["event_indicator"].to_numpy()
X.shape
(412, 5)

cross_validate fits a fresh model on each training fold and scores it on the held-out fold. Pass several scorers as a dict to get one column each. Here we use Uno’s C-index up to five years (higher is better) and the integrated Brier score over years 1 to 5 (lower is better).

horizons = np.linspace(1.0, 5.0, 20)
scorers = {
"uno_c": scoring.uno(horizon=5.0),
"ibs": scoring.integrated_brier(horizons),
}
def build():
return ts.trees.RandomSurvivalForest(n_estimators=100, seed=0)
cv = cross_validate(build, X, Y, E, scoring=scorers, cv=5, progress=False)
cv.scores
shape: (5, 4)
┌──────┬──────────┬──────────┬──────────┐
│ fold ┆ uno_c    ┆ ibs      ┆ seconds  │
│ ---  ┆ ---      ┆ ---      ┆ ---      │
│ i64  ┆ f64      ┆ f64      ┆ f64      │
╞══════╪══════════╪══════════╪══════════╡
│ 0    ┆ 0.878439 ┆ 0.088935 ┆ 0.020828 │
│ 1    ┆ 0.803553 ┆ 0.104775 ┆ 0.021235 │
│ 2    ┆ 0.828538 ┆ 0.112812 ┆ 0.020317 │
│ 3    ┆ 0.841373 ┆ 0.084965 ┆ 0.020432 │
│ 4    ┆ 0.887258 ┆ 0.088359 ┆ 0.022219 │
└──────┴──────────┴──────────┴──────────┘
cv.scores.select(pl.col("uno_c", "ibs").mean())
shape: (1, 2)
┌──────────┬──────────┐
│ uno_c    ┆ ibs      │
│ ---      ┆ ---      │
│ f64      ┆ f64      │
╞══════════╪══════════╡
│ 0.847832 ┆ 0.095969 │
└──────────┴──────────┘

The result also works as a predictor. On the study’s own data it predicts out of fold: each patient is predicted by the model that did not see them. For new patients, cv.ensemble averages the fold models.

oof_risk = cv.predict(X)
new_risk = cv.ensemble.predict(X[:3])

To tune, give build a trial argument and draw each hyperparameter where it is used, with Optuna’s trial.suggest_*. tune runs a Bayesian search in which each trial is scored by cross-validation, and by default refits the best settings on all the data.

def build_tuned(trial):
return ts.trees.RandomSurvivalForest(
n_estimators=100,
min_samples_leaf=trial.suggest_int("min_samples_leaf", 5, 50, log=True),
max_depth=trial.suggest_int("max_depth", 2, 12),
seed=0,
)
best = tune(
build_tuned,
X,
Y,
E,
scoring=scoring.uno(horizon=5.0),
cv=3,
n_trials=15,
progress=False,
)
best.params, round(best.score, 3)
({‘min_samples_leaf’: 5, ‘max_depth’: 11}, 0.85)

best.model is the refitted model, and best.trials has every trial with its parameters and score.

best.trials.sort("value", descending=True).head(5)
shape: (5, 5)
┌────────┬──────────┬──────────┬──────────────────┬───────────┐
│ number ┆ value    ┆ state    ┆ min_samples_leaf ┆ max_depth │
│ ---    ┆ ---      ┆ ---      ┆ ---              ┆ ---       │
│ i64    ┆ f64      ┆ str      ┆ i64              ┆ i64       │
╞════════╪══════════╪══════════╪══════════════════╪═══════════╡
│ 8      ┆ 0.84953  ┆ COMPLETE ┆ 5                ┆ 11        │
│ 11     ┆ 0.84953  ┆ COMPLETE ┆ 5                ┆ 11        │
│ 12     ┆ 0.84953  ┆ COMPLETE ┆ 5                ┆ 11        │
│ 10     ┆ 0.84856  ┆ COMPLETE ┆ 6                ┆ 8         │
│ 13     ┆ 0.845648 ┆ COMPLETE ┆ 8                ┆ 12        │
└────────┴──────────┴──────────┴──────────────────┴───────────┘

best.score is optimistic: the search picked the settings that happened to score highest on these folds, so it has partly fitted the folds’ noise. To estimate how well the whole procedure, search included, generalises, nested_cv reruns the search inside every outer training fold and scores the winner on the outer test fold, which the search never saw.

ncv = nested_cv(
build_tuned,
X,
Y,
E,
scoring=scoring.uno(horizon=5.0),
outer=5,
inner=3,
n_trials=15,
progress=False,
)
ncv.scores
shape: (5, 3)
┌──────┬──────────┬──────────┐
│ fold ┆ uno_c    ┆ seconds  │
│ ---  ┆ ---      ┆ ---      │
│ i64  ┆ f64      ┆ f64      │
╞══════╪══════════╪══════════╡
│ 0    ┆ 0.876735 ┆ 0.721894 │
│ 1    ┆ 0.806448 ┆ 0.737986 │
│ 2    ┆ 0.826109 ┆ 0.716834 │
│ 3    ┆ 0.83961  ┆ 0.707611 │
│ 4    ┆ 0.879594 ┆ 0.727918 │
└──────┴──────────┴──────────┘
ncv.scores["uno_c"].mean()
0.8456994261401874

Report this mean, not best.score, as the tuned model’s performance. ncv.params shows the settings each outer fold chose; when they vary a lot, the score is flat across that range and the exact value matters little.

pl.DataFrame(ncv.params)
shape: (5, 2)
┌──────────────────┬───────────┐
│ min_samples_leaf ┆ max_depth │
│ ---              ┆ ---       │
│ i64              ┆ i64       │
╞══════════════════╪═══════════╡
│ 5                ┆ 11        │
│ 5                ┆ 11        │
│ 17               ┆ 9         │
│ 18               ┆ 12        │
│ 8                ┆ 4         │
└──────────────────┴───────────┘
  • Comparing fixed models: use cross_validate.
  • Fitting a final model: use tune, and ship best.model.
  • Reporting the tuned model’s performance: use nested_cv, and report its outer scores.

A nested study is outer * (n_trials * inner + 1) fits, 230 here, so keep n_trials small while you iterate. Pass storage="sqlite:///study.db" to tune to make a long search resumable, and use .save(path) on any result to keep the fitted models.