""" main_nba_mid_graphv6_v2.py — MID on NBA with V6-style future interaction graph. DESIGN PRINCIPLE ---------------- Only the denoising network changes from main_nba_mid.py: - Past context encoder : NBAEncoder (GRU + SocialTransformer) — UNCHANGED - Diffusion schedule : DiffusionTrajGraph (DDPM, 100 steps) — UNCHANGED - Denoising network : GraphDenoiserNet (NEW) TransformerConcatLinear backbone → eps_hat (noise prediction, same as baseline) x_0 derived from eps for graph → geometry for neighbor selection FutureInteractionGraphV6 (K=1) → refines eps_hat using inter-agent geometry GRAPH INTEGRATION (same manner as V6) -------------------------------------- Agent selection — RAG-style: q_i = W_q([y0_emb_i, 0]) per-node query (no sigma: first pass only) k_j = W_k([y0_emb_j, 0]) per-node key semantic_score = (q_i · k_j) / √D_s geo_bias = geo_mlp([mean_rel, std_rel, min_dist, heading_mean]) score = semantic + geo → top-N per target agent Pairwise encoding — RelTrajEncoder: edge_feat = RelTrajEncoder([rel_pos(2), heading_diff(1)]) over T steps → GNN message passing → gated residual refinement of x_0_hat Usage: python main_nba_mid_graphv6_v2.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 import torch.nn.functional as F from torch.utils.data import Dataset, DataLoader from torch.utils.tensorboard import SummaryWriter # tbX-broken from tqdm.auto import tqdm # --------------------------------------------------------------------------- # MoFlow graph module on sys.path # --------------------------------------------------------------------------- MOFLOW_ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), '..', 'MoFlow')) sys.path.insert(0, MOFLOW_ROOT) from models.graph_interaction_nba_v6 import FutureInteractionGraphV6 from models.context_encoder.mtr_encoder import SinusoidalPosEmb # MID diffusion components (unchanged) from models.diffusion import VarianceSchedule, TransformerConcatLinear # --------------------------------------------------------------------------- # Constants # --------------------------------------------------------------------------- OBS_LEN = 10 PRED_LEN = 20 NUM_AGENTS = 11 TRAJ_SCALE = 5.0 TRAJ_MEAN = torch.FloatTensor([14.0, 7.5]) K_EVAL = 20 # --------------------------------------------------------------------------- # Dataset (identical to main_nba_mid.py) # --------------------------------------------------------------------------- class NBADatasetMID(Dataset): def __init__(self, data_dir: str, training: bool = True): super().__init__() fname = 'nba_train.npy' if training else 'nba_test.npy' trajs = np.load(os.path.join(data_dir, fname)).astype(np.float32) trajs /= (94.0 / 28.0) trajs = torch.from_numpy(trajs).permute(0, 2, 1, 3) # (N, 11, 30, 2) 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]) # (B, 11, 10, 2) fut = torch.stack([b[1] for b in batch]) # (B, 11, 20, 2) return pre, fut # --------------------------------------------------------------------------- # Data pre-processing (identical to main_nba_mid.py) # --------------------------------------------------------------------------- def preprocess_batch(pre_motion, fut_motion, device): """ Returns: past_6ch: [B*A, T_obs, 6] (÷ TRAJ_SCALE) fut_rel: [B*A, T_fut, 2] (÷ TRAJ_SCALE, relative to last obs) social_mask: [B*A, B*A] last_obs: [B*A, 1, 2] un-scaled """ B, A = pre_motion.shape[:2] traj_mean = TRAJ_MEAN.to(device) pre = pre_motion.reshape(B * A, OBS_LEN, 2) fut = fut_motion.reshape(B * A, PRED_LEN, 2) last_obs = pre[:, -1:, :] abs_xy = (pre - traj_mean) / TRAJ_SCALE rel_xy = (pre - last_obs) / TRAJ_SCALE vel_xy = torch.cat([rel_xy[:, 1:] - rel_xy[:, :-1], torch.zeros_like(rel_xy[:, :1])], dim=1) past_6ch = torch.cat([abs_xy, rel_xy, vel_xy], dim=-1) fut_rel = (fut - last_obs) / TRAJ_SCALE mask = torch.full((B * A, B * A), float('-inf'), device=device) for i in range(B): s, e = i * A, (i + 1) * A mask[s:e, s:e] = 0.0 return past_6ch, fut_rel, mask, last_obs # --------------------------------------------------------------------------- # NBA Encoder (identical to main_nba_mid.py) # --------------------------------------------------------------------------- class _STEncoder(nn.Module): def __init__(self, in_channels=6, hidden=256): super().__init__() self.conv = nn.Conv1d(in_channels, 32, kernel_size=3, padding=1) self.relu = nn.ReLU() self.gru = nn.GRU(32, hidden, num_layers=1, batch_first=True) nn.init.kaiming_normal_(self.conv.weight) nn.init.kaiming_normal_(self.gru.weight_ih_l0) nn.init.kaiming_normal_(self.gru.weight_hh_l0) nn.init.zeros_(self.conv.bias) nn.init.zeros_(self.gru.bias_ih_l0) nn.init.zeros_(self.gru.bias_hh_l0) def forward(self, x): h = self.relu(self.conv(x.transpose(1, 2))) _, state = self.gru(h.transpose(1, 2)) return state.squeeze(0) class _SocialTransformer(nn.Module): def __init__(self, past_len=OBS_LEN, hidden=256): super().__init__() self.proj = nn.Linear(past_len * 6, hidden, bias=False) layer = nn.TransformerEncoderLayer( d_model=hidden, nhead=2, dim_feedforward=hidden, batch_first=False) self.encoder = nn.TransformerEncoder(layer, num_layers=2) def forward(self, x_flat, mask): h = self.proj(x_flat).unsqueeze(1) h = h + self.encoder(h, mask) return h.squeeze(1) class NBAEncoder(nn.Module): def __init__(self, encoder_dim=256, past_len=OBS_LEN): super().__init__() self.ego_encoder = _STEncoder(in_channels=6, hidden=256) self.social_encoder = _SocialTransformer(past_len=past_len, hidden=256) self.fusion = nn.Linear(512, encoder_dim) def forward(self, past_6ch, social_mask): ego = self.ego_encoder(past_6ch) social = self.social_encoder( past_6ch.reshape(past_6ch.size(0), -1), social_mask) return self.fusion(torch.cat([ego, social], dim=-1)) # --------------------------------------------------------------------------- # Graph-augmented denoising network # --------------------------------------------------------------------------- class GraphDenoiserNet(nn.Module): """TransformerConcatLinear backbone + V6-style future interaction graph. Predicts epsilon (noise), same as original MID baseline. Two-pass design (mirrors V6): Pass 1 (no_grad): backbone(x_t, beta, context) → eps_geom [N, T, 2] x_0_geom = (x_t - c1*eps_geom) / c0 (detached) Used ONLY for graph geometry — no gradient. Pass 2 (grad-enabled): backbone(x_t, beta, context) → eps_hat [N, T, 2] (grad-connected) x_0_hat = (x_t - c1*eps_hat) / c0 (grad through eps_hat) node_proj(x_0_hat) → y_emb [B, 1, A, D] (grad-connected) FutureInteractionGraphV6 (K=1): geometry : x_0_geom * TRAJ_SCALE + last_obs (detached from pass 1) nodes : y_emb (grad-connected from pass 2) → RAG scoring (W_q/W_k + geo_bias) → top-N selection → RelTrajEncoder([rel_pos, heading_diff]) → GNN + gated residual → refined [B, 1, A, D] refine_head(refined) → delta [N, T, 2] output: eps_hat (pass 2) + delta Gradient flow: - Backbone receives gradients from eps_hat directly and via x_0_hat → node_proj → graph → refine_head → delta. - Edge geometry (positions) is detached → topology selection has no gradient. """ def __init__(self, context_dim: int = 256, tf_layer: int = 3, T: int = PRED_LEN, num_agents: int = NUM_AGENTS, graph_hidden: int = 128, top_n_neighbors: int = 5, rel_traj_hidden: int = 32, y0_score_dim: int = 32, num_gnn_layers: int = 2, graph_dropout: float = 0.1): super().__init__() self.T = T self.A = num_agents self.D = graph_hidden # -- Backbone: standard MID denoising network ---------------------- self.backbone = TransformerConcatLinear( point_dim = 2, context_dim = context_dim, tf_layer = tf_layer, residual = False, ) # -- Node projection: x_0_hat trajectory → graph embedding -------- self.node_proj = nn.Sequential( nn.Linear(T * 2, graph_hidden), nn.ReLU(inplace=True), ) # -- Time embedding for GNN: beta (noise level) → [D] -------------- self.time_mlp = nn.Sequential( SinusoidalPosEmb(graph_hidden), nn.Linear(graph_hidden, graph_hidden), nn.ReLU(), nn.Linear(graph_hidden, graph_hidden), ) # -- V6-style future interaction graph (K=1) ----------------------- self.future_graph = FutureInteractionGraphV6( embed_dim = graph_hidden, future_steps = T, num_agents = num_agents, num_heads = 4, dropout = graph_dropout, num_gnn_layers = num_gnn_layers, time_dim = graph_hidden, top_n_neighbors = top_n_neighbors, rel_traj_hidden = rel_traj_hidden, y0_score_dim = y0_score_dim, ) # -- Refinement head: embedding → trajectory correction ------------ self.refine_head = nn.Linear(graph_hidden, T * 2) # Zero-init: graph starts as an identity (no correction), training adds it nn.init.zeros_(self.refine_head.weight) nn.init.zeros_(self.refine_head.bias) def forward(self, x_t: torch.Tensor, # [N=B*A, T, 2] noisy trajectory beta: torch.Tensor, # [N] noise level (same per scene) context: torch.Tensor, # [N, context_dim] B: int, A: int, last_obs: torch.Tensor, # [N, 1, 2] last observed position (un-scaled) tau: torch.Tensor, # [B] normalised timestep ∈ [0, 1] alpha_bar: torch.Tensor, # [N] cumulative alpha for x_0 recovery ) -> torch.Tensor: # [N, T, 2] predicted noise (epsilon) N, T, _ = x_t.shape D = self.D c0 = alpha_bar.sqrt().view(N, 1, 1) # [N, 1, 1] c1 = (1 - alpha_bar).sqrt().view(N, 1, 1) # ---- Pass 1 (no_grad): predict eps → derive x_0 for geometry ----- with torch.no_grad(): eps_geom = self.backbone(x_t, beta=beta, context=context) # [N, T, 2] x_0_geom = (x_t - c1 * eps_geom) / c0 # [N, T, 2] # Absolute positions (detached — geometry has no gradient) x_0_abs = x_0_geom * TRAJ_SCALE + last_obs # [N, T, 2] y_abs = x_0_abs.view(B, A, T, 2).unsqueeze(1) # [B, 1, A, T, 2] # ---- Pass 2 (grad-enabled): predict eps + derive x_0 for nodes --- eps_hat = self.backbone(x_t, beta=beta, context=context) # [N, T, 2] x_0_hat = (x_t - c1 * eps_hat) / c0 # grad through eps_hat # ---- Node embeddings from derived x_0 (grad-connected) ----------- y_emb = self.node_proj(x_0_hat.view(N, T * 2)) # [N, D] y_emb = y_emb.view(B, 1, A, D) # [B, 1, A, D] # ---- 4. Time embedding (per scene, from per-agent beta) ----------- beta_scene = beta.view(B, A)[:, 0] # [B] t_emb = self.time_mlp(beta_scene) # [B, D] # ---- 5. V6 future interaction graph (K=1, sigma=None) ------------ refined = self.future_graph( y_emb, y_abs, t_emb, tau, sigma_agent=None, ) # [B, 1, A, D] # ---- 6. Decode and add residual correction ------------------------ delta = self.refine_head( refined.squeeze(1).reshape(N, D) ).view(N, T, 2) # [N, T, 2] return eps_hat + delta # --------------------------------------------------------------------------- # Graph-aware diffusion wrapper # --------------------------------------------------------------------------- class DiffusionTrajGraph(nn.Module): """DDPM wrapper for GraphDenoiserNet. Key differences from DiffusionTraj: - Predicts epsilon (noise), same as original MID - Derives x_0 from epsilon internally for graph geometry - Samples one timestep per scene (not per agent) → all 11 agents in a scene share the same noise level, giving the graph a consistent geometric signal - Passes B, A, last_obs, tau, alpha_bar to the net at every step """ def __init__(self, net: GraphDenoiserNet, var_sched: VarianceSchedule): super().__init__() self.net = net self.var_sched = var_sched def get_loss(self, x_0: torch.Tensor, # [B*A, T, 2] ground-truth relative future context: torch.Tensor, # [B*A, enc_dim] B: int, A: int, last_obs: torch.Tensor, # [B*A, 1, 2] un-scaled last obs ) -> torch.Tensor: device = x_0.device N = B * A # One timestep per scene, repeated for each of its A agents t_scene = self.var_sched.uniform_sample_t(B) # list of B ints ∈ [1, T] t = [ts for ts in t_scene for _ in range(A)] # B*A ints alpha_bar = self.var_sched.alpha_bars[t].to(device) # [N] beta = self.var_sched.betas[t].to(device) # [N] c0 = alpha_bar.sqrt().view(N, 1, 1) # [N, 1, 1] c1 = (1 - alpha_bar).sqrt().view(N, 1, 1) e_rand = torch.randn_like(x_0) x_t = c0 * x_0 + c1 * e_rand # [N, T, 2] tau = torch.tensor( [ts / self.var_sched.num_steps for ts in t_scene], dtype=torch.float32, device=device, ) # [B] ∈ [0, 1] eps_pred = self.net(x_t, beta, context, B, A, last_obs, tau, alpha_bar) return F.mse_loss(eps_pred, e_rand) @torch.no_grad() def sample(self, num_points: int, context: torch.Tensor, # [B*A, enc_dim] B: int, A: int, last_obs: torch.Tensor, # [B*A, 1, 2] sample: int = K_EVAL, bestof: bool = True, sampling: str = 'ddim', step: int = 10, ) -> torch.Tensor: # [K, B*A, T, 2] device = context.device N = B * A stride = self.var_sched.num_steps // step traj_list = [] for _ in range(sample): x_T = (torch.randn(N, num_points, 2, device=device) if bestof else torch.zeros(N, num_points, 2, device=device)) x_t = x_T for t in range(self.var_sched.num_steps, 0, -stride): alpha_bar = self.var_sched.alpha_bars[t].to(device) alpha_bar_prev = self.var_sched.alpha_bars[t - stride].to(device) beta_t = self.var_sched.betas[t].to(device) beta_batch = beta_t.expand(N) alpha_bar_vec = alpha_bar.expand(N) tau = torch.full((B,), t / self.var_sched.num_steps, dtype=torch.float32, device=device) e_theta = self.net(x_t, beta_batch, context, B, A, last_obs, tau, alpha_bar_vec) if sampling == 'ddim': x0_t = (x_t - (1 - alpha_bar).sqrt() * e_theta) / alpha_bar.sqrt() x_t = (alpha_bar_prev.sqrt() * x0_t + (1 - alpha_bar_prev).sqrt() * e_theta) elif sampling == 'ddpm': sigma = self.var_sched.get_sigmas(t, flexibility=0.0) z = torch.randn_like(x_t) if t > stride else torch.zeros_like(x_t) c0 = 1.0 / self.var_sched.alphas[t].sqrt() c1 = (1 - self.var_sched.alphas[t]) / (1 - alpha_bar).sqrt() x_t = c0 * (x_t - c1 * e_theta) + sigma * z else: raise ValueError(f"Unknown sampling: {sampling}") traj_list.append(x_t) # [N, T, 2] return torch.stack(traj_list) # [K, N, T, 2] # --------------------------------------------------------------------------- # 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_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_data(self): train_dset = NBADatasetMID(self.args.data_dir, training=True) test_dset = NBADatasetMID(self.args.data_dir, training=False) self.train_loader = DataLoader( train_dset, batch_size=self.args.batch_size, shuffle=True, num_workers=4, collate_fn=nba_collate, pin_memory=True) self.test_loader = DataLoader( test_dset, batch_size=self.args.eval_batch_size, shuffle=False, num_workers=4, collate_fn=nba_collate, pin_memory=True) self.log.info( f"Train: {len(train_dset)} scenes Test: {len(test_dset)} scenes") def _build_model(self): # ---- Past context encoder (unchanged from main_nba_mid.py) -------- self.encoder = NBAEncoder( encoder_dim = self.args.encoder_dim, past_len = OBS_LEN, ).to(self.device) # ---- Denoising network with V6 graph interaction ------------------ net = GraphDenoiserNet( context_dim = self.args.encoder_dim, tf_layer = self.args.tf_layer, T = PRED_LEN, num_agents = NUM_AGENTS, graph_hidden = self.args.graph_hidden, top_n_neighbors = self.args.top_n_neighbors, rel_traj_hidden = self.args.rel_traj_hidden, y0_score_dim = self.args.y0_score_dim, num_gnn_layers = self.args.graph_gnn_layers, graph_dropout = self.args.graph_dropout, ) self.diffusion = DiffusionTrajGraph( net = net, var_sched = VarianceSchedule( num_steps = 100, beta_T = 5e-2, mode = 'linear', ), ).to(self.device) n_enc = sum(p.numel() for p in self.encoder.parameters()) n_back = sum(p.numel() for p in net.backbone.parameters()) n_grph = sum(p.numel() for p in net.future_graph.parameters()) n_diff = sum(p.numel() for p in self.diffusion.parameters()) self.log.info(f"Encoder params: {n_enc:,}") self.log.info(f"Backbone params: {n_back:,}") self.log.info(f"Graph params: {n_grph:,}") self.log.info(f"Total diffusion params:{n_diff:,}") def _build_optimizer(self): params = (list(self.encoder.parameters()) + list(self.diffusion.parameters())) self.optimizer = torch.optim.Adam(params, lr=self.args.lr) self.scheduler = torch.optim.lr_scheduler.ExponentialLR( self.optimizer, gamma=0.98) # ------------------------------------------------------------------ def train(self): for epoch in range(1, self.args.epochs + 1): self.encoder.train() self.diffusion.train() total_loss, count = 0.0, 0 pbar = tqdm(self.train_loader, ncols=90) for pre, fut in pbar: pre = pre.to(self.device) fut = fut.to(self.device) B = pre.size(0) past_6ch, fut_rel, mask, last_obs = preprocess_batch( pre, fut, self.device) context = self.encoder(past_6ch, mask) loss = self.diffusion.get_loss( fut_rel, context, B=B, A=NUM_AGENTS, last_obs=last_obs) self.optimizer.zero_grad() loss.backward() nn.utils.clip_grad_norm_( list(self.encoder.parameters()) + list(self.diffusion.parameters()), 1.0) self.optimizer.step() total_loss += loss.item() count += 1 pbar.set_description( f"Epoch {epoch} loss={total_loss/count:.4f}") self.scheduler.step() avg_loss = total_loss / count self.tb_log.add_scalar('loss/train', avg_loss, epoch) self.log.info(f"Epoch {epoch} train_loss={avg_loss:.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} ADE={ade:.4f} FDE={fde:.4f}") torch.save({ 'encoder': self.encoder.state_dict(), 'diffusion': self.diffusion.state_dict(), 'epoch': epoch, }, os.path.join(self.exp_dir, f'nba_epoch{epoch:04d}.pt')) @torch.no_grad() def evaluate(self): self.encoder.eval() self.diffusion.eval() ade_sum, fde_sum, n_agents = 0.0, 0.0, 0 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) past_6ch, _, mask, last_obs = preprocess_batch( pre, fut, self.device) context = self.encoder(past_6ch, mask) # [B*11, enc_dim] # pred_rel: [K, B*11, T, 2] (in ÷TRAJ_SCALE space) pred_rel = self.diffusion.sample( num_points = PRED_LEN, context = context, B = B, A = NUM_AGENTS, last_obs = last_obs, sample = K_EVAL, bestof = True, sampling = self.args.sampling, step = self.args.sampling_step, ) # absolute positions pred_abs = pred_rel * TRAJ_SCALE + last_obs.unsqueeze(0) # [K, B*11, T, 2] fut_abs = fut.reshape(B * NUM_AGENTS, PRED_LEN, 2) # [B*11, T, 2] dist = (pred_abs - fut_abs.unsqueeze(0)).norm(dim=-1) # [K, B*11, T] # minADE: pick best single trajectory (mean over T, then min over K) ade_per_mode = dist.mean(dim=-1) # [K, B*11] ade_sum += ade_per_mode.min(dim=0).values.sum().item() # minFDE: pick best mode at last frame fde_sum += dist[:, :, -1].min(dim=0).values.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_graphv6v2') # Training p.add_argument('--epochs', type=int, default=100) p.add_argument('--batch_size', type=int, default=32) p.add_argument('--eval_batch_size', type=int, default=64) p.add_argument('--lr', type=float, default=1e-3) p.add_argument('--eval_every', type=int, default=5) # MID backbone p.add_argument('--encoder_dim', type=int, default=256) p.add_argument('--tf_layer', type=int, default=3) # Graph p.add_argument('--graph_hidden', type=int, default=128) 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) # Sampling p.add_argument('--sampling', type=str, default='ddim', choices=['ddpm', 'ddim']) p.add_argument('--sampling_step', type=int, default=10) return p.parse_args() # --------------------------------------------------------------------------- # Main # --------------------------------------------------------------------------- if __name__ == '__main__': args = parse_args() trainer = Trainer(args) trainer.train()