| """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 |
|
|
| |
| 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') |
|
|
| |
| 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() |
| |
| 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}') |
|
|