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