sra-trajectory-code / MID /eval_sdd_graph.py
po03087's picture
SRA: MID/LED/MoFlow code + RUNNING.md instructions (code only, no data/ckpts)
d4cbafd verified
Raw
History Blame Contribute Delete
5.89 kB
"""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}')