File size: 5,891 Bytes
d4cbafd | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 | """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}')
|