| """ |
| main_sdd_mid_graphv5_sigma.py — MID + V6 graph + learned σ on SDD. |
| |
| SDD format: per-pedestrian (past[8,2], future[12,2], neighbors[20,N,2]). |
| Combine target + neighbors into variable-A scene, batch_size=1. |
| Graph skipped when A<2 (46% of samples have no neighbors). |
| Eval only on target agent (index 0). |
| |
| Coordinates in pixels → TRAJ_SCALE=100. |
| |
| Usage: |
| python main_sdd_mid_graphv5_sigma.py --gpu 0 |
| """ |
|
|
| import os, sys, time, pickle, logging, 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.context_encoder.mtr_encoder import SinusoidalPosEmb |
| from models.diffusion import VarianceSchedule |
| from models.common import PositionalEncoding, ConcatSquashLinear |
|
|
| OBS_LEN = 8 |
| PRED_LEN = 12 |
| TRAJ_SCALE = 100.0 |
| K_EVAL = 20 |
| HORIZONS = {'1.6s': 4, '3.2s': 8, '4.8s': 12} |
| DATA_ROOT = '/mnt/jaewoo4tb/srtp/MoFlow/data/sdd/original' |
|
|
|
|
| class SDDDataset(Dataset): |
| def __init__(self, split='train'): |
| super().__init__() |
| path = os.path.join(DATA_ROOT, f'sdd_{split}.pkl') |
| with open(path, 'rb') as f: |
| raw = pickle.load(f) |
| self.scenes = [] |
| for past, fut, neigh in raw: |
| past = past.astype(np.float32) |
| fut = fut.astype(np.float32) |
| neigh = neigh.astype(np.float32) |
| N = neigh.shape[1] |
| traj_target = np.concatenate([past, fut], axis=0)[None] |
| if N > 0: |
| traj_neigh = neigh.transpose(1, 0, 2) |
| traj_all = np.concatenate([traj_target, traj_neigh], axis=0) |
| else: |
| traj_all = traj_target |
| self.scenes.append(torch.from_numpy(traj_all)) |
| a = np.array([len(x) for x in self.scenes]) |
| print(f'[SDDDataset] {split}: {len(self.scenes)} samples, ' |
| f'A min/mean/max = {a.min()}/{a.mean():.1f}/{a.max()}') |
|
|
| def __len__(self): return len(self.scenes) |
| def __getitem__(self, i): |
| x = self.scenes[i] |
| return x[:, :OBS_LEN], x[:, OBS_LEN:] |
|
|
|
|
| def collate_bs1(batch): |
| assert len(batch) == 1 |
| return batch[0] |
|
|
|
|
| def preprocess_scene(pre, fut, device): |
| pre = pre.to(device); fut = fut.to(device) |
| last_obs = pre[:, -1:, :] |
| rel = (pre - last_obs) / TRAJ_SCALE |
| vel = torch.cat([rel[:, 1:] - rel[:, :-1], |
| torch.zeros_like(rel[:, :1])], dim=1) |
| past_6ch = torch.cat([rel, rel, vel], dim=-1) |
| fut_rel = ((fut - last_obs) / TRAJ_SCALE).contiguous() |
| A = pre.size(0) |
| mask = torch.zeros(A, A, device=device) |
| 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))) |
| _, s = self.gru(h.transpose(1, 2)) |
| return s.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 SDDEncoder(nn.Module): |
| def __init__(self, encoder_dim=256, past_len=OBS_LEN): |
| super().__init__() |
| self.ego_encoder = _STEncoder(6, 256) |
| self.social_encoder = _SocialTransformer(past_len=past_len, hidden=256) |
| self.fusion = nn.Linear(512, encoder_dim) |
| def forward(self, past_6ch, mask): |
| ego = self.ego_encoder(past_6ch) |
| soc = self.social_encoder(past_6ch.reshape(past_6ch.size(0), -1), mask) |
| return self.fusion(torch.cat([ego, soc], dim=-1)) |
|
|
|
|
| def _rebuild_graph_for_A(graph, A, max_top_n, device): |
| graph.num_agents = A |
| graph._E0 = A * (A - 1) |
| graph.top_n = max(1, min(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) |
| graph._single_edge_index = torch.tensor([src, dst], dtype=torch.long, device=device) |
|
|
|
|
| class GraphDenoiserNet(nn.Module): |
| def __init__(self, context_dim=256, tf_layer=3, T=PRED_LEN, |
| max_agents=16, graph_hidden=128, |
| top_n_neighbors=5, rel_traj_hidden=32, |
| y0_score_dim=32, num_gnn_layers=2, |
| graph_dropout=0.1): |
| super().__init__() |
| self.T, self.D = T, graph_hidden |
| self.max_top_n = top_n_neighbors |
| 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_pool_query = nn.Parameter(torch.randn(1, 1, hid) * 0.02) |
| 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 = FutureInteractionGraphV6( |
| embed_dim=graph_hidden, future_steps=T, num_agents=max_agents, |
| num_heads=4, dropout=graph_dropout, num_gnn_layers=num_gnn_layers, |
| time_dim=graph_hidden, top_n_neighbors=min(top_n_neighbors, max_agents - 1), |
| rel_traj_hidden=rel_traj_hidden, y0_score_dim=y0_score_dim) |
|
|
| self.graph_out_proj = nn.Linear(graph_hidden, hid) |
| nn.init.xavier_uniform_(self.graph_out_proj.weight, gain=0.1) |
| nn.init.zeros_(self.graph_out_proj.bias) |
| self.graph_gate = nn.Parameter(torch.tensor(0.1)) |
|
|
| self.logvar_head = nn.Sequential( |
| nn.Linear(hid, hid // 2), nn.ReLU(inplace=True), |
| nn.Linear(hid // 2, 1)) |
|
|
| def _build_ctx_emb(self, beta, context): |
| N = beta.size(0) |
| beta_v = beta.view(N, 1, 1) |
| ctx_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, ctx_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, beta, context, A, tau, |
| y_0_for_graph=None, skip_graph=False): |
| N, T, _ = x_t.shape |
| D = self.D |
| ctx_emb = self._build_ctx_emb(beta, context) |
| trans = self._encode(x_t, ctx_emb) |
| logvar = self.logvar_head(trans).squeeze(-1).clamp(min=-5, max=5) |
|
|
| if (not skip_graph) and (y_0_for_graph is not None) and A >= 2: |
| _rebuild_graph_for_A(self.future_graph, A, self.max_top_n, x_t.device) |
| y_abs = y_0_for_graph.view(1, 1, A, T, 2) |
| sigma_agent = logvar.view(1, 1, A, T) |
| attn = (self.node_pool_query * trans).sum(-1, keepdim=True).softmax(dim=1) |
| node = (trans * attn).sum(dim=1) |
| y_emb = self.node_proj(node).view(1, 1, A, D) |
| beta_scene = beta.view(1, A)[:, 0] |
| t_emb = self.time_mlp(beta_scene) |
| y_emb_out = self.future_graph( |
| y_emb, y_abs, t_emb, tau, sigma_agent=sigma_agent) |
| graph_out = self.graph_out_proj(y_emb_out.squeeze(1).reshape(N, D)) |
| trans = trans + self.graph_gate * graph_out.unsqueeze(1) |
|
|
| return self._decode(trans, ctx_emb), logvar |
|
|
|
|
| class DiffusionTrajGraph(nn.Module): |
| def __init__(self, net, var_sched, train_mode='two_pass', |
| uncertainty_weight=0.01): |
| super().__init__() |
| self.net = net; self.var_sched = var_sched |
| self.train_mode = train_mode |
| self.uncertainty_weight = uncertainty_weight |
|
|
| def get_loss(self, x_0, context, A, last_obs): |
| device = x_0.device |
| N = A |
| t_scene = self.var_sched.uniform_sample_t(1) |
| t = [t_scene[0]] * 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([t_scene[0] / self.var_sched.num_steps], |
| dtype=torch.float32, device=device) |
|
|
| if self.train_mode == 'gt': |
| y_0_abs = x_0 * TRAJ_SCALE + last_obs |
| eps_pred, logvar = self.net(x_t, beta, context, A, tau, |
| y_0_for_graph=y_0_abs) |
| else: |
| with torch.no_grad(): |
| eps_geom, _ = self.net(x_t, beta, context, A, tau, |
| skip_graph=True) |
| x_0_geom = (x_t - c1 * eps_geom) / c0 |
| y_0_abs = x_0_geom * TRAJ_SCALE + last_obs |
| eps_pred, logvar = self.net(x_t, beta, context, A, tau, |
| y_0_for_graph=y_0_abs) |
|
|
| mse = F.mse_loss(eps_pred, e_rand) |
| eps_err_sq = (eps_pred.detach() - e_rand).pow(2).mean(dim=-1) |
| nll = 0.5 * (logvar + eps_err_sq / logvar.exp()).mean() |
| return mse + self.uncertainty_weight * nll |
|
|
| @torch.no_grad() |
| def sample(self, num_points, context, A, last_obs, sample=K_EVAL, |
| bestof=True, sampling='ddim', step=10): |
| device, N = context.device, A |
| stride = self.var_sched.num_steps // step |
| out = [] |
| 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((1,), t / self.var_sched.num_steps, |
| dtype=torch.float32, device=device) |
| e_theta, _ = self.net(x_t, beta_batch, context, 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 |
| else: |
| 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 |
| out.append(x_t) |
| return torch.stack(out) |
|
|
|
|
| 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, |
| f'sdd_{time.strftime("%Y-%m-%d-%H-%M")}.log') |
| 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 = SDDDataset(split='train') |
| test_dset = SDDDataset(split='test') |
| self.train_loader = DataLoader(train_dset, batch_size=1, shuffle=True, |
| num_workers=2, collate_fn=collate_bs1) |
| self.test_loader = DataLoader(test_dset, batch_size=1, shuffle=False, |
| num_workers=2, collate_fn=collate_bs1) |
| self.log.info(f'Train={len(train_dset)} Test={len(test_dset)}') |
|
|
| def _build_model(self): |
| self.encoder = SDDEncoder(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, max_agents=self.args.max_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'), |
| train_mode=self.args.train_mode, |
| uncertainty_weight=self.args.uncertainty_weight).to(self.device) |
| n_enc = sum(p.numel() for p in self.encoder.parameters()) |
| n_diff = sum(p.numel() for p in self.diffusion.parameters()) |
| self.log.info(f'Encoder: {n_enc:,} Diffusion: {n_diff:,}') |
|
|
| def _build_optimizer(self): |
| net = self.diffusion.net |
| graph_names = {'future_graph', 'node_proj', 'time_mlp', |
| 'graph_out_proj', 'graph_gate', 'node_pool_query'} |
| graph_params, other_params = [], [] |
| for n, p in net.named_parameters(): |
| (graph_params if any(k in n for k in graph_names) else other_params).append(p) |
| self.optimizer = torch.optim.Adam([ |
| {'params': list(self.encoder.parameters()) + other_params, 'lr': self.args.lr}, |
| {'params': graph_params, 'lr': self.args.lr * self.args.graph_lr_mult}, |
| ]) |
| self.scheduler = torch.optim.lr_scheduler.CosineAnnealingLR( |
| self.optimizer, T_max=self.args.epochs, eta_min=1e-6) |
|
|
| def train(self): |
| best_ade = float('inf') |
| accum = self.args.grad_accum |
| for epoch in range(1, self.args.epochs + 1): |
| self.encoder.train(); self.diffusion.train() |
| total, count = 0.0, 0 |
| self.optimizer.zero_grad() |
| for i, (pre, fut) in enumerate(tqdm(self.train_loader, ncols=90, desc=f'E{epoch}')): |
| A = pre.size(0) |
| if A < 1: continue |
| past_6ch, fut_rel, mask, last_obs = preprocess_scene(pre, fut, self.device) |
| context = self.encoder(past_6ch, mask) |
| loss = self.diffusion.get_loss(fut_rel, context, A=A, last_obs=last_obs) |
| (loss / accum).backward() |
| if (i + 1) % accum == 0: |
| nn.utils.clip_grad_norm_( |
| list(self.encoder.parameters()) + |
| list(self.diffusion.parameters()), 1.0) |
| self.optimizer.step(); self.optimizer.zero_grad() |
| total += loss.item(); count += 1 |
| self.optimizer.step(); self.optimizer.zero_grad() |
| self.scheduler.step() |
| avg = total / max(count, 1) |
| self.tb_log.add_scalar('loss/train', avg, epoch) |
| self.log.info(f'Epoch {epoch} train_loss={avg:.4f} ' |
| f'lr={self.scheduler.get_last_lr()[0]:.6f}') |
|
|
| if epoch % self.args.eval_every == 0: |
| m = self.evaluate() |
| for k, v in m.items(): |
| self.tb_log.add_scalar(f'metric/{k}', v, epoch) |
| self.log.info('Epoch %d ' % epoch + ' '.join( |
| f'ADE({h})={m[f"ADE_{h}"]:.4f}/FDE={m[f"FDE_{h}"]:.4f}' |
| for h in HORIZONS)) |
| ade = m['ADE_4.8s'] |
| if ade < best_ade: |
| best_ade = ade |
| torch.save({'encoder': self.encoder.state_dict(), |
| 'diffusion': self.diffusion.state_dict(), |
| 'epoch': epoch, 'metrics': m}, |
| os.path.join(self.exp_dir, 'best.pt')) |
| self.log.info(f' ** New best ADE(4.8s)={ade:.4f}') |
|
|
| @torch.no_grad() |
| def evaluate(self): |
| """SDD eval: only the TARGET agent (index 0) counts.""" |
| self.encoder.eval(); self.diffusion.eval() |
| sums = {f'{k}_{h}': 0.0 for h in HORIZONS for k in ('ADE', 'FDE')} |
| n_target = 0 |
| for pre, fut in tqdm(self.test_loader, ncols=90, desc='Eval'): |
| A = pre.size(0) |
| if A < 1: continue |
| past_6ch, _, mask, last_obs = preprocess_scene(pre, fut, self.device) |
| context = self.encoder(past_6ch, mask) |
| pred_rel = self.diffusion.sample( |
| num_points=PRED_LEN, context=context, A=A, 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.to(self.device) |
| dist = (pred_abs[:, 0] - fut_abs[0].unsqueeze(0)).norm(dim=-1) |
| for h, end in HORIZONS.items(): |
| sums[f'ADE_{h}'] += dist[:, :end].mean(dim=-1).min().item() |
| sums[f'FDE_{h}'] += dist[:, end - 1].min().item() |
| n_target += 1 |
| return {k: v / n_target for k, v in sums.items()} |
|
|
|
|
| def parse_args(): |
| p = argparse.ArgumentParser() |
| p.add_argument('--exp_name', type=str, default='mid_sdd_graphv5_sigma') |
| p.add_argument('--gpu', type=int, default=0) |
| p.add_argument('--epochs', type=int, default=100) |
| p.add_argument('--grad_accum', type=int, default=32) |
| p.add_argument('--lr', type=float, default=1e-3) |
| p.add_argument('--graph_lr_mult', type=float, default=1.0) |
| p.add_argument('--eval_every', type=int, default=1) |
| p.add_argument('--encoder_dim', type=int, default=256) |
| p.add_argument('--tf_layer', type=int, default=3) |
| p.add_argument('--max_agents', type=int, default=16) |
| 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('--train_mode', type=str, default='two_pass', |
| choices=['gt', 'two_pass']) |
| p.add_argument('--sampling', type=str, default='ddim') |
| p.add_argument('--sampling_step', type=int, default=10) |
| p.add_argument('--uncertainty_weight', type=float, default=0.01) |
| return p.parse_args() |
|
|
|
|
| if __name__ == '__main__': |
| args = parse_args() |
| Trainer(args).train() |
|
|