""" MID + Graph (sigma) on SDD, using the original Trajectron++ data pipeline. Training iterates per-scene to preserve agent grouping for the graph module. Eval follows the original MID protocol (ADE/FDE × 50). """ import os, sys, time, logging, argparse import numpy as np import torch import torch.nn as nn import torch.nn.functional as F from torch import optim from torch.utils.tensorboard import SummaryWriter # tbX-broken from tqdm.auto import tqdm import dill from dataset import EnvironmentDataset, collate, get_timesteps_data, restore from models.autoencoder import AutoEncoder from models.trajectron import Trajectron from utils.model_registrar import ModelRegistrar from utils.trajectron_hypers import get_traj_hypers from models.diffusion import DiffusionTraj, VarianceSchedule, TransformerConcatLinear import evaluation 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 class GraphDenoiserWrapper(nn.Module): """Wraps the base TransformerConcatLinear denoiser + adds graph module. During forward: two-pass (skip_graph → with_graph). Graph operates on y0_hat estimates with scene-level agent grouping.""" def __init__(self, base_net, encoder_dim=256, pred_len=12, graph_hidden=128, top_n=5, num_gnn_layers=2, graph_dropout=0.1): super().__init__() self.base_net = base_net self.pred_len = pred_len self.graph_hidden = graph_hidden self.max_top_n = top_n self.node_proj = nn.Sequential( nn.Linear(encoder_dim, 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 = FutureInteractionGraphV6( embed_dim=graph_hidden, future_steps=pred_len, num_agents=64, num_heads=4, dropout=graph_dropout, num_gnn_layers=num_gnn_layers, time_dim=graph_hidden, top_n_neighbors=min(top_n, 63), rel_traj_hidden=32, y0_score_dim=32) self.graph_out_proj = nn.Sequential( nn.Linear(graph_hidden, graph_hidden), nn.ReLU(inplace=True), nn.Linear(graph_hidden, pred_len * 2)) nn.init.zeros_(self.graph_out_proj[-1].weight) nn.init.zeros_(self.graph_out_proj[-1].bias) self.graph_gate = nn.Parameter(torch.tensor(0.1)) self.logvar_head = nn.Sequential( nn.Linear(encoder_dim, encoder_dim // 2), nn.ReLU(inplace=True), nn.Linear(encoder_dim // 2, 1)) def _rebuild_graph(self, A, device): self.future_graph.num_agents = A self.future_graph._E0 = A * (A - 1) self.future_graph.top_n = max(1, min(self.max_top_n, A - 1)) src, dst = [], [] for i in range(A): for j in range(A): if i != j: src.append(j); dst.append(i) self.future_graph._single_edge_index = torch.tensor( [src, dst], dtype=torch.long, device=device) def forward(self, x_t, beta, context, y_0_for_graph=None, skip_graph=False): """x_t: [N, T, 2], beta: [N], context: [N, encoder_dim]""" eps_pred = self.base_net(x_t, beta=beta, context=context) N = x_t.size(0) T = self.pred_len if (not skip_graph) and y_0_for_graph is not None and N >= 2: self._rebuild_graph(N, x_t.device) D = self.graph_hidden node_emb = self.node_proj(context).view(1, 1, N, D) y_abs = y_0_for_graph.view(1, 1, N, T, 2) logvar = self.logvar_head(context).clamp(-5, 5) # [N, 1] sigma_agent = logvar.view(1, 1, N, 1).expand(-1, -1, -1, T) beta_scene = beta[:1] tau = (beta_scene / 0.05).clamp(0, 1) t_emb = self.time_mlp(beta_scene) y_emb_out = self.future_graph( node_emb, y_abs, t_emb, tau, sigma_agent=sigma_agent) graph_out = self.graph_out_proj( y_emb_out.squeeze(1).squeeze(0)) # [N, T*2] delta = graph_out.view(N, T, 2) eps_pred = eps_pred + self.graph_gate * delta return eps_pred class MIDGraph: def __init__(self, config): self.config = config torch.backends.cudnn.benchmark = True self._build() def _build(self): self.model_dir = os.path.join("./experiments", self.config.exp_name) self.log_writer = SummaryWriter(log_dir=self.model_dir) os.makedirs(self.model_dir, exist_ok=True) log_name = f"sdd_{time.strftime('%Y-%m-%d-%H-%M')}.log" self.log = logging.getLogger(self.config.exp_name) self.log.setLevel(logging.INFO) self.log.addHandler(logging.FileHandler(os.path.join(self.model_dir, log_name))) self.log.addHandler(logging.StreamHandler()) self.log.info(f"Config: {self.config}") self.train_data_path = os.path.join(self.config.data_dir, "sdd_train.pkl") self.eval_data_path = os.path.join(self.config.data_dir, "sdd_test.pkl") self.hyperparams = get_traj_hypers() self.hyperparams['enc_rnn_dim_edge'] = self.config.encoder_dim // 2 self.hyperparams['enc_rnn_dim_edge_influence'] = self.config.encoder_dim // 2 self.hyperparams['enc_rnn_dim_history'] = self.config.encoder_dim // 2 self.hyperparams['enc_rnn_dim_future'] = self.config.encoder_dim // 2 self.registrar = ModelRegistrar(self.model_dir, "cuda") with open(self.train_data_path, 'rb') as f: self.train_env = dill.load(f, encoding='latin1') with open(self.eval_data_path, 'rb') as f: self.eval_env = dill.load(f, encoding='latin1') self.encoder = Trajectron(self.registrar, self.hyperparams, "cuda") self.encoder.set_environment(self.train_env) self.encoder.set_annealing_params() base_net = TransformerConcatLinear( point_dim=2, context_dim=self.config.encoder_dim, tf_layer=self.config.tf_layer, residual=False) self.graph_net = GraphDenoiserWrapper( base_net, encoder_dim=self.config.encoder_dim, pred_len=12, graph_hidden=128, top_n=self.config.top_n_neighbors, num_gnn_layers=self.config.graph_gnn_layers, graph_dropout=self.config.graph_dropout).cuda() # Optional: override graph gate init and dropout if configured if hasattr(self.config, 'graph_gate_init'): with torch.no_grad(): self.graph_net.graph_gate.fill_(self.config.graph_gate_init) self.var_sched = VarianceSchedule(num_steps=100, beta_T=5e-2, mode='linear') self.optimizer = optim.Adam([ {'params': self.registrar.get_all_but_name_match('map_encoder').parameters()}, {'params': self.graph_net.parameters()}, ], lr=self.config.lr) self.scheduler = optim.lr_scheduler.ExponentialLR(self.optimizer, gamma=0.98) self.train_scenes = self.train_env.scenes self.eval_scenes = self.eval_env.scenes self.log.info(f"Train scenes: {len(self.train_scenes)}, Eval scenes: {len(self.eval_scenes)}") def _get_loss(self, batch, node_type): (first_history_index, x_t, y_t, x_st_t, y_st_t, neighbors_data_st, neighbors_edge_value, robot_traj_st_t, map_) = batch context = self.encoder.get_latent(batch, node_type) # [N, enc_dim] y_0 = y_t.cuda() # [N, 12, 2] N = y_0.size(0) t = self.var_sched.uniform_sample_t(N) alpha_bar = self.var_sched.alpha_bars[t].cuda() beta = self.var_sched.betas[t].cuda() c0 = alpha_bar.sqrt().view(N, 1, 1) c1 = (1 - alpha_bar).sqrt().view(N, 1, 1) e_rand = torch.randn_like(y_0) x_noisy = c0 * y_0 + c1 * e_rand if N >= 2: with torch.no_grad(): eps_geom = self.graph_net(x_noisy, beta, context, skip_graph=True) y_0_hat = (x_noisy - c1 * eps_geom) / c0 eps_pred = self.graph_net(x_noisy, beta, context, y_0_for_graph=y_0_hat) else: eps_pred = self.graph_net(x_noisy, beta, context, skip_graph=True) return F.mse_loss(eps_pred.reshape(-1, 2), e_rand.reshape(-1, 2)) def train(self): node_type = "PEDESTRIAN" ph = self.hyperparams['prediction_horizon'] max_hl = self.hyperparams['maximum_history_length'] for epoch in range(1, self.config.epochs + 1): self.graph_net.train() total_loss, n_batches = 0.0, 0 for scene in self.train_scenes: for t in range(0, scene.timesteps, 10): timesteps = np.arange(t, t + 10) batch = get_timesteps_data( env=self.train_env, scene=scene, t=timesteps, node_type=node_type, state=self.hyperparams['state'], pred_state=self.hyperparams['pred_state'], edge_types=self.train_env.get_edge_types(), min_ht=1, max_ht=max_hl, min_ft=12, max_ft=12, hyperparams=self.hyperparams) if batch is None: continue loss = self._get_loss(batch[0], node_type) self.optimizer.zero_grad() loss.backward() nn.utils.clip_grad_norm_(self.graph_net.parameters(), 1.0) self.optimizer.step() total_loss += loss.item(); n_batches += 1 self.scheduler.step() avg = total_loss / max(1, n_batches) self.log.info(f"Epoch {epoch} train_loss={avg:.4f}") self.log_writer.add_scalar('loss/train', avg, epoch) if epoch % self.config.eval_every == 0: ade, fde = self._eval(node_type, ph, max_hl) ade *= 50; fde *= 50 self.log.info(f"Epoch {epoch} Best Of 20: ADE: {ade:.4f} FDE: {fde:.4f}") self.log_writer.add_scalar('metric/ADE', ade, epoch) self.log_writer.add_scalar('metric/FDE', fde, epoch) torch.save({ 'encoder': self.registrar.model_dict, 'graph_net': self.graph_net.state_dict(), }, os.path.join(self.model_dir, f"sdd_epoch{epoch}.pt")) @torch.no_grad() def _eval(self, node_type, ph, max_hl): self.graph_net.eval() ade_errors, fde_errors = [], [] for scene in self.eval_scenes: for t in range(0, scene.timesteps, 10): timesteps = np.arange(t, t + 10) batch = get_timesteps_data( env=self.eval_env, scene=scene, t=timesteps, node_type=node_type, state=self.hyperparams['state'], pred_state=self.hyperparams['pred_state'], edge_types=self.eval_env.get_edge_types(), min_ht=7, max_ht=max_hl, min_ft=12, max_ft=12, hyperparams=self.hyperparams) if batch is None: continue test_batch, nodes, timesteps_o = batch context = self.encoder.get_latent(test_batch, node_type) dynamics = self.encoder.node_models_dict[node_type].dynamic N = context.size(0) # Sample K=20 trajectories with graph preds = self._sample_with_graph(context, N, num_points=12, K=20) predicted_y_pos = dynamics.integrate_samples(preds) predictions = predicted_y_pos.cpu().numpy() predictions_dict = {} for i, ts in enumerate(timesteps_o): if ts not in predictions_dict: predictions_dict[ts] = {} predictions_dict[ts][nodes[i]] = np.transpose(predictions[:, [i]], (1, 0, 2, 3)) batch_error = evaluation.compute_batch_statistics( predictions_dict, scene.dt, max_hl=max_hl, ph=ph, node_type_enum=self.eval_env.NodeType, kde=False, map=None, best_of=True, prune_ph_to_future=True) ade_errors = np.hstack((ade_errors, batch_error[node_type]['ade'])) fde_errors = np.hstack((fde_errors, batch_error[node_type]['fde'])) return np.mean(ade_errors), np.mean(fde_errors) def _sample_with_graph(self, context, N, num_points=12, K=20): traj_list = [] stride = 5 # 100/20 = 5 steps (ddim-like with 20 steps) for _ in range(K): x_t = torch.randn(N, num_points, 2, device=context.device) y_0_prev = None for t in range(self.var_sched.num_steps, 0, -stride): alpha_bar = self.var_sched.alpha_bars[t] alpha_bar_next = self.var_sched.alpha_bars[t - stride] beta = self.var_sched.betas[[t] * N].cuda() if y_0_prev is not None and N >= 2: eps = self.graph_net(x_t, beta, context, y_0_for_graph=y_0_prev) else: eps = self.graph_net(x_t, beta, context, skip_graph=True) x0_pred = (x_t - (1 - alpha_bar).sqrt() * eps) / alpha_bar.sqrt() y_0_prev = x0_pred x_t = alpha_bar_next.sqrt() * x0_pred + (1 - alpha_bar_next).sqrt() * eps traj_list.append(x_t) return torch.stack(traj_list) # [K, N, T, 2] def main(): p = argparse.ArgumentParser() p.add_argument('--data_dir', default='processed_data') p.add_argument('--exp_name', default='mid_sdd_graph_sigma') p.add_argument('--gpu', type=int, default=0) p.add_argument('--epochs', type=int, default=90) p.add_argument('--lr', type=float, default=1e-3) p.add_argument('--eval_every', type=int, default=30) p.add_argument('--encoder_dim', type=int, default=256) p.add_argument('--tf_layer', type=int, default=3) p.add_argument('--top_n_neighbors', type=int, default=5) p.add_argument('--graph_gnn_layers', type=int, default=2) p.add_argument('--graph_dropout', type=float, default=0.1) p.add_argument('--graph_gate_init', type=float, default=0.1) config = p.parse_args() torch.cuda.set_device(config.gpu) MIDGraph(config).train() if __name__ == '__main__': main()