| """ |
| main_nba_mid_graphv6_v3.py — MID on NBA with V6-style future interaction graph. |
| |
| V3: Follows MoFlow V7 protocol. |
| |
| TRAINING — single pass with GT teacher forcing: |
| GraphDenoiserNet.forward(x_t, ..., y_0_for_graph=GT_abs) |
| → encoder → graph(GT geometry) enriches intermediate → decoder → eps |
| No wasted no_grad pass. |
| |
| INFERENCE — chaining across DDIM steps: |
| Step 0: forward(skip_graph=True) → eps → derive y_0 → save as y_0_prev |
| Step 1+: forward(y_0_for_graph=y_0_prev) → eps → derive y_0 → update y_0_prev |
| |
| GRAPH — same as MoFlow V6: |
| FutureInteractionGraphV6 with internal gated residual: |
| out = y_emb + sigmoid(gate_proj([y_emb, GNN_nodes])) * out_proj(GNN_nodes) |
| Graph output REPLACES y_emb (not adds delta to final output). |
| Enriched y_emb flows through the output head to produce final prediction. |
| |
| Usage: |
| python main_nba_mid_graphv6_v3.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 |
| 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 models.graph_interaction_nba_v6 import FutureInteractionGraphV6 |
| from models.interaction_baselines import build_interaction_module |
| from models.context_encoder.mtr_encoder import SinusoidalPosEmb |
|
|
| |
| from models.diffusion import VarianceSchedule |
| from models.common import PositionalEncoding, ConcatSquashLinear |
|
|
|
|
| |
| |
| |
|
|
| OBS_LEN = 10 |
| PRED_LEN = 20 |
| NUM_AGENTS = 11 |
| TRAJ_SCALE = 5.0 |
| TRAJ_MEAN = torch.FloatTensor([14.0, 7.5]) |
| K_EVAL = 20 |
|
|
|
|
| |
| |
| |
|
|
| 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) |
| 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(pre_motion, fut_motion, device): |
| 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 |
|
|
|
|
| |
| |
| |
|
|
| 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)) |
|
|
|
|
| |
| |
| |
|
|
| class GraphDenoiserNet(nn.Module): |
| """TransformerConcatLinear with V6 graph inserted mid-network. |
| |
| Single forward pass. Graph geometry provided externally: |
| - Training: GT future (teacher forcing) |
| - Inference: chained y_0_prev from previous sampling step |
| |
| Architecture: |
| Encoder: concat1 → pos_emb → transformer_encoder → trans [N, T, hid] |
| Graph: pool(trans) → node_proj → V6 graph(y_0_for_graph) → graph_out_proj |
| trans = trans + graph_enrichment (broadcast over T) |
| Decoder: concat3 → concat4 → linear → eps [N, T, 2] |
| |
| The V6 graph has internal gated residual: |
| out = y_emb + sigmoid(gate([y_emb, GNN_nodes])) * out_proj(GNN_nodes) |
| """ |
|
|
| 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 |
|
|
| hid = 2 * context_dim |
| ctx = context_dim + 3 |
|
|
| |
| self.pos_emb = PositionalEncoding(d_model=hid, dropout=0.1, max_len=24) |
| self.concat1 = ConcatSquashLinear(2, hid, ctx) |
| layer = nn.TransformerEncoderLayer( |
| d_model=hid, nhead=4, dim_feedforward=4 * context_dim) |
| self.transformer_encoder = nn.TransformerEncoder( |
| layer, num_layers=tf_layer) |
|
|
| |
| self.concat3 = ConcatSquashLinear(hid, context_dim, ctx) |
| self.concat4 = ConcatSquashLinear(context_dim, context_dim // 2, ctx) |
| self.out_linear = ConcatSquashLinear(context_dim // 2, 2, ctx) |
|
|
| |
| self.node_proj = nn.Sequential( |
| nn.Linear(hid, graph_hidden), |
| nn.ReLU(inplace=True), |
| ) |
|
|
| |
| self.time_mlp = nn.Sequential( |
| SinusoidalPosEmb(graph_hidden), |
| nn.Linear(graph_hidden, graph_hidden), |
| nn.ReLU(), |
| nn.Linear(graph_hidden, graph_hidden), |
| ) |
|
|
| |
| self.future_graph = build_interaction_module( |
| 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, |
| ) |
|
|
| |
| |
| self.graph_out_proj = nn.Linear(graph_hidden, hid) |
| nn.init.zeros_(self.graph_out_proj.weight) |
| nn.init.zeros_(self.graph_out_proj.bias) |
|
|
| def _build_ctx_emb(self, beta, context): |
| N = beta.size(0) |
| beta_v = beta.view(N, 1, 1) |
| context_v = context.view(N, 1, -1) |
| time_emb = torch.cat([beta_v, torch.sin(beta_v), torch.cos(beta_v)], dim=-1) |
| return torch.cat([time_emb, context_v], dim=-1) |
|
|
| def _encode(self, x_t, ctx_emb): |
| h = self.concat1(ctx_emb, x_t) |
| h = self.pos_emb(h.permute(1, 0, 2)) |
| return self.transformer_encoder(h).permute(1, 0, 2) |
|
|
| def _decode(self, trans, ctx_emb): |
| h = self.concat3(ctx_emb, trans) |
| h = self.concat4(ctx_emb, h) |
| return self.out_linear(ctx_emb, h) |
|
|
| def forward(self, |
| x_t: torch.Tensor, |
| beta: torch.Tensor, |
| context: torch.Tensor, |
| B: int, |
| A: int, |
| tau: torch.Tensor, |
| y_0_for_graph: torch.Tensor = None, |
| skip_graph: bool = False, |
| ) -> torch.Tensor: |
| N, T, _ = x_t.shape |
| D = self.D |
|
|
| ctx_emb = self._build_ctx_emb(beta, context) |
|
|
| |
| trans = self._encode(x_t, ctx_emb) |
|
|
| |
| if not skip_graph and y_0_for_graph is not None: |
| y_abs = y_0_for_graph.view(B, A, T, 2).unsqueeze(1) |
|
|
| |
| node_emb = self.node_proj(trans.mean(dim=1)) |
| y_emb = node_emb.view(B, 1, A, D) |
|
|
| |
| beta_scene = beta.view(B, A)[:, 0] |
| t_emb = self.time_mlp(beta_scene) |
|
|
| |
| y_emb_graph = self.future_graph( |
| y_emb, y_abs, t_emb, tau, sigma_agent=None) |
|
|
| |
| graph_out = self.graph_out_proj( |
| y_emb_graph.squeeze(1).reshape(N, D)) |
| trans = trans + graph_out.unsqueeze(1) |
|
|
| |
| return self._decode(trans, ctx_emb) |
|
|
|
|
| |
| |
| |
|
|
| class DiffusionTrajGraph(nn.Module): |
| """DDPM wrapper — epsilon prediction with V7 train/inference protocol.""" |
|
|
| def __init__(self, net: GraphDenoiserNet, var_sched: VarianceSchedule): |
| super().__init__() |
| self.net = net |
| self.var_sched = var_sched |
|
|
| def get_loss(self, x_0, context, B, A, last_obs): |
| """Training loss — single pass with GT teacher forcing.""" |
| device = x_0.device |
| N = B * A |
|
|
| |
| t_scene = self.var_sched.uniform_sample_t(B) |
| t = [ts for ts in t_scene for _ in range(A)] |
|
|
| alpha_bar = self.var_sched.alpha_bars[t].to(device) |
| beta = self.var_sched.betas[t].to(device) |
|
|
| c0 = alpha_bar.sqrt().view(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 |
|
|
| tau = torch.tensor( |
| [ts / self.var_sched.num_steps for ts in t_scene], |
| dtype=torch.float32, device=device) |
|
|
| |
| y_0_gt_abs = x_0 * TRAJ_SCALE + last_obs |
|
|
| eps_pred = self.net(x_t, beta, context, B, A, tau, |
| y_0_for_graph=y_0_gt_abs) |
| return F.mse_loss(eps_pred, e_rand) |
|
|
| @torch.no_grad() |
| def sample(self, num_points, context, B, A, last_obs, |
| sample=K_EVAL, bestof=True, sampling='ddim', step=10): |
| """Inference — chaining y_0_prev across DDIM steps.""" |
| 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)) |
|
|
| y_0_prev = None |
|
|
| 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_batch = self.var_sched.betas[t].to(device).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, tau, |
| y_0_for_graph=y_0_prev, |
| skip_graph=(y_0_prev is None)) |
|
|
| |
| x_0_pred = (x_t - (1 - alpha_bar).sqrt() * e_theta) / alpha_bar.sqrt() |
| y_0_prev = x_0_pred * TRAJ_SCALE + last_obs |
|
|
| |
| if sampling == 'ddim': |
| x_t = (alpha_bar_prev.sqrt() * x_0_pred |
| + (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 |
|
|
| traj_list.append(x_t) |
|
|
| return torch.stack(traj_list) |
|
|
|
|
| |
| |
| |
|
|
| 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): |
| self.encoder = NBAEncoder( |
| encoder_dim = self.args.encoder_dim, |
| past_len = OBS_LEN, |
| ).to(self.device) |
|
|
| 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_enc_part = sum(p.numel() for n, p in net.named_parameters() |
| if not any(k in n for k in ['future_graph', 'node_proj', |
| 'time_mlp', 'graph_out_proj'])) |
| n_grph = (sum(p.numel() for p in net.future_graph.parameters()) + |
| sum(p.numel() for p in net.node_proj.parameters()) + |
| sum(p.numel() for p in net.time_mlp.parameters()) + |
| sum(p.numel() for p in net.graph_out_proj.parameters())) |
| self.log.info(f"Encoder params: {n_enc:,}") |
| self.log.info(f"Backbone params: {n_enc_part:,}") |
| self.log.info(f"Graph params: {n_grph:,}") |
| self.log.info(f"Total diffusion params:{sum(p.numel() for p in self.diffusion.parameters()):,}") |
|
|
| 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) |
|
|
| 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, |
| ) |
|
|
| pred_abs = pred_rel * TRAJ_SCALE + last_obs.unsqueeze(0) |
| fut_abs = fut.reshape(B * NUM_AGENTS, PRED_LEN, 2) |
|
|
| dist = (pred_abs - fut_abs.unsqueeze(0)).norm(dim=-1) |
|
|
| |
| ade_per_mode = dist.mean(dim=-1) |
| ade_sum += ade_per_mode.min(dim=0).values.sum().item() |
|
|
| |
| 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 |
|
|
|
|
| |
| |
| |
|
|
| 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_graphv6v3') |
| 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) |
| p.add_argument('--encoder_dim', type=int, default=256) |
| p.add_argument('--tf_layer', type=int, default=3) |
| 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) |
| p.add_argument('--sampling', type=str, default='ddim', |
| choices=['ddpm', 'ddim']) |
| p.add_argument('--sampling_step', type=int, default=10) |
| return p.parse_args() |
|
|
|
|
| if __name__ == '__main__': |
| args = parse_args() |
| trainer = Trainer(args) |
| trainer.train() |
|
|