sra-trajectory-code / LED /main_led_nba_grpo.py
po03087's picture
SRA: MID/LED/MoFlow code + RUNNING.md instructions (code only, no data/ckpts)
d4cbafd verified
Raw
History Blame Contribute Delete
18.4 kB
"""
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()