From 98a8f529e17b13fd92ad38236d0fcd376b392b35 Mon Sep 17 00:00:00 2001 From: cedric Date: Fri, 27 Mar 2026 08:51:31 +0000 Subject: [PATCH] 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)