Files
DeepHealth/export_death_burden.py

523 lines
18 KiB
Python
Raw Normal View History

2026-09-01 13:47:51 +08:00
"""Compute conditional fixed-horizon death burden from Weibull exports.
For patient ``i``, landmark age ``a`` and horizon ``H``:
B_ia(H) = P(T_death <= H | alive at age a)
= 1 - exp(-((H / scale_ia) ** shape_ia)).
The input is the unified HDF5 file written by
``export_weibull_parameters.py``. That source file contains only patients from
its configured test EID file, and each landmark group contains patients alive
and otherwise eligible at that age. The output preserves those rows exactly.
"""
from __future__ import annotations
import argparse
import importlib
import json
from pathlib import Path
from typing import Any, Dict, Iterable, List, Sequence
import numpy as np
FORMAT_VERSION = 1
DEFAULT_HORIZONS = (1.0, 5.0, 10.0)
DEFAULT_COMPRESSION_LEVEL = 4
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description=(
"Compute conditional death probabilities at fixed horizons from "
"an exported Weibull shape/scale HDF5 file."
)
)
parser.add_argument(
"--run_path",
required=True,
help="Run directory containing 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>/death_burden_test_age40_80_step2.h5."
),
)
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 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 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 age_group_name(age: float) -> str:
text = f"{age:g}".replace("-", "minus_").replace(".", "p")
return f"age_{text}"
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 find_death_column(source_file: Any) -> int:
required_paths = (
"tokens/column",
"tokens/token_id",
"tokens/label_code",
"tokens/label_text",
"tokens/outcome_type",
)
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}")
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.")
columns = np.asarray(source_file["tokens/column"][...], dtype=np.int64)
if not np.array_equal(columns, np.arange(token_count, dtype=np.int64)):
raise ValueError("Input HDF5 /tokens/column is not zero-based and ordered.")
codes = decode_strings(source_file["tokens/label_code"][...])
outcome_types = decode_strings(source_file["tokens/outcome_type"][...])
death_columns = [
column
for column, (code, outcome_type) in enumerate(zip(codes, outcome_types))
if code.lower() == "death" and outcome_type.lower() == "death"
]
if len(death_columns) != 1:
raise ValueError(
"Input HDF5 must contain exactly one token labelled Death with "
f"outcome_type='death'; found {len(death_columns)}."
)
return int(death_columns[0])
def death_probability(
shape: np.ndarray,
scale: np.ndarray,
horizons: np.ndarray,
) -> np.ndarray:
"""Return P(death within H | alive at the landmark)."""
shape64 = np.asarray(shape, dtype=np.float64).reshape(-1)
scale64 = np.asarray(scale, dtype=np.float64).reshape(-1)
horizons64 = validate_horizons(horizons)
if shape64.shape != scale64.shape:
raise ValueError("Death shape and scale vectors are not aligned.")
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, :] - 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 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 process_age_group(
*,
source_group: Any,
output_group: Any,
death_column: int,
horizons: np.ndarray,
rows_per_chunk: int,
compression_level: int,
) -> Dict[str, Any]:
landmark_age = float(source_group.attrs["age"])
n_rows = int(source_group["eid"].shape[0])
n_tokens = int(source_group.attrs["n_tokens"])
expected_shape = (n_rows, n_tokens)
for name in ("shape", "scale"):
if source_group[name].shape != expected_shape:
raise ValueError(
f"Age {landmark_age:g} {name} matrix has unexpected dimensions."
)
output_group.attrs["age"] = landmark_age
output_group.attrs["n_rows"] = n_rows
output_group.attrs["n_horizons"] = int(horizons.size)
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(
"death_burden", shape=(0, horizons.size), dtype=np.float32
)
return {
"age": landmark_age,
"n_rows": 0,
"nonfinite_parameter_values": 0,
"nonfinite_burden_values": 0,
}
row_chunk = min(rows_per_chunk, n_rows)
burden_dataset = output_group.create_dataset(
"death_burden",
shape=(n_rows, horizons.size),
dtype=np.float32,
chunks=(row_chunk, 1),
compression="gzip",
compression_opts=compression_level,
shuffle=True,
)
nonfinite_parameters = 0
nonfinite_burden = 0
for row_slice in iter_slices(n_rows, rows_per_chunk):
shape = np.asarray(
source_group["shape"][row_slice, death_column], dtype=np.float32
)
scale = np.asarray(
source_group["scale"][row_slice, death_column], dtype=np.float32
)
burden = death_probability(shape, scale, horizons)
burden_dataset[row_slice, :] = burden
valid_parameters = (
np.isfinite(shape)
& np.isfinite(scale)
& (shape > 0.0)
& (scale > 0.0)
)
nonfinite_parameters += int((~valid_parameters).sum())
nonfinite_burden += int((~np.isfinite(burden)).sum())
output_group.attrs["nonfinite_parameter_values"] = nonfinite_parameters
output_group.attrs["nonfinite_burden_values"] = nonfinite_burden
return {
"age": landmark_age,
"n_rows": n_rows,
"nonfinite_parameter_values": nonfinite_parameters,
"nonfinite_burden_values": nonfinite_burden,
}
def validate_source_file(source_file: Any) -> None:
required_paths = (
"ages",
"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.")
def validate_output_file(
output_file: Any,
*,
ages: np.ndarray,
horizons: np.ndarray,
summaries: Sequence[Dict[str, Any]],
) -> None:
for path in (
"ages",
"horizons",
"death_token",
"test_population",
"landmarks",
"age_summary",
):
if path not in output_file:
raise RuntimeError(f"Output HDF5 is missing /{path}.")
if len(summaries) != int(ages.size):
raise RuntimeError("Output age summary length 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))}"]
expected_shape = (int(summary["n_rows"]), int(horizons.size))
if group["death_burden"].shape != expected_shape:
raise RuntimeError(
f"Death burden matrix has {group['death_burden'].shape}, "
f"expected {expected_shape}."
)
def main() -> None:
args = parse_args()
run_path = Path(args.run_path).resolve()
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)
output_path = (
Path(args.output_path).resolve()
if args.output_path
else run_path / "death_burden_test_age40_80_step2.h5"
)
if output_path == input_path:
raise ValueError("The output path must differ from the input path.")
if output_path.exists():
raise FileExistsError(
f"Output file already exists: {output_path}. Choose a new --output_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.parent.mkdir(parents=True, exist_ok=True)
temporary_output = output_path.with_name(f".{output_path.name}.partial")
if temporary_output.exists():
raise FileExistsError(
f"Partial output already exists: {temporary_output}. Remove or rename it."
)
h5py = require_h5py()
with h5py.File(input_path, "r") as source_file:
validate_source_file(source_file)
death_column = find_death_column(source_file)
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.")
death_token_id = int(source_file["tokens/token_id"][death_column])
death_code = decode_strings(
source_file["tokens/label_code"][death_column : death_column + 1]
)[0]
death_text = decode_strings(
source_file["tokens/label_text"][death_column : death_column + 1]
)[0]
metadata: Dict[str, Any] = {
"format_version": FORMAT_VERSION,
"complete": False,
"definition": (
"B_ia(H) = P(T_death <= H | alive at landmark age a) = "
"1 - exp(-((H / scale_ia)^shape_ia))"
),
"source_weibull_path": str(input_path),
"source_run_path": str(run_path),
"source_population": "test_population",
"death_source_column": death_column,
"death_token_id": death_token_id,
"ages": [float(value) for value in ages.tolist()],
"horizons_years": [float(value) for value in horizons.tolist()],
"matrix_dtype": "float32",
"hdf5_layout": {
"death_token": "/death_token/{source_column,token_id,label_code,label_text}",
"test_population": "/test_population/{eid,dataset_index}",
"landmarks": (
"/landmarks/age_*/{eid,dataset_index,sex,age,death_burden}"
),
"death_burden_dimensions": ["landmark_row", "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["source_population"] = "test"
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)
)
death_group = output_file.create_group("death_token")
death_group.create_dataset(
"source_column", data=np.asarray(death_column, dtype=np.int64)
)
death_group.create_dataset(
"token_id", data=np.asarray(death_token_id, dtype=np.int64)
)
death_group.create_dataset(
"label_code", data=death_code, dtype=string_dtype
)
death_group.create_dataset(
"label_text", data=death_text, dtype=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,
death_column=death_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']} death-burden rows")
summary_group = output_file.create_group("age_summary")
for name, dtype in (
("age", np.float32),
("n_rows", np.int64),
("nonfinite_parameter_values", np.int64),
("nonfinite_burden_values", np.int64),
):
summary_group.create_dataset(
name,
data=np.asarray([row[name] for row in summaries], dtype=dtype),
)
validate_output_file(
output_file,
ages=ages,
horizons=horizons,
summaries=summaries,
)
metadata["total_exported_query_rows"] = sum(
int(row["n_rows"]) for row in summaries
)
metadata["nonfinite_parameter_values"] = sum(
int(row["nonfinite_parameter_values"]) for row in summaries
)
metadata["nonfinite_burden_values"] = sum(
int(row["nonfinite_burden_values"]) for row in 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()
temporary_output.replace(output_path)
print(f"Saved conditional death-burden file to: {output_path}")
if __name__ == "__main__":
main()