Files
hand-motion-pipeline/scripts/rl_track_reference.py
T
liyang ae28d55f81 Update to 2026-09-17 pipeline snapshot; add weights, L20 assets and recording via Git LFS
Source: RGB-D -> Dyn-HaMR -> L20 retargeting -> FoundationPose -> reference repair -> SPIDER,
documented in docs/PIPELINE_LATEST.md and docs/SETUP_AND_WEIGHTS.md. Adds FoundationPose and
nvdiffrast upstream snapshots, requirements/pipeline_venv.txt and the FoundationPose weight
manifest/downloader.

Assets (Git LFS): weights/ (WiLoR detector, HandFlow denoiser, UniDepth-L), FoundationPose
checkpoints, HaMeR checkpoint, Dyn-HaMR HMP model and BMC constraints, L20 URDF/meshes, the
20260915_171525 D405 recording and the two box CADs. MANO models are not redistributed
(third_party/hamer/_DATA/data/mano/README.txt). Environments, caches and run outputs excluded.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
2026-09-17 11:43:37 +08:00

195 lines
20 KiB
Python

"""PPO tracking policy: kinematic reference as PD targets + residual actions, MuJoCo Warp parallel physics.
Reward (DexMachina/ManipTrans style): object pose tracking + landmark imitation + joint-space bc + contact-point term
- action-rate/magnitude penalties. Virtual object controller (VOC) with decaying gains assists early training.
Reference state initialization: episodes start at random demo frames. Eval: full demo from t=0, VOC off, deterministic.
"""
import os, sys, json, time, math, argparse
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
os.environ['WARP_CACHE_PATH'] = str(ROOT / '.spider_cache/warp'); os.environ.setdefault('MUJOCO_GL', 'osmesa')
import numpy as np, torch, torch.nn as nn, mujoco, mujoco_warp as mw, warp as wp
p = argparse.ArgumentParser()
p.add_argument('--out', type=Path, default=ROOT / 'output/rl_track_20260915'); p.add_argument('--nworld', type=int, default=512); p.add_argument('--iters', type=int, default=1500)
p.add_argument('--horizon', type=int, default=24); p.add_argument('--decimation', type=int, default=8); p.add_argument('--eval-every', type=int, default=50); p.add_argument('--seed', type=int, default=0)
p.add_argument('--voc-kp', type=float, default=200.); p.add_argument('--resume', type=Path); p.add_argument('--smoke', action='store_true')
p.add_argument('--phase-control', action='store_true', help='extra action: reference playback rate in [0,1.5]'); p.add_argument('--wrist-pos-scale', type=float, default=.03); p.add_argument('--wrist-rot-scale', type=float, default=.15); p.add_argument('--finger-scale', type=float, default=.4); p.add_argument('--prog-weight', type=float, default=.2); p.add_argument('--rate-max', type=float, default=1.5); p.add_argument('--seg-start', type=int, default=-1, help='segment mode: episodes start in [seg-start, seg-start+seg-jitter] (400 Hz index)'); p.add_argument('--seg-end', type=int, default=-1); p.add_argument('--seg-jitter', type=int, default=80); p.add_argument('--prog-mode', choices=['bonus', 'scale'], default='bonus', help='bonus: + w*(rate-1); scale: tracking reward multiplied by rate')
p.add_argument('--voc-start', type=float, default=1.0); p.add_argument('--voc-zero-iters', type=int, default=0, help='iterations (from start) until assistance is forced to zero; 0 = 70%% of --iters'); p.add_argument('--contact-start-frac', type=float, default=0.0); p.add_argument('--vel-penalty', type=float, default=0.0)
A = p.parse_args(); A.out.mkdir(parents=True, exist_ok=True); torch.manual_seed(A.seed); np.random.seed(A.seed)
dev = torch.device('cuda:0'); wp.init()
# ---------------- reference ----------------
R = dict(np.load(A.out / 'reference.npz')); N = len(R['qpos']); ref = {k: torch.as_tensor(R[k], device=dev, dtype=torch.float32) for k in ['qpos', 'qvel', 'ctrl', 'contact', 'obj_pos', 'obj_quat', 'tip_pos', 'landmark_pos', 'cp_local']}
cp_obj = torch.as_tensor(R['cp_obj'], device=dev, dtype=torch.long); OBJ = [int(x) for x in R['obj_body_ids']]; TIPS = [int(x) for x in R['tip_site_ids']]; LM = [int(x) for x in R['landmark_site_ids']]; WRIST = [int(x) for x in R['wrist_qpos_adr']]
m = mujoco.MjModel.from_xml_path(str(R['scene'])); phys = json.loads(Path(str(R['physics_parameters'])).read_text()); m.opt.iterations = phys['solver_iterations']; m.opt.ls_iterations = phys['ls_iterations']
assert not m.actuator_gainprm[44:].any(); NQ, NU = m.nq, m.nu; HAND_Q = 54; NA_H = 44; NA = NA_H + (1 if A.phase_control else 0)
ctrl_lo = torch.as_tensor(m.actuator_ctrlrange[:NA_H, 0], device=dev); ctrl_hi = torch.as_tensor(m.actuator_ctrlrange[:NA_H, 1], device=dev); ctrl_lim = torch.as_tensor(m.actuator_ctrllimited[:NA_H].astype(bool), device=dev)
ACT_SCALE = torch.tensor(([A.wrist_pos_scale] * 3 + [A.wrist_rot_scale] * 3 + [A.finger_scale] * 16) * 2, device=dev)
d0 = mujoco.MjData(m); d0.qpos[:] = R['qpos'][0]; mujoco.mj_forward(m, d0)
with wp.ScopedDevice('cuda:0'):
wm = mw.put_model(m); wd = mw.put_data(m, d0, nworld=A.nworld, nconmax=1024, njmax=3072)
with wp.ScopedCapture() as cap: mw.step(wm, wd)
graph = cap.graph
T_ = lambda a: wp.to_torch(a)
qpos, qvel, ctrl, xfrc, site_xpos, xpos, xquat, cvel = T_(wd.qpos), T_(wd.qvel), T_(wd.ctrl), T_(wd.xfrc_applied), T_(wd.site_xpos), T_(wd.xpos), T_(wd.xquat), T_(wd.cvel)
EXTRA_RESET = [T_(getattr(wd, k)) for k in ['qacc_warmstart', 'qacc', 'qfrc_applied', 'act', 'act_dot', 'efc_force'] if hasattr(wd, k) and getattr(wd, k) is not None and getattr(wd, k).shape[0] == A.nworld]
NW = A.nworld; DEC = A.decimation; CTRL_STEPS = N // DEC
def kin():
torch.cuda.synchronize(); mw.kinematics(wm, wd); wp.synchronize()
def sim_steps(n):
torch.cuda.synchronize()
for _ in range(n): wp.capture_launch(graph)
wp.synchronize()
# ---------------- helpers ----------------
def quat_mul(a, b):
w1, x1, y1, z1 = a.unbind(-1); w2, x2, y2, z2 = b.unbind(-1)
return torch.stack([w1 * w2 - x1 * x2 - y1 * y2 - z1 * z2, w1 * x2 + x1 * w2 + y1 * z2 - z1 * y2, w1 * y2 - x1 * z2 + y1 * w2 + z1 * x2, w1 * z2 + x1 * y2 - y1 * x2 + z1 * w2], -1)
def quat_conj(q): return q * torch.tensor([1, -1, -1, -1], device=q.device)
def quat_rotate(q, v):
qv = torch.cat([torch.zeros_like(v[..., :1]), v], -1); return quat_mul(quat_mul(q, qv), quat_conj(q))[..., 1:]
def rotvec_err(q_ref, q): # rotation vector taking q to q_ref (world frame)
dq = quat_mul(q_ref, quat_conj(q)); dq = dq * torch.sign(dq[..., :1] + 1e-12); ang = 2 * torch.acos(dq[..., 0].clamp(-1, 1)); s = torch.sqrt((1 - dq[..., 0] ** 2).clamp_min(1e-9)); return dq[..., 1:] / s[..., None] * ang[..., None], ang
# ---------------- env state ----------------
t_idx = torch.zeros(NW, dtype=torch.long, device=dev); t_pos = torch.zeros(NW, device=dev); max_len = torch.zeros(NW, dtype=torch.long, device=dev); ep_len = torch.zeros(NW, dtype=torch.long, device=dev); prev_a = torch.zeros(NW, NA, device=dev)
voc_gain = torch.full((NW,), float(A.voc_start), device=dev); voc_gain[:16] = 0. # first 16 worlds never get object assistance
OBJ_QADR = torch.arange(54, 66, device=dev); OBJ_VADR = OBJ_QADR - 0 # nv == nq here (all hinge/slide)
assert m.nv == m.nq
CONTACT_FRAMES = torch.nonzero(ref['contact'][:N - 2 * DEC * 60].sum(-1) > 0).squeeze(-1)
def reset(mask, start=None):
idx = torch.nonzero(mask).squeeze(-1)
if len(idx) == 0: return
if start is None and A.seg_start >= 0:
s = A.seg_start + torch.randint(0, max(1, A.seg_jitter), (len(idx),), device=dev)
elif start is None:
s = torch.randint(0, N - 2 * DEC * 60, (len(idx),), device=dev)
if A.contact_start_frac > 0 and len(CONTACT_FRAMES):
pick = torch.rand(len(idx), device=dev) < A.contact_start_frac; s = torch.where(pick, CONTACT_FRAMES[torch.randint(0, len(CONTACT_FRAMES), (len(idx),), device=dev)], s)
else: s = torch.full((len(idx),), int(start), device=dev, dtype=torch.long)
qpos[idx] = ref['qpos'][s]; qvel[idx] = ref['qvel'][s] * 0.; ctrl[idx] = ref['ctrl'][s]; xfrc[idx] = 0
for arr in EXTRA_RESET: arr[idx] = 0
qpos[idx] = ref['qpos'][s]; t_idx[idx] = s; t_pos[idx] = s.float(); max_len[idx] = (((A.seg_end if A.seg_end > 0 else N) - s).float() / DEC * 1.5).long() + 5; ep_len[idx] = 0; prev_a[idx] = 0
def obj_state():
op = xpos[:, OBJ]; oq = xquat[:, OBJ]; ov = qvel[:, OBJ_QADR].view(NW, 2, 6); return op, oq, ov[..., :3], ov[..., 3:]
def apply_voc(t):
op, oq, lv, av = obj_state(); pr = ref['obj_pos'][t]; qr = ref['obj_quat'][t]; g = voc_gain[:, None, None]
F = (A.voc_kp * (pr - op) - 10. * lv).clamp(-20, 20) * g; rv, _ = rotvec_err(qr, oq); Tq = (2. * rv - .1 * av).clamp(-1, 1) * g
xfrc[:, OBJ, :3] = F; xfrc[:, OBJ, 3:] = Tq
def observe(t):
hq = qpos[:, :HAND_Q]; hv = qvel[:, :HAND_Q]; hq_ref = ref['qpos'][t, :HAND_Q]; op, oq, lv, av = obj_state(); pr = ref['obj_pos'][t]; qr = ref['obj_quat'][t]
rv, ang = rotvec_err(qr, oq); t2 = (t + 20).clamp(max=N - 1); t3 = (t + 40).clamp(max=N - 1)
tips = site_xpos[:, TIPS]; rel = (tips[:, :, None, :] - op[:, None, :, :]).reshape(NW, -1)
obs = torch.cat([hq, hv * .05, hq - hq_ref, (op - pr).reshape(NW, -1), oq.reshape(NW, -1), rv.reshape(NW, -1), lv.reshape(NW, -1) * .2, av.reshape(NW, -1) * .05, (ref['obj_pos'][t2] - op).reshape(NW, -1), (ref['obj_pos'][t3] - op).reshape(NW, -1), rel, ref['contact'][t], (t.float() / N)[:, None], voc_gain[:, None], prev_a], -1)
return obs
def reward(t, a):
op, oq, lv, av = obj_state(); pr = ref['obj_pos'][t]; qr = ref['obj_quat'][t]; ep = (op - pr).norm(dim=-1); _, ang = rotvec_err(qr, oq)
r_obj = (torch.exp(-20 * ep) * torch.exp(-3 * ang)).mean(-1)
lm = site_xpos[:, LM]; e_lm = (lm - ref['landmark_pos'][t]).norm(dim=-1).mean(-1); r_imi = torch.exp(-20 * e_lm)
dq = qpos[:, :HAND_Q] - ref['qpos'][t, :HAND_Q]; w = torch.ones(HAND_Q, device=dev); w[WRIST[:3]] = 10; w[WRIST[6:9]] = 10; r_bc = torch.exp(-2 * (dq.abs() * w).mean(-1))
con = ref['contact'][t] > .5; cpl = ref['cp_local'][t]; co = cp_obj[t]; oq_sel = torch.gather(oq, 1, co[..., None].expand(-1, -1, 4)); op_sel = torch.gather(op, 1, co[..., None].expand(-1, -1, 3))
target = quat_rotate(oq_sel, cpl) + op_sel; dtip = (site_xpos[:, TIPS] - target).norm(dim=-1); rc = torch.exp(-50 * dtip); n_on = con.sum(-1); r_con = torch.where(n_on > 0, (rc * con).sum(-1) / n_on.clamp_min(1), torch.full_like(r_obj, .5))
pen = .05 * ((a - prev_a[:, :NA_H]) ** 2).mean(-1) + .01 * (a ** 2).mean(-1)
if A.vel_penalty > 0:
tn = (t + DEC).clamp(max=N - 1); v_ref = (ref['obj_pos'][tn] - ref['obj_pos'][t]) / (DEC * float(m.opt.timestep)); dv = ((lv - v_ref) ** 2).sum(-1).clamp(max=4.).mean(-1); pen = pen + A.vel_penalty * dv
track = 1.0 * r_obj + .3 * r_imi + .1 * r_bc + .5 * r_con; r = track - pen
return r, dict(r_obj=r_obj, r_imi=r_imi, r_bc=r_bc, r_con=r_con, obj_err_cm=ep.mean(-1) * 100, obj_rot_deg=ang.mean(-1) * 57.3, lm_err_cm=e_lm * 100, track=track, pen=pen, rot_each=ang * 57.3, err_each=ep * 100)
FAIL_THR = torch.tensor(.15, device=dev)
def step(a):
t = t_idx; hand = ref['ctrl'][t, :NA_H] + a[:, :NA_H] * ACT_SCALE; hand = torch.where(ctrl_lim, hand.clamp(ctrl_lo, ctrl_hi), hand); ctrl[:, :NA_H] = hand; ctrl[:, NA_H:] = ref['ctrl'][t, NA_H:]
rate = (1. + .5 * a[:, NA_H]).clamp(0., A.rate_max) if A.phase_control else torch.ones(NW, device=dev)
apply_voc(t); sim_steps(DEC); t_pos.add_(DEC * rate); t_idx[:] = t_pos.round().long().clamp(max=N - 1); ep_len.add_(1); t = t_idx
r, info = reward(t, a[:, :NA_H]); op, oq, _, _ = obj_state(); ep = (op - ref['obj_pos'][t]).norm(dim=-1).max(-1).values
if A.phase_control:
if A.prog_mode == 'scale': r = info['track'] * rate - info['pen']
else: r = r + A.prog_weight * (rate - 1.)
info['rate'] = rate
wrist_err = torch.stack([(qpos[:, WRIST[:3]] - ref['qpos'][t][:, WRIST[:3]]).norm(dim=-1), (qpos[:, WRIST[6:9]] - ref['qpos'][t][:, WRIST[6:9]]).norm(dim=-1)], -1).max(-1).values
seg_end = A.seg_end if A.seg_end > 0 else N - 1
bad = ~torch.isfinite(qpos).all(-1) | (ep > FAIL_THR) | (wrist_err > .25) | (ep_len >= max_len); trunc = (t_pos >= seg_end) & ~bad
r = torch.where(bad, r - 1., r); done = bad | trunc; prev_a[:] = a; info['fail'] = bad.float(); info['trunc'] = trunc.float(); info['nan_worlds'] = (~torch.isfinite(qpos).all(-1)).float().sum().expand(1)
return r, done, trunc, info
# ---------------- PPO ----------------
class RunningNorm:
def __init__(s, n): s.mean = torch.zeros(n, device=dev); s.var = torch.ones(n, device=dev); s.count = 1e-4
def update(s, x):
x = x[torch.isfinite(x).all(-1)]
if len(x) < 2: return
bm, bv, bc = x.mean(0), x.var(0, unbiased=False), x.shape[0]; delta = bm - s.mean; tot = s.count + bc
s.mean = s.mean + delta * bc / tot; s.var = (s.var * s.count + bv * bc + delta ** 2 * s.count * bc / tot) / tot; s.count = tot
def __call__(s, x): return torch.nan_to_num((x - s.mean) / torch.sqrt(s.var + 1e-6), nan=0., posinf=10., neginf=-10.).clamp(-10, 10)
def state(s): return dict(mean=s.mean, var=s.var, count=s.count)
def mlp(i, o, h=(512, 256, 128)):
layers = []; last = i
for hh in h: layers += [nn.Linear(last, hh), nn.ELU()]; last = hh
return nn.Sequential(*layers, nn.Linear(last, o))
reset(torch.ones(NW, dtype=torch.bool, device=dev)); kin(); OBS = observe(t_idx).shape[1]
actor, critic = mlp(OBS, NA).to(dev), mlp(OBS, 1).to(dev); log_std = nn.Parameter(torch.full((NA,), math.log(.5), device=dev))
opt = torch.optim.Adam(list(actor.parameters()) + list(critic.parameters()) + [log_std], lr=3e-4); norm = RunningNorm(OBS); lr = 3e-4
start_iter = 0
if A.resume:
ck = torch.load(A.resume, map_location=dev); actor.load_state_dict(ck['actor']); critic.load_state_dict(ck['critic']); log_std.data[:] = ck['log_std']; norm.mean, norm.var, norm.count = ck['norm']['mean'], ck['norm']['var'], ck['norm']['count']; voc_gain[16:] = min(ck.get('voc', 1.), float(A.voc_start)); start_iter = ck.get('iter', 0); lr = ck.get('lr', lr)
for g_ in opt.param_groups: g_['lr'] = lr
GAMMA, LAM, CLIP, EPOCHS, MB = .99, .95, .2, 5, 4; H = A.horizon
from torch.utils.tensorboard import SummaryWriter
tb = SummaryWriter(str(A.out / 'tb')); csv = open(A.out / 'train_log.csv', 'a')
def evaluate(it):
res_all = {}
for mode in ['det', 'sto']:
res_all[mode] = evaluate_mode(it, mode)
res = res_all['det']; res.update({k + '_sto': v for k, v in res_all['sto'].items() if k not in ('iter',)})
(A.out / 'eval_log.jsonl').open('a').write(json.dumps(res) + '\n'); print('EVAL', json.dumps(res), flush=True)
torch.save(dict(actor=actor.state_dict(), critic=critic.state_dict(), log_std=log_std.data, norm=norm.state(), voc=float(voc_gain[16:].mean()), iter=it, lr=lr, obs_dim=OBS), A.out / f'ckpt_iter{it:05d}.pt'); torch.save(dict(actor=actor.state_dict(), critic=critic.state_dict(), log_std=log_std.data, norm=norm.state(), voc=float(voc_gain[16:].mean()), iter=it, lr=lr, obs_dim=OBS), A.out / 'ckpt_latest.pt')
reset(torch.ones(NW, dtype=torch.bool, device=dev)); kin()
def evaluate_mode(it, mode):
reset(torch.ones(NW, dtype=torch.bool, device=dev), start=max(0, A.seg_start)); kin(); g_save = voc_gain.clone(); voc_gain[:] = 0; errs = []; rots = []; rot_each = []; err_each = []; traj = []; times = []; alive = torch.ones(NW, dtype=torch.bool, device=dev); seg_end = A.seg_end if A.seg_end > 0 else N - 1
with torch.no_grad():
for k in range(int(CTRL_STEPS * 1.5) + 5):
o = norm(observe(t_idx.clamp(max=N - 1))); mu = actor(o); a = (mu if mode == 'det' else torch.distributions.Normal(mu, log_std.exp()).sample()).clamp(-3, 3); r, done, trunc, info = step(a)
alive &= ~(info['fail'] > 0); errs.append(info['obj_err_cm']); rots.append(info['obj_rot_deg']); rot_each.append(info['rot_each']); err_each.append(info['err_each']); traj.append(qpos[0].clone()); times.append(float(k * DEC * m.opt.timestep))
if t_pos[0] >= seg_end or (t_pos >= seg_end).all(): break
E = torch.stack(errs); Rr = torch.stack(rots); RE = torch.stack(rot_each); EE = torch.stack(err_each); voc_gain[:] = g_save; ok = torch.isfinite(E).all(0) & torch.isfinite(Rr).all(0); E = E[:, ok]; Rr = Rr[:, ok]; alive = alive[ok]; RE = RE[:, ok]; EE = EE[:, ok]
nlast = max(1, len(errs) // 10); final_rot = RE[-nlast:].mean(0); final_err = EE[-nlast:].mean(0) # per world, per object, last 10% of the rollout
res = dict(iter=it, finite_worlds=int(ok.sum()), final_rot_deg_red_median=float(final_rot[:, 0].median()), final_rot_deg_blue_median=float(final_rot[:, 1].median()), flip_success_blue_frac=float(((final_rot[:, 1] < 30) & (final_err[:, 1] < 10)).float().mean()), red_ok_frac=float(((final_rot[:, 0] < 30) & (final_err[:, 0] < 10)).float().mean()), phase_reached_frac=float((t_pos[0] - max(0, A.seg_start)) / (seg_end - max(0, A.seg_start))), eval_steps=len(errs), obj_err_cm_mean=float(E.mean()), obj_err_cm_max_over_time_mean=float(E.max(0).values.mean()), obj_rot_deg_mean=float(Rr.mean()), survived_frac=float(alive.float().mean()), success_5cm_frac=float((E.max(0).values < 5).float().mean()), success_10cm_frac=float((E.max(0).values < 10).float().mean()))
np.savez_compressed(A.out / (f'eval_iter{it:05d}.npz' if mode == 'det' else f'eval_iter{it:05d}_sto.npz'), qpos=torch.stack(traj).cpu().numpy(), time=np.array(times), metrics=json.dumps(res), final_rot_deg=final_rot.cpu().numpy(), final_err_cm=final_err.cpu().numpy())
return res
robj_hist = []
iters = 3 if A.smoke else A.iters
for it in range(start_iter, start_iter + iters):
t0 = time.perf_counter(); obs_b = torch.zeros(H, NW, OBS, device=dev); act_b = torch.zeros(H, NW, NA, device=dev); mu_b = torch.zeros(H, NW, NA, device=dev); std_old = log_std.exp().detach().clone(); logp_b = torch.zeros(H, NW, device=dev); rew_b = torch.zeros(H, NW, device=dev); done_b = torch.zeros(H, NW, device=dev); val_b = torch.zeros(H + 1, NW, device=dev); agg = {}
with torch.no_grad():
for k in range(H):
o_raw = observe(t_idx.clamp(max=N - 1)); norm.update(o_raw); o = norm(o_raw); mu = actor(o); std = log_std.exp(); dist = torch.distributions.Normal(mu, std); a = dist.sample().clamp(-3, 3); logp = dist.log_prob(a).sum(-1); v = critic(o).squeeze(-1)
r, done, trunc, info = step(a); r = torch.nan_to_num(r, nan=-1., posinf=1., neginf=-1.)
# bootstrap truncated episodes with value of next state
if trunc.any():
kin(); vn = critic(norm(observe(t_idx.clamp(max=N - 1)))).squeeze(-1); r = torch.where(trunc & ~(info['fail'] > 0), r + GAMMA * vn, r)
obs_b[k], act_b[k], logp_b[k], rew_b[k], done_b[k], val_b[k], mu_b[k] = o, a, logp, r, done.float(), v, mu
for kk, vv in info.items(): agg.setdefault(kk, []).append(vv.mean().item())
reset(done); kin()
val_b[H] = critic(norm(observe(t_idx.clamp(max=N - 1)))).squeeze(-1)
adv = torch.zeros(H, NW, device=dev); last = torch.zeros(NW, device=dev)
for k in reversed(range(H)):
nd = 1. - done_b[k]; delta = rew_b[k] + GAMMA * val_b[k + 1] * nd - val_b[k]; last = delta + GAMMA * LAM * nd * last; adv[k] = last
ret = adv + val_b[:H]; adv_n = (adv - adv.mean()) / (adv.std() + 1e-8)
B = H * NW; ob, ab, lpb, rb, advb, vb, mub = obs_b.reshape(B, -1), act_b.reshape(B, -1), logp_b.reshape(B), ret.reshape(B), adv_n.reshape(B), val_b[:H].reshape(B), mu_b.reshape(B, -1); kls = []
for ep in range(EPOCHS):
perm = torch.randperm(B, device=dev)
for mb in range(MB):
i = perm[mb * B // MB:(mb + 1) * B // MB]; mu_new = actor(ob[i]); dist = torch.distributions.Normal(mu_new, log_std.exp()); lp = dist.log_prob(ab[i]).sum(-1); ratio = torch.exp(lp - lpb[i])
s1 = ratio * advb[i]; s2 = ratio.clamp(1 - CLIP, 1 + CLIP) * advb[i]; pl = -torch.min(s1, s2).mean(); vpred = critic(ob[i]).squeeze(-1); vcl = vb[i] + (vpred - vb[i]).clamp(-CLIP, CLIP); vl = torch.max((vpred - rb[i]) ** 2, (vcl - rb[i]) ** 2).mean()
loss = pl + .5 * vl; opt.zero_grad(); loss.backward(); nn.utils.clip_grad_norm_(list(actor.parameters()) + list(critic.parameters()) + [log_std], 1.); opt.step()
with torch.no_grad(): # analytic KL(old || new) between diagonal Gaussians (rsl_rl style)
s_new = log_std.exp(); kl = (torch.log(s_new / std_old) + (std_old ** 2 + (mub[i] - mu_new) ** 2) / (2 * s_new ** 2) - .5).sum(-1).mean().item(); kls.append(kl)
if kl > .02: lr = max(1e-5, lr / 1.5)
elif kl < .005: lr = min(1e-3, lr * 1.5)
for g_ in opt.param_groups: g_['lr'] = lr
# VOC curriculum: decay when object tracking is good; hard schedule guarantees zero by 70% of training
robj_hist.append(np.mean(agg['r_obj'])); gated = len(robj_hist) >= 10 and np.mean(robj_hist[-10:]) > .6
zero_iters = A.voc_zero_iters if A.voc_zero_iters > 0 else int(.7 * A.iters); sched = max(0., A.voc_start * (1. - (it - start_iter) / zero_iters)); g = voc_gain[16:].mean().item()
if gated: g *= .97
g = min(g, sched); g = 0. if g < .02 else g; voc_gain[16:] = g
fps = B * DEC / (time.perf_counter() - t0); row = dict(iter=it, fps=int(fps), rew=float(rew_b.mean()), ep_len=float(ep_len.float().mean()), voc=g, lr=lr, kl=float(np.mean(kls)), std=float(log_std.exp().mean()), **{k: float(np.mean(v)) for k, v in agg.items()})
for k, v in row.items(): tb.add_scalar(k, v, it)
csv.write(json.dumps(row) + '\n'); csv.flush()
if it % 5 == 0 or A.smoke: print(' '.join(f'{k}={v:.3g}' if isinstance(v, float) else f'{k}={v}' for k, v in row.items()), flush=True)
if (it + 1) % A.eval_every == 0 or A.smoke and it == start_iter + iters - 1: evaluate(it + 1)
print('TRAIN_DONE', flush=True)