1
Fork 0

fix: avoids local ruff mismatch

This commit is contained in:
cedric 2026-03-27 08:51:31 +00:00
parent a452912b36
commit 98a8f529e1
3 changed files with 24 additions and 33 deletions

View file

@ -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:

View file

@ -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
]

View file

@ -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)