File size: 18,428 Bytes
d4cbafd | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 | """
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()
|