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
+7 -3
View File
@@ -218,7 +218,11 @@ class BraxGymnaxWrapper:
self.reward_scaling = reward_scaling
def reset(self, key):
state = self.env.reset(key)
# Handle both single key and batched keys
if key.ndim > 1: # Batched keys
state = jax.vmap(self.env.reset)(key)
else: # Single key
state = self.env.reset(key)
return state.obs, state
def step(self, key, state, action):
@@ -232,7 +236,7 @@ class BraxGymnaxWrapper:
{},
)
def observation_space(self):
def observation_space(self, params=None):
return spaces.Box(
low=-jnp.inf,
high=jnp.inf,
@@ -243,7 +247,7 @@ class BraxGymnaxWrapper:
shape=(self.env.observation_size,),
)
def action_space(self):
def action_space(self, params=None):
return spaces.Box(
low=-1.0,
high=1.0,