Report AUCs in Delphi2M format

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

View File

@@ -3,6 +3,9 @@
This script supports DeepHealth fixed-horizon risk scores for exponential,
Weibull, and mixed all-future distributions.
The default horizons are 0.1, 1, 5, and 10 years. As in Delphi2M, 0.1 years
is reported as the no-gap evaluation.
Landmark querying depends on the model target mode saved in train_config.json:
- next_token: insert a <NO_EVENT> token at landmark age and read it out;
- all_future: pass landmark age directly as t_query.
@@ -28,6 +31,10 @@ from torch.utils.data import DataLoader, Dataset
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
from models import DeepHealth
from readouts import build_readout
@@ -324,44 +331,6 @@ def _first_existing_column(df: pd.DataFrame, candidates: Sequence[str]) -> Optio
return None
def build_metadata_for_merge(dataset: HealthDataset, labels_meta: Optional[pd.DataFrame]) -> pd.DataFrame:
base_rows = []
for token, code in dataset.label_id_to_code.items():
token = int(token)
code_text = str(code)
if token in SPECIAL_TOKENS or code_text.startswith("<"):
continue
base_rows.append({"token": token, "label_code": code_text})
base = pd.DataFrame(base_rows)
if labels_meta is None or labels_meta.empty:
return base
meta = labels_meta.copy()
code_col = _first_existing_column(
meta, ["Name", "code", "ICD10", "icd10", "label", "token", "disease_code"])
if code_col is not None:
meta["_label_code"] = meta[code_col].astype(
str).map(lambda s: s.split()[0].strip())
merged = base.merge(meta, left_on="label_code",
right_on="_label_code", how="left")
return merged.drop(columns=["_label_code"], errors="ignore")
if "index" in meta.columns:
idx = pd.to_numeric(meta["index"], errors="coerce")
has_no_event = (
NO_EVENT_IDX in dataset.label_id_to_code
and dataset.label_id_to_code.get(NO_EVENT_IDX) == "<NO_EVENT>"
)
if has_no_event:
idx = idx.where(idx < NO_EVENT_IDX, idx + 1)
meta["_index_int"] = idx.astype("Int64")
merged = base.merge(meta, left_on="token",
right_on="_index_int", how="left")
return merged.drop(columns=["_index_int"], errors="ignore")
return base
def _metadata_count_map(dataset: HealthDataset, labels_meta: Optional[pd.DataFrame]) -> Dict[int, float]:
if labels_meta is None or labels_meta.empty or "count" not in labels_meta.columns:
return {}
@@ -1101,7 +1070,6 @@ def evaluate_landmark_auc(
loader: DataLoader,
landmark_dataset: LandmarkDataset,
output_path: Path,
labels_meta: Optional[pd.DataFrame],
disease_ids: Sequence[int],
disease_chunk_size: int,
score_mode: str,
@@ -1118,7 +1086,6 @@ def evaluate_landmark_auc(
use_amp: bool,
hidden_cache_dtype: str,
logit_batch_size: int,
meta_info: Dict[str, Any],
) -> Tuple[pd.DataFrame, pd.DataFrame]:
model.eval().to(device)
@@ -1235,54 +1202,21 @@ def evaluate_landmark_auc(
df_unpooled["label_code"] = df_unpooled["token"].map(
landmark_dataset.dataset.label_id_to_code)
for k, v in meta_info.items():
df_unpooled[k] = v
meta_table = build_metadata_for_merge(landmark_dataset.dataset, labels_meta)
df_unpooled = df_unpooled.merge(
meta_table, on=["token", "label_code"], how="left")
grouped = df_unpooled.groupby(
["token", "label_code", "horizon"], dropna=False, as_index=False)
df_merged = 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 landmark-age "
"strata, reported separately for Female and Male."
)
df_merged["auc_variance_delong"] = (
df_merged["auc_variance_sum"]
/ (df_merged["n_strata"].clip(lower=1).astype(np.float64) ** 2)
df_report = build_delphi2m_auc_report(
df_unpooled,
period_col="horizon",
)
df_merged = df_merged.drop(columns=["auc_variance_sum"])
keep_meta = [
c for c in [
"model_ckpt_path",
"config_path",
"target_mode",
"model_target_mode",
"dist_mode",
"time_mode",
"attn_mask_mode",
"readout_name",
"landmark_query_mode",
"landmark_token_mode",
"score_mode",
"eval_split",
]
if c in df_unpooled.columns
]
for col in keep_meta:
df_merged[col] = meta_info[col]
output_path.mkdir(parents=True, exist_ok=True)
df_unpooled.to_csv(
output_path / "df_auc_landmark_unpooled.csv", index=False)
df_merged.to_csv(output_path / "df_auc_landmark.csv", index=False)
report_path = output_path / "df_auc_landmark_delphi2m_report.csv"
df_report.to_csv(report_path, index=False)
print(f"Saved Delphi2M-style landmark AUC report: {report_path}")
return df_unpooled, df_merged
return df_unpooled, df_report
def main() -> None:
@@ -1308,7 +1242,12 @@ def main() -> 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", type=str, default=None)
parser.add_argument(
"--horizons",
type=str,
default=None,
help="Comma-separated horizons in years; defaults to 0.1,1,5,10, where 0.1 is Delphi2M no gap.",
)
parser.add_argument("--min_cases", type=int, default=None)
parser.add_argument("--min_history_events", type=int, default=None)
@@ -1428,8 +1367,9 @@ def main() -> None:
"Landmark ages are empty. Check landmark_start/landmark_stop/landmark_step.")
horizons = np.asarray(
parse_float_list(cfg_get(args, cfg, "horizons", "1,5,10")) or [
1.0, 5.0, 10.0],
parse_float_list(
cfg_get(args, cfg, "horizons", "0.1,1,5,10")
) or list(DEFAULT_DELPHI2M_PERIODS_YEARS),
dtype=np.float32,
)
if horizons.size == 0:
@@ -1520,8 +1460,6 @@ def main() -> None:
if model_target_mode == "next_token"
else "direct_t_query"
)
score_mode_out = f"{landmark_query_mode}_{score_mode}"
num_workers_auc = int(
cfg_get(args, cfg, "num_workers_auc", max(1, (os.cpu_count() or 2) - 1)))
auc_task_chunk_size = int(cfg_get(args, cfg, "auc_task_chunk_size", 0))
@@ -1553,27 +1491,11 @@ def main() -> None:
print(f"AUC workers: {num_workers_auc}")
print(f"Output path: {output_path}")
meta_info = {
"score_mode": score_mode_out,
"eval_split": eval_split,
"model_ckpt_path": str(model_ckpt_path),
"config_path": str(config_path),
"target_mode": str(target_mode),
"model_target_mode": str(model_target_mode),
"dist_mode": str(dist_mode),
"time_mode": str(time_mode),
"attn_mask_mode": str(attn_mask_mode),
"readout_name": str(readout_name),
"landmark_query_mode": landmark_query_mode,
"landmark_token_mode": "no_event" if model_target_mode == "next_token" else "none",
}
evaluate_landmark_auc(
model=model,
loader=loader,
landmark_dataset=landmark_dataset,
output_path=output_path,
labels_meta=labels_meta,
disease_ids=disease_ids,
disease_chunk_size=disease_chunk_size,
score_mode=score_mode,
@@ -1590,7 +1512,6 @@ def main() -> None:
use_amp=use_amp,
hidden_cache_dtype=hidden_cache_dtype,
logit_batch_size=logit_batch_size,
meta_info=meta_info,
)