Improve all-future first-onset training
This commit is contained in:
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user