sra-trajectory-code / MID /main_nba_mid_graphv6.py
po03087's picture
SRA: MID/LED/MoFlow code + RUNNING.md instructions (code only, no data/ckpts)
d4cbafd verified
Raw
History Blame Contribute Delete
15.1 kB
"""
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()