"""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 from concurrent.futures import ThreadPoolExecutor, as_completed from pathlib import Path from typing import Any, Dict, List, Optional, Sequence, Tuple 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 PAD_IDX, RESERVED_IDX 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) parser.add_argument( "--num_workers_calibration", type=int, default=None, help=( "CPU threads for per-disease calibration statistics. " "0 uses all logical CPUs." ), ) 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.""" 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.""" probabilities = np.asarray(probabilities, dtype=np.float64) outcomes = np.asarray(outcomes, dtype=np.float64) weights = np.asarray(weights, dtype=np.float64) 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") valid = ( np.isfinite(probabilities) & np.isfinite(outcomes) & np.isfinite(weights) & (weights > 0) ) 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) logit_p = np.log(p) - np.log1p(-p) 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() for _ in range(max_iter): if not np.any(intercept_active): break 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 for _ in range(max_iter): if not np.any(beta_active): break 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 def _censoring_km( observed_times: np.ndarray, censor_events: np.ndarray, ) -> tuple[np.ndarray, np.ndarray]: """Estimate censoring survival in O(N log N) time.""" 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) 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), ) 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) 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.""" 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, ) event_times = np.asarray(event_times, dtype=np.float64) censor_times = np.asarray(censor_times, dtype=np.float64) 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]) if n_at_risk == 0: return [None] * int(horizons.size) observed_event = event_times <= censor_times event_by_horizon = ( observed_event[None, :] & (event_times[None, :] <= horizons[:, None]) ) 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) 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, event_times, before=True, ) control_g = _km_value( km_times, km_survival, horizons, before=False, ) 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, max_ipcw_weight, ) control_weight_values, _ = _cap_ipcw( raw_control_weights, max_ipcw_weight, ) 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) ) 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) 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 = ( fit_weighted_logistic_calibration_batch( probabilities=p, outcomes=outcomes, weights=metric_weights, ) ) 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 ) + ( 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 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) 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, ) for bin_index in np.flatnonzero(n_predictions).tolist(): key = (outcome, sex, float(horizon), int(bin_index)) row = accumulator[key] 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 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 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" 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 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, num_workers_calibration: int, ) -> 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) 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}") metric_rows: List[Dict[str, Any]] = [] curve_accumulator: Dict[ tuple[Any, ...], Dict[str, float] ] = defaultdict(_empty_curve_values) executor = ( ThreadPoolExecutor( max_workers=worker_count, thread_name_prefix="calibration", ) if worker_count > 1 else None ) try: for chunk_index, chunk_ids in enumerate( tqdm(chunks, desc="Disease chunks", dynamic_ncols=True) ): 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, ) 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, ) ) 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, ) 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, ) del logits_chunk, rho_chunk finally: if executor is not None: executor.shutdown(wait=True, cancel_futures=True) if not metric_rows: raise RuntimeError( "No calibration rows were produced. Check split, landmark ages, " "horizons, min_cases, and disease selection." ) metrics = ( pd.DataFrame(metric_rows) .sort_values( ["token", "sex", "landmark_age", "horizon"], kind="stable", ) .reset_index(drop=True) ) 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} 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) 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"], ) 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"], ) else: raise ValueError(f"Unsupported dist_mode: {dist_mode!r}") 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 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.") death_token_ids = _get_death_token_ids(dataset, labels_meta) 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) ) 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)) 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, num_workers_calibration=num_workers_calibration, ) 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, "num_workers_calibration": num_workers_calibration, "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()