@nnx.jit
def _ppo_step(policy, opt, obs, act, adv, logp_old, mask, epsilon,
entropy_coef, use_clip):
def loss_fn(policy):
logp_all = jax.nn.log_softmax(policy(obs), axis=-1)
logp = jnp.take_along_axis(logp_all, act[:, None], -1).squeeze(-1)
rho = jnp.exp(logp - logp_old)
surr = jnp.where(use_clip, jnp.minimum(
rho * adv, jnp.clip(rho, 1 - epsilon, 1 + epsilon) * adv),
rho * adv)
entropy = -(jnp.exp(logp_all) * logp_all).sum(-1)
loss = -(mask * (surr + entropy_coef * entropy)).sum() / mask.sum()
return loss, (rho, logp, entropy)
(_, (rho, logp, entropy)), grads = nnx.value_and_grad(
loss_fn, has_aux=True)(policy)
opt.update(policy, grads)
n = mask.sum()
return ((mask * (jnp.abs(rho - 1) > epsilon)).sum() / n,
(mask * (logp_old - logp)).sum() / n, (mask * entropy).sum() / n)
def ppo_epochs(ac, batch, adv, logp_old, epsilon, num_epochs,
entropy_coef=0.01, use_clip=True):
"""num_epochs clipped-surrogate passes on one frozen batch; returns
[num_epochs, 3] numpy diagnostics: fraction of ratios outside the
band, approximate KL, mean policy entropy."""
size = 1 << max(6, (len(adv) - 1).bit_length())
mask = jnp.asarray((np.arange(size) < len(adv)).astype(np.float32))
obs, act, adv, logp_old = (_pad(np.asarray(x), size) for x in
(batch.obs, batch.act, adv, logp_old))
step = nnx.cached_partial(_ppo_step, ac.policy, ac.opt_pi)
return np.array([step(obs, act, adv, logp_old, mask, epsilon,
entropy_coef, use_clip)
for _ in range(num_epochs)])