Implement TrajMixer block

This commit is contained in:
2026-07-22 11:52:44 +08:00
parent f6bde7e167
commit db0947ce9d
8 changed files with 575 additions and 24 deletions

View File

@@ -27,7 +27,7 @@ from tqdm.auto import tqdm
from dataset import AllFutureHealthDataset, all_future_collate_fn
from losses import build_loss
from models import DeepHealth
from models import TRAJ_MIXER_ARCHITECTURE, DeepHealth
from targets import CHECKUP_IDX, PAD_IDX
from train_util import (
configure_torch_for_training,
@@ -89,6 +89,7 @@ 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("--hidden_group", type=int, default=20)
parser.add_argument("--batch_size", type=int, default=128)
parser.add_argument("--base_lr", type=float, default=3e-4)
@@ -159,6 +160,7 @@ def build_model(args: argparse.Namespace, dataset: AllFutureHealthDataset) -> De
time_mode=args.time_mode,
dist_mode=args.dist_mode,
dropout=args.dropout,
hidden_group=args.hidden_group,
)
@@ -298,6 +300,7 @@ def build_metadata(
"dataset_class": "AllFutureHealthDataset",
"collate_fn": "all_future_collate_fn",
"model_class": "DeepHealth",
"model_architecture": TRAJ_MIXER_ARCHITECTURE,
"model_target_mode": "all_future",
"target_mode": "all_future",
"dist_mode": args.dist_mode,