Unify FFN and TrajMixer model architectures

This commit is contained in:
2026-07-25 12:59:09 +08:00
parent 4526191fe1
commit b13db5e407
18 changed files with 1019 additions and 52 deletions

View File

@@ -27,12 +27,17 @@ from tqdm.auto import tqdm
from dataset import AllFutureHealthDataset, all_future_collate_fn
from losses import build_loss
from model_architectures import (
DEFAULT_MODEL_ARCHITECTURE,
SUPPORTED_MODEL_ARCHITECTURES,
)
from models import DeepHealth
from targets import CHECKUP_IDX, PAD_IDX
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,
@@ -64,6 +69,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)
@@ -79,8 +85,7 @@ 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"])
@@ -89,6 +94,12 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--dist_mode", type=str, default="exponential",
choices=["exponential", "weibull", "mixed"])
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("--batch_size", type=int, default=128)
parser.add_argument("--base_lr", type=float, default=3e-4)
@@ -147,8 +158,7 @@ def build_model(args: argparse.Namespace, dataset: AllFutureHealthDataset) -> De
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,
@@ -159,6 +169,7 @@ def build_model(args: argparse.Namespace, dataset: AllFutureHealthDataset) -> De
time_mode=args.time_mode,
dist_mode=args.dist_mode,
dropout=args.dropout,
model_architecture=args.model_architecture,
)
@@ -298,6 +309,7 @@ def build_metadata(
"dataset_class": "AllFutureHealthDataset",
"collate_fn": "all_future_collate_fn",
"model_class": "DeepHealth",
"model_architecture": args.model_architecture,
"model_target_mode": "all_future",
"target_mode": "all_future",
"dist_mode": args.dist_mode,
@@ -335,12 +347,14 @@ def main() -> None:
configure_torch_for_training(device)
run_dir, run_name = create_unique_run_dir(
lambda timestamp: f"{args.time_mode}_{args.dist_mode}_all_future_pure_disease_{timestamp}"
lambda timestamp: f"{args.time_mode}_{args.dist_mode}_all_future_pure_disease_{timestamp}",
runs_root=Path(args.runs_root) / args.model_architecture,
)
logger = setup_logging(run_dir)
logger.info(f"Starting all-future 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("Loading all-future datasets...")
@@ -434,6 +448,12 @@ def main() -> None:
)
model = build_model(args, train_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']:,}"
)
optimizer = AdamW(
model.parameters(),
lr=args.base_lr,
@@ -443,10 +463,14 @@ def main() -> None:
criterion = build_criterion(args, train_dataset)
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
)
train_metadata.update(parameter_counts)
save_config(
args,
run_dir / "train_config.json",
extra=build_metadata(args, train_dataset, run_name, train_subset, val_subset, test_subset),
extra=train_metadata,
)
best_val = float("inf")