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:
ys1087@partner.kit.edu
2025-07-22 18:47:43 +02:00
parent 6e3ecb95ff
commit 1caaa9d01f
5 changed files with 175 additions and 19 deletions
+17 -6
View File
@@ -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,