1
Fork 0

feat: value loss clipping

This commit is contained in:
JibrilExe 2026-04-14 19:23:20 +02:00
parent 1f8fdbdc59
commit 7a87facc8a

View file

@ -143,7 +143,19 @@ def ppo_loss(
pg_loss1 = -mb_advantages * ratio
pg_loss2 = -mb_advantages * jnp.clip(ratio, 1 - args.clip_coef, 1 + args.clip_coef)
pg_loss = jnp.maximum(pg_loss1, pg_loss2).mean()
v_loss = 0.5 * ((newvalue - mb_returns) ** 2).mean()
old_value = mb_returns - mb_advantages
v_clipped = old_value + jnp.clip(
newvalue - old_value,
-args.clip_coef,
args.clip_coef
)
v_loss_unclipped = (newvalue - mb_returns) ** 2
v_loss_clipped = (v_clipped - mb_returns) ** 2
v_loss = 0.5 * jnp.maximum(v_loss_unclipped, v_loss_clipped).mean()
entropy_loss = entropy.mean()
loss = pg_loss - args.ent_coef * entropy_loss + v_loss * args.vf_coef
return loss, (pg_loss, v_loss, entropy_loss, jax.lax.stop_gradient(approx_kl))