"""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()