Improve all-future first-onset training

This commit is contained in:
2026-08-21 13:48:59 +08:00
parent 9ccc6b56ec
commit d728d8585c
9 changed files with 1002 additions and 84 deletions

View File

@@ -5,8 +5,8 @@ Training samples are patient-level. For each patient and each __getitem__ call,
AllFutureHealthDataset randomly samples a query time t_query, uses events at or
before t_query as history, and uses events after t_query as the future target set.
Validation/test samples are deterministic query points built from future event
times, then split by patient.
All splits use the same patient/interval/time-uniform query distribution.
Validation/test keep one deterministic query draw per patient.
"""
from __future__ import annotations
@@ -37,13 +37,15 @@ from model_architectures import (
SUPPORTED_MODEL_ARCHITECTURES,
)
from models import DeepHealth
from targets import PAD_IDX, RESERVED_IDX
from targets import NO_EVENT_IDX, PAD_IDX, RESERVED_IDX
from train_util import (
AllFutureBaselineStats,
ContinuousRobustScalerStats,
configure_torch_for_training,
create_unique_run_dir,
format_extra_info_types,
fit_continuous_robust_scaler,
fit_all_future_baseline,
get_lr,
get_model_parameter_counts,
load_extra_info_types_file,
@@ -113,6 +115,15 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--dist_mode", type=str, default="exponential",
choices=["exponential", "weibull"])
parser.add_argument("--dropout", type=float, default=0.0)
parser.add_argument(
"--risk_head_bias",
action=argparse.BooleanOptionalAction,
default=True,
help=(
"Initialize a learnable all-future output bias from training-only "
"marginal first-onset rates"
),
)
parser.add_argument(
"--model_architecture",
type=str,
@@ -179,6 +190,7 @@ def build_model(
args: argparse.Namespace,
dataset: AllFutureHealthDataset,
scaler_stats: ContinuousRobustScalerStats,
baseline_stats: AllFutureBaselineStats | None,
) -> DeepHealth:
if tuple(int(x) for x in dataset.cont_type_ids) != scaler_stats.cont_type_ids:
raise ValueError(
@@ -186,6 +198,8 @@ def build_model(
)
center = scaler_stats.center if dataset.n_cont_types > 0 else None
scale = scaler_stats.scale if dataset.n_cont_types > 0 else None
if args.risk_head_bias and baseline_stats is None:
raise ValueError("baseline_stats is required when risk_head_bias is enabled")
return DeepHealth(
vocab_size=dataset.vocab_size,
n_embd=args.n_embd,
@@ -204,11 +218,17 @@ def build_model(
dist_mode=args.dist_mode,
dropout=args.dropout,
model_architecture=args.model_architecture,
risk_head_bias=args.risk_head_bias,
risk_head_bias_init=(
baseline_stats.bias
if args.risk_head_bias and baseline_stats is not None
else None
),
)
def build_criterion(args: argparse.Namespace):
ignored_idx = {PAD_IDX, RESERVED_IDX}
ignored_idx = {PAD_IDX, RESERVED_IDX, NO_EVENT_IDX}
if args.dist_mode == "exponential":
return build_loss("exponential", ignored_idx=ignored_idx)
if args.dist_mode == "weibull":
@@ -224,9 +244,7 @@ def compute_all_future_loss(
device: torch.device,
) -> torch.Tensor:
required_keys = set(MODEL_INPUT_KEYS)
required_keys.update(("future_targets", "exposure"))
if args.dist_mode == "weibull":
required_keys.add("future_dt")
required_keys.update(("future_targets", "future_dt", "exposure"))
batch = move_batch_to_device(
{key: batch[key] for key in required_keys},
device,
@@ -250,6 +268,8 @@ def compute_all_future_loss(
logits=logits,
targets=batch["future_targets"],
exposure=batch["exposure"],
dt=batch["future_dt"],
history=batch["event_seq"],
)
elif args.dist_mode == "weibull":
loss = criterion(
@@ -258,6 +278,7 @@ def compute_all_future_loss(
targets=batch["future_targets"],
dt=batch["future_dt"],
exposure=batch["exposure"],
history=batch["event_seq"],
)
else:
raise ValueError(f"Unknown dist_mode: {args.dist_mode}")
@@ -325,6 +346,7 @@ def build_metadata(
val_subset,
test_subset,
scaler_stats: ContinuousRobustScalerStats,
baseline_stats: AllFutureBaselineStats | None,
) -> Dict[str, Any]:
scaler_metadata = scaler_stats.as_metadata()
return {
@@ -342,6 +364,17 @@ def build_metadata(
"all_future_min_history_events": int(args.min_history_events),
"all_future_min_future_events": int(args.min_future_events),
"all_future_validation_query_seed": int(args.validation_query_seed),
"all_future_query_distribution": "patient_interval_time_uniform",
"all_future_queries_per_validation_patient": 1,
"all_future_likelihood": "first_onset_survival_v2",
"all_future_prevalent_outcomes_excluded": True,
"risk_head_bias": bool(args.risk_head_bias),
"risk_head_bias_weight_decay": 0.0 if args.risk_head_bias else None,
"risk_head_baseline": (
baseline_stats.as_metadata()
if args.risk_head_bias and baseline_stats is not None
else {"method": "none"}
),
"extra_info_types_file": (
Path(args.extra_info_types_file).name
if args.extra_info_types_file is not None
@@ -476,6 +509,27 @@ def main() -> None:
f"max_per_feature={int(scaler_stats.observation_count.max()):,}"
)
baseline_stats = None
if args.risk_head_bias:
logger.info(
"Fitting all-future output baseline on one seeded query per "
f"training patient: patients={len(train_subset):,}"
)
baseline_stats = fit_all_future_baseline(
train_dataset,
train_subset,
seed=args.seed,
)
baseline_summary = baseline_stats.as_metadata()
rate_summary = baseline_summary["rate_per_year"]
logger.info(
"All-future baseline rates/year: "
f"min={rate_summary['min']:.8g}, "
f"median={rate_summary['median']:.8g}, "
f"max={rate_summary['max']:.8g}, "
f"first_onsets={baseline_summary['observed_first_onsets']:,}"
)
train_loader = DataLoader(
train_subset,
batch_size=args.batch_size,
@@ -507,15 +561,37 @@ def main() -> None:
prefetch_factor=2 if args.num_workers > 0 else None,
)
model = build_model(args, train_dataset, scaler_stats=scaler_stats).to(device)
model = build_model(
args,
train_dataset,
scaler_stats=scaler_stats,
baseline_stats=baseline_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']:,}"
)
if model.risk_head.bias is None:
optimizer_parameters = model.parameters()
else:
optimizer_parameters = [
{
"params": [
parameter
for parameter in model.parameters()
if parameter is not model.risk_head.bias
],
"weight_decay": args.weight_decay,
},
{
"params": [model.risk_head.bias],
"weight_decay": 0.0,
},
]
optimizer = AdamW(
model.parameters(),
optimizer_parameters,
lr=args.base_lr,
betas=tuple(args.betas),
weight_decay=args.weight_decay,
@@ -531,6 +607,7 @@ def main() -> None:
val_subset,
test_subset,
scaler_stats,
baseline_stats,
)
train_metadata.update(parameter_counts)
save_config(