"""Eval MID SDD graph checkpoint at given stride.""" import argparse, os, sys, time, numpy as np, torch, logging, dill import torch.nn as nn from tqdm.auto import tqdm import evaluation from dataset import get_timesteps_data from models.trajectron import Trajectron from models.diffusion import TransformerConcatLinear, VarianceSchedule from utils.model_registrar import ModelRegistrar from utils.trajectron_hypers import get_traj_hypers # Re-use the graph wrapper from training script from mid_sdd_graph import GraphDenoiserWrapper if __name__ == '__main__': p = argparse.ArgumentParser() p.add_argument('--exp_name', default='mid_sdd_graph_sigma_v2') p.add_argument('--data_dir', default='processed_data') p.add_argument('--epoch', type=int, required=True) p.add_argument('--stride', type=int, default=10) 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('--gpu', type=int, default=0) args = p.parse_args() torch.cuda.set_device(args.gpu) model_dir = os.path.join("./experiments", args.exp_name) ckpt_path = os.path.join(model_dir, f"sdd_epoch{args.epoch}.pt") print(f'Loading {ckpt_path}') ckpt = torch.load(ckpt_path, map_location='cpu') # Build eval env with open(os.path.join(args.data_dir, 'sdd_test.pkl'), 'rb') as f: eval_env = dill.load(f, encoding='latin1') with open(os.path.join(args.data_dir, 'sdd_train.pkl'), 'rb') as f: train_env = dill.load(f, encoding='latin1') hyperparams = get_traj_hypers() for k in ['enc_rnn_dim_edge','enc_rnn_dim_edge_influence','enc_rnn_dim_history','enc_rnn_dim_future']: hyperparams[k] = args.encoder_dim // 2 registrar = ModelRegistrar(model_dir, "cuda") encoder = Trajectron(registrar, hyperparams, "cuda") encoder.set_environment(train_env) encoder.set_annealing_params() registrar.load_models(ckpt['encoder']) base_net = TransformerConcatLinear( point_dim=2, context_dim=args.encoder_dim, tf_layer=args.tf_layer, residual=False) graph_net = GraphDenoiserWrapper( base_net, encoder_dim=args.encoder_dim, pred_len=12, graph_hidden=128, top_n=args.top_n_neighbors, num_gnn_layers=args.graph_gnn_layers, graph_dropout=args.graph_dropout).cuda() # _single_edge_index is a dynamic buffer rebuilt per scene — skip when loading sd = {k: v for k, v in ckpt['graph_net'].items() if '_single_edge_index' not in k} graph_net.load_state_dict(sd, strict=False) graph_net.eval() var_sched = VarianceSchedule(num_steps=100, beta_T=5e-2, mode='linear') def sample_with_graph(context, N, num_points=12, K=20, stride=10): traj_list = [] for _ in range(K): x_t = torch.randn(N, num_points, 2, device=context.device) y_0_prev = None for t in range(var_sched.num_steps, 0, -stride): alpha_bar = var_sched.alpha_bars[t] alpha_bar_next = var_sched.alpha_bars[t - stride] beta = var_sched.betas[[t] * N].cuda() if y_0_prev is not None and N >= 2: eps = graph_net(x_t, beta, context, y_0_for_graph=y_0_prev) else: eps = graph_net(x_t, beta, context, skip_graph=True) x0 = (x_t - (1 - alpha_bar).sqrt() * eps) / alpha_bar.sqrt() y_0_prev = x0 x_t = alpha_bar_next.sqrt() * x0 + (1 - alpha_bar_next).sqrt() * eps traj_list.append(x_t) return torch.stack(traj_list) ph = hyperparams['prediction_horizon'] max_hl = hyperparams['maximum_history_length'] node_type = "PEDESTRIAN" ade_errors, fde_errors = [], [] with torch.no_grad(): for i, scene in enumerate(eval_env.scenes): print(f'Scene {i+1}/{len(eval_env.scenes)}') for t in tqdm(range(0, scene.timesteps, 10)): timesteps = np.arange(t, t + 10) batch = get_timesteps_data( env=eval_env, scene=scene, t=timesteps, node_type=node_type, state=hyperparams['state'], pred_state=hyperparams['pred_state'], edge_types=eval_env.get_edge_types(), min_ht=7, max_ht=max_hl, min_ft=12, max_ft=12, hyperparams=hyperparams) if batch is None: continue test_batch, nodes, timesteps_o = batch context = encoder.get_latent(test_batch, node_type) dynamics = encoder.node_models_dict[node_type].dynamic N = context.size(0) preds_vel = sample_with_graph(context, N, 12, 20, args.stride) preds_pos = dynamics.integrate_samples(preds_vel) predictions = preds_pos.cpu().numpy() predictions_dict = {} for i_n, ts in enumerate(timesteps_o): if ts not in predictions_dict: predictions_dict[ts] = {} predictions_dict[ts][nodes[i_n]] = np.transpose(predictions[:, [i_n]], (1, 0, 2, 3)) be = evaluation.compute_batch_statistics( predictions_dict, scene.dt, max_hl=max_hl, ph=ph, node_type_enum=eval_env.NodeType, kde=False, map=None, best_of=True, prune_ph_to_future=True) ade_errors = np.hstack((ade_errors, be[node_type]['ade'])) fde_errors = np.hstack((fde_errors, be[node_type]['fde'])) ade = np.mean(ade_errors) * 50 fde = np.mean(fde_errors) * 50 print(f'Epoch {args.epoch} stride={args.stride}: ADE={ade:.4f} FDE={fde:.4f}')