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

@@ -24,7 +24,7 @@ from tqdm.auto import tqdm
from dataset import HealthDataset, collate_fn
from losses import build_loss
from models import DeepHealth, DeepHealthOutput
from models import TRAJ_MIXER_ARCHITECTURE, DeepHealth, DeepHealthOutput
from readouts import build_readout
from targets import CHECKUP_IDX, NO_EVENT_IDX, PAD_IDX
from train_util import (
@@ -83,6 +83,7 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--time_mode", type=str, default="relative",
choices=["relative", "absolute"])
parser.add_argument("--dropout", type=float, default=0.0)
parser.add_argument("--hidden_group", type=int, default=20)
parser.add_argument("--target_mode", type=str, default="uts",
choices=["delphi2m", "uts"])
@@ -164,6 +165,7 @@ def build_model(args: argparse.Namespace, dataset: HealthDataset) -> DeepHealth:
time_mode=args.time_mode,
dist_mode="exponential",
dropout=args.dropout,
hidden_group=args.hidden_group,
)
@@ -484,6 +486,7 @@ def build_metadata(
"dataset_class": "NextStepHealthDataset",
"collate_fn": "next_step_collate_fn",
"model_class": "DeepHealth",
"model_architecture": TRAJ_MIXER_ARCHITECTURE,
"model_target_mode": "next_token",
"target_mode": args.target_mode,
"dist_mode": "exponential",