feat: value loss clipping
This commit is contained in:
parent
1f8fdbdc59
commit
7a87facc8a
1 changed files with 13 additions and 1 deletions
|
|
@ -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))
|
||||
|
|
|
|||
Reference in a new issue