"""Reproduce Delphi2M with absolute-time next-token supervision.""" from __future__ import annotations import argparse import json import logging import math import time from pathlib import Path from typing import Any, Dict import numpy as np import torch from torch.nn.utils import clip_grad_norm_ from torch.optim import AdamW from torch.utils.data import DataLoader, RandomSampler from tqdm.auto import tqdm from dataset import HealthDataset, collate_fn from losses import build_loss from model_architectures import ( DEFAULT_MODEL_ARCHITECTURE, SUPPORTED_MODEL_ARCHITECTURES, ) from models import DeepHealth, DeepHealthOutput from targets import CHECKUP_IDX, PAD_IDX from train_util import ( configure_torch_for_training, create_unique_run_dir, format_extra_info_types, get_lr, get_model_parameter_counts, load_extra_info_types_file, move_batch_to_device, resolve_device, save_checkpoint, save_config, set_optimizer_lr, set_seed, setup_logging, split_dataset, split_dataset_by_eid_files, ) MODEL_INPUT_KEYS = ( "event_seq", "time_seq", "sex", "padding_mask", "other_type", "other_value", "other_value_kind", "other_time", ) def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser( description="Train DeepHealth with next-token/point supervision") parser.add_argument("--data_prefix", type=str, default="ukb") parser.add_argument("--labels_file", type=str, default="labels.csv") parser.add_argument("--runs_root", type=str, default="runs") parser.add_argument("--seed", type=int, default=42) parser.add_argument("--extra_info_types_file", type=str, default=None) parser.add_argument("--no_event_interval_years", type=float, default=5.0) parser.add_argument("--train_ratio", type=float, default=0.7) parser.add_argument("--val_ratio", type=float, default=0.15) parser.add_argument("--test_ratio", type=float, default=0.15) parser.add_argument("--train_eid_file", type=str, default="ukb_train_eid.csv") parser.add_argument("--val_eid_file", type=str, default="ukb_val_eid.csv") parser.add_argument("--test_eid_file", type=str, default="ukb_test_eid.csv") parser.add_argument("--n_embd", type=int, default=120) parser.add_argument("--n_head", type=int, default=10) parser.add_argument("--n_layer", type=int, default=12) parser.add_argument("--n_bins", type=int, default=16) parser.add_argument("--extra_pool_reduce", type=str, default="mean", choices=["mean", "sum"]) parser.add_argument("--dropout", type=float, default=0.0) parser.add_argument( "--model_architecture", type=str, default=DEFAULT_MODEL_ARCHITECTURE, choices=SUPPORTED_MODEL_ARCHITECTURES, ) parser.add_argument("--t_min", type=float, default=0.0027378507871321013) parser.add_argument("--max_exp_input", type=float, default=60.0) parser.add_argument("--ce_weight", type=float, default=1.0) parser.add_argument("--time_weight", type=float, default=1.0) parser.add_argument("--batch_size", type=int, default=128) parser.add_argument("--base_lr", type=float, default=3e-4) parser.add_argument("--weight_decay", type=float, default=0.1) parser.add_argument("--betas", type=float, nargs=2, default=(0.9, 0.99)) parser.add_argument("--grad_clip", type=float, default=1.0) parser.add_argument("--max_epochs", type=int, default=200) parser.add_argument("--warmup_epochs", type=int, default=10) parser.add_argument("--patience", type=int, default=15) parser.add_argument("--min_lr_ratio", type=float, default=0.1) parser.add_argument("--num_workers", type=int, default=4) parser.add_argument("--device", type=str, default="cuda") parser.add_argument("--progress_interval", type=int, default=20) args = parser.parse_args() use_eid_split = all( getattr(args, name) for name in ("train_eid_file", "val_eid_file", "test_eid_file") ) if not use_eid_split and not np.isclose(args.train_ratio + args.val_ratio + args.test_ratio, 1.0): raise ValueError("train_ratio + val_ratio + test_ratio must equal 1.0") args.extra_info_types = ( load_extra_info_types_file(args.extra_info_types_file) if args.extra_info_types_file is not None else None ) return args def build_model(args: argparse.Namespace, dataset: HealthDataset) -> DeepHealth: return DeepHealth( vocab_size=dataset.vocab_size, n_embd=args.n_embd, n_head=args.n_head, n_layer=args.n_layer, n_types=dataset.n_types, n_cont_types=dataset.n_cont_types, n_categories=dataset.n_categories, cont_type_ids=dataset.cont_type_ids, n_bins=args.n_bins, extra_pool_reduce=args.extra_pool_reduce, target_mode="next_token", time_mode="absolute", dist_mode="exponential", dropout=args.dropout, model_architecture=args.model_architecture, ) def build_next_step_loss(args: argparse.Namespace): return build_loss( "delphi2m", ignored_tokens={PAD_IDX, CHECKUP_IDX}, t_min=args.t_min, max_exp_input=args.max_exp_input, ce_weight=args.ce_weight, time_weight=args.time_weight, ) def build_augmented_next_step_targets( batch_cpu: Dict[str, torch.Tensor], model_out: DeepHealthOutput, ) -> Dict[str, torch.Tensor]: hidden_len = model_out.hidden.size(1) event_len = int(model_out.event_len) extra_len = hidden_len - event_len device = model_out.hidden.device non_blocking = device.type == "cuda" if extra_len <= 0: return { "target_event_seq": batch_cpu["target_event_seq"].to(device, non_blocking=non_blocking), "target_time_seq": batch_cpu["target_time_seq"].to(device, non_blocking=non_blocking), } bsz = batch_cpu["target_event_seq"].size(0) other_valid = batch_cpu["other_type"] > 0 extra_time = batch_cpu["other_time"].new_zeros(bsz, extra_len) extra_mask = torch.zeros(bsz, extra_len, dtype=torch.bool) for b in range(bsz): unique_time = torch.unique(batch_cpu["other_time"][b, other_valid[b]], sorted=True) n_time = min(int(unique_time.numel()), extra_len) if n_time > 0: extra_time[b, :n_time] = unique_time[:n_time] extra_mask[b, :n_time] = True target_event_seq = torch.cat( [ batch_cpu["target_event_seq"], torch.full( (bsz, extra_len), PAD_IDX, dtype=batch_cpu["target_event_seq"].dtype, ), ], dim=1, ) target_time_seq = torch.cat( [ batch_cpu["target_time_seq"], torch.zeros( bsz, extra_len, dtype=batch_cpu["target_time_seq"].dtype, ), ], dim=1, ) for b in range(bsz): valid_event = batch_cpu["padding_mask"][b].bool() if not valid_event.any(): continue n_event = int(valid_event.sum().item()) events = torch.cat( [ batch_cpu["event_seq"][b, :n_event], batch_cpu["target_event_seq"][b, n_event - 1:n_event], ] ) times = torch.cat( [ batch_cpu["time_seq"][b, :n_event], batch_cpu["target_time_seq"][b, n_event - 1:n_event], ] ) valid_full = events > PAD_IDX events = events[valid_full] times = times[valid_full] if events.numel() == 0: continue for j in range(extra_len): if not bool(extra_mask[b, j]): continue pos = event_len + j t = extra_time[b, j] future = times > t if not future.any(): continue first_idx = int(torch.nonzero(future, as_tuple=False)[0].item()) next_time = times[first_idx] next_event = events[first_idx] target_event_seq[b, pos] = next_event target_time_seq[b, pos] = next_time return { "target_event_seq": target_event_seq.to(device, non_blocking=non_blocking), "target_time_seq": target_time_seq.to(device, non_blocking=non_blocking), } def compute_next_step_loss( args: argparse.Namespace, model: DeepHealth, criterion, batch: Dict[str, torch.Tensor], device: torch.device, ) -> tuple[torch.Tensor, Dict[str, torch.Tensor]]: batch_cpu = batch batch = move_batch_to_device( {key: batch_cpu[key] for key in MODEL_INPUT_KEYS}, device, ) model_out = model( event_seq=batch["event_seq"], time_seq=batch["time_seq"], sex=batch["sex"], padding_mask=batch["padding_mask"], other_type=batch["other_type"], other_value=batch["other_value"], other_value_kind=batch["other_value_kind"], other_time=batch["other_time"], return_output=True, ) if not isinstance(model_out, DeepHealthOutput): raise TypeError("DeepHealth return_output=True must return DeepHealthOutput") targets = build_augmented_next_step_targets( batch_cpu=batch_cpu, model_out=model_out, ) logits = model.calc_risk(model_out.hidden) loss, parts = criterion( logits=logits, target_events=targets["target_event_seq"], target_times=targets["target_time_seq"], current_times=model_out.time_seq, padding_mask=model_out.padding_mask, return_components=True, ) if not torch.isfinite(loss): raise RuntimeError(f"Loss is not finite: {float(loss.detach().cpu())}") return loss, parts def run_epoch( logger: logging.Logger, args: argparse.Namespace, model: DeepHealth, criterion, loader: DataLoader, optimizer: AdamW | None, device: torch.device, is_train: bool, ) -> float: model.train(is_train) total = torch.zeros((), device=device) n_batches = 0 skipped = 0 parts_sum: Dict[str, torch.Tensor] = {} desc = "train" if is_train else "val" progress_interval = max(1, int(args.progress_interval)) progress = tqdm(loader, desc=desc, leave=False, dynamic_ncols=True) for batch_idx, batch in enumerate(progress): try: loss, parts = compute_next_step_loss( args, model, criterion, batch, device ) if is_train: if optimizer is None: raise ValueError("optimizer is required for training") optimizer.zero_grad(set_to_none=True) loss.backward() if args.grad_clip > 0: clip_grad_norm_(model.parameters(), args.grad_clip) optimizer.step() total = total + loss.detach() n_batches += 1 for name, value in parts.items(): parts_sum[name] = parts_sum.get(name, torch.zeros((), device=device)) + value.detach() if (batch_idx + 1) % progress_interval == 0: avg = total / max(1, n_batches) postfix = { "loss": f"{float(loss.detach().cpu()):.4f}", "avg": f"{float(avg.detach().cpu()):.4f}", "skipped": skipped, } for name, value in parts_sum.items(): postfix[name] = f"{float((value / max(1, n_batches)).detach().cpu()):.4f}" progress.set_postfix(postfix) except RuntimeError as exc: if "Loss is not finite" not in str(exc): raise skipped += 1 logger.warning(f"Batch {batch_idx} skipped: {str(exc)[:120]}") if skipped: logger.info(f"Skipped {skipped} batches due to non-finite loss") return float((total / max(1, n_batches)).detach().cpu()) if n_batches else float("inf") def build_metadata( args: argparse.Namespace, dataset: HealthDataset, run_name: str, train_subset, val_subset, test_subset, ) -> Dict[str, Any]: return { "run_name": run_name, "dataset_class": "NextStepHealthDataset", "collate_fn": "next_step_collate_fn", "model_class": "DeepHealth", "model_architecture": args.model_architecture, "model_target_mode": "next_token", "target_mode": "delphi2m", "time_mode": "absolute", "dist_mode": "exponential", "extra_info_types_file": ( Path(args.extra_info_types_file).name if args.extra_info_types_file is not None else None ), "extra_info_types": [int(x) for x in dataset.extra_info_types], "dataset_metadata": { "vocab_size": int(dataset.vocab_size), "n_types": int(dataset.n_types), "n_cont_types": int(dataset.n_cont_types), "n_categories": int(dataset.n_categories), "cont_type_ids": [int(x) for x in dataset.cont_type_ids], "extra_info_types": [int(x) for x in dataset.extra_info_types], }, "split_sizes": { "train": int(len(train_subset)), "val": int(len(val_subset)), "test": int(len(test_subset)), }, "resolved_readout_name": "token", "resolved_loss_name": "delphi2m", } def main() -> None: args = parse_args() set_seed(args.seed) device = resolve_device(args.device) configure_torch_for_training(device) run_dir, run_name = create_unique_run_dir( lambda timestamp: ( "absolute_exponential_next_token_delphi2m_" f"gap_{args.no_event_interval_years:g}y_{timestamp}" ), runs_root=Path(args.runs_root) / args.model_architecture, ) logger = setup_logging(run_dir) logger.info(f"Starting next-step training run: {run_name}") logger.info(f"Device: {device}") logger.info(f"Model architecture: {args.model_architecture}") logger.info(f"extra_info_types: {format_extra_info_types(args.extra_info_types)}") logger.info("time_mode=absolute, readout=token, target_mode=delphi2m") dataset = HealthDataset( data_prefix=args.data_prefix, labels_file=args.labels_file, no_event_interval_years=args.no_event_interval_years, extra_info_types=args.extra_info_types, ) if args.train_eid_file and args.val_eid_file and args.test_eid_file: train_subset, val_subset, test_subset = split_dataset_by_eid_files( dataset=dataset, train_eid_file=args.train_eid_file, val_eid_file=args.val_eid_file, test_eid_file=args.test_eid_file, ) logger.info( "Using eid split files: " f"train={args.train_eid_file}, val={args.val_eid_file}, test={args.test_eid_file}" ) else: train_subset, val_subset, test_subset = split_dataset( dataset=dataset, train_ratio=args.train_ratio, val_ratio=args.val_ratio, test_ratio=args.test_ratio, seed=args.seed, ) logger.info( f"Using random ratio split: train={args.train_ratio}, " f"val={args.val_ratio}, test={args.test_ratio}, seed={args.seed}" ) logger.info( f"Samples: train={len(train_subset)}, val={len(val_subset)}, test={len(test_subset)}" ) train_loader = DataLoader( train_subset, batch_size=args.batch_size, sampler=RandomSampler(train_subset, generator=torch.Generator().manual_seed(args.seed)), collate_fn=collate_fn, num_workers=args.num_workers, pin_memory=device.type == "cuda", persistent_workers=args.num_workers > 0, prefetch_factor=2 if args.num_workers > 0 else None, ) val_loader = DataLoader( val_subset, batch_size=args.batch_size, shuffle=False, collate_fn=collate_fn, num_workers=args.num_workers, pin_memory=device.type == "cuda", persistent_workers=args.num_workers > 0, prefetch_factor=2 if args.num_workers > 0 else None, ) test_loader = DataLoader( test_subset, batch_size=args.batch_size, shuffle=False, collate_fn=collate_fn, num_workers=args.num_workers, pin_memory=device.type == "cuda", persistent_workers=args.num_workers > 0, prefetch_factor=2 if args.num_workers > 0 else None, ) model = build_model(args, dataset).to(device) parameter_counts = get_model_parameter_counts(model) logger.info( "Model parameters: " f"total={parameter_counts['model_parameter_count']:,}, " f"trainable={parameter_counts['trainable_parameter_count']:,}" ) criterion = build_next_step_loss(args) optimizer = AdamW( model.parameters(), lr=args.base_lr, betas=tuple(args.betas), weight_decay=args.weight_decay, ) adaptive_lr = args.base_lr * math.sqrt(args.batch_size / 128) train_metadata = build_metadata( args, dataset, run_name, train_subset, val_subset, test_subset ) train_metadata.update(parameter_counts) save_config( args, run_dir / "train_config.json", extra=train_metadata, ) best_val = float("inf") patience = 0 history = [] best_model_path = run_dir / "best_model.pt" start = time.time() for epoch in range(args.max_epochs): lr = get_lr(epoch, args, adaptive_lr) set_optimizer_lr(optimizer, lr) train_loss = run_epoch( logger, args, model, criterion, train_loader, optimizer, device, True ) with torch.no_grad(): val_loss = run_epoch( logger, args, model, criterion, val_loader, None, device, False ) is_best = val_loss < best_val if is_best: best_val = val_loss patience = 0 save_checkpoint(model, best_model_path) else: patience += 1 logger.info( f"Epoch {epoch + 1}/{args.max_epochs} | lr={lr:.6f} | " f"train_loss={train_loss:.6f} | val_loss={val_loss:.6f} | " f"best_val_loss={best_val:.6f} | patience={patience}/{args.patience} | " f"elapsed={time.time() - start:.1f}s" ) history.append({ "epoch": epoch + 1, "lr": lr, "train_loss": train_loss, "val_loss": val_loss, "best_val_loss": best_val, "is_best": int(is_best), }) if patience >= args.patience: logger.info(f"Early stopping triggered at epoch {epoch + 1}") break with (run_dir / "history.json").open("w", encoding="utf-8") as f: json.dump(history, f, indent=2) logger.info("Evaluating best model on next-step test split...") model.load_state_dict(torch.load(best_model_path, map_location=device)) with torch.no_grad(): test_loss = run_epoch( logger, args, model, criterion, test_loader, None, device, False ) logger.info(f"Test loss: {test_loss:.6f}") logger.info(f"Best checkpoint: {best_model_path}") if __name__ == "__main__": main()