Remove legacy event and mixed distribution paths
This commit is contained in:
@@ -23,10 +23,12 @@ from model_architectures import (
|
||||
SUPPORTED_MODEL_ARCHITECTURES,
|
||||
)
|
||||
from models import DeepHealth, DeepHealthOutput
|
||||
from targets import CHECKUP_IDX, PAD_IDX
|
||||
from targets import PAD_IDX, RESERVED_IDX
|
||||
from train_util import (
|
||||
ContinuousRobustScalerStats,
|
||||
configure_torch_for_training,
|
||||
create_unique_run_dir,
|
||||
fit_continuous_robust_scaler,
|
||||
format_extra_info_types,
|
||||
get_lr,
|
||||
get_model_parameter_counts,
|
||||
@@ -120,7 +122,17 @@ def parse_args() -> argparse.Namespace:
|
||||
return args
|
||||
|
||||
|
||||
def build_model(args: argparse.Namespace, dataset: HealthDataset) -> DeepHealth:
|
||||
def build_model(
|
||||
args: argparse.Namespace,
|
||||
dataset: HealthDataset,
|
||||
scaler_stats: ContinuousRobustScalerStats,
|
||||
) -> DeepHealth:
|
||||
if tuple(int(x) for x in dataset.cont_type_ids) != scaler_stats.cont_type_ids:
|
||||
raise ValueError(
|
||||
"RobustScale statistics are not aligned with dataset.cont_type_ids"
|
||||
)
|
||||
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,
|
||||
@@ -131,6 +143,8 @@ def build_model(args: argparse.Namespace, dataset: HealthDataset) -> DeepHealth:
|
||||
n_categories=dataset.n_categories,
|
||||
cont_type_ids=dataset.cont_type_ids,
|
||||
n_bins=args.n_bins,
|
||||
continuous_value_center=center,
|
||||
continuous_value_scale=scale,
|
||||
extra_pool_reduce=args.extra_pool_reduce,
|
||||
target_mode="next_token",
|
||||
time_mode="absolute",
|
||||
@@ -143,7 +157,7 @@ def build_model(args: argparse.Namespace, dataset: HealthDataset) -> DeepHealth:
|
||||
def build_next_step_loss(args: argparse.Namespace):
|
||||
return build_loss(
|
||||
"delphi2m",
|
||||
ignored_tokens={PAD_IDX, CHECKUP_IDX},
|
||||
ignored_tokens={PAD_IDX, RESERVED_IDX},
|
||||
t_min=args.t_min,
|
||||
max_exp_input=args.max_exp_input,
|
||||
ce_weight=args.ce_weight,
|
||||
@@ -244,7 +258,6 @@ def build_augmented_next_step_targets(
|
||||
|
||||
|
||||
def compute_next_step_loss(
|
||||
args: argparse.Namespace,
|
||||
model: DeepHealth,
|
||||
criterion,
|
||||
batch: Dict[str, torch.Tensor],
|
||||
@@ -309,7 +322,7 @@ def run_epoch(
|
||||
for batch_idx, batch in enumerate(progress):
|
||||
try:
|
||||
loss, parts = compute_next_step_loss(
|
||||
args, model, criterion, batch, device
|
||||
model, criterion, batch, device
|
||||
)
|
||||
if is_train:
|
||||
if optimizer is None:
|
||||
@@ -352,6 +365,7 @@ def build_metadata(
|
||||
train_subset,
|
||||
val_subset,
|
||||
test_subset,
|
||||
scaler_stats: ContinuousRobustScalerStats,
|
||||
) -> Dict[str, Any]:
|
||||
return {
|
||||
"run_name": run_name,
|
||||
@@ -361,6 +375,8 @@ def build_metadata(
|
||||
"model_architecture": args.model_architecture,
|
||||
"model_target_mode": "next_token",
|
||||
"target_mode": "delphi2m",
|
||||
"event_stream_version": "disease_death_only_v1",
|
||||
"uses_assessment_event_token": False,
|
||||
"time_mode": "absolute",
|
||||
"dist_mode": "exponential",
|
||||
"extra_info_types_file": (
|
||||
@@ -369,6 +385,8 @@ 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_stats.as_metadata(),
|
||||
"dataset_metadata": {
|
||||
"vocab_size": int(dataset.vocab_size),
|
||||
"n_types": int(dataset.n_types),
|
||||
@@ -376,6 +394,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)),
|
||||
@@ -406,6 +426,7 @@ 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("Continuous value scaling: RobustScale (required)")
|
||||
logger.info("time_mode=absolute, readout=token, target_mode=delphi2m")
|
||||
|
||||
dataset = HealthDataset(
|
||||
@@ -441,6 +462,20 @@ def main() -> None:
|
||||
f"Samples: train={len(train_subset)}, val={len(val_subset)}, test={len(test_subset)}"
|
||||
)
|
||||
|
||||
if dataset.n_cont_types > 0:
|
||||
logger.info(
|
||||
"Fitting continuous RobustScaler on the complete training subset: "
|
||||
f"patients={len(train_subset):,}, features={dataset.n_cont_types}"
|
||||
)
|
||||
scaler_stats = fit_continuous_robust_scaler(dataset, train_subset)
|
||||
if dataset.n_cont_types > 0:
|
||||
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,
|
||||
@@ -472,7 +507,7 @@ def main() -> None:
|
||||
prefetch_factor=2 if args.num_workers > 0 else None,
|
||||
)
|
||||
|
||||
model = build_model(args, dataset).to(device)
|
||||
model = build_model(args, dataset, scaler_stats).to(device)
|
||||
parameter_counts = get_model_parameter_counts(model)
|
||||
logger.info(
|
||||
"Model parameters: "
|
||||
@@ -489,7 +524,13 @@ def main() -> None:
|
||||
adaptive_lr = args.base_lr * math.sqrt(args.batch_size / 128)
|
||||
|
||||
train_metadata = build_metadata(
|
||||
args, dataset, run_name, train_subset, val_subset, test_subset
|
||||
args,
|
||||
dataset,
|
||||
run_name,
|
||||
train_subset,
|
||||
val_subset,
|
||||
test_subset,
|
||||
scaler_stats,
|
||||
)
|
||||
train_metadata.update(parameter_counts)
|
||||
save_config(
|
||||
|
||||
Reference in New Issue
Block a user