| """ |
| GRPO fine-tuning of LED + SRA graph on NBA. |
| |
| Formulation (single-step bandit on the initializer): |
| LED's 20-mode diversity comes entirely from the initializer; the 5-step |
| leapfrog denoising is near-deterministic (its DDPM noise is x1e-5). So we |
| treat the initializer's mode set `loc` [B*A, K, T, 2] as the ACTION: |
| |
| loc = initializer(past) # deterministic mean |
| loc_s = loc + init_noise * z # sampled action, z ~ N(0,I) |
| pred = leapfrog_decode(loc_s) # fixed decoder (graph+denoiser frozen) |
| reward = per-agent accuracy (ADE/FDE) + joint (JADE/JFDE) |
| |
| This sidesteps the long-chain credit-assignment / high-variance problem that |
| limited the MoFlow (10-step ODE) experiment: the whole trajectory-set is one |
| action, log-prob factorizes over (agent, mode), and GRPO's group-relative |
| advantage is taken over the K modes per agent. |
| |
| Only the initializer is trained (graph + core denoiser frozen); KL anchor to a |
| frozen copy of the warm-start initializer. Eval uses the near-deterministic |
| decoder (matches the LED baseline) and reports ADE/FDE/JADE/JFDE. |
| """ |
|
|
| import os |
| import sys |
| import argparse |
| import math |
| import copy |
|
|
| import numpy as np |
| import torch |
|
|
| from trainer.train_led_graph import Trainer as LEDTrainer, NUM_Tau |
|
|
|
|
| |
| def compute_reward_agentwise(pred, gt, init_pos, *, w_ade=1.0, w_fde=1.0, |
| w_jade=1.0, w_jfde=1.0, w_col=0.0, w_kin=0.0, |
| d_min=0.4, a_max=1.0, ball_idx=None): |
| B, K, A, T, _ = pred.shape |
| err = (pred - gt.unsqueeze(1)).norm(dim=-1) |
| ade = err.mean(dim=-1) |
| fde = err[..., -1] |
| r_marg = -(w_ade * ade + w_fde * fde) |
| jade = ade.mean(dim=2, keepdim=True) |
| jfde = fde.mean(dim=2, keepdim=True) |
| r_joint = -(w_jade * jade + w_jfde * jfde) |
| reward = r_marg + r_joint |
| info = {'ade': ade.detach(), 'jade': jade.squeeze(-1).detach(), |
| 'ade_bestk': ade.min(dim=1).values.mean().detach(), |
| 'jade_bestk': jade.squeeze(-1).min(dim=1).values.mean().detach()} |
| return reward, info |
|
|
|
|
| def group_advantage(reward, eps=1e-4): |
| mean = reward.mean(dim=1, keepdim=True) |
| std = reward.std(dim=1, keepdim=True) |
| return (reward - mean) / (std + eps) |
|
|
|
|
| def _player_mask(A, ball_idx, device): |
| pm = ~torch.eye(A, dtype=torch.bool, device=device) |
| if ball_idx is not None: |
| pm[ball_idx, :] = False |
| pm[:, ball_idx] = False |
| return pm |
|
|
|
|
| def compute_reward_collision(pred, gt, init_pos, *, w_ade=0.3, w_col=1.0, |
| d_min=0.4, ball_idx=10): |
| """Non-differentiable HARD collision-count reward (the objective the |
| supervised min-of-K loss cannot optimize) + a soft ADE term to hold accuracy. |
| pred/gt in RELATIVE metric (court) units; init_pos absolute [B,A,2]. |
| Returns reward [B,K,A], info.""" |
| B, K, A, T, _ = pred.shape |
| err = (pred - gt.unsqueeze(1)).norm(dim=-1) |
| ade = err.mean(dim=-1) |
| abs_p = pred + init_pos[:, None, :, None, :] |
| mind = (abs_p.unsqueeze(3) - abs_p.unsqueeze(2)).norm(dim=-1).min(dim=-1).values |
| pm = _player_mask(A, ball_idx, pred.device) |
| hard = ((mind < d_min) & pm).float() |
| coll_count = hard.sum(dim=-1) |
| reward = -(w_ade * ade) - (w_col * coll_count) |
| info = {'ade_bestk': ade.min(dim=1).values.mean().detach(), |
| 'jade_bestk': ade.mean(dim=2).min(dim=1).values.mean().detach(), |
| 'coll_count': coll_count.mean().detach(), |
| 'coll_rate': (coll_count > 0).float().mean().detach()} |
| return reward, info |
|
|
|
|
| class LEDGRPOTrainer(LEDTrainer): |
| def __init__(self, config): |
| |
| config.use_v6_graph = True |
| config.edge_mode = getattr(config, 'edge_mode', 'relpos_only') |
| config.neighbor_mode = getattr(config, 'neighbor_mode', 'rag') |
| config.top_n = getattr(config, 'top_n', 5) |
| config.use_sigma = False |
| config.residual_on = getattr(config, 'residual_on', 'eps') |
| super().__init__(config) |
|
|
| |
| ck = torch.load(config.warm_ckpt, map_location='cpu') |
| self.model_initializer.load_state_dict(ck['model_initializer_dict']) |
| self.interaction_graph.load_state_dict(ck['interaction_graph_dict']) |
| print(f'[LED-GRPO] warm-started from {config.warm_ckpt}') |
|
|
| |
| for p in self.interaction_graph.parameters(): |
| p.requires_grad_(False) |
| for p in self.model.parameters(): |
| p.requires_grad_(False) |
| self.interaction_graph.eval() |
| self.model.eval() |
|
|
| |
| if getattr(config, 'batch', 0): |
| from data.dataloader_nba import NBADataset, seq_collate |
| from torch.utils.data import DataLoader |
| tr = NBADataset(obs_len=self.cfg.past_frames, pred_len=self.cfg.future_frames, training=True) |
| self.train_loader = DataLoader(tr, batch_size=config.batch, shuffle=True, |
| num_workers=4, collate_fn=seq_collate, pin_memory=True, drop_last=True) |
|
|
| |
| self.ref_initializer = copy.deepcopy(self.model_initializer).cuda().eval() |
| for p in self.ref_initializer.parameters(): |
| p.requires_grad_(False) |
|
|
| |
| self.opt = torch.optim.AdamW(self.model_initializer.parameters(), lr=config.grpo_lr) |
|
|
| |
| self.G = 20 |
| self.init_noise = float(config.init_noise) |
| self.kl_beta = float(config.kl_beta) |
| self.clip_eps = float(config.clip_eps) |
| self.inner_epochs = int(config.inner_epochs) |
| self.grpo_iters = int(config.grpo_iters) |
| self.eval_every = int(config.eval_every) |
| self.logratio_clip = 10.0 |
| self.max_eval_batches = int(getattr(config, 'max_eval_batches', 0)) |
| self.rw = dict(w_ade=config.w_ade, w_fde=config.w_fde, |
| w_jade=config.w_jade, w_jfde=config.w_jfde, |
| w_col=0.0, w_kin=0.0, ball_idx=None) |
| self.reward_mode = getattr(config, 'reward_mode', 'accuracy') |
| self.rw_coll = dict(w_ade=getattr(config, 'w_ade_soft', 0.3), |
| w_col=getattr(config, 'w_col', 1.0), |
| d_min=getattr(config, 'd_min', 0.4), ball_idx=10) |
| self.d_min_eval = getattr(config, 'd_min', 0.4) |
| self.ade_tol = getattr(config, 'ade_tol', 0.80) |
| self.best_sum = float('inf') |
| self.best_coll = float('inf') |
|
|
| |
| def get_loc(self, past_traj, traj_mask): |
| """Initializer -> deterministic mode set loc [B*A, K, T, 2].""" |
| guess_var, guess_mean, guess_scale = self.model_initializer(past_traj, traj_mask) |
| sp = (torch.exp(guess_scale / 2)[..., None, None] * guess_var |
| / guess_var.std(dim=1).mean(dim=(1, 2))[:, None, None, None]) |
| return sp + guess_mean[:, None] |
|
|
| def get_loc_from(self, initializer, past_traj, traj_mask): |
| guess_var, guess_mean, guess_scale = initializer(past_traj, traj_mask) |
| sp = (torch.exp(guess_scale / 2)[..., None, None] * guess_var |
| / guess_var.std(dim=1).mean(dim=(1, 2))[:, None, None, None]) |
| return sp + guess_mean[:, None] |
|
|
| @staticmethod |
| def _logp(action, mean, std): |
| var = std * std |
| lp = -0.5 * (((action - mean) ** 2) / var + math.log(2 * math.pi * var)) |
| return lp.sum(dim=(2, 3)) |
|
|
| def _to_bkat(self, x_ba_k, B, A): |
| """[B*A, K, T, 2] -> [B, K, A, T, 2]""" |
| K, T = x_ba_k.shape[1], x_ba_k.shape[2] |
| return x_ba_k.view(B, A, K, T, 2).permute(0, 2, 1, 3, 4) |
|
|
| |
| def train(self): |
| A = 11 |
| self.eval_grpo(-1) |
| self.model_initializer.train() |
| dl = self._cycle(self.train_loader) |
| for it in range(self.grpo_iters): |
| data = next(dl) |
| B, traj_mask, past, fut = self.data_preprocess(data) |
|
|
| |
| with torch.no_grad(): |
| loc = self.get_loc(past, traj_mask) |
| z = torch.randn_like(loc) |
| loc_s = loc + self.init_noise * z |
| logp_old = self._logp(loc_s, loc, self.init_noise) |
| pred = self.p_sample_loop_accelerate(past, traj_mask, loc_s) |
| loc_ref = self.get_loc_from(self.ref_initializer, past, traj_mask) |
|
|
| |
| pred_m = self._to_bkat(pred, B, A) * self.traj_scale |
| gt_m = fut.view(B, A, fut.shape[1], 2) * self.traj_scale |
| if self.reward_mode == 'collision': |
| init_pos = data['pre_motion_3D'].cuda()[:, :, -1, :] |
| reward, info = compute_reward_collision(pred_m, gt_m, init_pos, **self.rw_coll) |
| else: |
| init_pos = torch.zeros(B, A, 2, device=pred.device) |
| reward, info = compute_reward_agentwise(pred_m, gt_m, init_pos, **self.rw) |
| |
| adv = group_advantage(reward) |
| |
| adv_bak = adv.permute(0, 2, 1).reshape(B * A, self.G) |
|
|
| logp_old_flat = logp_old |
| loc_s_c = loc_s |
|
|
| |
| stats = {} |
| for _ in range(self.inner_epochs): |
| self.opt.zero_grad() |
| loc_new = self.get_loc(past, traj_mask) |
| logp_new = self._logp(loc_s_c, loc_new, self.init_noise) |
| logratio = (logp_new - logp_old_flat).clamp(-self.logratio_clip, self.logratio_clip) |
| ratio = logratio.exp() |
| unclipped = ratio * adv_bak |
| clipped = ratio.clamp(1 - self.clip_eps, 1 + self.clip_eps) * adv_bak |
| pg = -torch.min(unclipped, clipped).mean() |
| kl = (0.5 * ((loc_new - loc_ref) ** 2) / (self.init_noise ** 2)).sum(dim=(2, 3)).mean() |
| loss = pg + self.kl_beta * kl |
| loss.backward() |
| torch.nn.utils.clip_grad_norm_(self.model_initializer.parameters(), 1.0) |
| self.opt.step() |
| stats = dict(pg=pg.item(), kl=kl.item(), ratio=ratio.mean().item(), |
| clipfrac=((ratio - 1).abs() > self.clip_eps).float().mean().item()) |
|
|
| if it % 10 == 0: |
| extra = (f'collrate={info["coll_rate"].item():.3f} cnt={info["coll_count"].item():.3f} ' |
| if 'coll_rate' in info else f'JADE*={info["jade_bestk"].item():.4f} ') |
| print(f'[LED-GRPO {it}/{self.grpo_iters}] R={reward.mean().item():.4f} ' |
| f'ADE*={info["ade_bestk"].item():.4f} {extra}' |
| f'| pg={stats["pg"]:.4f} kl={stats["kl"]:.5f} ' |
| f'ratio={stats["ratio"]:.3f} clipfrac={stats["clipfrac"]:.3f}', flush=True) |
|
|
| if (it + 1) % self.eval_every == 0: |
| self.eval_grpo(it) |
| self.model_initializer.train() |
|
|
| |
| @torch.no_grad() |
| def eval_grpo(self, it): |
| self.model_initializer.eval() |
| A = 11 |
| perf = {'ADE': [0.]*4, 'FDE': [0.]*4, 'JADE': [0.]*4, 'JFDE': [0.]*4} |
| coll_thr = (0.2, 0.3, 0.4) |
| collP = {th: 0. for th in coll_thr}; collG = {th: 0. for th in coll_thr} |
| nP, nG = 0, 0 |
| n_ag, n_sc = 0, 0 |
| for bi, data in enumerate(self.test_loader): |
| if self.max_eval_batches and bi >= self.max_eval_batches: |
| break |
| B, traj_mask, past, fut = self.data_preprocess(data) |
| loc = self.get_loc(past, traj_mask) |
| pred = self.p_sample_loop_accelerate(past, traj_mask, loc) |
| |
| ipos = data['pre_motion_3D'].cuda()[:, :, -1, :] |
| Tf = fut.shape[1] |
| absP = self._to_bkat(pred, B, A) * self.traj_scale + ipos[:, None, :, None, :] |
| absG = (fut.view(B, A, Tf, 2) * self.traj_scale + ipos[:, :, None, :]).unsqueeze(1) |
| pm = _player_mask(A, 10, absP.device) |
| cpP = ((absP.unsqueeze(3) - absP.unsqueeze(2)).norm(dim=-1).min(dim=-1).values |
| .masked_fill(~pm, 1e9).reshape(B, self.G, -1).min(-1).values) |
| cpG = ((absG.unsqueeze(3) - absG.unsqueeze(2)).norm(dim=-1).min(dim=-1).values |
| .masked_fill(~pm, 1e9).reshape(B, 1, -1).min(-1).values) |
| for th in coll_thr: |
| collP[th] += (cpP < th).float().sum().item() |
| collG[th] += (cpG < th).float().sum().item() |
| nP += B * self.G; nG += B |
| fut_r = fut.unsqueeze(1).repeat(1, self.G, 1, 1) |
| d = (fut_r - pred).norm(dim=-1) * self.traj_scale |
| dB = d.view(B, A, self.G, d.shape[-1]) |
| for ti in range(1, 5): |
| e = 5 * ti |
| |
| ade = d[..., :e].mean(-1).min(dim=1)[0].sum() |
| fde = d[..., e-1].min(dim=1)[0].sum() |
| |
| jade = dB[..., :e].mean(-1).mean(dim=1).min(dim=1)[0].sum() |
| jfde = dB[..., e-1].mean(dim=1).min(dim=1)[0].sum() |
| perf['ADE'][ti-1] += ade.item(); perf['FDE'][ti-1] += fde.item() |
| perf['JADE'][ti-1] += jade.item(); perf['JFDE'][ti-1] += jfde.item() |
| n_ag += B * A; n_sc += B |
| ade4 = perf['ADE'][3]/n_ag; fde4 = perf['FDE'][3]/n_ag |
| jade4 = perf['JADE'][3]/n_sc; jfde4 = perf['JFDE'][3]/n_sc |
| s = ade4 + fde4 + jade4 + jfde4 |
| cstr = ' '.join(f'@{th}:{collP[th]/nP*100:.1f}%(GT{collG[th]/nG*100:.1f})' for th in coll_thr) |
| print(f'[LED-GRPO eval @ {it}] ADE={ade4:.4f} FDE={fde4:.4f} ' |
| f'JADE={jade4:.4f} JFDE={jfde4:.4f} | coll[pred(GT)]: {cstr}', flush=True) |
| |
| if self.reward_mode == 'collision': |
| c = collP[self.d_min_eval] / nP if self.d_min_eval in collP else collP[0.4] / nP |
| if ade4 <= getattr(self, 'ade_tol', 0.80) and c < self.best_coll: |
| self.best_coll = c |
| torch.save({'model_initializer_dict': self.model_initializer.state_dict(), |
| 'interaction_graph_dict': self.interaction_graph.state_dict()}, |
| os.path.join(self.cfg.log_dir, 'grpo_best.p')) |
| print(f' new best coll@{self.d_min_eval}={c*100:.2f}% at ADE={ade4:.4f} -> grpo_best.p', flush=True) |
| elif s < self.best_sum: |
| self.best_sum = s |
| torch.save({'model_initializer_dict': self.model_initializer.state_dict(), |
| 'interaction_graph_dict': self.interaction_graph.state_dict()}, |
| os.path.join(self.cfg.log_dir, 'grpo_best.p')) |
| print(f' new best sum={s:.4f} -> grpo_best.p', flush=True) |
|
|
| @staticmethod |
| def _cycle(dl): |
| while True: |
| for d in dl: |
| yield d |
|
|
|
|
| def parse_config(): |
| p = argparse.ArgumentParser() |
| p.add_argument('--cfg', default='led_augment') |
| p.add_argument('--info', default='grpo', type=str) |
| p.add_argument('--gpu', type=int, default=0) |
| p.add_argument('--cuda', default=True) |
| p.add_argument('--learning_rate', type=float, default=0.002) |
| p.add_argument('--warm_ckpt', type=str, |
| default='./results/led_augment/graph_v6_edge_relpos/models/model_0036.p') |
| p.add_argument('--edge_mode', default='relpos_only', type=str) |
| |
| p.add_argument('--batch', type=int, default=64) |
| p.add_argument('--grpo_lr', type=float, default=1e-4) |
| p.add_argument('--init_noise', type=float, default=0.1) |
| p.add_argument('--kl_beta', type=float, default=0.0) |
| p.add_argument('--clip_eps', type=float, default=0.2) |
| p.add_argument('--inner_epochs', type=int, default=2) |
| p.add_argument('--grpo_iters', type=int, default=1000) |
| p.add_argument('--eval_every', type=int, default=50) |
| p.add_argument('--max_eval_batches', type=int, default=5) |
| p.add_argument('--w_ade', type=float, default=1.0) |
| p.add_argument('--w_fde', type=float, default=1.0) |
| p.add_argument('--w_jade', type=float, default=1.0) |
| p.add_argument('--w_jfde', type=float, default=1.0) |
| |
| p.add_argument('--reward_mode', default='accuracy', choices=['accuracy', 'collision']) |
| p.add_argument('--w_ade_soft', type=float, default=0.3, help='soft ADE weight (hold accuracy)') |
| p.add_argument('--w_col', type=float, default=1.0, help='hard collision-count weight') |
| p.add_argument('--d_min', type=float, default=0.4) |
| p.add_argument('--ade_tol', type=float, default=0.80) |
| return p.parse_args() |
|
|
|
|
| def main(): |
| cfg = parse_config() |
| torch.cuda.set_device(cfg.gpu) |
| t = LEDGRPOTrainer(cfg) |
| t.train() |
|
|
|
|
| if __name__ == '__main__': |
| main() |
|
|