| """ |
| 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 |
| from tqdm.auto import tqdm |
|
|
| |
| |
| |
|
|
| 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 |
|
|
|
|
| |
| |
| |
|
|
| OBS_LEN = 10 |
| PRED_LEN = 20 |
| NUM_AGENTS = 11 |
| K_EVAL = 20 |
| TRAJ_MEAN = torch.FloatTensor([14.0, 7.5]) |
|
|
|
|
| |
| |
| |
|
|
| 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) |
| trajs /= (94.0 / 28.0) |
|
|
| trajs = torch.from_numpy(trajs).permute(0, 2, 1, 3) |
| self.pre = trajs[:, :, :OBS_LEN, :] |
| self.fut = trajs[:, :, OBS_LEN:, :] |
|
|
| def __len__(self): |
| return len(self.pre) |
|
|
| def __getitem__(self, idx): |
| return self.pre[idx], self.fut[idx] |
|
|
|
|
| def nba_collate(batch): |
| pre = torch.stack([b[0] for b in batch]) |
| fut = torch.stack([b[1] for b in batch]) |
| return pre, fut |
|
|
|
|
| |
| |
| |
|
|
| 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:, :] |
|
|
| abs_xy = pre_motion - traj_mean |
| rel_xy = pre_motion - last_obs |
| vel_xy = torch.cat( |
| [rel_xy[:, :, 1:] - rel_xy[:, :, :-1], |
| torch.zeros_like(rel_xy[:, :, :1])], dim=2) |
|
|
| past_6ch = torch.cat([abs_xy, rel_xy, vel_xy], dim=-1) |
| fut_rel = fut_motion - last_obs |
|
|
| x_data = { |
| 'batch_size': B, |
| 'fut_traj': fut_rel, |
| 'past_traj_original_scale': past_6ch, |
| } |
| return x_data, last_obs |
|
|
|
|
| |
| |
| |
|
|
| 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) |
|
|
| |
| cfg.device = 'cuda' if torch.cuda.is_available() else 'cpu' |
|
|
| |
| cfg.data_norm = 'original' |
|
|
| |
| 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 |
| cfg.tied_noise = True |
|
|
| |
| cfg.drop_method = 'emb' |
| cfg.drop_logi_k = 20.0 |
| cfg.drop_logi_m = 0.5 |
|
|
| |
| 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 |
|
|
| |
| 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 |
|
|
| |
| 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 |
|
|
|
|
| |
| |
| |
|
|
| 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) |
|
|
| |
| _, y_data_at_t_ls, _, _, _ = self.denoiser.sample(x_data, num_trajs=K) |
|
|
| |
| pred_rel = y_data_at_t_ls[:, -1].view(B, K, NUM_AGENTS, PRED_LEN, 2) |
|
|
| |
| |
| pred_abs = pred_rel + last_obs.unsqueeze(1) |
| fut_abs = fut |
|
|
| |
| dist = (pred_abs - fut_abs.unsqueeze(1)).norm(dim=-1) |
| best_k = dist.min(dim=1).values |
|
|
| ade_sum += best_k.mean(dim=-1).sum().item() |
| 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 |
|
|
|
|
| |
| |
| |
|
|
| def parse_args(): |
| p = argparse.ArgumentParser() |
|
|
| |
| p.add_argument('--data_dir', type=str, default='../data/nba/original') |
|
|
| |
| p.add_argument('--exp_name', type=str, default='mid_nba_graphv6') |
|
|
| |
| 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) |
|
|
| |
| p.add_argument('--sampling_steps', type=int, default=10) |
|
|
| |
| 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() |
|
|
|
|
| |
| |
| |
|
|
| if __name__ == '__main__': |
| args = parse_args() |
| trainer = Trainer(args) |
| trainer.train() |
|
|