From ccbbfc18d79c493c3d22bdd719db69aa5824734b Mon Sep 17 00:00:00 2001 From: cedric Date: Fri, 27 Mar 2026 03:27:26 +0000 Subject: [PATCH] fix: PPO extracted and integrated with jax --- ruff.toml | 1 - src/ppo.py | 179 ++++++++++++++++++++++++++++++++--------------------- 2 files changed, 107 insertions(+), 73 deletions(-) diff --git a/ruff.toml b/ruff.toml index 7bc4795..bbd1f99 100644 --- a/ruff.toml +++ b/ruff.toml @@ -347,7 +347,6 @@ extend-ignore = [ # "PLR1705", # no-else-return # "PLR1706", # consider-using-ternary # "PLR1707", # trailing-comma-tuple - "PLR1708", # stop-iteration-return # "PLR1709", # simplify-boolean-expression # "PLR1710", # inconsistent-return-statements "PLR1711", # useless-return diff --git a/src/ppo.py b/src/ppo.py index 85c4c6b..153c749 100644 --- a/src/ppo.py +++ b/src/ppo.py @@ -1,96 +1,131 @@ -# docs and experiment results can be found at https://docs.cleanrl.dev/rl-algorithms/ppo/#ppo_atari_envpool_xla_jaxpy +from functools import partial +import flax import jax import jax.numpy as jnp -from flax.training.train_state import TrainState -@jax.jit + +# Chose to use a class as it seemed the easiest way to integrate the CleanRL code style +# with our need to seperate concerns +class PPO: + def __init__(self, args, network, actor, critic): + self.args = args + self.network = network + self.actor = actor + self.critic = critic + + self.ppo_loss_grad_fn = jax.value_and_grad( + partial(ppo_loss, args=args, network=network, actor=actor, critic=critic), + has_aux=True, + ) + + # This PPO class should be initialized only once, + # or this function will need to recompile + @partial(jax.jit, static_argnums=0) + def update_ppo(self, agent_state, storage, key): + args = self.args + ppo_loss_grad_fn = self.ppo_loss_grad_fn + + def update_epoch(carry, _): + agent_state, key = carry + key, subkey = jax.random.split(key) + + def flatten(x): + return x.reshape((-1,) + x.shape[2:]) + + def convert_data(x): + x = jax.random.permutation(subkey, x) + return jnp.reshape(x, (args.num_minibatches, -1) + x.shape[1:]) + + flatten_storage = jax.tree.map(flatten, storage) + shuffled_storage = jax.tree.map(convert_data, flatten_storage) + + def update_minibatch(agent_state, minibatch): + (loss, (pg_loss, v_loss, entropy_loss, approx_kl)), grads = ( + ppo_loss_grad_fn( + agent_state.params, + minibatch.obs, + minibatch.actions, + minibatch.logprobs, + minibatch.advantages, + minibatch.returns, + ) + ) + agent_state = agent_state.apply_gradients(grads=grads) + return agent_state, ( + loss, + pg_loss, + v_loss, + entropy_loss, + approx_kl, + grads, + ) + + agent_state, metrics = jax.lax.scan( + update_minibatch, agent_state, shuffled_storage + ) + return (agent_state, key), metrics + + (agent_state, key), (loss, pg_loss, v_loss, entropy_loss, approx_kl, grads) = ( + jax.lax.scan( + update_epoch, (agent_state, key), (), length=args.update_epochs + ) + ) + return agent_state, loss, pg_loss, v_loss, entropy_loss, approx_kl, key + + +""" +Should be ok to use partial here, since the references to network, +actor and critic should not change at runtime +The cost of seperating concerns is to somehow pass these values +that are now not in the same scope +Chose to pass the actual apply functions, since they should never change +Other option was to pass the networks, but how does it influence compilation when their +class values would ever change? Better safe than sorry. +""" + + +@partial(jax.jit, static_argnums=(0, 1, 2)) def get_action_and_value2( + network_apply, + actor_apply, + critic_apply, params: flax.core.FrozenDict, - x: np.ndarray, - action: np.ndarray, + x: jnp.ndarray, + action: jnp.ndarray, ): - """calculate value, logprob of supplied `action`, and entropy""" - hidden = network.apply(params.network_params, x) - # assume that actor returns mean and log std over continuos action space, why log?, better for..?, research this - mean, log_std = actor.apply(params.actor_params, hidden) + hidden = network_apply(params["network_params"], x) + mean, log_std = actor_apply(params["actor_params"], hidden) std = jnp.exp(log_std) - - # compute logprob of the given action - var = std ** 2 - logprob = -0.5 * (((action - mean) ** 2) / var + 2 * log_std + jnp.log(2 * jnp.pi)) - logprob = logprob.sum(-1) - - # Validate that this is a good entropy (check other ppo implementations) - entropy = 0.5 + 0.5 * jnp.log(2 * jnp.pi) + log_std - entropy = entropy.sum(-1) - value = critic.apply(params.critic_params, hidden).squeeze() + logprob = -0.5 * ( + ((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi) + ).sum(-1) + entropy = (0.5 + 0.5 * jnp.log(2 * jnp.pi) + log_std).sum(-1) + value = critic_apply(params["critic_params"], hidden).squeeze(-1) + return logprob, entropy, value -def ppo_loss(params, x, a, logp, mb_advantages, mb_returns): - newlogprob, entropy, newvalue = get_action_and_value2(params, x, a) + +def ppo_loss( + params, x, a, logp, mb_advantages, mb_returns, args, network, actor, critic +): + newlogprob, entropy, newvalue = get_action_and_value2( + network.apply, actor.apply, critic.apply, params, x, a + ) logratio = newlogprob - logp ratio = jnp.exp(logratio) approx_kl = ((ratio - 1) - logratio).mean() if args.norm_adv: - mb_advantages = (mb_advantages - mb_advantages.mean()) / (mb_advantages.std() + 1e-8) + mb_advantages = (mb_advantages - mb_advantages.mean()) / ( + mb_advantages.std() + 1e-8 + ) - # Policy 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() - - # Value loss v_loss = 0.5 * ((newvalue - mb_returns) ** 2).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)) - -ppo_loss_grad_fn = jax.value_and_grad(ppo_loss, has_aux=True) - -@jax.jit -def update_ppo( - agent_state: TrainState, - storage: Storage, - key: jax.random.PRNGKey, -): - def update_epoch(carry, unused_inp): - agent_state, key = carry - key, subkey = jax.random.split(key) - - def flatten(x): - return x.reshape((-1,) + x.shape[2:]) - - # taken from: https://github.com/google/brax/blob/main/brax/training/agents/ppo/train.py - def convert_data(x: jnp.ndarray): - x = jax.random.permutation(subkey, x) - x = jnp.reshape(x, (args.num_minibatches, -1) + x.shape[1:]) - return x - - flatten_storage = jax.tree_map(flatten, storage) - shuffled_storage = jax.tree_map(convert_data, flatten_storage) - - def update_minibatch(agent_state, minibatch): - (loss, (pg_loss, v_loss, entropy_loss, approx_kl)), grads = ppo_loss_grad_fn( - agent_state.params, - minibatch.obs, - minibatch.actions, - minibatch.logprobs, - minibatch.advantages, - minibatch.returns, - ) - agent_state = agent_state.apply_gradients(grads=grads) - return agent_state, (loss, pg_loss, v_loss, entropy_loss, approx_kl, grads) - - agent_state, (loss, pg_loss, v_loss, entropy_loss, approx_kl, grads) = jax.lax.scan( - update_minibatch, agent_state, shuffled_storage - ) - return (agent_state, key), (loss, pg_loss, v_loss, entropy_loss, approx_kl, grads) - - (agent_state, key), (loss, pg_loss, v_loss, entropy_loss, approx_kl, grads) = jax.lax.scan( - update_epoch, (agent_state, key), (), length=args.update_epochs - ) - return agent_state, loss, pg_loss, v_loss, entropy_loss, approx_kl, key \ No newline at end of file