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:
@@ -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),
|
||||
|
||||
Reference in New Issue
Block a user