Add experiment infrastructure and production scripts
- Fix 6 critical bugs in original REPPO repository - Add comprehensive README documentation - Create production SLURM script for accelerated partition - Add experiment submission script for batch jobs - Algorithm now runs successfully with strong performance - Ready for paper replication experiments on Brax suite
This commit is contained in:
@@ -218,15 +218,26 @@ class BraxGymnaxWrapper:
|
||||
self.reward_scaling = reward_scaling
|
||||
|
||||
def reset(self, key):
|
||||
# Handle both single key and batched keys
|
||||
if key.ndim > 1: # Batched keys
|
||||
state = jax.vmap(self.env.reset)(key)
|
||||
else: # Single key
|
||||
# Handle both single keys and vectorized keys
|
||||
if key.ndim > 1:
|
||||
# Vectorized reset - use vmap
|
||||
reset_fn = jax.vmap(self.env.reset)
|
||||
state = reset_fn(key)
|
||||
else:
|
||||
# Single environment reset
|
||||
state = self.env.reset(key)
|
||||
return state.obs, state
|
||||
# Return obs, critic_obs, env_state (critic_obs = obs for Brax)
|
||||
return state.obs, state.obs, state
|
||||
|
||||
def step(self, key, state, action):
|
||||
next_state = self.env.step(state, action)
|
||||
# Handle both single and vectorized operations
|
||||
if key.ndim > 1:
|
||||
# Vectorized step - use vmap
|
||||
step_fn = jax.vmap(self.env.step, in_axes=(0, 0))
|
||||
next_state = step_fn(state, action)
|
||||
else:
|
||||
# Single environment step
|
||||
next_state = self.env.step(state, action)
|
||||
return (
|
||||
next_state.obs,
|
||||
next_state.obs,
|
||||
|
||||
Reference in New Issue
Block a user