refactor: isolate Delphi2M next-token pipeline
This commit is contained in:
@@ -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}")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user