| """ |
| All-agents SDD evaluation for LED checkpoints — matches MoFlow's protocol. |
| |
| Protocol (mirrors MoFlow trainer's compute_ADE_FDE): |
| For each scene with A agents, run LED sampling → pred [A, K=20, T, 2]. |
| Per-horizon buckets: 1.2s (frame 3), 2.4s (6), 3.6s (9), 4.8s (12). |
| ADE_min(H) = mean_{t=1..H} ‖pred - gt‖ → min over K → sum over A agents |
| FDE_min(H) = ‖pred[H-1] - gt[H-1]‖ → min over K → sum over A agents |
| ADE_avg(H) = mean over K instead of min |
| Report in pixels (× 50). |
| |
| Usage: |
| python eval_sdd_led_allagents.py --exp baseline_v2 --epoch 40 |
| python eval_sdd_led_allagents.py --exp graph_sigma_v2 --epoch 28 --use_graph --use_v6_graph |
| """ |
| import argparse, os, sys, random, torch, numpy as np |
| from torch.utils.data import DataLoader |
| from data.dataloader_sdd import SDDDataset, sdd_seq_collate |
| from models.model_led_initializer import LEDInitializer as InitializationModel |
| from models.model_diffusion import TransformerDenoisingModel as CoreDenoisingModel |
| from trainer.train_sdd_led import NUM_Tau |
| from utils.config import Config |
|
|
|
|
| def build_models(cfg, use_graph, use_v6_graph, ckpt_path, device): |
| model = CoreDenoisingModel(past_len=cfg.past_frames).to(device) |
| core_ckpt = torch.load(cfg.pretrained_core_denoising_model, map_location='cpu') |
| model.load_state_dict(core_ckpt['model_dict']) |
| model.eval() |
|
|
| init = InitializationModel( |
| t_h=cfg.past_frames, d_h=6, |
| t_f=cfg.future_frames, d_f=2, k_pred=20).to(device) |
|
|
| ckpt = torch.load(ckpt_path, map_location='cpu') |
| init.load_state_dict(ckpt['model_initializer_dict']) |
| init.eval() |
|
|
| graph = None |
| if use_graph: |
| from models.future_interaction_graph_v6 import FutureInteractionGraphV6Wrapper |
| graph = FutureInteractionGraphV6Wrapper( |
| num_agents=64, future_steps=cfg.future_frames, |
| past_steps=cfg.past_frames, past_channels=6, |
| node_dim=128, top_n=5, num_denoise_steps=NUM_Tau).to(device) |
| sd = {k: v for k, v in ckpt['interaction_graph_dict'].items() |
| if '_single_edge_index' not in k} |
| graph.load_state_dict(sd, strict=False) |
| graph.eval() |
|
|
| return model, init, graph |
|
|
|
|
| def make_beta_schedule(n=100, start=1e-4, end=5e-2): |
| return torch.linspace(start, end, n) |
|
|
|
|
| def extract(a, t, x): |
| out = torch.gather(a, 0, t.to(a.device)) |
| return out.reshape(t.shape[0], *([1] * (len(x.shape) - 1))) |
|
|
|
|
| @torch.no_grad() |
| def p_sample_accelerate(x, mask, cur_y, t, model, graph, use_v6_graph, sigma, |
| betas, alphas, alphas_bar_sqrt, one_minus_alphas_bar_sqrt): |
| t_tensor = torch.tensor([int(t)]).to(x.device) |
| eps_factor = ((1 - extract(alphas, t_tensor, cur_y)) |
| / extract(one_minus_alphas_bar_sqrt, t_tensor, cur_y)) |
| beta = extract(betas, t_tensor.repeat(x.shape[0]), cur_y) |
| eps_theta = model.generate_accelerate(cur_y, beta, x, mask) |
|
|
| if graph is not None: |
| abs_t = extract(alphas_bar_sqrt, t_tensor, cur_y) |
| am1_t = extract(one_minus_alphas_bar_sqrt, t_tensor, cur_y) |
| y0_hat = (cur_y - am1_t * eps_theta) / abs_t |
| delta = graph(y0_hat, x, int(t), sigma=sigma, A_override=x.size(0)) |
| eps_theta = eps_theta - (abs_t / am1_t) * delta |
|
|
| mean = (1 / extract(alphas, t_tensor, cur_y).sqrt()) \ |
| * (cur_y - eps_factor * eps_theta) |
| z = torch.randn_like(cur_y) |
| sigma_t = extract(betas, t_tensor, cur_y).sqrt() |
| return mean + sigma_t * z * 0.00001 |
|
|
|
|
| |
| HORIZON_FRAMES = {'1.2s': 3, '2.4s': 6, '3.6s': 9, '4.8s': 12} |
|
|
|
|
| @torch.no_grad() |
| def run(args): |
| device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') |
| cfg = Config(args.cfg, args.exp) |
|
|
| test_dset = SDDDataset(obs_len=cfg.past_frames, |
| pred_len=cfg.future_frames, split='test') |
| loader = DataLoader(test_dset, batch_size=1, shuffle=False, |
| num_workers=2, collate_fn=sdd_seq_collate) |
|
|
| ckpt_path = cfg.model_path % args.epoch |
| print(f'Loading checkpoint: {ckpt_path}') |
|
|
| model, init, graph = build_models( |
| cfg, use_graph=args.use_graph, use_v6_graph=args.use_v6_graph, |
| ckpt_path=ckpt_path, device=device) |
|
|
| betas = make_beta_schedule().to(device) |
| alphas = 1 - betas |
| alphas_prod = torch.cumprod(alphas, 0) |
| abs_sqrt = torch.sqrt(alphas_prod) |
| one_minus_abs_sqrt = torch.sqrt(1 - alphas_prod) |
|
|
| traj_mean = torch.FloatTensor(cfg.traj_mean).to(device).view(1, 1, 1, 2) |
| traj_scale = float(cfg.traj_scale) |
|
|
| np.random.seed(0); random.seed(0) |
| torch.manual_seed(0); torch.cuda.manual_seed_all(0) |
|
|
| |
| sums = {f'{k}_{h}': 0.0 |
| for h in HORIZON_FRAMES for k in ['ADE_min', 'FDE_min', 'ADE_avg', 'FDE_avg']} |
| n_agents = 0 |
| T = cfg.future_frames |
|
|
| for data in loader: |
| pre = data['pre_motion_3D'].to(device) |
| fut = data['fut_motion_3D'].to(device) |
| A = pre.size(1) |
| initial_pos = pre[:, :, -1:] |
| past_abs = ((pre - traj_mean) / traj_scale).contiguous().view(-1, cfg.past_frames, 2) |
| past_rel = ((pre - initial_pos) / traj_scale).contiguous().view(-1, cfg.past_frames, 2) |
| past_vel = torch.cat([past_rel[:, 1:] - past_rel[:, :-1], |
| torch.zeros_like(past_rel[:, -1:])], dim=1) |
| past = torch.cat([past_abs, past_rel, past_vel], dim=-1) |
| fut_rel = ((fut - initial_pos) / traj_scale).contiguous().view(-1, T, 2) |
| mask = torch.ones(A, A).to(device) |
|
|
| sp, me, ve = init(past, mask) |
| ve = ve.clamp(min=-5, max=5) |
| sp = torch.exp(ve / 2)[..., None, None] * sp \ |
| / (sp.std(dim=1).mean(dim=(1, 2))[:, None, None, None] + 1e-6) |
| loc = sp + me[:, None] |
| sigma_in = ve if args.use_v6_graph else None |
|
|
| |
| cur_y = loc[:, :10] |
| for i in reversed(range(NUM_Tau)): |
| cur_y = p_sample_accelerate( |
| past, mask, cur_y, i, model, graph, args.use_v6_graph, sigma_in, |
| betas, alphas, abs_sqrt, one_minus_abs_sqrt) |
| cur_y_ = loc[:, 10:] |
| for i in reversed(range(NUM_Tau)): |
| cur_y_ = p_sample_accelerate( |
| past, mask, cur_y_, i, model, graph, args.use_v6_graph, sigma_in, |
| betas, alphas, abs_sqrt, one_minus_abs_sqrt) |
| pred = torch.cat((cur_y_, cur_y), dim=1) |
|
|
| |
| dist = torch.norm(pred - fut_rel.unsqueeze(1), dim=-1) * traj_scale |
| for h_name, h_end in HORIZON_FRAMES.items(): |
| |
| ade_min = dist[..., :h_end].mean(dim=-1).min(dim=-1)[0] |
| fde_min = dist[..., h_end - 1].min(dim=-1)[0] |
| |
| ade_avg = dist[..., :h_end].mean(dim=-1).mean(dim=-1) |
| fde_avg = dist[..., h_end - 1].mean(dim=-1) |
| sums[f'ADE_min_{h_name}'] += ade_min.sum().item() |
| sums[f'FDE_min_{h_name}'] += fde_min.sum().item() |
| sums[f'ADE_avg_{h_name}'] += ade_avg.sum().item() |
| sums[f'FDE_avg_{h_name}'] += fde_avg.sum().item() |
| n_agents += A |
|
|
| |
| print(f'\n{args.exp} @ epoch {args.epoch} (n_agents={n_agents}, all-agents protocol)') |
| print('--- pixels ---') |
| for h in HORIZON_FRAMES: |
| am = sums[f'ADE_min_{h}'] / n_agents * 50.0 |
| fm = sums[f'FDE_min_{h}'] / n_agents * 50.0 |
| aa = sums[f'ADE_avg_{h}'] / n_agents * 50.0 |
| fa = sums[f'FDE_avg_{h}'] / n_agents * 50.0 |
| print(f' ADE_min({h})={am:8.4f} FDE_min({h})={fm:8.4f} ' |
| f'ADE_avg({h})={aa:8.4f} FDE_avg({h})={fa:8.4f}') |
|
|
|
|
| if __name__ == '__main__': |
| p = argparse.ArgumentParser() |
| p.add_argument('--cfg', default='sdd/sdd') |
| p.add_argument('--exp', required=True, help='info tag, e.g. baseline_v2') |
| p.add_argument('--epoch', type=int, required=True) |
| p.add_argument('--use_graph', action='store_true') |
| p.add_argument('--use_v6_graph', action='store_true') |
| args = p.parse_args() |
| run(args) |
|
|