Report AUCs in Delphi2M format
This commit is contained in:
@@ -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
|
||||
from readouts import build_readout
|
||||
@@ -1158,30 +1163,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
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -1231,8 +1229,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()
|
||||
@@ -1280,9 +1288,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)
|
||||
|
||||
Reference in New Issue
Block a user