refactor: isolate Delphi2M next-token pipeline

This commit is contained in:
2026-07-25 14:22:36 +08:00
parent 15ace878f4
commit 315f552301
17 changed files with 330 additions and 1817 deletions

View File

@@ -1,10 +1,4 @@
"""
Train DeepHealth with next-token / next-time-point supervision.
The next-step dataset uses observed event histories, including CHECKUP state
tokens, plus optional gap <NO_EVENT> imputation. UTS training reads out only
same-time group ends.
"""
"""Reproduce Delphi2M with absolute-time next-token supervision."""
from __future__ import annotations
import argparse
@@ -29,14 +23,15 @@ from model_architectures import (
SUPPORTED_MODEL_ARCHITECTURES,
)
from models import DeepHealth, DeepHealthOutput
from readouts import build_readout
from targets import CHECKUP_IDX, NO_EVENT_IDX, PAD_IDX
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,
@@ -70,7 +65,6 @@ def parse_args() -> argparse.Namespace:
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("--include_no_event_in_uts_target", action="store_true")
parser.add_argument("--train_ratio", type=float, default=0.7)
parser.add_argument("--val_ratio", type=float, default=0.15)
@@ -85,8 +79,6 @@ def parse_args() -> argparse.Namespace:
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("--time_mode", type=str, default="relative",
choices=["relative", "absolute"])
parser.add_argument("--dropout", type=float, default=0.0)
parser.add_argument(
"--model_architecture",
@@ -95,17 +87,10 @@ def parse_args() -> argparse.Namespace:
choices=SUPPORTED_MODEL_ARCHITECTURES,
)
parser.add_argument("--target_mode", type=str, default="uts",
choices=["delphi2m", "uts"])
parser.add_argument("--readout_name", type=str, default=None,
choices=["token", "same_time_group_end", "last_valid"])
parser.add_argument("--readout_reduce", type=str, default="mean",
choices=["mean", "sum"])
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("--ignore_no_event_in_delphi2m", action="store_true")
parser.add_argument("--batch_size", type=int, default=128)
parser.add_argument("--base_lr", type=float, default=3e-4)
@@ -127,11 +112,6 @@ def parse_args() -> argparse.Namespace:
)
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")
if args.target_mode == "uts":
args.readout_name = args.readout_name or "same_time_group_end"
args.include_no_event_in_uts_target = True
else:
args.readout_name = args.readout_name or "token"
args.extra_info_types = (
load_extra_info_types_file(args.extra_info_types_file)
if args.extra_info_types_file is not None
@@ -140,24 +120,6 @@ def parse_args() -> argparse.Namespace:
return args
def get_lr(epoch: int, args: argparse.Namespace, 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 build_model(args: argparse.Namespace, dataset: HealthDataset) -> DeepHealth:
return DeepHealth(
vocab_size=dataset.vocab_size,
@@ -171,44 +133,27 @@ def build_model(args: argparse.Namespace, dataset: HealthDataset) -> DeepHealth:
n_bins=args.n_bins,
extra_pool_reduce=args.extra_pool_reduce,
target_mode="next_token",
time_mode=args.time_mode,
time_mode="absolute",
dist_mode="exponential",
dropout=args.dropout,
model_architecture=args.model_architecture,
)
def build_next_step_readout(args: argparse.Namespace):
if args.readout_name == "same_time_group_end":
return build_readout("same_time_group_end", reduce=args.readout_reduce)
return build_readout(args.readout_name)
def build_next_step_loss(args: argparse.Namespace):
if args.target_mode == "delphi2m":
ignored_tokens = {PAD_IDX, CHECKUP_IDX}
if args.ignore_no_event_in_delphi2m:
ignored_tokens.add(NO_EVENT_IDX)
return build_loss(
"delphi2m",
ignored_tokens=ignored_tokens,
t_min=args.t_min,
max_exp_input=args.max_exp_input,
ce_weight=args.ce_weight,
time_weight=args.time_weight,
)
return build_loss(
"uts",
ignored_idx={PAD_IDX, CHECKUP_IDX},
"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,
include_uts_targets: bool,
) -> Dict[str, torch.Tensor]:
hidden_len = model_out.hidden.size(1)
event_len = int(model_out.event_len)
@@ -216,26 +161,12 @@ def build_augmented_next_step_targets(
device = model_out.hidden.device
non_blocking = device.type == "cuda"
if extra_len <= 0:
targets = {
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),
"readout_mask": batch_cpu["readout_mask"].to(device, non_blocking=non_blocking),
}
if include_uts_targets:
targets["target_dt_unique"] = batch_cpu["target_dt_unique"].to(
device, non_blocking=non_blocking
)
targets["target_multi_hot"] = batch_cpu["target_multi_hot"].to(
device, non_blocking=non_blocking
)
return targets
bsz = batch_cpu["target_event_seq"].size(0)
vocab_size = (
batch_cpu["target_multi_hot"].size(2)
if include_uts_targets
else None
)
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)
@@ -268,34 +199,6 @@ def build_augmented_next_step_targets(
],
dim=1,
)
readout_mask = torch.cat([batch_cpu["readout_mask"], extra_mask], dim=1)
target_dt_unique = None
target_multi_hot = None
if include_uts_targets:
target_dt_unique = torch.cat(
[
batch_cpu["target_dt_unique"],
torch.zeros(
bsz,
extra_len,
dtype=batch_cpu["target_dt_unique"].dtype,
),
],
dim=1,
)
target_multi_hot = torch.cat(
[
batch_cpu["target_multi_hot"],
torch.zeros(
bsz,
extra_len,
vocab_size,
dtype=batch_cpu["target_multi_hot"].dtype,
),
],
dim=1,
)
for b in range(bsz):
valid_event = batch_cpu["padding_mask"][b].bool()
if not valid_event.any():
@@ -326,7 +229,6 @@ def build_augmented_next_step_targets(
t = extra_time[b, j]
future = times > t
if not future.any():
readout_mask[b, pos] = False
continue
first_idx = int(torch.nonzero(future, as_tuple=False)[0].item())
@@ -335,35 +237,15 @@ def build_augmented_next_step_targets(
target_event_seq[b, pos] = next_event
target_time_seq[b, pos] = next_time
if not include_uts_targets:
continue
same_next_time = times == next_time
next_events = events[same_next_time]
valid_next_events = next_events[
(next_events > PAD_IDX) & (next_events < vocab_size)
].long()
if valid_next_events.numel() == 0:
readout_mask[b, pos] = False
continue
target_multi_hot[b, pos, valid_next_events] = True
target_dt_unique[b, pos] = next_time - t
targets = {
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),
"readout_mask": readout_mask.to(device, non_blocking=non_blocking),
}
if include_uts_targets:
targets["target_dt_unique"] = target_dt_unique.to(device, non_blocking=non_blocking)
targets["target_multi_hot"] = target_multi_hot.to(device, non_blocking=non_blocking)
return targets
def compute_next_step_loss(
args: argparse.Namespace,
model: DeepHealth,
readout,
criterion,
batch: Dict[str, torch.Tensor],
device: torch.device,
@@ -382,7 +264,6 @@ def compute_next_step_loss(
other_value=batch["other_value"],
other_value_kind=batch["other_value_kind"],
other_time=batch["other_time"],
target_mode="next_token",
return_output=True,
)
if not isinstance(model_out, DeepHealthOutput):
@@ -390,35 +271,17 @@ def compute_next_step_loss(
targets = build_augmented_next_step_targets(
batch_cpu=batch_cpu,
model_out=model_out,
include_uts_targets=args.target_mode == "uts",
)
readout_out = readout(
hidden=model_out.hidden,
time_seq=model_out.time_seq,
padding_mask=model_out.padding_mask,
readout_mask=targets["readout_mask"]
if args.readout_name == "same_time_group_end"
else None,
)
logits = model.calc_risk(readout_out.hidden)
logits = model.calc_risk(model_out.hidden)
if args.target_mode == "delphi2m":
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=readout_out.readout_mask,
return_components=True,
)
else:
loss, parts = criterion(
logits=logits,
target_multi_hot=targets["target_multi_hot"],
target_dt_unique=targets["target_dt_unique"],
readout_mask=readout_out.readout_mask,
return_components=True,
)
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
@@ -428,7 +291,6 @@ def run_epoch(
logger: logging.Logger,
args: argparse.Namespace,
model: DeepHealth,
readout,
criterion,
loader: DataLoader,
optimizer: AdamW | None,
@@ -436,7 +298,6 @@ def run_epoch(
is_train: bool,
) -> float:
model.train(is_train)
readout.train(is_train)
total = torch.zeros((), device=device)
n_batches = 0
skipped = 0
@@ -447,7 +308,9 @@ def run_epoch(
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, readout, criterion, batch, device)
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")
@@ -497,7 +360,8 @@ def build_metadata(
"model_class": "DeepHealth",
"model_architecture": args.model_architecture,
"model_target_mode": "next_token",
"target_mode": args.target_mode,
"target_mode": "delphi2m",
"time_mode": "absolute",
"dist_mode": "exponential",
"extra_info_types_file": (
Path(args.extra_info_types_file).name
@@ -518,8 +382,8 @@ def build_metadata(
"val": int(len(val_subset)),
"test": int(len(test_subset)),
},
"resolved_readout_name": args.readout_name,
"resolved_loss_name": args.target_mode,
"resolved_readout_name": "token",
"resolved_loss_name": "delphi2m",
}
@@ -531,7 +395,7 @@ def main() -> None:
run_dir, run_name = create_unique_run_dir(
lambda timestamp: (
f"{args.time_mode}_exponential_next_token_{args.target_mode}_"
"absolute_exponential_next_token_delphi2m_"
f"gap_{args.no_event_interval_years:g}y_{timestamp}"
),
runs_root=Path(args.runs_root) / args.model_architecture,
@@ -542,13 +406,12 @@ def main() -> None:
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(f"readout={args.readout_name}, target_mode={args.target_mode}")
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,
include_no_event_in_uts_target=args.include_no_event_in_uts_target,
extra_info_types=args.extra_info_types,
)
if args.train_eid_file and args.val_eid_file and args.test_eid_file:
@@ -616,7 +479,6 @@ def main() -> None:
f"total={parameter_counts['model_parameter_count']:,}, "
f"trainable={parameter_counts['trainable_parameter_count']:,}"
)
readout = build_next_step_readout(args).to(device)
criterion = build_next_step_loss(args)
optimizer = AdamW(
model.parameters(),
@@ -646,9 +508,13 @@ def main() -> None:
lr = get_lr(epoch, args, adaptive_lr)
set_optimizer_lr(optimizer, lr)
train_loss = run_epoch(logger, args, model, readout, criterion, train_loader, optimizer, device, True)
train_loss = run_epoch(
logger, args, model, criterion, train_loader, optimizer, device, True
)
with torch.no_grad():
val_loss = run_epoch(logger, args, model, readout, criterion, val_loader, None, device, False)
val_loss = run_epoch(
logger, args, model, criterion, val_loader, None, device, False
)
is_best = val_loss < best_val
if is_best:
@@ -682,7 +548,9 @@ def main() -> None:
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, readout, criterion, test_loader, None, device, False)
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}")