| """ |
| Stage-2 LED trainer for SDD. Variable-A scenes, batch_size=1, grad_accum. |
| Optional graph module (--use_graph --use_v6_graph). |
| Eval: only target agent (index 0) counts toward ADE/FDE. |
| """ |
|
|
| import os, time, torch, random, numpy as np |
| import torch.nn as nn |
| from utils.config import Config |
| from utils.utils import print_log |
| from torch.utils.data import DataLoader |
| from torch.utils.tensorboard import SummaryWriter |
| 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 |
|
|
| NUM_Tau = 5 |
|
|
|
|
| class Trainer: |
| def __init__(self, config): |
| if torch.cuda.is_available(): |
| torch.cuda.set_device(config.gpu) |
| self.device = torch.device('cuda') if config.cuda else torch.device('cpu') |
| self.cfg = Config(config.cfg, config.info) |
| self.use_graph = bool(getattr(config, 'use_graph', False)) |
| self.use_v6_graph = bool(getattr(config, 'use_v6_graph', False)) |
| self.residual_on = getattr(config, 'residual_on', 'y0') |
| self.grad_accum = getattr(config, 'grad_accum', 16) |
|
|
| train_dset = SDDDataset(obs_len=self.cfg.past_frames, |
| pred_len=self.cfg.future_frames, split='train') |
| test_dset = SDDDataset(obs_len=self.cfg.past_frames, |
| pred_len=self.cfg.future_frames, split='test') |
| self.train_loader = DataLoader(train_dset, batch_size=1, shuffle=True, |
| num_workers=2, collate_fn=sdd_seq_collate) |
| self.test_loader = DataLoader(test_dset, batch_size=1, shuffle=False, |
| num_workers=2, collate_fn=sdd_seq_collate) |
|
|
| self.traj_mean = torch.FloatTensor(self.cfg.traj_mean).cuda().unsqueeze(0).unsqueeze(0).unsqueeze(0) |
| self.traj_scale = float(self.cfg.traj_scale) |
|
|
| self.n_steps = self.cfg.diffusion.steps |
| self.betas = self._make_beta_schedule( |
| self.cfg.diffusion.beta_schedule, self.n_steps, |
| self.cfg.diffusion.beta_start, self.cfg.diffusion.beta_end).cuda() |
| self.alphas = 1 - self.betas |
| self.alphas_prod = torch.cumprod(self.alphas, 0) |
| self.alphas_bar_sqrt = torch.sqrt(self.alphas_prod) |
| self.one_minus_alphas_bar_sqrt = torch.sqrt(1 - self.alphas_prod) |
|
|
| self.model = CoreDenoisingModel(past_len=self.cfg.past_frames).cuda() |
| ckpt_path = self.cfg.pretrained_core_denoising_model |
| if not os.path.isfile(ckpt_path): |
| raise FileNotFoundError(f'Missing pretrained denoiser: {ckpt_path}') |
| self.model.load_state_dict(torch.load(ckpt_path, map_location='cpu')['model_dict']) |
|
|
| self.model_initializer = InitializationModel( |
| t_h=self.cfg.past_frames, d_h=6, |
| t_f=self.cfg.future_frames, d_f=2, k_pred=20).cuda() |
|
|
| params = list(self.model_initializer.parameters()) |
| self.interaction_graph = None |
| if self.use_graph: |
| from models.future_interaction_graph_v6 import FutureInteractionGraphV6Wrapper |
| self.interaction_graph = FutureInteractionGraphV6Wrapper( |
| num_agents=64, future_steps=self.cfg.future_frames, |
| past_steps=self.cfg.past_frames, past_channels=6, |
| node_dim=128, top_n=5, num_denoise_steps=NUM_Tau).cuda() |
| params += list(self.interaction_graph.parameters()) |
|
|
| self.opt = torch.optim.AdamW(params, lr=config.learning_rate) |
| self.scheduler = torch.optim.lr_scheduler.StepLR( |
| self.opt, step_size=self.cfg.decay_step, gamma=self.cfg.decay_gamma) |
|
|
| self.log = open(os.path.join(self.cfg.log_dir, 'log.txt'), 'a+') |
| self.tb = SummaryWriter(log_dir=os.path.join(self.cfg.log_dir, 'tb')) |
| self.global_step = 0 |
| self._print_param(self.model, 'Core Denoiser') |
| self._print_param(self.model_initializer, 'Initializer') |
| if self.interaction_graph: |
| self._print_param(self.interaction_graph, 'Graph') |
|
|
| T = self.cfg.future_frames |
| self.temporal_reweight = torch.FloatTensor( |
| [(T + 1) - i for i in range(1, T + 1)]).cuda().unsqueeze(0).unsqueeze(0) / (T / 2) |
|
|
| def _print_param(self, m, name): |
| t = sum(p.numel() for p in m.parameters()) |
| tr = sum(p.numel() for p in m.parameters() if p.requires_grad) |
| print_log(f'[{name}] {tr}/{t}', self.log) |
|
|
| def _make_beta_schedule(self, schedule, n, start, end): |
| if schedule == 'linear': return torch.linspace(start, end, n) |
| return torch.linspace(start, end, n) |
|
|
| def _extract(self, a, t, x): |
| out = torch.gather(a, 0, t.to(a.device)) |
| return out.reshape(t.shape[0], *([1] * (len(x.shape) - 1))) |
|
|
| def p_sample_accelerate(self, x, mask, cur_y, t, sigma=None): |
| t_tensor = torch.tensor([int(t)]).cuda() |
| eps_factor = ((1 - self._extract(self.alphas, t_tensor, cur_y)) |
| / self._extract(self.one_minus_alphas_bar_sqrt, t_tensor, cur_y)) |
| beta = self._extract(self.betas, t_tensor.repeat(x.shape[0]), cur_y) |
| eps_theta = self.model.generate_accelerate(cur_y, beta, x, mask) |
|
|
| if self.interaction_graph is not None: |
| abs_t = self._extract(self.alphas_bar_sqrt, t_tensor, cur_y) |
| am1_t = self._extract(self.one_minus_alphas_bar_sqrt, t_tensor, cur_y) |
| y0_hat = (cur_y - am1_t * eps_theta) / abs_t |
| delta = self.interaction_graph( |
| y0_hat, x, int(t), sigma=sigma, A_override=x.size(0)) |
| eps_theta = eps_theta - (abs_t / am1_t) * delta |
|
|
| mean = (1 / self._extract(self.alphas, t_tensor, cur_y).sqrt()) \ |
| * (cur_y - eps_factor * eps_theta) |
| z = torch.randn_like(cur_y) |
| sigma_t = self._extract(self.betas, t_tensor, cur_y).sqrt() |
| return mean + sigma_t * z * 0.00001 |
|
|
| def p_sample_loop_accelerate(self, x, mask, loc, sigma=None): |
| cur_y = loc[:, :10] |
| for i in reversed(range(NUM_Tau)): |
| cur_y = self.p_sample_accelerate(x, mask, cur_y, i, sigma=sigma) |
| cur_y_ = loc[:, 10:] |
| for i in reversed(range(NUM_Tau)): |
| cur_y_ = self.p_sample_accelerate(x, mask, cur_y_, i, sigma=sigma) |
| return torch.cat((cur_y_, cur_y), dim=1) |
|
|
| def data_preprocess(self, data): |
| pre = data['pre_motion_3D'].cuda() |
| fut = data['fut_motion_3D'].cuda() |
| A = pre.size(1) |
| initial_pos = pre[:, :, -1:] |
| past_abs = ((pre - self.traj_mean) / self.traj_scale).contiguous().view(-1, self.cfg.past_frames, 2) |
| past_rel = ((pre - initial_pos) / self.traj_scale).contiguous().view(-1, self.cfg.past_frames, 2) |
| past_vel = torch.cat([past_rel[:, 1:] - past_rel[:, :-1], |
| torch.zeros_like(past_rel[:, -1:])], dim=1) |
| past_traj = torch.cat([past_abs, past_rel, past_vel], dim=-1) |
| fut_traj = ((fut - initial_pos) / self.traj_scale).contiguous().view(-1, self.cfg.future_frames, 2) |
| mask = torch.ones(A, A).cuda() |
| return A, mask, past_traj, fut_traj |
|
|
| def fit(self): |
| for epoch in range(self.cfg.num_epochs): |
| lt, ld, lu = self._train_epoch(epoch) |
| print_log(f'[{time.strftime("%Y-%m-%d %H:%M:%S")}] Epoch: {epoch}\t' |
| f'Loss: {lt:.6f}\tDist: {ld:.6f}\tUnc: {lu:.6f}', self.log) |
| self.tb.add_scalar('train/loss', lt, epoch) |
| self.tb.add_scalar('train/loss_dist', ld, epoch) |
|
|
| if (epoch + 1) % self.cfg.test_interval == 0: |
| perf, n = self._test_epoch() |
| |
| ade_px = perf['ADE'] / n * 50.0 |
| fde_px = perf['FDE'] / n * 50.0 |
| print_log(f'Epoch {epoch} Best Of 20: ADE: {ade_px:.4f} FDE: {fde_px:.4f}', self.log) |
| self.tb.add_scalar('val/ADE_px', ade_px, epoch) |
| self.tb.add_scalar('val/FDE_px', fde_px, epoch) |
|
|
| cp = {'model_initializer_dict': self.model_initializer.state_dict()} |
| if self.interaction_graph: |
| cp['interaction_graph_dict'] = self.interaction_graph.state_dict() |
| torch.save(cp, self.cfg.model_path % (epoch + 1)) |
| self.scheduler.step() |
| self.tb.flush(); self.tb.close() |
|
|
| def _train_epoch(self, epoch): |
| self.model.train(); self.model_initializer.train() |
| if self.interaction_graph: self.interaction_graph.train() |
| lt, ld, lu, cnt = 0, 0, 0, 0 |
| self.opt.zero_grad() |
| for i, data in enumerate(self.train_loader): |
| A, mask, past, fut = self.data_preprocess(data) |
| sp, me, ve = self.model_initializer(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 self.use_v6_graph else None |
| gen = self.p_sample_loop_accelerate(past, mask, loc, sigma=sigma_in) |
| loss_d = ((gen - fut.unsqueeze(1)).norm(p=2, dim=-1) |
| * self.temporal_reweight).mean(dim=-1).min(dim=1)[0].mean() |
| loss_u = (torch.exp(-ve) |
| * (gen - fut.unsqueeze(1)).norm(p=2, dim=-1).mean(dim=(1, 2)) |
| + ve).mean() |
| loss = loss_d * 50 + loss_u |
| (loss / self.grad_accum).backward() |
| if (i + 1) % self.grad_accum == 0: |
| params = list(self.model_initializer.parameters()) |
| if self.interaction_graph: params += list(self.interaction_graph.parameters()) |
| nn.utils.clip_grad_norm_(params, 1.0) |
| self.opt.step(); self.opt.zero_grad() |
| lt += loss.item(); ld += loss_d.item() * 50; lu += loss_u.item(); cnt += 1 |
| self.global_step += 1 |
| self.opt.step(); self.opt.zero_grad() |
| return lt / cnt, ld / cnt, lu / cnt |
|
|
| def _test_epoch(self): |
| """MID-style SDD protocol: |
| per-pedestrian full-horizon ADE (mean L2 over 12 future frames) and |
| FDE (L2 at final frame), best_of_20 per pedestrian, then average |
| across all evaluated pedestrians. Coordinates are already in |
| MID's ÷50 mean-centered space, so the final ADE/FDE is multiplied |
| by 50 to report in pixels. |
| Each scene in the SDD dataloader corresponds to one target |
| pedestrian (index 0) + its neighbors; we evaluate only the target |
| per scene so each pedestrian is counted exactly once (matches |
| MID's get_timesteps_data qualification intent). |
| """ |
| T = self.cfg.future_frames |
| perf = {'ADE': 0.0, 'FDE': 0.0} |
| n = 0 |
| np.random.seed(0); random.seed(0) |
| torch.manual_seed(0); torch.cuda.manual_seed_all(0) |
| self.model_initializer.eval() |
| if self.interaction_graph: self.interaction_graph.eval() |
| with torch.no_grad(): |
| for data in self.test_loader: |
| A, mask, past, fut = self.data_preprocess(data) |
| sp, me, ve = self.model_initializer(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 self.use_v6_graph else None |
| pred = self.p_sample_loop_accelerate(past, mask, loc, sigma=sigma_in) |
| |
| pred_0 = pred[0:1] |
| fut_0 = fut[0:1] |
| dist = torch.norm(fut_0.unsqueeze(1) - pred_0, dim=-1) * self.traj_scale |
| |
| ade_per_ped = dist.mean(dim=-1).min(dim=-1)[0] |
| fde_per_ped = dist[:, :, -1].min(dim=-1)[0] |
| perf['ADE'] += ade_per_ped.sum().item() |
| perf['FDE'] += fde_per_ped.sum().item() |
| n += 1 |
| return perf, n |
|
|