Remove obsolete evaluation and batch scripts
This commit is contained in:
@@ -1,816 +0,0 @@
|
|||||||
"""Compute landmark future death and incident system-disease risks.
|
|
||||||
|
|
||||||
For each selected patient and landmark age, this script computes:
|
|
||||||
|
|
||||||
* future death risk within tau years;
|
|
||||||
* future incident disease risk for each ICD-10 chapter-derived system;
|
|
||||||
* model attribution of each historical organ/system disease set to predicted
|
|
||||||
mortality risk, computed by deleting that system's historical disease tokens
|
|
||||||
and re-querying the model;
|
|
||||||
* historical modeled-disease count;
|
|
||||||
* historical modeled-disease count within each ICD-10 chapter-derived system.
|
|
||||||
|
|
||||||
Death is always token vocab_size - 1. Disease groups are read from
|
|
||||||
icd10_chapter_organ_mapping.csv.
|
|
||||||
"""
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import argparse
|
|
||||||
import json
|
|
||||||
from pathlib import Path
|
|
||||||
from typing import Any, Dict, List, Optional, Sequence
|
|
||||||
|
|
||||||
import numpy as np
|
|
||||||
import pandas as pd
|
|
||||||
import torch
|
|
||||||
from torch.nn.utils.rnn import pad_sequence
|
|
||||||
from torch.utils.data import DataLoader, Dataset
|
|
||||||
from tqdm.auto import tqdm
|
|
||||||
|
|
||||||
from dataset import HealthDataset
|
|
||||||
from eval_data import load_sequence_eval_dataset
|
|
||||||
from evaluate_auc_v2 import (
|
|
||||||
LandmarkDataset,
|
|
||||||
build_model_from_dataset,
|
|
||||||
cfg_get,
|
|
||||||
load_checkpoint_state_dict,
|
|
||||||
load_json_config,
|
|
||||||
load_model_state,
|
|
||||||
make_eval_indices,
|
|
||||||
resolve_dist_mode_for_checkpoint,
|
|
||||||
resolve_eval_device,
|
|
||||||
validate_dataset_metadata,
|
|
||||||
)
|
|
||||||
from future_risk import (
|
|
||||||
death_risk_from_probabilities,
|
|
||||||
new_disease_risk_from_probabilities,
|
|
||||||
probabilities_from_logits,
|
|
||||||
)
|
|
||||||
from models import DeepHealth
|
|
||||||
from readouts import build_readout
|
|
||||||
from targets import CHECKUP_IDX, NO_EVENT_IDX, PAD_IDX
|
|
||||||
from train_util import load_eid_file, load_extra_info_types_file
|
|
||||||
|
|
||||||
|
|
||||||
SPECIAL_TOKENS = {PAD_IDX, CHECKUP_IDX, NO_EVENT_IDX}
|
|
||||||
|
|
||||||
|
|
||||||
def parse_int_list(value: Any) -> Optional[List[int]]:
|
|
||||||
if value is None:
|
|
||||||
return None
|
|
||||||
if isinstance(value, (list, tuple, np.ndarray)):
|
|
||||||
return [int(x) for x in value]
|
|
||||||
text = str(value).strip()
|
|
||||||
if text == "":
|
|
||||||
return None
|
|
||||||
if text.startswith("["):
|
|
||||||
values = json.loads(text)
|
|
||||||
if not isinstance(values, list):
|
|
||||||
raise ValueError(
|
|
||||||
f"Expected a JSON list, got {type(values).__name__}")
|
|
||||||
return [int(x) for x in values]
|
|
||||||
return [int(x.strip()) for x in text.split(",") if x.strip()]
|
|
||||||
|
|
||||||
|
|
||||||
def load_extra_info_types(value: Any) -> Optional[List[int]]:
|
|
||||||
if value is None:
|
|
||||||
return None
|
|
||||||
text = str(value)
|
|
||||||
path = Path(text)
|
|
||||||
if path.exists():
|
|
||||||
return load_extra_info_types_file(text)
|
|
||||||
return parse_int_list(value)
|
|
||||||
|
|
||||||
|
|
||||||
def make_landmark_ages(start: float, stop: float, step: float) -> np.ndarray:
|
|
||||||
if step <= 0:
|
|
||||||
raise ValueError("landmark_step must be positive")
|
|
||||||
if stop < start:
|
|
||||||
raise ValueError("landmark_stop must be >= landmark_start")
|
|
||||||
# Include stop when it lands on the grid, e.g. 40,45,...,80.
|
|
||||||
return np.arange(start, stop + step * 0.5, step, dtype=np.float32)
|
|
||||||
|
|
||||||
|
|
||||||
def build_first_occurrence_maps_for_landmarks(
|
|
||||||
dataset: HealthDataset,
|
|
||||||
subset_indices: np.ndarray,
|
|
||||||
) -> Dict[int, tuple[np.ndarray, np.ndarray]]:
|
|
||||||
first_lists: Dict[int, list[tuple[int, float]]] = {}
|
|
||||||
for patient_id, dataset_index in enumerate(np.asarray(subset_indices, dtype=np.int64).tolist()):
|
|
||||||
s = dataset.samples[int(dataset_index)]
|
|
||||||
seq_event = np.asarray(s["event_seq"], dtype=np.int64)
|
|
||||||
seq_time = np.asarray(s["time_seq"], dtype=np.float32)
|
|
||||||
tgt_event = np.asarray(s["target_event_seq"], dtype=np.int64)
|
|
||||||
tgt_time = np.asarray(s["target_time_seq"], dtype=np.float32)
|
|
||||||
if seq_event.size == 0 or tgt_event.size == 0:
|
|
||||||
continue
|
|
||||||
|
|
||||||
full_event = np.concatenate([seq_event, tgt_event[-1:]])
|
|
||||||
full_time = np.concatenate([seq_time, tgt_time[-1:]])
|
|
||||||
uniq_tokens, first_idx = np.unique(full_event, return_index=True)
|
|
||||||
for token, idx in zip(uniq_tokens.tolist(), first_idx.tolist()):
|
|
||||||
token = int(token)
|
|
||||||
if token in SPECIAL_TOKENS:
|
|
||||||
continue
|
|
||||||
first_lists.setdefault(token, []).append(
|
|
||||||
(patient_id, float(full_time[int(idx)])))
|
|
||||||
|
|
||||||
return {
|
|
||||||
int(token): (
|
|
||||||
np.asarray([p for p, _ in pairs], dtype=np.int32),
|
|
||||||
np.asarray([t for _, t in pairs], dtype=np.float32),
|
|
||||||
)
|
|
||||||
for token, pairs in first_lists.items()
|
|
||||||
if pairs
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def normalize_eval_split(args: argparse.Namespace, cfg: Dict[str, Any]) -> str:
|
|
||||||
eval_split = str(cfg_get(args, cfg, "eval_split", "test")).lower()
|
|
||||||
if eval_split in {"valid", "validation"}:
|
|
||||||
return "val"
|
|
||||||
if eval_split not in {"train", "val", "test", "all"}:
|
|
||||||
raise ValueError(f"Unsupported eval_split={eval_split!r}")
|
|
||||||
return eval_split
|
|
||||||
|
|
||||||
|
|
||||||
def load_eval_sequence_dataset(
|
|
||||||
args: argparse.Namespace,
|
|
||||||
cfg: Dict[str, Any],
|
|
||||||
) -> tuple[Any, np.ndarray, str, str]:
|
|
||||||
eval_split = normalize_eval_split(args, cfg)
|
|
||||||
model_target_mode = str(cfg.get("model_target_mode", "next_token")).lower()
|
|
||||||
data_prefix = str(cfg.get("data_prefix", "ukb"))
|
|
||||||
labels_file = str(cfg.get("labels_file", "labels.csv"))
|
|
||||||
no_event_interval_years = float(cfg.get("no_event_interval_years", 5.0))
|
|
||||||
include_no_event_in_uts_target = bool(
|
|
||||||
cfg.get("include_no_event_in_uts_target", False))
|
|
||||||
extra_info_types = load_extra_info_types(args.extra_info_types)
|
|
||||||
if extra_info_types is None:
|
|
||||||
extra_info_types = parse_int_list(cfg.get("extra_info_types", None))
|
|
||||||
|
|
||||||
print("Loading one sequence eval dataset...")
|
|
||||||
dataset = load_sequence_eval_dataset(
|
|
||||||
model_target_mode=model_target_mode,
|
|
||||||
data_prefix=data_prefix,
|
|
||||||
labels_file=labels_file,
|
|
||||||
no_event_interval_years=no_event_interval_years,
|
|
||||||
include_no_event_in_uts_target=include_no_event_in_uts_target,
|
|
||||||
min_history_events=int(cfg.get("all_future_min_history_events", 1)),
|
|
||||||
min_future_events=int(cfg.get("all_future_min_future_events", 1)),
|
|
||||||
extra_info_types=extra_info_types,
|
|
||||||
)
|
|
||||||
|
|
||||||
train_eid_file = cfg_get(args, cfg, "train_eid_file", "ukb_train_eid.csv")
|
|
||||||
val_eid_file = cfg_get(args, cfg, "val_eid_file", "ukb_val_eid.csv")
|
|
||||||
test_eid_file = cfg_get(args, cfg, "test_eid_file", "ukb_test_eid.csv")
|
|
||||||
split_files_exist = all(
|
|
||||||
Path(str(path)).exists()
|
|
||||||
for path in (train_eid_file, val_eid_file, test_eid_file)
|
|
||||||
)
|
|
||||||
|
|
||||||
if eval_split != "all" and split_files_exist:
|
|
||||||
split_files = {
|
|
||||||
"train": train_eid_file,
|
|
||||||
"val": val_eid_file,
|
|
||||||
"test": test_eid_file,
|
|
||||||
}
|
|
||||||
selected_eids = load_eid_file(split_files[eval_split])
|
|
||||||
out = np.asarray(
|
|
||||||
[
|
|
||||||
idx
|
|
||||||
for idx, sample in enumerate(dataset.samples)
|
|
||||||
if int(sample["eid"]) in selected_eids
|
|
||||||
],
|
|
||||||
dtype=np.int64,
|
|
||||||
)
|
|
||||||
if out.size == 0:
|
|
||||||
raise ValueError(
|
|
||||||
f"No samples found for eval_split={eval_split!r} using {split_files[eval_split]}"
|
|
||||||
)
|
|
||||||
split_source = "eid_files"
|
|
||||||
else:
|
|
||||||
if eval_split == "all":
|
|
||||||
out = np.arange(len(dataset.samples), dtype=np.int64)
|
|
||||||
split_source = "all"
|
|
||||||
else:
|
|
||||||
out = make_eval_indices(dataset, args, cfg)
|
|
||||||
split_source = "ratio_split"
|
|
||||||
|
|
||||||
subset_size = cfg_get(args, cfg, "dataset_subset_size", None)
|
|
||||||
if subset_size is not None and int(subset_size) > 0:
|
|
||||||
out = out[: int(subset_size)]
|
|
||||||
return dataset, np.asarray(out, dtype=np.int64), eval_split, split_source
|
|
||||||
|
|
||||||
|
|
||||||
def load_organ_groups(
|
|
||||||
path: Path,
|
|
||||||
*,
|
|
||||||
vocab_size: int,
|
|
||||||
) -> tuple[dict[str, list[int]], dict[str, str], dict[int, str]]:
|
|
||||||
table = pd.read_csv(path)
|
|
||||||
required = {"token_id", "organ_system", "organ_system_label", "is_death"}
|
|
||||||
missing = required - set(table.columns)
|
|
||||||
if missing:
|
|
||||||
raise ValueError(f"{path} is missing columns: {sorted(missing)}")
|
|
||||||
|
|
||||||
death_idx = int(vocab_size) - 1
|
|
||||||
groups: dict[str, list[int]] = {}
|
|
||||||
labels: dict[str, str] = {}
|
|
||||||
token_to_group: dict[int, str] = {}
|
|
||||||
for row in table.itertuples(index=False):
|
|
||||||
token = int(getattr(row, "token_id"))
|
|
||||||
if token in SPECIAL_TOKENS or token == death_idx:
|
|
||||||
continue
|
|
||||||
if token < 0 or token >= int(vocab_size):
|
|
||||||
continue
|
|
||||||
if int(getattr(row, "is_death")) == 1:
|
|
||||||
continue
|
|
||||||
group = str(getattr(row, "organ_system"))
|
|
||||||
label = str(getattr(row, "organ_system_label"))
|
|
||||||
groups.setdefault(group, []).append(token)
|
|
||||||
labels[group] = label
|
|
||||||
token_to_group[token] = group
|
|
||||||
|
|
||||||
groups = {k: sorted(set(v)) for k, v in groups.items() if v}
|
|
||||||
return groups, labels, token_to_group
|
|
||||||
|
|
||||||
|
|
||||||
class IndexedLandmarkDataset(Dataset):
|
|
||||||
def __init__(self, base: LandmarkDataset) -> None:
|
|
||||||
self.base = base
|
|
||||||
|
|
||||||
def __len__(self) -> int:
|
|
||||||
return len(self.base)
|
|
||||||
|
|
||||||
def __getitem__(self, idx: int) -> Dict[str, torch.Tensor]:
|
|
||||||
item = dict(self.base[idx])
|
|
||||||
item["row_idx"] = torch.tensor(int(idx), dtype=torch.long)
|
|
||||||
return item
|
|
||||||
|
|
||||||
|
|
||||||
def collate_indexed_landmark_fn(batch: List[Dict[str, torch.Tensor]]) -> Dict[str, torch.Tensor]:
|
|
||||||
event_seq = pad_sequence(
|
|
||||||
[x["event_seq"] for x in batch], batch_first=True, padding_value=PAD_IDX
|
|
||||||
)
|
|
||||||
time_seq = pad_sequence(
|
|
||||||
[x["time_seq"] for x in batch], batch_first=True, padding_value=0.0
|
|
||||||
)
|
|
||||||
readout_mask = pad_sequence(
|
|
||||||
[x["readout_mask"] for x in batch], batch_first=True, padding_value=False
|
|
||||||
)
|
|
||||||
other_type = pad_sequence(
|
|
||||||
[x["other_type"] for x in batch], batch_first=True, padding_value=0
|
|
||||||
)
|
|
||||||
other_value = pad_sequence(
|
|
||||||
[x["other_value"] for x in batch], batch_first=True, padding_value=0.0
|
|
||||||
)
|
|
||||||
other_value_kind = pad_sequence(
|
|
||||||
[x["other_value_kind"] for x in batch], batch_first=True, padding_value=0
|
|
||||||
)
|
|
||||||
other_time = pad_sequence(
|
|
||||||
[x["other_time"] for x in batch], batch_first=True, padding_value=0.0
|
|
||||||
)
|
|
||||||
return {
|
|
||||||
"event_seq": event_seq,
|
|
||||||
"time_seq": time_seq,
|
|
||||||
"padding_mask": event_seq > PAD_IDX,
|
|
||||||
"readout_mask": readout_mask,
|
|
||||||
"sex": torch.stack([x["sex"] for x in batch]),
|
|
||||||
"other_type": other_type,
|
|
||||||
"other_value": other_value,
|
|
||||||
"other_value_kind": other_value_kind,
|
|
||||||
"other_time": other_time,
|
|
||||||
"landmark_pos": torch.stack([x["landmark_pos"] for x in batch]),
|
|
||||||
"t_query": torch.stack([x["t_query"] for x in batch]),
|
|
||||||
"patient_id": torch.stack([x["patient_id"] for x in batch]),
|
|
||||||
"landmark_age": torch.stack([x["landmark_age"] for x in batch]),
|
|
||||||
"followup_end_time": torch.stack([x["followup_end_time"] for x in batch]),
|
|
||||||
"death_time": torch.stack([x["death_time"] for x in batch]),
|
|
||||||
"row_idx": torch.stack([x["row_idx"] for x in batch]),
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def build_group_ablated_slice(
|
|
||||||
batch: Dict[str, torch.Tensor],
|
|
||||||
token_ids: Sequence[int],
|
|
||||||
row_indices: torch.Tensor,
|
|
||||||
) -> Dict[str, torch.Tensor]:
|
|
||||||
"""Build one fixed-width ablated slice without rebuilding variable-length rows."""
|
|
||||||
event_seq = batch["event_seq"]
|
|
||||||
|
|
||||||
out: Dict[str, torch.Tensor] = {}
|
|
||||||
out["event_seq"] = event_seq[row_indices].clone()
|
|
||||||
out["time_seq"] = batch["time_seq"][row_indices]
|
|
||||||
out["readout_mask"] = batch["readout_mask"][row_indices].clone()
|
|
||||||
out["padding_mask"] = batch["padding_mask"][row_indices].bool().clone()
|
|
||||||
out["landmark_pos"] = batch["landmark_pos"][row_indices].clone()
|
|
||||||
|
|
||||||
seq_len = int(event_seq.shape[1])
|
|
||||||
positions = torch.arange(seq_len, device=event_seq.device)[None, :]
|
|
||||||
ids = torch.as_tensor(token_ids, dtype=event_seq.dtype,
|
|
||||||
device=event_seq.device)
|
|
||||||
remove = torch.isin(out["event_seq"], ids) & out["padding_mask"]
|
|
||||||
out["event_seq"] = torch.where(
|
|
||||||
remove,
|
|
||||||
torch.full_like(out["event_seq"], PAD_IDX),
|
|
||||||
out["event_seq"],
|
|
||||||
)
|
|
||||||
out["padding_mask"] &= ~remove
|
|
||||||
out["readout_mask"] &= ~remove
|
|
||||||
|
|
||||||
has_valid = out["padding_mask"].any(dim=1)
|
|
||||||
if not bool(has_valid.all().item()):
|
|
||||||
empty_rows = torch.nonzero(~has_valid, as_tuple=False).flatten()
|
|
||||||
out["event_seq"][empty_rows, 0] = CHECKUP_IDX
|
|
||||||
out["time_seq"][empty_rows, 0] = batch["t_query"][row_indices[empty_rows]].to(
|
|
||||||
dtype=out["time_seq"].dtype
|
|
||||||
)
|
|
||||||
out["padding_mask"][empty_rows, 0] = True
|
|
||||||
out["readout_mask"][empty_rows, 0] = True
|
|
||||||
out["landmark_pos"][empty_rows] = 0
|
|
||||||
|
|
||||||
has_readout = out["readout_mask"].any(dim=1)
|
|
||||||
if not bool(has_readout.all().item()):
|
|
||||||
rows = torch.nonzero(~has_readout, as_tuple=False).flatten()
|
|
||||||
local_valid = out["padding_mask"][rows]
|
|
||||||
last_pos = torch.where(
|
|
||||||
local_valid,
|
|
||||||
positions.expand(local_valid.shape[0], -1),
|
|
||||||
torch.zeros_like(positions.expand(local_valid.shape[0], -1)),
|
|
||||||
).amax(dim=1)
|
|
||||||
out["readout_mask"][rows] = False
|
|
||||||
out["readout_mask"][rows, last_pos] = True
|
|
||||||
out["landmark_pos"][rows] = last_pos.to(dtype=out["landmark_pos"].dtype)
|
|
||||||
|
|
||||||
repeated_keys = (
|
|
||||||
"sex",
|
|
||||||
"other_type",
|
|
||||||
"other_value",
|
|
||||||
"other_value_kind",
|
|
||||||
"other_time",
|
|
||||||
"t_query",
|
|
||||||
"patient_id",
|
|
||||||
"landmark_age",
|
|
||||||
"followup_end_time",
|
|
||||||
"death_time",
|
|
||||||
"row_idx",
|
|
||||||
)
|
|
||||||
for key in repeated_keys:
|
|
||||||
out[key] = batch[key][row_indices]
|
|
||||||
return out
|
|
||||||
|
|
||||||
|
|
||||||
def concat_tensor_batches(chunks: Sequence[Dict[str, torch.Tensor]]) -> Dict[str, torch.Tensor]:
|
|
||||||
return {
|
|
||||||
key: torch.cat([chunk[key] for chunk in chunks], dim=0)
|
|
||||||
for key in chunks[0]
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def iter_group_ablated_batches(
|
|
||||||
batch: Dict[str, torch.Tensor],
|
|
||||||
group_names: Sequence[str],
|
|
||||||
organ_groups: dict[str, list[int]],
|
|
||||||
occurred: torch.Tensor,
|
|
||||||
max_batch_size: int,
|
|
||||||
):
|
|
||||||
"""Yield ablated chunks as soon as enough rows are available for a forward pass."""
|
|
||||||
pending_batches: list[Dict[str, torch.Tensor]] = []
|
|
||||||
pending_groups: list[str] = []
|
|
||||||
pending_rows: list[int] = []
|
|
||||||
pending_n = 0
|
|
||||||
|
|
||||||
for group in group_names:
|
|
||||||
ids = torch.as_tensor(
|
|
||||||
organ_groups[group], dtype=torch.long, device=occurred.device)
|
|
||||||
if ids.numel() == 0:
|
|
||||||
continue
|
|
||||||
active_rows = torch.nonzero(
|
|
||||||
occurred[:, ids].any(dim=1), as_tuple=False).flatten()
|
|
||||||
if active_rows.numel() == 0:
|
|
||||||
continue
|
|
||||||
|
|
||||||
row_offset = 0
|
|
||||||
while row_offset < int(active_rows.numel()):
|
|
||||||
capacity = int(max_batch_size) - pending_n
|
|
||||||
row_stop = min(int(active_rows.numel()), row_offset + capacity)
|
|
||||||
row_indices = active_rows[row_offset:row_stop].to(
|
|
||||||
device=batch["event_seq"].device)
|
|
||||||
chunk = build_group_ablated_slice(
|
|
||||||
batch=batch,
|
|
||||||
token_ids=organ_groups[group],
|
|
||||||
row_indices=row_indices,
|
|
||||||
)
|
|
||||||
chunk_n = int(row_indices.numel())
|
|
||||||
pending_batches.append(chunk)
|
|
||||||
pending_groups.extend([group] * chunk_n)
|
|
||||||
pending_rows.extend(int(x)
|
|
||||||
for x in row_indices.detach().cpu().tolist())
|
|
||||||
pending_n += chunk_n
|
|
||||||
row_offset = row_stop
|
|
||||||
|
|
||||||
if pending_n >= int(max_batch_size):
|
|
||||||
yield concat_tensor_batches(pending_batches), pending_groups, pending_rows
|
|
||||||
pending_batches = []
|
|
||||||
pending_groups = []
|
|
||||||
pending_rows = []
|
|
||||||
pending_n = 0
|
|
||||||
|
|
||||||
if pending_batches:
|
|
||||||
yield concat_tensor_batches(pending_batches), pending_groups, pending_rows
|
|
||||||
|
|
||||||
|
|
||||||
@torch.no_grad()
|
|
||||||
def infer_landmark_hidden(
|
|
||||||
*,
|
|
||||||
model: DeepHealth,
|
|
||||||
batch: Dict[str, torch.Tensor],
|
|
||||||
device: torch.device,
|
|
||||||
model_target_mode: str,
|
|
||||||
readout_name: str,
|
|
||||||
readout_reduce: str,
|
|
||||||
) -> torch.Tensor:
|
|
||||||
batch_dev = {
|
|
||||||
k: (v.to(device, non_blocking=True) if isinstance(v, torch.Tensor) else v)
|
|
||||||
for k, v in batch.items()
|
|
||||||
}
|
|
||||||
if model_target_mode == "all_future":
|
|
||||||
return model(
|
|
||||||
event_seq=batch_dev["event_seq"].long(),
|
|
||||||
time_seq=batch_dev["time_seq"].float(),
|
|
||||||
sex=batch_dev["sex"].long(),
|
|
||||||
padding_mask=batch_dev["padding_mask"].bool(),
|
|
||||||
t_query=batch_dev["t_query"].float(),
|
|
||||||
other_type=batch_dev["other_type"].long(),
|
|
||||||
other_value=batch_dev["other_value"].float(),
|
|
||||||
other_value_kind=batch_dev["other_value_kind"].long(),
|
|
||||||
other_time=batch_dev["other_time"].float(),
|
|
||||||
target_mode="all_future",
|
|
||||||
)
|
|
||||||
|
|
||||||
hidden = model(
|
|
||||||
event_seq=batch_dev["event_seq"].long(),
|
|
||||||
time_seq=batch_dev["time_seq"].float(),
|
|
||||||
sex=batch_dev["sex"].long(),
|
|
||||||
padding_mask=batch_dev["padding_mask"].bool(),
|
|
||||||
other_type=batch_dev["other_type"].long(),
|
|
||||||
other_value=batch_dev["other_value"].float(),
|
|
||||||
other_value_kind=batch_dev["other_value_kind"].long(),
|
|
||||||
other_time=batch_dev["other_time"].float(),
|
|
||||||
target_mode="next_token",
|
|
||||||
)
|
|
||||||
readout = build_readout(readout_name, reduce=readout_reduce)
|
|
||||||
readout_out = readout(
|
|
||||||
hidden=hidden,
|
|
||||||
time_seq=batch_dev["time_seq"].float(),
|
|
||||||
padding_mask=batch_dev["padding_mask"].bool(),
|
|
||||||
readout_mask=batch_dev["readout_mask"].bool(),
|
|
||||||
)
|
|
||||||
return readout_out.hidden.gather(
|
|
||||||
1,
|
|
||||||
batch_dev["landmark_pos"].long()[:, None, None].expand(
|
|
||||||
-1, 1, readout_out.hidden.shape[-1]
|
|
||||||
),
|
|
||||||
).squeeze(1)
|
|
||||||
|
|
||||||
|
|
||||||
def make_occurred_mask(
|
|
||||||
event_seq: torch.Tensor,
|
|
||||||
*,
|
|
||||||
vocab_size: int,
|
|
||||||
device: torch.device,
|
|
||||||
) -> torch.Tensor:
|
|
||||||
occurred = torch.zeros(event_seq.shape[0], int(
|
|
||||||
vocab_size), dtype=torch.bool, device=device)
|
|
||||||
valid = (event_seq >= 0) & (event_seq < int(vocab_size))
|
|
||||||
safe = event_seq.clamp(min=0, max=int(vocab_size) - 1).to(device)
|
|
||||||
occurred.scatter_(1, safe, valid.to(device))
|
|
||||||
return occurred
|
|
||||||
|
|
||||||
|
|
||||||
def mortality_hazard_from_risk(risk: torch.Tensor, eps: float = 1e-7) -> torch.Tensor:
|
|
||||||
return -torch.log1p(-risk.clamp(0.0, 1.0 - float(eps)))
|
|
||||||
|
|
||||||
|
|
||||||
def death_risk_for_batch(
|
|
||||||
*,
|
|
||||||
model: DeepHealth,
|
|
||||||
batch: Dict[str, torch.Tensor],
|
|
||||||
device: torch.device,
|
|
||||||
model_target_mode: str,
|
|
||||||
readout_name: str,
|
|
||||||
readout_reduce: str,
|
|
||||||
dist_mode: str,
|
|
||||||
tau: float,
|
|
||||||
) -> torch.Tensor:
|
|
||||||
hidden = infer_landmark_hidden(
|
|
||||||
model=model,
|
|
||||||
batch=batch,
|
|
||||||
device=device,
|
|
||||||
model_target_mode=model_target_mode,
|
|
||||||
readout_name=readout_name,
|
|
||||||
readout_reduce=readout_reduce,
|
|
||||||
)
|
|
||||||
logits = model.calc_risk(hidden)
|
|
||||||
rho = model.calc_weibull_rho(hidden) if dist_mode == "weibull" else None
|
|
||||||
death_rho = model.calc_death_rho(hidden) if dist_mode == "mixed" else None
|
|
||||||
probabilities = probabilities_from_logits(
|
|
||||||
logits,
|
|
||||||
tau,
|
|
||||||
dist_mode=dist_mode,
|
|
||||||
rho=rho,
|
|
||||||
death_rho=death_rho,
|
|
||||||
)
|
|
||||||
return death_risk_from_probabilities(probabilities)
|
|
||||||
|
|
||||||
|
|
||||||
def historical_counts_by_group(
|
|
||||||
tokens: np.ndarray,
|
|
||||||
*,
|
|
||||||
death_idx: int,
|
|
||||||
token_to_group: dict[int, str],
|
|
||||||
group_names: Sequence[str],
|
|
||||||
) -> tuple[int, dict[str, int]]:
|
|
||||||
unique_tokens = {
|
|
||||||
int(token)
|
|
||||||
for token in np.asarray(tokens, dtype=np.int64).tolist()
|
|
||||||
if int(token) not in SPECIAL_TOKENS and int(token) != int(death_idx)
|
|
||||||
}
|
|
||||||
total = len(unique_tokens)
|
|
||||||
out = {group: 0 for group in group_names}
|
|
||||||
for token in unique_tokens:
|
|
||||||
group = token_to_group.get(token)
|
|
||||||
if group in out:
|
|
||||||
out[group] += 1
|
|
||||||
return total, out
|
|
||||||
|
|
||||||
|
|
||||||
def output_name_for_run(run_path: Path, eval_split: str, tau: float) -> Path:
|
|
||||||
return run_path / f"future_risk_{eval_split}_tau{tau:g}y.csv"
|
|
||||||
|
|
||||||
|
|
||||||
def parse_args() -> argparse.Namespace:
|
|
||||||
parser = argparse.ArgumentParser(
|
|
||||||
description="Compute landmark death and incident system-disease risks."
|
|
||||||
)
|
|
||||||
parser.add_argument("--run_path", type=str, required=True)
|
|
||||||
parser.add_argument("--output_path", type=str, default=None)
|
|
||||||
parser.add_argument("--organ_mapping_path", type=str,
|
|
||||||
default="icd10_chapter_organ_mapping.csv")
|
|
||||||
parser.add_argument("--eval_split", type=str, default=None)
|
|
||||||
parser.add_argument("--dataset_subset_size", type=int, default=None)
|
|
||||||
parser.add_argument("--train_eid_file", type=str, default=None)
|
|
||||||
parser.add_argument("--val_eid_file", type=str, default=None)
|
|
||||||
parser.add_argument("--test_eid_file", type=str, default=None)
|
|
||||||
parser.add_argument("--landmark_start", type=float, default=40.0)
|
|
||||||
parser.add_argument("--landmark_stop", type=float, default=80.0)
|
|
||||||
parser.add_argument("--landmark_step", type=float, default=5.0)
|
|
||||||
parser.add_argument("--tau", type=float, default=5.0)
|
|
||||||
parser.add_argument("--min_history_events", type=int, default=None)
|
|
||||||
parser.add_argument("--batch_size", type=int, default=None)
|
|
||||||
parser.add_argument(
|
|
||||||
"--attribution_batch_size",
|
|
||||||
type=int,
|
|
||||||
default=None,
|
|
||||||
help="Forward batch size for expanded organ/system ablation queries.",
|
|
||||||
)
|
|
||||||
parser.add_argument("--num_workers", type=int, default=None)
|
|
||||||
parser.add_argument("--device", type=str, default=None)
|
|
||||||
parser.add_argument("--extra_info_types", type=str, default=None)
|
|
||||||
return parser.parse_args()
|
|
||||||
|
|
||||||
|
|
||||||
def main() -> None:
|
|
||||||
args = parse_args()
|
|
||||||
run_path = Path(args.run_path)
|
|
||||||
config_path = run_path / "train_config.json"
|
|
||||||
checkpoint_path = run_path / "best_model.pt"
|
|
||||||
if not config_path.exists():
|
|
||||||
raise FileNotFoundError(f"train_config.json not found: {config_path}")
|
|
||||||
if not checkpoint_path.exists():
|
|
||||||
raise FileNotFoundError(f"best_model.pt not found: {checkpoint_path}")
|
|
||||||
|
|
||||||
cfg = load_json_config(config_path)
|
|
||||||
model_target_mode = str(cfg.get("model_target_mode", "next_token")).lower()
|
|
||||||
if model_target_mode not in {"next_token", "all_future"}:
|
|
||||||
raise ValueError(
|
|
||||||
f"Unsupported model_target_mode: {model_target_mode!r}")
|
|
||||||
|
|
||||||
target_mode = str(cfg.get("target_mode", "uts"))
|
|
||||||
attn_mask_mode = str(
|
|
||||||
cfg.get("attn_mask_mode", "non_strict_time" if target_mode ==
|
|
||||||
"uts" else "target_aware")
|
|
||||||
)
|
|
||||||
readout_name = str(cfg.get(
|
|
||||||
"readout_name", "same_time_group_end" if target_mode == "uts" else "token"))
|
|
||||||
readout_reduce = str(cfg.get("readout_reduce", "mean"))
|
|
||||||
|
|
||||||
dataset, subset_indices, eval_split, split_source = load_eval_sequence_dataset(
|
|
||||||
args,
|
|
||||||
cfg,
|
|
||||||
)
|
|
||||||
validate_dataset_metadata(dataset, cfg)
|
|
||||||
|
|
||||||
landmark_ages = make_landmark_ages(
|
|
||||||
float(args.landmark_start),
|
|
||||||
float(args.landmark_stop),
|
|
||||||
float(args.landmark_step),
|
|
||||||
)
|
|
||||||
tau = float(args.tau)
|
|
||||||
if tau < 0:
|
|
||||||
raise ValueError("tau must be non-negative")
|
|
||||||
|
|
||||||
first_occurrence_by_token = build_first_occurrence_maps_for_landmarks(
|
|
||||||
dataset,
|
|
||||||
subset_indices,
|
|
||||||
)
|
|
||||||
death_idx = int(dataset.vocab_size) - 1
|
|
||||||
landmark_dataset = LandmarkDataset(
|
|
||||||
dataset=dataset,
|
|
||||||
subset_indices=subset_indices,
|
|
||||||
landmark_ages=landmark_ages,
|
|
||||||
attn_mask_mode=attn_mask_mode,
|
|
||||||
model_target_mode=model_target_mode,
|
|
||||||
min_history_events=int(cfg_get(args, cfg, "min_history_events", 1)),
|
|
||||||
first_occurrence_by_token=first_occurrence_by_token,
|
|
||||||
death_token_ids=[death_idx],
|
|
||||||
)
|
|
||||||
|
|
||||||
organ_groups, organ_labels, token_to_group = load_organ_groups(
|
|
||||||
Path(args.organ_mapping_path),
|
|
||||||
vocab_size=int(dataset.vocab_size),
|
|
||||||
)
|
|
||||||
group_names = sorted(organ_groups)
|
|
||||||
|
|
||||||
state_dict = load_checkpoint_state_dict(checkpoint_path, map_location="cpu")
|
|
||||||
dist_mode = resolve_dist_mode_for_checkpoint(
|
|
||||||
str(cfg.get("dist_mode", "exponential")), state_dict)
|
|
||||||
cfg_model = dict(cfg)
|
|
||||||
cfg_model["dist_mode"] = dist_mode
|
|
||||||
device = resolve_eval_device(args.device)
|
|
||||||
model = build_model_from_dataset(
|
|
||||||
args, cfg_model, dataset, state_dict=state_dict
|
|
||||||
).to(device)
|
|
||||||
load_model_state(model, state_dict)
|
|
||||||
model.eval()
|
|
||||||
|
|
||||||
batch_size = int(cfg_get(args, cfg, "batch_size", 128))
|
|
||||||
attribution_batch_size = int(
|
|
||||||
cfg_get(args, cfg, "attribution_batch_size",
|
|
||||||
max(batch_size * 4, batch_size))
|
|
||||||
)
|
|
||||||
if attribution_batch_size <= 0:
|
|
||||||
raise ValueError("attribution_batch_size must be positive")
|
|
||||||
num_workers = int(cfg_get(args, cfg, "num_workers", 4))
|
|
||||||
loader = DataLoader(
|
|
||||||
IndexedLandmarkDataset(landmark_dataset),
|
|
||||||
batch_size=batch_size,
|
|
||||||
shuffle=False,
|
|
||||||
collate_fn=collate_indexed_landmark_fn,
|
|
||||||
num_workers=num_workers,
|
|
||||||
pin_memory=device.type == "cuda",
|
|
||||||
persistent_workers=num_workers > 0,
|
|
||||||
prefetch_factor=2 if num_workers > 0 else None,
|
|
||||||
)
|
|
||||||
|
|
||||||
output_path = Path(args.output_path) if args.output_path else output_name_for_run(
|
|
||||||
run_path, eval_split, tau)
|
|
||||||
output_path.parent.mkdir(parents=True, exist_ok=True)
|
|
||||||
|
|
||||||
print(f"Eval split: {eval_split}")
|
|
||||||
print(f"Split source: {split_source}")
|
|
||||||
print(f"Selected patients: {len(subset_indices)}")
|
|
||||||
print(f"Landmark ages: {landmark_ages.tolist()}")
|
|
||||||
print(f"Tau: {tau:g} years")
|
|
||||||
print(f"Dist mode: {dist_mode}")
|
|
||||||
print(f"Device: {device}")
|
|
||||||
print(f"Death token: {death_idx}")
|
|
||||||
print(f"Organ/system groups: {len(group_names)}")
|
|
||||||
print(f"Landmark rows: {len(landmark_dataset)}")
|
|
||||||
print(f"Attribution batch size: {attribution_batch_size}")
|
|
||||||
print(f"Output: {output_path}")
|
|
||||||
|
|
||||||
rows: list[dict[str, Any]] = []
|
|
||||||
for batch in tqdm(loader, desc="Future risks", dynamic_ncols=True):
|
|
||||||
hidden = infer_landmark_hidden(
|
|
||||||
model=model,
|
|
||||||
batch=batch,
|
|
||||||
device=device,
|
|
||||||
model_target_mode=model_target_mode,
|
|
||||||
readout_name=readout_name,
|
|
||||||
readout_reduce=readout_reduce,
|
|
||||||
)
|
|
||||||
logits = model.calc_risk(hidden)
|
|
||||||
rho = model.calc_weibull_rho(hidden) if dist_mode == "weibull" else None
|
|
||||||
death_rho = model.calc_death_rho(
|
|
||||||
hidden) if dist_mode == "mixed" else None
|
|
||||||
probabilities = probabilities_from_logits(
|
|
||||||
logits,
|
|
||||||
tau,
|
|
||||||
dist_mode=dist_mode,
|
|
||||||
rho=rho,
|
|
||||||
death_rho=death_rho,
|
|
||||||
)
|
|
||||||
occurred = make_occurred_mask(
|
|
||||||
batch["event_seq"].to(device),
|
|
||||||
vocab_size=int(dataset.vocab_size),
|
|
||||||
device=device,
|
|
||||||
)
|
|
||||||
|
|
||||||
death_risk_tensor = death_risk_from_probabilities(probabilities)
|
|
||||||
death_hazard_tensor = mortality_hazard_from_risk(death_risk_tensor)
|
|
||||||
death_risk = death_risk_tensor.detach().cpu().numpy()
|
|
||||||
|
|
||||||
group_risk: dict[str, np.ndarray] = {}
|
|
||||||
for group in group_names:
|
|
||||||
group_risk[group] = new_disease_risk_from_probabilities(
|
|
||||||
probabilities,
|
|
||||||
occurred,
|
|
||||||
organ_groups[group],
|
|
||||||
).detach().cpu().numpy()
|
|
||||||
|
|
||||||
group_mortality_attr_prob: dict[str, np.ndarray] = {}
|
|
||||||
group_mortality_attr_hazard: dict[str, np.ndarray] = {}
|
|
||||||
batch_n = int(batch["event_seq"].shape[0])
|
|
||||||
zeros = np.zeros(batch_n, dtype=np.float32)
|
|
||||||
for group in group_names:
|
|
||||||
group_mortality_attr_prob[group] = zeros.copy()
|
|
||||||
group_mortality_attr_hazard[group] = zeros.copy()
|
|
||||||
|
|
||||||
for ablated_chunk, chunk_groups, chunk_rows in iter_group_ablated_batches(
|
|
||||||
batch=batch,
|
|
||||||
group_names=group_names,
|
|
||||||
organ_groups=organ_groups,
|
|
||||||
occurred=occurred,
|
|
||||||
max_batch_size=attribution_batch_size,
|
|
||||||
):
|
|
||||||
ablated_death_risk = death_risk_for_batch(
|
|
||||||
model=model,
|
|
||||||
batch=ablated_chunk,
|
|
||||||
device=device,
|
|
||||||
model_target_mode=model_target_mode,
|
|
||||||
readout_name=readout_name,
|
|
||||||
readout_reduce=readout_reduce,
|
|
||||||
dist_mode=dist_mode,
|
|
||||||
tau=tau,
|
|
||||||
)
|
|
||||||
row_tensor = torch.as_tensor(
|
|
||||||
chunk_rows, dtype=torch.long, device=device)
|
|
||||||
ablated_death_hazard = mortality_hazard_from_risk(
|
|
||||||
ablated_death_risk)
|
|
||||||
attr_prob = (
|
|
||||||
death_risk_tensor[row_tensor] - ablated_death_risk
|
|
||||||
).detach().cpu().numpy()
|
|
||||||
attr_hazard = (
|
|
||||||
death_hazard_tensor[row_tensor] - ablated_death_hazard
|
|
||||||
).detach().cpu().numpy()
|
|
||||||
for local_idx, (group, row_idx) in enumerate(zip(chunk_groups, chunk_rows)):
|
|
||||||
group_mortality_attr_prob[group][row_idx] = attr_prob[local_idx]
|
|
||||||
group_mortality_attr_hazard[group][row_idx] = attr_hazard[local_idx]
|
|
||||||
|
|
||||||
row_indices = batch["row_idx"].cpu().numpy().astype(np.int64)
|
|
||||||
for j, row_idx in enumerate(row_indices.tolist()):
|
|
||||||
meta = landmark_dataset.rows[int(row_idx)]
|
|
||||||
dataset_index = int(meta["dataset_index"])
|
|
||||||
sample = dataset.samples[dataset_index]
|
|
||||||
hist_tokens = np.asarray(meta["event_seq"], dtype=np.int64)
|
|
||||||
total_count, group_counts = historical_counts_by_group(
|
|
||||||
hist_tokens,
|
|
||||||
death_idx=death_idx,
|
|
||||||
token_to_group=token_to_group,
|
|
||||||
group_names=group_names,
|
|
||||||
)
|
|
||||||
|
|
||||||
out: dict[str, Any] = {
|
|
||||||
"patient_id": int(meta["patient_id"]),
|
|
||||||
"dataset_index": dataset_index,
|
|
||||||
"eid": int(sample.get("eid", -1)),
|
|
||||||
"sex": int(meta["sex"]),
|
|
||||||
"landmark_age": float(meta["landmark_age"]),
|
|
||||||
"tau": tau,
|
|
||||||
"followup_end_time": float(meta["followup_end_time"]),
|
|
||||||
"history_disease_count": int(total_count),
|
|
||||||
"death_risk": float(death_risk[j]),
|
|
||||||
}
|
|
||||||
for group in group_names:
|
|
||||||
out[f"history_count__{group}"] = int(group_counts[group])
|
|
||||||
out[f"new_disease_risk__{group}"] = float(group_risk[group][j])
|
|
||||||
if int(group_counts[group]) == 0:
|
|
||||||
group_mortality_attr_prob[group][j] = 0.0
|
|
||||||
group_mortality_attr_hazard[group][j] = 0.0
|
|
||||||
out[f"mortality_attribution_probability__{group}"] = float(
|
|
||||||
group_mortality_attr_prob[group][j]
|
|
||||||
)
|
|
||||||
out[f"mortality_attribution_hazard__{group}"] = float(
|
|
||||||
group_mortality_attr_hazard[group][j]
|
|
||||||
)
|
|
||||||
rows.append(out)
|
|
||||||
|
|
||||||
df = pd.DataFrame(rows)
|
|
||||||
df.to_csv(output_path, index=False)
|
|
||||||
print(f"Wrote {len(df)} rows to {output_path}")
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
main()
|
|
||||||
@@ -1,952 +0,0 @@
|
|||||||
"""Evaluate extra-info attribution to death and disease distribution parameters.
|
|
||||||
|
|
||||||
For each landmark query, this script scans selected extra-info types that are
|
|
||||||
available at or before the query age. For each such type it re-runs the model
|
|
||||||
with that extra-info type removed and summarizes:
|
|
||||||
|
|
||||||
* death distribution parameters before and after ablation;
|
|
||||||
* disease distribution parameters before and after ablation, by ICD-10
|
|
||||||
chapter-derived organ/system groups.
|
|
||||||
|
|
||||||
Death is always token vocab_size - 1.
|
|
||||||
"""
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import argparse
|
|
||||||
import json
|
|
||||||
import re
|
|
||||||
from concurrent.futures import ProcessPoolExecutor, as_completed
|
|
||||||
from pathlib import Path
|
|
||||||
from typing import Any, Sequence
|
|
||||||
|
|
||||||
import numpy as np
|
|
||||||
import pandas as pd
|
|
||||||
import torch
|
|
||||||
import torch.nn.functional as F
|
|
||||||
from torch.utils.data import DataLoader
|
|
||||||
from tqdm.auto import tqdm
|
|
||||||
|
|
||||||
from evaluate_auc_v2 import (
|
|
||||||
build_model_from_dataset,
|
|
||||||
cfg_get,
|
|
||||||
load_checkpoint_state_dict,
|
|
||||||
load_json_config,
|
|
||||||
load_model_state,
|
|
||||||
resolve_dist_mode_for_checkpoint,
|
|
||||||
resolve_eval_device,
|
|
||||||
validate_dataset_metadata,
|
|
||||||
)
|
|
||||||
from landmark_eval_utils import (
|
|
||||||
IndexedLandmarkDataset,
|
|
||||||
LandmarkDataset,
|
|
||||||
build_first_occurrence_maps_for_landmarks,
|
|
||||||
collate_indexed_landmark_fn,
|
|
||||||
infer_landmark_hidden,
|
|
||||||
load_eval_sequence_dataset,
|
|
||||||
load_organ_groups,
|
|
||||||
make_landmark_ages,
|
|
||||||
)
|
|
||||||
|
|
||||||
EXTRA_KEY_COLUMNS = [
|
|
||||||
"selected_extra_info_type_id",
|
|
||||||
"selected_extra_info_var_name",
|
|
||||||
"selected_extra_info_full_name",
|
|
||||||
"landmark_age",
|
|
||||||
"sex",
|
|
||||||
]
|
|
||||||
|
|
||||||
DEATH_PARAMETER_COLUMNS = [
|
|
||||||
"original_death_lambda",
|
|
||||||
"ablated_death_lambda",
|
|
||||||
"original_death_scale",
|
|
||||||
"ablated_death_scale",
|
|
||||||
"original_death_shape",
|
|
||||||
"ablated_death_shape",
|
|
||||||
]
|
|
||||||
|
|
||||||
DISEASE_PARAMETER_KEY_COLUMNS = [
|
|
||||||
*EXTRA_KEY_COLUMNS,
|
|
||||||
"target_group",
|
|
||||||
"target_group_label",
|
|
||||||
]
|
|
||||||
|
|
||||||
DISEASE_PARAMETER_COLUMNS = [
|
|
||||||
"original_disease_lambda",
|
|
||||||
"ablated_disease_lambda",
|
|
||||||
"original_disease_scale",
|
|
||||||
"ablated_disease_scale",
|
|
||||||
"original_disease_shape",
|
|
||||||
"ablated_disease_shape",
|
|
||||||
]
|
|
||||||
|
|
||||||
|
|
||||||
def parse_int_list(value: Any) -> list[int] | None:
|
|
||||||
if value is None:
|
|
||||||
return None
|
|
||||||
if isinstance(value, (list, tuple, np.ndarray)):
|
|
||||||
return [int(x) for x in value]
|
|
||||||
text = str(value).strip()
|
|
||||||
if text == "":
|
|
||||||
return None
|
|
||||||
if text.startswith("["):
|
|
||||||
raw = json.loads(text)
|
|
||||||
if not isinstance(raw, list):
|
|
||||||
raise ValueError("Expected JSON list for integer list")
|
|
||||||
return [int(x) for x in raw]
|
|
||||||
return [int(x.strip()) for x in re.split(r"[,;\s]+", text) if x.strip()]
|
|
||||||
|
|
||||||
|
|
||||||
def load_extra_info_metadata(
|
|
||||||
*,
|
|
||||||
dataset_extra_info_types: Sequence[int],
|
|
||||||
search_root: Path = Path("."),
|
|
||||||
) -> dict[int, dict[str, Any]]:
|
|
||||||
metadata: dict[int, dict[str, Any]] = {
|
|
||||||
int(type_id): {
|
|
||||||
"type_id": int(type_id),
|
|
||||||
"var_name": f"extra_info_{int(type_id)}",
|
|
||||||
"full_name": f"extra-info type {int(type_id)}",
|
|
||||||
}
|
|
||||||
for type_id in dataset_extra_info_types
|
|
||||||
}
|
|
||||||
|
|
||||||
line_re = re.compile(r"^\s*(\d+)\s*#\s*([^|#]+?)(?:\s*\|\s*(.*?))?\s*$")
|
|
||||||
for path in sorted(search_root.glob("extra_info_types*.txt")):
|
|
||||||
for line in path.read_text(encoding="utf-8").splitlines():
|
|
||||||
match = line_re.match(line)
|
|
||||||
if not match:
|
|
||||||
continue
|
|
||||||
type_id = int(match.group(1))
|
|
||||||
if type_id not in metadata:
|
|
||||||
continue
|
|
||||||
var_name = match.group(2).strip()
|
|
||||||
full_name = (match.group(3) or var_name).strip()
|
|
||||||
metadata[type_id] = {
|
|
||||||
"type_id": type_id,
|
|
||||||
"var_name": var_name,
|
|
||||||
"full_name": full_name,
|
|
||||||
}
|
|
||||||
|
|
||||||
return metadata
|
|
||||||
|
|
||||||
|
|
||||||
def resolve_extra_info_types(
|
|
||||||
value: str | None,
|
|
||||||
*,
|
|
||||||
dataset_extra_info_types: Sequence[int],
|
|
||||||
metadata: dict[int, dict[str, Any]],
|
|
||||||
) -> list[int]:
|
|
||||||
available = [int(x) for x in dataset_extra_info_types]
|
|
||||||
if value is None or str(value).strip() == "":
|
|
||||||
return available
|
|
||||||
|
|
||||||
out: list[int] = []
|
|
||||||
seen: set[int] = set()
|
|
||||||
by_var = {
|
|
||||||
str(meta.get("var_name", "")).lower(): int(type_id)
|
|
||||||
for type_id, meta in metadata.items()
|
|
||||||
}
|
|
||||||
by_full = {
|
|
||||||
str(meta.get("full_name", "")).lower(): int(type_id)
|
|
||||||
for type_id, meta in metadata.items()
|
|
||||||
}
|
|
||||||
for part in re.split(r"[,;]+", str(value)):
|
|
||||||
text = part.strip()
|
|
||||||
if not text:
|
|
||||||
continue
|
|
||||||
if text.isdigit() or (text.startswith("-") and text[1:].isdigit()):
|
|
||||||
type_id = int(text)
|
|
||||||
else:
|
|
||||||
lower = text.lower()
|
|
||||||
if lower in by_var:
|
|
||||||
type_id = by_var[lower]
|
|
||||||
elif lower in by_full:
|
|
||||||
type_id = by_full[lower]
|
|
||||||
else:
|
|
||||||
matches = [
|
|
||||||
int(t)
|
|
||||||
for t, meta in metadata.items()
|
|
||||||
if lower in str(meta.get("var_name", "")).lower()
|
|
||||||
or lower in str(meta.get("full_name", "")).lower()
|
|
||||||
]
|
|
||||||
if len(matches) != 1:
|
|
||||||
raise ValueError(
|
|
||||||
f"--extra_info={text!r} matched {len(matches)} types; "
|
|
||||||
"use a type id or exact variable name."
|
|
||||||
)
|
|
||||||
type_id = matches[0]
|
|
||||||
if type_id not in available:
|
|
||||||
raise ValueError(
|
|
||||||
f"extra-info type {type_id} is not available in this dataset/run"
|
|
||||||
)
|
|
||||||
if type_id not in seen:
|
|
||||||
out.append(type_id)
|
|
||||||
seen.add(type_id)
|
|
||||||
return out
|
|
||||||
|
|
||||||
|
|
||||||
def death_distribution_parameters(
|
|
||||||
model,
|
|
||||||
hidden: torch.Tensor,
|
|
||||||
*,
|
|
||||||
dist_mode: str,
|
|
||||||
eps: float = 1e-8,
|
|
||||||
) -> tuple[str, torch.Tensor]:
|
|
||||||
logits = model.calc_risk(hidden)
|
|
||||||
death_idx = int(logits.shape[1]) - 1
|
|
||||||
death_lambda = F.softplus(logits[:, death_idx]) + float(eps)
|
|
||||||
|
|
||||||
if dist_mode == "exponential":
|
|
||||||
nan = torch.full_like(death_lambda, float("nan"))
|
|
||||||
return "exponential", torch.stack([death_lambda, nan, nan], dim=1)
|
|
||||||
|
|
||||||
if dist_mode == "weibull":
|
|
||||||
rho = model.calc_weibull_rho(hidden)[:, death_idx].to(dtype=death_lambda.dtype)
|
|
||||||
elif dist_mode == "mixed":
|
|
||||||
rho = model.calc_death_rho(hidden).to(dtype=death_lambda.dtype)
|
|
||||||
else:
|
|
||||||
raise ValueError(f"Unsupported dist_mode={dist_mode!r}")
|
|
||||||
|
|
||||||
shape = rho.clamp_min(float(eps))
|
|
||||||
scale = torch.pow(death_lambda.clamp_min(float(eps)), -1.0 / shape)
|
|
||||||
nan = torch.full_like(death_lambda, float("nan"))
|
|
||||||
return "weibull", torch.stack([nan, scale, shape], dim=1)
|
|
||||||
|
|
||||||
|
|
||||||
def parameter_pair_block(original: torch.Tensor, ablated: torch.Tensor) -> torch.Tensor:
|
|
||||||
return torch.stack(
|
|
||||||
[
|
|
||||||
original[:, 0],
|
|
||||||
ablated[:, 0],
|
|
||||||
original[:, 1],
|
|
||||||
ablated[:, 1],
|
|
||||||
original[:, 2],
|
|
||||||
ablated[:, 2],
|
|
||||||
],
|
|
||||||
dim=1,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def all_disease_parameter_pair_block(
|
|
||||||
*,
|
|
||||||
original_logits: torch.Tensor,
|
|
||||||
ablated_logits: torch.Tensor,
|
|
||||||
dist_mode: str,
|
|
||||||
original_rho: torch.Tensor | None = None,
|
|
||||||
ablated_rho: torch.Tensor | None = None,
|
|
||||||
eps: float = 1e-8,
|
|
||||||
) -> torch.Tensor:
|
|
||||||
original_lambda = F.softplus(original_logits) + float(eps)
|
|
||||||
ablated_lambda = F.softplus(ablated_logits) + float(eps)
|
|
||||||
|
|
||||||
if dist_mode in {"exponential", "mixed"}:
|
|
||||||
nan = torch.full_like(original_lambda, float("nan"))
|
|
||||||
return torch.stack(
|
|
||||||
[
|
|
||||||
original_lambda,
|
|
||||||
ablated_lambda,
|
|
||||||
nan,
|
|
||||||
nan,
|
|
||||||
nan,
|
|
||||||
nan,
|
|
||||||
],
|
|
||||||
dim=2,
|
|
||||||
)
|
|
||||||
|
|
||||||
if dist_mode == "weibull":
|
|
||||||
if original_rho is None or ablated_rho is None:
|
|
||||||
raise ValueError("rho tensors are required for weibull disease parameters")
|
|
||||||
original_shape = original_rho.to(dtype=original_lambda.dtype).clamp_min(float(eps))
|
|
||||||
ablated_shape = ablated_rho.to(dtype=ablated_lambda.dtype).clamp_min(float(eps))
|
|
||||||
original_scale = torch.pow(original_lambda.clamp_min(float(eps)), -1.0 / original_shape)
|
|
||||||
ablated_scale = torch.pow(ablated_lambda.clamp_min(float(eps)), -1.0 / ablated_shape)
|
|
||||||
nan = torch.full_like(original_lambda, float("nan"))
|
|
||||||
return torch.stack(
|
|
||||||
[
|
|
||||||
nan,
|
|
||||||
nan,
|
|
||||||
original_scale,
|
|
||||||
ablated_scale,
|
|
||||||
original_shape,
|
|
||||||
ablated_shape,
|
|
||||||
],
|
|
||||||
dim=2,
|
|
||||||
)
|
|
||||||
|
|
||||||
raise ValueError(f"Unsupported dist_mode={dist_mode!r}")
|
|
||||||
|
|
||||||
|
|
||||||
def grouped_parameter_stats(
|
|
||||||
values: torch.Tensor,
|
|
||||||
group_token_mask: torch.Tensor,
|
|
||||||
) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
|
|
||||||
finite = torch.isfinite(values)
|
|
||||||
values64 = values.to(dtype=torch.float64)
|
|
||||||
safe_values = torch.where(finite, values64, torch.zeros_like(values64))
|
|
||||||
mask = group_token_mask.to(device=values.device, dtype=torch.float64)
|
|
||||||
sums = torch.einsum("nvc,gv->ngc", safe_values, mask)
|
|
||||||
sumsq = torch.einsum("nvc,gv->ngc", safe_values * safe_values, mask)
|
|
||||||
counts = torch.einsum("nvc,gv->ngc", finite.to(dtype=torch.float64), mask)
|
|
||||||
return (
|
|
||||||
sums.detach().cpu().numpy().astype(np.float64, copy=False),
|
|
||||||
sumsq.detach().cpu().numpy().astype(np.float64, copy=False),
|
|
||||||
counts.detach().cpu().numpy().astype(np.float64, copy=False),
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def build_extra_info_ablated_slice(
|
|
||||||
batch: dict[str, torch.Tensor],
|
|
||||||
*,
|
|
||||||
row_indices: torch.Tensor,
|
|
||||||
extra_info_type_id: int,
|
|
||||||
) -> dict[str, torch.Tensor]:
|
|
||||||
out: dict[str, torch.Tensor] = {}
|
|
||||||
repeated_keys = (
|
|
||||||
"event_seq",
|
|
||||||
"time_seq",
|
|
||||||
"padding_mask",
|
|
||||||
"readout_mask",
|
|
||||||
"sex",
|
|
||||||
"landmark_pos",
|
|
||||||
"t_query",
|
|
||||||
"patient_id",
|
|
||||||
"landmark_age",
|
|
||||||
"followup_end_time",
|
|
||||||
"death_time",
|
|
||||||
"row_idx",
|
|
||||||
)
|
|
||||||
for key in repeated_keys:
|
|
||||||
out[key] = batch[key][row_indices]
|
|
||||||
|
|
||||||
out["other_type"] = batch["other_type"][row_indices].clone()
|
|
||||||
out["other_value"] = batch["other_value"][row_indices].clone()
|
|
||||||
out["other_value_kind"] = batch["other_value_kind"][row_indices].clone()
|
|
||||||
out["other_time"] = batch["other_time"][row_indices].clone()
|
|
||||||
|
|
||||||
remove = out["other_type"] == int(extra_info_type_id)
|
|
||||||
out["other_type"][remove] = 0
|
|
||||||
out["other_value"][remove] = 0
|
|
||||||
out["other_value_kind"][remove] = 0
|
|
||||||
out["other_time"][remove] = 0
|
|
||||||
return out
|
|
||||||
|
|
||||||
|
|
||||||
def concat_tensor_batches(chunks: Sequence[dict[str, torch.Tensor]]) -> dict[str, torch.Tensor]:
|
|
||||||
return {key: torch.cat([chunk[key] for chunk in chunks], dim=0) for key in chunks[0]}
|
|
||||||
|
|
||||||
|
|
||||||
def iter_extra_info_ablated_batches(
|
|
||||||
batch: dict[str, torch.Tensor],
|
|
||||||
*,
|
|
||||||
selected_extra_info_types: Sequence[int],
|
|
||||||
max_batch_size: int,
|
|
||||||
):
|
|
||||||
pending_batches: list[dict[str, torch.Tensor]] = []
|
|
||||||
pending_types: list[int] = []
|
|
||||||
pending_rows: list[int] = []
|
|
||||||
pending_n = 0
|
|
||||||
|
|
||||||
other_type = batch["other_type"]
|
|
||||||
visible = other_type > 0
|
|
||||||
visible &= batch["other_time"] <= batch["t_query"][:, None].to(batch["other_time"].dtype)
|
|
||||||
|
|
||||||
for type_id in selected_extra_info_types:
|
|
||||||
active_rows = torch.nonzero(
|
|
||||||
((other_type == int(type_id)) & visible).any(dim=1),
|
|
||||||
as_tuple=False,
|
|
||||||
).flatten()
|
|
||||||
if active_rows.numel() == 0:
|
|
||||||
continue
|
|
||||||
|
|
||||||
row_offset = 0
|
|
||||||
while row_offset < int(active_rows.numel()):
|
|
||||||
capacity = int(max_batch_size) - pending_n
|
|
||||||
row_stop = min(int(active_rows.numel()), row_offset + capacity)
|
|
||||||
row_indices = active_rows[row_offset:row_stop].to(device=batch["event_seq"].device)
|
|
||||||
chunk = build_extra_info_ablated_slice(
|
|
||||||
batch,
|
|
||||||
row_indices=row_indices,
|
|
||||||
extra_info_type_id=int(type_id),
|
|
||||||
)
|
|
||||||
chunk_n = int(row_indices.numel())
|
|
||||||
pending_batches.append(chunk)
|
|
||||||
pending_types.extend([int(type_id)] * chunk_n)
|
|
||||||
pending_rows.extend(int(x) for x in row_indices.detach().cpu().tolist())
|
|
||||||
pending_n += chunk_n
|
|
||||||
row_offset = row_stop
|
|
||||||
|
|
||||||
if pending_n >= int(max_batch_size):
|
|
||||||
yield concat_tensor_batches(pending_batches), pending_types, pending_rows
|
|
||||||
pending_batches = []
|
|
||||||
pending_types = []
|
|
||||||
pending_rows = []
|
|
||||||
pending_n = 0
|
|
||||||
|
|
||||||
if pending_batches:
|
|
||||||
yield concat_tensor_batches(pending_batches), pending_types, pending_rows
|
|
||||||
|
|
||||||
|
|
||||||
def finite_float64(values: Any) -> np.ndarray:
|
|
||||||
arr = np.asarray(values, dtype=np.float64)
|
|
||||||
return arr[np.isfinite(arr)]
|
|
||||||
|
|
||||||
|
|
||||||
def update_death_summary(
|
|
||||||
summary: dict[tuple[Any, ...], dict[str, float]],
|
|
||||||
*,
|
|
||||||
key_rows: pd.DataFrame,
|
|
||||||
values: np.ndarray,
|
|
||||||
) -> None:
|
|
||||||
if key_rows.empty:
|
|
||||||
return
|
|
||||||
table = key_rows.copy()
|
|
||||||
for idx, column in enumerate(DEATH_PARAMETER_COLUMNS):
|
|
||||||
table[column] = values[:, idx]
|
|
||||||
|
|
||||||
for key, group in table.groupby(EXTRA_KEY_COLUMNS, dropna=False, sort=False):
|
|
||||||
if not isinstance(key, tuple):
|
|
||||||
key = (key,)
|
|
||||||
acc = summary.setdefault(
|
|
||||||
key,
|
|
||||||
{
|
|
||||||
"n": 0.0,
|
|
||||||
**{f"count__{col}": 0.0 for col in DEATH_PARAMETER_COLUMNS},
|
|
||||||
**{f"sum__{col}": 0.0 for col in DEATH_PARAMETER_COLUMNS},
|
|
||||||
**{f"sumsq__{col}": 0.0 for col in DEATH_PARAMETER_COLUMNS},
|
|
||||||
},
|
|
||||||
)
|
|
||||||
acc["n"] += float(len(group))
|
|
||||||
for column in DEATH_PARAMETER_COLUMNS:
|
|
||||||
vals = finite_float64(pd.to_numeric(group[column], errors="coerce"))
|
|
||||||
acc[f"count__{column}"] += float(vals.size)
|
|
||||||
acc[f"sum__{column}"] += float(vals.sum())
|
|
||||||
acc[f"sumsq__{column}"] += float(np.square(vals).sum())
|
|
||||||
|
|
||||||
|
|
||||||
def update_disease_parameter_summary_from_group_stats(
|
|
||||||
summary: dict[tuple[Any, ...], dict[str, float]],
|
|
||||||
*,
|
|
||||||
key_rows: pd.DataFrame,
|
|
||||||
group_names: Sequence[str],
|
|
||||||
group_labels: Sequence[str],
|
|
||||||
sums: np.ndarray,
|
|
||||||
sumsq: np.ndarray,
|
|
||||||
counts: np.ndarray,
|
|
||||||
) -> None:
|
|
||||||
if key_rows.empty or sums.size == 0:
|
|
||||||
return
|
|
||||||
rows = key_rows.reset_index(drop=True)
|
|
||||||
for row_idx, row in rows.iterrows():
|
|
||||||
base_key = tuple(row[column] for column in EXTRA_KEY_COLUMNS)
|
|
||||||
for group_idx, (group, label) in enumerate(zip(group_names, group_labels)):
|
|
||||||
count_row = counts[int(row_idx), int(group_idx)]
|
|
||||||
n_add = float(np.nanmax(count_row)) if count_row.size else 0.0
|
|
||||||
if n_add <= 0:
|
|
||||||
continue
|
|
||||||
full_key = (*base_key, str(group), str(label))
|
|
||||||
acc = summary.setdefault(
|
|
||||||
full_key,
|
|
||||||
{
|
|
||||||
"n": 0.0,
|
|
||||||
**{f"count__{col}": 0.0 for col in DISEASE_PARAMETER_COLUMNS},
|
|
||||||
**{f"sum__{col}": 0.0 for col in DISEASE_PARAMETER_COLUMNS},
|
|
||||||
**{f"sumsq__{col}": 0.0 for col in DISEASE_PARAMETER_COLUMNS},
|
|
||||||
},
|
|
||||||
)
|
|
||||||
acc["n"] += n_add
|
|
||||||
for col_idx, column in enumerate(DISEASE_PARAMETER_COLUMNS):
|
|
||||||
count = float(counts[int(row_idx), int(group_idx), int(col_idx)])
|
|
||||||
if count <= 0:
|
|
||||||
continue
|
|
||||||
acc[f"count__{column}"] += count
|
|
||||||
acc[f"sum__{column}"] += float(sums[int(row_idx), int(group_idx), int(col_idx)])
|
|
||||||
acc[f"sumsq__{column}"] += float(sumsq[int(row_idx), int(group_idx), int(col_idx)])
|
|
||||||
|
|
||||||
|
|
||||||
def merge_summary_dict(
|
|
||||||
dst: dict[tuple[Any, ...], dict[str, float]],
|
|
||||||
src: dict[tuple[Any, ...], dict[str, float]],
|
|
||||||
) -> None:
|
|
||||||
for key, src_acc in src.items():
|
|
||||||
dst_acc = dst.setdefault(key, {name: 0.0 for name in src_acc})
|
|
||||||
for name, value in src_acc.items():
|
|
||||||
dst_acc[name] = dst_acc.get(name, 0.0) + float(value)
|
|
||||||
|
|
||||||
|
|
||||||
def reduce_attribution_chunk_bundle(
|
|
||||||
payload: tuple[
|
|
||||||
list[tuple[pd.DataFrame, np.ndarray]],
|
|
||||||
list[tuple[pd.DataFrame, np.ndarray, np.ndarray, np.ndarray]],
|
|
||||||
list[str],
|
|
||||||
list[str],
|
|
||||||
],
|
|
||||||
) -> tuple[dict[tuple[Any, ...], dict[str, float]], dict[tuple[Any, ...], dict[str, float]]]:
|
|
||||||
death_items, disease_items, group_names, group_labels = payload
|
|
||||||
death_summary: dict[tuple[Any, ...], dict[str, float]] = {}
|
|
||||||
disease_summary: dict[tuple[Any, ...], dict[str, float]] = {}
|
|
||||||
|
|
||||||
for key_rows, values in death_items:
|
|
||||||
update_death_summary(
|
|
||||||
death_summary,
|
|
||||||
key_rows=key_rows,
|
|
||||||
values=values,
|
|
||||||
)
|
|
||||||
|
|
||||||
for key_rows, sums, sumsq, counts in disease_items:
|
|
||||||
update_disease_parameter_summary_from_group_stats(
|
|
||||||
disease_summary,
|
|
||||||
key_rows=key_rows,
|
|
||||||
group_names=group_names,
|
|
||||||
group_labels=group_labels,
|
|
||||||
sums=sums,
|
|
||||||
sumsq=sumsq,
|
|
||||||
counts=counts,
|
|
||||||
)
|
|
||||||
|
|
||||||
return death_summary, disease_summary
|
|
||||||
|
|
||||||
|
|
||||||
def reduce_attribution_chunks(
|
|
||||||
*,
|
|
||||||
death_key_chunks: list[pd.DataFrame],
|
|
||||||
death_value_chunks: list[np.ndarray],
|
|
||||||
disease_stat_chunks: list[tuple[pd.DataFrame, np.ndarray, np.ndarray, np.ndarray]],
|
|
||||||
group_names: list[str],
|
|
||||||
group_labels: list[str],
|
|
||||||
cpu_reduce_workers: int,
|
|
||||||
) -> tuple[dict[tuple[Any, ...], dict[str, float]], dict[tuple[Any, ...], dict[str, float]]]:
|
|
||||||
n_chunks = max(len(death_key_chunks), len(disease_stat_chunks))
|
|
||||||
if n_chunks == 0:
|
|
||||||
return {}, {}
|
|
||||||
|
|
||||||
worker_count = max(1, min(int(cpu_reduce_workers), n_chunks))
|
|
||||||
if worker_count == 1:
|
|
||||||
return reduce_attribution_chunk_bundle(
|
|
||||||
(
|
|
||||||
list(zip(death_key_chunks, death_value_chunks)),
|
|
||||||
disease_stat_chunks,
|
|
||||||
group_names,
|
|
||||||
group_labels,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
bundles = []
|
|
||||||
for worker_idx in range(worker_count):
|
|
||||||
start = worker_idx * n_chunks // worker_count
|
|
||||||
stop = (worker_idx + 1) * n_chunks // worker_count
|
|
||||||
if start >= stop:
|
|
||||||
continue
|
|
||||||
death_items = [
|
|
||||||
(death_key_chunks[i], death_value_chunks[i])
|
|
||||||
for i in range(start, min(stop, len(death_key_chunks)))
|
|
||||||
]
|
|
||||||
disease_items = disease_stat_chunks[start:min(stop, len(disease_stat_chunks))]
|
|
||||||
bundles.append((death_items, disease_items, group_names, group_labels))
|
|
||||||
|
|
||||||
merged_death: dict[tuple[Any, ...], dict[str, float]] = {}
|
|
||||||
merged_disease: dict[tuple[Any, ...], dict[str, float]] = {}
|
|
||||||
with ProcessPoolExecutor(max_workers=len(bundles)) as executor:
|
|
||||||
futures = [executor.submit(reduce_attribution_chunk_bundle, bundle) for bundle in bundles]
|
|
||||||
for future in tqdm(as_completed(futures), total=len(futures), desc="CPU summary reduction", dynamic_ncols=True):
|
|
||||||
death_part, disease_part = future.result()
|
|
||||||
merge_summary_dict(merged_death, death_part)
|
|
||||||
merge_summary_dict(merged_disease, disease_part)
|
|
||||||
|
|
||||||
return merged_death, merged_disease
|
|
||||||
|
|
||||||
|
|
||||||
def write_death_summary_csv(
|
|
||||||
path: Path,
|
|
||||||
summary: dict[tuple[Any, ...], dict[str, float]],
|
|
||||||
*,
|
|
||||||
death_distribution: str,
|
|
||||||
) -> int:
|
|
||||||
rows: list[dict[str, Any]] = []
|
|
||||||
for key, acc in summary.items():
|
|
||||||
n = int(acc["n"])
|
|
||||||
row = {column: value for column, value in zip(EXTRA_KEY_COLUMNS, key)}
|
|
||||||
row["n"] = n
|
|
||||||
row["death_distribution"] = death_distribution
|
|
||||||
for column in DEATH_PARAMETER_COLUMNS:
|
|
||||||
count = int(acc[f"count__{column}"])
|
|
||||||
mean = acc[f"sum__{column}"] / count if count > 0 else np.nan
|
|
||||||
second = acc[f"sumsq__{column}"] / count if count > 0 else np.nan
|
|
||||||
row[f"mean__{column}"] = mean
|
|
||||||
row[f"var__{column}"] = second - mean * mean if count > 0 else np.nan
|
|
||||||
rows.append(row)
|
|
||||||
columns = [
|
|
||||||
*EXTRA_KEY_COLUMNS,
|
|
||||||
"n",
|
|
||||||
"death_distribution",
|
|
||||||
*[
|
|
||||||
name
|
|
||||||
for column in DEATH_PARAMETER_COLUMNS
|
|
||||||
for name in (f"mean__{column}", f"var__{column}")
|
|
||||||
],
|
|
||||||
]
|
|
||||||
pd.DataFrame(rows, columns=columns).sort_values(
|
|
||||||
["selected_extra_info_type_id", "landmark_age", "sex"],
|
|
||||||
kind="mergesort",
|
|
||||||
).to_csv(path, index=False)
|
|
||||||
return len(rows)
|
|
||||||
|
|
||||||
|
|
||||||
def write_disease_parameter_summary_csv(
|
|
||||||
path: Path,
|
|
||||||
summary: dict[tuple[Any, ...], dict[str, float]],
|
|
||||||
) -> int:
|
|
||||||
rows: list[dict[str, Any]] = []
|
|
||||||
for key, acc in summary.items():
|
|
||||||
n = int(acc["n"])
|
|
||||||
row = {column: value for column, value in zip(DISEASE_PARAMETER_KEY_COLUMNS, key)}
|
|
||||||
row["n"] = n
|
|
||||||
for column in DISEASE_PARAMETER_COLUMNS:
|
|
||||||
count = int(acc[f"count__{column}"])
|
|
||||||
mean = acc[f"sum__{column}"] / count if count > 0 else np.nan
|
|
||||||
second = acc[f"sumsq__{column}"] / count if count > 0 else np.nan
|
|
||||||
row[f"mean__{column}"] = mean
|
|
||||||
row[f"var__{column}"] = second - mean * mean if count > 0 else np.nan
|
|
||||||
rows.append(row)
|
|
||||||
columns = [
|
|
||||||
*DISEASE_PARAMETER_KEY_COLUMNS,
|
|
||||||
"n",
|
|
||||||
*[
|
|
||||||
name
|
|
||||||
for column in DISEASE_PARAMETER_COLUMNS
|
|
||||||
for name in (f"mean__{column}", f"var__{column}")
|
|
||||||
],
|
|
||||||
]
|
|
||||||
pd.DataFrame(rows, columns=columns).sort_values(
|
|
||||||
["selected_extra_info_type_id", "target_group", "landmark_age", "sex"],
|
|
||||||
kind="mergesort",
|
|
||||||
).to_csv(path, index=False)
|
|
||||||
return len(rows)
|
|
||||||
|
|
||||||
|
|
||||||
def parse_args() -> argparse.Namespace:
|
|
||||||
parser = argparse.ArgumentParser(
|
|
||||||
description="Compute extra-info ablation attribution for death and disease distribution parameters."
|
|
||||||
)
|
|
||||||
parser.add_argument("--run_path", type=str, required=True)
|
|
||||||
parser.add_argument(
|
|
||||||
"--extra_info",
|
|
||||||
type=str,
|
|
||||||
default=None,
|
|
||||||
help=(
|
|
||||||
"Optional type id, variable name, exact full name, or comma-separated list. "
|
|
||||||
"If omitted, scan all extra-info types available in the run."
|
|
||||||
),
|
|
||||||
)
|
|
||||||
parser.add_argument("--output_dir", type=str, default=None)
|
|
||||||
parser.add_argument("--organ_mapping_path", type=str, default="icd10_chapter_organ_mapping.csv")
|
|
||||||
parser.add_argument("--eval_split", type=str, default=None)
|
|
||||||
parser.add_argument("--dataset_subset_size", type=int, default=None)
|
|
||||||
parser.add_argument("--train_eid_file", type=str, default=None)
|
|
||||||
parser.add_argument("--val_eid_file", type=str, default=None)
|
|
||||||
parser.add_argument("--test_eid_file", type=str, default=None)
|
|
||||||
parser.add_argument("--landmark_start", type=float, default=40.0)
|
|
||||||
parser.add_argument("--landmark_stop", type=float, default=80.0)
|
|
||||||
parser.add_argument("--landmark_step", type=float, default=5.0)
|
|
||||||
parser.add_argument("--min_history_events", type=int, default=None)
|
|
||||||
parser.add_argument("--batch_size", type=int, default=None)
|
|
||||||
parser.add_argument(
|
|
||||||
"--attribution_batch_size",
|
|
||||||
type=int,
|
|
||||||
default=None,
|
|
||||||
help="Forward batch size for expanded extra-info ablation queries.",
|
|
||||||
)
|
|
||||||
parser.add_argument("--num_workers", type=int, default=None)
|
|
||||||
parser.add_argument(
|
|
||||||
"--cpu_reduce_workers",
|
|
||||||
type=int,
|
|
||||||
default=None,
|
|
||||||
help="Worker processes for CPU-side summary reduction. Defaults to --num_workers.",
|
|
||||||
)
|
|
||||||
parser.add_argument("--device", type=str, default=None)
|
|
||||||
return parser.parse_args()
|
|
||||||
|
|
||||||
|
|
||||||
def main() -> None:
|
|
||||||
args = parse_args()
|
|
||||||
# Dataset extra-info types must reproduce the checkpoint training config.
|
|
||||||
# --extra_info only filters which already-trained types are ablated.
|
|
||||||
args.extra_info_types = None
|
|
||||||
run_path = Path(args.run_path)
|
|
||||||
config_path = run_path / "train_config.json"
|
|
||||||
checkpoint_path = run_path / "best_model.pt"
|
|
||||||
if not config_path.exists():
|
|
||||||
raise FileNotFoundError(f"train_config.json not found: {config_path}")
|
|
||||||
if not checkpoint_path.exists():
|
|
||||||
raise FileNotFoundError(f"best_model.pt not found: {checkpoint_path}")
|
|
||||||
|
|
||||||
cfg = load_json_config(config_path)
|
|
||||||
model_target_mode = str(cfg.get("model_target_mode", "next_token")).lower()
|
|
||||||
if model_target_mode not in {"next_token", "all_future"}:
|
|
||||||
raise ValueError(f"Unsupported model_target_mode: {model_target_mode!r}")
|
|
||||||
|
|
||||||
target_mode = str(cfg.get("target_mode", "uts"))
|
|
||||||
attn_mask_mode = str(
|
|
||||||
cfg.get("attn_mask_mode", "non_strict_time" if target_mode == "uts" else "target_aware")
|
|
||||||
)
|
|
||||||
readout_name = str(
|
|
||||||
cfg.get("readout_name", "same_time_group_end" if target_mode == "uts" else "token")
|
|
||||||
)
|
|
||||||
readout_reduce = str(cfg.get("readout_reduce", "mean"))
|
|
||||||
|
|
||||||
dataset, subset_indices, eval_split, split_source = load_eval_sequence_dataset(args, cfg)
|
|
||||||
validate_dataset_metadata(dataset, cfg)
|
|
||||||
|
|
||||||
extra_metadata = load_extra_info_metadata(
|
|
||||||
dataset_extra_info_types=dataset.extra_info_types,
|
|
||||||
search_root=Path("."),
|
|
||||||
)
|
|
||||||
selected_extra_info_types = resolve_extra_info_types(
|
|
||||||
args.extra_info,
|
|
||||||
dataset_extra_info_types=dataset.extra_info_types,
|
|
||||||
metadata=extra_metadata,
|
|
||||||
)
|
|
||||||
if not selected_extra_info_types:
|
|
||||||
raise ValueError("No extra-info types selected for attribution")
|
|
||||||
|
|
||||||
landmark_ages = make_landmark_ages(
|
|
||||||
float(args.landmark_start),
|
|
||||||
float(args.landmark_stop),
|
|
||||||
float(args.landmark_step),
|
|
||||||
)
|
|
||||||
first_occurrence_by_token = build_first_occurrence_maps_for_landmarks(
|
|
||||||
dataset,
|
|
||||||
subset_indices,
|
|
||||||
)
|
|
||||||
death_idx = int(dataset.vocab_size) - 1
|
|
||||||
landmark_dataset = LandmarkDataset(
|
|
||||||
dataset=dataset,
|
|
||||||
subset_indices=subset_indices,
|
|
||||||
landmark_ages=landmark_ages,
|
|
||||||
attn_mask_mode=attn_mask_mode,
|
|
||||||
model_target_mode=model_target_mode,
|
|
||||||
min_history_events=int(cfg_get(args, cfg, "min_history_events", 1)),
|
|
||||||
first_occurrence_by_token=first_occurrence_by_token,
|
|
||||||
death_token_ids=[death_idx],
|
|
||||||
)
|
|
||||||
|
|
||||||
organ_groups, organ_labels, _token_to_group = load_organ_groups(
|
|
||||||
Path(args.organ_mapping_path),
|
|
||||||
vocab_size=int(dataset.vocab_size),
|
|
||||||
)
|
|
||||||
all_disease_tokens = sorted(
|
|
||||||
{
|
|
||||||
int(token)
|
|
||||||
for tokens in organ_groups.values()
|
|
||||||
for token in tokens
|
|
||||||
if int(token) != death_idx
|
|
||||||
}
|
|
||||||
)
|
|
||||||
risk_groups = {
|
|
||||||
"all_modeled_diseases": all_disease_tokens,
|
|
||||||
**{group: tokens for group, tokens in sorted(organ_groups.items())},
|
|
||||||
}
|
|
||||||
risk_group_labels = {
|
|
||||||
"all_modeled_diseases": "All modeled diseases",
|
|
||||||
**organ_labels,
|
|
||||||
}
|
|
||||||
group_names = list(risk_groups.keys())
|
|
||||||
group_labels = [str(risk_group_labels[group]) for group in group_names]
|
|
||||||
|
|
||||||
state_dict = load_checkpoint_state_dict(checkpoint_path, map_location="cpu")
|
|
||||||
dist_mode = resolve_dist_mode_for_checkpoint(str(cfg.get("dist_mode", "exponential")), state_dict)
|
|
||||||
death_distribution_name = "exponential" if dist_mode == "exponential" else "weibull"
|
|
||||||
cfg_model = dict(cfg)
|
|
||||||
cfg_model["dist_mode"] = dist_mode
|
|
||||||
device = resolve_eval_device(args.device)
|
|
||||||
model = build_model_from_dataset(
|
|
||||||
args, cfg_model, dataset, state_dict=state_dict
|
|
||||||
).to(device)
|
|
||||||
load_model_state(model, state_dict)
|
|
||||||
model.eval()
|
|
||||||
|
|
||||||
group_token_mask = torch.zeros(
|
|
||||||
(len(group_names), int(dataset.vocab_size)),
|
|
||||||
dtype=torch.float32,
|
|
||||||
device=device,
|
|
||||||
)
|
|
||||||
for group_idx, group in enumerate(group_names):
|
|
||||||
valid_tokens = [
|
|
||||||
int(token)
|
|
||||||
for token in risk_groups[group]
|
|
||||||
if 0 <= int(token) < int(dataset.vocab_size) and int(token) != death_idx
|
|
||||||
]
|
|
||||||
if valid_tokens:
|
|
||||||
group_token_mask[group_idx, torch.as_tensor(valid_tokens, dtype=torch.long, device=device)] = 1.0
|
|
||||||
|
|
||||||
batch_size = int(cfg_get(args, cfg, "batch_size", 128))
|
|
||||||
attribution_batch_size = int(
|
|
||||||
cfg_get(args, cfg, "attribution_batch_size", max(batch_size * 32, 4096))
|
|
||||||
)
|
|
||||||
if attribution_batch_size <= 0:
|
|
||||||
raise ValueError("attribution_batch_size must be positive")
|
|
||||||
|
|
||||||
num_workers = int(cfg_get(args, cfg, "num_workers", 4))
|
|
||||||
cpu_reduce_workers = int(
|
|
||||||
args.cpu_reduce_workers
|
|
||||||
if args.cpu_reduce_workers is not None
|
|
||||||
else max(1, num_workers)
|
|
||||||
)
|
|
||||||
if cpu_reduce_workers <= 0:
|
|
||||||
raise ValueError("--cpu_reduce_workers must be positive")
|
|
||||||
loader = DataLoader(
|
|
||||||
IndexedLandmarkDataset(landmark_dataset),
|
|
||||||
batch_size=batch_size,
|
|
||||||
shuffle=False,
|
|
||||||
collate_fn=collate_indexed_landmark_fn,
|
|
||||||
num_workers=num_workers,
|
|
||||||
pin_memory=device.type == "cuda",
|
|
||||||
persistent_workers=num_workers > 0,
|
|
||||||
prefetch_factor=2 if num_workers > 0 else None,
|
|
||||||
)
|
|
||||||
|
|
||||||
output_dir = (
|
|
||||||
Path(args.output_dir)
|
|
||||||
if args.output_dir
|
|
||||||
else run_path / f"extra_info_attribution_{eval_split}"
|
|
||||||
)
|
|
||||||
output_dir.mkdir(parents=True, exist_ok=True)
|
|
||||||
|
|
||||||
print(f"Eval split: {eval_split}")
|
|
||||||
print(f"Split source: {split_source}")
|
|
||||||
print(f"Selected patients: {len(subset_indices)}")
|
|
||||||
print(f"Landmark ages: {landmark_ages.tolist()}")
|
|
||||||
print(f"Dist mode: {dist_mode}")
|
|
||||||
print(f"Device: {device}")
|
|
||||||
print(f"Death token: {death_idx}")
|
|
||||||
print(f"Extra-info types: {selected_extra_info_types}")
|
|
||||||
print(f"Landmark rows: {len(landmark_dataset)}")
|
|
||||||
print(f"Attribution batch size: {attribution_batch_size}")
|
|
||||||
print(f"CPU reduce workers: {cpu_reduce_workers}")
|
|
||||||
print(f"Output directory: {output_dir}")
|
|
||||||
|
|
||||||
death_key_chunks: list[pd.DataFrame] = []
|
|
||||||
death_value_chunks: list[np.ndarray] = []
|
|
||||||
disease_stat_chunks: list[tuple[pd.DataFrame, np.ndarray, np.ndarray, np.ndarray]] = []
|
|
||||||
|
|
||||||
for batch in tqdm(loader, desc="Extra-info attribution", dynamic_ncols=True):
|
|
||||||
batch_dev = {
|
|
||||||
k: (v.to(device, non_blocking=True) if isinstance(v, torch.Tensor) else v)
|
|
||||||
for k, v in batch.items()
|
|
||||||
}
|
|
||||||
with torch.no_grad():
|
|
||||||
hidden = infer_landmark_hidden(
|
|
||||||
model=model,
|
|
||||||
batch=batch_dev,
|
|
||||||
device=device,
|
|
||||||
model_target_mode=model_target_mode,
|
|
||||||
readout_name=readout_name,
|
|
||||||
readout_reduce=readout_reduce,
|
|
||||||
)
|
|
||||||
_death_distribution, original_death_params = death_distribution_parameters(
|
|
||||||
model,
|
|
||||||
hidden,
|
|
||||||
dist_mode=dist_mode,
|
|
||||||
)
|
|
||||||
original_logits = model.calc_risk(hidden)
|
|
||||||
original_rho = model.calc_weibull_rho(hidden) if dist_mode == "weibull" else None
|
|
||||||
|
|
||||||
for ablated_batch, type_ids, local_rows in iter_extra_info_ablated_batches(
|
|
||||||
batch_dev,
|
|
||||||
selected_extra_info_types=selected_extra_info_types,
|
|
||||||
max_batch_size=attribution_batch_size,
|
|
||||||
):
|
|
||||||
row_tensor = torch.as_tensor(local_rows, dtype=torch.long, device=device)
|
|
||||||
with torch.no_grad():
|
|
||||||
ablated_hidden = infer_landmark_hidden(
|
|
||||||
model=model,
|
|
||||||
batch=ablated_batch,
|
|
||||||
device=device,
|
|
||||||
model_target_mode=model_target_mode,
|
|
||||||
readout_name=readout_name,
|
|
||||||
readout_reduce=readout_reduce,
|
|
||||||
)
|
|
||||||
_ablated_distribution, ablated_death_params = death_distribution_parameters(
|
|
||||||
model,
|
|
||||||
ablated_hidden,
|
|
||||||
dist_mode=dist_mode,
|
|
||||||
)
|
|
||||||
ablated_logits = model.calc_risk(ablated_hidden)
|
|
||||||
ablated_rho = model.calc_weibull_rho(ablated_hidden) if dist_mode == "weibull" else None
|
|
||||||
|
|
||||||
key_rows = []
|
|
||||||
for type_id, local_row in zip(type_ids, local_rows):
|
|
||||||
meta = extra_metadata[int(type_id)]
|
|
||||||
key_rows.append(
|
|
||||||
{
|
|
||||||
"selected_extra_info_type_id": int(type_id),
|
|
||||||
"selected_extra_info_var_name": str(meta.get("var_name", "")),
|
|
||||||
"selected_extra_info_full_name": str(meta.get("full_name", "")),
|
|
||||||
"landmark_age": float(batch["landmark_age"][int(local_row)].item()),
|
|
||||||
"sex": int(batch["sex"][int(local_row)].item()),
|
|
||||||
}
|
|
||||||
)
|
|
||||||
key_table = pd.DataFrame(key_rows, columns=EXTRA_KEY_COLUMNS)
|
|
||||||
value_block = parameter_pair_block(
|
|
||||||
original_death_params[row_tensor],
|
|
||||||
ablated_death_params,
|
|
||||||
).detach().cpu().numpy()
|
|
||||||
death_key_chunks.append(key_table)
|
|
||||||
death_value_chunks.append(value_block)
|
|
||||||
|
|
||||||
disease_values = all_disease_parameter_pair_block(
|
|
||||||
original_logits=original_logits[row_tensor],
|
|
||||||
ablated_logits=ablated_logits,
|
|
||||||
dist_mode=dist_mode,
|
|
||||||
original_rho=None if original_rho is None else original_rho[row_tensor],
|
|
||||||
ablated_rho=ablated_rho,
|
|
||||||
)
|
|
||||||
sums, sumsq, counts = grouped_parameter_stats(
|
|
||||||
disease_values,
|
|
||||||
group_token_mask,
|
|
||||||
)
|
|
||||||
disease_stat_chunks.append((key_table, sums, sumsq, counts))
|
|
||||||
|
|
||||||
death_summary, disease_parameter_summary = reduce_attribution_chunks(
|
|
||||||
death_key_chunks=death_key_chunks,
|
|
||||||
death_value_chunks=death_value_chunks,
|
|
||||||
disease_stat_chunks=disease_stat_chunks,
|
|
||||||
group_names=group_names,
|
|
||||||
group_labels=group_labels,
|
|
||||||
cpu_reduce_workers=cpu_reduce_workers,
|
|
||||||
)
|
|
||||||
|
|
||||||
death_summary_path = output_dir / "summary_extra_info_death_parameters.csv"
|
|
||||||
disease_summary_path = output_dir / "summary_extra_info_disease_parameters.csv"
|
|
||||||
death_rows = write_death_summary_csv(
|
|
||||||
death_summary_path,
|
|
||||||
death_summary,
|
|
||||||
death_distribution=death_distribution_name,
|
|
||||||
)
|
|
||||||
disease_rows = write_disease_parameter_summary_csv(
|
|
||||||
disease_summary_path,
|
|
||||||
disease_parameter_summary,
|
|
||||||
)
|
|
||||||
manifest = {
|
|
||||||
"death_summary_file": death_summary_path.name,
|
|
||||||
"disease_parameter_summary_file": disease_summary_path.name,
|
|
||||||
"death_summary_rows": int(death_rows),
|
|
||||||
"disease_parameter_summary_rows": int(disease_rows),
|
|
||||||
"eval_split": eval_split,
|
|
||||||
"split_source": split_source,
|
|
||||||
"dist_mode": dist_mode,
|
|
||||||
"landmark_start": float(args.landmark_start),
|
|
||||||
"landmark_stop": float(args.landmark_stop),
|
|
||||||
"landmark_step": float(args.landmark_step),
|
|
||||||
"selected_extra_info_types": [
|
|
||||||
extra_metadata[int(type_id)] for type_id in selected_extra_info_types
|
|
||||||
],
|
|
||||||
}
|
|
||||||
with (output_dir / "manifest.json").open("w", encoding="utf-8") as f:
|
|
||||||
json.dump(manifest, f, ensure_ascii=False, indent=2)
|
|
||||||
|
|
||||||
print(f"Wrote {death_rows} death summary rows to {death_summary_path}")
|
|
||||||
print(f"Wrote {disease_rows} disease-parameter summary rows to {disease_summary_path}")
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
main()
|
|
||||||
@@ -1,7 +0,0 @@
|
|||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from evaluate_auc_v2 import main
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
main()
|
|
||||||
@@ -1,819 +0,0 @@
|
|||||||
"""Compute per-disease attribution to predicted mortality distribution parameters.
|
|
||||||
|
|
||||||
For each selected patient and landmark age, this script keeps only rows where
|
|
||||||
each scanned disease token has already occurred in the history. It then deletes
|
|
||||||
that historical disease token, re-queries the model, and reports the original
|
|
||||||
and ablated fitted death distribution parameters. If --disease is omitted, all
|
|
||||||
disease tokens in the mapping are scanned.
|
|
||||||
|
|
||||||
Death is always token vocab_size - 1.
|
|
||||||
"""
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import argparse
|
|
||||||
import json
|
|
||||||
from pathlib import Path
|
|
||||||
from typing import Any, Dict
|
|
||||||
|
|
||||||
import numpy as np
|
|
||||||
import pandas as pd
|
|
||||||
import torch
|
|
||||||
import torch.nn.functional as F
|
|
||||||
from torch.utils.data import DataLoader
|
|
||||||
from tqdm.auto import tqdm
|
|
||||||
|
|
||||||
from evaluate_auc_v2 import (
|
|
||||||
build_model_from_dataset,
|
|
||||||
cfg_get,
|
|
||||||
load_checkpoint_state_dict,
|
|
||||||
load_json_config,
|
|
||||||
load_model_state,
|
|
||||||
resolve_dist_mode_for_checkpoint,
|
|
||||||
resolve_eval_device,
|
|
||||||
validate_dataset_metadata,
|
|
||||||
)
|
|
||||||
from landmark_eval_utils import (
|
|
||||||
IndexedLandmarkDataset,
|
|
||||||
LandmarkDataset,
|
|
||||||
build_first_occurrence_maps_for_landmarks,
|
|
||||||
collate_indexed_landmark_fn,
|
|
||||||
historical_counts_by_group,
|
|
||||||
infer_landmark_hidden,
|
|
||||||
load_eval_sequence_dataset,
|
|
||||||
load_organ_groups,
|
|
||||||
make_landmark_ages,
|
|
||||||
)
|
|
||||||
from targets import CHECKUP_IDX, PAD_IDX
|
|
||||||
|
|
||||||
|
|
||||||
OUTPUT_COLUMNS = [
|
|
||||||
"patient_id",
|
|
||||||
"dataset_index",
|
|
||||||
"eid",
|
|
||||||
"sex",
|
|
||||||
"landmark_age",
|
|
||||||
"followup_end_time",
|
|
||||||
"history_disease_count",
|
|
||||||
"selected_disease_history_count",
|
|
||||||
"selected_disease_token_id",
|
|
||||||
"selected_disease_code",
|
|
||||||
"selected_disease_name",
|
|
||||||
"selected_disease_organ_system",
|
|
||||||
"selected_disease_organ_system_label",
|
|
||||||
"history_count__selected_organ_system",
|
|
||||||
"death_distribution",
|
|
||||||
"original_death_lambda",
|
|
||||||
"ablated_death_lambda",
|
|
||||||
"original_death_scale",
|
|
||||||
"ablated_death_scale",
|
|
||||||
"original_death_shape",
|
|
||||||
"ablated_death_shape",
|
|
||||||
]
|
|
||||||
|
|
||||||
SUMMARY_KEY_COLUMNS = [
|
|
||||||
"selected_disease_token_id",
|
|
||||||
"selected_disease_code",
|
|
||||||
"selected_disease_name",
|
|
||||||
"selected_disease_organ_system",
|
|
||||||
"selected_disease_organ_system_label",
|
|
||||||
"landmark_age",
|
|
||||||
"sex",
|
|
||||||
]
|
|
||||||
|
|
||||||
SUMMARY_MEAN_COLUMNS = [
|
|
||||||
"history_disease_count",
|
|
||||||
"selected_disease_history_count",
|
|
||||||
"history_count__selected_organ_system",
|
|
||||||
]
|
|
||||||
|
|
||||||
SUMMARY_PARAMETER_COLUMNS = [
|
|
||||||
"original_death_lambda",
|
|
||||||
"ablated_death_lambda",
|
|
||||||
"original_death_scale",
|
|
||||||
"ablated_death_scale",
|
|
||||||
"original_death_shape",
|
|
||||||
"ablated_death_shape",
|
|
||||||
]
|
|
||||||
|
|
||||||
|
|
||||||
def write_compressed_npz_table(path: Path, table: pd.DataFrame) -> int:
|
|
||||||
table = table.reindex(columns=OUTPUT_COLUMNS)
|
|
||||||
arrays: dict[str, np.ndarray] = {
|
|
||||||
"__columns__": np.asarray(OUTPUT_COLUMNS, dtype="U"),
|
|
||||||
}
|
|
||||||
for column in OUTPUT_COLUMNS:
|
|
||||||
values = table[column] if column in table else pd.Series([], dtype=object)
|
|
||||||
if values.dtype == object:
|
|
||||||
arrays[column] = values.fillna("").astype(str).to_numpy(dtype="U")
|
|
||||||
else:
|
|
||||||
arrays[column] = values.to_numpy()
|
|
||||||
np.savez_compressed(path, **arrays)
|
|
||||||
return int(len(table))
|
|
||||||
|
|
||||||
|
|
||||||
def normalize_output_dir(path: Path) -> Path:
|
|
||||||
if path.suffix:
|
|
||||||
return path.with_suffix(path.suffix + "_shards")
|
|
||||||
return path
|
|
||||||
|
|
||||||
|
|
||||||
def write_manifest(
|
|
||||||
output_dir: Path,
|
|
||||||
*,
|
|
||||||
rows: int,
|
|
||||||
shards: list[dict[str, Any]],
|
|
||||||
summary_file: str,
|
|
||||||
scanned_diseases: list[dict[str, Any]],
|
|
||||||
eval_split: str,
|
|
||||||
dist_mode: str,
|
|
||||||
landmark_start: float,
|
|
||||||
landmark_stop: float,
|
|
||||||
landmark_step: float,
|
|
||||||
) -> None:
|
|
||||||
payload = {
|
|
||||||
"format": "compressed_npz_shards",
|
|
||||||
"columns": OUTPUT_COLUMNS,
|
|
||||||
"rows": int(rows),
|
|
||||||
"shards": shards,
|
|
||||||
"summary_file": summary_file,
|
|
||||||
"scanned_diseases": scanned_diseases,
|
|
||||||
"eval_split": eval_split,
|
|
||||||
"dist_mode": str(dist_mode),
|
|
||||||
"landmark_start": float(landmark_start),
|
|
||||||
"landmark_stop": float(landmark_stop),
|
|
||||||
"landmark_step": float(landmark_step),
|
|
||||||
}
|
|
||||||
with (output_dir / "manifest.json").open("w", encoding="utf-8") as f:
|
|
||||||
json.dump(payload, f, ensure_ascii=False, indent=2)
|
|
||||||
|
|
||||||
|
|
||||||
def update_summary_accumulator(
|
|
||||||
summary: dict[tuple[Any, ...], dict[str, float]],
|
|
||||||
table: pd.DataFrame,
|
|
||||||
) -> None:
|
|
||||||
if table.empty:
|
|
||||||
return
|
|
||||||
grouped = table.groupby(SUMMARY_KEY_COLUMNS, dropna=False, sort=False)
|
|
||||||
for key, group in grouped:
|
|
||||||
if not isinstance(key, tuple):
|
|
||||||
key = (key,)
|
|
||||||
acc = summary.setdefault(
|
|
||||||
key,
|
|
||||||
{
|
|
||||||
"n": 0.0,
|
|
||||||
**{column: 0.0 for column in SUMMARY_MEAN_COLUMNS},
|
|
||||||
**{f"count__{column}": 0.0 for column in SUMMARY_PARAMETER_COLUMNS},
|
|
||||||
**{f"sum__{column}": 0.0 for column in SUMMARY_PARAMETER_COLUMNS},
|
|
||||||
**{f"sumsq__{column}": 0.0 for column in SUMMARY_PARAMETER_COLUMNS},
|
|
||||||
},
|
|
||||||
)
|
|
||||||
n = int(len(group))
|
|
||||||
acc["n"] += float(n)
|
|
||||||
for column in SUMMARY_MEAN_COLUMNS:
|
|
||||||
acc[column] += float(pd.to_numeric(group[column], errors="coerce").sum())
|
|
||||||
for column in SUMMARY_PARAMETER_COLUMNS:
|
|
||||||
values = pd.to_numeric(group[column], errors="coerce").dropna()
|
|
||||||
acc[f"count__{column}"] += float(len(values))
|
|
||||||
acc[f"sum__{column}"] += float(values.sum())
|
|
||||||
acc[f"sumsq__{column}"] += float((values * values).sum())
|
|
||||||
|
|
||||||
|
|
||||||
def write_summary_csv(
|
|
||||||
path: Path,
|
|
||||||
summary: dict[tuple[Any, ...], dict[str, float]],
|
|
||||||
) -> int:
|
|
||||||
rows: list[dict[str, Any]] = []
|
|
||||||
for key, acc in summary.items():
|
|
||||||
n = int(acc["n"])
|
|
||||||
out = {column: value for column, value in zip(SUMMARY_KEY_COLUMNS, key)}
|
|
||||||
out["n"] = n
|
|
||||||
for column in SUMMARY_MEAN_COLUMNS:
|
|
||||||
out[f"mean__{column}"] = acc[column] / n if n > 0 else np.nan
|
|
||||||
for column in SUMMARY_PARAMETER_COLUMNS:
|
|
||||||
count = int(acc[f"count__{column}"])
|
|
||||||
mean = acc[f"sum__{column}"] / count if count > 0 else np.nan
|
|
||||||
second = acc[f"sumsq__{column}"] / count if count > 0 else np.nan
|
|
||||||
out[f"mean__{column}"] = mean
|
|
||||||
out[f"var__{column}"] = second - mean * mean if count > 0 else np.nan
|
|
||||||
rows.append(out)
|
|
||||||
|
|
||||||
columns = [
|
|
||||||
*SUMMARY_KEY_COLUMNS,
|
|
||||||
"n",
|
|
||||||
*[f"mean__{column}" for column in SUMMARY_MEAN_COLUMNS],
|
|
||||||
*[
|
|
||||||
name
|
|
||||||
for column in SUMMARY_PARAMETER_COLUMNS
|
|
||||||
for name in (f"mean__{column}", f"var__{column}")
|
|
||||||
],
|
|
||||||
]
|
|
||||||
pd.DataFrame(rows, columns=columns).sort_values(
|
|
||||||
["selected_disease_token_id", "landmark_age", "sex"],
|
|
||||||
kind="mergesort",
|
|
||||||
).to_csv(path, index=False)
|
|
||||||
return len(rows)
|
|
||||||
|
|
||||||
|
|
||||||
def build_disease_ablated_slice(
|
|
||||||
batch: Dict[str, torch.Tensor],
|
|
||||||
row_indices: torch.Tensor,
|
|
||||||
token_ids: torch.Tensor,
|
|
||||||
) -> Dict[str, torch.Tensor]:
|
|
||||||
"""Build an ablated slice for aligned (row, disease_token) pairs."""
|
|
||||||
event_seq = batch["event_seq"]
|
|
||||||
row_indices = row_indices.to(device=event_seq.device, dtype=torch.long)
|
|
||||||
token_ids = token_ids.to(device=event_seq.device, dtype=event_seq.dtype)
|
|
||||||
|
|
||||||
out: Dict[str, torch.Tensor] = {}
|
|
||||||
out["event_seq"] = event_seq[row_indices].clone()
|
|
||||||
out["time_seq"] = batch["time_seq"][row_indices]
|
|
||||||
out["readout_mask"] = batch["readout_mask"][row_indices].clone()
|
|
||||||
out["padding_mask"] = batch["padding_mask"][row_indices].bool().clone()
|
|
||||||
out["landmark_pos"] = batch["landmark_pos"][row_indices].clone()
|
|
||||||
|
|
||||||
seq_len = int(event_seq.shape[1])
|
|
||||||
positions = torch.arange(seq_len, device=event_seq.device)[None, :]
|
|
||||||
remove = (out["event_seq"] == token_ids[:, None]) & out["padding_mask"]
|
|
||||||
out["event_seq"] = torch.where(
|
|
||||||
remove,
|
|
||||||
torch.full_like(out["event_seq"], PAD_IDX),
|
|
||||||
out["event_seq"],
|
|
||||||
)
|
|
||||||
out["padding_mask"] &= ~remove
|
|
||||||
out["readout_mask"] &= ~remove
|
|
||||||
|
|
||||||
has_valid = out["padding_mask"].any(dim=1)
|
|
||||||
empty_rows = ~has_valid
|
|
||||||
out["event_seq"][empty_rows, 0] = CHECKUP_IDX
|
|
||||||
out["time_seq"][empty_rows, 0] = batch["t_query"][row_indices[empty_rows]].to(
|
|
||||||
dtype=out["time_seq"].dtype
|
|
||||||
)
|
|
||||||
out["padding_mask"][empty_rows, 0] = True
|
|
||||||
out["readout_mask"][empty_rows, 0] = True
|
|
||||||
out["landmark_pos"][empty_rows] = 0
|
|
||||||
|
|
||||||
has_readout = out["readout_mask"].any(dim=1)
|
|
||||||
missing_readout = ~has_readout
|
|
||||||
local_valid = out["padding_mask"]
|
|
||||||
last_pos = torch.where(
|
|
||||||
local_valid,
|
|
||||||
positions.expand(local_valid.shape[0], -1),
|
|
||||||
torch.zeros_like(positions.expand(local_valid.shape[0], -1)),
|
|
||||||
).amax(dim=1)
|
|
||||||
out["readout_mask"][missing_readout] = False
|
|
||||||
out["readout_mask"][missing_readout, last_pos[missing_readout]] = True
|
|
||||||
out["landmark_pos"][missing_readout] = last_pos[missing_readout].to(
|
|
||||||
dtype=out["landmark_pos"].dtype
|
|
||||||
)
|
|
||||||
|
|
||||||
repeated_keys = (
|
|
||||||
"sex",
|
|
||||||
"other_type",
|
|
||||||
"other_value",
|
|
||||||
"other_value_kind",
|
|
||||||
"other_time",
|
|
||||||
"t_query",
|
|
||||||
"patient_id",
|
|
||||||
"landmark_age",
|
|
||||||
"followup_end_time",
|
|
||||||
"death_time",
|
|
||||||
"row_idx",
|
|
||||||
)
|
|
||||||
for key in repeated_keys:
|
|
||||||
out[key] = batch[key][row_indices]
|
|
||||||
return out
|
|
||||||
|
|
||||||
|
|
||||||
def load_disease_metadata(
|
|
||||||
mapping_path: Path,
|
|
||||||
*,
|
|
||||||
vocab_size: int,
|
|
||||||
) -> dict[int, dict[str, Any]]:
|
|
||||||
if not mapping_path.exists():
|
|
||||||
raise FileNotFoundError(f"Disease mapping file not found: {mapping_path}")
|
|
||||||
table = pd.read_csv(mapping_path)
|
|
||||||
required = {"token_id", "code", "name", "is_death"}
|
|
||||||
missing = required - set(table.columns)
|
|
||||||
if missing:
|
|
||||||
raise ValueError(f"{mapping_path} is missing columns: {sorted(missing)}")
|
|
||||||
|
|
||||||
death_idx = int(vocab_size) - 1
|
|
||||||
out: dict[int, dict[str, Any]] = {}
|
|
||||||
for row in table.itertuples(index=False):
|
|
||||||
token = int(getattr(row, "token_id"))
|
|
||||||
if token < 0 or token >= int(vocab_size) or token == death_idx:
|
|
||||||
continue
|
|
||||||
if int(getattr(row, "is_death")) == 1:
|
|
||||||
continue
|
|
||||||
meta = {
|
|
||||||
"token_id": token,
|
|
||||||
"code": str(getattr(row, "code")),
|
|
||||||
"name": str(getattr(row, "name")),
|
|
||||||
}
|
|
||||||
for column in (
|
|
||||||
"icd10_chapter",
|
|
||||||
"icd10_chapter_title",
|
|
||||||
"organ_system",
|
|
||||||
"organ_system_label",
|
|
||||||
):
|
|
||||||
if hasattr(row, column):
|
|
||||||
meta[column] = str(getattr(row, column))
|
|
||||||
out[token] = meta
|
|
||||||
return out
|
|
||||||
|
|
||||||
|
|
||||||
def resolve_disease_token(
|
|
||||||
value: str,
|
|
||||||
metadata: dict[int, dict[str, Any]],
|
|
||||||
) -> tuple[int, dict[str, Any]]:
|
|
||||||
text = str(value).strip()
|
|
||||||
if text == "":
|
|
||||||
raise ValueError("--disease must not be empty")
|
|
||||||
|
|
||||||
if text.isdigit() or (text.startswith("-") and text[1:].isdigit()):
|
|
||||||
token = int(text)
|
|
||||||
if token not in metadata:
|
|
||||||
raise ValueError(f"Disease token_id {token} was not found in the mapping")
|
|
||||||
return token, metadata[token]
|
|
||||||
|
|
||||||
lower = text.lower()
|
|
||||||
exact = [
|
|
||||||
(token, meta)
|
|
||||||
for token, meta in metadata.items()
|
|
||||||
if str(meta.get("code", "")).lower() == lower
|
|
||||||
or str(meta.get("name", "")).lower() == lower
|
|
||||||
]
|
|
||||||
if len(exact) == 1:
|
|
||||||
return exact[0]
|
|
||||||
if len(exact) > 1:
|
|
||||||
raise ValueError(f"--disease={value!r} matched multiple diseases exactly")
|
|
||||||
|
|
||||||
contains = [
|
|
||||||
(token, meta)
|
|
||||||
for token, meta in metadata.items()
|
|
||||||
if lower in str(meta.get("code", "")).lower()
|
|
||||||
or lower in str(meta.get("name", "")).lower()
|
|
||||||
]
|
|
||||||
if len(contains) == 1:
|
|
||||||
return contains[0]
|
|
||||||
if not contains:
|
|
||||||
raise ValueError(f"--disease={value!r} did not match any disease token")
|
|
||||||
preview = ", ".join(
|
|
||||||
f"{token}:{meta.get('code')} ({meta.get('name')})"
|
|
||||||
for token, meta in contains[:10]
|
|
||||||
)
|
|
||||||
raise ValueError(
|
|
||||||
f"--disease={value!r} matched {len(contains)} diseases; use token_id or code. "
|
|
||||||
f"First matches: {preview}"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def resolve_disease_tokens(
|
|
||||||
value: str | None,
|
|
||||||
metadata: dict[int, dict[str, Any]],
|
|
||||||
) -> list[tuple[int, dict[str, Any]]]:
|
|
||||||
if value is None or str(value).strip() == "":
|
|
||||||
return [(token, metadata[token]) for token in sorted(metadata)]
|
|
||||||
out: list[tuple[int, dict[str, Any]]] = []
|
|
||||||
seen: set[int] = set()
|
|
||||||
for part in str(value).split(","):
|
|
||||||
token, meta = resolve_disease_token(part, metadata)
|
|
||||||
if token not in seen:
|
|
||||||
out.append((token, meta))
|
|
||||||
seen.add(token)
|
|
||||||
return out
|
|
||||||
|
|
||||||
|
|
||||||
def death_distribution_parameters(
|
|
||||||
model,
|
|
||||||
hidden: torch.Tensor,
|
|
||||||
*,
|
|
||||||
dist_mode: str,
|
|
||||||
eps: float = 1e-8,
|
|
||||||
) -> tuple[str, torch.Tensor]:
|
|
||||||
"""Return death distribution parameters with columns matching PARAMETER_VALUE_COLUMNS."""
|
|
||||||
logits = model.calc_risk(hidden)
|
|
||||||
death_idx = int(logits.shape[1]) - 1
|
|
||||||
death_lambda = F.softplus(logits[:, death_idx]) + float(eps)
|
|
||||||
|
|
||||||
if dist_mode == "exponential":
|
|
||||||
nan = torch.full_like(death_lambda, float("nan"))
|
|
||||||
return "exponential", torch.stack([death_lambda, nan, nan], dim=1)
|
|
||||||
|
|
||||||
if dist_mode == "weibull":
|
|
||||||
rho = model.calc_weibull_rho(hidden)[:, death_idx].to(dtype=death_lambda.dtype)
|
|
||||||
elif dist_mode == "mixed":
|
|
||||||
rho = model.calc_death_rho(hidden).to(dtype=death_lambda.dtype)
|
|
||||||
else:
|
|
||||||
raise ValueError(f"Unsupported dist_mode={dist_mode!r}")
|
|
||||||
|
|
||||||
shape = rho.clamp_min(float(eps))
|
|
||||||
scale = torch.pow(death_lambda.clamp_min(float(eps)), -1.0 / shape)
|
|
||||||
nan = torch.full_like(death_lambda, float("nan"))
|
|
||||||
return "weibull", torch.stack([nan, scale, shape], dim=1)
|
|
||||||
|
|
||||||
|
|
||||||
def parameter_pair_block(original: torch.Tensor, ablated: torch.Tensor) -> torch.Tensor:
|
|
||||||
return torch.stack(
|
|
||||||
[
|
|
||||||
original[:, 0],
|
|
||||||
ablated[:, 0],
|
|
||||||
original[:, 1],
|
|
||||||
ablated[:, 1],
|
|
||||||
original[:, 2],
|
|
||||||
ablated[:, 2],
|
|
||||||
],
|
|
||||||
dim=1,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def output_name_for_run(run_path: Path, eval_split: str, *, all_diseases: bool) -> Path:
|
|
||||||
scope = "all_diseases" if all_diseases else "selected_diseases"
|
|
||||||
return run_path / f"single_disease_mortality_parameters_{eval_split}_{scope}"
|
|
||||||
|
|
||||||
|
|
||||||
def parse_args() -> argparse.Namespace:
|
|
||||||
parser = argparse.ArgumentParser(
|
|
||||||
description="Compute per-disease model attribution to mortality distribution parameters."
|
|
||||||
)
|
|
||||||
parser.add_argument("--run_path", type=str, required=True)
|
|
||||||
parser.add_argument(
|
|
||||||
"--disease",
|
|
||||||
type=str,
|
|
||||||
default=None,
|
|
||||||
help=(
|
|
||||||
"Optional disease token_id, ICD-10 code, exact name, unambiguous name "
|
|
||||||
"substring, or comma-separated list. If omitted, scan all disease tokens."
|
|
||||||
),
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
|
||||||
"--output_path",
|
|
||||||
type=str,
|
|
||||||
default=None,
|
|
||||||
help="Output directory for compressed .npz shards.",
|
|
||||||
)
|
|
||||||
parser.add_argument("--organ_mapping_path", type=str, default="icd10_chapter_organ_mapping.csv")
|
|
||||||
parser.add_argument("--eval_split", type=str, default=None)
|
|
||||||
parser.add_argument("--dataset_subset_size", type=int, default=None)
|
|
||||||
parser.add_argument("--train_eid_file", type=str, default=None)
|
|
||||||
parser.add_argument("--val_eid_file", type=str, default=None)
|
|
||||||
parser.add_argument("--test_eid_file", type=str, default=None)
|
|
||||||
parser.add_argument("--landmark_start", type=float, default=40.0)
|
|
||||||
parser.add_argument("--landmark_stop", type=float, default=80.0)
|
|
||||||
parser.add_argument("--landmark_step", type=float, default=5.0)
|
|
||||||
parser.add_argument("--min_history_events", type=int, default=None)
|
|
||||||
parser.add_argument("--batch_size", type=int, default=None)
|
|
||||||
parser.add_argument(
|
|
||||||
"--attribution_batch_size",
|
|
||||||
type=int,
|
|
||||||
default=None,
|
|
||||||
help="Forward batch size for disease-token ablation queries.",
|
|
||||||
)
|
|
||||||
parser.add_argument("--num_workers", type=int, default=None)
|
|
||||||
parser.add_argument("--device", type=str, default=None)
|
|
||||||
parser.add_argument("--extra_info_types", type=str, default=None)
|
|
||||||
parser.add_argument(
|
|
||||||
"--shard_rows",
|
|
||||||
type=int,
|
|
||||||
default=200_000,
|
|
||||||
help="Approximate number of detailed rows to buffer before writing one .npz shard.",
|
|
||||||
)
|
|
||||||
return parser.parse_args()
|
|
||||||
|
|
||||||
|
|
||||||
def main() -> None:
|
|
||||||
args = parse_args()
|
|
||||||
run_path = Path(args.run_path)
|
|
||||||
config_path = run_path / "train_config.json"
|
|
||||||
checkpoint_path = run_path / "best_model.pt"
|
|
||||||
if not config_path.exists():
|
|
||||||
raise FileNotFoundError(f"train_config.json not found: {config_path}")
|
|
||||||
if not checkpoint_path.exists():
|
|
||||||
raise FileNotFoundError(f"best_model.pt not found: {checkpoint_path}")
|
|
||||||
|
|
||||||
cfg = load_json_config(config_path)
|
|
||||||
model_target_mode = str(cfg.get("model_target_mode", "next_token")).lower()
|
|
||||||
if model_target_mode not in {"next_token", "all_future"}:
|
|
||||||
raise ValueError(f"Unsupported model_target_mode: {model_target_mode!r}")
|
|
||||||
|
|
||||||
target_mode = str(cfg.get("target_mode", "uts"))
|
|
||||||
attn_mask_mode = str(
|
|
||||||
cfg.get("attn_mask_mode", "non_strict_time" if target_mode == "uts" else "target_aware")
|
|
||||||
)
|
|
||||||
readout_name = str(
|
|
||||||
cfg.get("readout_name", "same_time_group_end" if target_mode == "uts" else "token")
|
|
||||||
)
|
|
||||||
readout_reduce = str(cfg.get("readout_reduce", "mean"))
|
|
||||||
|
|
||||||
dataset, subset_indices, eval_split, split_source = load_eval_sequence_dataset(args, cfg)
|
|
||||||
validate_dataset_metadata(dataset, cfg)
|
|
||||||
|
|
||||||
metadata = load_disease_metadata(
|
|
||||||
Path(args.organ_mapping_path),
|
|
||||||
vocab_size=int(dataset.vocab_size),
|
|
||||||
)
|
|
||||||
scanned_disease_items = resolve_disease_tokens(args.disease, metadata)
|
|
||||||
if not scanned_disease_items:
|
|
||||||
raise ValueError("No diseases selected for attribution")
|
|
||||||
scanned_disease_tokens = [token for token, _meta in scanned_disease_items]
|
|
||||||
|
|
||||||
landmark_ages = make_landmark_ages(
|
|
||||||
float(args.landmark_start),
|
|
||||||
float(args.landmark_stop),
|
|
||||||
float(args.landmark_step),
|
|
||||||
)
|
|
||||||
|
|
||||||
first_occurrence_by_token = build_first_occurrence_maps_for_landmarks(
|
|
||||||
dataset,
|
|
||||||
subset_indices,
|
|
||||||
)
|
|
||||||
death_idx = int(dataset.vocab_size) - 1
|
|
||||||
landmark_dataset = LandmarkDataset(
|
|
||||||
dataset=dataset,
|
|
||||||
subset_indices=subset_indices,
|
|
||||||
landmark_ages=landmark_ages,
|
|
||||||
attn_mask_mode=attn_mask_mode,
|
|
||||||
model_target_mode=model_target_mode,
|
|
||||||
min_history_events=int(cfg_get(args, cfg, "min_history_events", 1)),
|
|
||||||
first_occurrence_by_token=first_occurrence_by_token,
|
|
||||||
death_token_ids=[death_idx],
|
|
||||||
)
|
|
||||||
|
|
||||||
organ_groups, _organ_labels, token_to_group = load_organ_groups(
|
|
||||||
Path(args.organ_mapping_path),
|
|
||||||
vocab_size=int(dataset.vocab_size),
|
|
||||||
)
|
|
||||||
group_names = sorted(organ_groups)
|
|
||||||
|
|
||||||
state_dict = load_checkpoint_state_dict(checkpoint_path, map_location="cpu")
|
|
||||||
dist_mode = resolve_dist_mode_for_checkpoint(str(cfg.get("dist_mode", "exponential")), state_dict)
|
|
||||||
death_distribution_name = "exponential" if dist_mode == "exponential" else "weibull"
|
|
||||||
cfg_model = dict(cfg)
|
|
||||||
cfg_model["dist_mode"] = dist_mode
|
|
||||||
device = resolve_eval_device(args.device)
|
|
||||||
selected_token_mask = np.zeros(int(dataset.vocab_size), dtype=bool)
|
|
||||||
selected_token_mask[np.asarray(scanned_disease_tokens, dtype=np.int64)] = True
|
|
||||||
model = build_model_from_dataset(
|
|
||||||
args, cfg_model, dataset, state_dict=state_dict
|
|
||||||
).to(device)
|
|
||||||
load_model_state(model, state_dict)
|
|
||||||
model.eval()
|
|
||||||
|
|
||||||
batch_size = int(cfg_get(args, cfg, "batch_size", 128))
|
|
||||||
attribution_batch_size = int(
|
|
||||||
cfg_get(args, cfg, "attribution_batch_size", max(batch_size * 32, 4096))
|
|
||||||
)
|
|
||||||
if attribution_batch_size <= 0:
|
|
||||||
raise ValueError("attribution_batch_size must be positive")
|
|
||||||
if int(args.shard_rows) <= 0:
|
|
||||||
raise ValueError("--shard_rows must be positive")
|
|
||||||
|
|
||||||
num_workers = int(cfg_get(args, cfg, "num_workers", 4))
|
|
||||||
loader = DataLoader(
|
|
||||||
IndexedLandmarkDataset(landmark_dataset),
|
|
||||||
batch_size=batch_size,
|
|
||||||
shuffle=False,
|
|
||||||
collate_fn=collate_indexed_landmark_fn,
|
|
||||||
num_workers=num_workers,
|
|
||||||
pin_memory=device.type == "cuda",
|
|
||||||
persistent_workers=num_workers > 0,
|
|
||||||
prefetch_factor=2 if num_workers > 0 else None,
|
|
||||||
)
|
|
||||||
|
|
||||||
output_path = (
|
|
||||||
Path(args.output_path)
|
|
||||||
if args.output_path
|
|
||||||
else output_name_for_run(
|
|
||||||
run_path,
|
|
||||||
eval_split,
|
|
||||||
all_diseases=args.disease is None or str(args.disease).strip() == "",
|
|
||||||
)
|
|
||||||
)
|
|
||||||
output_dir = normalize_output_dir(output_path)
|
|
||||||
output_dir.mkdir(parents=True, exist_ok=True)
|
|
||||||
|
|
||||||
print(f"Eval split: {eval_split}")
|
|
||||||
print(f"Split source: {split_source}")
|
|
||||||
print(f"Selected patients: {len(subset_indices)}")
|
|
||||||
print(f"Landmark ages: {landmark_ages.tolist()}")
|
|
||||||
print(f"Dist mode: {dist_mode}")
|
|
||||||
print(f"Device: {device}")
|
|
||||||
print(f"Death token: {death_idx}")
|
|
||||||
if len(scanned_disease_items) == len(metadata):
|
|
||||||
print(f"Diseases: all mapped diseases ({len(scanned_disease_items)})")
|
|
||||||
else:
|
|
||||||
preview = ", ".join(
|
|
||||||
f"{token}:{meta.get('code')}" for token, meta in scanned_disease_items[:10]
|
|
||||||
)
|
|
||||||
print(f"Diseases: {len(scanned_disease_items)} selected ({preview})")
|
|
||||||
print(f"Landmark rows: {len(landmark_dataset)}")
|
|
||||||
print(f"Attribution batch size: {attribution_batch_size}")
|
|
||||||
print(f"Output directory: {output_dir}")
|
|
||||||
|
|
||||||
written_rows = 0
|
|
||||||
shard_index = 0
|
|
||||||
shards: list[dict[str, Any]] = []
|
|
||||||
row_base_cache: dict[int, dict[str, Any]] = {}
|
|
||||||
result_row_idx_chunks: list[np.ndarray] = []
|
|
||||||
result_disease_token_chunks: list[np.ndarray] = []
|
|
||||||
result_value_chunks: list[np.ndarray] = []
|
|
||||||
|
|
||||||
def get_row_base(row_idx: int) -> dict[str, Any]:
|
|
||||||
cached = row_base_cache.get(row_idx)
|
|
||||||
if cached is not None:
|
|
||||||
return cached
|
|
||||||
|
|
||||||
meta = landmark_dataset.rows[int(row_idx)]
|
|
||||||
dataset_index = int(meta["dataset_index"])
|
|
||||||
sample = dataset.samples[dataset_index]
|
|
||||||
hist_tokens = np.asarray(meta["event_seq"], dtype=np.int64)
|
|
||||||
unique_tokens, token_counts = np.unique(hist_tokens, return_counts=True)
|
|
||||||
total_count, group_counts = historical_counts_by_group(
|
|
||||||
hist_tokens,
|
|
||||||
death_idx=death_idx,
|
|
||||||
token_to_group=token_to_group,
|
|
||||||
group_names=group_names,
|
|
||||||
)
|
|
||||||
cached = {
|
|
||||||
"patient_id": int(meta["patient_id"]),
|
|
||||||
"dataset_index": dataset_index,
|
|
||||||
"eid": int(sample.get("eid", -1)),
|
|
||||||
"sex": int(meta["sex"]),
|
|
||||||
"landmark_age": float(meta["landmark_age"]),
|
|
||||||
"followup_end_time": float(meta["followup_end_time"]),
|
|
||||||
"history_disease_count": int(total_count),
|
|
||||||
"_hist_tokens": hist_tokens,
|
|
||||||
"_token_counts": {
|
|
||||||
int(token): int(count)
|
|
||||||
for token, count in zip(unique_tokens.tolist(), token_counts.tolist())
|
|
||||||
},
|
|
||||||
"_group_counts": group_counts,
|
|
||||||
}
|
|
||||||
row_base_cache[row_idx] = cached
|
|
||||||
return cached
|
|
||||||
|
|
||||||
for batch in tqdm(loader, desc="Per-disease mortality attribution", dynamic_ncols=True):
|
|
||||||
batch_dev = {
|
|
||||||
k: (v.to(device, non_blocking=True) if isinstance(v, torch.Tensor) else v)
|
|
||||||
for k, v in batch.items()
|
|
||||||
}
|
|
||||||
hidden = infer_landmark_hidden(
|
|
||||||
model=model,
|
|
||||||
batch=batch_dev,
|
|
||||||
device=device,
|
|
||||||
model_target_mode=model_target_mode,
|
|
||||||
readout_name=readout_name,
|
|
||||||
readout_reduce=readout_reduce,
|
|
||||||
)
|
|
||||||
with torch.no_grad():
|
|
||||||
_death_distribution, original_params = death_distribution_parameters(
|
|
||||||
model,
|
|
||||||
hidden,
|
|
||||||
dist_mode=dist_mode,
|
|
||||||
)
|
|
||||||
event_np = batch["event_seq"].numpy()
|
|
||||||
valid_event = (event_np >= 0) & (event_np < int(dataset.vocab_size))
|
|
||||||
selected_event = np.zeros_like(valid_event, dtype=bool)
|
|
||||||
selected_event[valid_event] = selected_token_mask[event_np[valid_event]]
|
|
||||||
pair_row_np, pair_pos_np = np.nonzero(selected_event)
|
|
||||||
if pair_row_np.size == 0:
|
|
||||||
continue
|
|
||||||
pair_disease_np = event_np[pair_row_np, pair_pos_np].astype(np.int64, copy=False)
|
|
||||||
|
|
||||||
pair_offset = 0
|
|
||||||
while pair_offset < int(pair_row_np.shape[0]):
|
|
||||||
pair_stop = min(int(pair_row_np.shape[0]), pair_offset + int(attribution_batch_size))
|
|
||||||
local_rows_np = pair_row_np[pair_offset:pair_stop].astype(np.int64, copy=False)
|
|
||||||
disease_tokens_np = pair_disease_np[pair_offset:pair_stop]
|
|
||||||
local_rows = torch.as_tensor(local_rows_np, dtype=torch.long, device=device)
|
|
||||||
disease_token_ids = torch.as_tensor(disease_tokens_np, dtype=torch.long, device=device)
|
|
||||||
ablated_chunk = build_disease_ablated_slice(
|
|
||||||
batch=batch_dev,
|
|
||||||
row_indices=local_rows,
|
|
||||||
token_ids=disease_token_ids,
|
|
||||||
)
|
|
||||||
with torch.no_grad():
|
|
||||||
ablated_hidden = infer_landmark_hidden(
|
|
||||||
model=model,
|
|
||||||
batch=ablated_chunk,
|
|
||||||
device=device,
|
|
||||||
model_target_mode=model_target_mode,
|
|
||||||
readout_name=readout_name,
|
|
||||||
readout_reduce=readout_reduce,
|
|
||||||
)
|
|
||||||
_ablated_distribution, ablated_params = death_distribution_parameters(
|
|
||||||
model,
|
|
||||||
ablated_hidden,
|
|
||||||
dist_mode=dist_mode,
|
|
||||||
)
|
|
||||||
value_block = parameter_pair_block(
|
|
||||||
original_params[local_rows],
|
|
||||||
ablated_params,
|
|
||||||
).detach().cpu().numpy()
|
|
||||||
row_ids = batch["row_idx"][local_rows_np].numpy().astype(np.int64, copy=False)
|
|
||||||
disease_tokens_list = disease_tokens_np
|
|
||||||
result_row_idx_chunks.append(row_ids)
|
|
||||||
result_disease_token_chunks.append(disease_tokens_list)
|
|
||||||
result_value_chunks.append(value_block)
|
|
||||||
pair_offset = pair_stop
|
|
||||||
|
|
||||||
if result_value_chunks:
|
|
||||||
all_row_ids = np.concatenate(result_row_idx_chunks).astype(np.int64, copy=False)
|
|
||||||
all_disease_tokens = np.concatenate(result_disease_token_chunks).astype(
|
|
||||||
np.int64,
|
|
||||||
copy=False,
|
|
||||||
)
|
|
||||||
all_values = np.concatenate(result_value_chunks, axis=0)
|
|
||||||
|
|
||||||
rows: list[dict[str, Any]] = []
|
|
||||||
for i, (row_idx, disease_token) in enumerate(
|
|
||||||
zip(all_row_ids.tolist(), all_disease_tokens.tolist())
|
|
||||||
):
|
|
||||||
disease_token = int(disease_token)
|
|
||||||
disease_meta = metadata[disease_token]
|
|
||||||
row_base = get_row_base(int(row_idx))
|
|
||||||
group_counts = row_base["_group_counts"]
|
|
||||||
disease_history_count = int(row_base["_token_counts"].get(disease_token, 0))
|
|
||||||
if disease_history_count <= 0:
|
|
||||||
raise RuntimeError(
|
|
||||||
"Internal mismatch: occurred mask selected disease "
|
|
||||||
f"{disease_token} for row {row_idx}, but cached history has count 0"
|
|
||||||
)
|
|
||||||
|
|
||||||
rows.append(
|
|
||||||
{
|
|
||||||
"patient_id": row_base["patient_id"],
|
|
||||||
"dataset_index": row_base["dataset_index"],
|
|
||||||
"eid": row_base["eid"],
|
|
||||||
"sex": row_base["sex"],
|
|
||||||
"landmark_age": row_base["landmark_age"],
|
|
||||||
"followup_end_time": row_base["followup_end_time"],
|
|
||||||
"history_disease_count": row_base["history_disease_count"],
|
|
||||||
"selected_disease_history_count": disease_history_count,
|
|
||||||
"selected_disease_token_id": int(disease_token),
|
|
||||||
"selected_disease_code": str(disease_meta.get("code", "")),
|
|
||||||
"selected_disease_name": str(disease_meta.get("name", "")),
|
|
||||||
"selected_disease_organ_system": str(disease_meta.get("organ_system", "")),
|
|
||||||
"selected_disease_organ_system_label": str(
|
|
||||||
disease_meta.get("organ_system_label", "")
|
|
||||||
),
|
|
||||||
"history_count__selected_organ_system": int(
|
|
||||||
group_counts.get(str(disease_meta.get("organ_system", "")), 0)
|
|
||||||
),
|
|
||||||
"death_distribution": death_distribution_name,
|
|
||||||
"original_death_lambda": float(all_values[i, 0]),
|
|
||||||
"ablated_death_lambda": float(all_values[i, 1]),
|
|
||||||
"original_death_scale": float(all_values[i, 2]),
|
|
||||||
"ablated_death_scale": float(all_values[i, 3]),
|
|
||||||
"original_death_shape": float(all_values[i, 4]),
|
|
||||||
"ablated_death_shape": float(all_values[i, 5]),
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
result_table = pd.DataFrame(rows).reindex(columns=OUTPUT_COLUMNS)
|
|
||||||
written_rows = int(len(result_table))
|
|
||||||
|
|
||||||
summary_accumulator: dict[tuple[Any, ...], dict[str, float]] = {}
|
|
||||||
update_summary_accumulator(summary_accumulator, result_table)
|
|
||||||
|
|
||||||
for start in range(0, written_rows, int(args.shard_rows)):
|
|
||||||
stop = min(written_rows, start + int(args.shard_rows))
|
|
||||||
shard_name = f"part-{shard_index:06d}.npz"
|
|
||||||
shard_path = output_dir / shard_name
|
|
||||||
shard_rows = write_compressed_npz_table(
|
|
||||||
shard_path,
|
|
||||||
result_table.iloc[start:stop],
|
|
||||||
)
|
|
||||||
shards.append({"file": shard_name, "rows": int(shard_rows)})
|
|
||||||
shard_index += 1
|
|
||||||
else:
|
|
||||||
result_table = pd.DataFrame(columns=OUTPUT_COLUMNS)
|
|
||||||
summary_accumulator = {}
|
|
||||||
|
|
||||||
if not shards:
|
|
||||||
empty_path = output_dir / "part-000000.npz"
|
|
||||||
write_compressed_npz_table(empty_path, pd.DataFrame(columns=OUTPUT_COLUMNS))
|
|
||||||
shards.append({"file": empty_path.name, "rows": 0})
|
|
||||||
summary_path = output_dir / "summary_by_disease_age_sex.csv"
|
|
||||||
summary_rows = write_summary_csv(summary_path, summary_accumulator)
|
|
||||||
write_manifest(
|
|
||||||
output_dir,
|
|
||||||
rows=written_rows,
|
|
||||||
shards=shards,
|
|
||||||
summary_file=summary_path.name,
|
|
||||||
scanned_diseases=[
|
|
||||||
{"token_id": int(token), **{k: v for k, v in meta.items() if k != "token_id"}}
|
|
||||||
for token, meta in scanned_disease_items
|
|
||||||
],
|
|
||||||
eval_split=eval_split,
|
|
||||||
dist_mode=dist_mode,
|
|
||||||
landmark_start=float(args.landmark_start),
|
|
||||||
landmark_stop=float(args.landmark_stop),
|
|
||||||
landmark_step=float(args.landmark_step),
|
|
||||||
)
|
|
||||||
print(f"Wrote {written_rows} rows in {len(shards)} shard(s) to {output_dir}")
|
|
||||||
print(f"Wrote {summary_rows} summary rows to {summary_path}")
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
main()
|
|
||||||
@@ -1,7 +0,0 @@
|
|||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from evaluate_auc import main
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
main()
|
|
||||||
@@ -1,327 +0,0 @@
|
|||||||
"""Export landmark risk logits and hidden states for t_query ages.
|
|
||||||
|
|
||||||
This script follows evaluate_event_free_survival.py's data loading,
|
|
||||||
landmark construction, checkpoint loading, and readout logic, but only exports:
|
|
||||||
|
|
||||||
* all token/disease risk logits from ``model.calc_risk(hidden)``;
|
|
||||||
* the corresponding landmark hidden state.
|
|
||||||
|
|
||||||
The two large arrays are saved separately as .npy files. Row metadata is saved
|
|
||||||
as a CSV with matching row order.
|
|
||||||
"""
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import argparse
|
|
||||||
import json
|
|
||||||
from pathlib import Path
|
|
||||||
from typing import Any, Optional
|
|
||||||
|
|
||||||
import numpy as np
|
|
||||||
import pandas as pd
|
|
||||||
import torch
|
|
||||||
from torch.utils.data import DataLoader
|
|
||||||
from tqdm.auto import tqdm
|
|
||||||
|
|
||||||
from evaluate_auc_v2 import (
|
|
||||||
LandmarkDataset,
|
|
||||||
build_model_from_dataset,
|
|
||||||
cfg_get,
|
|
||||||
load_checkpoint_state_dict,
|
|
||||||
load_json_config,
|
|
||||||
load_model_state,
|
|
||||||
resolve_dist_mode_for_checkpoint,
|
|
||||||
resolve_eval_device,
|
|
||||||
validate_dataset_metadata,
|
|
||||||
)
|
|
||||||
from evaluate_event_free_survival import (
|
|
||||||
IndexedLandmarkDataset,
|
|
||||||
build_first_occurrence_maps_for_landmarks,
|
|
||||||
collate_indexed_landmark_fn,
|
|
||||||
infer_landmark_hidden,
|
|
||||||
load_eval_sequence_dataset,
|
|
||||||
make_landmark_ages,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def numpy_float_dtype(name: str) -> np.dtype:
|
|
||||||
key = str(name).lower()
|
|
||||||
if key in {"float16", "fp16", "half"}:
|
|
||||||
return np.dtype(np.float16)
|
|
||||||
if key in {"float32", "fp32", "single"}:
|
|
||||||
return np.dtype(np.float32)
|
|
||||||
raise ValueError(f"dtype must be float16 or float32, got {name!r}")
|
|
||||||
|
|
||||||
|
|
||||||
def output_paths_for_run(
|
|
||||||
run_path: Path,
|
|
||||||
eval_split: str,
|
|
||||||
landmark_start: float,
|
|
||||||
landmark_stop: float,
|
|
||||||
landmark_step: float,
|
|
||||||
output_dir: Optional[str],
|
|
||||||
) -> tuple[Path, Path, Path, Path]:
|
|
||||||
suffix = f"{eval_split}_t{landmark_start:g}-{landmark_stop:g}_step{landmark_step:g}"
|
|
||||||
base_dir = Path(output_dir) if output_dir else run_path
|
|
||||||
return (
|
|
||||||
base_dir / f"tquery_logits_{suffix}.npy",
|
|
||||||
base_dir / f"tquery_hidden_{suffix}.npy",
|
|
||||||
base_dir / f"tquery_metadata_{suffix}.csv",
|
|
||||||
base_dir / f"tquery_export_config_{suffix}.json",
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def parse_args() -> argparse.Namespace:
|
|
||||||
parser = argparse.ArgumentParser(
|
|
||||||
description="Export landmark risk logits and hidden states for t_query ages."
|
|
||||||
)
|
|
||||||
parser.add_argument("--run_path", type=str, required=True)
|
|
||||||
parser.add_argument(
|
|
||||||
"--output_dir",
|
|
||||||
type=str,
|
|
||||||
default=None,
|
|
||||||
help="Directory for output files. Defaults to run_path.",
|
|
||||||
)
|
|
||||||
parser.add_argument("--logits_path", type=str, default=None)
|
|
||||||
parser.add_argument("--hidden_path", type=str, default=None)
|
|
||||||
parser.add_argument("--metadata_path", type=str, default=None)
|
|
||||||
parser.add_argument("--export_config_path", type=str, default=None)
|
|
||||||
parser.add_argument("--eval_split", type=str, default=None)
|
|
||||||
parser.add_argument("--dataset_subset_size", type=int, default=None)
|
|
||||||
parser.add_argument("--train_eid_file", type=str, default=None)
|
|
||||||
parser.add_argument("--val_eid_file", type=str, default=None)
|
|
||||||
parser.add_argument("--test_eid_file", type=str, default=None)
|
|
||||||
parser.add_argument("--landmark_start", type=float, default=40.0)
|
|
||||||
parser.add_argument("--landmark_stop", type=float, default=80.0)
|
|
||||||
parser.add_argument(
|
|
||||||
"--landmark_step",
|
|
||||||
type=float,
|
|
||||||
default=1.0,
|
|
||||||
help="t_query grid step in years. Default exports every integer age 40..80.",
|
|
||||||
)
|
|
||||||
parser.add_argument("--min_history_events", type=int, default=None)
|
|
||||||
parser.add_argument("--batch_size", type=int, default=None)
|
|
||||||
parser.add_argument("--num_workers", type=int, default=None)
|
|
||||||
parser.add_argument("--device", type=str, default=None)
|
|
||||||
parser.add_argument("--extra_info_types", type=str, default=None)
|
|
||||||
parser.add_argument(
|
|
||||||
"--logits_dtype",
|
|
||||||
type=str,
|
|
||||||
default="float32",
|
|
||||||
choices=["float16", "float32"],
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
|
||||||
"--hidden_dtype",
|
|
||||||
type=str,
|
|
||||||
default="float32",
|
|
||||||
choices=["float16", "float32"],
|
|
||||||
)
|
|
||||||
return parser.parse_args()
|
|
||||||
|
|
||||||
|
|
||||||
def main() -> None:
|
|
||||||
args = parse_args()
|
|
||||||
run_path = Path(args.run_path)
|
|
||||||
config_path = run_path / "train_config.json"
|
|
||||||
checkpoint_path = run_path / "best_model.pt"
|
|
||||||
if not config_path.exists():
|
|
||||||
raise FileNotFoundError(f"train_config.json not found: {config_path}")
|
|
||||||
if not checkpoint_path.exists():
|
|
||||||
raise FileNotFoundError(f"best_model.pt not found: {checkpoint_path}")
|
|
||||||
|
|
||||||
cfg = load_json_config(config_path)
|
|
||||||
model_target_mode = str(cfg.get("model_target_mode", "next_token")).lower()
|
|
||||||
if model_target_mode not in {"next_token", "all_future"}:
|
|
||||||
raise ValueError(f"Unsupported model_target_mode: {model_target_mode!r}")
|
|
||||||
|
|
||||||
target_mode = str(cfg.get("target_mode", "uts"))
|
|
||||||
attn_mask_mode = str(
|
|
||||||
cfg.get(
|
|
||||||
"attn_mask_mode",
|
|
||||||
"non_strict_time" if target_mode == "uts" else "target_aware",
|
|
||||||
)
|
|
||||||
)
|
|
||||||
readout_name = str(
|
|
||||||
cfg.get("readout_name", "same_time_group_end" if target_mode == "uts" else "token")
|
|
||||||
)
|
|
||||||
readout_reduce = str(cfg.get("readout_reduce", "mean"))
|
|
||||||
|
|
||||||
dataset, subset_indices, eval_split, split_source = load_eval_sequence_dataset(
|
|
||||||
args,
|
|
||||||
cfg,
|
|
||||||
)
|
|
||||||
validate_dataset_metadata(dataset, cfg)
|
|
||||||
|
|
||||||
landmark_ages = make_landmark_ages(
|
|
||||||
float(args.landmark_start),
|
|
||||||
float(args.landmark_stop),
|
|
||||||
float(args.landmark_step),
|
|
||||||
)
|
|
||||||
first_occurrence_by_token = build_first_occurrence_maps_for_landmarks(
|
|
||||||
dataset,
|
|
||||||
subset_indices,
|
|
||||||
)
|
|
||||||
death_idx = int(dataset.vocab_size) - 1
|
|
||||||
landmark_dataset = LandmarkDataset(
|
|
||||||
dataset=dataset,
|
|
||||||
subset_indices=subset_indices,
|
|
||||||
landmark_ages=landmark_ages,
|
|
||||||
attn_mask_mode=attn_mask_mode,
|
|
||||||
model_target_mode=model_target_mode,
|
|
||||||
min_history_events=int(cfg_get(args, cfg, "min_history_events", 1)),
|
|
||||||
first_occurrence_by_token=first_occurrence_by_token,
|
|
||||||
death_token_ids=[death_idx],
|
|
||||||
)
|
|
||||||
|
|
||||||
state_dict = load_checkpoint_state_dict(checkpoint_path, map_location="cpu")
|
|
||||||
dist_mode = resolve_dist_mode_for_checkpoint(
|
|
||||||
str(cfg.get("dist_mode", "exponential")),
|
|
||||||
state_dict,
|
|
||||||
)
|
|
||||||
cfg_model = dict(cfg)
|
|
||||||
cfg_model["dist_mode"] = dist_mode
|
|
||||||
device = resolve_eval_device(args.device)
|
|
||||||
model = build_model_from_dataset(
|
|
||||||
args, cfg_model, dataset, state_dict=state_dict
|
|
||||||
).to(device)
|
|
||||||
load_model_state(model, state_dict)
|
|
||||||
model.eval()
|
|
||||||
|
|
||||||
default_logits_path, default_hidden_path, default_metadata_path, default_config_path = (
|
|
||||||
output_paths_for_run(
|
|
||||||
run_path=run_path,
|
|
||||||
eval_split=eval_split,
|
|
||||||
landmark_start=float(args.landmark_start),
|
|
||||||
landmark_stop=float(args.landmark_stop),
|
|
||||||
landmark_step=float(args.landmark_step),
|
|
||||||
output_dir=args.output_dir,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
logits_path = Path(args.logits_path) if args.logits_path else default_logits_path
|
|
||||||
hidden_path = Path(args.hidden_path) if args.hidden_path else default_hidden_path
|
|
||||||
metadata_path = Path(args.metadata_path) if args.metadata_path else default_metadata_path
|
|
||||||
export_config_path = (
|
|
||||||
Path(args.export_config_path) if args.export_config_path else default_config_path
|
|
||||||
)
|
|
||||||
for path in (logits_path, hidden_path, metadata_path, export_config_path):
|
|
||||||
path.parent.mkdir(parents=True, exist_ok=True)
|
|
||||||
|
|
||||||
n_rows = len(landmark_dataset)
|
|
||||||
vocab_size = int(dataset.vocab_size)
|
|
||||||
hidden_dim = int(getattr(model, "n_embd", cfg_get(args, cfg_model, "n_embd", 120)))
|
|
||||||
logits_dtype = numpy_float_dtype(args.logits_dtype)
|
|
||||||
hidden_dtype = numpy_float_dtype(args.hidden_dtype)
|
|
||||||
|
|
||||||
logits_memmap = np.lib.format.open_memmap(
|
|
||||||
logits_path,
|
|
||||||
mode="w+",
|
|
||||||
dtype=logits_dtype,
|
|
||||||
shape=(n_rows, vocab_size),
|
|
||||||
)
|
|
||||||
hidden_memmap = np.lib.format.open_memmap(
|
|
||||||
hidden_path,
|
|
||||||
mode="w+",
|
|
||||||
dtype=hidden_dtype,
|
|
||||||
shape=(n_rows, hidden_dim),
|
|
||||||
)
|
|
||||||
|
|
||||||
batch_size = int(cfg_get(args, cfg, "batch_size", 128))
|
|
||||||
num_workers = int(cfg_get(args, cfg, "num_workers", 4))
|
|
||||||
loader = DataLoader(
|
|
||||||
IndexedLandmarkDataset(landmark_dataset),
|
|
||||||
batch_size=batch_size,
|
|
||||||
shuffle=False,
|
|
||||||
collate_fn=collate_indexed_landmark_fn,
|
|
||||||
num_workers=num_workers,
|
|
||||||
pin_memory=device.type == "cuda",
|
|
||||||
persistent_workers=num_workers > 0,
|
|
||||||
prefetch_factor=2 if num_workers > 0 else None,
|
|
||||||
)
|
|
||||||
|
|
||||||
print(f"Eval split: {eval_split}")
|
|
||||||
print(f"Split source: {split_source}")
|
|
||||||
print(f"Selected patients: {len(subset_indices)}")
|
|
||||||
print(f"t_query ages: {landmark_ages.tolist()}")
|
|
||||||
print(f"Dist mode: {dist_mode}")
|
|
||||||
print(f"Device: {device}")
|
|
||||||
print(f"Landmark rows: {n_rows}")
|
|
||||||
print(f"Logits: {logits_path} shape={(n_rows, vocab_size)} dtype={logits_dtype}")
|
|
||||||
print(f"Hidden: {hidden_path} shape={(n_rows, hidden_dim)} dtype={hidden_dtype}")
|
|
||||||
print(f"Metadata: {metadata_path}")
|
|
||||||
|
|
||||||
meta_rows: list[dict[str, Any]] = []
|
|
||||||
written = 0
|
|
||||||
with torch.no_grad():
|
|
||||||
for batch in tqdm(loader, desc="Export logits/hidden", dynamic_ncols=True):
|
|
||||||
hidden = infer_landmark_hidden(
|
|
||||||
model=model,
|
|
||||||
batch=batch,
|
|
||||||
device=device,
|
|
||||||
model_target_mode=model_target_mode,
|
|
||||||
readout_name=readout_name,
|
|
||||||
readout_reduce=readout_reduce,
|
|
||||||
)
|
|
||||||
logits = model.calc_risk(hidden)
|
|
||||||
row_indices = batch["row_idx"].detach().cpu().numpy().astype(np.int64)
|
|
||||||
if not np.array_equal(row_indices, np.arange(written, written + len(row_indices))):
|
|
||||||
raise RuntimeError("DataLoader row order changed; export requires shuffle=False.")
|
|
||||||
|
|
||||||
batch_n = int(logits.shape[0])
|
|
||||||
logits_memmap[written : written + batch_n] = (
|
|
||||||
logits.detach().cpu().numpy().astype(logits_dtype, copy=False)
|
|
||||||
)
|
|
||||||
hidden_memmap[written : written + batch_n] = (
|
|
||||||
hidden.detach().cpu().numpy().astype(hidden_dtype, copy=False)
|
|
||||||
)
|
|
||||||
|
|
||||||
for row_idx in row_indices.tolist():
|
|
||||||
meta = landmark_dataset.rows[int(row_idx)]
|
|
||||||
sample = dataset.samples[int(meta["dataset_index"])]
|
|
||||||
meta_rows.append(
|
|
||||||
{
|
|
||||||
"row_index": int(row_idx),
|
|
||||||
"patient_id": int(meta["patient_id"]),
|
|
||||||
"dataset_index": int(meta["dataset_index"]),
|
|
||||||
"eid": int(sample.get("eid", -1)),
|
|
||||||
"sex": int(meta["sex"]),
|
|
||||||
"t_query": float(meta["t_query"]),
|
|
||||||
"landmark_age": float(meta["landmark_age"]),
|
|
||||||
"followup_end_time": float(meta["followup_end_time"]),
|
|
||||||
"death_time": float(meta["death_time"]),
|
|
||||||
}
|
|
||||||
)
|
|
||||||
written += batch_n
|
|
||||||
|
|
||||||
logits_memmap.flush()
|
|
||||||
hidden_memmap.flush()
|
|
||||||
pd.DataFrame(meta_rows).to_csv(metadata_path, index=False)
|
|
||||||
|
|
||||||
export_config = {
|
|
||||||
"run_path": str(run_path),
|
|
||||||
"eval_split": eval_split,
|
|
||||||
"split_source": split_source,
|
|
||||||
"model_target_mode": model_target_mode,
|
|
||||||
"target_mode": target_mode,
|
|
||||||
"attn_mask_mode": attn_mask_mode,
|
|
||||||
"readout_name": readout_name,
|
|
||||||
"readout_reduce": readout_reduce,
|
|
||||||
"dist_mode": dist_mode,
|
|
||||||
"landmark_ages": [float(x) for x in landmark_ages.tolist()],
|
|
||||||
"n_rows": int(n_rows),
|
|
||||||
"vocab_size": int(vocab_size),
|
|
||||||
"hidden_dim": int(hidden_dim),
|
|
||||||
"death_token": int(death_idx),
|
|
||||||
"logits_path": str(logits_path),
|
|
||||||
"hidden_path": str(hidden_path),
|
|
||||||
"metadata_path": str(metadata_path),
|
|
||||||
"logits_dtype": str(logits_dtype),
|
|
||||||
"hidden_dtype": str(hidden_dtype),
|
|
||||||
}
|
|
||||||
with export_config_path.open("w", encoding="utf-8") as f:
|
|
||||||
json.dump(export_config, f, indent=2)
|
|
||||||
|
|
||||||
print(f"Wrote {written} rows.")
|
|
||||||
print(f"Wrote export config: {export_config_path}")
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
main()
|
|
||||||
@@ -1,519 +0,0 @@
|
|||||||
"""Export Weibull shape-parameter statistics on the test split.
|
|
||||||
|
|
||||||
The script is intended for all_future checkpoints with dist_mode="weibull" or
|
|
||||||
dist_mode="mixed". For full Weibull models it reads rho_head[Death]; for mixed
|
|
||||||
models it reads rho_death_head. For full Weibull models it also exports disease
|
|
||||||
token rho summaries, which are the main evidence for whether risk/hazard changes
|
|
||||||
with horizon instead of following an exponential shape.
|
|
||||||
"""
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import argparse
|
|
||||||
import contextlib
|
|
||||||
import json
|
|
||||||
from pathlib import Path
|
|
||||||
from typing import Any, Dict, Iterable, List, Optional
|
|
||||||
|
|
||||||
import numpy as np
|
|
||||||
import pandas as pd
|
|
||||||
import torch
|
|
||||||
import torch.nn.functional as F
|
|
||||||
import torch.multiprocessing as torch_mp
|
|
||||||
from torch.utils.data import DataLoader
|
|
||||||
from tqdm.auto import tqdm
|
|
||||||
|
|
||||||
from eval_data import load_sequence_eval_dataset
|
|
||||||
from evaluate_auc_v2 import (
|
|
||||||
LandmarkDataset,
|
|
||||||
_build_first_occurrence_maps,
|
|
||||||
_get_death_token_ids,
|
|
||||||
build_model_from_dataset,
|
|
||||||
cfg_get,
|
|
||||||
collate_landmark_fn,
|
|
||||||
load_checkpoint_state_dict,
|
|
||||||
load_json_config,
|
|
||||||
load_model_state,
|
|
||||||
make_eval_indices,
|
|
||||||
parse_float_list,
|
|
||||||
parse_int_list,
|
|
||||||
resolve_dist_mode_for_checkpoint,
|
|
||||||
resolve_eval_device,
|
|
||||||
validate_dataset_metadata,
|
|
||||||
)
|
|
||||||
|
|
||||||
try:
|
|
||||||
torch_mp.set_sharing_strategy("file_system")
|
|
||||||
except RuntimeError:
|
|
||||||
pass
|
|
||||||
|
|
||||||
|
|
||||||
def quantile_summary(df: pd.DataFrame, group_cols: List[str], value_cols: List[str]) -> pd.DataFrame:
|
|
||||||
probs = [0.01, 0.05, 0.25, 0.50, 0.75, 0.95, 0.99]
|
|
||||||
rows: List[Dict[str, Any]] = []
|
|
||||||
grouped = [((), df)] if not group_cols else df.groupby(group_cols, dropna=False)
|
|
||||||
|
|
||||||
for key, g in grouped:
|
|
||||||
if not isinstance(key, tuple):
|
|
||||||
key = (key,)
|
|
||||||
base = {col: val for col, val in zip(group_cols, key)}
|
|
||||||
base["n"] = int(len(g))
|
|
||||||
for col in value_cols:
|
|
||||||
x = pd.to_numeric(g[col], errors="coerce").to_numpy(dtype=np.float64)
|
|
||||||
x = x[np.isfinite(x)]
|
|
||||||
if x.size == 0:
|
|
||||||
continue
|
|
||||||
row = dict(base)
|
|
||||||
row["variable"] = col
|
|
||||||
row["mean"] = float(np.mean(x))
|
|
||||||
row["std"] = float(np.std(x, ddof=1)) if x.size > 1 else 0.0
|
|
||||||
row["min"] = float(np.min(x))
|
|
||||||
row["max"] = float(np.max(x))
|
|
||||||
for p in probs:
|
|
||||||
row[f"p{int(p * 100):02d}"] = float(np.quantile(x, p))
|
|
||||||
rows.append(row)
|
|
||||||
return pd.DataFrame(rows)
|
|
||||||
|
|
||||||
|
|
||||||
def load_labels_meta(path: Optional[str]) -> Optional[pd.DataFrame]:
|
|
||||||
if path is None:
|
|
||||||
return None
|
|
||||||
fp = Path(path)
|
|
||||||
if not fp.exists():
|
|
||||||
return None
|
|
||||||
return pd.read_csv(fp)
|
|
||||||
|
|
||||||
|
|
||||||
@torch.inference_mode()
|
|
||||||
def infer_landmark_hidden_local(
|
|
||||||
model,
|
|
||||||
loader: DataLoader,
|
|
||||||
device: torch.device,
|
|
||||||
use_amp: bool,
|
|
||||||
hidden_cache_dtype: str,
|
|
||||||
) -> tuple[np.ndarray, Dict[str, np.ndarray]]:
|
|
||||||
"""Minimal all_future landmark hidden inference for parameter export."""
|
|
||||||
out_dtype = np.float32 if str(hidden_cache_dtype).lower() == "float32" else np.float16
|
|
||||||
hidden_parts: List[np.ndarray] = []
|
|
||||||
arrays: Dict[str, List[np.ndarray]] = {
|
|
||||||
"patient_id": [],
|
|
||||||
"sex": [],
|
|
||||||
"landmark_age": [],
|
|
||||||
"followup_end_time": [],
|
|
||||||
"death_time": [],
|
|
||||||
}
|
|
||||||
amp_enabled = bool(use_amp and device.type == "cuda")
|
|
||||||
|
|
||||||
for batch in tqdm(loader, desc="Landmark hidden", dynamic_ncols=True):
|
|
||||||
batch_dev = {
|
|
||||||
k: (v.to(device, non_blocking=True) if isinstance(v, torch.Tensor) else v)
|
|
||||||
for k, v in batch.items()
|
|
||||||
}
|
|
||||||
amp_ctx = (
|
|
||||||
torch.autocast(device_type=device.type, dtype=torch.float16)
|
|
||||||
if amp_enabled
|
|
||||||
else contextlib.nullcontext()
|
|
||||||
)
|
|
||||||
with amp_ctx:
|
|
||||||
hidden = model(
|
|
||||||
event_seq=batch_dev["event_seq"],
|
|
||||||
time_seq=batch_dev["time_seq"],
|
|
||||||
sex=batch_dev["sex"],
|
|
||||||
padding_mask=batch_dev["padding_mask"],
|
|
||||||
t_query=batch_dev["t_query"],
|
|
||||||
other_type=batch_dev["other_type"],
|
|
||||||
other_value=batch_dev["other_value"],
|
|
||||||
other_value_kind=batch_dev["other_value_kind"],
|
|
||||||
other_time=batch_dev["other_time"],
|
|
||||||
target_mode="all_future",
|
|
||||||
)
|
|
||||||
hidden_parts.append(hidden.detach().cpu().numpy().astype(out_dtype, copy=False))
|
|
||||||
for key in arrays:
|
|
||||||
arrays[key].append(batch[key].cpu().numpy())
|
|
||||||
|
|
||||||
hidden_all = np.concatenate(hidden_parts, axis=0)
|
|
||||||
row_arrays = {key: np.concatenate(parts, axis=0) for key, parts in arrays.items()}
|
|
||||||
return hidden_all, row_arrays
|
|
||||||
|
|
||||||
|
|
||||||
@torch.inference_mode()
|
|
||||||
def project_death_params(
|
|
||||||
model,
|
|
||||||
hidden_all: np.ndarray,
|
|
||||||
dist_mode: str,
|
|
||||||
device: torch.device,
|
|
||||||
batch_size: int,
|
|
||||||
use_amp: bool,
|
|
||||||
) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
|
|
||||||
death_idx = int(getattr(model, "death_idx", getattr(model, "vocab_size", hidden_all.shape[0]) - 1))
|
|
||||||
if not hasattr(model, "vocab_size"):
|
|
||||||
death_idx = int(model.risk_head.out_features - 1)
|
|
||||||
|
|
||||||
compute_dtype = torch.float16 if (device.type == "cuda" and use_amp) else torch.float32
|
|
||||||
risk_w = model.risk_head.weight[death_idx: death_idx + 1].detach().to(device=device, dtype=compute_dtype)
|
|
||||||
risk_b = None
|
|
||||||
if model.risk_head.bias is not None:
|
|
||||||
risk_b = model.risk_head.bias[death_idx: death_idx + 1].detach().to(device=device, dtype=compute_dtype)
|
|
||||||
|
|
||||||
if dist_mode == "weibull":
|
|
||||||
rho_w = model.rho_head.weight[death_idx: death_idx + 1].detach().to(device=device, dtype=compute_dtype)
|
|
||||||
rho_b = model.rho_head.bias[death_idx: death_idx + 1].detach().to(device=device, dtype=compute_dtype)
|
|
||||||
elif dist_mode == "mixed":
|
|
||||||
rho_w = model.rho_death_head.weight.detach().to(device=device, dtype=compute_dtype)
|
|
||||||
rho_b = model.rho_death_head.bias.detach().to(device=device, dtype=compute_dtype)
|
|
||||||
else:
|
|
||||||
raise ValueError("Death Weibull parameter export requires dist_mode='weibull' or 'mixed'.")
|
|
||||||
|
|
||||||
logits_out: List[np.ndarray] = []
|
|
||||||
rate_out: List[np.ndarray] = []
|
|
||||||
rho_out: List[np.ndarray] = []
|
|
||||||
|
|
||||||
for start in tqdm(range(0, hidden_all.shape[0], batch_size), desc="Death eta/rho", dynamic_ncols=True):
|
|
||||||
end = min(start + batch_size, hidden_all.shape[0])
|
|
||||||
h = torch.from_numpy(hidden_all[start:end]).to(device=device, dtype=compute_dtype, non_blocking=True)
|
|
||||||
logits = F.linear(h, risk_w, risk_b).squeeze(-1)
|
|
||||||
rate = F.softplus(logits) + 1e-8
|
|
||||||
rho = F.softplus(F.linear(h, rho_w, rho_b).squeeze(-1)) + 1e-6
|
|
||||||
logits_out.append(logits.float().cpu().numpy())
|
|
||||||
rate_out.append(rate.float().cpu().numpy())
|
|
||||||
rho_out.append(rho.float().cpu().numpy())
|
|
||||||
del h, logits, rate, rho
|
|
||||||
|
|
||||||
return (
|
|
||||||
np.concatenate(logits_out).astype(np.float32, copy=False),
|
|
||||||
np.concatenate(rate_out).astype(np.float32, copy=False),
|
|
||||||
np.concatenate(rho_out).astype(np.float32, copy=False),
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@torch.inference_mode()
|
|
||||||
def export_all_token_rho_summary(
|
|
||||||
model,
|
|
||||||
hidden_all: np.ndarray,
|
|
||||||
dataset,
|
|
||||||
device: torch.device,
|
|
||||||
output_dir: Path,
|
|
||||||
token_chunk_size: int,
|
|
||||||
row_batch_size: int,
|
|
||||||
use_amp: bool,
|
|
||||||
horizons: np.ndarray,
|
|
||||||
) -> None:
|
|
||||||
if not hasattr(model, "rho_head"):
|
|
||||||
print("[INFO] Skipping all-token rho summary because this is not a full Weibull model.")
|
|
||||||
return
|
|
||||||
|
|
||||||
special = {0, 1, 2}
|
|
||||||
token_ids = [
|
|
||||||
int(t)
|
|
||||||
for t, code in dataset.label_id_to_code.items()
|
|
||||||
if int(t) not in special and not str(code).startswith("<")
|
|
||||||
]
|
|
||||||
token_ids = sorted(set(token_ids))
|
|
||||||
death_idx = int(getattr(model, "death_idx", getattr(model, "vocab_size", len(token_ids)) - 1))
|
|
||||||
if not hasattr(model, "vocab_size"):
|
|
||||||
death_idx = int(model.risk_head.out_features - 1)
|
|
||||||
compute_dtype = torch.float16 if (device.type == "cuda" and use_amp) else torch.float32
|
|
||||||
rows: List[Dict[str, Any]] = []
|
|
||||||
|
|
||||||
for chunk_start in tqdm(range(0, len(token_ids), token_chunk_size), desc="All-token rho chunks", dynamic_ncols=True):
|
|
||||||
chunk = token_ids[chunk_start: chunk_start + token_chunk_size]
|
|
||||||
w = model.rho_head.weight[chunk].detach().to(device=device, dtype=compute_dtype)
|
|
||||||
b = model.rho_head.bias[chunk].detach().to(device=device, dtype=compute_dtype)
|
|
||||||
vals_parts: List[np.ndarray] = []
|
|
||||||
for row_start in range(0, hidden_all.shape[0], row_batch_size):
|
|
||||||
row_end = min(row_start + row_batch_size, hidden_all.shape[0])
|
|
||||||
h = torch.from_numpy(hidden_all[row_start:row_end]).to(device=device, dtype=compute_dtype, non_blocking=True)
|
|
||||||
rho = F.softplus(F.linear(h, w, b)) + 1e-6
|
|
||||||
vals_parts.append(rho.float().cpu().numpy())
|
|
||||||
del h, rho
|
|
||||||
vals = np.concatenate(vals_parts, axis=0)
|
|
||||||
for j, token in enumerate(chunk):
|
|
||||||
x = vals[:, j].astype(np.float64, copy=False)
|
|
||||||
row = {
|
|
||||||
"token": int(token),
|
|
||||||
"label_code": dataset.label_id_to_code.get(int(token), ""),
|
|
||||||
"endpoint_type": "death" if int(token) == int(death_idx) else "disease",
|
|
||||||
"n_landmark_rows": int(x.size),
|
|
||||||
"rho_mean": float(np.mean(x)),
|
|
||||||
"rho_std": float(np.std(x, ddof=1)) if x.size > 1 else 0.0,
|
|
||||||
"rho_minus_one_mean": float(np.mean(x - 1.0)),
|
|
||||||
"frac_rho_gt_1": float(np.mean(x > 1.0)),
|
|
||||||
"frac_rho_lt_1": float(np.mean(x < 1.0)),
|
|
||||||
"frac_rho_gt_1_1": float(np.mean(x > 1.1)),
|
|
||||||
"frac_rho_lt_0_9": float(np.mean(x < 0.9)),
|
|
||||||
"rho_p01": float(np.quantile(x, 0.01)),
|
|
||||||
"rho_p05": float(np.quantile(x, 0.05)),
|
|
||||||
"rho_p25": float(np.quantile(x, 0.25)),
|
|
||||||
"rho_p50": float(np.quantile(x, 0.50)),
|
|
||||||
"rho_p75": float(np.quantile(x, 0.75)),
|
|
||||||
"rho_p95": float(np.quantile(x, 0.95)),
|
|
||||||
"rho_p99": float(np.quantile(x, 0.99)),
|
|
||||||
}
|
|
||||||
for horizon in horizons.tolist():
|
|
||||||
h = float(horizon)
|
|
||||||
if h <= 0:
|
|
||||||
continue
|
|
||||||
# Shape-only time scaling. For rho=1 this equals 1, i.e. an
|
|
||||||
# exponential model with constant instantaneous hazard.
|
|
||||||
inst_scale = np.power(h, x - 1.0)
|
|
||||||
cumhaz_scale = np.power(h, x)
|
|
||||||
row[f"instant_hazard_scale_h{h:g}y_vs_1y_mean"] = float(np.mean(inst_scale))
|
|
||||||
row[f"instant_hazard_scale_h{h:g}y_vs_1y_p50"] = float(np.quantile(inst_scale, 0.50))
|
|
||||||
row[f"cumhaz_scale_h{h:g}y_mean"] = float(np.mean(cumhaz_scale))
|
|
||||||
row[f"cumhaz_scale_h{h:g}y_p50"] = float(np.quantile(cumhaz_scale, 0.50))
|
|
||||||
rows.append(row)
|
|
||||||
del vals, vals_parts
|
|
||||||
|
|
||||||
out = pd.DataFrame(rows)
|
|
||||||
out.to_csv(output_dir / "all_token_weibull_shape_summary.csv", index=False)
|
|
||||||
out[out["endpoint_type"] == "disease"].to_csv(
|
|
||||||
output_dir / "disease_token_weibull_shape_summary.csv", index=False
|
|
||||||
)
|
|
||||||
out[out["endpoint_type"] == "death"].to_csv(
|
|
||||||
output_dir / "death_token_weibull_shape_summary.csv", index=False
|
|
||||||
)
|
|
||||||
|
|
||||||
disease = out[out["endpoint_type"] == "disease"].copy()
|
|
||||||
if not disease.empty:
|
|
||||||
pd.DataFrame([
|
|
||||||
{
|
|
||||||
"n_tokens": int(len(disease)),
|
|
||||||
"rho_mean_across_tokens": float(disease["rho_mean"].mean()),
|
|
||||||
"rho_median_across_tokens": float(disease["rho_p50"].median()),
|
|
||||||
"tokens_with_mean_rho_gt_1": int((disease["rho_mean"] > 1.0).sum()),
|
|
||||||
"tokens_with_mean_rho_lt_1": int((disease["rho_mean"] < 1.0).sum()),
|
|
||||||
"frac_tokens_with_mean_rho_gt_1": float((disease["rho_mean"] > 1.0).mean()),
|
|
||||||
"frac_tokens_with_mean_rho_lt_1": float((disease["rho_mean"] < 1.0).mean()),
|
|
||||||
"tokens_with_mean_rho_gt_1_1": int((disease["rho_mean"] > 1.1).sum()),
|
|
||||||
"tokens_with_mean_rho_lt_0_9": int((disease["rho_mean"] < 0.9).sum()),
|
|
||||||
}
|
|
||||||
]).to_csv(output_dir / "disease_weibull_shape_overall_summary.csv", index=False)
|
|
||||||
|
|
||||||
|
|
||||||
def main() -> None:
|
|
||||||
parser = argparse.ArgumentParser(description="Export test-split Weibull shape parameter statistics.")
|
|
||||||
parser.add_argument("--run_path", type=str, required=True)
|
|
||||||
parser.add_argument("--output_path", type=str, default=None)
|
|
||||||
parser.add_argument("--eval_split", type=str, default="test", choices=["test", "val", "valid", "validation", "train", "all"])
|
|
||||||
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("--batch_size", type=int, default=None)
|
|
||||||
parser.add_argument(
|
|
||||||
"--num_workers",
|
|
||||||
type=int,
|
|
||||||
default=0,
|
|
||||||
help=(
|
|
||||||
"DataLoader workers. Default 0 avoids Linux multiprocessing "
|
|
||||||
"'received 0 items of ancdata' failures on shared filesystems."
|
|
||||||
),
|
|
||||||
)
|
|
||||||
parser.add_argument("--device", type=str, default=None)
|
|
||||||
parser.add_argument("--use_amp", action=argparse.BooleanOptionalAction, default=None)
|
|
||||||
parser.add_argument("--hidden_cache_dtype", type=str, default="float32", choices=["float16", "float32"])
|
|
||||||
parser.add_argument(
|
|
||||||
"--include_all_token_rho_summary",
|
|
||||||
action=argparse.BooleanOptionalAction,
|
|
||||||
default=True,
|
|
||||||
help=(
|
|
||||||
"For full Weibull models, export disease/death token rho summaries. "
|
|
||||||
"Use --no-include_all_token_rho_summary to skip the heavier token projection."
|
|
||||||
),
|
|
||||||
)
|
|
||||||
parser.add_argument("--token_chunk_size", type=int, default=32)
|
|
||||||
parser.add_argument("--row_batch_size", type=int, default=512)
|
|
||||||
args = parser.parse_args()
|
|
||||||
|
|
||||||
run_path = Path(args.run_path)
|
|
||||||
config_path = run_path / "train_config.json"
|
|
||||||
ckpt_path = run_path / "best_model.pt"
|
|
||||||
if not config_path.exists():
|
|
||||||
raise FileNotFoundError(config_path)
|
|
||||||
if not ckpt_path.exists():
|
|
||||||
raise FileNotFoundError(ckpt_path)
|
|
||||||
|
|
||||||
cfg = load_json_config(config_path)
|
|
||||||
model_target_mode = str(cfg.get("model_target_mode", "next_token")).lower()
|
|
||||||
if model_target_mode != "all_future":
|
|
||||||
raise ValueError("This export is intended for all_future checkpoints.")
|
|
||||||
|
|
||||||
data_prefix = cfg.get("data_prefix", "ukb")
|
|
||||||
labels_file = cfg.get("labels_file", "labels.csv")
|
|
||||||
no_event_interval_years = cfg.get("no_event_interval_years", 5.0)
|
|
||||||
include_no_event_in_uts_target = cfg.get("include_no_event_in_uts_target", False)
|
|
||||||
|
|
||||||
dataset = load_sequence_eval_dataset(
|
|
||||||
model_target_mode=model_target_mode,
|
|
||||||
data_prefix=data_prefix,
|
|
||||||
labels_file=labels_file,
|
|
||||||
no_event_interval_years=float(no_event_interval_years),
|
|
||||||
include_no_event_in_uts_target=bool(include_no_event_in_uts_target),
|
|
||||||
min_history_events=int(cfg.get("all_future_min_history_events", 1)),
|
|
||||||
min_future_events=int(cfg.get("all_future_min_future_events", 1)),
|
|
||||||
extra_info_types=parse_int_list(cfg.get("extra_info_types", None)),
|
|
||||||
)
|
|
||||||
validate_dataset_metadata(dataset, cfg)
|
|
||||||
|
|
||||||
subset_indices = make_eval_indices(dataset, args, cfg)
|
|
||||||
first_occurrence_by_token, _, _, _ = _build_first_occurrence_maps(dataset, subset_indices)
|
|
||||||
|
|
||||||
landmark_start = float(cfg_get(args, cfg, "landmark_start", 40.0))
|
|
||||||
landmark_stop = float(cfg_get(args, cfg, "landmark_stop", 80.0))
|
|
||||||
landmark_step = float(cfg_get(args, cfg, "landmark_step", 5.0))
|
|
||||||
landmark_ages = np.arange(landmark_start, landmark_stop, landmark_step, dtype=np.float32)
|
|
||||||
if landmark_ages.size == 0:
|
|
||||||
raise ValueError("No landmark ages produced.")
|
|
||||||
|
|
||||||
horizons = np.asarray(
|
|
||||||
parse_float_list(cfg_get(args, cfg, "horizons", "1,5,10")) or [1.0, 5.0, 10.0],
|
|
||||||
dtype=np.float32,
|
|
||||||
)
|
|
||||||
if horizons.size == 0:
|
|
||||||
raise ValueError("No horizons provided.")
|
|
||||||
|
|
||||||
state_dict = load_checkpoint_state_dict(ckpt_path, map_location="cpu")
|
|
||||||
dist_mode = resolve_dist_mode_for_checkpoint(str(cfg.get("dist_mode", "exponential")), state_dict)
|
|
||||||
if dist_mode not in {"weibull", "mixed"}:
|
|
||||||
raise ValueError(
|
|
||||||
f"Resolved dist_mode={dist_mode!r}; expected 'weibull' or 'mixed' for Weibull shape export."
|
|
||||||
)
|
|
||||||
|
|
||||||
cfg_model = dict(cfg)
|
|
||||||
cfg_model["dist_mode"] = dist_mode
|
|
||||||
device = resolve_eval_device(args.device)
|
|
||||||
model = build_model_from_dataset(
|
|
||||||
args, cfg_model, dataset, state_dict=state_dict
|
|
||||||
).to(device)
|
|
||||||
load_model_state(model, state_dict)
|
|
||||||
model.eval()
|
|
||||||
|
|
||||||
death_token_ids = _get_death_token_ids(dataset, None)
|
|
||||||
death_idx = int(death_token_ids[0])
|
|
||||||
attn_mask_mode = str(cfg.get("attn_mask_mode", "target_aware"))
|
|
||||||
landmark_dataset = LandmarkDataset(
|
|
||||||
dataset=dataset,
|
|
||||||
subset_indices=subset_indices,
|
|
||||||
landmark_ages=landmark_ages,
|
|
||||||
attn_mask_mode=attn_mask_mode,
|
|
||||||
model_target_mode=model_target_mode,
|
|
||||||
min_history_events=int(cfg_get(args, cfg, "min_history_events", 1)),
|
|
||||||
first_occurrence_by_token=first_occurrence_by_token,
|
|
||||||
death_token_ids=death_token_ids,
|
|
||||||
)
|
|
||||||
|
|
||||||
batch_size = int(cfg_get(args, cfg, "batch_size", 128))
|
|
||||||
num_workers = int(cfg_get(args, cfg, "num_workers", 0))
|
|
||||||
loader_kwargs = {
|
|
||||||
"batch_size": batch_size,
|
|
||||||
"shuffle": False,
|
|
||||||
"collate_fn": collate_landmark_fn,
|
|
||||||
"num_workers": num_workers,
|
|
||||||
"pin_memory": device.type == "cuda",
|
|
||||||
}
|
|
||||||
if num_workers > 0:
|
|
||||||
loader_kwargs["persistent_workers"] = True
|
|
||||||
loader_kwargs["prefetch_factor"] = 2
|
|
||||||
loader = DataLoader(landmark_dataset, **loader_kwargs)
|
|
||||||
|
|
||||||
use_amp = bool(cfg_get(args, cfg, "use_amp", False))
|
|
||||||
hidden_all, row_arrays = infer_landmark_hidden_local(
|
|
||||||
model=model,
|
|
||||||
loader=loader,
|
|
||||||
device=device,
|
|
||||||
use_amp=use_amp,
|
|
||||||
hidden_cache_dtype=str(args.hidden_cache_dtype),
|
|
||||||
)
|
|
||||||
eta, rate, rho = project_death_params(
|
|
||||||
model=model,
|
|
||||||
hidden_all=hidden_all,
|
|
||||||
dist_mode=dist_mode,
|
|
||||||
device=device,
|
|
||||||
batch_size=int(args.row_batch_size),
|
|
||||||
use_amp=use_amp,
|
|
||||||
)
|
|
||||||
|
|
||||||
rows = pd.DataFrame({
|
|
||||||
"patient_id": row_arrays["patient_id"].astype(np.int64),
|
|
||||||
"sex": row_arrays["sex"].astype(np.int64),
|
|
||||||
"sex_label": np.where(row_arrays["sex"].astype(np.int64) == 0, "female", "male"),
|
|
||||||
"landmark_age": row_arrays["landmark_age"].astype(np.float32),
|
|
||||||
"followup_end_time": row_arrays["followup_end_time"].astype(np.float32),
|
|
||||||
"death_time": row_arrays["death_time"].astype(np.float32),
|
|
||||||
"death_eta": eta,
|
|
||||||
"death_rate": rate,
|
|
||||||
"death_rho": rho,
|
|
||||||
})
|
|
||||||
for horizon in horizons.tolist():
|
|
||||||
h = float(horizon)
|
|
||||||
cumulative_hazard = rows["death_rate"].to_numpy(dtype=np.float64) * np.power(h, rows["death_rho"].to_numpy(dtype=np.float64))
|
|
||||||
rows[f"death_cumhaz_h{h:g}y"] = cumulative_hazard
|
|
||||||
rows[f"death_risk_h{h:g}y"] = -np.expm1(-cumulative_hazard)
|
|
||||||
rows[f"death_observed_h{h:g}y"] = (
|
|
||||||
(rows["death_time"].to_numpy(dtype=np.float64) > rows["landmark_age"].to_numpy(dtype=np.float64))
|
|
||||||
& (rows["death_time"].to_numpy(dtype=np.float64) <= rows["landmark_age"].to_numpy(dtype=np.float64) + h)
|
|
||||||
).astype(np.int8)
|
|
||||||
|
|
||||||
output_dir = Path(args.output_path) if args.output_path else run_path / "weibull_death_parameter_stats_test"
|
|
||||||
output_dir.mkdir(parents=True, exist_ok=True)
|
|
||||||
|
|
||||||
rows.to_csv(output_dir / "death_weibull_parameters_by_landmark.csv", index=False)
|
|
||||||
value_cols = ["death_eta", "death_rate", "death_rho"]
|
|
||||||
for horizon in horizons.tolist():
|
|
||||||
h = float(horizon)
|
|
||||||
value_cols.extend([f"death_cumhaz_h{h:g}y", f"death_risk_h{h:g}y"])
|
|
||||||
|
|
||||||
quantile_summary(rows, [], value_cols).to_csv(output_dir / "death_weibull_parameter_summary_overall.csv", index=False)
|
|
||||||
quantile_summary(rows, ["landmark_age"], value_cols).to_csv(output_dir / "death_weibull_parameter_summary_by_landmark_age.csv", index=False)
|
|
||||||
quantile_summary(rows, ["sex_label"], value_cols).to_csv(output_dir / "death_weibull_parameter_summary_by_sex.csv", index=False)
|
|
||||||
quantile_summary(rows, ["sex_label", "landmark_age"], value_cols).to_csv(output_dir / "death_weibull_parameter_summary_by_sex_landmark_age.csv", index=False)
|
|
||||||
|
|
||||||
metadata = {
|
|
||||||
"run_path": str(run_path),
|
|
||||||
"config_path": str(config_path),
|
|
||||||
"checkpoint_path": str(ckpt_path),
|
|
||||||
"eval_split": str(args.eval_split),
|
|
||||||
"model_target_mode": model_target_mode,
|
|
||||||
"time_mode": str(cfg.get("time_mode")),
|
|
||||||
"dist_mode_config": str(cfg.get("dist_mode")),
|
|
||||||
"dist_mode_resolved": dist_mode,
|
|
||||||
"extra_info_types": cfg.get("extra_info_types"),
|
|
||||||
"death_token_id": death_idx,
|
|
||||||
"death_label_code": dataset.label_id_to_code.get(death_idx, "Death"),
|
|
||||||
"n_selected_patients": int(len(subset_indices)),
|
|
||||||
"n_landmark_rows": int(len(rows)),
|
|
||||||
"landmark_ages": [float(x) for x in landmark_ages.tolist()],
|
|
||||||
"horizons": [float(x) for x in horizons.tolist()],
|
|
||||||
}
|
|
||||||
with (output_dir / "metadata.json").open("w", encoding="utf-8") as f:
|
|
||||||
json.dump(metadata, f, indent=2)
|
|
||||||
|
|
||||||
if args.include_all_token_rho_summary and dist_mode == "weibull":
|
|
||||||
export_all_token_rho_summary(
|
|
||||||
model=model,
|
|
||||||
hidden_all=hidden_all,
|
|
||||||
dataset=dataset,
|
|
||||||
device=device,
|
|
||||||
output_dir=output_dir,
|
|
||||||
token_chunk_size=int(args.token_chunk_size),
|
|
||||||
row_batch_size=int(args.row_batch_size),
|
|
||||||
use_amp=use_amp,
|
|
||||||
horizons=horizons,
|
|
||||||
)
|
|
||||||
elif dist_mode == "mixed":
|
|
||||||
pd.DataFrame([
|
|
||||||
{
|
|
||||||
"dist_mode": dist_mode,
|
|
||||||
"disease_shape_available": False,
|
|
||||||
"death_shape_available": True,
|
|
||||||
"note": (
|
|
||||||
"The mixed model uses Weibull rho only for Death. "
|
|
||||||
"Non-death disease hazards are exponential, equivalent to fixed rho=1."
|
|
||||||
),
|
|
||||||
}
|
|
||||||
]).to_csv(output_dir / "disease_shape_not_available_for_mixed_model.csv", index=False)
|
|
||||||
|
|
||||||
print(f"Wrote Weibull shape parameter statistics to: {output_dir}")
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
main()
|
|
||||||
@@ -1,7 +0,0 @@
|
|||||||
"""Compatibility entry point for Weibull shape-parameter export."""
|
|
||||||
|
|
||||||
from export_weibull_death_parameter_stats import main
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
main()
|
|
||||||
@@ -1,511 +0,0 @@
|
|||||||
"""Shared landmark evaluation helpers for attribution scripts."""
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import argparse
|
|
||||||
import json
|
|
||||||
from pathlib import Path
|
|
||||||
from typing import Any, Dict, List, Optional, Sequence
|
|
||||||
|
|
||||||
import numpy as np
|
|
||||||
import pandas as pd
|
|
||||||
import torch
|
|
||||||
from torch.nn.utils.rnn import pad_sequence
|
|
||||||
from torch.utils.data import Dataset
|
|
||||||
|
|
||||||
from dataset import HealthDataset
|
|
||||||
from eval_data import load_sequence_eval_dataset
|
|
||||||
from evaluate_auc_v2 import (
|
|
||||||
LandmarkDataset,
|
|
||||||
build_model_from_dataset,
|
|
||||||
cfg_get,
|
|
||||||
make_eval_indices,
|
|
||||||
)
|
|
||||||
from models import DeepHealth
|
|
||||||
from readouts import build_readout
|
|
||||||
from targets import CHECKUP_IDX, NO_EVENT_IDX, PAD_IDX
|
|
||||||
from train_util import load_eid_file, load_extra_info_types_file
|
|
||||||
|
|
||||||
|
|
||||||
SPECIAL_TOKENS = {PAD_IDX, CHECKUP_IDX, NO_EVENT_IDX}
|
|
||||||
|
|
||||||
|
|
||||||
def parse_int_list(value: Any) -> Optional[List[int]]:
|
|
||||||
if value is None:
|
|
||||||
return None
|
|
||||||
if isinstance(value, (list, tuple, np.ndarray)):
|
|
||||||
return [int(x) for x in value]
|
|
||||||
text = str(value).strip()
|
|
||||||
if text == "":
|
|
||||||
return None
|
|
||||||
if text.startswith("["):
|
|
||||||
values = json.loads(text)
|
|
||||||
if not isinstance(values, list):
|
|
||||||
raise ValueError(f"Expected a JSON list, got {type(values).__name__}")
|
|
||||||
return [int(x) for x in values]
|
|
||||||
return [int(x.strip()) for x in text.split(",") if x.strip()]
|
|
||||||
|
|
||||||
|
|
||||||
def load_extra_info_types(value: Any) -> Optional[List[int]]:
|
|
||||||
if value is None:
|
|
||||||
return None
|
|
||||||
text = str(value)
|
|
||||||
path = Path(text)
|
|
||||||
if path.exists():
|
|
||||||
return load_extra_info_types_file(text)
|
|
||||||
return parse_int_list(value)
|
|
||||||
|
|
||||||
|
|
||||||
def make_landmark_ages(start: float, stop: float, step: float) -> np.ndarray:
|
|
||||||
if step <= 0:
|
|
||||||
raise ValueError("landmark_step must be positive")
|
|
||||||
if stop < start:
|
|
||||||
raise ValueError("landmark_stop must be >= landmark_start")
|
|
||||||
# Include stop when it lands on the grid, e.g. 40,45,...,80.
|
|
||||||
return np.arange(start, stop + step * 0.5, step, dtype=np.float32)
|
|
||||||
|
|
||||||
|
|
||||||
def build_first_occurrence_maps_for_landmarks(
|
|
||||||
dataset: HealthDataset,
|
|
||||||
subset_indices: np.ndarray,
|
|
||||||
) -> Dict[int, tuple[np.ndarray, np.ndarray]]:
|
|
||||||
first_lists: Dict[int, list[tuple[int, float]]] = {}
|
|
||||||
for patient_id, dataset_index in enumerate(np.asarray(subset_indices, dtype=np.int64).tolist()):
|
|
||||||
s = dataset.samples[int(dataset_index)]
|
|
||||||
seq_event = np.asarray(s["event_seq"], dtype=np.int64)
|
|
||||||
seq_time = np.asarray(s["time_seq"], dtype=np.float32)
|
|
||||||
tgt_event = np.asarray(s["target_event_seq"], dtype=np.int64)
|
|
||||||
tgt_time = np.asarray(s["target_time_seq"], dtype=np.float32)
|
|
||||||
if seq_event.size == 0 or tgt_event.size == 0:
|
|
||||||
continue
|
|
||||||
|
|
||||||
full_event = np.concatenate([seq_event, tgt_event[-1:]])
|
|
||||||
full_time = np.concatenate([seq_time, tgt_time[-1:]])
|
|
||||||
uniq_tokens, first_idx = np.unique(full_event, return_index=True)
|
|
||||||
for token, idx in zip(uniq_tokens.tolist(), first_idx.tolist()):
|
|
||||||
token = int(token)
|
|
||||||
if token in SPECIAL_TOKENS:
|
|
||||||
continue
|
|
||||||
first_lists.setdefault(token, []).append((patient_id, float(full_time[int(idx)])))
|
|
||||||
|
|
||||||
return {
|
|
||||||
int(token): (
|
|
||||||
np.asarray([p for p, _ in pairs], dtype=np.int32),
|
|
||||||
np.asarray([t for _, t in pairs], dtype=np.float32),
|
|
||||||
)
|
|
||||||
for token, pairs in first_lists.items()
|
|
||||||
if pairs
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def normalize_eval_split(args: argparse.Namespace, cfg: Dict[str, Any]) -> str:
|
|
||||||
eval_split = str(cfg_get(args, cfg, "eval_split", "test")).lower()
|
|
||||||
if eval_split in {"valid", "validation"}:
|
|
||||||
return "val"
|
|
||||||
if eval_split not in {"train", "val", "test", "all"}:
|
|
||||||
raise ValueError(f"Unsupported eval_split={eval_split!r}")
|
|
||||||
return eval_split
|
|
||||||
|
|
||||||
|
|
||||||
def load_eval_sequence_dataset(
|
|
||||||
args: argparse.Namespace,
|
|
||||||
cfg: Dict[str, Any],
|
|
||||||
) -> tuple[Any, np.ndarray, str, str]:
|
|
||||||
eval_split = normalize_eval_split(args, cfg)
|
|
||||||
model_target_mode = str(cfg.get("model_target_mode", "next_token")).lower()
|
|
||||||
data_prefix = str(cfg.get("data_prefix", "ukb"))
|
|
||||||
labels_file = str(cfg.get("labels_file", "labels.csv"))
|
|
||||||
no_event_interval_years = float(cfg.get("no_event_interval_years", 5.0))
|
|
||||||
include_no_event_in_uts_target = bool(cfg.get("include_no_event_in_uts_target", False))
|
|
||||||
extra_info_types = load_extra_info_types(args.extra_info_types)
|
|
||||||
if extra_info_types is None:
|
|
||||||
extra_info_types = parse_int_list(cfg.get("extra_info_types", None))
|
|
||||||
|
|
||||||
print("Loading one sequence eval dataset...")
|
|
||||||
dataset = load_sequence_eval_dataset(
|
|
||||||
model_target_mode=model_target_mode,
|
|
||||||
data_prefix=data_prefix,
|
|
||||||
labels_file=labels_file,
|
|
||||||
no_event_interval_years=no_event_interval_years,
|
|
||||||
include_no_event_in_uts_target=include_no_event_in_uts_target,
|
|
||||||
min_history_events=int(cfg.get("all_future_min_history_events", 1)),
|
|
||||||
min_future_events=int(cfg.get("all_future_min_future_events", 1)),
|
|
||||||
extra_info_types=extra_info_types,
|
|
||||||
)
|
|
||||||
|
|
||||||
train_eid_file = cfg_get(args, cfg, "train_eid_file", "ukb_train_eid.csv")
|
|
||||||
val_eid_file = cfg_get(args, cfg, "val_eid_file", "ukb_val_eid.csv")
|
|
||||||
test_eid_file = cfg_get(args, cfg, "test_eid_file", "ukb_test_eid.csv")
|
|
||||||
split_files_exist = all(
|
|
||||||
Path(str(path)).exists()
|
|
||||||
for path in (train_eid_file, val_eid_file, test_eid_file)
|
|
||||||
)
|
|
||||||
|
|
||||||
if eval_split != "all" and split_files_exist:
|
|
||||||
split_files = {
|
|
||||||
"train": train_eid_file,
|
|
||||||
"val": val_eid_file,
|
|
||||||
"test": test_eid_file,
|
|
||||||
}
|
|
||||||
selected_eids = load_eid_file(split_files[eval_split])
|
|
||||||
out = np.asarray(
|
|
||||||
[
|
|
||||||
idx
|
|
||||||
for idx, sample in enumerate(dataset.samples)
|
|
||||||
if int(sample["eid"]) in selected_eids
|
|
||||||
],
|
|
||||||
dtype=np.int64,
|
|
||||||
)
|
|
||||||
if out.size == 0:
|
|
||||||
raise ValueError(
|
|
||||||
f"No samples found for eval_split={eval_split!r} using {split_files[eval_split]}"
|
|
||||||
)
|
|
||||||
split_source = "eid_files"
|
|
||||||
else:
|
|
||||||
if eval_split == "all":
|
|
||||||
out = np.arange(len(dataset.samples), dtype=np.int64)
|
|
||||||
split_source = "all"
|
|
||||||
else:
|
|
||||||
out = make_eval_indices(dataset, args, cfg)
|
|
||||||
split_source = "ratio_split"
|
|
||||||
|
|
||||||
subset_size = cfg_get(args, cfg, "dataset_subset_size", None)
|
|
||||||
if subset_size is not None and int(subset_size) > 0:
|
|
||||||
out = out[: int(subset_size)]
|
|
||||||
return dataset, np.asarray(out, dtype=np.int64), eval_split, split_source
|
|
||||||
|
|
||||||
|
|
||||||
def load_organ_groups(
|
|
||||||
path: Path,
|
|
||||||
*,
|
|
||||||
vocab_size: int,
|
|
||||||
) -> tuple[dict[str, list[int]], dict[str, str], dict[int, str]]:
|
|
||||||
table = pd.read_csv(path)
|
|
||||||
required = {"token_id", "organ_system", "organ_system_label", "is_death"}
|
|
||||||
missing = required - set(table.columns)
|
|
||||||
if missing:
|
|
||||||
raise ValueError(f"{path} is missing columns: {sorted(missing)}")
|
|
||||||
|
|
||||||
death_idx = int(vocab_size) - 1
|
|
||||||
groups: dict[str, list[int]] = {}
|
|
||||||
labels: dict[str, str] = {}
|
|
||||||
token_to_group: dict[int, str] = {}
|
|
||||||
for row in table.itertuples(index=False):
|
|
||||||
token = int(getattr(row, "token_id"))
|
|
||||||
if token in SPECIAL_TOKENS or token == death_idx:
|
|
||||||
continue
|
|
||||||
if token < 0 or token >= int(vocab_size):
|
|
||||||
continue
|
|
||||||
if int(getattr(row, "is_death")) == 1:
|
|
||||||
continue
|
|
||||||
group = str(getattr(row, "organ_system"))
|
|
||||||
label = str(getattr(row, "organ_system_label"))
|
|
||||||
groups.setdefault(group, []).append(token)
|
|
||||||
labels[group] = label
|
|
||||||
token_to_group[token] = group
|
|
||||||
|
|
||||||
groups = {k: sorted(set(v)) for k, v in groups.items() if v}
|
|
||||||
return groups, labels, token_to_group
|
|
||||||
|
|
||||||
|
|
||||||
class IndexedLandmarkDataset(Dataset):
|
|
||||||
def __init__(self, base: LandmarkDataset) -> None:
|
|
||||||
self.base = base
|
|
||||||
|
|
||||||
def __len__(self) -> int:
|
|
||||||
return len(self.base)
|
|
||||||
|
|
||||||
def __getitem__(self, idx: int) -> Dict[str, torch.Tensor]:
|
|
||||||
item = dict(self.base[idx])
|
|
||||||
item["row_idx"] = torch.tensor(int(idx), dtype=torch.long)
|
|
||||||
return item
|
|
||||||
|
|
||||||
|
|
||||||
def collate_indexed_landmark_fn(batch: List[Dict[str, torch.Tensor]]) -> Dict[str, torch.Tensor]:
|
|
||||||
event_seq = pad_sequence(
|
|
||||||
[x["event_seq"] for x in batch], batch_first=True, padding_value=PAD_IDX
|
|
||||||
)
|
|
||||||
time_seq = pad_sequence(
|
|
||||||
[x["time_seq"] for x in batch], batch_first=True, padding_value=0.0
|
|
||||||
)
|
|
||||||
readout_mask = pad_sequence(
|
|
||||||
[x["readout_mask"] for x in batch], batch_first=True, padding_value=False
|
|
||||||
)
|
|
||||||
other_type = pad_sequence(
|
|
||||||
[x["other_type"] for x in batch], batch_first=True, padding_value=0
|
|
||||||
)
|
|
||||||
other_value = pad_sequence(
|
|
||||||
[x["other_value"] for x in batch], batch_first=True, padding_value=0.0
|
|
||||||
)
|
|
||||||
other_value_kind = pad_sequence(
|
|
||||||
[x["other_value_kind"] for x in batch], batch_first=True, padding_value=0
|
|
||||||
)
|
|
||||||
other_time = pad_sequence(
|
|
||||||
[x["other_time"] for x in batch], batch_first=True, padding_value=0.0
|
|
||||||
)
|
|
||||||
return {
|
|
||||||
"event_seq": event_seq,
|
|
||||||
"time_seq": time_seq,
|
|
||||||
"padding_mask": event_seq > PAD_IDX,
|
|
||||||
"readout_mask": readout_mask,
|
|
||||||
"sex": torch.stack([x["sex"] for x in batch]),
|
|
||||||
"other_type": other_type,
|
|
||||||
"other_value": other_value,
|
|
||||||
"other_value_kind": other_value_kind,
|
|
||||||
"other_time": other_time,
|
|
||||||
"landmark_pos": torch.stack([x["landmark_pos"] for x in batch]),
|
|
||||||
"t_query": torch.stack([x["t_query"] for x in batch]),
|
|
||||||
"patient_id": torch.stack([x["patient_id"] for x in batch]),
|
|
||||||
"landmark_age": torch.stack([x["landmark_age"] for x in batch]),
|
|
||||||
"followup_end_time": torch.stack([x["followup_end_time"] for x in batch]),
|
|
||||||
"death_time": torch.stack([x["death_time"] for x in batch]),
|
|
||||||
"row_idx": torch.stack([x["row_idx"] for x in batch]),
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def build_group_ablated_slice(
|
|
||||||
batch: Dict[str, torch.Tensor],
|
|
||||||
token_ids: Sequence[int],
|
|
||||||
row_indices: torch.Tensor,
|
|
||||||
) -> Dict[str, torch.Tensor]:
|
|
||||||
"""Build one fixed-width ablated slice without rebuilding variable-length rows."""
|
|
||||||
event_seq = batch["event_seq"]
|
|
||||||
|
|
||||||
out: Dict[str, torch.Tensor] = {}
|
|
||||||
out["event_seq"] = event_seq[row_indices].clone()
|
|
||||||
out["time_seq"] = batch["time_seq"][row_indices]
|
|
||||||
out["readout_mask"] = batch["readout_mask"][row_indices].clone()
|
|
||||||
out["padding_mask"] = batch["padding_mask"][row_indices].bool().clone()
|
|
||||||
out["landmark_pos"] = batch["landmark_pos"][row_indices].clone()
|
|
||||||
|
|
||||||
seq_len = int(event_seq.shape[1])
|
|
||||||
positions = torch.arange(seq_len, device=event_seq.device)[None, :]
|
|
||||||
ids = torch.as_tensor(token_ids, dtype=event_seq.dtype, device=event_seq.device)
|
|
||||||
remove = torch.isin(out["event_seq"], ids) & out["padding_mask"]
|
|
||||||
out["event_seq"] = torch.where(
|
|
||||||
remove,
|
|
||||||
torch.full_like(out["event_seq"], PAD_IDX),
|
|
||||||
out["event_seq"],
|
|
||||||
)
|
|
||||||
out["padding_mask"] &= ~remove
|
|
||||||
out["readout_mask"] &= ~remove
|
|
||||||
|
|
||||||
has_valid = out["padding_mask"].any(dim=1)
|
|
||||||
if not bool(has_valid.all().item()):
|
|
||||||
empty_rows = torch.nonzero(~has_valid, as_tuple=False).flatten()
|
|
||||||
out["event_seq"][empty_rows, 0] = CHECKUP_IDX
|
|
||||||
out["time_seq"][empty_rows, 0] = batch["t_query"][row_indices[empty_rows]].to(
|
|
||||||
dtype=out["time_seq"].dtype
|
|
||||||
)
|
|
||||||
out["padding_mask"][empty_rows, 0] = True
|
|
||||||
out["readout_mask"][empty_rows, 0] = True
|
|
||||||
out["landmark_pos"][empty_rows] = 0
|
|
||||||
|
|
||||||
has_readout = out["readout_mask"].any(dim=1)
|
|
||||||
if not bool(has_readout.all().item()):
|
|
||||||
rows = torch.nonzero(~has_readout, as_tuple=False).flatten()
|
|
||||||
local_valid = out["padding_mask"][rows]
|
|
||||||
last_pos = torch.where(
|
|
||||||
local_valid,
|
|
||||||
positions.expand(local_valid.shape[0], -1),
|
|
||||||
torch.zeros_like(positions.expand(local_valid.shape[0], -1)),
|
|
||||||
).amax(dim=1)
|
|
||||||
out["readout_mask"][rows] = False
|
|
||||||
out["readout_mask"][rows, last_pos] = True
|
|
||||||
out["landmark_pos"][rows] = last_pos.to(dtype=out["landmark_pos"].dtype)
|
|
||||||
|
|
||||||
repeated_keys = (
|
|
||||||
"sex",
|
|
||||||
"other_type",
|
|
||||||
"other_value",
|
|
||||||
"other_value_kind",
|
|
||||||
"other_time",
|
|
||||||
"t_query",
|
|
||||||
"patient_id",
|
|
||||||
"landmark_age",
|
|
||||||
"followup_end_time",
|
|
||||||
"death_time",
|
|
||||||
"row_idx",
|
|
||||||
)
|
|
||||||
for key in repeated_keys:
|
|
||||||
out[key] = batch[key][row_indices]
|
|
||||||
return out
|
|
||||||
|
|
||||||
|
|
||||||
def concat_tensor_batches(chunks: Sequence[Dict[str, torch.Tensor]]) -> Dict[str, torch.Tensor]:
|
|
||||||
return {
|
|
||||||
key: torch.cat([chunk[key] for chunk in chunks], dim=0)
|
|
||||||
for key in chunks[0]
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def iter_group_ablated_batches(
|
|
||||||
batch: Dict[str, torch.Tensor],
|
|
||||||
group_names: Sequence[str],
|
|
||||||
organ_groups: dict[str, list[int]],
|
|
||||||
occurred: torch.Tensor,
|
|
||||||
max_batch_size: int,
|
|
||||||
):
|
|
||||||
"""Yield ablated chunks as soon as enough rows are available for a forward pass."""
|
|
||||||
pending_batches: list[Dict[str, torch.Tensor]] = []
|
|
||||||
pending_groups: list[str] = []
|
|
||||||
pending_rows: list[int] = []
|
|
||||||
pending_n = 0
|
|
||||||
|
|
||||||
for group in group_names:
|
|
||||||
ids = torch.as_tensor(organ_groups[group], dtype=torch.long, device=occurred.device)
|
|
||||||
if ids.numel() == 0:
|
|
||||||
continue
|
|
||||||
active_rows = torch.nonzero(occurred[:, ids].any(dim=1), as_tuple=False).flatten()
|
|
||||||
if active_rows.numel() == 0:
|
|
||||||
continue
|
|
||||||
|
|
||||||
row_offset = 0
|
|
||||||
while row_offset < int(active_rows.numel()):
|
|
||||||
capacity = int(max_batch_size) - pending_n
|
|
||||||
row_stop = min(int(active_rows.numel()), row_offset + capacity)
|
|
||||||
row_indices = active_rows[row_offset:row_stop].to(device=batch["event_seq"].device)
|
|
||||||
chunk = build_group_ablated_slice(
|
|
||||||
batch=batch,
|
|
||||||
token_ids=organ_groups[group],
|
|
||||||
row_indices=row_indices,
|
|
||||||
)
|
|
||||||
chunk_n = int(row_indices.numel())
|
|
||||||
pending_batches.append(chunk)
|
|
||||||
pending_groups.extend([group] * chunk_n)
|
|
||||||
pending_rows.extend(int(x) for x in row_indices.detach().cpu().tolist())
|
|
||||||
pending_n += chunk_n
|
|
||||||
row_offset = row_stop
|
|
||||||
|
|
||||||
if pending_n >= int(max_batch_size):
|
|
||||||
yield concat_tensor_batches(pending_batches), pending_groups, pending_rows
|
|
||||||
pending_batches = []
|
|
||||||
pending_groups = []
|
|
||||||
pending_rows = []
|
|
||||||
pending_n = 0
|
|
||||||
|
|
||||||
if pending_batches:
|
|
||||||
yield concat_tensor_batches(pending_batches), pending_groups, pending_rows
|
|
||||||
|
|
||||||
|
|
||||||
@torch.no_grad()
|
|
||||||
def infer_landmark_hidden(
|
|
||||||
*,
|
|
||||||
model: DeepHealth,
|
|
||||||
batch: Dict[str, torch.Tensor],
|
|
||||||
device: torch.device,
|
|
||||||
model_target_mode: str,
|
|
||||||
readout_name: str,
|
|
||||||
readout_reduce: str,
|
|
||||||
) -> torch.Tensor:
|
|
||||||
batch_dev = {
|
|
||||||
k: (v.to(device, non_blocking=True) if isinstance(v, torch.Tensor) else v)
|
|
||||||
for k, v in batch.items()
|
|
||||||
}
|
|
||||||
if model_target_mode == "all_future":
|
|
||||||
return model(
|
|
||||||
event_seq=batch_dev["event_seq"].long(),
|
|
||||||
time_seq=batch_dev["time_seq"].float(),
|
|
||||||
sex=batch_dev["sex"].long(),
|
|
||||||
padding_mask=batch_dev["padding_mask"].bool(),
|
|
||||||
t_query=batch_dev["t_query"].float(),
|
|
||||||
other_type=batch_dev["other_type"].long(),
|
|
||||||
other_value=batch_dev["other_value"].float(),
|
|
||||||
other_value_kind=batch_dev["other_value_kind"].long(),
|
|
||||||
other_time=batch_dev["other_time"].float(),
|
|
||||||
target_mode="all_future",
|
|
||||||
)
|
|
||||||
|
|
||||||
hidden = model(
|
|
||||||
event_seq=batch_dev["event_seq"].long(),
|
|
||||||
time_seq=batch_dev["time_seq"].float(),
|
|
||||||
sex=batch_dev["sex"].long(),
|
|
||||||
padding_mask=batch_dev["padding_mask"].bool(),
|
|
||||||
other_type=batch_dev["other_type"].long(),
|
|
||||||
other_value=batch_dev["other_value"].float(),
|
|
||||||
other_value_kind=batch_dev["other_value_kind"].long(),
|
|
||||||
other_time=batch_dev["other_time"].float(),
|
|
||||||
target_mode="next_token",
|
|
||||||
)
|
|
||||||
readout = build_readout(readout_name, reduce=readout_reduce)
|
|
||||||
readout_out = readout(
|
|
||||||
hidden=hidden,
|
|
||||||
time_seq=batch_dev["time_seq"].float(),
|
|
||||||
padding_mask=batch_dev["padding_mask"].bool(),
|
|
||||||
readout_mask=batch_dev["readout_mask"].bool(),
|
|
||||||
)
|
|
||||||
return readout_out.hidden.gather(
|
|
||||||
1,
|
|
||||||
batch_dev["landmark_pos"].long()[:, None, None].expand(
|
|
||||||
-1, 1, readout_out.hidden.shape[-1]
|
|
||||||
),
|
|
||||||
).squeeze(1)
|
|
||||||
|
|
||||||
|
|
||||||
def make_occurred_mask(
|
|
||||||
event_seq: torch.Tensor,
|
|
||||||
*,
|
|
||||||
vocab_size: int,
|
|
||||||
device: torch.device,
|
|
||||||
) -> torch.Tensor:
|
|
||||||
occurred = torch.zeros(event_seq.shape[0], int(vocab_size), dtype=torch.bool, device=device)
|
|
||||||
valid = (event_seq >= 0) & (event_seq < int(vocab_size))
|
|
||||||
safe = event_seq.clamp(min=0, max=int(vocab_size) - 1).to(device)
|
|
||||||
occurred.scatter_(1, safe, valid.to(device))
|
|
||||||
return occurred
|
|
||||||
|
|
||||||
|
|
||||||
def mortality_hazard_from_risk(risk: torch.Tensor, eps: float = 1e-7) -> torch.Tensor:
|
|
||||||
return -torch.log1p(-risk.clamp(0.0, 1.0 - float(eps)))
|
|
||||||
|
|
||||||
|
|
||||||
def death_risk_for_batch(
|
|
||||||
*,
|
|
||||||
model: DeepHealth,
|
|
||||||
batch: Dict[str, torch.Tensor],
|
|
||||||
device: torch.device,
|
|
||||||
model_target_mode: str,
|
|
||||||
readout_name: str,
|
|
||||||
readout_reduce: str,
|
|
||||||
dist_mode: str,
|
|
||||||
tau: float,
|
|
||||||
) -> torch.Tensor:
|
|
||||||
hidden = infer_landmark_hidden(
|
|
||||||
model=model,
|
|
||||||
batch=batch,
|
|
||||||
device=device,
|
|
||||||
model_target_mode=model_target_mode,
|
|
||||||
readout_name=readout_name,
|
|
||||||
readout_reduce=readout_reduce,
|
|
||||||
)
|
|
||||||
logits = model.calc_risk(hidden)
|
|
||||||
rho = model.calc_weibull_rho(hidden) if dist_mode == "weibull" else None
|
|
||||||
death_rho = model.calc_death_rho(hidden) if dist_mode == "mixed" else None
|
|
||||||
probabilities = probabilities_from_logits(
|
|
||||||
logits,
|
|
||||||
tau,
|
|
||||||
dist_mode=dist_mode,
|
|
||||||
rho=rho,
|
|
||||||
death_rho=death_rho,
|
|
||||||
)
|
|
||||||
return death_risk_from_probabilities(probabilities)
|
|
||||||
|
|
||||||
|
|
||||||
def historical_counts_by_group(
|
|
||||||
tokens: np.ndarray,
|
|
||||||
*,
|
|
||||||
death_idx: int,
|
|
||||||
token_to_group: dict[int, str],
|
|
||||||
group_names: Sequence[str],
|
|
||||||
) -> tuple[int, dict[str, int]]:
|
|
||||||
unique_tokens = {
|
|
||||||
int(token)
|
|
||||||
for token in np.asarray(tokens, dtype=np.int64).tolist()
|
|
||||||
if int(token) not in SPECIAL_TOKENS and int(token) != int(death_idx)
|
|
||||||
}
|
|
||||||
total = len(unique_tokens)
|
|
||||||
out = {group: 0 for group in group_names}
|
|
||||||
for token in unique_tokens:
|
|
||||||
group = token_to_group.get(token)
|
|
||||||
if group in out:
|
|
||||||
out[group] += 1
|
|
||||||
return total, out
|
|
||||||
@@ -1,553 +0,0 @@
|
|||||||
#!/usr/bin/env Rscript
|
|
||||||
|
|
||||||
# Paper-grade single-panel figures supporting the conclusion that fixed-landmark
|
|
||||||
# horizon evaluation favors all_future over next_token.
|
|
||||||
#
|
|
||||||
# Outputs are written as separate panel files. This script intentionally does not
|
|
||||||
# combine panels with plot_grid().
|
|
||||||
|
|
||||||
suppressPackageStartupMessages({
|
|
||||||
library(cowplot)
|
|
||||||
library(dplyr)
|
|
||||||
library(ggplot2)
|
|
||||||
library(jsonlite)
|
|
||||||
library(readr)
|
|
||||||
library(stringr)
|
|
||||||
library(tibble)
|
|
||||||
library(tidyr)
|
|
||||||
})
|
|
||||||
|
|
||||||
root_dir <- "."
|
|
||||||
runs_dir <- file.path(root_dir, "runs")
|
|
||||||
out_dir <- file.path(root_dir, "figures_next_token_to_all_future_absolute_smoking_alcohol_bmi")
|
|
||||||
dir.create(out_dir, showWarnings = FALSE, recursive = TRUE)
|
|
||||||
|
|
||||||
required_time_mode <- "absolute"
|
|
||||||
required_extra_info_types <- c(11L, 66L, 67L)
|
|
||||||
required_extra_info_signature <- paste(sort(required_extra_info_types), collapse = ",")
|
|
||||||
|
|
||||||
theme_set(
|
|
||||||
theme_cowplot(font_size = 9) +
|
|
||||||
theme(
|
|
||||||
plot.background = element_rect(fill = "white", color = NA),
|
|
||||||
panel.background = element_rect(fill = "white", color = NA),
|
|
||||||
legend.background = element_rect(fill = "white", color = NA),
|
|
||||||
legend.key = element_rect(fill = "white", color = NA)
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
target_cols <- c(
|
|
||||||
"next_token" = "#B54A3A",
|
|
||||||
"all_future" = "#2C7FB8"
|
|
||||||
)
|
|
||||||
|
|
||||||
dist_shapes <- c(
|
|
||||||
"exponential" = 16,
|
|
||||||
"weibull" = 17,
|
|
||||||
"mixed" = 15
|
|
||||||
)
|
|
||||||
|
|
||||||
read_run_config <- function(run_path) {
|
|
||||||
cfg_path <- file.path(run_path, "train_config.json")
|
|
||||||
if (!file.exists(cfg_path)) return(NULL)
|
|
||||||
cfg <- jsonlite::read_json(cfg_path, simplifyVector = TRUE)
|
|
||||||
extra_info_types <- cfg$extra_info_types %||% integer(0)
|
|
||||||
extra_info_signature <- paste(sort(as.integer(extra_info_types)), collapse = ",")
|
|
||||||
tibble(
|
|
||||||
run = basename(run_path),
|
|
||||||
model_target_mode = as.character(cfg$model_target_mode %||% NA_character_),
|
|
||||||
target_mode = as.character(cfg$target_mode %||% NA_character_),
|
|
||||||
dist_mode = as.character(cfg$dist_mode %||% NA_character_),
|
|
||||||
time_mode = as.character(cfg$time_mode %||% NA_character_),
|
|
||||||
readout_name = as.character(cfg$readout_name %||% NA_character_),
|
|
||||||
attn_mask_mode = as.character(cfg$attn_mask_mode %||% NA_character_),
|
|
||||||
extra_info_signature = extra_info_signature
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
`%||%` <- function(x, y) {
|
|
||||||
if (is.null(x) || length(x) == 0) y else x
|
|
||||||
}
|
|
||||||
|
|
||||||
load_one_result <- function(run_path, file_name, eval_family) {
|
|
||||||
cfg <- read_run_config(run_path)
|
|
||||||
if (is.null(cfg)) return(NULL)
|
|
||||||
fp <- file.path(run_path, file_name)
|
|
||||||
if (!file.exists(fp)) return(NULL)
|
|
||||||
|
|
||||||
df <- suppressMessages(readr::read_csv(fp, show_col_types = FALSE))
|
|
||||||
if (!("auc" %in% names(df)) || nrow(df) == 0) return(NULL)
|
|
||||||
|
|
||||||
out <- df %>%
|
|
||||||
mutate(
|
|
||||||
run = basename(run_path),
|
|
||||||
eval_family = eval_family,
|
|
||||||
auc = as.numeric(auc)
|
|
||||||
) %>%
|
|
||||||
left_join(cfg, by = "run", suffix = c("", "_cfg"))
|
|
||||||
|
|
||||||
coalesce_joined <- function(data, col) {
|
|
||||||
cfg_col <- paste0(col, "_cfg")
|
|
||||||
if (col %in% names(data) && cfg_col %in% names(data)) {
|
|
||||||
dplyr::coalesce(data[[col]], data[[cfg_col]])
|
|
||||||
} else if (col %in% names(data)) {
|
|
||||||
data[[col]]
|
|
||||||
} else if (cfg_col %in% names(data)) {
|
|
||||||
data[[cfg_col]]
|
|
||||||
} else {
|
|
||||||
rep(NA_character_, nrow(data))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
for (col in c("model_target_mode", "target_mode", "dist_mode", "time_mode", "readout_name", "attn_mask_mode")) {
|
|
||||||
out[[col]] <- coalesce_joined(out, col)
|
|
||||||
}
|
|
||||||
|
|
||||||
out %>%
|
|
||||||
select(-any_of(c(
|
|
||||||
"model_target_mode_cfg", "target_mode_cfg", "dist_mode_cfg",
|
|
||||||
"time_mode_cfg", "readout_name_cfg", "attn_mask_mode_cfg"
|
|
||||||
)))
|
|
||||||
}
|
|
||||||
|
|
||||||
run_paths <- list.dirs(runs_dir, recursive = FALSE, full.names = TRUE)
|
|
||||||
|
|
||||||
landmark_auc <- bind_rows(lapply(
|
|
||||||
run_paths,
|
|
||||||
load_one_result,
|
|
||||||
file_name = "df_auc_landmark.csv",
|
|
||||||
eval_family = "Fixed landmark + horizon"
|
|
||||||
)) %>%
|
|
||||||
filter(time_mode == "absolute")
|
|
||||||
|
|
||||||
token_auc <- bind_rows(lapply(
|
|
||||||
run_paths,
|
|
||||||
load_one_result,
|
|
||||||
file_name = "df_both.csv",
|
|
||||||
eval_family = "Delphi2M-style token"
|
|
||||||
)) %>%
|
|
||||||
filter(time_mode == "absolute")
|
|
||||||
|
|
||||||
if (nrow(landmark_auc) == 0) {
|
|
||||||
stop("No landmark AUC files found under runs/*/df_auc_landmark.csv")
|
|
||||||
}
|
|
||||||
if (nrow(token_auc) == 0) {
|
|
||||||
stop("No token AUC files found under runs/*/df_both.csv")
|
|
||||||
}
|
|
||||||
|
|
||||||
landmark_auc <- landmark_auc %>%
|
|
||||||
filter(
|
|
||||||
time_mode == required_time_mode,
|
|
||||||
extra_info_signature == required_extra_info_signature
|
|
||||||
)
|
|
||||||
|
|
||||||
token_auc <- token_auc %>%
|
|
||||||
filter(
|
|
||||||
time_mode == required_time_mode,
|
|
||||||
extra_info_signature == required_extra_info_signature
|
|
||||||
)
|
|
||||||
|
|
||||||
if (nrow(landmark_auc) == 0 || nrow(token_auc) == 0) {
|
|
||||||
stop(
|
|
||||||
"No AUC rows remain after filtering for time_mode='",
|
|
||||||
required_time_mode,
|
|
||||||
"' and extra_info_types='",
|
|
||||||
required_extra_info_signature,
|
|
||||||
"'."
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
message(
|
|
||||||
"Using runs with time_mode='", required_time_mode,
|
|
||||||
"' and extra_info_types='", required_extra_info_signature, "':"
|
|
||||||
)
|
|
||||||
print(sort(unique(landmark_auc$run)))
|
|
||||||
|
|
||||||
classify_endpoint <- function(data) {
|
|
||||||
data %>%
|
|
||||||
mutate(
|
|
||||||
endpoint_type = if_else(
|
|
||||||
str_to_lower(as.character(label_code)) == "death",
|
|
||||||
"Death",
|
|
||||||
"Non-death disease"
|
|
||||||
),
|
|
||||||
endpoint_type = factor(endpoint_type, levels = c("Non-death disease", "Death"))
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
landmark_auc <- classify_endpoint(landmark_auc)
|
|
||||||
token_auc <- classify_endpoint(token_auc)
|
|
||||||
|
|
||||||
landmark_auc_disease <- landmark_auc %>% filter(endpoint_type == "Non-death disease")
|
|
||||||
token_auc_disease <- token_auc %>% filter(endpoint_type == "Non-death disease")
|
|
||||||
landmark_auc_death <- landmark_auc %>% filter(endpoint_type == "Death")
|
|
||||||
token_auc_death <- token_auc %>% filter(endpoint_type == "Death")
|
|
||||||
|
|
||||||
if (nrow(landmark_auc_death) == 0 || nrow(token_auc_death) == 0) {
|
|
||||||
warning("Death rows were not found in one or both AUC tables.")
|
|
||||||
}
|
|
||||||
|
|
||||||
auc_all <- bind_rows(
|
|
||||||
landmark_auc_disease %>% mutate(horizon = as.numeric(horizon), offset = NA_real_),
|
|
||||||
token_auc_disease %>% mutate(horizon = NA_real_, offset = as.numeric(offset))
|
|
||||||
) %>%
|
|
||||||
mutate(
|
|
||||||
model_target_mode = factor(model_target_mode, levels = c("next_token", "all_future")),
|
|
||||||
eval_family = factor(eval_family, levels = c("Delphi2M-style token", "Fixed landmark + horizon")),
|
|
||||||
dist_mode = factor(dist_mode, levels = c("exponential", "weibull", "mixed")),
|
|
||||||
model_label = recode(
|
|
||||||
as.character(model_target_mode),
|
|
||||||
"next_token" = "next-token objective",
|
|
||||||
"all_future" = "all-future objective"
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
mean_ci <- function(x) {
|
|
||||||
x <- x[is.finite(x)]
|
|
||||||
n <- length(x)
|
|
||||||
m <- mean(x)
|
|
||||||
se <- sd(x) / sqrt(n)
|
|
||||||
tibble(mean = m, ymin = m - 1.96 * se, ymax = m + 1.96 * se, n = n)
|
|
||||||
}
|
|
||||||
|
|
||||||
save_panel <- function(plot, name, width = 3.6, height = 3.0) {
|
|
||||||
pdf_path <- file.path(out_dir, paste0(name, ".pdf"))
|
|
||||||
png_path <- file.path(out_dir, paste0(name, ".png"))
|
|
||||||
cowplot::save_plot(pdf_path, plot, base_width = width, base_height = height, bg = "white")
|
|
||||||
cowplot::save_plot(png_path, plot, base_width = width, base_height = height, dpi = 600, bg = "white")
|
|
||||||
message("Wrote: ", pdf_path)
|
|
||||||
message("Wrote: ", png_path)
|
|
||||||
}
|
|
||||||
|
|
||||||
# Panel 1: run-level mean AUC under the clinically aligned landmark/horizon task.
|
|
||||||
# Death is excluded here and plotted separately below.
|
|
||||||
landmark_run <- landmark_auc_disease %>%
|
|
||||||
mutate(model_target_mode = factor(model_target_mode, levels = c("next_token", "all_future"))) %>%
|
|
||||||
group_by(run, model_target_mode, dist_mode, time_mode, target_mode) %>%
|
|
||||||
summarise(mean_auc = mean(auc, na.rm = TRUE), median_auc = median(auc, na.rm = TRUE), .groups = "drop")
|
|
||||||
|
|
||||||
landmark_summary <- landmark_run %>%
|
|
||||||
group_by(model_target_mode) %>%
|
|
||||||
summarise(mean_ci(mean_auc), .groups = "drop")
|
|
||||||
|
|
||||||
p1 <- ggplot(landmark_run, aes(x = model_target_mode, y = mean_auc)) +
|
|
||||||
geom_point(
|
|
||||||
aes(color = model_target_mode, shape = dist_mode),
|
|
||||||
position = position_jitter(width = 0.09, height = 0, seed = 1),
|
|
||||||
size = 2.2,
|
|
||||||
alpha = 0.88
|
|
||||||
) +
|
|
||||||
geom_errorbar(
|
|
||||||
data = landmark_summary,
|
|
||||||
aes(x = model_target_mode, y = mean, ymin = ymin, ymax = ymax, color = model_target_mode),
|
|
||||||
width = 0.12,
|
|
||||||
linewidth = 0.55,
|
|
||||||
inherit.aes = FALSE
|
|
||||||
) +
|
|
||||||
geom_point(
|
|
||||||
data = landmark_summary,
|
|
||||||
aes(x = model_target_mode, y = mean, color = model_target_mode),
|
|
||||||
size = 3.4,
|
|
||||||
inherit.aes = FALSE
|
|
||||||
) +
|
|
||||||
scale_color_manual(values = target_cols, guide = "none") +
|
|
||||||
scale_shape_manual(values = dist_shapes, na.translate = FALSE) +
|
|
||||||
scale_x_discrete(labels = c("next_token", "all_future")) +
|
|
||||||
coord_cartesian(ylim = c(0.58, 0.78)) +
|
|
||||||
labs(
|
|
||||||
x = NULL,
|
|
||||||
y = "Mean AUC per run",
|
|
||||||
shape = "Risk head",
|
|
||||||
title = "Non-death landmark AUC (absolute time)"
|
|
||||||
) +
|
|
||||||
theme(
|
|
||||||
plot.title = element_text(face = "bold", size = 10),
|
|
||||||
axis.text.x = element_text(size = 9),
|
|
||||||
legend.position = c(0.72, 0.20),
|
|
||||||
legend.background = element_blank()
|
|
||||||
)
|
|
||||||
|
|
||||||
save_panel(p1, "panel_01_landmark_overall")
|
|
||||||
|
|
||||||
# Panel 2: landmark AUC by prediction horizon.
|
|
||||||
landmark_horizon_run <- landmark_auc_disease %>%
|
|
||||||
mutate(
|
|
||||||
horizon = as.numeric(horizon),
|
|
||||||
model_target_mode = factor(model_target_mode, levels = c("next_token", "all_future"))
|
|
||||||
) %>%
|
|
||||||
group_by(run, model_target_mode, horizon) %>%
|
|
||||||
summarise(mean_auc = mean(auc, na.rm = TRUE), .groups = "drop")
|
|
||||||
|
|
||||||
landmark_horizon_summary <- landmark_horizon_run %>%
|
|
||||||
group_by(model_target_mode, horizon) %>%
|
|
||||||
summarise(mean_ci(mean_auc), .groups = "drop")
|
|
||||||
|
|
||||||
p2 <- ggplot(landmark_horizon_run, aes(x = horizon, y = mean_auc, color = model_target_mode)) +
|
|
||||||
geom_line(aes(group = run), alpha = 0.18, linewidth = 0.35) +
|
|
||||||
geom_point(alpha = 0.32, size = 1.1) +
|
|
||||||
geom_ribbon(
|
|
||||||
data = landmark_horizon_summary,
|
|
||||||
aes(x = horizon, y = mean, ymin = ymin, ymax = ymax, fill = model_target_mode, group = model_target_mode),
|
|
||||||
alpha = 0.13,
|
|
||||||
color = NA,
|
|
||||||
inherit.aes = FALSE
|
|
||||||
) +
|
|
||||||
geom_line(data = landmark_horizon_summary, aes(y = mean), linewidth = 0.85) +
|
|
||||||
geom_point(data = landmark_horizon_summary, aes(y = mean), size = 2.0) +
|
|
||||||
scale_color_manual(
|
|
||||||
values = target_cols,
|
|
||||||
labels = c("next_token", "all_future"),
|
|
||||||
name = NULL
|
|
||||||
) +
|
|
||||||
scale_fill_manual(values = target_cols, guide = "none") +
|
|
||||||
scale_x_continuous(breaks = c(1, 5, 10)) +
|
|
||||||
coord_cartesian(ylim = c(0.58, 0.78)) +
|
|
||||||
labs(
|
|
||||||
x = "Prediction horizon, years",
|
|
||||||
y = "Mean AUC per run",
|
|
||||||
title = "Non-death landmark AUC across horizons"
|
|
||||||
) +
|
|
||||||
theme(
|
|
||||||
plot.title = element_text(face = "bold", size = 10),
|
|
||||||
legend.position = c(0.31, 0.20),
|
|
||||||
legend.background = element_blank()
|
|
||||||
)
|
|
||||||
|
|
||||||
save_panel(p2, "panel_02_landmark_by_horizon", width = 3.8, height = 3.0)
|
|
||||||
|
|
||||||
# Panel 3: Delphi2M-style token AUC by offset. This documents why the old
|
|
||||||
# evaluation can make next_token look competitive, especially near the event.
|
|
||||||
token_offset_run <- token_auc_disease %>%
|
|
||||||
mutate(
|
|
||||||
offset = as.numeric(offset),
|
|
||||||
model_target_mode = factor(model_target_mode, levels = c("next_token", "all_future"))
|
|
||||||
) %>%
|
|
||||||
group_by(run, model_target_mode, offset) %>%
|
|
||||||
summarise(mean_auc = mean(auc, na.rm = TRUE), .groups = "drop")
|
|
||||||
|
|
||||||
token_offset_summary <- token_offset_run %>%
|
|
||||||
group_by(model_target_mode, offset) %>%
|
|
||||||
summarise(mean_ci(mean_auc), .groups = "drop")
|
|
||||||
|
|
||||||
p3 <- ggplot(token_offset_run, aes(x = offset, y = mean_auc, color = model_target_mode)) +
|
|
||||||
geom_line(aes(group = run), alpha = 0.18, linewidth = 0.35) +
|
|
||||||
geom_point(alpha = 0.32, size = 1.1) +
|
|
||||||
geom_ribbon(
|
|
||||||
data = token_offset_summary,
|
|
||||||
aes(x = offset, y = mean, ymin = ymin, ymax = ymax, fill = model_target_mode, group = model_target_mode),
|
|
||||||
alpha = 0.13,
|
|
||||||
color = NA,
|
|
||||||
inherit.aes = FALSE
|
|
||||||
) +
|
|
||||||
geom_line(data = token_offset_summary, aes(y = mean), linewidth = 0.85) +
|
|
||||||
geom_point(data = token_offset_summary, aes(y = mean), size = 2.0) +
|
|
||||||
scale_color_manual(
|
|
||||||
values = target_cols,
|
|
||||||
labels = c("next_token", "all_future"),
|
|
||||||
name = NULL
|
|
||||||
) +
|
|
||||||
scale_fill_manual(values = target_cols, guide = "none") +
|
|
||||||
scale_x_continuous(breaks = c(0.1, 1, 5, 10), trans = "log10") +
|
|
||||||
coord_cartesian(ylim = c(0.55, 0.82)) +
|
|
||||||
labs(
|
|
||||||
x = "Minimum offset before event, years",
|
|
||||||
y = "Mean AUC per run",
|
|
||||||
title = "Non-death token AUC by offset"
|
|
||||||
) +
|
|
||||||
theme(
|
|
||||||
plot.title = element_text(face = "bold", size = 10),
|
|
||||||
legend.position = c(0.31, 0.20),
|
|
||||||
legend.background = element_blank()
|
|
||||||
)
|
|
||||||
|
|
||||||
save_panel(p3, "panel_03_token_auc_by_offset", width = 3.8, height = 3.0)
|
|
||||||
|
|
||||||
# Panel 4: within-run contrast between old token evaluation and landmark
|
|
||||||
# evaluation. Each run contributes one point per evaluation family.
|
|
||||||
run_eval_contrast <- auc_all %>%
|
|
||||||
group_by(run, model_target_mode, dist_mode, eval_family) %>%
|
|
||||||
summarise(mean_auc = mean(auc, na.rm = TRUE), .groups = "drop")
|
|
||||||
|
|
||||||
p4 <- ggplot(run_eval_contrast, aes(x = eval_family, y = mean_auc, color = model_target_mode)) +
|
|
||||||
geom_line(aes(group = run), alpha = 0.34, linewidth = 0.45) +
|
|
||||||
geom_point(aes(shape = dist_mode), size = 2.0, alpha = 0.84) +
|
|
||||||
stat_summary(
|
|
||||||
aes(group = model_target_mode),
|
|
||||||
fun = mean,
|
|
||||||
geom = "point",
|
|
||||||
size = 3.3,
|
|
||||||
shape = 18,
|
|
||||||
position = position_dodge(width = 0.16)
|
|
||||||
) +
|
|
||||||
scale_color_manual(
|
|
||||||
values = target_cols,
|
|
||||||
labels = c("next_token", "all_future"),
|
|
||||||
name = NULL
|
|
||||||
) +
|
|
||||||
scale_shape_manual(values = dist_shapes, na.translate = FALSE, name = "Risk head") +
|
|
||||||
coord_cartesian(ylim = c(0.58, 0.78)) +
|
|
||||||
labs(
|
|
||||||
x = NULL,
|
|
||||||
y = "Mean AUC per run",
|
|
||||||
title = "Evaluation choice changes the conclusion (absolute time)"
|
|
||||||
) +
|
|
||||||
theme(
|
|
||||||
plot.title = element_text(face = "bold", size = 10),
|
|
||||||
axis.text.x = element_text(angle = 18, hjust = 1),
|
|
||||||
legend.position = "right"
|
|
||||||
)
|
|
||||||
|
|
||||||
save_panel(p4, "panel_04_evaluation_contrast", width = 4.3, height = 3.1)
|
|
||||||
|
|
||||||
# Panel 5: disease-level distribution for the landmark task, pooled over
|
|
||||||
# horizons and runs. This shows the shift without hiding heterogeneity.
|
|
||||||
landmark_density <- landmark_auc_disease %>%
|
|
||||||
mutate(model_target_mode = factor(model_target_mode, levels = c("next_token", "all_future"))) %>%
|
|
||||||
filter(is.finite(auc))
|
|
||||||
|
|
||||||
p5 <- ggplot(landmark_density, aes(x = auc, fill = model_target_mode, color = model_target_mode)) +
|
|
||||||
geom_density(alpha = 0.20, linewidth = 0.65, adjust = 1.1) +
|
|
||||||
geom_vline(
|
|
||||||
data = landmark_density %>%
|
|
||||||
group_by(model_target_mode) %>%
|
|
||||||
summarise(mean_auc = mean(auc), .groups = "drop"),
|
|
||||||
aes(xintercept = mean_auc, color = model_target_mode),
|
|
||||||
linewidth = 0.75,
|
|
||||||
linetype = "22"
|
|
||||||
) +
|
|
||||||
scale_color_manual(values = target_cols, labels = c("next_token", "all_future"), name = NULL) +
|
|
||||||
scale_fill_manual(values = target_cols, labels = c("next_token", "all_future"), name = NULL) +
|
|
||||||
coord_cartesian(xlim = c(0.35, 1.0)) +
|
|
||||||
labs(
|
|
||||||
x = "AUC",
|
|
||||||
y = "Density",
|
|
||||||
title = "Non-death landmark AUC distribution"
|
|
||||||
) +
|
|
||||||
theme(
|
|
||||||
plot.title = element_text(face = "bold", size = 10),
|
|
||||||
legend.position = c(0.24, 0.82),
|
|
||||||
legend.background = element_blank()
|
|
||||||
)
|
|
||||||
|
|
||||||
save_panel(p5, "panel_05_landmark_auc_distribution", width = 3.8, height = 3.0)
|
|
||||||
|
|
||||||
# Panel 6: death-only fixed landmark + horizon AUC. Death has one endpoint token,
|
|
||||||
# so each line is a run trajectory across horizons.
|
|
||||||
death_landmark_run <- landmark_auc_death %>%
|
|
||||||
mutate(
|
|
||||||
horizon = as.numeric(horizon),
|
|
||||||
model_target_mode = factor(model_target_mode, levels = c("next_token", "all_future")),
|
|
||||||
dist_mode = factor(dist_mode, levels = c("exponential", "weibull", "mixed"))
|
|
||||||
) %>%
|
|
||||||
group_by(run, model_target_mode, dist_mode, horizon) %>%
|
|
||||||
summarise(mean_auc = mean(auc, na.rm = TRUE), .groups = "drop")
|
|
||||||
|
|
||||||
death_landmark_summary <- death_landmark_run %>%
|
|
||||||
group_by(model_target_mode, horizon) %>%
|
|
||||||
summarise(mean_ci(mean_auc), .groups = "drop")
|
|
||||||
|
|
||||||
p6 <- ggplot(death_landmark_run, aes(x = horizon, y = mean_auc, color = model_target_mode)) +
|
|
||||||
geom_line(aes(group = run), alpha = 0.42, linewidth = 0.45) +
|
|
||||||
geom_point(aes(shape = dist_mode), alpha = 0.9, size = 2.0) +
|
|
||||||
geom_line(data = death_landmark_summary, aes(y = mean, group = model_target_mode), linewidth = 0.9) +
|
|
||||||
geom_point(data = death_landmark_summary, aes(y = mean), size = 2.2) +
|
|
||||||
scale_color_manual(values = target_cols, labels = c("next_token", "all_future"), name = NULL) +
|
|
||||||
scale_shape_manual(values = dist_shapes, na.translate = FALSE, name = "Risk head") +
|
|
||||||
scale_x_continuous(breaks = c(1, 5, 10)) +
|
|
||||||
coord_cartesian(ylim = c(0.58, 0.95)) +
|
|
||||||
labs(
|
|
||||||
x = "Prediction horizon, years",
|
|
||||||
y = "AUC",
|
|
||||||
title = "Death-only landmark AUC"
|
|
||||||
) +
|
|
||||||
theme(
|
|
||||||
plot.title = element_text(face = "bold", size = 10),
|
|
||||||
legend.position = "right"
|
|
||||||
)
|
|
||||||
|
|
||||||
save_panel(p6, "panel_06_death_landmark_by_horizon", width = 3.9, height = 3.0)
|
|
||||||
|
|
||||||
# Panel 7: death-only Delphi2M-style token AUC by offset.
|
|
||||||
death_token_run <- token_auc_death %>%
|
|
||||||
mutate(
|
|
||||||
offset = as.numeric(offset),
|
|
||||||
model_target_mode = factor(model_target_mode, levels = c("next_token", "all_future")),
|
|
||||||
dist_mode = factor(dist_mode, levels = c("exponential", "weibull", "mixed"))
|
|
||||||
) %>%
|
|
||||||
group_by(run, model_target_mode, dist_mode, offset) %>%
|
|
||||||
summarise(mean_auc = mean(auc, na.rm = TRUE), .groups = "drop")
|
|
||||||
|
|
||||||
death_token_summary <- death_token_run %>%
|
|
||||||
group_by(model_target_mode, offset) %>%
|
|
||||||
summarise(mean_ci(mean_auc), .groups = "drop")
|
|
||||||
|
|
||||||
p7 <- ggplot(death_token_run, aes(x = offset, y = mean_auc, color = model_target_mode)) +
|
|
||||||
geom_line(aes(group = run), alpha = 0.42, linewidth = 0.45) +
|
|
||||||
geom_point(aes(shape = dist_mode), alpha = 0.9, size = 2.0) +
|
|
||||||
geom_line(data = death_token_summary, aes(y = mean, group = model_target_mode), linewidth = 0.9) +
|
|
||||||
geom_point(data = death_token_summary, aes(y = mean), size = 2.2) +
|
|
||||||
scale_color_manual(values = target_cols, labels = c("next_token", "all_future"), name = NULL) +
|
|
||||||
scale_shape_manual(values = dist_shapes, na.translate = FALSE, name = "Risk head") +
|
|
||||||
scale_x_continuous(breaks = c(0.1, 1, 5, 10), trans = "log10") +
|
|
||||||
coord_cartesian(ylim = c(0.58, 0.95)) +
|
|
||||||
labs(
|
|
||||||
x = "Minimum offset before event, years",
|
|
||||||
y = "AUC",
|
|
||||||
title = "Death-only token AUC"
|
|
||||||
) +
|
|
||||||
theme(
|
|
||||||
plot.title = element_text(face = "bold", size = 10),
|
|
||||||
legend.position = "right"
|
|
||||||
)
|
|
||||||
|
|
||||||
save_panel(p7, "panel_07_death_token_auc_by_offset", width = 3.9, height = 3.0)
|
|
||||||
|
|
||||||
# Panel 8: death-only contrast between the two evaluation families.
|
|
||||||
death_eval_contrast <- bind_rows(
|
|
||||||
landmark_auc_death %>% mutate(horizon = as.numeric(horizon), offset = NA_real_),
|
|
||||||
token_auc_death %>% mutate(horizon = NA_real_, offset = as.numeric(offset))
|
|
||||||
) %>%
|
|
||||||
mutate(
|
|
||||||
model_target_mode = factor(model_target_mode, levels = c("next_token", "all_future")),
|
|
||||||
eval_family = factor(eval_family, levels = c("Delphi2M-style token", "Fixed landmark + horizon")),
|
|
||||||
dist_mode = factor(dist_mode, levels = c("exponential", "weibull", "mixed"))
|
|
||||||
) %>%
|
|
||||||
group_by(run, model_target_mode, dist_mode, eval_family) %>%
|
|
||||||
summarise(mean_auc = mean(auc, na.rm = TRUE), .groups = "drop")
|
|
||||||
|
|
||||||
p8 <- ggplot(death_eval_contrast, aes(x = eval_family, y = mean_auc, color = model_target_mode)) +
|
|
||||||
geom_line(aes(group = run), alpha = 0.38, linewidth = 0.5) +
|
|
||||||
geom_point(aes(shape = dist_mode), size = 2.2, alpha = 0.9) +
|
|
||||||
stat_summary(
|
|
||||||
aes(group = model_target_mode),
|
|
||||||
fun = mean,
|
|
||||||
geom = "point",
|
|
||||||
size = 3.4,
|
|
||||||
shape = 18,
|
|
||||||
position = position_dodge(width = 0.16)
|
|
||||||
) +
|
|
||||||
scale_color_manual(values = target_cols, labels = c("next_token", "all_future"), name = NULL) +
|
|
||||||
scale_shape_manual(values = dist_shapes, na.translate = FALSE, name = "Risk head") +
|
|
||||||
coord_cartesian(ylim = c(0.58, 0.95)) +
|
|
||||||
labs(
|
|
||||||
x = NULL,
|
|
||||||
y = "Mean AUC per run",
|
|
||||||
title = "Death endpoint evaluated separately"
|
|
||||||
) +
|
|
||||||
theme(
|
|
||||||
plot.title = element_text(face = "bold", size = 10),
|
|
||||||
axis.text.x = element_text(angle = 18, hjust = 1),
|
|
||||||
legend.position = "right"
|
|
||||||
)
|
|
||||||
|
|
||||||
save_panel(p8, "panel_08_death_evaluation_contrast", width = 4.3, height = 3.1)
|
|
||||||
|
|
||||||
# Export the exact run-level summaries used by the figures.
|
|
||||||
readr::write_csv(landmark_run, file.path(out_dir, "landmark_run_summary.csv"))
|
|
||||||
readr::write_csv(token_offset_run, file.path(out_dir, "token_offset_run_summary.csv"))
|
|
||||||
readr::write_csv(run_eval_contrast, file.path(out_dir, "run_evaluation_contrast.csv"))
|
|
||||||
readr::write_csv(death_landmark_run, file.path(out_dir, "death_landmark_run_summary.csv"))
|
|
||||||
readr::write_csv(death_token_run, file.path(out_dir, "death_token_offset_run_summary.csv"))
|
|
||||||
readr::write_csv(death_eval_contrast, file.path(out_dir, "death_evaluation_contrast.csv"))
|
|
||||||
|
|
||||||
message("Done. Panels are in: ", normalizePath(out_dir, winslash = "/"))
|
|
||||||
@@ -1,221 +0,0 @@
|
|||||||
#!/usr/bin/env bash
|
|
||||||
set -euo pipefail
|
|
||||||
|
|
||||||
# Run all non-wrapper evaluation scripts for every completed current-format
|
|
||||||
# experiment under runs/. The script is written for Linux servers with bash 4.2.
|
|
||||||
|
|
||||||
cd "$(dirname "${BASH_SOURCE[0]}")"
|
|
||||||
shopt -s globstar nullglob
|
|
||||||
|
|
||||||
PYTHON_BIN="${PYTHON_BIN:-python}"
|
|
||||||
DEVICE="${DEVICE:-cuda}"
|
|
||||||
EVAL_SPLIT="${EVAL_SPLIT:-test}"
|
|
||||||
NUM_WORKERS="${NUM_WORKERS:-4}"
|
|
||||||
CPU_REDUCE_WORKERS="${CPU_REDUCE_WORKERS:-}"
|
|
||||||
NUM_WORKERS_AUC="${NUM_WORKERS_AUC:-}"
|
|
||||||
BATCH_SIZE="${BATCH_SIZE:-}"
|
|
||||||
DATASET_SUBSET_SIZE="${DATASET_SUBSET_SIZE:-}"
|
|
||||||
DRY_RUN="${DRY_RUN:-0}"
|
|
||||||
|
|
||||||
# These attribution jobs can be expensive, but they are part of the evaluation
|
|
||||||
# surface in this repository. Set either variable to 0 to leave that family out.
|
|
||||||
RUN_EXTRA_INFO_ATTRIBUTION="${RUN_EXTRA_INFO_ATTRIBUTION:-1}"
|
|
||||||
RUN_SINGLE_DISEASE_MORTALITY_ATTRIBUTION="${RUN_SINGLE_DISEASE_MORTALITY_ATTRIBUTION:-1}"
|
|
||||||
|
|
||||||
common_args_base() {
|
|
||||||
printf '%s\n' --run_path "$1" --eval_split "${EVAL_SPLIT}" --num_workers "${NUM_WORKERS}"
|
|
||||||
if [[ -n "${BATCH_SIZE}" ]]; then
|
|
||||||
printf '%s\n' --batch_size "${BATCH_SIZE}"
|
|
||||||
fi
|
|
||||||
if [[ -n "${DATASET_SUBSET_SIZE}" ]]; then
|
|
||||||
printf '%s\n' --dataset_subset_size "${DATASET_SUBSET_SIZE}"
|
|
||||||
fi
|
|
||||||
}
|
|
||||||
|
|
||||||
common_args_with_device() {
|
|
||||||
common_args_base "$1"
|
|
||||||
printf '%s\n' --device "${DEVICE}"
|
|
||||||
}
|
|
||||||
|
|
||||||
auc_args() {
|
|
||||||
if [[ -n "${NUM_WORKERS_AUC}" ]]; then
|
|
||||||
printf '%s\n' --num_workers_auc "${NUM_WORKERS_AUC}"
|
|
||||||
fi
|
|
||||||
}
|
|
||||||
|
|
||||||
cpu_reduce_args() {
|
|
||||||
if [[ -n "${CPU_REDUCE_WORKERS}" ]]; then
|
|
||||||
printf '%s\n' --cpu_reduce_workers "${CPU_REDUCE_WORKERS}"
|
|
||||||
fi
|
|
||||||
}
|
|
||||||
|
|
||||||
has_completed_dir() {
|
|
||||||
local dir="$1"
|
|
||||||
shift
|
|
||||||
[[ -d "${dir}" ]] || return 1
|
|
||||||
local required
|
|
||||||
for required in "$@"; do
|
|
||||||
[[ -s "${dir}/${required}" ]] || return 1
|
|
||||||
done
|
|
||||||
}
|
|
||||||
|
|
||||||
run_command() {
|
|
||||||
echo " run: $*"
|
|
||||||
if [[ "${DRY_RUN}" == "1" ]]; then
|
|
||||||
return 0
|
|
||||||
fi
|
|
||||||
"$@"
|
|
||||||
}
|
|
||||||
|
|
||||||
run_dir_result_if_missing() {
|
|
||||||
local label="$1"
|
|
||||||
local result_dir="$2"
|
|
||||||
local required_1="$3"
|
|
||||||
local required_2="$4"
|
|
||||||
shift 4
|
|
||||||
|
|
||||||
if has_completed_dir "${result_dir}" "${required_1}" "${required_2}"; then
|
|
||||||
echo " skip ${label}: found ${result_dir}"
|
|
||||||
return 0
|
|
||||||
fi
|
|
||||||
|
|
||||||
run_command "$@"
|
|
||||||
}
|
|
||||||
|
|
||||||
run_file_result_if_missing() {
|
|
||||||
local label="$1"
|
|
||||||
local result_dir="$2"
|
|
||||||
local required="$3"
|
|
||||||
shift 3
|
|
||||||
|
|
||||||
if [[ -s "${result_dir}/${required}" ]]; then
|
|
||||||
echo " skip ${label}: found ${result_dir}/${required}"
|
|
||||||
return 0
|
|
||||||
fi
|
|
||||||
|
|
||||||
run_command "$@"
|
|
||||||
}
|
|
||||||
|
|
||||||
run_has_extra_info() {
|
|
||||||
"${PYTHON_BIN}" - "$1" <<'PY'
|
|
||||||
import json
|
|
||||||
import sys
|
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
cfg_path = Path(sys.argv[1]) / "train_config.json"
|
|
||||||
try:
|
|
||||||
cfg = json.loads(cfg_path.read_text(encoding="utf-8"))
|
|
||||||
except Exception:
|
|
||||||
raise SystemExit(1)
|
|
||||||
|
|
||||||
extra = cfg.get("extra_info_types", [])
|
|
||||||
raise SystemExit(0 if isinstance(extra, list) and len(extra) > 0 else 1)
|
|
||||||
PY
|
|
||||||
}
|
|
||||||
|
|
||||||
run_is_all_future() {
|
|
||||||
"${PYTHON_BIN}" - "$1" <<'PY'
|
|
||||||
import json
|
|
||||||
import sys
|
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
cfg_path = Path(sys.argv[1]) / "train_config.json"
|
|
||||||
try:
|
|
||||||
cfg = json.loads(cfg_path.read_text(encoding="utf-8"))
|
|
||||||
except Exception:
|
|
||||||
raise SystemExit(1)
|
|
||||||
|
|
||||||
mode = str(cfg.get("model_target_mode", "next_token")).lower()
|
|
||||||
raise SystemExit(0 if mode == "all_future" else 1)
|
|
||||||
PY
|
|
||||||
}
|
|
||||||
|
|
||||||
run_has_current_model_config() {
|
|
||||||
"${PYTHON_BIN}" - "$1" <<'PY'
|
|
||||||
import json
|
|
||||||
import sys
|
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
cfg_path = Path(sys.argv[1]) / "train_config.json"
|
|
||||||
try:
|
|
||||||
cfg = json.loads(cfg_path.read_text(encoding="utf-8"))
|
|
||||||
n_layer = int(cfg.get("n_layer", 0))
|
|
||||||
except Exception:
|
|
||||||
raise SystemExit(1)
|
|
||||||
|
|
||||||
supported = {"transformer_ffn_v1", "traj_mixer_v5"}
|
|
||||||
raise SystemExit(
|
|
||||||
0
|
|
||||||
if cfg.get("model_architecture") in supported and n_layer >= 1
|
|
||||||
else 1
|
|
||||||
)
|
|
||||||
PY
|
|
||||||
}
|
|
||||||
|
|
||||||
for config_path in runs/**/train_config.json; do
|
|
||||||
run_path="${config_path%/train_config.json}"
|
|
||||||
|
|
||||||
echo "==> ${run_path}"
|
|
||||||
if [[ ! -f "${run_path}/train_config.json" ]]; then
|
|
||||||
echo " skip run: missing train_config.json"
|
|
||||||
continue
|
|
||||||
fi
|
|
||||||
if [[ ! -s "${run_path}/best_model.pt" ]]; then
|
|
||||||
echo " skip run: missing best_model.pt"
|
|
||||||
continue
|
|
||||||
fi
|
|
||||||
if ! run_has_current_model_config "${run_path}"; then
|
|
||||||
echo " skip run: config lacks current model_architecture/n_layer fields"
|
|
||||||
continue
|
|
||||||
fi
|
|
||||||
|
|
||||||
common=()
|
|
||||||
while IFS= read -r arg; do common+=("${arg}"); done < <(common_args_with_device "${run_path}")
|
|
||||||
|
|
||||||
auc_extra=()
|
|
||||||
while IFS= read -r arg; do auc_extra+=("${arg}"); done < <(auc_args)
|
|
||||||
|
|
||||||
cpu_reduce_extra=()
|
|
||||||
while IFS= read -r arg; do cpu_reduce_extra+=("${arg}"); done < <(cpu_reduce_args)
|
|
||||||
|
|
||||||
run_file_result_if_missing \
|
|
||||||
"evaluate_auc.py" \
|
|
||||||
"${run_path}" \
|
|
||||||
"df_auc_delphi2m_report.csv" \
|
|
||||||
"${PYTHON_BIN}" evaluate_auc.py "${common[@]}" "${auc_extra[@]}"
|
|
||||||
|
|
||||||
run_file_result_if_missing \
|
|
||||||
"evaluate_auc_v2.py" \
|
|
||||||
"${run_path}" \
|
|
||||||
"df_auc_landmark_delphi2m_report.csv" \
|
|
||||||
"${PYTHON_BIN}" evaluate_auc_v2.py "${common[@]}" "${auc_extra[@]}"
|
|
||||||
|
|
||||||
if ! run_is_all_future "${run_path}"; then
|
|
||||||
echo " skip attribution evaluations: model_target_mode is not all_future"
|
|
||||||
continue
|
|
||||||
fi
|
|
||||||
|
|
||||||
if [[ "${RUN_EXTRA_INFO_ATTRIBUTION}" == "1" ]]; then
|
|
||||||
if run_has_extra_info "${run_path}"; then
|
|
||||||
run_dir_result_if_missing \
|
|
||||||
"evaluate_extra_info_attribution.py" \
|
|
||||||
"${run_path}/extra_info_attribution_${EVAL_SPLIT}" \
|
|
||||||
"manifest.json" \
|
|
||||||
"summary_extra_info_disease_parameters.csv" \
|
|
||||||
"${PYTHON_BIN}" evaluate_extra_info_attribution.py "${common[@]}" "${cpu_reduce_extra[@]}"
|
|
||||||
else
|
|
||||||
echo " skip evaluate_extra_info_attribution.py: run has no extra-info types"
|
|
||||||
fi
|
|
||||||
fi
|
|
||||||
|
|
||||||
if [[ "${RUN_SINGLE_DISEASE_MORTALITY_ATTRIBUTION}" == "1" ]]; then
|
|
||||||
run_dir_result_if_missing \
|
|
||||||
"evaluate_single_disease_mortality_attribution.py" \
|
|
||||||
"${run_path}/single_disease_mortality_parameters_${EVAL_SPLIT}_all_diseases" \
|
|
||||||
"manifest.json" \
|
|
||||||
"summary_by_disease_age_sex.csv" \
|
|
||||||
"${PYTHON_BIN}" evaluate_single_disease_mortality_attribution.py "${common[@]}"
|
|
||||||
fi
|
|
||||||
done
|
|
||||||
|
|
||||||
echo "All missing evaluations are complete."
|
|
||||||
@@ -1,143 +0,0 @@
|
|||||||
#!/usr/bin/env bash
|
|
||||||
set -euo pipefail
|
|
||||||
|
|
||||||
# Linux bash 5.2+ training-only script.
|
|
||||||
#
|
|
||||||
# Based on the existing runs, the objective/time/death-distribution checks are
|
|
||||||
# already covered. The remaining gap for the current proof chain is the
|
|
||||||
# extra-info ablation under the final candidate model:
|
|
||||||
#
|
|
||||||
# all_future + relative time + mixed death/risk head
|
|
||||||
#
|
|
||||||
# This script only launches those missing training jobs. It intentionally does
|
|
||||||
# not call evaluate_*.py and does not add extra random seeds. Set
|
|
||||||
# MODEL_ARCHITECTURE=traj_mixer_v5 to run the TrajMixer variant.
|
|
||||||
|
|
||||||
cd "$(dirname "${BASH_SOURCE[0]}")"
|
|
||||||
|
|
||||||
PYTHON_BIN="${PYTHON_BIN:-python}"
|
|
||||||
DEVICE="${DEVICE:-cuda}"
|
|
||||||
NUM_WORKERS="${NUM_WORKERS:-4}"
|
|
||||||
PROGRESS_INTERVAL="${PROGRESS_INTERVAL:-20}"
|
|
||||||
MODEL_ARCHITECTURE="${MODEL_ARCHITECTURE:-transformer_ffn_v1}"
|
|
||||||
N_LAYER="${N_LAYER:-12}"
|
|
||||||
|
|
||||||
TIME_MODE="relative"
|
|
||||||
DIST_MODE="mixed"
|
|
||||||
SEED="42"
|
|
||||||
VALIDATION_QUERY_SEED="42"
|
|
||||||
|
|
||||||
COMMON_ARGS=(
|
|
||||||
--data_prefix ukb
|
|
||||||
--labels_file labels.csv
|
|
||||||
--seed "${SEED}"
|
|
||||||
--validation_query_seed "${VALIDATION_QUERY_SEED}"
|
|
||||||
--train_eid_file ukb_train_eid.csv
|
|
||||||
--val_eid_file ukb_val_eid.csv
|
|
||||||
--test_eid_file ukb_test_eid.csv
|
|
||||||
--min_history_events 1
|
|
||||||
--min_future_events 1
|
|
||||||
--n_embd 120
|
|
||||||
--n_head 10
|
|
||||||
--n_layer "${N_LAYER}"
|
|
||||||
--model_architecture "${MODEL_ARCHITECTURE}"
|
|
||||||
--n_bins 16
|
|
||||||
--extra_pool_reduce mean
|
|
||||||
--dropout 0.0
|
|
||||||
--batch_size 256
|
|
||||||
--base_lr 0.0003
|
|
||||||
--weight_decay 0.1
|
|
||||||
--betas 0.9 0.99
|
|
||||||
--grad_clip 1.0
|
|
||||||
--max_epochs 200
|
|
||||||
--warmup_epochs 10
|
|
||||||
--patience 15
|
|
||||||
--min_lr_ratio 0.1
|
|
||||||
--num_workers "${NUM_WORKERS}"
|
|
||||||
--device "${DEVICE}"
|
|
||||||
--progress_interval "${PROGRESS_INTERVAL}"
|
|
||||||
)
|
|
||||||
|
|
||||||
already_trained() {
|
|
||||||
local extra_file="$1"
|
|
||||||
"${PYTHON_BIN}" - "$TIME_MODE" "$DIST_MODE" "$extra_file" "$SEED" "$VALIDATION_QUERY_SEED" "$MODEL_ARCHITECTURE" "$N_LAYER" <<'PY'
|
|
||||||
import json
|
|
||||||
import sys
|
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
(
|
|
||||||
time_mode,
|
|
||||||
dist_mode,
|
|
||||||
extra_file,
|
|
||||||
seed,
|
|
||||||
validation_query_seed,
|
|
||||||
model_architecture,
|
|
||||||
n_layer,
|
|
||||||
) = sys.argv[1:8]
|
|
||||||
extra_name = Path(extra_file).name
|
|
||||||
|
|
||||||
for config_path in Path("runs").rglob("train_config.json"):
|
|
||||||
try:
|
|
||||||
cfg = json.loads(config_path.read_text(encoding="utf-8"))
|
|
||||||
except Exception:
|
|
||||||
continue
|
|
||||||
|
|
||||||
observed_query_seed = cfg.get(
|
|
||||||
"all_future_validation_query_seed",
|
|
||||||
cfg.get("validation_query_seed", -1),
|
|
||||||
)
|
|
||||||
|
|
||||||
if (
|
|
||||||
cfg.get("model_target_mode") == "all_future"
|
|
||||||
and cfg.get("model_architecture") == model_architecture
|
|
||||||
and int(cfg.get("n_layer", -1)) == int(n_layer)
|
|
||||||
and cfg.get("time_mode") == time_mode
|
|
||||||
and cfg.get("dist_mode") == dist_mode
|
|
||||||
and Path(str(cfg.get("extra_info_types_file", ""))).name == extra_name
|
|
||||||
and int(cfg.get("seed", -1)) == int(seed)
|
|
||||||
and int(observed_query_seed) == int(validation_query_seed)
|
|
||||||
):
|
|
||||||
print(config_path.parent)
|
|
||||||
raise SystemExit(0)
|
|
||||||
|
|
||||||
raise SystemExit(1)
|
|
||||||
PY
|
|
||||||
}
|
|
||||||
|
|
||||||
train_if_missing() {
|
|
||||||
local label="$1"
|
|
||||||
local extra_file="$2"
|
|
||||||
|
|
||||||
if [[ ! -f "${extra_file}" ]]; then
|
|
||||||
echo "ERROR: missing extra-info type file: ${extra_file}" >&2
|
|
||||||
return 2
|
|
||||||
fi
|
|
||||||
|
|
||||||
echo "==> Checking ${label}: ${MODEL_ARCHITECTURE} n_layer=${N_LAYER} ${TIME_MODE} ${DIST_MODE} all_future with ${extra_file}"
|
|
||||||
if existing_run="$(already_trained "$extra_file")"; then
|
|
||||||
echo " skip: already trained at ${existing_run}"
|
|
||||||
return 0
|
|
||||||
fi
|
|
||||||
|
|
||||||
echo " train: ${label}"
|
|
||||||
"${PYTHON_BIN}" train_all_future.py \
|
|
||||||
"${COMMON_ARGS[@]}" \
|
|
||||||
--time_mode "${TIME_MODE}" \
|
|
||||||
--dist_mode "${DIST_MODE}" \
|
|
||||||
--extra_info_types_file "${extra_file}"
|
|
||||||
}
|
|
||||||
|
|
||||||
# Already present in runs/:
|
|
||||||
# - next-token objective checks under SAB, plus older absolute extra ablations
|
|
||||||
# - all-future absolute/relative x exponential/weibull/mixed under SAB
|
|
||||||
#
|
|
||||||
# Still needed:
|
|
||||||
# - final all-future relative+mixed extra-info ablations beyond the existing
|
|
||||||
# SAB baseline. These close the disease-only question without expanding seed
|
|
||||||
# count or running downstream evaluation.
|
|
||||||
train_if_missing "true_disease_only" "extra_info_types_none.txt"
|
|
||||||
train_if_missing "assessment_only_extra" "extra_info_types_assessment_only.txt"
|
|
||||||
train_if_missing "exposure_only_extra" "extra_info_types_exposure_only.txt"
|
|
||||||
train_if_missing "all_extra_info" "extra_info_types_all.txt"
|
|
||||||
|
|
||||||
echo "All requested training-only missing configurations are done."
|
|
||||||
@@ -1,76 +0,0 @@
|
|||||||
#!/usr/bin/env bash
|
|
||||||
|
|
||||||
# Export Weibull shape-parameter summaries for the all_future models trained
|
|
||||||
# with smoking/alcohol/BMI extra information.
|
|
||||||
#
|
|
||||||
# Bash 4.2 compatible. Run from the DeepHealth repository root on the Linux
|
|
||||||
# server, for example:
|
|
||||||
#
|
|
||||||
# bash run_weibull_shape_exports.sh
|
|
||||||
#
|
|
||||||
# Optional overrides:
|
|
||||||
# PYTHON=python3 DEVICE=cuda BATCH_SIZE=128 NUM_WORKERS=0 ROW_BATCH_SIZE=512 \
|
|
||||||
# bash run_weibull_shape_exports.sh
|
|
||||||
|
|
||||||
set -euo pipefail
|
|
||||||
|
|
||||||
PYTHON="${PYTHON:-python}"
|
|
||||||
DEVICE="${DEVICE:-cuda}"
|
|
||||||
BATCH_SIZE="${BATCH_SIZE:-128}"
|
|
||||||
NUM_WORKERS="${NUM_WORKERS:-0}"
|
|
||||||
ROW_BATCH_SIZE="${ROW_BATCH_SIZE:-512}"
|
|
||||||
LANDMARK_START="${LANDMARK_START:-40}"
|
|
||||||
LANDMARK_STOP="${LANDMARK_STOP:-80}"
|
|
||||||
LANDMARK_STEP="${LANDMARK_STEP:-5}"
|
|
||||||
HORIZONS="${HORIZONS:-1,5,10}"
|
|
||||||
|
|
||||||
RUNS=(
|
|
||||||
"runs/relative_weibull_all_future_pure_disease_20260620_095229"
|
|
||||||
"runs/relative_mixed_all_future_pure_disease_20260620_132415"
|
|
||||||
"runs/absolute_weibull_all_future_pure_disease_20260620_114816"
|
|
||||||
"runs/absolute_mixed_all_future_pure_disease_20260620_161804"
|
|
||||||
)
|
|
||||||
|
|
||||||
echo "Python: ${PYTHON}"
|
|
||||||
echo "Device: ${DEVICE}"
|
|
||||||
echo "Batch size: ${BATCH_SIZE}"
|
|
||||||
echo "Workers: ${NUM_WORKERS}"
|
|
||||||
echo "Row batch size: ${ROW_BATCH_SIZE}"
|
|
||||||
echo "Horizons: ${HORIZONS}"
|
|
||||||
echo
|
|
||||||
|
|
||||||
for run_path in "${RUNS[@]}"; do
|
|
||||||
if [[ ! -d "${run_path}" ]]; then
|
|
||||||
echo "[ERROR] Missing run directory: ${run_path}" >&2
|
|
||||||
exit 1
|
|
||||||
fi
|
|
||||||
if [[ ! -f "${run_path}/best_model.pt" ]]; then
|
|
||||||
echo "[ERROR] Missing checkpoint: ${run_path}/best_model.pt" >&2
|
|
||||||
exit 1
|
|
||||||
fi
|
|
||||||
if [[ ! -f "${run_path}/train_config.json" ]]; then
|
|
||||||
echo "[ERROR] Missing config: ${run_path}/train_config.json" >&2
|
|
||||||
exit 1
|
|
||||||
fi
|
|
||||||
|
|
||||||
output_path="${run_path}/weibull_shape_parameter_stats_test"
|
|
||||||
echo "=== Exporting Weibull shape stats: ${run_path} ==="
|
|
||||||
"${PYTHON}" export_weibull_shape_parameter_stats.py \
|
|
||||||
--run_path "${run_path}" \
|
|
||||||
--output_path "${output_path}" \
|
|
||||||
--eval_split test \
|
|
||||||
--device "${DEVICE}" \
|
|
||||||
--batch_size "${BATCH_SIZE}" \
|
|
||||||
--num_workers "${NUM_WORKERS}" \
|
|
||||||
--row_batch_size "${ROW_BATCH_SIZE}" \
|
|
||||||
--hidden_cache_dtype float32 \
|
|
||||||
--landmark_start "${LANDMARK_START}" \
|
|
||||||
--landmark_stop "${LANDMARK_STOP}" \
|
|
||||||
--landmark_step "${LANDMARK_STEP}" \
|
|
||||||
--horizons "${HORIZONS}" \
|
|
||||||
--include_all_token_rho_summary
|
|
||||||
echo "Wrote: ${output_path}"
|
|
||||||
echo
|
|
||||||
done
|
|
||||||
|
|
||||||
echo "All Weibull shape exports completed."
|
|
||||||
Reference in New Issue
Block a user