sra-trajectory-code / LED /viz_denoising_process.py
po03087's picture
SRA: MID/LED/MoFlow code + RUNNING.md instructions (code only, no data/ckpts)
d4cbafd verified
Raw
History Blame Contribute Delete
13.6 kB
"""
Visualize the full LED denoising process with/without uncertainty.
Generates side-by-side plots showing:
- Left: Without uncertainty (graph only)
- Right: With uncertainty (graph + sigma coloring)
For each sample:
- Past trajectories (black lines)
- GT future (green dashed)
- Denoising steps τ=4,3,2,1,0 with progressively refined predictions
- For sigma version: trajectory color indicates uncertainty (red=high, blue=low)
"""
import os
import sys
import torch
import random
import numpy as np
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
import matplotlib.cm as cm
from matplotlib.colors import Normalize
sys.path.insert(0, os.path.dirname(__file__))
from utils.config import Config
from data.dataloader_nba import NBADataset, seq_collate
from torch.utils.data import DataLoader
from models.model_led_initializer import LEDInitializer as InitializationModel
from models.model_diffusion import TransformerDenoisingModel as CoreDenoisingModel
from models.future_interaction_graph_v6 import FutureInteractionGraphV6Wrapper
NUM_Tau = 5
TRAJ_SCALE = 94.0 / 28.0 # to convert back to feet
def load_models(ckpt_path, use_sigma=True, use_v6_graph=True, edge_mode='relpos_only',
top_n=5, device='cuda'):
"""Load LED models from checkpoint."""
cfg = Config('led_augment', 'viz')
model = CoreDenoisingModel().to(device)
model_cp = torch.load(cfg.pretrained_core_denoising_model, map_location='cpu')
model.load_state_dict(model_cp['model_dict'])
model.eval()
model_init = InitializationModel(t_h=10, d_h=6, t_f=20, d_f=2, k_pred=20).to(device)
graph = FutureInteractionGraphV6Wrapper(
num_agents=11, future_steps=20, past_steps=10,
past_channels=6, node_dim=128, top_n=top_n,
num_denoise_steps=NUM_Tau, edge_mode=edge_mode,
).to(device)
ckpt = torch.load(ckpt_path, map_location='cpu')
model_init.load_state_dict(ckpt['model_initializer_dict'])
graph.load_state_dict(ckpt['interaction_graph_dict'])
model_init.eval()
graph.eval()
return cfg, model, model_init, graph
def make_beta_schedule(n_timesteps=100, start=1e-5, end=1e-2):
return torch.linspace(start, end, n_timesteps).cuda()
def denoise_with_intermediates(model, graph, past_traj, traj_mask, loc,
betas, alphas_prod, alphas_bar_sqrt,
one_minus_alphas_bar_sqrt, alphas,
use_sigma=False, sigma=None):
"""Run denoising and return intermediate predictions at each step."""
intermediates = [] # list of (y0_hat, sigma_val) at each step
cur_y = loc[:, :10]
for i in reversed(range(NUM_Tau)):
t = torch.tensor([i]).cuda()
eps_factor = ((1 - alphas[i]) / one_minus_alphas_bar_sqrt[i])
beta = betas[i].repeat(past_traj.shape[0]).unsqueeze(-1).unsqueeze(-1).unsqueeze(-1)
eps_theta = model.generate_accelerate(cur_y, beta.squeeze(-1).squeeze(-1), past_traj, traj_mask)
alpha_bar_sqrt_t = alphas_bar_sqrt[i]
one_minus_abs_t = one_minus_alphas_bar_sqrt[i]
y0_hat = (cur_y - one_minus_abs_t * eps_theta) / alpha_bar_sqrt_t
sigma_input = sigma if use_sigma else None
delta = graph(y0_hat, past_traj, i, sigma=sigma_input)
eps_theta = eps_theta + delta
# Store y0_hat after graph correction
y0_corrected = (cur_y - one_minus_abs_t * eps_theta) / alpha_bar_sqrt_t
intermediates.append({
'step': i,
'y0_hat': y0_corrected.detach().cpu(),
'sigma': sigma.detach().cpu() if sigma is not None else None,
})
mean = (1 / alphas[i].sqrt()) * (cur_y - eps_factor * eps_theta)
z = torch.randn_like(cur_y)
sigma_t = betas[i].sqrt()
cur_y = mean + sigma_t * z * 0.00001
# Second half (modes 10-19)
cur_y_ = loc[:, 10:]
for i in reversed(range(NUM_Tau)):
t = torch.tensor([i]).cuda()
eps_factor = ((1 - alphas[i]) / one_minus_alphas_bar_sqrt[i])
beta = betas[i].repeat(past_traj.shape[0]).unsqueeze(-1).unsqueeze(-1).unsqueeze(-1)
eps_theta = model.generate_accelerate(cur_y_, beta.squeeze(-1).squeeze(-1), past_traj, traj_mask)
alpha_bar_sqrt_t = alphas_bar_sqrt[i]
one_minus_abs_t = one_minus_alphas_bar_sqrt[i]
y0_hat = (cur_y_ - one_minus_abs_t * eps_theta) / alpha_bar_sqrt_t
sigma_input = sigma if use_sigma else None
delta = graph(y0_hat, past_traj, i, sigma=sigma_input)
eps_theta = eps_theta + delta
mean = (1 / alphas[i].sqrt()) * (cur_y_ - eps_factor * eps_theta)
z = torch.randn_like(cur_y_)
sigma_t = betas[i].sqrt()
cur_y_ = mean + sigma_t * z * 0.00001
final_pred = torch.cat((cur_y_, cur_y), dim=1)
return intermediates, final_pred
def plot_single_sample(ax, past, gt_fut, intermediates, final_pred,
initial_pos, traj_mean, traj_scale,
sigma_vals=None, title='', show_uncertainty=False):
"""Plot one sample's denoising process on an axis."""
A = 11
K = 20
# Convert past to absolute court positions
past_abs = past.reshape(A, 10, 6)[:, :, :2] # abs_xy channels
past_abs = past_abs * traj_scale + traj_mean # back to court coords
# GT future
gt_abs = gt_fut.reshape(A, 20, 2) * traj_scale + initial_pos.reshape(A, 1, 2)
# Colors for agents
agent_colors = plt.cm.tab10(np.linspace(0, 1, A))
# Draw court
ax.set_xlim(-2, 30)
ax.set_ylim(-2, 17)
ax.set_aspect('equal')
ax.set_facecolor('#2d5016') # dark green court
# Draw past trajectories
for a in range(A):
ax.plot(past_abs[a, :, 0], past_abs[a, :, 1], '-',
color='white', alpha=0.5, linewidth=1)
ax.plot(past_abs[a, -1, 0], past_abs[a, -1, 1], 'o',
color='white', markersize=4)
# Draw GT future
for a in range(A):
ax.plot(gt_abs[a, :, 0], gt_abs[a, :, 1], '--',
color='lime', alpha=0.4, linewidth=1)
# Draw denoising steps (from noisy to clean)
step_alphas = [0.15, 0.25, 0.35, 0.5, 0.7]
for idx, inter in enumerate(intermediates):
y0 = inter['y0_hat'] # [B*A, K, T, 2]
step = inter['step']
alpha = step_alphas[min(idx, len(step_alphas) - 1)]
# Take best mode (mode 0 for simplicity)
y0_mode0 = y0[:A, 0, :, :] # [A, T, 2]
y0_abs = y0_mode0 * traj_scale + initial_pos.reshape(A, 1, 2)
if show_uncertainty and inter['sigma'] is not None:
sigma_a = inter['sigma'][:A, 0].numpy()
norm = Normalize(vmin=sigma_a.min(), vmax=sigma_a.max())
cmap = cm.coolwarm # blue=certain, red=uncertain
for a in range(A):
color = cmap(norm(sigma_a[a]))
ax.plot(y0_abs[a, :, 0], y0_abs[a, :, 1], '-',
color=color, alpha=alpha, linewidth=1.5)
else:
for a in range(A):
ax.plot(y0_abs[a, :, 0], y0_abs[a, :, 1], '-',
color=agent_colors[a], alpha=alpha, linewidth=1)
# Draw final prediction (best mode)
if final_pred is not None:
pred = final_pred[:A] # [A, K, T, 2]
pred_abs = pred.numpy() * traj_scale + initial_pos.reshape(A, 1, 1, 2)
gt_exp = np.expand_dims(gt_abs, 1).repeat(pred_abs.shape[1], axis=1)
ade_per_mode = np.linalg.norm(pred_abs - gt_exp, axis=-1).mean(axis=-1) # [A, K]
best_modes = ade_per_mode.argmin(axis=-1) # [A]
for a in range(A):
best = pred_abs[a, best_modes[a]]
ax.plot(best[:, 0], best[:, 1], '-',
color=agent_colors[a], alpha=0.9, linewidth=2)
ax.plot(best[-1, 0], best[-1, 1], '*',
color=agent_colors[a], markersize=6)
ax.set_title(title, fontsize=10, color='white')
ax.tick_params(colors='gray')
def main():
device = 'cuda:0' # CUDA_VISIBLE_DEVICES remaps to 0
torch.cuda.set_device(0)
# Paths
ckpt_sigma = '/mnt/jaewoo4tb/srtp/LED/results/led_augment/graph_v6_edge_relpos/models/model_0052.p'
ckpt_nosigma = '/mnt/jaewoo4tb/srtp/LED/results/led_augment/graph_v6_nosigma_n3/models/model_0084.p'
out_dir = '/mnt/jaewoo4tb/srtp/LED/visualizations'
os.makedirs(out_dir, exist_ok=True)
# Load models
cfg_s, model_s, init_s, graph_s = load_models(
ckpt_sigma, use_sigma=True, edge_mode='relpos_only', top_n=5, device=device)
cfg_n, model_n, init_n, graph_n = load_models(
ckpt_nosigma, use_sigma=False, edge_mode='full', top_n=3, device=device)
# Diffusion schedule
betas = torch.linspace(1e-5, 1e-2, 100).to(device)
alphas = 1 - betas
alphas_prod = torch.cumprod(alphas, 0)
alphas_bar_sqrt = torch.sqrt(alphas_prod)
one_minus_alphas_bar_sqrt = torch.sqrt(1 - alphas_prod)
traj_mean = torch.FloatTensor(cfg_s.traj_mean).to(device).unsqueeze(0).unsqueeze(0).unsqueeze(0)
traj_scale = cfg_s.traj_scale
# Load test data
test_dset = NBADataset(obs_len=10, pred_len=20, training=False)
test_loader = DataLoader(test_dset, batch_size=1, shuffle=False, collate_fn=seq_collate)
# Set seed for reproducibility
np.random.seed(42)
random.seed(42)
torch.manual_seed(42)
num_samples = 5
sample_indices = sorted(random.sample(range(len(test_dset)), num_samples))
with torch.no_grad():
for sample_idx, data in enumerate(test_loader):
if sample_idx not in sample_indices:
continue
if sample_idx > max(sample_indices):
break
batch_size = 1
traj_mask = torch.ones(11, 11).to(device)
initial_pos = data['pre_motion_3D'].to(device)[:, :, -1:]
past_traj_abs = ((data['pre_motion_3D'].to(device) - traj_mean) / traj_scale).view(-1, 10, 2)
past_traj_rel = ((data['pre_motion_3D'].to(device) - initial_pos) / traj_scale).view(-1, 10, 2)
past_traj_vel = torch.cat((past_traj_rel[:, 1:] - past_traj_rel[:, :-1],
torch.zeros_like(past_traj_rel[:, :1])), dim=1)
past_traj = torch.cat((past_traj_abs, past_traj_rel, past_traj_vel), dim=-1)
fut_traj = ((data['fut_motion_3D'].to(device) - initial_pos) / traj_scale).view(-1, 20, 2)
# --- With sigma ---
sample_pred_s, mean_s, var_s = init_s(past_traj, traj_mask)
sample_pred_s = (torch.exp(var_s / 2)[..., None, None]
* sample_pred_s / sample_pred_s.std(dim=1).mean(dim=(1, 2))[:, None, None, None])
loc_s = sample_pred_s + mean_s[:, None]
intermediates_s, final_s = denoise_with_intermediates(
model_s, graph_s, past_traj, traj_mask, loc_s,
betas, alphas_prod, alphas_bar_sqrt, one_minus_alphas_bar_sqrt, alphas,
use_sigma=True, sigma=var_s)
# --- Without sigma ---
sample_pred_n, mean_n, var_n = init_n(past_traj, traj_mask)
sample_pred_n = (torch.exp(var_n / 2)[..., None, None]
* sample_pred_n / sample_pred_n.std(dim=1).mean(dim=(1, 2))[:, None, None, None])
loc_n = sample_pred_n + mean_n[:, None]
intermediates_n, final_n = denoise_with_intermediates(
model_n, graph_n, past_traj, traj_mask, loc_n,
betas, alphas_prod, alphas_bar_sqrt, one_minus_alphas_bar_sqrt, alphas,
use_sigma=False, sigma=None)
# --- Plot ---
fig, axes = plt.subplots(1, 2, figsize=(20, 8))
fig.patch.set_facecolor('#1a1a1a')
init_pos_cpu = initial_pos.cpu().squeeze(0)
traj_mean_cpu = traj_mean.cpu().squeeze(0).squeeze(0)
plot_single_sample(
axes[0], past_traj.cpu().numpy(), fut_traj.cpu().numpy(),
intermediates_n, final_n.cpu(),
init_pos_cpu.numpy(), traj_mean_cpu.numpy(), traj_scale,
title=f'Without Uncertainty (sample {sample_idx})')
plot_single_sample(
axes[1], past_traj.cpu().numpy(), fut_traj.cpu().numpy(),
intermediates_s, final_s.cpu(),
init_pos_cpu.numpy(), traj_mean_cpu.numpy(), traj_scale,
sigma_vals=var_s.cpu(), show_uncertainty=True,
title=f'With Uncertainty (sample {sample_idx})')
# Add colorbar for uncertainty
sm = plt.cm.ScalarMappable(cmap=cm.coolwarm)
sm.set_array([])
cbar = fig.colorbar(sm, ax=axes[1], shrink=0.6, pad=0.02)
cbar.set_label('Uncertainty (σ)', color='white')
cbar.ax.yaxis.set_tick_params(color='white')
plt.setp(plt.getp(cbar.ax.axes, 'yticklabels'), color='white')
plt.suptitle(f'LED Denoising Process — Sample {sample_idx}\n'
f'Fading lines: τ=4→0 (noisy→clean). '
f'Green dashed: GT. Stars: final endpoints.',
color='white', fontsize=12)
plt.tight_layout()
save_path = os.path.join(out_dir, f'denoising_sample_{sample_idx:04d}.png')
plt.savefig(save_path, dpi=150, bbox_inches='tight',
facecolor=fig.get_facecolor())
plt.close()
print(f'Saved: {save_path}')
print(f'\nAll visualizations saved to {out_dir}/')
if __name__ == '__main__':
main()