Optimize calibration evaluation

This commit is contained in:
2026-07-30 18:29:15 +08:00
parent e471de030d
commit 4580058687
3 changed files with 1121 additions and 278 deletions

File diff suppressed because it is too large Load Diff

View File

@@ -29,6 +29,7 @@ LOG_ROOT="$SCRIPT_DIR/batch_logs/evaluate_calibration_all_runs"
GPU_CSV="0" GPU_CSV="0"
PYTHON_BIN="${PYTHON_BIN:-python}" PYTHON_BIN="${PYTHON_BIN:-python}"
NUM_WORKERS=4 NUM_WORKERS=4
NUM_WORKERS_CALIBRATION=0
BATCH_SIZE=128 BATCH_SIZE=128
DISEASE_CHUNK_SIZE=64 DISEASE_CHUNK_SIZE=64
HORIZONS="" HORIZONS=""
@@ -52,6 +53,8 @@ Options:
--python PATH Python executable --python PATH Python executable
(default: $PYTHON_BIN or python). (default: $PYTHON_BIN or python).
--num-workers N DataLoader workers per job (default: 4). --num-workers N DataLoader workers per job (default: 4).
--num-workers-calibration N CPU calibration workers per job. Default: 0,
which divides all logical CPUs across GPUs.
--batch-size N Evaluation batch size (default: 128). --batch-size N Evaluation batch size (default: 128).
--disease-chunk-size N Disease projection chunk size (default: 64). --disease-chunk-size N Disease projection chunk size (default: 64).
--horizons LIST Optional comma-separated horizons in years. --horizons LIST Optional comma-separated horizons in years.
@@ -107,6 +110,14 @@ while (($# > 0)); do
NUM_WORKERS="$2" NUM_WORKERS="$2"
shift 2 shift 2
;; ;;
--num-workers-calibration)
[[ $# -ge 2 ]] || {
echo "ERROR: --num-workers-calibration requires a value." >&2
exit 2
}
NUM_WORKERS_CALIBRATION="$2"
shift 2
;;
--batch-size) --batch-size)
[[ $# -ge 2 ]] || { [[ $# -ge 2 ]] || {
echo "ERROR: --batch-size requires a value." >&2 echo "ERROR: --batch-size requires a value." >&2
@@ -163,6 +174,10 @@ done
echo "ERROR: --num-workers must be a non-negative integer." >&2 echo "ERROR: --num-workers must be a non-negative integer." >&2
exit 2 exit 2
} }
[[ "$NUM_WORKERS_CALIBRATION" =~ ^[0-9]+$ ]] || {
echo "ERROR: --num-workers-calibration must be a non-negative integer." >&2
exit 2
}
[[ "$BATCH_SIZE" =~ ^[1-9][0-9]*$ ]] || { [[ "$BATCH_SIZE" =~ ^[1-9][0-9]*$ ]] || {
echo "ERROR: --batch-size must be a positive integer." >&2 echo "ERROR: --batch-size must be a positive integer." >&2
exit 2 exit 2
@@ -203,6 +218,14 @@ for gpu in "${GPU_IDS[@]}"; do
SEEN_GPUS["$gpu"]=1 SEEN_GPUS["$gpu"]=1
done done
if ((NUM_WORKERS_CALIBRATION == 0)); then
TOTAL_CPUS="$(nproc)"
NUM_WORKERS_CALIBRATION=$(( (TOTAL_CPUS + ${#GPU_IDS[@]} - 1) / ${#GPU_IDS[@]} ))
if ((NUM_WORKERS_CALIBRATION < 1)); then
NUM_WORKERS_CALIBRATION=1
fi
fi
declare -a JOB_RUN_DIRS=() declare -a JOB_RUN_DIRS=()
declare -a JOB_LOG_FILES=() declare -a JOB_LOG_FILES=()
@@ -281,6 +304,7 @@ run_job() {
--eval_split test --eval_split test
--device cuda --device cuda
--num_workers "$NUM_WORKERS" --num_workers "$NUM_WORKERS"
--num_workers_calibration "$NUM_WORKERS_CALIBRATION"
--batch_size "$BATCH_SIZE" --batch_size "$BATCH_SIZE"
--disease_chunk_size "$DISEASE_CHUNK_SIZE" --disease_chunk_size "$DISEASE_CHUNK_SIZE"
) )
@@ -341,6 +365,7 @@ worker() {
echo "Runs root: $RUNS_ROOT" echo "Runs root: $RUNS_ROOT"
echo "GPUs: ${GPU_IDS[*]}" echo "GPUs: ${GPU_IDS[*]}"
echo "Calibration CPU workers per GPU job: $NUM_WORKERS_CALIBRATION"
echo "Runs discovered: $run_count" echo "Runs discovered: $run_count"
echo "Incomplete runs skipped: $incomplete_count" echo "Incomplete runs skipped: $incomplete_count"
echo "next_token runs skipped: $next_token_count" echo "next_token runs skipped: $next_token_count"

View File

@@ -1,17 +1,66 @@
import math import math
import unittest import unittest
from concurrent.futures import ThreadPoolExecutor
from types import SimpleNamespace
from unittest.mock import patch
import numpy as np import numpy as np
import pandas as pd import pandas as pd
import torch
from evaluate_calibration import ( from evaluate_calibration import (
_censoring_km,
_evaluate_calibration_token,
_risk_probability_matrix,
aggregate_metric_rows, aggregate_metric_rows,
compute_ipcw_cell, compute_ipcw_cell,
compute_ipcw_horizons,
evaluate_landmark_calibration,
fit_weighted_logistic_calibration, fit_weighted_logistic_calibration,
fit_weighted_logistic_calibration_batch,
) )
from evaluate_auc_v2 import _score_to_probability
class IPCWCalibrationMetricTests(unittest.TestCase): class IPCWCalibrationMetricTests(unittest.TestCase):
@staticmethod
def _naive_censoring_km(observed_times, censor_events):
observed_times = np.asarray(observed_times, dtype=np.float64)
censor_events = np.asarray(censor_events, dtype=bool)
event_times = np.unique(observed_times[censor_events])
survival = 1.0
survival_after = []
for time_value in event_times:
at_risk = np.sum(observed_times >= time_value)
censored = np.sum(
censor_events & (observed_times == time_value)
)
survival *= 1.0 - float(censored) / float(at_risk)
survival_after.append(survival)
return event_times, np.asarray(survival_after)
def test_sorted_censoring_km_matches_naive_reference(self):
rng = np.random.RandomState(12)
observed_times = rng.randint(1, 20, size=500).astype(np.float64)
censor_events = rng.uniform(size=500) < 0.4
expected_times, expected_survival = self._naive_censoring_km(
observed_times,
censor_events,
)
actual_times, actual_survival = _censoring_km(
observed_times,
censor_events,
)
np.testing.assert_array_equal(actual_times, expected_times)
np.testing.assert_allclose(
actual_survival,
expected_survival,
rtol=1e-14,
atol=1e-14,
)
def test_no_censoring_matches_binary_metrics(self): def test_no_censoring_matches_binary_metrics(self):
result = compute_ipcw_cell( result = compute_ipcw_cell(
probabilities=np.asarray([0.2, 0.8]), probabilities=np.asarray([0.2, 0.8]),
@@ -83,6 +132,278 @@ class IPCWCalibrationMetricTests(unittest.TestCase):
self.assertAlmostEqual(intercept, 0.0, places=7) self.assertAlmostEqual(intercept, 0.0, places=7)
self.assertAlmostEqual(slope, 1.0, places=7) self.assertAlmostEqual(slope, 1.0, places=7)
def test_batched_calibration_fits_multiple_horizons(self):
probabilities = np.vstack(
[
np.repeat([0.1, 0.3, 0.7, 0.9], 100),
np.repeat([0.2, 0.4, 0.6, 0.8], 100),
]
)
outcomes = np.vstack(
[
np.concatenate(
[
np.r_[np.ones(10), np.zeros(90)],
np.r_[np.ones(30), np.zeros(70)],
np.r_[np.ones(70), np.zeros(30)],
np.r_[np.ones(90), np.zeros(10)],
]
),
np.concatenate(
[
np.r_[np.ones(20), np.zeros(80)],
np.r_[np.ones(40), np.zeros(60)],
np.r_[np.ones(60), np.zeros(40)],
np.r_[np.ones(80), np.zeros(20)],
]
),
]
)
calibration_in_large, intercept, slope = (
fit_weighted_logistic_calibration_batch(
probabilities,
outcomes,
np.ones_like(probabilities),
)
)
np.testing.assert_allclose(calibration_in_large, 0.0, atol=1e-7)
np.testing.assert_allclose(intercept, 0.0, atol=1e-7)
np.testing.assert_allclose(slope, 1.0, atol=1e-7)
def test_all_horizons_reuse_one_censoring_km(self):
probabilities = np.asarray(
[
[0.05, 0.10, 0.15, 0.20, 0.25],
[0.10, 0.20, 0.30, 0.40, 0.50],
[0.20, 0.35, 0.50, 0.65, 0.80],
]
)
event_times = np.asarray([0.5, 1.5, 4.0, np.inf, np.inf])
censor_times = np.asarray([5.0, 5.0, 5.0, 2.5, 5.0])
with patch(
"evaluate_calibration._censoring_km",
wraps=_censoring_km,
) as km:
results = compute_ipcw_horizons(
probabilities=probabilities,
event_times=event_times,
censor_times=censor_times,
horizons=np.asarray([1.0, 2.0, 5.0]),
min_cases=1,
min_controls=1,
max_ipcw_weight=0.0,
)
self.assertEqual(km.call_count, 1)
self.assertEqual(len(results), 3)
self.assertTrue(all(result is not None for result in results))
self.assertEqual([result[0]["n_events"] for result in results], [1, 2, 3])
def test_batched_risk_probabilities_match_scalar_reference(self):
logits = np.asarray([-2.0, -0.5, 0.2, 1.5], dtype=np.float32)
rho = np.asarray([0.8, 1.0, 1.2, 1.5], dtype=np.float32)
horizons = np.asarray([0.1, 1.0, 5.0], dtype=np.float32)
for dist_mode, token, death_idx, selected_rho in (
("exponential", 4, 9, None),
("weibull", 4, 9, rho),
("mixed", 9, 9, rho),
("mixed", 4, 9, None),
):
actual = _risk_probability_matrix(
logits=logits,
rho=selected_rho,
horizons=horizons,
dist_mode=dist_mode,
token=token,
death_idx=death_idx,
)
expected = np.vstack(
[
_score_to_probability(
logits,
selected_rho,
score_mode="risk",
horizon=float(horizon),
dist_mode=dist_mode,
token=token,
death_idx=death_idx,
)
for horizon in horizons
]
)
np.testing.assert_allclose(actual, expected, rtol=1e-6, atol=1e-7)
def test_per_disease_worker_is_thread_safe(self):
logits_chunk = np.asarray(
[
[-2.0, -1.5],
[-1.0, -0.5],
[0.0, 0.5],
[0.5, 1.0],
[1.0, 1.5],
[1.5, 2.0],
],
dtype=np.float32,
)
common = {
"logits_chunk": logits_chunk,
"rho_chunk": None,
"strata": [("Female", 50.0, np.arange(6, dtype=np.int64))],
"row_patient_id": np.arange(6, dtype=np.int32),
"row_followup_end": np.full(6, 65.0, dtype=np.float32),
"row_death_time": np.full(6, np.inf, dtype=np.float32),
"first_occurrence_by_token": {
4: (
np.asarray([0, 1], dtype=np.int32),
np.asarray([50.5, 52.0], dtype=np.float32),
),
5: (
np.asarray([2, 3], dtype=np.int32),
np.asarray([50.7, 53.0], dtype=np.float32),
),
},
"patient_count": 6,
"death_tokens": set(),
"label_id_to_code": {4: "D4", 5: "D5"},
"dist_mode": "exponential",
"horizons": np.asarray([1.0, 5.0], dtype=np.float32),
"death_index": 9,
"min_cases": 1,
"min_controls": 1,
"max_ipcw_weight": 0.0,
"exclude_death_competing": True,
"probability_bins": np.asarray([0.0, 0.5, 1.0]),
}
tasks = [
{"column_index": 0, "token": 4, **common},
{"column_index": 1, "token": 5, **common},
]
serial = [
_evaluate_calibration_token(**task)
for task in tasks
]
with ThreadPoolExecutor(max_workers=2) as executor:
parallel = list(
executor.map(
lambda task: _evaluate_calibration_token(**task),
tasks,
)
)
for (serial_rows, serial_curve), (
parallel_rows,
parallel_curve,
) in zip(serial, parallel):
pd.testing.assert_frame_equal(
pd.DataFrame(serial_rows),
pd.DataFrame(parallel_rows),
)
self.assertEqual(set(serial_curve), set(parallel_curve))
for key in serial_curve:
self.assertEqual(
set(serial_curve[key]),
set(parallel_curve[key]),
)
np.testing.assert_allclose(
list(serial_curve[key].values()),
list(parallel_curve[key].values()),
equal_nan=True,
)
def test_landmark_evaluation_parallel_matches_serial(self):
logits_chunk = np.asarray(
[
[-2.0, -1.5],
[-1.0, -0.5],
[0.0, 0.5],
[0.5, 1.0],
[1.0, 1.5],
[1.5, 2.0],
],
dtype=np.float32,
)
row_arrays = {
"patient_id": np.arange(6, dtype=np.int32),
"sex": np.zeros(6, dtype=np.int8),
"landmark_age": np.full(6, 50.0, dtype=np.float32),
"followup_end_time": np.full(6, 65.0, dtype=np.float32),
"death_time": np.full(6, np.inf, dtype=np.float32),
}
landmark_dataset = SimpleNamespace(
subset_indices=np.arange(6, dtype=np.int64),
death_token_ids=[],
first_occurrence_by_token={
4: (
np.asarray([0, 1], dtype=np.int32),
np.asarray([50.5, 52.0], dtype=np.float32),
),
5: (
np.asarray([2, 3], dtype=np.int32),
np.asarray([50.7, 53.0], dtype=np.float32),
),
},
dataset=SimpleNamespace(
label_id_to_code={4: "D4", 5: "D5"}
),
)
class FakeModel:
death_idx = 9
vocab_size = 10
def eval(self):
return self
def to(self, _device):
return self
common = {
"model": FakeModel(),
"loader": [],
"landmark_dataset": landmark_dataset,
"disease_ids": [4, 5],
"dist_mode": "exponential",
"horizons": np.asarray([1.0, 5.0], dtype=np.float32),
"device": torch.device("cpu"),
"use_amp": False,
"hidden_cache_dtype": "float16",
"logit_batch_size": 8,
"disease_chunk_size": 2,
"min_cases": 1,
"min_controls": 1,
"max_ipcw_weight": 0.0,
"exclude_death_competing": True,
"probability_bins": np.asarray([0.0, 0.5, 1.0]),
}
with (
patch(
"evaluate_calibration.infer_landmark_hidden",
return_value=(
np.zeros((6, 4), dtype=np.float16),
row_arrays,
),
),
patch(
"evaluate_calibration.project_distribution_chunk",
return_value=(logits_chunk, None),
),
):
serial_metrics, serial_curve = evaluate_landmark_calibration(
**common,
num_workers_calibration=1,
)
parallel_metrics, parallel_curve = evaluate_landmark_calibration(
**common,
num_workers_calibration=2,
)
pd.testing.assert_frame_equal(serial_metrics, parallel_metrics)
pd.testing.assert_frame_equal(serial_curve, parallel_curve)
def test_metric_aggregation_uses_contribution_sums(self): def test_metric_aggregation_uses_contribution_sums(self):
metrics = pd.DataFrame( metrics = pd.DataFrame(
[ [