sra-trajectory-code / LED /eval_sdd_led_mid_protocol.py
po03087's picture
SRA: MID/LED/MoFlow code + RUNNING.md instructions (code only, no data/ckpts)
d4cbafd verified
Raw
History Blame Contribute Delete
7 kB
"""
Re-evaluate an LED SDD checkpoint using the MID evaluation protocol:
per-pedestrian full-horizon (12-frame) ADE / final-frame FDE,
best_of_20 per pedestrian, averaged across all target pedestrians,
×50 to report in pixels.
Usage:
python eval_sdd_led_mid_protocol.py --exp baseline_v2 --epoch 60
python eval_sdd_led_mid_protocol.py --exp graph_sigma_v2 --epoch 60 --use_graph
"""
import argparse, os, sys, random, torch, numpy as np
from torch.utils.data import DataLoader
from data.dataloader_sdd import SDDDataset, sdd_seq_collate
from models.model_led_initializer import LEDInitializer as InitializationModel
from models.model_diffusion import TransformerDenoisingModel as CoreDenoisingModel
from trainer.train_sdd_led import NUM_Tau
from utils.config import Config
def build_models(cfg, use_graph, use_v6_graph, ckpt_path, device):
model = CoreDenoisingModel(past_len=cfg.past_frames).to(device)
core_ckpt = torch.load(cfg.pretrained_core_denoising_model, map_location='cpu')
model.load_state_dict(core_ckpt['model_dict'])
model.eval()
init = InitializationModel(
t_h=cfg.past_frames, d_h=6,
t_f=cfg.future_frames, d_f=2, k_pred=20).to(device)
ckpt = torch.load(ckpt_path, map_location='cpu')
init.load_state_dict(ckpt['model_initializer_dict'])
init.eval()
graph = None
if use_graph:
from models.future_interaction_graph_v6 import FutureInteractionGraphV6Wrapper
graph = FutureInteractionGraphV6Wrapper(
num_agents=64, future_steps=cfg.future_frames,
past_steps=cfg.past_frames, past_channels=6,
node_dim=128, top_n=5, num_denoise_steps=NUM_Tau).to(device)
sd = {k: v for k, v in ckpt['interaction_graph_dict'].items()
if '_single_edge_index' not in k}
graph.load_state_dict(sd, strict=False)
graph.eval()
return model, init, graph
def make_beta_schedule(n=100, start=1e-4, end=5e-2):
return torch.linspace(start, end, n)
def extract(a, t, x):
out = torch.gather(a, 0, t.to(a.device))
return out.reshape(t.shape[0], *([1] * (len(x.shape) - 1)))
@torch.no_grad()
def p_sample_accelerate(x, mask, cur_y, t, model, graph, use_v6_graph, sigma,
betas, alphas, alphas_bar_sqrt, one_minus_alphas_bar_sqrt):
t_tensor = torch.tensor([int(t)]).to(x.device)
eps_factor = ((1 - extract(alphas, t_tensor, cur_y))
/ extract(one_minus_alphas_bar_sqrt, t_tensor, cur_y))
beta = extract(betas, t_tensor.repeat(x.shape[0]), cur_y)
eps_theta = model.generate_accelerate(cur_y, beta, x, mask)
if graph is not None:
abs_t = extract(alphas_bar_sqrt, t_tensor, cur_y)
am1_t = extract(one_minus_alphas_bar_sqrt, t_tensor, cur_y)
y0_hat = (cur_y - am1_t * eps_theta) / abs_t
delta = graph(y0_hat, x, int(t), sigma=sigma, A_override=x.size(0))
eps_theta = eps_theta - (abs_t / am1_t) * delta
mean = (1 / extract(alphas, t_tensor, cur_y).sqrt()) \
* (cur_y - eps_factor * eps_theta)
z = torch.randn_like(cur_y)
sigma_t = extract(betas, t_tensor, cur_y).sqrt()
return mean + sigma_t * z * 0.00001
@torch.no_grad()
def run(args):
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
cfg = Config(args.cfg, args.exp)
test_dset = SDDDataset(obs_len=cfg.past_frames,
pred_len=cfg.future_frames, split='test')
loader = DataLoader(test_dset, batch_size=1, shuffle=False,
num_workers=2, collate_fn=sdd_seq_collate)
exp_dir = cfg.model_dir # results_sdd/sdd/sdd/<info>/models
ckpt_path = cfg.model_path % args.epoch
print(f'Loading checkpoint: {ckpt_path}')
model, init, graph = build_models(
cfg, use_graph=args.use_graph, use_v6_graph=args.use_v6_graph,
ckpt_path=ckpt_path, device=device)
betas = make_beta_schedule().to(device)
alphas = 1 - betas
alphas_prod = torch.cumprod(alphas, 0)
abs_sqrt = torch.sqrt(alphas_prod)
one_minus_abs_sqrt = torch.sqrt(1 - alphas_prod)
traj_mean = torch.FloatTensor(cfg.traj_mean).to(device).view(1, 1, 1, 2)
traj_scale = float(cfg.traj_scale)
np.random.seed(0); random.seed(0)
torch.manual_seed(0); torch.cuda.manual_seed_all(0)
total_ade, total_fde, n = 0.0, 0.0, 0
T = cfg.future_frames
for data in loader:
pre = data['pre_motion_3D'].to(device) # [1, A, 8, 2]
fut = data['fut_motion_3D'].to(device)
A = pre.size(1)
initial_pos = pre[:, :, -1:]
past_abs = ((pre - traj_mean) / traj_scale).contiguous().view(-1, cfg.past_frames, 2)
past_rel = ((pre - initial_pos) / traj_scale).contiguous().view(-1, cfg.past_frames, 2)
past_vel = torch.cat([past_rel[:, 1:] - past_rel[:, :-1],
torch.zeros_like(past_rel[:, -1:])], dim=1)
past = torch.cat([past_abs, past_rel, past_vel], dim=-1)
fut_rel = ((fut - initial_pos) / traj_scale).contiguous().view(-1, T, 2)
mask = torch.ones(A, A).to(device)
sp, me, ve = init(past, mask)
ve = ve.clamp(min=-5, max=5)
sp = torch.exp(ve / 2)[..., None, None] * sp \
/ (sp.std(dim=1).mean(dim=(1, 2))[:, None, None, None] + 1e-6)
loc = sp + me[:, None]
sigma_in = ve if args.use_v6_graph else None
# leapfrog: 20 modes = 10+10 two halves, each 5 reverse steps
cur_y = loc[:, :10]
for i in reversed(range(NUM_Tau)):
cur_y = p_sample_accelerate(
past, mask, cur_y, i, model, graph, args.use_v6_graph, sigma_in,
betas, alphas, abs_sqrt, one_minus_abs_sqrt)
cur_y_ = loc[:, 10:]
for i in reversed(range(NUM_Tau)):
cur_y_ = p_sample_accelerate(
past, mask, cur_y_, i, model, graph, args.use_v6_graph, sigma_in,
betas, alphas, abs_sqrt, one_minus_abs_sqrt)
pred = torch.cat((cur_y_, cur_y), dim=1) # [A, 20, T, 2]
# target only (index 0) per scene, full-horizon ADE / final FDE, best_of_20
pred_0 = pred[0:1]
fut_0 = fut_rel[0:1]
dist = torch.norm(fut_0.unsqueeze(1) - pred_0, dim=-1) * traj_scale
ade = dist.mean(dim=-1).min(dim=-1)[0]
fde = dist[:, :, -1].min(dim=-1)[0]
total_ade += ade.sum().item()
total_fde += fde.sum().item()
n += 1
ade_px = total_ade / n * 50.0
fde_px = total_fde / n * 50.0
print(f'Epoch {args.epoch} MID-protocol: ADE={ade_px:.4f} FDE={fde_px:.4f} (n={n})')
if __name__ == '__main__':
p = argparse.ArgumentParser()
p.add_argument('--cfg', default='sdd/sdd')
p.add_argument('--exp', required=True, help='info tag e.g. baseline_v2')
p.add_argument('--epoch', type=int, required=True)
p.add_argument('--use_graph', action='store_true')
p.add_argument('--use_v6_graph', action='store_true')
args = p.parse_args()
run(args)