Report AUCs in Delphi2M format
This commit is contained in:
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user