diff --git a/export_weibull_parameters.py b/export_weibull_parameters.py new file mode 100644 index 0000000..f87c2a7 --- /dev/null +++ b/export_weibull_parameters.py @@ -0,0 +1,661 @@ +"""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 " + "/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()