Fix 6 critical bugs in REPPO repository preventing execution

- Fix missing MUON optimizer by replacing with optax.adam
- Fix Hydra configuration parameter paths (env.name instead of env_name)
- Fix BraxGymnaxWrapper method signatures to accept params argument
- Fix training loop division by zero with proper total_time_steps
- Fix incorrect algorithm name in wandb (reppo instead of sac)
- Fix JAX key batching error in BraxGymnaxWrapper reset method
- Add comprehensive HoReKa SLURM integration with wandb logging
- Update README with detailed bug documentation and fixes
This commit is contained in:
ys1087@partner.kit.edu
2025-07-22 17:26:43 +02:00
parent 137b9e80c9
commit b240a19ceb
7 changed files with 120 additions and 17 deletions
+5 -5
View File
@@ -24,7 +24,7 @@ from reppo_alg.env_utils.jax_wrappers import (
MjxGymnaxWrapper,
NormalizeVec,
)
from reppo_alg.jaxrl import utils, muon
from reppo_alg.jaxrl import utils
from reppo_alg.network_utils.jax_models import (
CategoricalCriticNetwork,
CriticNetwork,
@@ -239,15 +239,15 @@ def make_init(
if cfg.max_grad_norm is not None:
actor_optimizer = optax.chain(
optax.clip_by_global_norm(cfg.max_grad_norm),
muon.muon(lr), # optax.adam(lr) optax.adam(lr)
optax.adam(lr), # optax.adam(lr) optax.adam(lr)
)
critic_optimizer = optax.chain(
optax.clip_by_global_norm(cfg.max_grad_norm),
muon.muon(lr), # optax.adam(lr) optax.adam(lr)
optax.adam(lr), # optax.adam(lr) optax.adam(lr)
)
else:
actor_optimizer = muon.muon(lr) # optax.adam(lr)
critic_optimizer = muon.muon(lr) # optax.adam(lr)
actor_optimizer = optax.adam(lr) # optax.adam(lr)
critic_optimizer = optax.adam(lr) # optax.adam(lr)
actor_trainstate = nnx.TrainState.create(
graphdef=nnx.graphdef(actor_networks),