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