Files
DeepHealth/export_weibull_parameters.py

662 lines
22 KiB
Python
Raw Normal View History

"""Export individual Weibull parameters for every disease/death token.
The test population is read from ``test_eid_file`` in the specified run's
``train_config.json``. Each eligible patient is queried at ages
40, 42, ..., 80 using only disease history observed by that age.
The model parameterization is
S(t) = exp(-rate * t ** shape)
and the standard Weibull parameters exported by this script are
shape = rho
scale = rate ** (-1 / shape)
Output is sharded compressed NPZ rather than one monolithic matrix. Every NPZ
contains aligned ``shape`` and ``scale`` matrices with rows for patients and
columns for tokens. ``tokens.csv`` defines the column order and
``manifest.csv`` defines the shard/row order for each age.
"""
from __future__ import annotations
import argparse
import contextlib
import csv
import json
import math
from pathlib import Path
from typing import Any, Dict, Iterable, List, Sequence
import numpy as np
import torch
import torch.nn.functional as F
from torch.utils.data import DataLoader
from tqdm.auto import tqdm
from dataset import (
DISEASE_HISTORY_MODE_TIMED,
NO_EVENT_IDX,
PAD_IDX,
RESERVED_IDX,
normalize_disease_history_mode,
)
from eval_data import (
build_model_from_dataset,
load_json_config,
load_sequence_eval_dataset,
resolve_eval_device,
select_indices_by_eid_file,
validate_dataset_metadata,
validate_training_mode_config,
)
from evaluate_auc_v2 import (
LandmarkDataset,
collate_landmark_fn,
load_checkpoint_state_dict,
load_model_state,
resolve_dist_mode_for_checkpoint,
)
from model_architectures import resolve_model_architecture
SPECIAL_TOKENS = {PAD_IDX, RESERVED_IDX, NO_EVENT_IDX}
MANIFEST_FIELDS = [
"age",
"shard",
"file",
"row_start",
"row_stop",
"n_rows",
"n_tokens",
"nonfinite_shape_values",
"nonfinite_scale_values",
]
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description=(
"Export each test patient's complete disease/death Weibull "
"shape and scale matrices at two-year age landmarks."
)
)
parser.add_argument(
"--run_path",
required=True,
help="Run directory containing train_config.json and best_model.pt.",
)
parser.add_argument(
"--output_path",
default=None,
help=(
"Output directory. Defaults to "
"<run_path>/weibull_parameters_test_age40_80_step2."
),
)
parser.add_argument(
"--test_eid_file",
default=None,
help=(
"Optional override. By default use test_eid_file from the run's "
"train_config.json, or ukb_test_eid.csv if the field is absent."
),
)
parser.add_argument("--age_start", type=float, default=40.0)
parser.add_argument("--age_stop", type=float, default=80.0)
parser.add_argument("--age_step", type=float, default=2.0)
parser.add_argument("--batch_size", type=int, default=128)
parser.add_argument(
"--rows_per_shard",
type=int,
default=4096,
help=(
"Approximate patient rows per compressed NPZ shard. This is "
"independent of inference batch size."
),
)
parser.add_argument("--num_workers", type=int, default=4)
parser.add_argument(
"--device",
default=None,
help="For example cpu, cuda, or cuda:1. Defaults to CUDA when available.",
)
parser.add_argument(
"--use_amp",
action=argparse.BooleanOptionalAction,
default=False,
help="Use float16 autocast during CUDA inference.",
)
parser.add_argument(
"--dataset_subset_size",
type=int,
default=None,
help="Use only the first N matched test patients for a smoke test.",
)
return parser.parse_args()
def build_age_grid(start: float, stop: float, step: float) -> np.ndarray:
"""Build an inclusive age grid and reject a stop not aligned to the step."""
if not all(math.isfinite(x) for x in (start, stop, step)):
raise ValueError("Age start/stop/step must be finite.")
if step <= 0:
raise ValueError("age_step must be > 0.")
if stop < start:
raise ValueError("age_stop must be >= age_start.")
count = int(math.floor((stop - start) / step + 1e-9)) + 1
ages = start + step * np.arange(count, dtype=np.float64)
if ages.size == 0 or not math.isclose(
float(ages[-1]), stop, rel_tol=0.0, abs_tol=1e-7
):
raise ValueError(
"age_stop must lie on the grid defined by age_start and age_step."
)
ages[-1] = stop
return ages.astype(np.float32)
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 not text:
return None
if text.startswith("["):
parsed = json.loads(text)
if not isinstance(parsed, list):
raise ValueError("extra_info_types must be a list of integers.")
return [int(x) for x in parsed]
return [int(x.strip()) for x in text.split(",") if x.strip()]
def select_outcome_tokens(dataset: Any) -> List[int]:
"""Select every real outcome token, including Death, in token-id order."""
tokens = sorted(
int(token)
for token, code in dataset.label_id_to_code.items()
if int(token) not in SPECIAL_TOKENS and not str(code).startswith("<")
)
if not tokens:
raise RuntimeError("The dataset contains no disease/death outcome tokens.")
return tokens
def resolve_project_file(path_value: str | Path) -> Path:
path = Path(path_value)
if path.is_absolute():
return path
direct = Path.cwd() / path
return direct if direct.is_file() else Path(__file__).resolve().parent / path
def load_label_text(labels_file: str | Path) -> Dict[str, str]:
result: Dict[str, str] = {}
with resolve_project_file(labels_file).open("r", encoding="utf-8") as handle:
for line in handle:
text = line.strip()
if text:
result[text.split()[0]] = text
return result
def write_csv(
path: Path,
fieldnames: Sequence[str],
rows: Iterable[Dict[str, Any]],
) -> None:
temporary = path.with_name(f".{path.name}.tmp")
with temporary.open("w", newline="", encoding="utf-8-sig") as handle:
writer = csv.DictWriter(handle, fieldnames=fieldnames)
writer.writeheader()
writer.writerows(rows)
temporary.replace(path)
def write_json(path: Path, payload: Dict[str, Any]) -> None:
temporary = path.with_name(f".{path.name}.tmp")
temporary.write_text(
json.dumps(payload, ensure_ascii=False, indent=2),
encoding="utf-8",
)
temporary.replace(path)
def age_directory_name(age: float) -> str:
text = f"{age:g}".replace("-", "minus_").replace(".", "p")
return f"age_{text}"
def write_empty_age_export(
*,
age: float,
age_dir: Path,
token_count: int,
) -> Dict[str, Any]:
"""Represent an age with no eligible test patients by one empty shard."""
shard_path = age_dir / "shard_000000.npz"
np.savez_compressed(
shard_path,
eid=np.empty(0, dtype=np.int64),
dataset_index=np.empty(0, dtype=np.int64),
sex=np.empty(0, dtype=np.int8),
age=np.empty(0, dtype=np.float32),
shape=np.empty((0, token_count), dtype=np.float32),
scale=np.empty((0, token_count), dtype=np.float32),
)
return {
"age": float(age),
"shard": 0,
"file": str(shard_path.relative_to(age_dir.parent.parent)),
"row_start": 0,
"row_stop": 0,
"n_rows": 0,
"n_tokens": int(token_count),
"nonfinite_shape_values": 0,
"nonfinite_scale_values": 0,
}
@torch.inference_mode()
def export_age(
*,
age: float,
model: Any,
loader: DataLoader,
tokens: Sequence[int],
device: torch.device,
use_amp: bool,
dataset: Any,
subset_indices: np.ndarray,
age_dir: Path,
output_path: Path,
rows_per_shard: int,
) -> List[Dict[str, Any]]:
"""Export all patient-by-token matrices for one age in row-order shards."""
token_index = torch.as_tensor(tokens, dtype=torch.long, device=device)
patient_eids = np.asarray(
[int(dataset.samples[int(index)]["eid"]) for index in subset_indices],
dtype=np.int64,
)
amp_enabled = bool(use_amp and device.type == "cuda")
manifest_rows: List[Dict[str, Any]] = []
row_start = 0
shape_parts: List[np.ndarray] = []
scale_parts: List[np.ndarray] = []
eid_parts: List[np.ndarray] = []
dataset_index_parts: List[np.ndarray] = []
sex_parts: List[np.ndarray] = []
age_parts: List[np.ndarray] = []
buffered_rows = 0
def flush_shard() -> None:
nonlocal row_start, buffered_rows
if buffered_rows == 0:
return
shape_matrix = np.concatenate(shape_parts, axis=0)
scale_matrix = np.concatenate(scale_parts, axis=0)
eid = np.concatenate(eid_parts)
dataset_index = np.concatenate(dataset_index_parts)
sex = np.concatenate(sex_parts)
age_values = np.concatenate(age_parts)
shard_index = len(manifest_rows)
n_rows = int(shape_matrix.shape[0])
row_stop = row_start + n_rows
shard_path = age_dir / f"shard_{shard_index:06d}.npz"
np.savez_compressed(
shard_path,
eid=eid,
dataset_index=dataset_index,
sex=sex,
age=age_values,
shape=shape_matrix,
scale=scale_matrix,
)
manifest_rows.append(
{
"age": float(age),
"shard": int(shard_index),
"file": str(shard_path.relative_to(output_path)),
"row_start": int(row_start),
"row_stop": int(row_stop),
"n_rows": n_rows,
"n_tokens": int(shape_matrix.shape[1]),
"nonfinite_shape_values": int(
(~np.isfinite(shape_matrix)).sum()
),
"nonfinite_scale_values": int(
(~np.isfinite(scale_matrix)).sum()
),
}
)
row_start = row_stop
buffered_rows = 0
shape_parts.clear()
scale_parts.clear()
eid_parts.clear()
dataset_index_parts.clear()
sex_parts.clear()
age_parts.clear()
for batch in tqdm(loader, desc=f"Age {age:g}", dynamic_ncols=True):
batch_device = {
key: (
value.to(device, non_blocking=True)
if isinstance(value, torch.Tensor)
else value
)
for key, value in batch.items()
}
amp_context = (
torch.autocast(device_type="cuda", dtype=torch.float16)
if amp_enabled
else contextlib.nullcontext()
)
with amp_context:
hidden = model(
event_seq=batch_device["event_seq"],
time_seq=batch_device["time_seq"],
sex=batch_device["sex"],
padding_mask=batch_device["padding_mask"],
t_query=batch_device["t_query"],
other_type=batch_device["other_type"],
other_value=batch_device["other_value"],
other_value_kind=batch_device["other_value_kind"],
other_time=batch_device["other_time"],
)
logits = model.calc_risk(hidden).index_select(1, token_index).float()
shape = (
model.calc_weibull_rho(hidden)
.index_select(1, token_index)
.float()
)
rate = F.softplus(logits) + 1e-8
scale = torch.exp(-torch.log(rate) / shape)
shape_np = shape.cpu().numpy().astype(np.float32, copy=False)
scale_np = scale.cpu().numpy().astype(np.float32, copy=False)
patient_id = batch["patient_id"].cpu().numpy().astype(np.int64)
dataset_index = subset_indices[patient_id]
n_rows = int(shape_np.shape[0])
shape_parts.append(shape_np)
scale_parts.append(scale_np)
eid_parts.append(patient_eids[patient_id])
dataset_index_parts.append(
dataset_index.astype(np.int64, copy=False)
)
sex_parts.append(
batch["sex"].cpu().numpy().astype(np.int8, copy=False)
)
age_parts.append(
batch["landmark_age"]
.cpu()
.numpy()
.astype(np.float32, copy=False)
)
buffered_rows += n_rows
if buffered_rows >= rows_per_shard:
flush_shard()
flush_shard()
return manifest_rows
def main() -> None:
args = parse_args()
run_path = Path(args.run_path).resolve()
config_path = run_path / "train_config.json"
checkpoint_path = run_path / "best_model.pt"
if not config_path.is_file():
raise FileNotFoundError(config_path)
if not checkpoint_path.is_file():
raise FileNotFoundError(checkpoint_path)
if args.batch_size <= 0:
raise ValueError("batch_size must be > 0.")
if args.rows_per_shard <= 0:
raise ValueError("rows_per_shard must be > 0.")
if args.num_workers < 0:
raise ValueError("num_workers must be >= 0.")
if args.dataset_subset_size is not None and args.dataset_subset_size <= 0:
raise ValueError("dataset_subset_size must be > 0.")
cfg = load_json_config(config_path)
validate_training_mode_config(cfg)
model_target_mode = str(cfg.get("model_target_mode", "next_token")).lower()
if model_target_mode != "all_future":
raise ValueError(
"This exporter requires model_target_mode='all_future'; got "
f"{model_target_mode!r}."
)
ages = build_age_grid(args.age_start, args.age_stop, args.age_step)
output_path = (
Path(args.output_path).resolve()
if args.output_path
else run_path / "weibull_parameters_test_age40_80_step2"
)
if output_path.exists() and any(output_path.iterdir()):
raise FileExistsError(
f"Output directory is not empty: {output_path}. Choose a new "
"--output_path so shards from different exports cannot be mixed."
)
output_path.mkdir(parents=True, exist_ok=True)
shards_root = output_path / "shards"
shards_root.mkdir(parents=True, exist_ok=True)
data_prefix = str(cfg.get("data_prefix", "ukb"))
labels_file = str(cfg.get("labels_file", "labels.csv"))
disease_history_mode = normalize_disease_history_mode(
cfg.get("disease_history_mode", DISEASE_HISTORY_MODE_TIMED)
)
min_history_events = int(
cfg.get("all_future_min_history_events", cfg.get("min_history_events", 1))
)
min_future_events = int(
cfg.get("all_future_min_future_events", cfg.get("min_future_events", 1))
)
print("Loading dataset...")
dataset = load_sequence_eval_dataset(
model_target_mode=model_target_mode,
data_prefix=data_prefix,
labels_file=labels_file,
no_event_interval_years=float(cfg.get("no_event_interval_years", 5.0)),
min_history_events=min_history_events,
min_future_events=min_future_events,
extra_info_types=parse_int_list(cfg.get("extra_info_types")),
disease_history_mode=disease_history_mode,
)
validate_dataset_metadata(dataset, cfg)
test_eid_file = args.test_eid_file or cfg.get(
"test_eid_file", "ukb_test_eid.csv"
)
if not test_eid_file:
raise ValueError(
"No test_eid_file is defined. Provide --test_eid_file explicitly."
)
subset_indices, resolved_eid_file = select_indices_by_eid_file(
dataset, str(test_eid_file)
)
if args.dataset_subset_size is not None:
subset_indices = subset_indices[: args.dataset_subset_size]
if subset_indices.size == 0:
raise RuntimeError("The selected test subset is empty.")
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
)
if dist_mode != "weibull":
raise ValueError(
f"The specified run uses dist_mode={dist_mode!r}, not 'weibull'."
)
cfg_model = dict(cfg)
cfg_model["dist_mode"] = dist_mode
cfg_model["model_architecture"] = resolve_model_architecture(
cfg_model, state_dict
)
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()
tokens = select_outcome_tokens(dataset)
label_text = load_label_text(labels_file)
death_tokens = [
token
for token in tokens
if str(dataset.label_id_to_code[token]).lower() == "death"
]
if not death_tokens:
raise RuntimeError("Death token was not found in the outcome vocabulary.")
token_rows = [
{
"column": column,
"token_id": token,
"label_code": str(dataset.label_id_to_code[token]),
"label_text": label_text.get(
str(dataset.label_id_to_code[token]),
str(dataset.label_id_to_code[token]),
),
"outcome_type": (
"death"
if str(dataset.label_id_to_code[token]).lower() == "death"
else "disease"
),
}
for column, token in enumerate(tokens)
]
write_csv(
output_path / "tokens.csv",
["column", "token_id", "label_code", "label_text", "outcome_type"],
token_rows,
)
metadata: Dict[str, Any] = {
"format_version": 1,
"complete": False,
"run_path": str(run_path),
"checkpoint": str(checkpoint_path),
"test_eid_file": str(resolved_eid_file),
"n_selected_test_patients": int(subset_indices.size),
"ages": [float(x) for x in ages],
"n_tokens": len(tokens),
"token_columns_file": "tokens.csv",
"manifest_file": "manifest.csv",
"matrix_dtype": "float32",
"npz_arrays": {
"eid": "int64 [n_rows]",
"dataset_index": "int64 [n_rows]",
"sex": "int8 [n_rows], 0=female and 1=male",
"age": "float32 [n_rows]",
"shape": "float32 [n_rows, n_tokens]",
"scale": "float32 [n_rows, n_tokens]",
},
"parameterization": {
"survival": "S(t) = exp(-rate * t^shape)",
"shape": "rho = softplus(rho_logit) + 1e-6",
"rate": "softplus(risk_logit) + 1e-8",
"scale": "rate^(-1/shape)",
"time_unit": "years after landmark age",
},
"eligibility": (
"At each age: follow-up extends beyond the landmark, the patient "
"is alive at the landmark, and the configured minimum disease "
"history is available. All disease and death token parameters are "
"exported, including tokens already prevalent by the landmark."
),
}
write_json(output_path / "metadata.json", metadata)
manifest_rows: List[Dict[str, Any]] = []
eligible_queries_by_age: Dict[str, int] = {}
for age_value in ages.tolist():
age = float(age_value)
age_dir = shards_root / age_directory_name(age)
age_dir.mkdir(parents=True, exist_ok=True)
try:
landmark_dataset = LandmarkDataset(
dataset=dataset,
subset_indices=subset_indices,
landmark_ages=np.asarray([age], dtype=np.float32),
model_target_mode=model_target_mode,
min_history_events=min_history_events,
first_occurrence_by_token={},
death_token_ids=death_tokens,
disease_history_mode=disease_history_mode,
)
except RuntimeError as exc:
if "No eligible landmark query samples" not in str(exc):
raise
age_rows = [
write_empty_age_export(
age=age,
age_dir=age_dir,
token_count=len(tokens),
)
]
eligible_count = 0
else:
loader = DataLoader(
landmark_dataset,
batch_size=int(args.batch_size),
shuffle=False,
collate_fn=collate_landmark_fn,
num_workers=int(args.num_workers),
pin_memory=device.type == "cuda",
persistent_workers=args.num_workers > 0,
prefetch_factor=2 if args.num_workers > 0 else None,
)
age_rows = export_age(
age=age,
model=model,
loader=loader,
tokens=tokens,
device=device,
use_amp=bool(args.use_amp),
dataset=dataset,
subset_indices=subset_indices,
age_dir=age_dir,
output_path=output_path,
rows_per_shard=int(args.rows_per_shard),
)
eligible_count = len(landmark_dataset)
manifest_rows.extend(age_rows)
eligible_queries_by_age[f"{age:g}"] = int(eligible_count)
write_csv(output_path / "manifest.csv", MANIFEST_FIELDS, manifest_rows)
metadata["eligible_queries_by_age"] = eligible_queries_by_age
write_json(output_path / "metadata.json", metadata)
print(f"Age {age:g}: exported {eligible_count} patient rows")
nonfinite_shape = sum(
int(row["nonfinite_shape_values"]) for row in manifest_rows
)
nonfinite_scale = sum(
int(row["nonfinite_scale_values"]) for row in manifest_rows
)
metadata["total_exported_query_rows"] = sum(
int(row["n_rows"]) for row in manifest_rows
)
metadata["nonfinite_shape_values"] = nonfinite_shape
metadata["nonfinite_scale_values"] = nonfinite_scale
metadata["complete"] = True
write_json(output_path / "metadata.json", metadata)
if nonfinite_shape or nonfinite_scale:
print(
"WARNING: exported non-finite values: "
f"shape={nonfinite_shape}, scale={nonfinite_scale}."
)
print(f"Saved individual Weibull parameter matrices to: {output_path}")
if __name__ == "__main__":
main()