| """ |
| MID + Graph on SDD — v4 (output-level). |
| |
| Combines v3's correct fixes with v2's stable training recipe: |
| |
| From v3 (bug fixes): |
| 1. Velocity→position: y_pos = cumsum(vel)*dt + init_pos |
| 2. Position normalization: y_abs / pos_scale (SDD pixels ~±500, scale=100) |
| 3. Trajectory-aware nodes: node_proj([context, y_vel_flat]) |
| 4. No sigma/logvar (removed untrained uncertainty head) |
| 5. Cosine LR schedule |
| |
| From v2 (stability): |
| 1. Per-agent independent diffusion t (not shared — MID batches mix agents |
| from different timesteps, shared t degrades base training) |
| 2. Two-pass training: pass1 base-only → y_0_hat, pass2 base+graph with |
| y_0_hat edges. Avoids GT train/eval gap. |
| 3. delta_scale=0.1, gate_init=0.1 (moderate — v2's 0.05 too weak, v3's 0.3 too strong) |
| 4. graph_warmup=5 epochs |
| |
| Target: beat previous graph best ADE=8.53 (baseline=8.27). |
| """ |
| import os, sys, time, logging, argparse, math, random |
| 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 |
| 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 GraphDenoiserWrapperV4(nn.Module): |
| 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, dt=0.4, pos_scale=100.0): |
| super().__init__() |
| self.base_net = base_net |
| self.pred_len = pred_len |
| self.graph_hidden = graph_hidden |
| self.max_top_n = top_n |
| self.dt = dt |
| self.pos_scale = pos_scale |
|
|
| self.node_proj = nn.Sequential( |
| nn.Linear(encoder_dim + pred_len * 2, 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.Tanh()) |
| nn.init.zeros_(self.graph_out_proj[-2].weight) |
| nn.init.zeros_(self.graph_out_proj[-2].bias) |
| gi = 0.1 |
| self.raw_gate = nn.Parameter(torch.tensor(math.log(gi / (1.0 - gi)))) |
| self.register_buffer('delta_scale', torch.tensor(0.1)) |
|
|
| def forward(self, x_t, beta, context, |
| y_vel=None, init_pos=None, skip_graph=False): |
| """ |
| x_t: [N, T, 2] noisy trajectory (velocity space) |
| beta: [N] diffusion beta |
| context: [N, enc_dim] encoder output |
| y_vel: [N, T, 2] velocity for graph edge features (y_0_hat from pass1) |
| init_pos: [N, 2] last observed absolute position |
| """ |
| 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_vel is not None and init_pos is not None and N >= 2: |
| D = self.graph_hidden |
|
|
| y_vel_flat = y_vel.view(N, T * 2) |
| node_input = torch.cat([context, y_vel_flat], dim=-1) |
| node_emb = self.node_proj(node_input).view(1, 1, N, D) |
|
|
| y_pos = (torch.cumsum(y_vel.view(N, T, 2), dim=1) * self.dt |
| + init_pos.unsqueeze(1)) |
| y_abs = (y_pos / self.pos_scale).view(1, 1, N, T, 2) |
|
|
| beta_scene = beta.mean().unsqueeze(0) |
| 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=None) |
| graph_out = self.graph_out_proj( |
| y_emb_out.squeeze(1).squeeze(0)) |
| delta = graph_out.view(N, T, 2) * self.delta_scale |
| gate = torch.sigmoid(self.raw_gate) |
| eps_pred = eps_pred + gate * delta |
|
|
| return eps_pred |
|
|
|
|
| class MIDGraphV4: |
| def __init__(self, config): |
| self.config = config |
| torch.backends.cudnn.benchmark = True |
| self._build() |
|
|
| def _build(self): |
| self._skip_graph_override = False |
| 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 = GraphDenoiserWrapperV4( |
| 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, |
| dt=self.config.dt, |
| pos_scale=self.config.pos_scale).cuda() |
| if hasattr(self.config, 'graph_gate_init') and self.config.graph_gate_init is not None: |
| gi = float(max(min(self.config.graph_gate_init, 0.999), 1e-4)) |
| with torch.no_grad(): |
| self.graph_net.raw_gate.fill_(math.log(gi / (1.0 - gi))) |
|
|
| self.var_sched = VarianceSchedule(num_steps=100, beta_T=5e-2, mode='linear') |
|
|
| graph_keys = ('future_graph', 'node_proj', 'time_mlp', |
| 'graph_out_proj', 'raw_gate') |
| self._graph_params = [p for n, p in self.graph_net.named_parameters() |
| if any(k in n for k in graph_keys)] |
| self._base_params = [p for n, p in self.graph_net.named_parameters() |
| if not any(k in n for k in graph_keys)] |
|
|
| self.optimizer = optim.Adam([ |
| {'params': self.registrar.get_all_but_name_match('map_encoder').parameters()}, |
| {'params': self.graph_net.parameters()}, |
| ], lr=self.config.lr) |
| warm = max(1, int(getattr(self.config, 'lr_warmup_epochs', 2))) |
| from torch.optim.lr_scheduler import LambdaLR, CosineAnnealingLR, SequentialLR |
| warm_sched = LambdaLR(self.optimizer, |
| lr_lambda=lambda e: min(1.0, (e + 1) / warm)) |
| cosine_sched = CosineAnnealingLR( |
| self.optimizer, T_max=self.config.epochs - warm, eta_min=1e-5) |
| self.scheduler = SequentialLR( |
| self.optimizer, schedulers=[warm_sched, cosine_sched], milestones=[warm]) |
|
|
| self.train_scenes = self.train_env.scenes |
| self.eval_scenes = self.eval_env.scenes |
| self.log.info(f"Train scenes: {len(self.train_scenes)}, " |
| f"Eval scenes: {len(self.eval_scenes)}") |
|
|
| def _get_loss(self, batch, node_type): |
| (first_history_index, x_t_raw, 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) |
| y_0 = y_t.cuda() |
| N = y_0.size(0) |
| init_pos = x_t_raw[:, -1, 0:2].cuda() |
|
|
| |
| 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 self._skip_graph_override or N < 2: |
| eps_pred = self.graph_net(x_noisy, beta, context, skip_graph=True) |
| else: |
| |
| with torch.no_grad(): |
| eps_base = self.graph_net(x_noisy, beta, context, skip_graph=True) |
| y_0_hat = (x_noisy - c1 * eps_base) / c0 |
| eps_pred = self.graph_net(x_noisy, beta, context, |
| y_vel=y_0_hat, init_pos=init_pos) |
|
|
| 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'] |
| graph_warm = int(getattr(self.config, 'graph_warmup_epochs', 5)) |
|
|
| for epoch in range(1, self.config.epochs + 1): |
| self.graph_net.train() |
| total_loss, n_batches = 0.0, 0 |
| self._skip_graph_override = (epoch <= graph_warm) |
|
|
| for scene in self.train_scenes: |
| for t_start in range(0, scene.timesteps, 10): |
| timesteps = np.arange(t_start, t_start + 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_params, 0.1) |
| nn.utils.clip_grad_norm_(self._base_params, 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: " |
| f"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_start in range(0, scene.timesteps, 10): |
| timesteps = np.arange(t_start, t_start + 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) |
|
|
| _, x_t_raw, *_ = test_batch |
| init_pos = x_t_raw[:, -1, 0:2].cuda() |
|
|
| preds = self._sample_with_graph( |
| context, N, init_pos, 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, init_pos, |
| num_points=12, K=20): |
| traj_list = [] |
| stride = 5 |
| 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_vel=y_0_prev, init_pos=init_pos) |
| 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) |
|
|
|
|
| def main(): |
| p = argparse.ArgumentParser() |
| p.add_argument('--data_dir', default='processed_data') |
| p.add_argument('--exp_name', default='mid_sdd_graph_v4') |
| p.add_argument('--gpu', type=int, default=0) |
| p.add_argument('--epochs', type=int, default=100) |
| p.add_argument('--lr', type=float, default=1e-3) |
| p.add_argument('--eval_every', type=int, default=3) |
| 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) |
| p.add_argument('--graph_warmup_epochs', type=int, default=5) |
| p.add_argument('--lr_warmup_epochs', type=int, default=2) |
| p.add_argument('--dt', type=float, default=0.4, |
| help='Scene timestep (SDD: 0.4s at 2.5Hz)') |
| p.add_argument('--pos_scale', type=float, default=100.0, |
| help='Divide positions by this before graph') |
| config = p.parse_args() |
| torch.cuda.set_device(config.gpu) |
| MIDGraphV4(config).train() |
|
|
|
|
| if __name__ == '__main__': |
| main() |
|
|