Files
DeepHealth/evaluate_calibration.py

1981 lines
66 KiB
Python
Raw Normal View History

2026-07-30 17:21:57 +08:00
"""Evaluate calibration, Brier score, and NLL for all-future runs.
This evaluator intentionally does not support the Delphi2M/next-token branch.
Outputs written to each run directory:
* ``df_calibration_landmark_metrics.csv``
Per disease, sex, landmark age, and horizon IPCW metrics.
* ``df_calibration_metrics.csv``
Per disease, sex, and horizon aggregation across landmark ages.
* ``df_calibration_summary.csv``
Disease/death summaries by sex and horizon.
* ``df_calibration_curve.csv``
Fixed probability-bin calibration curves.
* ``df_point_process_nll.csv``
Exact all-future test-query NLL using the training likelihood.
* ``calibration_evaluation_summary.json``
Completion marker and evaluation metadata. It is written last.
"""
from __future__ import annotations
import argparse
import contextlib
import json
import os
from collections import defaultdict
2026-07-30 18:29:15 +08:00
from concurrent.futures import ThreadPoolExecutor, as_completed
2026-07-30 17:21:57 +08:00
from pathlib import Path
from typing import Any, Dict, List, Optional, Sequence, Tuple
2026-07-30 17:21:57 +08:00
import numpy as np
import pandas as pd
import torch
from torch.utils.data import DataLoader, Subset
from tqdm.auto import tqdm
from dataset import (
DISEASE_HISTORY_MODE_TIMED,
AllFutureHealthDataset,
all_future_collate_fn,
normalize_disease_history_mode,
)
from eval_data import (
build_first_occurrence_map,
build_model_from_dataset,
cfg_get,
load_json_config,
load_sequence_eval_dataset,
resolve_eval_device,
validate_dataset_metadata,
validate_training_mode_config,
)
from evaluate_auc_v2 import (
LandmarkDataset,
_get_death_token_ids,
collate_landmark_fn,
infer_landmark_hidden,
load_checkpoint_state_dict,
load_model_state,
parse_float_list,
parse_int_list,
project_distribution_chunk,
resolve_dist_mode_for_checkpoint,
select_disease_tokens,
)
from losses import build_loss
from model_architectures import resolve_model_architecture
from targets import NO_EVENT_IDX, PAD_IDX, RESERVED_IDX
2026-07-30 17:21:57 +08:00
from train_util import load_eid_file
PROJECT_ROOT = Path(__file__).resolve().parent
COMPLETION_FILE = "calibration_evaluation_summary.json"
LANDMARK_METRICS_FILE = "df_calibration_landmark_metrics.csv"
TOKEN_METRICS_FILE = "df_calibration_metrics.csv"
SUMMARY_FILE = "df_calibration_summary.csv"
CALIBRATION_CURVE_FILE = "df_calibration_curve.csv"
POINT_PROCESS_NLL_FILE = "df_point_process_nll.csv"
DEFAULT_HORIZONS = [0.1, 1.0, 5.0, 10.0]
DEFAULT_PROBABILITY_BINS = [
0.0,
0.001,
0.002,
0.005,
0.01,
0.02,
0.05,
0.1,
0.2,
0.5,
1.0,
]
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(
description=(
"Evaluate all-future landmark calibration, IPCW Brier score, "
"fixed-horizon NLL, and point-process NLL."
)
)
parser.add_argument("--run_path", required=True)
parser.add_argument("--output_path", default=None)
parser.add_argument(
"--eval_split",
default="test",
choices=["val", "valid", "validation", "test"],
)
parser.add_argument("--dataset_subset_size", type=int, default=None)
parser.add_argument("--batch_size", type=int, default=None)
parser.add_argument("--num_workers", type=int, default=None)
parser.add_argument("--device", default=None)
parser.add_argument(
"--use_amp",
action=argparse.BooleanOptionalAction,
default=None,
)
parser.add_argument(
"--hidden_cache_dtype",
choices=["float16", "float32"],
default=None,
)
parser.add_argument("--logit_batch_size", type=int, default=None)
parser.add_argument("--disease_chunk_size", type=int, default=None)
2026-07-30 18:29:15 +08:00
parser.add_argument(
"--num_workers_calibration",
type=int,
default=None,
help=(
"CPU threads for per-disease calibration statistics. "
"0 uses all logical CPUs."
),
)
2026-07-30 17:21:57 +08:00
parser.add_argument("--filter_min_total", type=int, default=None)
parser.add_argument("--diseases_of_interest", default=None)
parser.add_argument("--labels_meta_path", default=None)
parser.add_argument("--landmark_start", type=float, default=None)
parser.add_argument("--landmark_stop", type=float, default=None)
parser.add_argument("--landmark_step", type=float, default=None)
parser.add_argument(
"--horizons",
default=None,
help="Comma-separated fixed horizons in years. Default: 0.1,1,5,10.",
)
parser.add_argument("--min_history_events", type=int, default=None)
parser.add_argument(
"--min_cases",
type=int,
default=None,
help="Minimum observed events in a landmark cell. Default: 2.",
)
parser.add_argument(
"--min_controls",
type=int,
default=1,
help="Minimum known event-free controls in a landmark cell.",
)
parser.add_argument(
"--probability_bins",
default=None,
help=(
"Comma-separated fixed probability-bin edges. "
"Default: 0,.001,.002,.005,.01,.02,.05,.1,.2,.5,1."
),
)
parser.add_argument(
"--max_ipcw_weight",
type=float,
default=0.0,
help="Optional IPCW cap; 0 means no cap.",
)
parser.add_argument(
"--exclude_death_in_window_without_disease",
action=argparse.BooleanOptionalAction,
default=None,
help=(
"For non-death outcomes, treat death before disease as censoring. "
"Default: true."
),
)
parser.add_argument(
"--point_process_nll",
action=argparse.BooleanOptionalAction,
default=True,
help="Also evaluate the exact all-future training likelihood.",
)
parser.add_argument(
"--force",
action="store_true",
help="Recompute even when the completion marker already exists.",
)
return parser
def _normalise_eval_split(value: str) -> str:
split = str(value).lower()
if split in {"valid", "validation"}:
return "val"
if split not in {"val", "test"}:
raise ValueError(f"eval_split must be val or test, got {value!r}")
return split
def _resolve_project_path(value: str | Path) -> Path:
path = Path(value)
if path.is_absolute():
return path
direct = Path.cwd() / path
if direct.exists():
return direct
return PROJECT_ROOT / path
def _configured_eid_file(cfg: Dict[str, Any], eval_split: str) -> Optional[Path]:
key = "val_eid_file" if eval_split == "val" else "test_eid_file"
raw = cfg.get(key)
if raw in {None, ""}:
return None
path = _resolve_project_path(str(raw))
if not path.is_file():
raise FileNotFoundError(
f"Configured {key} does not exist: {path}"
)
return path
def select_sequence_eval_indices(
dataset: Any,
cfg: Dict[str, Any],
eval_split: str,
subset_size: Optional[int],
) -> tuple[np.ndarray, str, Optional[str]]:
"""Select landmark patients using the same patient split as training."""
eval_split = _normalise_eval_split(eval_split)
eid_path = _configured_eid_file(cfg, eval_split)
if eid_path is not None:
selected_eids = load_eid_file(eid_path)
indices = np.asarray(
[
index
for index, sample in enumerate(dataset.samples)
if int(sample["eid"]) in selected_eids
],
dtype=np.int64,
)
method = "eid_file"
source = str(eid_path)
else:
n_patients = len(dataset)
train_ratio = float(cfg.get("train_ratio", 0.7))
val_ratio = float(cfg.get("val_ratio", 0.15))
test_ratio = float(cfg.get("test_ratio", 0.15))
total = train_ratio + val_ratio + test_ratio
if not np.isclose(total, 1.0, atol=1e-6):
raise ValueError(f"train/val/test ratios must sum to 1, got {total}")
order = np.random.RandomState(int(cfg.get("seed", 42))).permutation(
n_patients
)
n_train = int(n_patients * train_ratio)
n_val = int(n_patients * val_ratio)
indices = (
order[n_train:n_train + n_val]
if eval_split == "val"
else order[n_train + n_val:]
).astype(np.int64, copy=False)
method = "random_patient_ratio"
source = None
if subset_size is not None and int(subset_size) > 0:
indices = indices[: int(subset_size)]
if indices.size == 0:
raise RuntimeError("Selected landmark evaluation split is empty.")
return indices, method, source
def build_point_process_eval_subset(
cfg: Dict[str, Any],
eval_split: str,
disease_history_mode: str,
subset_size: Optional[int],
) -> tuple[AllFutureHealthDataset, Subset, str, Optional[str]]:
dataset_split = "valid" if _normalise_eval_split(eval_split) == "val" else "test"
dataset = AllFutureHealthDataset(
data_prefix=str(cfg.get("data_prefix", "ukb")),
labels_file=str(cfg.get("labels_file", "labels.csv")),
split=dataset_split,
min_history_events=int(cfg.get("all_future_min_history_events", 1)),
min_future_events=int(cfg.get("all_future_min_future_events", 1)),
validation_query_seed=int(
cfg.get("all_future_validation_query_seed", 42)
),
extra_info_types=parse_int_list(cfg.get("extra_info_types")),
disease_history_mode=disease_history_mode,
)
validate_dataset_metadata(dataset, cfg)
eval_split = _normalise_eval_split(eval_split)
eid_path = _configured_eid_file(cfg, eval_split)
if eid_path is not None:
selected_eids = load_eid_file(eid_path)
query_indices = [
query_index
for query_index, (patient_index, _query_time) in enumerate(
dataset.valid_queries
)
if int(dataset.patients[int(patient_index)]["eid"]) in selected_eids
]
method = "eid_file"
source = str(eid_path)
else:
patient_count = len(dataset.patients)
train_ratio = float(cfg.get("train_ratio", 0.7))
val_ratio = float(cfg.get("val_ratio", 0.15))
test_ratio = float(cfg.get("test_ratio", 0.15))
total = train_ratio + val_ratio + test_ratio
if not np.isclose(total, 1.0, atol=1e-6):
raise ValueError(f"train/val/test ratios must sum to 1, got {total}")
order = np.random.RandomState(int(cfg.get("seed", 42))).permutation(
patient_count
)
n_train = int(patient_count * train_ratio)
n_val = int(patient_count * val_ratio)
patient_set = set(
int(value)
for value in (
order[n_train:n_train + n_val]
if eval_split == "val"
else order[n_train + n_val:]
)
)
query_indices = [
query_index
for query_index, (patient_index, _query_time) in enumerate(
dataset.valid_queries
)
if int(patient_index) in patient_set
]
method = "random_patient_ratio"
source = None
if subset_size is not None and int(subset_size) > 0:
query_indices = query_indices[: int(subset_size)]
if not query_indices:
raise RuntimeError("Selected point-process evaluation split is empty.")
subset = Subset(dataset, np.asarray(query_indices, dtype=np.int64))
return dataset, subset, method, source
def _sigmoid(values: np.ndarray) -> np.ndarray:
values = np.asarray(values, dtype=np.float64)
result = np.empty_like(values, dtype=np.float64)
nonnegative = values >= 0
result[nonnegative] = 1.0 / (
1.0 + np.exp(-values[nonnegative])
)
exp_values = np.exp(values[~nonnegative])
result[~nonnegative] = exp_values / (1.0 + exp_values)
return result
def fit_weighted_logistic_calibration(
probabilities: np.ndarray,
outcomes: np.ndarray,
weights: np.ndarray,
eps: float = 1e-6,
max_iter: int = 100,
) -> tuple[float, float, float]:
"""Return calibration-in-the-large, free intercept, and free slope."""
2026-07-30 18:29:15 +08:00
calibration_in_large, intercept, slope = (
fit_weighted_logistic_calibration_batch(
probabilities=np.asarray(probabilities, dtype=np.float64)[None, :],
outcomes=np.asarray(outcomes, dtype=np.float64)[None, :],
weights=np.asarray(weights, dtype=np.float64)[None, :],
eps=eps,
max_iter=max_iter,
)
)
return (
float(calibration_in_large[0]),
float(intercept[0]),
float(slope[0]),
)
def fit_weighted_logistic_calibration_batch(
probabilities: np.ndarray,
outcomes: np.ndarray,
weights: np.ndarray,
eps: float = 1e-6,
max_iter: int = 100,
) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
"""Fit all horizon-specific calibration models in one vectorized pass."""
2026-07-30 17:21:57 +08:00
probabilities = np.asarray(probabilities, dtype=np.float64)
outcomes = np.asarray(outcomes, dtype=np.float64)
weights = np.asarray(weights, dtype=np.float64)
2026-07-30 18:29:15 +08:00
if probabilities.ndim != 2:
raise ValueError("probabilities must have shape [horizons, samples]")
if not (
probabilities.shape == outcomes.shape == weights.shape
):
raise ValueError("probabilities, outcomes, and weights must align")
2026-07-30 17:21:57 +08:00
valid = (
np.isfinite(probabilities)
& np.isfinite(outcomes)
& np.isfinite(weights)
& (weights > 0)
)
2026-07-30 18:29:15 +08:00
p = np.where(
valid,
np.clip(probabilities, eps, 1.0 - eps),
0.5,
)
y = np.where(valid, outcomes, 0.0)
w = np.where(valid, weights, 0.0)
2026-07-30 17:21:57 +08:00
logit_p = np.log(p) - np.log1p(-p)
2026-07-30 18:29:15 +08:00
valid_count = valid.sum(axis=1)
event_count = (valid & (y > 0.5)).sum(axis=1)
control_count = (valid & (y <= 0.5)).sum(axis=1)
probability_min = np.min(
np.where(valid, p, np.inf),
axis=1,
)
probability_max = np.max(
np.where(valid, p, -np.inf),
axis=1,
)
eligible = (
(valid_count >= 3)
& (event_count > 0)
& (control_count > 0)
& ((probability_max - probability_min) > 1e-12)
)
horizon_count = int(probabilities.shape[0])
intercept_only = np.zeros(horizon_count, dtype=np.float64)
intercept_active = eligible.copy()
2026-07-30 17:21:57 +08:00
for _ in range(max_iter):
2026-07-30 18:29:15 +08:00
if not np.any(intercept_active):
2026-07-30 17:21:57 +08:00
break
2026-07-30 18:29:15 +08:00
fitted = _sigmoid(logit_p + intercept_only[:, None])
gradient = np.sum(w * (y - fitted), axis=1)
information = np.sum(
w * fitted * (1.0 - fitted),
axis=1,
)
solvable = (
intercept_active
& np.isfinite(gradient)
& np.isfinite(information)
& (information > 1e-12)
)
failed = intercept_active & ~solvable
intercept_only[failed] = np.nan
intercept_active[failed] = False
step = np.zeros(horizon_count, dtype=np.float64)
step[solvable] = gradient[solvable] / information[solvable]
intercept_only[solvable] += step[solvable]
diverged = solvable & (
~np.isfinite(intercept_only)
| (np.abs(intercept_only) > 1e6)
)
intercept_only[diverged] = np.nan
intercept_active[diverged] = False
intercept_active[solvable & (np.abs(step) < 1e-9)] = False
beta_intercept = np.zeros(horizon_count, dtype=np.float64)
beta_slope = np.ones(horizon_count, dtype=np.float64)
beta_active = eligible.copy()
ridge = 1e-9
2026-07-30 17:21:57 +08:00
for _ in range(max_iter):
2026-07-30 18:29:15 +08:00
if not np.any(beta_active):
2026-07-30 17:21:57 +08:00
break
2026-07-30 18:29:15 +08:00
fitted = _sigmoid(
beta_intercept[:, None] + beta_slope[:, None] * logit_p
)
variance = np.clip(fitted * (1.0 - fitted), 1e-12, None)
residual = w * (y - fitted)
weighted_variance = w * variance
gradient_0 = np.sum(residual, axis=1)
gradient_1 = np.sum(residual * logit_p, axis=1)
information_00 = np.sum(weighted_variance, axis=1) + ridge
information_01 = np.sum(
weighted_variance * logit_p,
axis=1,
)
information_11 = np.sum(
weighted_variance * np.square(logit_p),
axis=1,
) + ridge
determinant = (
information_00 * information_11
- np.square(information_01)
)
solvable = (
beta_active
& np.isfinite(gradient_0)
& np.isfinite(gradient_1)
& np.isfinite(determinant)
& (determinant > 1e-18)
)
failed = beta_active & ~solvable
beta_intercept[failed] = np.nan
beta_slope[failed] = np.nan
beta_active[failed] = False
step_0 = np.zeros(horizon_count, dtype=np.float64)
step_1 = np.zeros(horizon_count, dtype=np.float64)
step_0[solvable] = (
gradient_0[solvable] * information_11[solvable]
- gradient_1[solvable] * information_01[solvable]
) / determinant[solvable]
step_1[solvable] = (
information_00[solvable] * gradient_1[solvable]
- information_01[solvable] * gradient_0[solvable]
) / determinant[solvable]
beta_intercept[solvable] += step_0[solvable]
beta_slope[solvable] += step_1[solvable]
diverged = solvable & (
~np.isfinite(beta_intercept)
| ~np.isfinite(beta_slope)
| (np.abs(beta_intercept) > 1e6)
| (np.abs(beta_slope) > 1e6)
)
beta_intercept[diverged] = np.nan
beta_slope[diverged] = np.nan
beta_active[diverged] = False
converged = solvable & (
np.maximum(np.abs(step_0), np.abs(step_1)) < 1e-9
)
beta_active[converged] = False
intercept_only[~eligible] = np.nan
beta_intercept[~eligible] = np.nan
beta_slope[~eligible] = np.nan
return intercept_only, beta_intercept, beta_slope
2026-07-30 17:21:57 +08:00
def _censoring_km(
observed_times: np.ndarray,
censor_events: np.ndarray,
) -> tuple[np.ndarray, np.ndarray]:
2026-07-30 18:29:15 +08:00
"""Estimate censoring survival in O(N log N) time."""
2026-07-30 17:21:57 +08:00
observed_times = np.asarray(observed_times, dtype=np.float64)
censor_events = np.asarray(censor_events, dtype=bool)
if observed_times.size == 0:
return np.empty(0, dtype=np.float64), np.empty(0, dtype=np.float64)
2026-07-30 18:29:15 +08:00
if observed_times.shape != censor_events.shape:
raise ValueError("observed_times and censor_events must align")
order = np.argsort(observed_times, kind="stable")
sorted_times = observed_times[order]
sorted_censor_events = censor_events[order].astype(
np.int64,
copy=False,
)
unique_times, first_indices = np.unique(
sorted_times,
return_index=True,
)
censor_counts = np.add.reduceat(
sorted_censor_events,
first_indices,
)
censor_mask = censor_counts > 0
if not np.any(censor_mask):
return (
np.empty(0, dtype=np.float64),
np.empty(0, dtype=np.float64),
2026-07-30 17:21:57 +08:00
)
2026-07-30 18:29:15 +08:00
event_times = unique_times[censor_mask]
at_risk = (
int(observed_times.size) - first_indices[censor_mask]
).astype(np.float64, copy=False)
survival_factors = (
1.0
- censor_counts[censor_mask].astype(np.float64, copy=False)
/ at_risk
)
survival_after = np.cumprod(survival_factors, dtype=np.float64)
2026-07-30 17:21:57 +08:00
return event_times, survival_after
def _km_value(
event_times: np.ndarray,
survival_after: np.ndarray,
query_times: np.ndarray | float,
*,
before: bool,
) -> np.ndarray:
values = np.asarray(query_times, dtype=np.float64)
if event_times.size == 0:
return np.ones_like(values, dtype=np.float64)
side = "left" if before else "right"
positions = np.searchsorted(event_times, values, side=side) - 1
result = np.ones_like(values, dtype=np.float64)
valid = positions >= 0
result[valid] = survival_after[positions[valid]]
return result
def _cap_ipcw(weights: np.ndarray, maximum: float) -> tuple[np.ndarray, int]:
weights = np.asarray(weights, dtype=np.float64)
if maximum <= 0:
return weights, 0
clipped = weights > maximum
return np.minimum(weights, maximum), int(clipped.sum())
def compute_ipcw_cell(
probabilities: np.ndarray,
event_times: np.ndarray,
censor_times: np.ndarray,
horizon: float,
min_cases: int,
min_controls: int,
max_ipcw_weight: float,
) -> Optional[tuple[Dict[str, Any], Dict[str, np.ndarray]]]:
"""Compute IPCW fixed-horizon metrics for one at-risk landmark cell."""
2026-07-30 18:29:15 +08:00
results = compute_ipcw_horizons(
probabilities=np.asarray(probabilities, dtype=np.float64)[None, :],
event_times=event_times,
censor_times=censor_times,
horizons=np.asarray([horizon], dtype=np.float64),
min_cases=min_cases,
min_controls=min_controls,
max_ipcw_weight=max_ipcw_weight,
)
return results[0]
def compute_ipcw_horizons(
probabilities: np.ndarray,
event_times: np.ndarray,
censor_times: np.ndarray,
horizons: np.ndarray,
min_cases: int,
min_controls: int,
max_ipcw_weight: float,
) -> List[Optional[tuple[Dict[str, Any], Dict[str, np.ndarray]]]]:
"""Compute all fixed horizons with one censoring KM and batched fits."""
p = np.clip(
np.asarray(probabilities, dtype=np.float64),
1e-8,
1.0 - 1e-8,
)
2026-07-30 17:21:57 +08:00
event_times = np.asarray(event_times, dtype=np.float64)
censor_times = np.asarray(censor_times, dtype=np.float64)
2026-07-30 18:29:15 +08:00
horizons = np.asarray(horizons, dtype=np.float64)
if p.ndim != 2:
raise ValueError("probabilities must have shape [horizons, samples]")
if horizons.ndim != 1 or p.shape[0] != horizons.size:
raise ValueError("probabilities and horizons must align")
if not (
event_times.shape == censor_times.shape == (p.shape[1],)
):
raise ValueError(
"probabilities, event_times, and censor_times must align"
)
n_at_risk = int(p.shape[1])
2026-07-30 17:21:57 +08:00
if n_at_risk == 0:
2026-07-30 18:29:15 +08:00
return [None] * int(horizons.size)
2026-07-30 17:21:57 +08:00
observed_event = event_times <= censor_times
2026-07-30 18:29:15 +08:00
event_by_horizon = (
observed_event[None, :]
& (event_times[None, :] <= horizons[:, None])
2026-07-30 17:21:57 +08:00
)
2026-07-30 18:29:15 +08:00
controls = (
(event_times[None, :] > horizons[:, None])
& (censor_times[None, :] >= horizons[:, None])
)
n_events = event_by_horizon.sum(axis=1).astype(np.int64)
n_controls = controls.sum(axis=1).astype(np.int64)
eligible = (
(n_events >= int(min_cases))
& (n_controls >= int(min_controls))
)
if not np.any(eligible):
return [None] * int(horizons.size)
2026-07-30 17:21:57 +08:00
observed_times = np.minimum(event_times, censor_times)
censor_events = censor_times < event_times
km_times, km_survival = _censoring_km(observed_times, censor_events)
event_g = _km_value(
km_times,
km_survival,
2026-07-30 18:29:15 +08:00
event_times,
2026-07-30 17:21:57 +08:00
before=True,
)
control_g = _km_value(
km_times,
km_survival,
2026-07-30 18:29:15 +08:00
horizons,
2026-07-30 17:21:57 +08:00
before=False,
)
2026-07-30 18:29:15 +08:00
raw_event_weights = 1.0 / np.clip(event_g, 1e-8, None)
raw_control_weights = 1.0 / np.clip(control_g, 1e-8, None)
event_weight_values, _ = _cap_ipcw(
raw_event_weights,
2026-07-30 17:21:57 +08:00
max_ipcw_weight,
)
2026-07-30 18:29:15 +08:00
control_weight_values, _ = _cap_ipcw(
raw_control_weights,
2026-07-30 17:21:57 +08:00
max_ipcw_weight,
)
2026-07-30 18:29:15 +08:00
event_clipped = (
raw_event_weights > max_ipcw_weight
if max_ipcw_weight > 0
else np.zeros(n_at_risk, dtype=bool)
)
control_clipped = (
raw_control_weights > max_ipcw_weight
if max_ipcw_weight > 0
else np.zeros(horizons.size, dtype=bool)
)
2026-07-30 17:21:57 +08:00
2026-07-30 18:29:15 +08:00
metric_weights = np.where(
event_by_horizon,
event_weight_values[None, :],
0.0,
)
metric_weights = np.where(
controls,
control_weight_values[:, None],
metric_weights,
)
outcomes = event_by_horizon.astype(np.float64)
2026-07-30 17:21:57 +08:00
brier_contribution = metric_weights * np.square(outcomes - p)
nll_contribution = -metric_weights * (
outcomes * np.log(p) + (1.0 - outcomes) * np.log1p(-p)
)
known = event_by_horizon | controls
calibration_in_large, calibration_intercept, calibration_slope = (
2026-07-30 18:29:15 +08:00
fit_weighted_logistic_calibration_batch(
probabilities=p,
outcomes=outcomes,
weights=metric_weights,
2026-07-30 17:21:57 +08:00
)
)
2026-07-30 18:29:15 +08:00
prediction_sum = p.sum(axis=1)
event_weight_sum = (
metric_weights * outcomes
).sum(axis=1)
brier_sum = brier_contribution.sum(axis=1)
nll_sum = nll_contribution.sum(axis=1)
known_count = known.sum(axis=1)
complete_case_brier = (
np.where(known, np.square(outcomes - p), 0.0).sum(axis=1)
/ np.maximum(known_count, 1)
)
complete_case_nll = (
-np.where(
known,
outcomes * np.log(p)
+ (1.0 - outcomes) * np.log1p(-p),
0.0,
).sum(axis=1)
/ np.maximum(known_count, 1)
)
results: List[
Optional[tuple[Dict[str, Any], Dict[str, np.ndarray]]]
] = []
for horizon_index in range(int(horizons.size)):
if not bool(eligible[horizon_index]):
results.append(None)
continue
observed_rate = (
float(event_weight_sum[horizon_index]) / float(n_at_risk)
)
predicted_mean = (
float(prediction_sum[horizon_index]) / float(n_at_risk)
)
clipped_count = int(
np.sum(
event_by_horizon[horizon_index] & event_clipped
2026-07-30 17:21:57 +08:00
)
2026-07-30 18:29:15 +08:00
+ (
int(n_controls[horizon_index])
if bool(control_clipped[horizon_index])
else 0
)
)
row = {
"n_at_risk": n_at_risk,
"n_events": int(n_events[horizon_index]),
"n_controls": int(n_controls[horizon_index]),
"n_censored_before_horizon": int(
n_at_risk
- n_events[horizon_index]
- n_controls[horizon_index]
),
"known_fraction": float(
known_count[horizon_index] / float(n_at_risk)
),
"prediction_sum": float(prediction_sum[horizon_index]),
"event_weight_sum": float(event_weight_sum[horizon_index]),
"brier_ipcw_sum": float(brier_sum[horizon_index]),
"nll_ipcw_sum": float(nll_sum[horizon_index]),
"predicted_mean": predicted_mean,
"observed_rate_ipcw": observed_rate,
"expected_observed_ratio": (
predicted_mean / observed_rate
if observed_rate > 0
else np.nan
),
"brier_ipcw": (
float(brier_sum[horizon_index]) / float(n_at_risk)
),
"nll_ipcw": (
float(nll_sum[horizon_index]) / float(n_at_risk)
),
"brier_complete_case": float(
complete_case_brier[horizon_index]
),
"nll_complete_case": float(
complete_case_nll[horizon_index]
),
"calibration_in_the_large": float(
calibration_in_large[horizon_index]
),
"calibration_intercept": float(
calibration_intercept[horizon_index]
),
"calibration_slope": float(
calibration_slope[horizon_index]
),
"ipcw_weight_max": float(
metric_weights[horizon_index].max()
),
"ipcw_weight_mean_known": float(
metric_weights[horizon_index].sum()
/ known_count[horizon_index]
),
"ipcw_weights_clipped": clipped_count,
}
arrays = {
"probabilities": p[horizon_index],
"outcomes": outcomes[horizon_index],
"metric_weights": metric_weights[horizon_index],
"brier_contribution": brier_contribution[horizon_index],
"nll_contribution": nll_contribution[horizon_index],
"known": known[horizon_index],
"event_by_horizon": event_by_horizon[horizon_index],
"controls": controls[horizon_index],
}
results.append((row, arrays))
return results
2026-07-30 17:21:57 +08:00
def _update_curve_accumulator(
accumulator: Dict[tuple[Any, ...], Dict[str, float]],
*,
outcome: str,
sex: str,
horizon: float,
probability_bins: np.ndarray,
arrays: Dict[str, np.ndarray],
) -> None:
probabilities = arrays["probabilities"]
bin_indices = np.searchsorted(
probability_bins,
probabilities,
side="right",
) - 1
bin_indices = np.clip(bin_indices, 0, len(probability_bins) - 2)
2026-07-30 18:29:15 +08:00
bin_count = int(len(probability_bins) - 1)
n_predictions = np.bincount(
bin_indices,
minlength=bin_count,
)
prediction_sum = np.bincount(
bin_indices,
weights=probabilities,
minlength=bin_count,
)
event_weight_sum = np.bincount(
bin_indices,
weights=arrays["metric_weights"] * arrays["outcomes"],
minlength=bin_count,
)
known_weight_sum = np.bincount(
bin_indices,
weights=arrays["metric_weights"],
minlength=bin_count,
)
brier_sum = np.bincount(
bin_indices,
weights=arrays["brier_contribution"],
minlength=bin_count,
)
nll_sum = np.bincount(
bin_indices,
weights=arrays["nll_contribution"],
minlength=bin_count,
)
n_events = np.bincount(
bin_indices,
weights=arrays["event_by_horizon"].astype(np.int64),
minlength=bin_count,
)
n_controls = np.bincount(
bin_indices,
weights=arrays["controls"].astype(np.int64),
minlength=bin_count,
)
2026-07-30 17:21:57 +08:00
2026-07-30 18:29:15 +08:00
for bin_index in np.flatnonzero(n_predictions).tolist():
2026-07-30 17:21:57 +08:00
key = (outcome, sex, float(horizon), int(bin_index))
row = accumulator[key]
2026-07-30 18:29:15 +08:00
row["n_predictions"] += int(n_predictions[bin_index])
row["prediction_sum"] += float(prediction_sum[bin_index])
row["event_weight_sum"] += float(event_weight_sum[bin_index])
row["known_weight_sum"] += float(known_weight_sum[bin_index])
row["brier_ipcw_sum"] += float(brier_sum[bin_index])
row["nll_ipcw_sum"] += float(nll_sum[bin_index])
row["n_events"] += int(n_events[bin_index])
row["n_controls"] += int(n_controls[bin_index])
def _empty_curve_values() -> Dict[str, float]:
return {
"n_predictions": 0,
"prediction_sum": 0.0,
"event_weight_sum": 0.0,
"known_weight_sum": 0.0,
"brier_ipcw_sum": 0.0,
"nll_ipcw_sum": 0.0,
"n_events": 0,
"n_controls": 0,
}
def _merge_curve_accumulator(
destination: Dict[tuple[Any, ...], Dict[str, float]],
source: Dict[tuple[Any, ...], Dict[str, float]],
) -> None:
for key, source_values in source.items():
destination_values = destination[key]
for field, value in source_values.items():
destination_values[field] += value
2026-07-30 17:21:57 +08:00
def _first_time_array(
first_occurrence_by_token: Dict[int, Tuple[np.ndarray, np.ndarray]],
token: int,
patient_count: int,
) -> np.ndarray:
result = np.full(patient_count, np.inf, dtype=np.float32)
pairs = first_occurrence_by_token.get(int(token))
if pairs is not None:
patient_ids, times = pairs
result[
np.asarray(patient_ids, dtype=np.int64)
] = np.asarray(times, dtype=np.float32)
return result
2026-07-30 18:29:15 +08:00
def _risk_probability_matrix(
logits: np.ndarray,
rho: Optional[np.ndarray],
horizons: np.ndarray,
dist_mode: str,
) -> np.ndarray:
"""Convert one token's logits to all horizon risks at once."""
logits = np.asarray(logits, dtype=np.float32)
horizons = np.asarray(horizons, dtype=np.float32)
rate = (
np.log1p(np.exp(-np.abs(logits)))
+ np.maximum(logits, np.float32(0.0))
+ np.float32(1e-8)
)
use_weibull = str(dist_mode).lower() == "weibull"
2026-07-30 18:29:15 +08:00
if use_weibull:
if rho is None:
raise RuntimeError(
"Weibull risk scoring requires rho parameters."
)
exposure = np.power(
horizons[:, None],
np.asarray(rho, dtype=np.float32)[None, :],
)
else:
exposure = horizons[:, None]
return (
-np.expm1(-rate[None, :] * exposure)
).astype(np.float64, copy=False)
def _build_calibration_strata(
row_sex: np.ndarray,
row_landmark_age: np.ndarray,
) -> List[tuple[str, float, np.ndarray]]:
row_sex = np.asarray(row_sex)
row_landmark_age = np.asarray(row_landmark_age)
strata: List[tuple[str, float, np.ndarray]] = []
for sex_value, sex_name in ((0, "Female"), (1, "Male")):
sex_rows = row_sex == int(sex_value)
if not np.any(sex_rows):
continue
for landmark_age_raw in np.unique(
row_landmark_age[sex_rows]
).tolist():
landmark_age = float(landmark_age_raw)
row_indices = np.flatnonzero(
sex_rows
& (row_landmark_age == np.float32(landmark_age))
)
if row_indices.size:
strata.append((sex_name, landmark_age, row_indices))
return strata
def _evaluate_calibration_token(
*,
column_index: int,
token: int,
logits_chunk: np.ndarray,
rho_chunk: Optional[np.ndarray],
strata: Sequence[tuple[str, float, np.ndarray]],
row_patient_id: np.ndarray,
row_followup_end: np.ndarray,
row_death_time: np.ndarray,
first_occurrence_by_token: Dict[
int, Tuple[np.ndarray, np.ndarray]
],
patient_count: int,
death_tokens: set[int],
label_id_to_code: Dict[int, str],
dist_mode: str,
horizons: np.ndarray,
min_cases: int,
min_controls: int,
max_ipcw_weight: float,
exclude_death_competing: bool,
probability_bins: np.ndarray,
) -> tuple[
List[Dict[str, Any]],
Dict[tuple[Any, ...], Dict[str, float]],
]:
token = int(token)
token_logits = logits_chunk[:, int(column_index)]
token_rho = (
None
if rho_chunk is None
else rho_chunk[:, int(column_index)]
)
first_time = _first_time_array(
first_occurrence_by_token,
token,
patient_count,
)
label_code = str(label_id_to_code.get(token, token))
outcome_name = "Death" if token in death_tokens else "Disease"
metric_rows: List[Dict[str, Any]] = []
curve_accumulator: Dict[
tuple[Any, ...], Dict[str, float]
] = defaultdict(_empty_curve_values)
for sex_name, landmark_age, stratum_indices in strata:
patient_ids = row_patient_id[stratum_indices]
token_first_time = first_time[patient_ids].astype(
np.float64,
copy=False,
)
at_risk = token_first_time > landmark_age
if not np.any(at_risk):
continue
row_indices = stratum_indices[at_risk]
token_first_time = token_first_time[at_risk]
followup_end = row_followup_end[row_indices].astype(
np.float64,
copy=False,
)
death_time = row_death_time[row_indices].astype(
np.float64,
copy=False,
)
effective_censor = followup_end.copy()
if exclude_death_competing and token not in death_tokens:
death_before_disease = death_time < token_first_time
effective_censor = np.where(
death_before_disease,
np.minimum(effective_censor, death_time),
effective_censor,
)
probabilities = _risk_probability_matrix(
logits=token_logits[row_indices],
rho=(
None
if token_rho is None
else token_rho[row_indices]
),
horizons=horizons,
dist_mode=dist_mode,
)
results = compute_ipcw_horizons(
probabilities=probabilities,
event_times=token_first_time - landmark_age,
censor_times=effective_censor - landmark_age,
horizons=horizons,
min_cases=min_cases,
min_controls=min_controls,
max_ipcw_weight=max_ipcw_weight,
)
for horizon, result in zip(horizons.tolist(), results):
if result is None:
continue
row, arrays = result
metric_rows.append(
{
"token": token,
"label_code": label_code,
"outcome": outcome_name,
"sex": sex_name,
"landmark_age": landmark_age,
"horizon": float(horizon),
**row,
}
)
_update_curve_accumulator(
curve_accumulator,
outcome=outcome_name,
sex=sex_name,
horizon=float(horizon),
probability_bins=probability_bins,
arrays=arrays,
)
return metric_rows, curve_accumulator
2026-07-30 17:21:57 +08:00
def evaluate_landmark_calibration(
*,
model: Any,
loader: DataLoader,
landmark_dataset: LandmarkDataset,
disease_ids: Sequence[int],
dist_mode: str,
horizons: np.ndarray,
device: torch.device,
use_amp: bool,
hidden_cache_dtype: str,
logit_batch_size: int,
disease_chunk_size: int,
min_cases: int,
min_controls: int,
max_ipcw_weight: float,
exclude_death_competing: bool,
probability_bins: np.ndarray,
2026-07-30 18:29:15 +08:00
num_workers_calibration: int,
2026-07-30 17:21:57 +08:00
) -> tuple[pd.DataFrame, pd.DataFrame]:
model.eval().to(device)
hidden_all, row_arrays = infer_landmark_hidden(
model=model,
loader=loader,
device=device,
model_target_mode="all_future",
use_amp=use_amp,
hidden_cache_dtype=hidden_cache_dtype,
)
print(
f"Cached landmark hidden: shape={hidden_all.shape}, "
f"dtype={hidden_all.dtype}"
)
disease_ids = [int(token) for token in disease_ids]
disease_chunk_size = (
len(disease_ids)
if int(disease_chunk_size) <= 0
else int(disease_chunk_size)
)
chunks = [
disease_ids[start:start + disease_chunk_size]
for start in range(0, len(disease_ids), disease_chunk_size)
]
patient_count = len(landmark_dataset.subset_indices)
death_tokens = set(int(value) for value in landmark_dataset.death_token_ids)
2026-07-30 18:29:15 +08:00
strata = _build_calibration_strata(
row_arrays["sex"],
row_arrays["landmark_age"],
)
worker_count = max(
1,
min(
int(num_workers_calibration),
len(disease_ids),
disease_chunk_size,
),
)
print(f"Calibration CPU workers: {worker_count}")
2026-07-30 17:21:57 +08:00
metric_rows: List[Dict[str, Any]] = []
curve_accumulator: Dict[
tuple[Any, ...], Dict[str, float]
2026-07-30 18:29:15 +08:00
] = defaultdict(_empty_curve_values)
2026-07-30 17:21:57 +08:00
2026-07-30 18:29:15 +08:00
executor = (
ThreadPoolExecutor(
max_workers=worker_count,
thread_name_prefix="calibration",
2026-07-30 17:21:57 +08:00
)
2026-07-30 18:29:15 +08:00
if worker_count > 1
else None
)
try:
for chunk_index, chunk_ids in enumerate(
tqdm(chunks, desc="Disease chunks", dynamic_ncols=True)
2026-07-30 17:21:57 +08:00
):
2026-07-30 18:29:15 +08:00
logits_chunk, rho_chunk = project_distribution_chunk(
model=model,
hidden_all=hidden_all,
disease_ids=chunk_ids,
dist_mode=dist_mode,
device=device,
logit_batch_size=logit_batch_size,
use_amp=use_amp,
2026-07-30 17:21:57 +08:00
)
2026-07-30 18:29:15 +08:00
task_kwargs = [
{
"column_index": column_index,
"token": int(token),
"logits_chunk": logits_chunk,
"rho_chunk": rho_chunk,
"strata": strata,
"row_patient_id": row_arrays["patient_id"],
"row_followup_end": row_arrays[
"followup_end_time"
],
"row_death_time": row_arrays["death_time"],
"first_occurrence_by_token": (
landmark_dataset.first_occurrence_by_token
),
"patient_count": patient_count,
"death_tokens": death_tokens,
"label_id_to_code": (
landmark_dataset.dataset.label_id_to_code
),
"dist_mode": dist_mode,
"horizons": horizons,
"min_cases": min_cases,
"min_controls": min_controls,
"max_ipcw_weight": max_ipcw_weight,
"exclude_death_competing": (
exclude_death_competing
),
"probability_bins": probability_bins,
}
for column_index, token in enumerate(chunk_ids)
]
if executor is None:
chunk_results = (
_evaluate_calibration_token(**kwargs)
for kwargs in tqdm(
task_kwargs,
desc=f"Calibration chunk {chunk_index}",
leave=False,
dynamic_ncols=True,
)
2026-07-30 17:21:57 +08:00
)
2026-07-30 18:29:15 +08:00
for rows, local_curve in chunk_results:
metric_rows.extend(rows)
_merge_curve_accumulator(
curve_accumulator,
local_curve,
)
else:
futures = [
executor.submit(
_evaluate_calibration_token,
**kwargs,
2026-07-30 17:21:57 +08:00
)
2026-07-30 18:29:15 +08:00
for kwargs in task_kwargs
]
for future in tqdm(
as_completed(futures),
total=len(futures),
desc=f"Calibration chunk {chunk_index}",
leave=False,
dynamic_ncols=True,
):
rows, local_curve = future.result()
metric_rows.extend(rows)
_merge_curve_accumulator(
curve_accumulator,
local_curve,
2026-07-30 17:21:57 +08:00
)
2026-07-30 18:29:15 +08:00
del logits_chunk, rho_chunk
finally:
if executor is not None:
executor.shutdown(wait=True, cancel_futures=True)
2026-07-30 17:21:57 +08:00
if not metric_rows:
raise RuntimeError(
"No calibration rows were produced. Check split, landmark ages, "
"horizons, min_cases, and disease selection."
)
2026-07-30 18:29:15 +08:00
metrics = (
pd.DataFrame(metric_rows)
.sort_values(
["token", "sex", "landmark_age", "horizon"],
kind="stable",
)
.reset_index(drop=True)
)
2026-07-30 17:21:57 +08:00
curve_rows: List[Dict[str, Any]] = []
for (outcome, sex, horizon, bin_index), values in sorted(
curve_accumulator.items()
):
n_predictions = int(values["n_predictions"])
curve_rows.append(
{
"outcome": outcome,
"sex": sex,
"horizon": float(horizon),
"probability_bin": int(bin_index),
"probability_lower": float(probability_bins[int(bin_index)]),
"probability_upper": float(
probability_bins[int(bin_index) + 1]
),
"n_predictions": n_predictions,
"n_events": int(values["n_events"]),
"n_controls": int(values["n_controls"]),
"predicted_mean": (
float(values["prediction_sum"]) / n_predictions
),
"observed_rate_ipcw": (
float(values["event_weight_sum"]) / n_predictions
),
"brier_ipcw": (
float(values["brier_ipcw_sum"]) / n_predictions
),
"nll_ipcw": (
float(values["nll_ipcw_sum"]) / n_predictions
),
"known_weight_sum": float(values["known_weight_sum"]),
}
)
curve = pd.DataFrame(curve_rows)
curve_all = _aggregate_curve_sexes(curve)
curve = pd.concat([curve, curve_all], ignore_index=True)
return metrics, curve
def _aggregate_curve_sexes(curve: pd.DataFrame) -> pd.DataFrame:
rows: List[Dict[str, Any]] = []
group_columns = [
"outcome",
"horizon",
"probability_bin",
"probability_lower",
"probability_upper",
]
for keys, group in curve.groupby(group_columns, sort=True, dropna=False):
(
outcome,
horizon,
probability_bin,
probability_lower,
probability_upper,
) = keys
n_predictions = int(group["n_predictions"].sum())
prediction_sum = float(
(group["predicted_mean"] * group["n_predictions"]).sum()
)
event_weight_sum = float(
(group["observed_rate_ipcw"] * group["n_predictions"]).sum()
)
brier_sum = float(
(group["brier_ipcw"] * group["n_predictions"]).sum()
)
nll_sum = float(
(group["nll_ipcw"] * group["n_predictions"]).sum()
)
rows.append(
{
"outcome": outcome,
"sex": "All",
"horizon": float(horizon),
"probability_bin": int(probability_bin),
"probability_lower": float(probability_lower),
"probability_upper": float(probability_upper),
"n_predictions": n_predictions,
"n_events": int(group["n_events"].sum()),
"n_controls": int(group["n_controls"].sum()),
"predicted_mean": prediction_sum / n_predictions,
"observed_rate_ipcw": event_weight_sum / n_predictions,
"brier_ipcw": brier_sum / n_predictions,
"nll_ipcw": nll_sum / n_predictions,
"known_weight_sum": float(group["known_weight_sum"].sum()),
}
)
return pd.DataFrame(rows)
def aggregate_metric_rows(
metrics: pd.DataFrame,
group_columns: Sequence[str],
) -> pd.DataFrame:
rows: List[Dict[str, Any]] = []
for keys, group in metrics.groupby(
list(group_columns),
sort=True,
dropna=False,
):
if not isinstance(keys, tuple):
keys = (keys,)
n_at_risk = int(group["n_at_risk"].sum())
prediction_sum = float(group["prediction_sum"].sum())
event_weight_sum = float(group["event_weight_sum"].sum())
predicted_mean = prediction_sum / n_at_risk
observed_rate = event_weight_sum / n_at_risk
rows.append(
{
**dict(zip(group_columns, keys)),
"n_landmark_cells": int(len(group)),
"n_at_risk": n_at_risk,
"n_events": int(group["n_events"].sum()),
"n_controls": int(group["n_controls"].sum()),
"n_censored_before_horizon": int(
group["n_censored_before_horizon"].sum()
),
"predicted_mean": predicted_mean,
"observed_rate_ipcw": observed_rate,
"expected_observed_ratio": (
predicted_mean / observed_rate
if observed_rate > 0
else np.nan
),
"brier_ipcw": float(group["brier_ipcw_sum"].sum())
/ n_at_risk,
"nll_ipcw": float(group["nll_ipcw_sum"].sum())
/ n_at_risk,
"calibration_in_the_large_median": float(
group["calibration_in_the_large"].median()
),
"calibration_intercept_median": float(
group["calibration_intercept"].median()
),
"calibration_slope_median": float(
group["calibration_slope"].median()
),
"ipcw_weight_max": float(group["ipcw_weight_max"].max()),
"ipcw_weights_clipped": int(
group["ipcw_weights_clipped"].sum()
),
}
)
return pd.DataFrame(rows)
def build_calibration_summary(metrics: pd.DataFrame) -> pd.DataFrame:
by_sex = aggregate_metric_rows(
metrics,
group_columns=["outcome", "sex", "horizon"],
)
all_sexes = aggregate_metric_rows(
metrics,
group_columns=["outcome", "horizon"],
)
all_sexes.insert(1, "sex", "All")
return (
pd.concat([by_sex, all_sexes], ignore_index=True)
.sort_values(["outcome", "sex", "horizon"], kind="stable")
.reset_index(drop=True)
)
def _build_point_process_criterion(
dist_mode: str,
) -> Any:
ignored = {PAD_IDX, RESERVED_IDX, NO_EVENT_IDX}
2026-07-30 17:21:57 +08:00
if dist_mode == "exponential":
return build_loss("exponential", ignored_idx=ignored)
if dist_mode == "weibull":
return build_loss("weibull", ignored_idx=ignored)
raise ValueError(f"Unsupported dist_mode: {dist_mode!r}")
@torch.inference_mode()
def evaluate_point_process_nll(
*,
model: Any,
loader: DataLoader,
dist_mode: str,
device: torch.device,
use_amp: bool,
) -> Dict[str, Any]:
criterion = _build_point_process_criterion(dist_mode)
2026-07-30 17:21:57 +08:00
model.eval().to(device)
total_nll = 0.0
query_count = 0
future_event_count = 0
exposure_sum = 0.0
amp_enabled = bool(use_amp and device.type == "cuda")
for batch in tqdm(
loader,
desc="Point-process NLL",
leave=False,
dynamic_ncols=True,
):
batch_device = {
key: (
value.to(device, non_blocking=True)
if isinstance(value, torch.Tensor)
else value
)
for key, value in batch.items()
}
amp_context = (
torch.autocast(device_type=device.type, dtype=torch.float16)
if amp_enabled
else contextlib.nullcontext()
)
with amp_context:
hidden = model(
event_seq=batch_device["event_seq"],
time_seq=batch_device["time_seq"],
sex=batch_device["sex"],
padding_mask=batch_device["padding_mask"],
t_query=batch_device["t_query"],
other_type=batch_device["other_type"],
other_value=batch_device["other_value"],
other_value_kind=batch_device["other_value_kind"],
other_time=batch_device["other_time"],
)
logits = model.calc_risk(hidden)
if dist_mode == "exponential":
loss = criterion(
logits=logits,
targets=batch_device["future_targets"],
exposure=batch_device["exposure"],
dt=batch_device["future_dt"],
history=batch_device["event_seq"],
2026-07-30 17:21:57 +08:00
)
elif dist_mode == "weibull":
loss = criterion(
logits=logits,
weibull_rho=model.calc_weibull_rho(hidden),
targets=batch_device["future_targets"],
dt=batch_device["future_dt"],
exposure=batch_device["exposure"],
history=batch_device["event_seq"],
2026-07-30 17:21:57 +08:00
)
else:
raise ValueError(f"Unsupported dist_mode: {dist_mode!r}")
2026-07-30 17:21:57 +08:00
if not torch.isfinite(loss):
raise RuntimeError("Non-finite point-process NLL encountered.")
batch_size = int(batch_device["event_seq"].shape[0])
total_nll += float(loss.detach().cpu()) * batch_size
query_count += batch_size
valid_targets = batch["future_targets"] > PAD_IDX
valid_targets &= batch["future_targets"] != RESERVED_IDX
valid_targets &= batch["future_targets"] != NO_EVENT_IDX
2026-07-30 17:21:57 +08:00
future_event_count += int(valid_targets.sum().item())
exposure_sum += float(batch["exposure"].sum().item())
if query_count == 0:
raise RuntimeError("Point-process NLL loader produced no queries.")
return {
"query_count": query_count,
"future_event_count": future_event_count,
"total_point_process_nll": total_nll,
"mean_point_process_nll_per_query": total_nll / query_count,
"mean_point_process_nll_per_future_event": (
total_nll / future_event_count
if future_event_count > 0
else np.nan
),
"mean_exposure_years": exposure_sum / query_count,
}
def _atomic_csv(frame: pd.DataFrame, path: Path) -> None:
temporary = path.with_name(f".{path.name}.tmp")
frame.to_csv(temporary, index=False)
os.replace(temporary, path)
def _atomic_json(payload: Dict[str, Any], path: Path) -> None:
temporary = path.with_name(f".{path.name}.tmp")
temporary.write_text(
json.dumps(payload, ensure_ascii=False, indent=2),
encoding="utf-8",
)
os.replace(temporary, path)
def main() -> None:
parser = build_parser()
args = parser.parse_args()
run_path = Path(args.run_path).resolve()
config_path = run_path / "train_config.json"
checkpoint_path = run_path / "best_model.pt"
if not config_path.is_file():
raise FileNotFoundError(config_path)
if not checkpoint_path.is_file():
raise FileNotFoundError(checkpoint_path)
cfg = load_json_config(config_path)
validate_training_mode_config(cfg)
model_target_mode = str(
cfg.get("model_target_mode", "next_token")
).lower()
if model_target_mode != "all_future":
raise ValueError(
"evaluate_calibration.py supports all_future runs only; "
f"got model_target_mode={model_target_mode!r}"
)
output_path = Path(args.output_path or run_path).resolve()
output_path.mkdir(parents=True, exist_ok=True)
completion_path = output_path / COMPLETION_FILE
if completion_path.is_file() and completion_path.stat().st_size > 0:
if not args.force:
print(f"[SKIP] Existing completion marker: {completion_path}")
return
eval_split = _normalise_eval_split(args.eval_split)
subset_size = cfg_get(args, cfg, "dataset_subset_size", None)
disease_history_mode = normalize_disease_history_mode(
cfg.get("disease_history_mode", DISEASE_HISTORY_MODE_TIMED)
)
horizons = np.asarray(
parse_float_list(
cfg_get(args, cfg, "horizons", None)
)
or DEFAULT_HORIZONS,
dtype=np.float32,
)
if horizons.size == 0 or np.any(horizons <= 0):
raise ValueError("horizons must contain positive values")
probability_bins = np.asarray(
parse_float_list(args.probability_bins)
or DEFAULT_PROBABILITY_BINS,
dtype=np.float64,
)
if (
probability_bins.size < 2
or probability_bins[0] != 0.0
or probability_bins[-1] != 1.0
or np.any(np.diff(probability_bins) <= 0)
):
raise ValueError(
"probability_bins must be strictly increasing from 0 to 1"
)
print("Loading all-future sequence evaluation dataset...")
dataset = load_sequence_eval_dataset(
model_target_mode="all_future",
data_prefix=str(cfg.get("data_prefix", "ukb")),
labels_file=str(cfg.get("labels_file", "labels.csv")),
no_event_interval_years=float(
cfg.get("no_event_interval_years", 5.0)
),
min_history_events=int(
cfg.get("all_future_min_history_events", 1)
),
min_future_events=int(
cfg.get("all_future_min_future_events", 1)
),
extra_info_types=parse_int_list(cfg.get("extra_info_types")),
disease_history_mode=disease_history_mode,
)
validate_dataset_metadata(dataset, cfg)
subset_indices, split_method, split_source = select_sequence_eval_indices(
dataset,
cfg,
eval_split,
subset_size,
)
first_occurrence_by_token = build_first_occurrence_map(
dataset,
subset_indices,
)
labels_meta_path = cfg_get(args, cfg, "labels_meta_path", None)
if labels_meta_path is None:
labels_meta_path = cfg.get(
"labels_meta_path",
"delphi_labels_chapters_colours_icd.csv",
)
labels_meta = None
if labels_meta_path:
resolved_labels_meta = _resolve_project_path(str(labels_meta_path))
if resolved_labels_meta.is_file():
labels_meta = pd.read_csv(resolved_labels_meta)
requested_diseases = parse_int_list(args.diseases_of_interest)
disease_ids = select_disease_tokens(
dataset=dataset,
labels_meta=labels_meta,
requested_tokens=requested_diseases,
filter_min_total=int(
cfg_get(args, cfg, "filter_min_total", 0)
),
first_occurrence_by_token=first_occurrence_by_token,
)
if not disease_ids:
raise RuntimeError("No disease tokens selected.")
state_dict = load_checkpoint_state_dict(
checkpoint_path,
map_location="cpu",
)
dist_mode = resolve_dist_mode_for_checkpoint(
str(cfg.get("dist_mode", "exponential")),
state_dict,
)
cfg_model = dict(cfg)
cfg_model["dist_mode"] = dist_mode
cfg_model["model_architecture"] = resolve_model_architecture(
cfg_model,
state_dict,
)
device = resolve_eval_device(args.device)
model = build_model_from_dataset(
args,
cfg_model,
dataset,
state_dict=state_dict,
).to(device)
load_model_state(model, state_dict)
landmark_start = float(
cfg_get(args, cfg, "landmark_start", 40.0)
)
landmark_stop = float(
cfg_get(args, cfg, "landmark_stop", 80.0)
)
landmark_step = float(
cfg_get(args, cfg, "landmark_step", 5.0)
)
if landmark_step <= 0:
raise ValueError("landmark_step must be > 0")
landmark_ages = np.arange(
landmark_start,
landmark_stop,
landmark_step,
dtype=np.float32,
)
if landmark_ages.size == 0:
raise ValueError("No landmark ages selected.")
2026-08-03 08:33:21 +08:00
death_token_ids = _get_death_token_ids(dataset)
2026-07-30 17:21:57 +08:00
landmark_dataset = LandmarkDataset(
dataset=dataset,
subset_indices=subset_indices,
landmark_ages=landmark_ages,
model_target_mode="all_future",
min_history_events=int(
cfg_get(args, cfg, "min_history_events", 1)
),
first_occurrence_by_token=first_occurrence_by_token,
death_token_ids=death_token_ids,
disease_history_mode=disease_history_mode,
)
batch_size = int(cfg_get(args, cfg, "batch_size", 128))
num_workers = int(cfg_get(args, cfg, "num_workers", 4))
use_amp = bool(cfg_get(args, cfg, "use_amp", False))
hidden_cache_dtype = str(
cfg_get(args, cfg, "hidden_cache_dtype", "float16")
)
logit_batch_size = int(
cfg_get(args, cfg, "logit_batch_size", batch_size)
)
disease_chunk_size = int(
cfg_get(args, cfg, "disease_chunk_size", 64)
)
2026-07-30 18:29:15 +08:00
num_workers_calibration = int(
cfg_get(args, cfg, "num_workers_calibration", 0)
)
if num_workers_calibration < 0:
raise ValueError("num_workers_calibration must be >= 0")
if num_workers_calibration == 0:
num_workers_calibration = max(1, int(os.cpu_count() or 1))
2026-07-30 17:21:57 +08:00
min_cases = int(cfg_get(args, cfg, "min_cases", 2))
min_controls = int(args.min_controls)
if min_cases < 1:
raise ValueError("min_cases must be >= 1")
if min_controls < 1:
raise ValueError("min_controls must be >= 1")
if float(args.max_ipcw_weight) < 0:
raise ValueError("max_ipcw_weight must be >= 0")
exclude_death_competing = bool(
cfg_get(
args,
cfg,
"exclude_death_in_window_without_disease",
True,
)
)
landmark_loader = DataLoader(
landmark_dataset,
batch_size=batch_size,
shuffle=False,
collate_fn=collate_landmark_fn,
num_workers=num_workers,
pin_memory=device.type == "cuda",
persistent_workers=num_workers > 0,
prefetch_factor=2 if num_workers > 0 else None,
)
print(f"Run: {run_path}")
print(f"Eval split: {eval_split} ({split_method})")
print(f"Selected patients: {len(subset_indices)}")
print(f"Landmark queries: {len(landmark_dataset)}")
print(f"Disease tokens: {len(disease_ids)}")
print(f"Dist mode: {dist_mode}")
print(f"Disease history mode: {disease_history_mode}")
print(f"Horizons: {horizons.tolist()}")
landmark_metrics, calibration_curve = evaluate_landmark_calibration(
model=model,
loader=landmark_loader,
landmark_dataset=landmark_dataset,
disease_ids=disease_ids,
dist_mode=dist_mode,
horizons=horizons,
device=device,
use_amp=use_amp,
hidden_cache_dtype=hidden_cache_dtype,
logit_batch_size=logit_batch_size,
disease_chunk_size=disease_chunk_size,
min_cases=min_cases,
min_controls=min_controls,
max_ipcw_weight=float(args.max_ipcw_weight),
exclude_death_competing=exclude_death_competing,
probability_bins=probability_bins,
2026-07-30 18:29:15 +08:00
num_workers_calibration=num_workers_calibration,
2026-07-30 17:21:57 +08:00
)
token_metrics = aggregate_metric_rows(
landmark_metrics,
group_columns=[
"token",
"label_code",
"outcome",
"sex",
"horizon",
],
)
calibration_summary = build_calibration_summary(landmark_metrics)
point_process_row: Dict[str, Any] = {
"evaluated": False,
"eval_split": eval_split,
}
if args.point_process_nll:
(
point_dataset,
point_subset,
point_split_method,
point_split_source,
) = build_point_process_eval_subset(
cfg,
eval_split,
disease_history_mode,
subset_size,
)
point_loader = DataLoader(
point_subset,
batch_size=batch_size,
shuffle=False,
collate_fn=all_future_collate_fn,
num_workers=num_workers,
pin_memory=device.type == "cuda",
persistent_workers=num_workers > 0,
prefetch_factor=2 if num_workers > 0 else None,
)
point_process_row = {
"evaluated": True,
"eval_split": eval_split,
"split_method": point_split_method,
"split_source": point_split_source,
**evaluate_point_process_nll(
model=model,
loader=point_loader,
dist_mode=dist_mode,
device=device,
use_amp=use_amp,
),
}
del point_dataset, point_subset, point_loader
landmark_path = output_path / LANDMARK_METRICS_FILE
token_path = output_path / TOKEN_METRICS_FILE
summary_path = output_path / SUMMARY_FILE
curve_path = output_path / CALIBRATION_CURVE_FILE
nll_path = output_path / POINT_PROCESS_NLL_FILE
_atomic_csv(landmark_metrics, landmark_path)
_atomic_csv(token_metrics, token_path)
_atomic_csv(calibration_summary, summary_path)
_atomic_csv(calibration_curve, curve_path)
_atomic_csv(pd.DataFrame([point_process_row]), nll_path)
completion = {
"status": "complete",
"run_path": str(run_path),
"output_path": str(output_path),
"model_target_mode": model_target_mode,
"model_architecture": cfg_model["model_architecture"],
"dist_mode": dist_mode,
"time_mode": str(cfg.get("time_mode", "")),
"disease_history_mode": disease_history_mode,
"extra_info_types": parse_int_list(cfg.get("extra_info_types")) or [],
"seed": int(cfg.get("seed", 42)),
"eval_split": eval_split,
"split_method": split_method,
"split_source": split_source,
"selected_patient_count": int(len(subset_indices)),
"landmark_query_count": int(len(landmark_dataset)),
"disease_token_count": int(len(disease_ids)),
"landmark_ages": [float(value) for value in landmark_ages],
"horizons": [float(value) for value in horizons],
"min_cases": min_cases,
"min_controls": min_controls,
2026-07-30 18:29:15 +08:00
"num_workers_calibration": num_workers_calibration,
2026-07-30 17:21:57 +08:00
"exclude_death_competing": exclude_death_competing,
"landmark_estimand": (
"first-onset fixed-horizon risk; for non-death outcomes, "
"death before disease is treated as censoring"
if exclude_death_competing
else "first-onset fixed-horizon risk without death censoring"
),
"landmark_nll_definition": (
"IPCW binary negative log-likelihood at each fixed horizon"
),
"point_process_nll_definition": (
"exact all-future continuous-time training likelihood "
"on deterministic evaluation queries"
),
"max_ipcw_weight": float(args.max_ipcw_weight),
"probability_bins": [
float(value) for value in probability_bins
],
"landmark_metric_rows": int(len(landmark_metrics)),
"token_metric_rows": int(len(token_metrics)),
"summary_rows": int(len(calibration_summary)),
"calibration_curve_rows": int(len(calibration_curve)),
"point_process_nll": point_process_row,
"outputs": {
"landmark_metrics": str(landmark_path),
"token_metrics": str(token_path),
"summary": str(summary_path),
"calibration_curve": str(curve_path),
"point_process_nll": str(nll_path),
},
}
_atomic_json(completion, completion_path)
print(f"Saved: {landmark_path}")
print(f"Saved: {token_path}")
print(f"Saved: {summary_path}")
print(f"Saved: {curve_path}")
print(f"Saved: {nll_path}")
print(f"Completion marker: {completion_path}")
if __name__ == "__main__":
main()