Files
DeepHealth/train_next_step.py
2026-08-21 13:16:39 +08:00

601 lines
21 KiB
Python

"""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 PAD_IDX, RESERVED_IDX
from train_util import (
ContinuousRobustScalerStats,
configure_torch_for_training,
create_unique_run_dir,
fit_continuous_robust_scaler,
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,
scaler_stats: ContinuousRobustScalerStats,
) -> DeepHealth:
if tuple(int(x) for x in dataset.cont_type_ids) != scaler_stats.cont_type_ids:
raise ValueError(
"RobustScale statistics are not aligned with dataset.cont_type_ids"
)
center = scaler_stats.center if dataset.n_cont_types > 0 else None
scale = scaler_stats.scale if dataset.n_cont_types > 0 else None
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,
continuous_value_center=center,
continuous_value_scale=scale,
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, RESERVED_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(
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(
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,
scaler_stats: ContinuousRobustScalerStats,
) -> 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",
"event_stream_version": "disease_death_only_v1",
"uses_assessment_event_token": False,
"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],
"continuous_value_scaling": "robust",
"continuous_value_scaler": scaler_stats.as_metadata(),
"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],
"event_stream_version": "disease_death_only_v1",
"uses_assessment_event_token": False,
},
"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("Continuous value scaling: RobustScale (required)")
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)}"
)
if dataset.n_cont_types > 0:
logger.info(
"Fitting continuous RobustScaler on the complete training subset: "
f"patients={len(train_subset):,}, features={dataset.n_cont_types}"
)
scaler_stats = fit_continuous_robust_scaler(dataset, train_subset)
if dataset.n_cont_types > 0:
logger.info(
"Continuous RobustScaler fitted: "
f"observations={int(scaler_stats.observation_count.sum()):,}, "
f"min_per_feature={int(scaler_stats.observation_count.min()):,}, "
f"max_per_feature={int(scaler_stats.observation_count.max()):,}"
)
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, scaler_stats).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,
scaler_stats,
)
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()