Vectorize and accelerate environments with JAX primitives
mainBecause gymnax is built with JAX, you can use jit, vmap, and pmap to accelerate environment transitions and perform batch rollouts.
Common acceleration patterns:
- JIT acceleration: Use
jax.jitonenv.stepfor faster single-step transitions. - Batch rollouts (vmap across keys): Use
jax.vmapto run multiple environments in parallel using different random keys. - Meta-learning (vmap across parameters): Use
jax.vmapto run environments in parallel with different environment parameters (e.g., different pendulum lengths).
# Jit-accelerated step transition
jit_step = jax.jit(env.step)
# map (vmap/pmap) across random keys for batch rollouts
reset_key = jax.vmap(env.reset, in_axes=(0, None))
step_key = jax.vmap(env.step, in_axes=(0, 0, 0, None))
# map (vmap/pmap) across env parameters (e.g. for meta-learning)
reset_params = jax.vmap(env.reset, in_axes=(None, 0))
step_params = jax.vmap(env.step, in_axes=(None, 0, 0, 0))