From 25a6299eb5395761f46b9b6d476b7df72268ef2c Mon Sep 17 00:00:00 2001 From: JibrilExe Date: Sat, 21 Mar 2026 17:54:37 +0100 Subject: [PATCH 1/7] extracted ppo relevant code from cleanrl_atari, and changed to continuous output --- src/ppo.py | 96 ++++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 96 insertions(+) create mode 100644 src/ppo.py diff --git a/src/ppo.py b/src/ppo.py new file mode 100644 index 0000000..85c4c6b --- /dev/null +++ b/src/ppo.py @@ -0,0 +1,96 @@ +# docs and experiment results can be found at https://docs.cleanrl.dev/rl-algorithms/ppo/#ppo_atari_envpool_xla_jaxpy + +import jax +import jax.numpy as jnp +from flax.training.train_state import TrainState + +@jax.jit +def get_action_and_value2( + params: flax.core.FrozenDict, + x: np.ndarray, + action: np.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) + 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() + 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) + 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) + + # 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 From ccbbfc18d79c493c3d22bdd719db69aa5824734b Mon Sep 17 00:00:00 2001 From: cedric Date: Fri, 27 Mar 2026 03:27:26 +0000 Subject: [PATCH 2/7] 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 From 31c2d188b7ec7c2d5f4f920c8cd6d5db55766c74 Mon Sep 17 00:00:00 2001 From: cedric Date: Fri, 27 Mar 2026 04:09:56 +0000 Subject: [PATCH 3/7] feature: optional message passer function and clearer naming --- src/ppo.py | 70 ++++++++++++++++++++++++++++++++++++++++++++---------- 1 file changed, 58 insertions(+), 12 deletions(-) diff --git a/src/ppo.py b/src/ppo.py index 153c749..bce814e 100644 --- a/src/ppo.py +++ b/src/ppo.py @@ -8,14 +8,31 @@ import jax.numpy as jnp # 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): + """ + args: params for training + input_network: network that takes input state and returns hidden state + action_network: network that takes hidden state and returns action distribution + critic: value network used for ppo + message_passer: function that executes message passing and state agregation X times + """ + + def __init__( + self, args, input_network, action_network, critic, message_passer=None + ): self.args = args - self.network = network - self.actor = actor - self.critic = critic + + if not message_passer: + message_passer = identity self.ppo_loss_grad_fn = jax.value_and_grad( - partial(ppo_loss, args=args, network=network, actor=actor, critic=critic), + partial( + ppo_loss, + args=args, + input_network_apply=input_network.apply, + action_network_apply=action_network.apply, + critic_apply=critic.apply, + message_passer=message_passer, + ), has_aux=True, ) @@ -85,17 +102,19 @@ class values would ever change? Better safe than sorry. """ -@partial(jax.jit, static_argnums=(0, 1, 2)) +@partial(jax.jit, static_argnums=(0, 1, 2, 3)) def get_action_and_value2( - network_apply, - actor_apply, + input_apply, + action_apply, + message_passer, critic_apply, params: flax.core.FrozenDict, x: jnp.ndarray, action: jnp.ndarray, ): - hidden = network_apply(params["network_params"], x) - mean, log_std = actor_apply(params["actor_params"], hidden) + hidden = input_apply(params["network_params"], x) + hidden = message_passer(hidden) + mean, log_std = action_apply(params["actor_params"], hidden) std = jnp.exp(log_std) logprob = -0.5 * ( @@ -108,10 +127,26 @@ def get_action_and_value2( def ppo_loss( - params, x, a, logp, mb_advantages, mb_returns, args, network, actor, critic + params, + x, + a, + logp, + mb_advantages, + mb_returns, + args, + input_network_apply, + action_network_apply, + message_passer, + critic_apply, ): newlogprob, entropy, newvalue = get_action_and_value2( - network.apply, actor.apply, critic.apply, params, x, a + input_network_apply, + action_network_apply, + message_passer, + critic_apply, + params, + x, + a, ) logratio = newlogprob - logp ratio = jnp.exp(logratio) @@ -129,3 +164,14 @@ def ppo_loss( 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)) + + +def identity(hidden): + """ + Used for seamless jax integration, + avoids having branching inside jitted function, + used as message_passer in case it is not given, + (in case of centralized lvl) + """ + + return hidden From a452912b36d1624bc6640093419c94249dd212f3 Mon Sep 17 00:00:00 2001 From: cedric Date: Fri, 27 Mar 2026 04:40:13 +0000 Subject: [PATCH 4/7] removed unnecessary comments --- src/ppo.py | 11 ----------- 1 file changed, 11 deletions(-) diff --git a/src/ppo.py b/src/ppo.py index bce814e..b639eb2 100644 --- a/src/ppo.py +++ b/src/ppo.py @@ -8,14 +8,6 @@ import jax.numpy as jnp # 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: - """ - args: params for training - input_network: network that takes input state and returns hidden state - action_network: network that takes hidden state and returns action distribution - critic: value network used for ppo - message_passer: function that executes message passing and state agregation X times - """ - def __init__( self, args, input_network, action_network, critic, message_passer=None ): @@ -96,9 +88,6 @@ 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. """ From 98a8f529e17b13fd92ad38236d0fcd376b392b35 Mon Sep 17 00:00:00 2001 From: cedric Date: Fri, 27 Mar 2026 08:51:31 +0000 Subject: [PATCH 5/7] fix: avoids local ruff mismatch --- .pre-commit-config.yaml | 15 ++++++++------- ruff.toml | 4 +++- src/ppo.py | 38 +++++++++++++------------------------- 3 files changed, 24 insertions(+), 33 deletions(-) diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index c86e484..80c52e6 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -1,4 +1,12 @@ repos: + - repo: local + hooks: + - id: ruff-format + name: ruff format (pipx) + entry: pipx run ruff format + language: system + types: [python] + - repo: https://github.com/alessandrojcm/commitlint-pre-commit-hook rev: v9.16.0 hooks: @@ -6,13 +14,6 @@ repos: stages: [commit-msg] additional_dependencies: ["@commitlint/config-conventional"] - - repo: https://github.com/astral-sh/ruff-pre-commit - rev: v0.9.9 - hooks: - - id: ruff - args: [ --fix ] - - id: ruff-format - - repo: https://github.com/pre-commit/pre-commit-hooks rev: v5.0.0 hooks: diff --git a/ruff.toml b/ruff.toml index bbd1f99..cf728e1 100644 --- a/ruff.toml +++ b/ruff.toml @@ -1,3 +1,5 @@ +line-length = 100 + [lint] extend-select = [ # "PLC0103", # invalid-name @@ -347,6 +349,7 @@ 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 @@ -391,4 +394,3 @@ extend-ignore = [ # "PLW1404", # implicit-str-concat ] - diff --git a/src/ppo.py b/src/ppo.py index b639eb2..5cd92cb 100644 --- a/src/ppo.py +++ b/src/ppo.py @@ -8,9 +8,7 @@ import jax.numpy as jnp # 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, input_network, action_network, critic, message_passer=None - ): + def __init__(self, args, input_network, action_network, critic, message_passer=None): self.args = args if not message_passer: @@ -50,15 +48,13 @@ class PPO: 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, - ) + (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, ( @@ -70,15 +66,11 @@ class PPO: grads, ) - agent_state, metrics = jax.lax.scan( - update_minibatch, agent_state, shuffled_storage - ) + 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 - ) + (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 @@ -106,9 +98,7 @@ def get_action_and_value2( mean, log_std = action_apply(params["actor_params"], hidden) std = jnp.exp(log_std) - logprob = -0.5 * ( - ((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi) - ).sum(-1) + 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) @@ -142,9 +132,7 @@ def ppo_loss( 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) pg_loss1 = -mb_advantages * ratio pg_loss2 = -mb_advantages * jnp.clip(ratio, 1 - args.clip_coef, 1 + args.clip_coef) From d119bf6408f29dc26dc087e27bba3b17412ef8b3 Mon Sep 17 00:00:00 2001 From: cedric Date: Tue, 31 Mar 2026 16:03:03 +0000 Subject: [PATCH 6/7] fix: critic is no longer shared with actor network --- src/ppo.py | 19 +++++++++++++------ 1 file changed, 13 insertions(+), 6 deletions(-) diff --git a/src/ppo.py b/src/ppo.py index 5cd92cb..23c1e45 100644 --- a/src/ppo.py +++ b/src/ppo.py @@ -8,7 +8,9 @@ import jax.numpy as jnp # 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, input_network, action_network, critic, message_passer=None): + def __init__( + self, args, input_network, action_network, critic, critic_network, message_passer=None + ): self.args = args if not message_passer: @@ -21,6 +23,7 @@ class PPO: input_network_apply=input_network.apply, action_network_apply=action_network.apply, critic_apply=critic.apply, + critic_network_apply=critic_network.apply, message_passer=message_passer, ), has_aux=True, @@ -83,24 +86,26 @@ that are now not in the same scope """ -@partial(jax.jit, static_argnums=(0, 1, 2, 3)) +@partial(jax.jit, static_argnums=(0, 1, 2, 3, 4)) def get_action_and_value2( input_apply, action_apply, message_passer, critic_apply, + critic_network_apply, params: flax.core.FrozenDict, x: jnp.ndarray, action: jnp.ndarray, ): - hidden = input_apply(params["network_params"], x) - hidden = message_passer(hidden) - mean, log_std = action_apply(params["actor_params"], hidden) + hidden_network = input_apply(params["network_params"], x) + hidden_critic = critic_network_apply(params["critic_network_params"], x) + hidden_network = message_passer(hidden_network) + mean, log_std = action_apply(params["actor_params"], hidden_network) std = jnp.exp(log_std) 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) + value = critic_apply(params["critic_params"], hidden_critic).squeeze(-1) return logprob, entropy, value @@ -117,12 +122,14 @@ def ppo_loss( action_network_apply, message_passer, critic_apply, + critic_network_apply, ): newlogprob, entropy, newvalue = get_action_and_value2( input_network_apply, action_network_apply, message_passer, critic_apply, + critic_network_apply, params, x, a, From 4e4ab629bdb298406080fb346e6677009760fa14 Mon Sep 17 00:00:00 2001 From: cedric Date: Tue, 31 Mar 2026 21:11:16 +0000 Subject: [PATCH 7/7] fix: usage of uv for pre commit ruff format --- .pre-commit-config.yaml | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 80c52e6..b751166 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -2,8 +2,8 @@ repos: - repo: local hooks: - id: ruff-format - name: ruff format (pipx) - entry: pipx run ruff format + name: ruff format (uv) + entry: uv run ruff format language: system types: [python]