"""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 " "/weibull_parameters_test_age40_80_step2.h5." ), ) parser.add_argument( "--output_path", default=None, help=( "Output HDF5 path. Defaults to " "/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 " "_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()