Remove legacy event and mixed distribution paths

This commit is contained in:
2026-08-01 14:23:18 +08:00
parent dfb22adf2d
commit de6f9b75b9
22 changed files with 370 additions and 463 deletions

View File

@@ -37,7 +37,7 @@ from model_architectures import (
SUPPORTED_MODEL_ARCHITECTURES,
)
from models import DeepHealth
from targets import CHECKUP_IDX, PAD_IDX
from targets import PAD_IDX, RESERVED_IDX
from train_util import (
ContinuousRobustScalerStats,
configure_torch_for_training,
@@ -106,22 +106,12 @@ def parse_args() -> argparse.Namespace:
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(
"--continuous_value_scaling",
type=str,
default="robust",
choices=["none", "robust"],
help=(
"Continuous extra-info scaling. 'robust' fits the median and IQR "
"on the complete training subset and stores them in the checkpoint."
),
)
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("--dist_mode", type=str, default="exponential",
choices=["exponential", "weibull", "mixed"])
choices=["exponential", "weibull"])
parser.add_argument("--dropout", type=float, default=0.0)
parser.add_argument(
"--model_architecture",
@@ -188,19 +178,14 @@ def parse_args() -> argparse.Namespace:
def build_model(
args: argparse.Namespace,
dataset: AllFutureHealthDataset,
scaler_stats: ContinuousRobustScalerStats | None = None,
scaler_stats: ContinuousRobustScalerStats,
) -> DeepHealth:
if (
args.continuous_value_scaling == "robust"
and dataset.n_cont_types > 0
and scaler_stats is None
):
if tuple(int(x) for x in dataset.cont_type_ids) != scaler_stats.cont_type_ids:
raise ValueError(
"Robust continuous-value scaling requires statistics fitted on the "
"training subset"
"RobustScale statistics are not aligned with dataset.cont_type_ids"
)
center = None if scaler_stats is None else scaler_stats.center
scale = None if scaler_stats is None else scaler_stats.scale
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,
@@ -211,7 +196,6 @@ def build_model(
n_categories=dataset.n_categories,
cont_type_ids=dataset.cont_type_ids,
n_bins=args.n_bins,
continuous_value_scaling=args.continuous_value_scaling,
continuous_value_center=center,
continuous_value_scale=scale,
extra_pool_reduce=args.extra_pool_reduce,
@@ -223,18 +207,12 @@ def build_model(
)
def build_criterion(args: argparse.Namespace, dataset: AllFutureHealthDataset):
ignored_idx = {PAD_IDX, CHECKUP_IDX}
def build_criterion(args: argparse.Namespace):
ignored_idx = {PAD_IDX, RESERVED_IDX}
if args.dist_mode == "exponential":
return build_loss("exponential", ignored_idx=ignored_idx)
if args.dist_mode == "weibull":
return build_loss("weibull", ignored_idx=ignored_idx)
if args.dist_mode == "mixed":
return build_loss(
"mixed",
death_idx=dataset.vocab_size - 1,
ignored_idx=ignored_idx,
)
raise ValueError(f"Unknown dist_mode: {args.dist_mode}")
@@ -247,7 +225,7 @@ def compute_all_future_loss(
) -> torch.Tensor:
required_keys = set(MODEL_INPUT_KEYS)
required_keys.update(("future_targets", "exposure"))
if args.dist_mode in {"weibull", "mixed"}:
if args.dist_mode == "weibull":
required_keys.add("future_dt")
batch = move_batch_to_device(
{key: batch[key] for key in required_keys},
@@ -282,13 +260,7 @@ def compute_all_future_loss(
exposure=batch["exposure"],
)
else:
loss = criterion(
logits=logits,
death_rho=model.calc_death_rho(hidden),
targets=batch["future_targets"],
dt=batch["future_dt"],
exposure=batch["exposure"],
)
raise ValueError(f"Unknown dist_mode: {args.dist_mode}")
if not torch.isfinite(loss):
raise RuntimeError(f"Loss is not finite: {float(loss.detach().cpu())}")
@@ -352,16 +324,9 @@ def build_metadata(
train_subset,
val_subset,
test_subset,
scaler_stats: ContinuousRobustScalerStats | None,
scaler_stats: ContinuousRobustScalerStats,
) -> Dict[str, Any]:
scaler_metadata: Dict[str, Any]
if scaler_stats is None:
scaler_metadata = {
"method": "none",
"fitted_on": None,
}
else:
scaler_metadata = scaler_stats.as_metadata()
scaler_metadata = scaler_stats.as_metadata()
return {
"run_name": run_name,
"dataset_class": "AllFutureHealthDataset",
@@ -370,6 +335,8 @@ def build_metadata(
"model_architecture": args.model_architecture,
"model_target_mode": "all_future",
"target_mode": "all_future",
"event_stream_version": "disease_death_only_v1",
"uses_assessment_event_token": False,
"dist_mode": args.dist_mode,
"disease_history_mode": args.disease_history_mode,
"all_future_min_history_events": int(args.min_history_events),
@@ -381,6 +348,7 @@ def build_metadata(
else None
),
"extra_info_types": [int(x) for x in dataset.extra_info_types],
"continuous_value_scaling": "robust",
"continuous_value_scaler": scaler_metadata,
"dataset_metadata": {
"vocab_size": int(dataset.vocab_size),
@@ -389,6 +357,8 @@ def build_metadata(
"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)),
@@ -425,7 +395,7 @@ def main() -> None:
logger.info(f"Model architecture: {args.model_architecture}")
logger.info(f"Disease history mode: {args.disease_history_mode}")
logger.info(f"extra_info_types: {format_extra_info_types(args.extra_info_types)}")
logger.info(f"Continuous value scaling: {args.continuous_value_scaling}")
logger.info("Continuous value scaling: RobustScale (required)")
logger.info("Loading all-future datasets...")
train_dataset = AllFutureHealthDataset(
@@ -489,16 +459,16 @@ def main() -> None:
f"Patients/queries: train={len(train_subset)}, val={len(val_subset)}, test={len(test_subset)}"
)
scaler_stats = None
if args.continuous_value_scaling == "robust" and train_dataset.n_cont_types > 0:
if train_dataset.n_cont_types > 0:
logger.info(
"Fitting continuous RobustScaler on the complete training subset: "
f"patients={len(train_subset):,}, features={train_dataset.n_cont_types}"
)
scaler_stats = fit_continuous_robust_scaler(
train_dataset,
train_subset,
)
scaler_stats = fit_continuous_robust_scaler(
train_dataset,
train_subset,
)
if train_dataset.n_cont_types > 0:
logger.info(
"Continuous RobustScaler fitted: "
f"observations={int(scaler_stats.observation_count.sum()):,}, "
@@ -550,7 +520,7 @@ def main() -> None:
betas=tuple(args.betas),
weight_decay=args.weight_decay,
)
criterion = build_criterion(args, train_dataset)
criterion = build_criterion(args)
adaptive_lr = args.base_lr * math.sqrt(args.batch_size / 128)
train_metadata = build_metadata(