from __future__ import annotations import json import logging import sys import time import csv from dataclasses import dataclass from datetime import datetime import math from pathlib import Path from typing import Any, Dict, Iterable, Tuple import numpy as np import torch from torch.optim import AdamW from torch.utils.data import Subset from dataset import AllFutureHealthDataset, HealthDataset from models import DeepHealth from targets import NO_EVENT_IDX, PAD_IDX, RESERVED_IDX @dataclass(frozen=True) class ContinuousRobustScalerStats: """Train-split robust scaling statistics aligned to ``cont_type_ids``.""" cont_type_ids: tuple[int, ...] center: np.ndarray scale: np.ndarray observation_count: np.ndarray quantile_range: tuple[float, float] = (25.0, 75.0) def as_metadata(self) -> Dict[str, Any]: return { "method": "robust", "fitted_on": "train_subset", "quantile_range": [float(x) for x in self.quantile_range], "cont_type_ids": [int(x) for x in self.cont_type_ids], "observation_count": [int(x) for x in self.observation_count.tolist()], "center_buffer": "tokenizer.continuous_value_center", "scale_buffer": "tokenizer.continuous_value_scale", } @dataclass(frozen=True) class AllFutureBaselineStats: """Training-only marginal first-onset rates for output initialization.""" event_count: np.ndarray at_risk_exposure: np.ndarray rate: np.ndarray bias: np.ndarray query_count: int rate_floor: float rate_ceiling: float def as_metadata(self) -> Dict[str, Any]: positive = self.rate[self.at_risk_exposure > 0] if positive.size: rate_summary = { "min": float(positive.min()), "median": float(np.median(positive)), "max": float(positive.max()), } else: rate_summary = {"min": 0.0, "median": 0.0, "max": 0.0} return { "method": "training_marginal_first_onset_rate", "fitted_on": "one_seeded_query_per_training_patient", "query_distribution": "patient_interval_time_uniform", "query_count": int(self.query_count), "observed_first_onsets": int(self.event_count.sum()), "rate_floor": float(self.rate_floor), "rate_ceiling": float(self.rate_ceiling), "rate_per_year": rate_summary, "checkpoint_parameter": "risk_head.bias", } def fit_all_future_baseline( dataset: AllFutureHealthDataset, subset: Subset, *, seed: int, ignored_idx: Iterable[int] = (PAD_IDX, RESERVED_IDX, NO_EVENT_IDX), rate_floor: float = 1e-6, rate_ceiling: float = 10.0, ) -> AllFutureBaselineStats: """Fit marginal per-outcome rates under the training query distribution.""" if subset.dataset is not dataset: raise ValueError("subset must reference the all-future training dataset") if rate_floor <= 0 or rate_ceiling <= rate_floor: raise ValueError("Require 0 < rate_floor < rate_ceiling") subset_indices = np.asarray(subset.indices, dtype=np.int64) if subset_indices.ndim != 1 or subset_indices.size == 0: raise ValueError("training subset must contain at least one patient") vocab_size = int(dataset.vocab_size) ignored = { int(idx) for idx in ignored_idx if 0 <= int(idx) < vocab_size } event_count = np.zeros(vocab_size, dtype=np.int64) exposure_adjustment = np.zeros(vocab_size, dtype=np.float64) common_exposure = 0.0 rng = np.random.RandomState(int(seed)) for patient_index in subset_indices.tolist(): patient = dataset.patients[int(patient_index)] t_query = dataset.sample_query(patient, rng) times = np.asarray(patient["times"], dtype=np.float64) labels = np.asarray(patient["labels"], dtype=np.int64) censor_time = max(float(patient["t_obs"]) - t_query, 0.0) common_exposure += censor_time prevalent = { int(label) for label in labels[times <= t_query].tolist() if 0 <= int(label) < vocab_size and int(label) not in ignored } for label in prevalent: exposure_adjustment[label] -= censor_time # Keep the earliest future occurrence defensively; prepared disease # events are already deduplicated to their first occurrence. first_future_dt: Dict[int, float] = {} for label, event_time in zip(labels[times > t_query], times[times > t_query]): label = int(label) if label in ignored or label in prevalent or not (0 <= label < vocab_size): continue event_dt = max(float(event_time) - t_query, 0.0) previous = first_future_dt.get(label) if previous is None or event_dt < previous: first_future_dt[label] = event_dt for label, event_dt in first_future_dt.items(): event_count[label] += 1 exposure_adjustment[label] += event_dt - censor_time at_risk_exposure = common_exposure + exposure_adjustment at_risk_exposure = np.maximum(at_risk_exposure, 0.0) rate = np.full(vocab_size, float(rate_floor), dtype=np.float64) has_exposure = at_risk_exposure > 0 rate[has_exposure] = ( event_count[has_exposure].astype(np.float64) / at_risk_exposure[has_exposure] ) rate = np.clip(rate, float(rate_floor), float(rate_ceiling)) for idx in ignored: at_risk_exposure[idx] = 0.0 rate[idx] = 0.0 bias = np.zeros(vocab_size, dtype=np.float64) valid = np.ones(vocab_size, dtype=bool) if ignored: valid[np.asarray(sorted(ignored), dtype=np.int64)] = False bias[valid] = np.log(np.expm1(rate[valid])) if not np.isfinite(bias).all(): raise RuntimeError("All-future baseline initialization is non-finite") return AllFutureBaselineStats( event_count=event_count, at_risk_exposure=at_risk_exposure.astype(np.float32), rate=rate.astype(np.float32), bias=bias.astype(np.float32), query_count=int(subset_indices.size), rate_floor=float(rate_floor), rate_ceiling=float(rate_ceiling), ) def fit_continuous_robust_scaler( dataset: HealthDataset | AllFutureHealthDataset, subset: Subset, *, quantile_range: tuple[float, float] = (25.0, 75.0), scale_epsilon: float = 1e-6, ) -> ContinuousRobustScalerStats: """Fit median/IQR statistics using only patients in the training subset. The prepared arrays remain unchanged. Scaling is performed later inside the model tokenizer so the fitted center and scale can live in the checkpoint. """ if subset.dataset is not dataset: raise ValueError("subset must reference the dataset used to fit the scaler") low, high = (float(quantile_range[0]), float(quantile_range[1])) if not (0.0 <= low < high <= 100.0): raise ValueError( "quantile_range must satisfy 0 <= low < high <= 100, got " f"{quantile_range!r}" ) if scale_epsilon <= 0: raise ValueError("scale_epsilon must be > 0") cont_type_ids = tuple(int(x) for x in dataset.cont_type_ids) n_cont_types = len(cont_type_ids) if n_cont_types == 0: empty = np.zeros(0, dtype=np.float32) return ContinuousRobustScalerStats( cont_type_ids=cont_type_ids, center=empty.copy(), scale=empty.copy(), observation_count=np.zeros(0, dtype=np.int64), quantile_range=(low, high), ) subset_indices = np.asarray(subset.indices, dtype=np.int64) if subset_indices.ndim != 1 or subset_indices.size == 0: raise ValueError("training subset must contain at least one patient") type_to_column = np.full(int(dataset.n_types), -1, dtype=np.int64) for column, type_id in enumerate(cont_type_ids): if type_id <= 0 or type_id >= len(type_to_column): raise ValueError( f"continuous type id {type_id} is outside [1, {len(type_to_column)})" ) type_to_column[type_id] = column if hasattr(dataset, "patients"): records = dataset.patients elif hasattr(dataset, "samples"): records = dataset.samples else: raise TypeError( "dataset must expose patient records through .patients or .samples" ) values = np.full( (int(subset_indices.size), n_cont_types), np.nan, dtype=np.float32, ) for row, patient_index in enumerate(subset_indices.tolist()): patient = records[int(patient_index)] other_type = np.asarray(patient["other_type"], dtype=np.int64) other_value = np.asarray(patient["other_value"], dtype=np.float32) other_kind = np.asarray(patient["other_value_kind"], dtype=np.int64) continuous = other_kind == 1 if not np.any(continuous): continue selected_type = other_type[continuous] selected_value = other_value[continuous] valid_type = (selected_type > 0) & (selected_type < len(type_to_column)) columns = np.full(selected_type.shape, -1, dtype=np.int64) columns[valid_type] = type_to_column[selected_type[valid_type]] valid = (columns >= 0) & np.isfinite(selected_value) values[row, columns[valid]] = selected_value[valid] observation_count = np.isfinite(values).sum(axis=0).astype(np.int64) missing_types = [ type_id for type_id, count in zip(cont_type_ids, observation_count.tolist()) if count == 0 ] if missing_types: raise ValueError( "Training subset has no finite observations for continuous type ids: " f"{missing_types}" ) low_value, center, high_value = np.nanpercentile( values, [low, 50.0, high], axis=0, ) scale = high_value - low_value near_constant = (~np.isfinite(scale)) | (np.abs(scale) <= float(scale_epsilon)) scale[near_constant] = 1.0 center = np.asarray(center, dtype=np.float32) scale = np.asarray(scale, dtype=np.float32) if not np.isfinite(center).all(): raise RuntimeError("Robust scaler produced non-finite center values") if not np.isfinite(scale).all() or np.any(scale <= 0): raise RuntimeError("Robust scaler produced invalid scale values") return ContinuousRobustScalerStats( cont_type_ids=cont_type_ids, center=center, scale=scale, observation_count=observation_count, quantile_range=(low, high), ) def create_unique_run_dir(name_fn, runs_root: Path = Path("runs")) -> tuple[Path, str]: while True: timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") run_name = name_fn(timestamp) run_dir = runs_root / run_name try: run_dir.mkdir(parents=True, exist_ok=False) return run_dir, run_name except FileExistsError: time.sleep(1.0) def setup_logging(run_dir: Path) -> logging.Logger: run_dir.mkdir(parents=True, exist_ok=True) logger = logging.getLogger("DeepHealth") logger.setLevel(logging.INFO) logger.handlers.clear() formatter = logging.Formatter( "%(asctime)s - %(levelname)s - %(message)s", datefmt="%Y-%m-%d %H:%M:%S", ) console_handler = logging.StreamHandler(sys.stdout) console_handler.setLevel(logging.INFO) console_handler.setFormatter(formatter) logger.addHandler(console_handler) file_handler = logging.FileHandler(run_dir / "train.log", mode="w") file_handler.setLevel(logging.INFO) file_handler.setFormatter(formatter) logger.addHandler(file_handler) return logger def set_seed(seed: int) -> None: np.random.seed(seed) torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed(seed) def load_extra_info_types_file(path: str) -> list[int]: file_path = Path(path) if not file_path.is_file(): raise FileNotFoundError(f"extra_info_types_file not found: {path}") text = file_path.read_text(encoding="utf-8").strip() if not text: return [] if text.startswith("["): raw_items = json.loads(text) if not isinstance(raw_items, list): raise ValueError("extra_info_types_file JSON must be a list") else: raw_items = [] for line in text.splitlines(): line = line.split("#", 1)[0].strip() if line: raw_items.extend(line.replace(",", " ").replace(";", " ").split()) try: return [int(x) for x in raw_items] except (TypeError, ValueError) as exc: raise ValueError(f"Invalid extra info type id in {path}") from exc def format_extra_info_types(extra_info_types: Iterable[int] | None) -> str: if extra_info_types is None: return "all" values = [int(x) for x in extra_info_types] if not values: return "none" return str(values) def load_eid_file(path: str | Path) -> set[int]: file_path = Path(path) if not file_path.is_file(): raise FileNotFoundError(f"eid split file not found: {file_path}") with file_path.open(newline="", encoding="utf-8-sig") as f: reader = csv.DictReader(f) if reader.fieldnames is None or "eid" not in reader.fieldnames: raise ValueError( f"eid split file must contain an 'eid' column: {file_path}" ) out: set[int] = set() for row in reader: raw = (row.get("eid") or "").strip() if raw: out.add(int(raw)) if not out: raise ValueError(f"eid split file is empty: {file_path}") return out def configure_torch_for_training(device: torch.device) -> None: if device.type != "cuda": return torch.backends.cuda.matmul.allow_tf32 = True torch.backends.cudnn.allow_tf32 = True if hasattr(torch, "set_float32_matmul_precision"): torch.set_float32_matmul_precision("high") def resolve_device(device_arg: str) -> torch.device: requested = device_arg.strip().lower() if requested == "cpu": return torch.device("cpu") if requested == "cuda": return torch.device("cuda" if torch.cuda.is_available() else "cpu") if requested.startswith("cuda:"): if not torch.cuda.is_available(): return torch.device("cpu") index = int(requested.split(":", 1)[1]) if index < 0 or index >= torch.cuda.device_count(): raise ValueError(f"Requested CUDA device is out of range: {device_arg}") return torch.device(f"cuda:{index}") raise ValueError(f"Unsupported device: {device_arg}") def get_lr(epoch: int, args: Any, adaptive_lr: float) -> float: if epoch < args.warmup_epochs: return adaptive_lr * (epoch + 1) / args.warmup_epochs progress = (epoch - args.warmup_epochs) / max( 1, args.max_epochs - args.warmup_epochs ) cosine = 0.5 * (1 + math.cos(math.pi * progress)) return adaptive_lr * ( args.min_lr_ratio + cosine * (1 - args.min_lr_ratio) ) def move_batch_to_device( batch: Dict[str, torch.Tensor], device: torch.device, ) -> Dict[str, torch.Tensor]: non_blocking = device.type == "cuda" return { key: value.to(device, non_blocking=non_blocking) if isinstance(value, torch.Tensor) else value for key, value in batch.items() } def split_dataset( dataset: HealthDataset, train_ratio: float, val_ratio: float, test_ratio: float, seed: int, ) -> Tuple[Subset, Subset, Subset]: total = train_ratio + val_ratio + test_ratio if not np.isclose(total, 1.0, atol=1e-6): raise ValueError(f"train/val/test ratios must sum to 1.0, got {total}") indices = np.random.RandomState(seed).permutation(len(dataset)) n_train = int(len(dataset) * train_ratio) n_val = int(len(dataset) * val_ratio) return ( Subset(dataset, indices[:n_train]), Subset(dataset, indices[n_train:n_train + n_val]), Subset(dataset, indices[n_train + n_val:]), ) def split_dataset_by_eid_files( dataset: HealthDataset, train_eid_file: str | Path, val_eid_file: str | Path, test_eid_file: str | Path, ) -> Tuple[Subset, Subset, Subset]: split_sets = { "train": load_eid_file(train_eid_file), "val": load_eid_file(val_eid_file), "test": load_eid_file(test_eid_file), } overlaps = ( split_sets["train"] & split_sets["val"], split_sets["train"] & split_sets["test"], split_sets["val"] & split_sets["test"], ) if any(overlaps): raise ValueError("eid split files must be disjoint") split_indices: Dict[str, list[int]] = {"train": [], "val": [], "test": []} for idx, sample in enumerate(dataset.samples): eid = int(sample["eid"]) for split_name, eid_set in split_sets.items(): if eid in eid_set: split_indices[split_name].append(idx) break missing = [name for name, indices in split_indices.items() if not indices] if missing: raise ValueError(f"Empty dataset split(s) after eid filtering: {missing}") return ( Subset(dataset, np.asarray(split_indices["train"], dtype=np.int64)), Subset(dataset, np.asarray(split_indices["val"], dtype=np.int64)), Subset(dataset, np.asarray(split_indices["test"], dtype=np.int64)), ) def split_all_future_datasets( train_dataset: AllFutureHealthDataset, val_dataset: AllFutureHealthDataset, test_dataset: AllFutureHealthDataset, train_ratio: float, val_ratio: float, test_ratio: float, seed: int, ) -> Tuple[Subset, Subset, Subset]: total = train_ratio + val_ratio + test_ratio if not np.isclose(total, 1.0, atol=1e-6): raise ValueError(f"train/val/test ratios must sum to 1.0, got {total}") patient_indices = np.random.RandomState(seed).permutation(len(train_dataset.patients)) n_train = int(len(patient_indices) * train_ratio) n_val = int(len(patient_indices) * val_ratio) train_patient_idx = patient_indices[:n_train] val_patient_set = set(int(x) for x in patient_indices[n_train:n_train + n_val]) test_patient_set = set(int(x) for x in patient_indices[n_train + n_val:]) val_query_idx = [ i for i, (pidx, _t_query) in enumerate(val_dataset.valid_queries) if int(pidx) in val_patient_set ] test_query_idx = [ i for i, (pidx, _t_query) in enumerate(test_dataset.valid_queries) if int(pidx) in test_patient_set ] if not val_query_idx: raise ValueError("All-future validation split has no valid query samples.") if not test_query_idx: raise ValueError("All-future test split has no valid query samples.") return ( Subset(train_dataset, train_patient_idx), Subset(val_dataset, np.asarray(val_query_idx, dtype=np.int64)), Subset(test_dataset, np.asarray(test_query_idx, dtype=np.int64)), ) def split_all_future_datasets_by_eid_files( train_dataset: AllFutureHealthDataset, val_dataset: AllFutureHealthDataset, test_dataset: AllFutureHealthDataset, train_eid_file: str | Path, val_eid_file: str | Path, test_eid_file: str | Path, ) -> Tuple[Subset, Subset, Subset]: split_sets = { "train": load_eid_file(train_eid_file), "val": load_eid_file(val_eid_file), "test": load_eid_file(test_eid_file), } overlaps = ( split_sets["train"] & split_sets["val"], split_sets["train"] & split_sets["test"], split_sets["val"] & split_sets["test"], ) if any(overlaps): raise ValueError("eid split files must be disjoint") train_patient_idx = [ idx for idx, patient in enumerate(train_dataset.patients) if int(patient["eid"]) in split_sets["train"] ] val_query_idx = [ idx for idx, (pidx, _t_query) in enumerate(val_dataset.valid_queries) if int(val_dataset.patients[int(pidx)]["eid"]) in split_sets["val"] ] test_query_idx = [ idx for idx, (pidx, _t_query) in enumerate(test_dataset.valid_queries) if int(test_dataset.patients[int(pidx)]["eid"]) in split_sets["test"] ] if not train_patient_idx: raise ValueError("All-future training eid split has no patients.") if not val_query_idx: raise ValueError("All-future validation eid split has no valid query samples.") if not test_query_idx: raise ValueError("All-future test eid split has no valid query samples.") return ( Subset(train_dataset, np.asarray(train_patient_idx, dtype=np.int64)), Subset(val_dataset, np.asarray(val_query_idx, dtype=np.int64)), Subset(test_dataset, np.asarray(test_query_idx, dtype=np.int64)), ) def get_model_parameter_counts(model: torch.nn.Module) -> Dict[str, int]: """Return stable total and trainable parameter counts.""" return { "model_parameter_count": sum( parameter.numel() for parameter in model.parameters() ), "trainable_parameter_count": sum( parameter.numel() for parameter in model.parameters() if parameter.requires_grad ), } def set_optimizer_lr(optimizer: AdamW, lr: float) -> None: for param_group in optimizer.param_groups: param_group["lr"] = lr def save_checkpoint(model: DeepHealth, checkpoint_path: Path) -> None: torch.save(model.state_dict(), checkpoint_path) def save_config( args: Any, config_path: Path, extra: Dict[str, Any] | None = None, ) -> None: config: Dict[str, Any] = {} for key, value in vars(args).items(): if isinstance(value, tuple): config[key] = list(value) elif isinstance(value, list): config[key] = value elif isinstance(value, (int, float, str, bool, type(None))): config[key] = value else: config[key] = str(value) if extra: config.update(extra) config_path.write_text(json.dumps(config, indent=2), encoding="utf-8")