Export disease-state Xiad trajectories

This commit is contained in:
2026-09-01 10:59:09 +08:00
parent d728d8585c
commit d9d1398a3a
2 changed files with 1166 additions and 0 deletions

View File

@@ -0,0 +1,995 @@
"""Compute disease-state trajectories X_iad(H) from Weibull exports.
For patient ``i``, landmark age ``a``, disease ``d`` and horizon ``H``:
X_iad(H) = I(T_id <= a) + I(T_id > a) * P_iad(H)
where
P_iad(H) = 1 - exp(-((H / scale_iad) ** shape_iad)).
Diseases already observed by the landmark therefore have state 1. Diseases
not yet observed retain their predicted first-onset probability. Death is not
a disease state and is excluded from the output.
The input is the unified HDF5 file written by
``export_weibull_parameters.py``. The output preserves its landmark-row
alignment and writes one three-dimensional ``x`` dataset per age with shape
``(n_patients, n_diseases, n_horizons)``. Processing is chunked by patient row
so the multi-gigabyte parameter matrices are never loaded in full.
"""
from __future__ import annotations
import argparse
import csv
import importlib
import json
from pathlib import Path
from typing import Any, Dict, Iterable, List, Mapping, Sequence, Tuple
import numpy as np
from dataset import (
DISEASE_HISTORY_MODE_TIMED,
NO_EVENT_IDX,
normalize_disease_history_mode,
)
from eval_data import (
load_json_config,
load_sequence_eval_dataset,
validate_dataset_metadata,
validate_training_mode_config,
)
FORMAT_VERSION = 1
DEFAULT_HORIZONS = (1.0, 5.0, 10.0)
DEFAULT_COMPRESSION_LEVEL = 4
LABEL_OFFSET = NO_EVENT_IDX + 1
AXIS_MAPPING_FIELDS = (
"column",
"source_column",
"label_index",
"token_id",
"code",
"name",
"label_text",
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description=(
"Compute all disease states X_iad(H) from an exported Weibull "
"shape/scale HDF5 file."
)
)
parser.add_argument(
"--run_path",
required=True,
help="Run directory containing train_config.json and the Weibull export.",
)
parser.add_argument(
"--input_path",
default=None,
help=(
"HDF5 file produced by export_weibull_parameters.py. Defaults to "
"<run_path>/weibull_parameters_test_age40_80_step2.h5."
),
)
parser.add_argument(
"--output_path",
default=None,
help=(
"Output HDF5 path. Defaults to "
"<run_path>/disease_state_xiad_test_age40_80_step2.h5."
),
)
parser.add_argument(
"--mapping_output_path",
default=None,
help=(
"CSV describing the exact disease-axis order. Defaults to "
"<output_stem>_icd10_columns.csv."
),
)
parser.add_argument(
"--horizons",
nargs="+",
type=float,
default=list(DEFAULT_HORIZONS),
help="Positive horizons in years. Default: 1 5 10.",
)
parser.add_argument(
"--rows_per_chunk",
type=int,
default=256,
help="Patient rows processed and stored per HDF5 chunk. Default: 256.",
)
parser.add_argument(
"--compression_level",
type=int,
default=DEFAULT_COMPRESSION_LEVEL,
help="Gzip compression level from 0 to 9. Default: 4.",
)
return parser.parse_args()
def require_h5py() -> Any:
try:
return importlib.import_module("h5py")
except ImportError as exc:
raise RuntimeError(
"This script requires h5py in the project Miniconda environment."
) from exc
def parse_int_list(value: Any) -> List[int] | None:
if value is None:
return None
if isinstance(value, (list, tuple, np.ndarray)):
return [int(item) for item 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(item) for item in parsed]
return [int(item.strip()) for item in text.split(",") if item.strip()]
def validate_horizons(values: Iterable[float]) -> np.ndarray:
horizons = np.asarray(list(values), dtype=np.float64)
if horizons.ndim != 1 or horizons.size == 0:
raise ValueError("At least one horizon is required.")
if not np.all(np.isfinite(horizons)) or np.any(horizons <= 0.0):
raise ValueError("Every horizon must be finite and greater than zero.")
if np.unique(horizons).size != horizons.size:
raise ValueError("Horizons must not contain duplicates.")
return horizons
def decode_strings(values: np.ndarray) -> List[str]:
result: List[str] = []
for value in np.asarray(values).tolist():
if isinstance(value, bytes):
result.append(value.decode("utf-8"))
else:
result.append(str(value))
return result
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_axis_rows(labels_file: str | Path) -> List[Dict[str, Any]]:
"""Read labels.csv using the same line-index/token convention as dataset.py."""
rows: List[Dict[str, Any]] = []
labels_path = resolve_project_file(labels_file)
if not labels_path.is_file():
raise FileNotFoundError(f"Labels file not found: {labels_path}")
with labels_path.open("r", encoding="utf-8") as handle:
for label_index, raw in enumerate(handle):
label_text = raw.strip()
if not label_text:
continue
code = label_text.split()[0]
name = label_text[len(code):].strip()
if name.startswith("(") and name.endswith(")"):
name = name[1:-1]
rows.append(
{
"label_index": int(label_index),
"token_id": int(LABEL_OFFSET + label_index),
"code": code,
"name": name or code,
"label_text": label_text,
"outcome_type": (
"death" if code.lower() == "death" else "disease"
),
}
)
if not rows:
raise ValueError(f"Labels file contains no outcomes: {labels_path}")
return rows
def load_matching_dataset(run_path: Path) -> Tuple[Any, Dict[str, Any]]:
cfg = load_json_config(run_path / "train_config.json")
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(
"X_iad(H) requires an all-future export; the run configuration "
f"uses model_target_mode={model_target_mode!r}."
)
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))
)
disease_history_mode = normalize_disease_history_mode(
cfg.get("disease_history_mode", DISEASE_HISTORY_MODE_TIMED)
)
dataset = load_sequence_eval_dataset(
model_target_mode=model_target_mode,
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)),
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)
return dataset, cfg
def validate_source_file(source_file: Any) -> None:
required_paths = (
"ages",
"tokens/column",
"tokens/token_id",
"tokens/label_code",
"tokens/label_text",
"tokens/outcome_type",
"test_population/eid",
"test_population/dataset_index",
"landmarks",
)
missing = [path for path in required_paths if path not in source_file]
if missing:
raise ValueError(f"Input HDF5 is missing required paths: {missing}")
if not bool(source_file.attrs.get("complete", False)):
raise ValueError("Input HDF5 is not marked complete.")
token_count = int(source_file["tokens/token_id"].shape[0])
for path in (
"tokens/column",
"tokens/label_code",
"tokens/label_text",
"tokens/outcome_type",
):
if source_file[path].shape != (token_count,):
raise ValueError(f"/{path} is not aligned with /tokens/token_id.")
def validate_source_token_order(
source_file: Any,
label_rows: Sequence[Mapping[str, Any]],
) -> List[Dict[str, Any]]:
"""Validate source columns against labels.csv and return the disease axis."""
source_token_ids = np.asarray(
source_file["tokens/token_id"][...], dtype=np.int64
)
source_columns = np.asarray(
source_file["tokens/column"][...], dtype=np.int64
)
expected_columns = np.arange(source_token_ids.size, dtype=np.int64)
if not np.array_equal(source_columns, expected_columns):
raise ValueError("Input HDF5 /tokens/column is not zero-based and ordered.")
source_codes = decode_strings(source_file["tokens/label_code"][...])
source_text = decode_strings(source_file["tokens/label_text"][...])
outcome_types = decode_strings(source_file["tokens/outcome_type"][...])
if source_token_ids.size != len(label_rows):
raise ValueError(
"Input HDF5 token count does not match labels.csv: "
f"{source_token_ids.size} versus {len(label_rows)}."
)
disease_rows: List[Dict[str, Any]] = []
for source_column, label_row in enumerate(label_rows):
expected_token = int(label_row["token_id"])
expected_code = str(label_row["code"])
expected_text = str(label_row["label_text"])
expected_type = str(label_row["outcome_type"])
observed = (
int(source_token_ids[source_column]),
source_codes[source_column],
source_text[source_column],
outcome_types[source_column],
)
expected = (
expected_token,
expected_code,
expected_text,
expected_type,
)
if observed != expected:
raise ValueError(
"Input HDF5 disease order does not match labels.csv at source "
f"column {source_column}: observed={observed!r}, "
f"expected={expected!r}."
)
if expected_type == "disease":
disease_rows.append(
{
"column": len(disease_rows),
"source_column": int(source_column),
"label_index": int(label_row["label_index"]),
"token_id": expected_token,
"code": expected_code,
"name": str(label_row["name"]),
"label_text": expected_text,
"outcome_type": "disease",
}
)
if not disease_rows:
raise ValueError("The validated labels.csv contains no disease outcomes.")
return disease_rows
def validate_token_table(
dataset: Any,
disease_rows: Sequence[Mapping[str, Any]],
) -> None:
for row in disease_rows:
token = int(row["token_id"])
code = str(row["code"])
dataset_code = dataset.label_id_to_code.get(token)
if dataset_code is None:
raise ValueError(f"Disease token {token} is absent from the dataset vocabulary.")
if str(dataset_code) != str(code):
raise ValueError(
f"Disease token {token} mismatch: input HDF5 has {code!r}, "
f"dataset has {dataset_code!r}."
)
def write_axis_mapping_csv(
output_path: Path,
disease_rows: Sequence[Mapping[str, Any]],
) -> None:
with output_path.open("w", encoding="utf-8", newline="") as handle:
writer = csv.DictWriter(handle, fieldnames=list(AXIS_MAPPING_FIELDS))
writer.writeheader()
for row in disease_rows:
writer.writerow({field: row[field] for field in AXIS_MAPPING_FIELDS})
def validate_axis_mapping_csv(
mapping_path: Path,
disease_rows: Sequence[Mapping[str, Any]],
) -> None:
with mapping_path.open("r", encoding="utf-8", newline="") as handle:
observed_rows = list(csv.DictReader(handle))
if len(observed_rows) != len(disease_rows):
raise RuntimeError("Disease-axis CSV row count is incorrect.")
for expected, observed in zip(disease_rows, observed_rows):
expected_text = {
field: str(expected[field]) for field in AXIS_MAPPING_FIELDS
}
if observed != expected_text:
raise RuntimeError(
"Disease-axis CSV does not match the validated labels order at "
f"output column {expected['column']}."
)
def validate_population_alignment(source_file: Any, dataset: Any) -> None:
dataset_indices = np.asarray(
source_file["test_population/dataset_index"][...], dtype=np.int64
)
eids = np.asarray(source_file["test_population/eid"][...], dtype=np.int64)
if dataset_indices.shape != eids.shape:
raise ValueError("Input test-population EIDs and dataset indices are misaligned.")
if dataset_indices.size == 0:
raise ValueError("Input test population is empty.")
if np.any(dataset_indices < 0) or np.any(dataset_indices >= len(dataset.samples)):
raise ValueError("Input test population contains an invalid dataset index.")
expected_eids = np.asarray(
[int(dataset.samples[int(index)]["eid"]) for index in dataset_indices],
dtype=np.int64,
)
if not np.array_equal(eids, expected_eids):
mismatch = int(np.flatnonzero(eids != expected_eids)[0])
raise ValueError(
"Input test population does not match the dataset loaded from the "
f"run configuration; first mismatch is at row {mismatch}."
)
def full_disease_history(sample: Mapping[str, Any]) -> Tuple[np.ndarray, np.ndarray]:
events = np.asarray(sample["event_seq"], dtype=np.int64)
times = np.asarray(sample["time_seq"], dtype=np.float32)
target_events = np.asarray(sample["target_event_seq"], dtype=np.int64)
target_times = np.asarray(sample["target_time_seq"], dtype=np.float32)
if target_events.size > 0:
events = np.concatenate([events, target_events[-1:]])
times = np.concatenate([times, target_times[-1:]])
if events.shape != times.shape:
raise ValueError("Disease events and times are misaligned in the dataset sample.")
return events, times
def build_prevalent_mask(
*,
dataset: Any,
dataset_indices: np.ndarray,
landmark_age: float,
token_to_column: Mapping[int, int],
n_diseases: int,
) -> np.ndarray:
"""Return I(T_id <= a) for a chunk of landmark rows."""
prevalent = np.zeros((len(dataset_indices), n_diseases), dtype=bool)
for row, dataset_index in enumerate(np.asarray(dataset_indices).tolist()):
sample = dataset.samples[int(dataset_index)]
events, times = full_disease_history(sample)
historical_events = events[times <= np.float32(landmark_age)]
for token in np.unique(historical_events).tolist():
column = token_to_column.get(int(token))
if column is not None:
prevalent[row, column] = True
return prevalent
def weibull_probability(
shape: np.ndarray,
scale: np.ndarray,
horizons: np.ndarray,
) -> np.ndarray:
"""Calculate stable Weibull fixed-horizon probabilities."""
shape64 = np.asarray(shape, dtype=np.float64)
scale64 = np.asarray(scale, dtype=np.float64)
horizons64 = np.asarray(horizons, dtype=np.float64)
if shape64.shape != scale64.shape or shape64.ndim != 2:
raise ValueError("shape and scale must be aligned two-dimensional matrices.")
horizons64 = validate_horizons(horizons64)
valid = (
np.isfinite(shape64)
& np.isfinite(scale64)
& (shape64 > 0.0)
& (scale64 > 0.0)
)
with np.errstate(divide="ignore", invalid="ignore", over="ignore"):
log_cumulative_hazard = shape64[:, :, None] * (
np.log(horizons64)[None, None, :]
- np.log(scale64)[:, :, None]
)
cumulative_hazard = np.exp(
np.clip(log_cumulative_hazard, -87.0, 40.0)
)
probability = -np.expm1(-cumulative_hazard)
probability[~valid, :] = np.nan
return probability.astype(np.float32, copy=False)
def compute_xiad(
*,
shape: np.ndarray,
scale: np.ndarray,
horizons: np.ndarray,
prevalent: np.ndarray,
) -> np.ndarray:
probability = weibull_probability(shape, scale, horizons)
prevalent = np.asarray(prevalent, dtype=bool)
if prevalent.shape != probability.shape[:2]:
raise ValueError("The prevalent mask is not aligned with shape and scale.")
return np.where(prevalent[:, :, None], 1.0, probability).astype(
np.float32, copy=False
)
def iter_slices(n_rows: int, rows_per_chunk: int) -> Iterable[slice]:
for start in range(0, n_rows, rows_per_chunk):
yield slice(start, min(start + rows_per_chunk, n_rows))
def age_group_name(age: float) -> str:
text = f"{age:g}".replace("-", "minus_").replace(".", "p")
return f"age_{text}"
def copy_vector_dataset(
source_group: Any,
output_group: Any,
name: str,
rows_per_chunk: int,
compression_level: int,
) -> None:
source = source_group[name]
n_rows = int(source.shape[0])
options: Dict[str, Any] = {}
if n_rows > 0:
options = {
"chunks": (min(rows_per_chunk, n_rows),),
"compression": "gzip",
"compression_opts": compression_level,
"shuffle": True,
}
output_group.create_dataset(name, data=source[...], dtype=source.dtype, **options)
def write_string_dataset(
group: Any,
name: str,
values: Sequence[str],
string_dtype: Any,
) -> None:
group.create_dataset(
name,
data=np.asarray(list(values), dtype=object),
dtype=string_dtype,
)
def process_age_group(
*,
source_group: Any,
output_group: Any,
dataset: Any,
disease_source_columns: np.ndarray,
token_to_column: Mapping[int, int],
horizons: np.ndarray,
rows_per_chunk: int,
compression_level: int,
) -> Dict[str, Any]:
n_rows = int(source_group["eid"].shape[0])
n_diseases = int(disease_source_columns.size)
n_horizons = int(horizons.size)
landmark_age = float(source_group.attrs["age"])
expected_matrix_shape = (n_rows, int(source_group.attrs["n_tokens"]))
if source_group["shape"].shape != expected_matrix_shape:
raise ValueError(
f"Age {landmark_age:g} shape matrix has unexpected dimensions."
)
if source_group["scale"].shape != expected_matrix_shape:
raise ValueError(
f"Age {landmark_age:g} scale matrix has unexpected dimensions."
)
output_group.attrs["age"] = landmark_age
output_group.attrs["n_rows"] = n_rows
output_group.attrs["n_diseases"] = n_diseases
output_group.attrs["n_horizons"] = n_horizons
for name in ("eid", "dataset_index", "sex", "age"):
copy_vector_dataset(
source_group,
output_group,
name,
rows_per_chunk,
compression_level,
)
if n_rows == 0:
output_group.create_dataset(
"prevalent", shape=(0, n_diseases), dtype=np.uint8
)
output_group.create_dataset(
"x", shape=(0, n_diseases, n_horizons), dtype=np.float32
)
return {
"age": landmark_age,
"n_rows": 0,
"prevalent_values": 0,
"nonfinite_parameter_values": 0,
"nonfinite_x_values": 0,
}
row_chunk = min(rows_per_chunk, n_rows)
prevalent_dataset = output_group.create_dataset(
"prevalent",
shape=(n_rows, n_diseases),
dtype=np.uint8,
chunks=(row_chunk, n_diseases),
compression="gzip",
compression_opts=compression_level,
shuffle=True,
)
x_dataset = output_group.create_dataset(
"x",
shape=(n_rows, n_diseases, n_horizons),
dtype=np.float32,
chunks=(row_chunk, n_diseases, 1),
compression="gzip",
compression_opts=compression_level,
shuffle=True,
)
prevalent_values = 0
nonfinite_parameter_values = 0
nonfinite_x_values = 0
for row_slice in iter_slices(n_rows, rows_per_chunk):
dataset_indices = np.asarray(
source_group["dataset_index"][row_slice], dtype=np.int64
)
if np.any(dataset_indices < 0) or np.any(
dataset_indices >= len(dataset.samples)
):
raise ValueError(
f"Age {landmark_age:g} contains an invalid dataset index."
)
shape = np.asarray(
source_group["shape"][row_slice, disease_source_columns],
dtype=np.float32,
)
scale = np.asarray(
source_group["scale"][row_slice, disease_source_columns],
dtype=np.float32,
)
prevalent = build_prevalent_mask(
dataset=dataset,
dataset_indices=dataset_indices,
landmark_age=landmark_age,
token_to_column=token_to_column,
n_diseases=n_diseases,
)
x = compute_xiad(
shape=shape,
scale=scale,
horizons=horizons,
prevalent=prevalent,
)
prevalent_dataset[row_slice, :] = prevalent.astype(np.uint8, copy=False)
x_dataset[row_slice, :, :] = x
prevalent_values += int(prevalent.sum())
valid_parameters = (
np.isfinite(shape)
& np.isfinite(scale)
& (shape > 0.0)
& (scale > 0.0)
)
nonfinite_parameter_values += int((~valid_parameters).sum())
nonfinite_x_values += int((~np.isfinite(x)).sum())
output_group.attrs["prevalent_values"] = prevalent_values
output_group.attrs["nonfinite_parameter_values"] = nonfinite_parameter_values
output_group.attrs["nonfinite_x_values"] = nonfinite_x_values
return {
"age": landmark_age,
"n_rows": n_rows,
"prevalent_values": prevalent_values,
"nonfinite_parameter_values": nonfinite_parameter_values,
"nonfinite_x_values": nonfinite_x_values,
}
def validate_output_file(
output_file: Any,
*,
ages: np.ndarray,
n_diseases: int,
n_horizons: int,
disease_rows: Sequence[Mapping[str, Any]],
summaries: Sequence[Mapping[str, Any]],
) -> None:
for path in ("tokens", "test_population", "landmarks", "age_summary"):
if path not in output_file:
raise RuntimeError(f"Output HDF5 is missing /{path}.")
if output_file["tokens/token_id"].shape != (n_diseases,):
raise RuntimeError("Output disease token table has an unexpected length.")
expected_integer_columns = {
"column": np.asarray(
[row["column"] for row in disease_rows], dtype=np.int64
),
"source_column": np.asarray(
[row["source_column"] for row in disease_rows], dtype=np.int64
),
"label_index": np.asarray(
[row["label_index"] for row in disease_rows], dtype=np.int64
),
"token_id": np.asarray(
[row["token_id"] for row in disease_rows], dtype=np.int64
),
}
for name, expected in expected_integer_columns.items():
observed = np.asarray(output_file[f"tokens/{name}"][...], dtype=np.int64)
if not np.array_equal(observed, expected):
raise RuntimeError(f"Output /tokens/{name} order is incorrect.")
expected_string_columns = {
"label_code": [str(row["code"]) for row in disease_rows],
"name": [str(row["name"]) for row in disease_rows],
"label_text": [str(row["label_text"]) for row in disease_rows],
"outcome_type": ["disease"] * len(disease_rows),
}
for name, expected in expected_string_columns.items():
observed = decode_strings(output_file[f"tokens/{name}"][...])
if observed != expected:
raise RuntimeError(f"Output /tokens/{name} order is incorrect.")
if len(summaries) != int(ages.size):
raise RuntimeError("Output age summary does not match the age grid.")
for age_value, summary in zip(ages.tolist(), summaries):
group = output_file[f"landmarks/{age_group_name(float(age_value))}"]
n_rows = int(summary["n_rows"])
if group["prevalent"].shape != (n_rows, n_diseases):
raise RuntimeError("Output prevalent matrix has unexpected dimensions.")
if group["x"].shape != (n_rows, n_diseases, n_horizons):
raise RuntimeError("Output X_iad(H) matrix has unexpected dimensions.")
def main() -> None:
args = parse_args()
run_path = Path(args.run_path).resolve()
config_path = run_path / "train_config.json"
if not config_path.is_file():
raise FileNotFoundError(config_path)
input_path = (
Path(args.input_path).resolve()
if args.input_path
else run_path / "weibull_parameters_test_age40_80_step2.h5"
)
if not input_path.is_file():
raise FileNotFoundError(input_path)
horizons = validate_horizons(args.horizons)
if args.rows_per_chunk <= 0:
raise ValueError("rows_per_chunk must be greater than zero.")
if not 0 <= args.compression_level <= 9:
raise ValueError("compression_level must be between 0 and 9.")
output_path = (
Path(args.output_path).resolve()
if args.output_path
else run_path / "disease_state_xiad_test_age40_80_step2.h5"
)
mapping_output_path = (
Path(args.mapping_output_path).resolve()
if args.mapping_output_path
else output_path.with_name(f"{output_path.stem}_icd10_columns.csv")
)
if output_path == input_path:
raise ValueError("The output path must differ from the input path.")
if mapping_output_path in {input_path, output_path}:
raise ValueError("The ICD-10 mapping path must be a separate file.")
if output_path.exists():
raise FileExistsError(
f"Output file already exists: {output_path}. Choose a new --output_path."
)
if mapping_output_path.exists():
raise FileExistsError(
"Disease-axis mapping already exists: "
f"{mapping_output_path}. Choose a new --mapping_output_path."
)
output_path.parent.mkdir(parents=True, exist_ok=True)
mapping_output_path.parent.mkdir(parents=True, exist_ok=True)
temporary_output = output_path.with_name(f".{output_path.name}.partial")
temporary_mapping = mapping_output_path.with_name(
f".{mapping_output_path.name}.partial"
)
if temporary_output.exists():
raise FileExistsError(
f"Partial output already exists: {temporary_output}. Remove or rename it."
)
if temporary_mapping.exists():
raise FileExistsError(
f"Partial mapping already exists: {temporary_mapping}. Remove or rename it."
)
h5py = require_h5py()
print(f"Loading dataset from run configuration: {run_path}")
dataset, cfg = load_matching_dataset(run_path)
labels_file = str(cfg.get("labels_file", "labels.csv"))
labels_path = resolve_project_file(labels_file).resolve()
label_rows = load_label_axis_rows(labels_path)
with h5py.File(input_path, "r") as source_file:
validate_source_file(source_file)
validate_population_alignment(source_file, dataset)
disease_rows = validate_source_token_order(source_file, label_rows)
validate_token_table(dataset, disease_rows)
ages = np.asarray(source_file["ages"][...], dtype=np.float32)
if ages.ndim != 1 or ages.size == 0:
raise ValueError("The input age grid is empty or invalid.")
disease_source_columns = np.asarray(
[row["source_column"] for row in disease_rows], dtype=np.int64
)
disease_label_indices = np.asarray(
[row["label_index"] for row in disease_rows], dtype=np.int64
)
disease_token_ids = np.asarray(
[row["token_id"] for row in disease_rows], dtype=np.int64
)
disease_codes = [str(row["code"]) for row in disease_rows]
disease_names = [str(row["name"]) for row in disease_rows]
disease_text = [str(row["label_text"]) for row in disease_rows]
token_to_column = {
int(token): column
for column, token in enumerate(disease_token_ids.tolist())
}
metadata: Dict[str, Any] = {
"format_version": FORMAT_VERSION,
"complete": False,
"definition": (
"X_iad(H) = I(T_id <= a) + I(T_id > a) * "
"[1 - exp(-((H / scale_iad)^shape_iad))]"
),
"source_weibull_path": str(input_path),
"source_run_path": str(run_path),
"labels_file": str(labels_path),
"disease_axis_mapping_csv": str(mapping_output_path),
"disease_axis_order": (
"Exactly labels.csv line order with Death removed; output "
"column is the second dimension of every landmark x dataset."
),
"disease_axis_validated_against_labels": True,
"ages": [float(value) for value in ages.tolist()],
"horizons_years": [float(value) for value in horizons.tolist()],
"n_diseases": int(disease_token_ids.size),
"matrix_dtype": "float32",
"prevalent_dtype": "uint8",
"death_excluded": True,
"prevalence_boundary": "first observed disease time <= landmark age",
"hdf5_layout": {
"tokens": (
"/tokens/{column,source_column,label_index,token_id,"
"label_code,name,label_text,outcome_type}"
),
"horizons": "/horizons",
"test_population": "/test_population/{eid,dataset_index}",
"landmarks": (
"/landmarks/age_*/{eid,dataset_index,sex,age,prevalent,x}"
),
"x_dimensions": ["landmark_row", "disease", "horizon"],
},
"rows_per_chunk": int(args.rows_per_chunk),
"compression": "gzip",
"compression_level": int(args.compression_level),
}
string_dtype = h5py.string_dtype(encoding="utf-8")
summaries: List[Dict[str, Any]] = []
with h5py.File(temporary_output, "w") as output_file:
output_file.attrs["format_version"] = FORMAT_VERSION
output_file.attrs["complete"] = False
output_file.attrs["source_weibull_path"] = str(input_path)
output_file.attrs["run_path"] = str(run_path)
output_file.attrs["n_diseases"] = int(disease_token_ids.size)
output_file.attrs["n_horizons"] = int(horizons.size)
output_file.attrs["matrix_dtype"] = "float32"
metadata_dataset = output_file.create_dataset(
"metadata_json", shape=(), dtype=string_dtype
)
metadata_dataset[()] = json.dumps(
metadata, ensure_ascii=False, separators=(",", ":")
)
output_file.create_dataset("ages", data=ages)
output_file.create_dataset(
"horizons", data=horizons.astype(np.float32)
)
token_group = output_file.create_group("tokens")
token_group.create_dataset(
"column",
data=np.arange(disease_token_ids.size, dtype=np.int64),
)
token_group.create_dataset(
"source_column", data=disease_source_columns
)
token_group.create_dataset(
"label_index", data=disease_label_indices
)
token_group.create_dataset("token_id", data=disease_token_ids)
write_string_dataset(
token_group,
"label_code",
disease_codes,
string_dtype,
)
write_string_dataset(
token_group,
"name",
disease_names,
string_dtype,
)
write_string_dataset(
token_group,
"label_text",
disease_text,
string_dtype,
)
write_string_dataset(
token_group,
"outcome_type",
["disease"] * len(disease_rows),
string_dtype,
)
population_group = output_file.create_group("test_population")
for name in ("eid", "dataset_index"):
copy_vector_dataset(
source_file["test_population"],
population_group,
name,
args.rows_per_chunk,
args.compression_level,
)
landmark_root = output_file.create_group("landmarks")
for age_value in ages.tolist():
age = float(age_value)
group_name = age_group_name(age)
source_path = f"landmarks/{group_name}"
if source_path not in source_file:
raise ValueError(f"Input HDF5 is missing /{source_path}.")
output_group = landmark_root.create_group(group_name)
summary = process_age_group(
source_group=source_file[source_path],
output_group=output_group,
dataset=dataset,
disease_source_columns=disease_source_columns,
token_to_column=token_to_column,
horizons=horizons,
rows_per_chunk=args.rows_per_chunk,
compression_level=args.compression_level,
)
summaries.append(summary)
metadata["completed_ages"] = [
float(row["age"]) for row in summaries
]
metadata_dataset[()] = json.dumps(
metadata, ensure_ascii=False, separators=(",", ":")
)
output_file.flush()
print(
f"Age {age:g}: wrote {summary['n_rows']} rows, "
f"prevalent states={summary['prevalent_values']}"
)
age_summary = output_file.create_group("age_summary")
for name, dtype in (
("age", np.float32),
("n_rows", np.int64),
("prevalent_values", np.int64),
("nonfinite_parameter_values", np.int64),
("nonfinite_x_values", np.int64),
):
age_summary.create_dataset(
name,
data=np.asarray([row[name] for row in summaries], dtype=dtype),
)
metadata["total_exported_query_rows"] = sum(
int(row["n_rows"]) for row in summaries
)
metadata["total_prevalent_values"] = sum(
int(row["prevalent_values"]) for row in summaries
)
metadata["nonfinite_parameter_values"] = sum(
int(row["nonfinite_parameter_values"]) for row in summaries
)
metadata["nonfinite_x_values"] = sum(
int(row["nonfinite_x_values"]) for row in summaries
)
validate_output_file(
output_file,
ages=ages,
n_diseases=int(disease_token_ids.size),
n_horizons=int(horizons.size),
disease_rows=disease_rows,
summaries=summaries,
)
metadata["validated"] = True
metadata["complete"] = True
metadata_dataset[()] = json.dumps(
metadata, ensure_ascii=False, separators=(",", ":")
)
output_file.attrs["validated"] = True
output_file.attrs.modify("complete", True)
output_file.flush()
write_axis_mapping_csv(temporary_mapping, disease_rows)
validate_axis_mapping_csv(temporary_mapping, disease_rows)
temporary_output.replace(output_path)
temporary_mapping.replace(mapping_output_path)
print(f"Saved disease-state X_iad(H) file to: {output_path}")
print(f"Saved disease-axis ICD-10 mapping to: {mapping_output_path}")
if __name__ == "__main__":
main()