""" main_sdd_mid.py — MID baseline on SDD (Stanford Drone Dataset). Uses MoFlow's original pkl at `MoFlow/data/sdd/original/sdd_{train,test}.pkl`. Each sample is a tuple (past[8,2], future[12,2], neighbors[20,N,2]) with variable N. We combine target + neighbors into a scene of A=1+N agents. batch_size=1 (variable A), gradient accumulation. Standard SDD: 8 past + 12 future frames. Coordinates in pixels. We normalize per-agent last-obs-relative and divide by TRAJ_SCALE=100 (rough pixel scale) to keep values small. Usage: python main_sdd_mid.py --gpu 0 """ import os, sys, time, pickle, logging, argparse import numpy as np import torch import torch.nn as nn from torch import optim from torch.utils.data import Dataset, DataLoader from torch.utils.tensorboard import SummaryWriter # tbX-broken from tqdm.auto import tqdm from models.diffusion import DiffusionTraj, VarianceSchedule, TransformerConcatLinear 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): """Per-pedestrian dataset: each item is target + neighbors as a scene.""" 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) # [8, 2] fut = fut.astype(np.float32) # [12, 2] neigh = neigh.astype(np.float32) # [20, N, 2] N = neigh.shape[1] traj_target = np.concatenate([past, fut], axis=0)[None] # [1, 20, 2] if N > 0: traj_neigh = neigh.transpose(1, 0, 2) # [N, 20, 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, 20, 2] 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)) 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 = TransformerConcatLinear(point_dim=2, context_dim=self.args.encoder_dim, tf_layer=self.args.tf_layer, residual=False) self.diffusion = DiffusionTraj( 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_diff = sum(p.numel() for p in self.diffusion.parameters()) self.log.info(f'Encoder: {n_enc:,} Diffusion: {n_diff:,}') def _build_optimizer(self): params = list(self.encoder.parameters()) + list(self.diffusion.parameters()) self.optimizer = optim.Adam(params, lr=self.args.lr) self.scheduler = optim.lr_scheduler.ExponentialLR(self.optimizer, gamma=0.98) 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, _ = preprocess_scene(pre, fut, self.device) context = self.encoder(past_6ch, mask) loss = self.diffusion.get_loss(fut_rel, context) (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}') 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( f'Epoch {epoch} ADE(4.8s)={m["ADE_4.8s"]:.4f} FDE(4.8s)={m["FDE_4.8s"]:.4f}' f' ADE(1.6s)={m["ADE_1.6s"]:.4f} ADE(3.2s)={m["ADE_3.2s"]:.4f}') 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 toward ADE/FDE.""" 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, 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) # Only target agent (index 0) dist = (pred_abs[:, 0] - fut_abs[0].unsqueeze(0)).norm(dim=-1) # [K, T] 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_baseline') 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('--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('--sampling', type=str, default='ddim') p.add_argument('--sampling_step', type=int, default=10) return p.parse_args() if __name__ == '__main__': args = parse_args() Trainer(args).train()