| """ |
| 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 |
|
|
|
|
| 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 = [] |
|
|
| 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 |
|
|
| |
| 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 |
|
|
| |
| 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 |
|
|
| |
| past_abs = past.reshape(A, 10, 6)[:, :, :2] |
| past_abs = past_abs * traj_scale + traj_mean |
|
|
| |
| gt_abs = gt_fut.reshape(A, 20, 2) * traj_scale + initial_pos.reshape(A, 1, 2) |
|
|
| |
| agent_colors = plt.cm.tab10(np.linspace(0, 1, A)) |
|
|
| |
| ax.set_xlim(-2, 30) |
| ax.set_ylim(-2, 17) |
| ax.set_aspect('equal') |
| ax.set_facecolor('#2d5016') |
|
|
| |
| 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) |
|
|
| |
| for a in range(A): |
| ax.plot(gt_abs[a, :, 0], gt_abs[a, :, 1], '--', |
| color='lime', alpha=0.4, linewidth=1) |
|
|
| |
| step_alphas = [0.15, 0.25, 0.35, 0.5, 0.7] |
| for idx, inter in enumerate(intermediates): |
| y0 = inter['y0_hat'] |
| step = inter['step'] |
| alpha = step_alphas[min(idx, len(step_alphas) - 1)] |
|
|
| |
| y0_mode0 = y0[:A, 0, :, :] |
| 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 |
|
|
| 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) |
|
|
| |
| if final_pred is not None: |
| pred = final_pred[:A] |
| 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) |
| best_modes = ade_per_mode.argmin(axis=-1) |
|
|
| 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' |
| torch.cuda.set_device(0) |
|
|
| |
| 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) |
|
|
| |
| 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) |
|
|
| |
| 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 |
|
|
| |
| 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) |
|
|
| |
| 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) |
|
|
| |
| 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) |
|
|
| |
| 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) |
|
|
| |
| 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})') |
|
|
| |
| 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() |
|
|