""" main_nba_mid_graphv6.py — NBA trajectory prediction with MoFlow's V6 graph denoiser. Replaces MID's DiffusionTraj + TransformerConcatLinear with: FlowMatcher + MotionTransformerGraphV6 (RAG-style sparse interaction graph, two-pass denoising, flow matching). Data pipeline is identical to main_nba_mid.py. Config is loaded from MoFlow's cfg/nba/cor_fm.yml with data_norm='original' (no min-max normalization — model operates in court units). Usage: python main_nba_mid_graphv6.py --data_dir ../data/nba/original """ import os import sys import time import logging import argparse import numpy as np import torch import torch.nn as nn from torch.utils.data import Dataset, DataLoader from torch.utils.tensorboard import SummaryWriter # tbX-broken from tqdm.auto import tqdm # --------------------------------------------------------------------------- # MoFlow on sys.path # --------------------------------------------------------------------------- MOFLOW_ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), '..', 'MoFlow')) sys.path.insert(0, MOFLOW_ROOT) from utils.config import Config from models.flow_matching import FlowMatcher from models.backbone_graph_v6 import MotionTransformerGraphV6 from trainer.denoising_model_trainers import build_optimizer, build_scheduler # --------------------------------------------------------------------------- # Constants (match LED / MoFlow NBA convention) # --------------------------------------------------------------------------- OBS_LEN = 10 PRED_LEN = 20 NUM_AGENTS = 11 K_EVAL = 20 # best-of-K modes at eval TRAJ_MEAN = torch.FloatTensor([14.0, 7.5]) # court-space mean after /=(94/28) # --------------------------------------------------------------------------- # Dataset (identical to main_nba_mid.py) # --------------------------------------------------------------------------- class NBADatasetMID(Dataset): """Loads nba_{train,test}.npy (shape: N, 30, 11, 2).""" def __init__(self, data_dir: str, training: bool = True): super().__init__() fname = 'nba_train.npy' if training else 'nba_test.npy' path = os.path.join(data_dir, fname) trajs = np.load(path).astype(np.float32) # (N, 30, 11, 2) trajs /= (94.0 / 28.0) # normalise court units trajs = torch.from_numpy(trajs).permute(0, 2, 1, 3) # (N, 11, 30, 2) self.pre = trajs[:, :, :OBS_LEN, :] # (N, 11, 10, 2) self.fut = trajs[:, :, OBS_LEN:, :] # (N, 11, 20, 2) def __len__(self): return len(self.pre) def __getitem__(self, idx): return self.pre[idx], self.fut[idx] # each (11, T, 2) def nba_collate(batch): pre = torch.stack([b[0] for b in batch]) # (B, 11, 10, 2) fut = torch.stack([b[1] for b in batch]) # (B, 11, 20, 2) return pre, fut # --------------------------------------------------------------------------- # Data pre-processing (MoFlow format — no TRAJ_SCALE division) # --------------------------------------------------------------------------- def preprocess_batch_graph(pre_motion: torch.Tensor, fut_motion: torch.Tensor, device: torch.device): """Build MoFlow-compatible x_data dict with data_norm='original'. past_traj_original_scale channels (un-divided): 0-1 : abs_xy = pre - traj_mean (centered at court mean) 2-3 : rel_xy = pre - last_obs (relative to last obs) 4-5 : vel_xy = diff(rel_xy) (frame-to-frame velocity) fut_traj = fut - last_obs (relative to last observation) Both in court units (after /94*28), matching MoFlow's NBADatasetMinMax 'past_traj_original_scale' / 'fut_traj_original_scale' layout. Args: pre_motion: [B, A, T_obs, 2] fut_motion: [B, A, T_fut, 2] Returns: x_data: MoFlow-compatible dict last_obs: [B, A, 1, 2] for eval reconstruction """ B, A = pre_motion.shape[:2] traj_mean = TRAJ_MEAN.to(device) last_obs = pre_motion[:, :, -1:, :] # [B, A, 1, 2] abs_xy = pre_motion - traj_mean # [B, A, T, 2] rel_xy = pre_motion - last_obs # [B, A, T, 2] vel_xy = torch.cat( [rel_xy[:, :, 1:] - rel_xy[:, :, :-1], torch.zeros_like(rel_xy[:, :, :1])], dim=2) # [B, A, T, 2] past_6ch = torch.cat([abs_xy, rel_xy, vel_xy], dim=-1) # [B, A, T, 6] fut_rel = fut_motion - last_obs # [B, A, T_fut, 2] x_data = { 'batch_size': B, 'fut_traj': fut_rel, # [B, A, T, 2] 'past_traj_original_scale': past_6ch, # [B, A, T_obs, 6] } return x_data, last_obs # --------------------------------------------------------------------------- # MoFlow config builder # --------------------------------------------------------------------------- def build_moflow_cfg(args): """Load cor_fm.yml and apply all CLI overrides.""" cfg_path = os.path.join(MOFLOW_ROOT, 'cfg', 'nba', 'cor_fm.yml') cfg = Config(cfg_path, tag=args.exp_name) # device cfg.device = 'cuda' if torch.cuda.is_available() else 'cpu' # data normalisation: 'original' → no unnorm in loss/eval cfg.data_norm = 'original' # flow-matching schedule cfg.sampling_steps = args.sampling_steps cfg.t_schedule = 'logit_normal' cfg.logit_norm_mean = -0.5 cfg.logit_norm_std = 1.5 cfg.fm_wrapper = 'direct' cfg.fm_rew_sqrt = False cfg.fm_in_scaling = True # --fm_in_scaling cfg.tied_noise = True # --tied_noise # input dropout (emb-level, same as default MoFlow graph runs) cfg.drop_method = 'emb' cfg.drop_logi_k = 20.0 cfg.drop_logi_m = 0.5 # loss settings cfg.LOSS_NN_MODE = 'agent' cfg.LOSS_REG_REDUCTION = 'sum' cfg.LOSS_REG_SQUARED = False cfg.LOSS_VELOCITY = False cfg.uncertainty_weight = args.uncertainty_weight # graph hyper-parameters cfg.graph_gnn_layers = args.graph_gnn_layers cfg.graph_dropout = args.graph_dropout cfg.top_n_neighbors = args.top_n_neighbors cfg.rel_traj_hidden = args.rel_traj_hidden cfg.y0_score_dim = args.y0_score_dim # training schedule overrides if args.epochs is not None: cfg.OPTIMIZATION.NUM_EPOCHS = args.epochs if args.batch_size is not None: cfg.train_batch_size = args.batch_size return cfg # --------------------------------------------------------------------------- # Trainer # --------------------------------------------------------------------------- class Trainer: def __init__(self, args): self.args = args self.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') self._build_dirs() self._build_cfg() self._build_data() self._build_model() self._build_optimizer() def _build_dirs(self): self.exp_dir = os.path.join('experiments', self.args.exp_name) os.makedirs(self.exp_dir, exist_ok=True) self.tb_log = SummaryWriter(log_dir=self.exp_dir) log_path = os.path.join( self.exp_dir, 'nba_{}.log'.format(time.strftime('%Y-%m-%d-%H-%M'))) self.log = logging.getLogger(self.args.exp_name) self.log.setLevel(logging.INFO) self.log.addHandler(logging.FileHandler(log_path)) self.log.addHandler(logging.StreamHandler(sys.stdout)) self.log.info(f"Args: {self.args}") def _build_cfg(self): self.cfg = build_moflow_cfg(self.args) self.log.info( f"MoFlow cfg: epochs={self.cfg.OPTIMIZATION.NUM_EPOCHS} " f"batch={self.cfg.train_batch_size} " f"sampling_steps={self.cfg.sampling_steps} " f"data_norm={self.cfg.data_norm}" ) def _build_data(self): train_dset = NBADatasetMID(self.args.data_dir, training=True) test_dset = NBADatasetMID(self.args.data_dir, training=False) batch_size = self.args.batch_size eval_bs = self.args.eval_batch_size self.train_loader = DataLoader( train_dset, batch_size=batch_size, shuffle=True, num_workers=4, collate_fn=nba_collate, pin_memory=True) self.test_loader = DataLoader( test_dset, batch_size=eval_bs, shuffle=False, num_workers=4, collate_fn=nba_collate, pin_memory=True) self.log.info( f"Train: {len(train_dset)} scenes " f"Test: {len(test_dset)} scenes " f"train_bs={batch_size} eval_bs={eval_bs}") def _build_model(self): model = MotionTransformerGraphV6( model_config = self.cfg.MODEL, logger = self.log, config = self.cfg, graph_num_gnn_layers = self.args.graph_gnn_layers, graph_dropout = self.args.graph_dropout, top_n_neighbors = self.args.top_n_neighbors, rel_traj_hidden = self.args.rel_traj_hidden, y0_score_dim = self.args.y0_score_dim, ) self.denoiser = FlowMatcher(self.cfg, model, logger=self.log).to(self.device) n_params = sum(p.numel() for p in self.denoiser.parameters()) self.log.info(f"Total denoiser params: {n_params:,}") def _build_optimizer(self): self.optimizer = build_optimizer(self.denoiser, self.cfg.OPTIMIZATION) self.scheduler = build_scheduler( self.optimizer, self.cfg.OPTIMIZATION, total_iters_each_epoch=len(self.train_loader), ) # ------------------------------------------------------------------ def train(self): num_epochs = self.cfg.OPTIMIZATION.NUM_EPOCHS for epoch in range(1, num_epochs + 1): self.denoiser.train() total_loss, total_reg, count = 0.0, 0.0, 0 log_dict = {'cur_epoch': epoch} pbar = tqdm(self.train_loader, ncols=90) for pre, fut in pbar: pre = pre.to(self.device) fut = fut.to(self.device) x_data, _ = preprocess_batch_graph(pre, fut, self.device) loss, loss_reg, loss_cls, _, _ = self.denoiser.p_losses( x_data, log_dict=log_dict) self.optimizer.zero_grad() loss.backward() nn.utils.clip_grad_norm_( self.denoiser.parameters(), self.cfg.OPTIMIZATION.GRAD_NORM_CLIP) self.optimizer.step() if self.scheduler is not None: self.scheduler.step() total_loss += loss.item() total_reg += loss_reg.item() count += 1 pbar.set_description( f"E{epoch} loss={total_loss/count:.4f}" f" reg={total_reg/count:.4f}") avg_loss = total_loss / count avg_reg = total_reg / count self.tb_log.add_scalar('loss/train', avg_loss, epoch) self.tb_log.add_scalar('loss/train_reg', avg_reg, epoch) self.log.info( f"Epoch {epoch:3d} train_loss={avg_loss:.4f}" f" reg={avg_reg:.4f}") if epoch % self.args.eval_every == 0: ade, fde = self.evaluate() self.tb_log.add_scalar('metric/ADE', ade, epoch) self.tb_log.add_scalar('metric/FDE', fde, epoch) self.log.info( f"Epoch {epoch:3d} ADE={ade:.4f} FDE={fde:.4f}") torch.save({ 'denoiser': self.denoiser.state_dict(), 'optimizer': self.optimizer.state_dict(), 'epoch': epoch, }, os.path.join(self.exp_dir, f'ckpt_epoch{epoch:04d}.pt')) @torch.no_grad() def evaluate(self): self.denoiser.eval() ade_sum, fde_sum, n_agents = 0.0, 0.0, 0 K = K_EVAL for pre, fut in tqdm(self.test_loader, ncols=90, desc='Eval'): pre = pre.to(self.device) fut = fut.to(self.device) B = pre.size(0) x_data, last_obs = preprocess_batch_graph(pre, fut, self.device) # sample returns (y_t, y_data_at_t_ls [B,S,K,A,T*2], t_ls, y_t_ls, pred_score) _, y_data_at_t_ls, _, _, _ = self.denoiser.sample(x_data, num_trajs=K) # final step: [B, K, A, T*2] → [B, K, A, T, 2] pred_rel = y_data_at_t_ls[:, -1].view(B, K, NUM_AGENTS, PRED_LEN, 2) # absolute positions: pred_rel + last_obs # last_obs: [B, A, 1, 2] → unsqueeze → [B, 1, A, 1, 2] pred_abs = pred_rel + last_obs.unsqueeze(1) # [B, K, A, T, 2] fut_abs = fut # [B, A, T, 2] # minADE / minFDE over K modes dist = (pred_abs - fut_abs.unsqueeze(1)).norm(dim=-1) # [B, K, A, T] best_k = dist.min(dim=1).values # [B, A, T] ade_sum += best_k.mean(dim=-1).sum().item() # sum over B*A fde_sum += best_k[:, :, -1].sum().item() n_agents += B * NUM_AGENTS ade = ade_sum / n_agents fde = fde_sum / n_agents return ade, fde # --------------------------------------------------------------------------- # Argument parsing # --------------------------------------------------------------------------- def parse_args(): p = argparse.ArgumentParser() # Data p.add_argument('--data_dir', type=str, default='../data/nba/original') # Experiment p.add_argument('--exp_name', type=str, default='mid_nba_graphv6') # Training p.add_argument('--epochs', type=int, default=None, help='Override cfg NUM_EPOCHS (default: 150 from YAML)') p.add_argument('--batch_size', type=int, default=192) p.add_argument('--eval_batch_size', type=int, default=500) p.add_argument('--eval_every', type=int, default=5) # Sampling p.add_argument('--sampling_steps', type=int, default=10) # Graph p.add_argument('--graph_gnn_layers', type=int, default=2) p.add_argument('--graph_dropout', type=float, default=0.1) p.add_argument('--top_n_neighbors', type=int, default=5) p.add_argument('--rel_traj_hidden', type=int, default=32) p.add_argument('--y0_score_dim', type=int, default=32) p.add_argument('--uncertainty_weight', type=float, default=0.01) return p.parse_args() # --------------------------------------------------------------------------- # Main # --------------------------------------------------------------------------- if __name__ == '__main__': args = parse_args() trainer = Trainer(args) trainer.train()