Report AUCs in Delphi2M format

This commit is contained in:
2026-07-25 11:15:06 +08:00
parent 8d0d71292e
commit f7d6cda8b6
3 changed files with 239 additions and 127 deletions

View File

@@ -7,7 +7,8 @@ This script follows the logic of the Delphi evaluation script supplied by the us
at least `offset` years before the target time;
3. run model inference by disease chunks to avoid materializing all logits;
4. compute AUC separately by sex and age bracket;
5. aggregate age brackets with DeLong variance.
5. average age-bracket AUCs within each sex and write a Delphi2M-style
Female/Male report.
Efficiency notes:
- transformer/readout inference is executed once and cached;
@@ -39,6 +40,10 @@ from torch.utils.data import DataLoader, Subset
from tqdm.auto import tqdm
from dataset import HealthDataset
from delphi2m_auc_report import (
DEFAULT_DELPHI2M_PERIODS_YEARS,
build_delphi2m_auc_report,
)
from eval_data import load_sequence_eval_dataset, sequence_eval_collate_fn
from models import (
DeepHealth,
@@ -1164,30 +1169,23 @@ def evaluate_auc_pipeline(
df_auc_unpooled["label_code"] = df_auc_unpooled["token"].map(
dataset.label_id_to_code)
print("Using DeLong method to calculate AUC confidence intervals.")
grouped = df_auc_unpooled.groupby(
["token", "label_code", "offset"], dropna=False, as_index=False)
df_auc = grouped.agg(
auc=("auc_delong", "mean"),
n_strata=("auc_delong", "size"),
n_diseased=("n_diseased", "sum"),
n_healthy=("n_healthy", "sum"),
auc_variance_sum=("auc_variance_delong", "sum"),
print(
"Building Delphi2M-style report: mean AUC across age strata, "
"reported separately for Female and Male."
)
df_auc["auc_variance_delong"] = (
df_auc["auc_variance_sum"]
/ (df_auc["n_strata"].clip(lower=1).astype(np.float64) ** 2)
df_report = build_delphi2m_auc_report(
df_auc_unpooled,
period_col="offset",
)
df_auc = df_auc.drop(columns=["auc_variance_sum"])
if output_path is not None:
out_dir = Path(output_path)
out_dir.mkdir(parents=True, exist_ok=True)
df_auc.to_csv(out_dir / "df_both.csv", index=False)
df_auc_unpooled.to_csv(
out_dir / "df_auc_unpooled.csv", index=False)
report_path = out_dir / "df_auc_delphi2m_report.csv"
df_report.to_csv(report_path, index=False)
print(f"Saved Delphi2M-style AUC report: {report_path}")
return df_auc_unpooled, df_auc
return df_auc_unpooled, df_report
# ---------------------------------------------------------------------------
@@ -1237,8 +1235,18 @@ def make_auc_offsets(args: argparse.Namespace, cfg: Dict[str, Any]) -> List[floa
if explicit_offsets is not None:
base_offsets = explicit_offsets
else:
next_token_offset = float(cfg_get(args, cfg, "offset", 0.1))
base_offsets = [next_token_offset, 1.0, 5.0, 10.0]
next_token_offset = float(
cfg_get(
args,
cfg,
"offset",
DEFAULT_DELPHI2M_PERIODS_YEARS[0],
)
)
base_offsets = [
next_token_offset,
*DEFAULT_DELPHI2M_PERIODS_YEARS[1:],
]
offsets: List[float] = []
seen = set()
@@ -1286,9 +1294,9 @@ def main() -> None:
parser.add_argument("--filter_min_total", type=int, default=None,
help="Minimum metadata count for disease selection; default 0.")
parser.add_argument("--offset", type=float, default=None,
help="Next-token prediction offset in years; preserved and evaluated alongside 1, 5, and 10 years by default.")
help="Next-token prediction offset in years; 0.1 is Delphi2M no gap and is evaluated alongside 1, 5, and 10 years by default.")
parser.add_argument("--offsets", type=str, default=None,
help="Comma-separated prediction offsets in years. Overrides the default set of offset,1,5,10.")
help="Comma-separated prediction offsets in years. Overrides the default set of 0.1,1,5,10.")
parser.add_argument("--age_start", type=float, default=None)
parser.add_argument("--age_stop", type=float, default=None)
parser.add_argument("--age_step", type=float, default=None)