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