""" 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 # ---- GRPO reward (inlined from MoFlow grpo/rewards.py to avoid sys.path clash) ---- 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) # [B,K,A,T] ade = err.mean(dim=-1) # [B,K,A] fde = err[..., -1] # [B,K,A] r_marg = -(w_ade * ade + w_fde * fde) jade = ade.mean(dim=2, keepdim=True) # [B,K,1] 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) # [B,K,A,T] ade = err.mean(dim=-1) # [B,K,A] (soft, accuracy) abs_p = pred + init_pos[:, None, :, None, :] # absolute positions mind = (abs_p.unsqueeze(3) - abs_p.unsqueeze(2)).norm(dim=-1).min(dim=-1).values # [B,K,A,A] pm = _player_mask(A, ball_idx, pred.device) hard = ((mind < d_min) & pm).float() # HARD indicator (non-diff) coll_count = hard.sum(dim=-1) # [B,K,A] #collisions of agent a 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): # graph config for the warm-start checkpoint (edge_relpos, v6, no sigma) 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) # ---- warm-start initializer + graph ---- 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}') # freeze graph + core denoiser; train ONLY the initializer 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() # bigger rollout batch than LED's default (10) for stable GRPO advantages 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) # frozen reference initializer (KL anchor) self.ref_initializer = copy.deepcopy(self.model_initializer).cuda().eval() for p in self.ref_initializer.parameters(): p.requires_grad_(False) # optimizer over the initializer only self.opt = torch.optim.AdamW(self.model_initializer.parameters(), lr=config.grpo_lr) # GRPO hyperparams 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)) # [B*A, K] sum over (T, 2) 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) # same-subset baseline (before any update) 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) # ---- rollout: sample action loc_s, decode, reward ---- with torch.no_grad(): loc = self.get_loc(past, traj_mask) # [B*A,K,T,2] z = torch.randn_like(loc) loc_s = loc + self.init_noise * z logp_old = self._logp(loc_s, loc, self.init_noise) # [B*A,K] pred = self.p_sample_loop_accelerate(past, traj_mask, loc_s) # decode loc_ref = self.get_loc_from(self.ref_initializer, past, traj_mask) # reward in metric units ([B,K,A,T,2], scaled by traj_scale) 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, :] # [B,A,2] absolute 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) # advantage per (agent): reshape reward [B,K,A] -> per-agent group over K adv = group_advantage(reward) # [B,K,A] # map advantage back to [B*A, K] to match logp layout adv_bak = adv.permute(0, 2, 1).reshape(B * A, self.G) # [B*A,K] logp_old_flat = logp_old loc_s_c = loc_s # ---- PPO update (initializer only) ---- stats = {} for _ in range(self.inner_epochs): self.opt.zero_grad() loc_new = self.get_loc(past, traj_mask) # grad logp_new = self._logp(loc_s_c, loc_new, self.init_noise) # [B*A,K] 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) # deterministic pred = self.p_sample_loop_accelerate(past, traj_mask, loc) # --- collision: absolute positions, player-pairs, ball(10) excluded --- ipos = data['pre_motion_3D'].cuda()[:, :, -1, :] # [B,A,2] Tf = fut.shape[1] absP = self._to_bkat(pred, B, A) * self.traj_scale + ipos[:, None, :, None, :] # [B,K,A,T,2] absG = (fut.view(B, A, Tf, 2) * self.traj_scale + ipos[:, :, None, :]).unsqueeze(1) # [B,1,A,T,2] 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) # [B,K] 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) # [B,1] 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) # [B*A,K,T,2] d = (fut_r - pred).norm(dim=-1) * self.traj_scale # [B*A,K,T] dB = d.view(B, A, self.G, d.shape[-1]) # [B,A,K,T] for ti in range(1, 5): e = 5 * ti # marginal: per-agent min over K ade = d[..., :e].mean(-1).min(dim=1)[0].sum() fde = d[..., e-1].min(dim=1)[0].sum() # joint: per-scene, mean over agents then min over K 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) # checkpoint: collision mode -> best collision@d_min with ADE guard; else -> best sum 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) # unused (grpo_lr used) 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) # GRPO 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) # collision (non-differentiable) reward 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()