Unify FFN and TrajMixer model architectures
This commit is contained in:
@@ -24,6 +24,10 @@ from tqdm.auto import tqdm
|
||||
|
||||
from dataset import HealthDataset, collate_fn
|
||||
from losses import build_loss
|
||||
from model_architectures import (
|
||||
DEFAULT_MODEL_ARCHITECTURE,
|
||||
SUPPORTED_MODEL_ARCHITECTURES,
|
||||
)
|
||||
from models import DeepHealth, DeepHealthOutput
|
||||
from readouts import build_readout
|
||||
from targets import CHECKUP_IDX, NO_EVENT_IDX, PAD_IDX
|
||||
@@ -31,6 +35,7 @@ from train_util import (
|
||||
configure_torch_for_training,
|
||||
create_unique_run_dir,
|
||||
format_extra_info_types,
|
||||
get_model_parameter_counts,
|
||||
load_extra_info_types_file,
|
||||
resolve_device,
|
||||
save_checkpoint,
|
||||
@@ -61,6 +66,7 @@ def parse_args() -> argparse.Namespace:
|
||||
|
||||
parser.add_argument("--data_prefix", type=str, default="ukb")
|
||||
parser.add_argument("--labels_file", type=str, default="labels.csv")
|
||||
parser.add_argument("--runs_root", type=str, default="runs")
|
||||
parser.add_argument("--seed", type=int, default=42)
|
||||
parser.add_argument("--extra_info_types_file", type=str, default=None)
|
||||
parser.add_argument("--no_event_interval_years", type=float, default=5.0)
|
||||
@@ -75,14 +81,19 @@ def parse_args() -> argparse.Namespace:
|
||||
|
||||
parser.add_argument("--n_embd", type=int, default=120)
|
||||
parser.add_argument("--n_head", type=int, default=10)
|
||||
parser.add_argument("--n_hist_layer", type=int, default=12)
|
||||
parser.add_argument("--n_tab_layer", type=int, default=4)
|
||||
parser.add_argument("--n_layer", type=int, default=12)
|
||||
parser.add_argument("--n_bins", type=int, default=16)
|
||||
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("--dropout", type=float, default=0.0)
|
||||
parser.add_argument(
|
||||
"--model_architecture",
|
||||
type=str,
|
||||
default=DEFAULT_MODEL_ARCHITECTURE,
|
||||
choices=SUPPORTED_MODEL_ARCHITECTURES,
|
||||
)
|
||||
|
||||
parser.add_argument("--target_mode", type=str, default="uts",
|
||||
choices=["delphi2m", "uts"])
|
||||
@@ -152,8 +163,7 @@ def build_model(args: argparse.Namespace, dataset: HealthDataset) -> DeepHealth:
|
||||
vocab_size=dataset.vocab_size,
|
||||
n_embd=args.n_embd,
|
||||
n_head=args.n_head,
|
||||
n_hist_layer=args.n_hist_layer,
|
||||
n_tab_layer=args.n_tab_layer,
|
||||
n_layer=args.n_layer,
|
||||
n_types=dataset.n_types,
|
||||
n_cont_types=dataset.n_cont_types,
|
||||
n_categories=dataset.n_categories,
|
||||
@@ -164,6 +174,7 @@ def build_model(args: argparse.Namespace, dataset: HealthDataset) -> DeepHealth:
|
||||
time_mode=args.time_mode,
|
||||
dist_mode="exponential",
|
||||
dropout=args.dropout,
|
||||
model_architecture=args.model_architecture,
|
||||
)
|
||||
|
||||
|
||||
@@ -484,6 +495,7 @@ def build_metadata(
|
||||
"dataset_class": "NextStepHealthDataset",
|
||||
"collate_fn": "next_step_collate_fn",
|
||||
"model_class": "DeepHealth",
|
||||
"model_architecture": args.model_architecture,
|
||||
"model_target_mode": "next_token",
|
||||
"target_mode": args.target_mode,
|
||||
"dist_mode": "exponential",
|
||||
@@ -521,12 +533,14 @@ def main() -> None:
|
||||
lambda timestamp: (
|
||||
f"{args.time_mode}_exponential_next_token_{args.target_mode}_"
|
||||
f"gap_{args.no_event_interval_years:g}y_{timestamp}"
|
||||
)
|
||||
),
|
||||
runs_root=Path(args.runs_root) / args.model_architecture,
|
||||
)
|
||||
logger = setup_logging(run_dir)
|
||||
|
||||
logger.info(f"Starting next-step training run: {run_name}")
|
||||
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(f"readout={args.readout_name}, target_mode={args.target_mode}")
|
||||
|
||||
@@ -596,6 +610,12 @@ def main() -> None:
|
||||
)
|
||||
|
||||
model = build_model(args, dataset).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']:,}"
|
||||
)
|
||||
readout = build_next_step_readout(args).to(device)
|
||||
criterion = build_next_step_loss(args)
|
||||
optimizer = AdamW(
|
||||
@@ -606,10 +626,14 @@ 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
|
||||
)
|
||||
train_metadata.update(parameter_counts)
|
||||
save_config(
|
||||
args,
|
||||
run_dir / "train_config.json",
|
||||
extra=build_metadata(args, dataset, run_name, train_subset, val_subset, test_subset),
|
||||
extra=train_metadata,
|
||||
)
|
||||
|
||||
best_val = float("inf")
|
||||
|
||||
Reference in New Issue
Block a user