Add train-split robust scaling for continuous values

This commit is contained in:
2026-08-01 11:37:21 +08:00
parent 89dcf4b362
commit dfb22adf2d
6 changed files with 419 additions and 5 deletions

View File

@@ -39,9 +39,11 @@ from model_architectures import (
from models import DeepHealth
from targets import CHECKUP_IDX, PAD_IDX
from train_util import (
ContinuousRobustScalerStats,
configure_torch_for_training,
create_unique_run_dir,
format_extra_info_types,
fit_continuous_robust_scaler,
get_lr,
get_model_parameter_counts,
load_extra_info_types_file,
@@ -104,6 +106,16 @@ 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",
@@ -173,7 +185,22 @@ def parse_args() -> argparse.Namespace:
return args
def build_model(args: argparse.Namespace, dataset: AllFutureHealthDataset) -> DeepHealth:
def build_model(
args: argparse.Namespace,
dataset: AllFutureHealthDataset,
scaler_stats: ContinuousRobustScalerStats | None = None,
) -> DeepHealth:
if (
args.continuous_value_scaling == "robust"
and dataset.n_cont_types > 0
and scaler_stats is None
):
raise ValueError(
"Robust continuous-value scaling requires statistics fitted on the "
"training subset"
)
center = None if scaler_stats is None else scaler_stats.center
scale = None if scaler_stats is None else scaler_stats.scale
return DeepHealth(
vocab_size=dataset.vocab_size,
n_embd=args.n_embd,
@@ -184,6 +211,9 @@ def build_model(args: argparse.Namespace, dataset: AllFutureHealthDataset) -> De
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,
target_mode="all_future",
time_mode=args.time_mode,
@@ -322,7 +352,16 @@ def build_metadata(
train_subset,
val_subset,
test_subset,
scaler_stats: ContinuousRobustScalerStats | None,
) -> 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()
return {
"run_name": run_name,
"dataset_class": "AllFutureHealthDataset",
@@ -342,6 +381,7 @@ def build_metadata(
else None
),
"extra_info_types": [int(x) for x in dataset.extra_info_types],
"continuous_value_scaler": scaler_metadata,
"dataset_metadata": {
"vocab_size": int(dataset.vocab_size),
"n_types": int(dataset.n_types),
@@ -385,6 +425,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("Loading all-future datasets...")
train_dataset = AllFutureHealthDataset(
@@ -448,6 +489,23 @@ 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:
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,
)
logger.info(
"Continuous RobustScaler fitted: "
f"observations={int(scaler_stats.observation_count.sum()):,}, "
f"min_per_feature={int(scaler_stats.observation_count.min()):,}, "
f"max_per_feature={int(scaler_stats.observation_count.max()):,}"
)
train_loader = DataLoader(
train_subset,
batch_size=args.batch_size,
@@ -479,7 +537,7 @@ def main() -> None:
prefetch_factor=2 if args.num_workers > 0 else None,
)
model = build_model(args, train_dataset).to(device)
model = build_model(args, train_dataset, scaler_stats=scaler_stats).to(device)
parameter_counts = get_model_parameter_counts(model)
logger.info(
"Model parameters: "
@@ -496,7 +554,13 @@ def main() -> None:
adaptive_lr = args.base_lr * math.sqrt(args.batch_size / 128)
train_metadata = build_metadata(
args, train_dataset, run_name, train_subset, val_subset, test_subset
args,
train_dataset,
run_name,
train_subset,
val_subset,
test_subset,
scaler_stats,
)
train_metadata.update(parameter_counts)
save_config(