gymnax

repository·main·Indexed 21 days ago

https://github.com/roberttlange/gymnax

JAX implementations of OpenAI Gym environments designed for high-throughput, massively vectorized experiments. It enables the use of jit, vmap, and pmap for accelerated environment transitions and batch rollouts. Supported environments include Classic Control (e.g., Pendulum-v1, CartPole-v1), MinAtar, Bsuite, and various miscellaneous tasks. The library follows a functional pattern, requiring explicit passing of random keys and environment parameters to maintain state.

Tokens
10.1K
Snippets
25
Records
29
Agent score
75%

What's inside gymnax

  1. Vectorize and accelerate environments with JAX primitives

    main

    Because 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.jit on env.step for faster single-step transitions.
    • Batch rollouts (vmap across keys): Use jax.vmap to run multiple environments in parallel using different random keys.
    • Meta-learning (vmap across parameters): Use jax.vmap to 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))
  2. Parallelize environments using jax.vmap

    main

    Because gymnax environments are pure functions, you can use jax.vmap to run multiple environment instances in parallel on a single device (e.g., a single GPU/TPU). This is highly efficient for batch rollouts.

    To parallelize, vmap the reset and step functions, specifying which axes correspond to the batch of random keys and environment states.

    num_envs = 8
    vmap_keys = jax.random.split(key, num_envs)
    
    # vmap reset: map over keys (axis 0), keep env_params constant (None)
    vmap_reset = jax.vmap(env.reset, in_axes=(0, None))
    
    # vmap step: map over keys, states, and actions (axis 0), keep env_params constant (None)
    vmap_step = jax.vmap(env.step, in_axes=(0, 0, 0, None))
    
    obs, state = vmap_reset(vmap_keys, env_params)
    obs, state, reward, done, _ = vmap_step(
        vmap_keys, state, jnp.zeros(num_envs), env_params
    )
  3. Perform batch rollouts using vmap over random keys

    main

    A key advantage of gymnax is the ability to collect data from multiple actors in parallel using JAX's vmap.

    To implement a batch rollout:

    1. Define a rollout function that performs a single episode using jax.lax.scan (or hk.scan) for efficiency.
    2. Transform the rollout function using hk.transform.
    3. Use jax.vmap to map the transformed function over a batch of random keys.

    This pattern is highly efficient for on-policy algorithms like A2C that require high data throughput.

    # 1. Define rollout with lax.scan
    def rollout(key_input, gamma, env_params, steps_in_episode):
        # ... (env.reset, hk.scan loop, etc.)
        return log_probs, advantages, entropies, jnp.sum(rewards), jnp.sum(regrets)
    
    # 2. Transform and vmap
    rollout_fn = hk.without_apply_key(hk.transform(rollout))
    batch_rollout_fn = jax.vmap(rollout_fn.apply, in_axes=(None, 0, None, None, None))
    
    # 3. Use in training
    log_probs, advantages, entropies, reward, regrets = rollout_fn.apply(
        net_params, key, 0.8, env_params, 100
    )
  4. Basic `gymnax` API Usage

    main

    The gymnax API follows a functional pattern designed for JAX. Instead of maintaining internal state, environment methods require explicit passing of random keys and environment parameters. This allows for seamless jit, vmap, and pmap acceleration.

    Key workflow steps:

    1. Instantiate: Use gymnax.make(env_name) to get the environment object and its parameters.
    2. Reset: Call env.reset(key, env_params) to get the initial observation and state.
    3. Action: Use env.action_space(env_params).sample(key) to get a random action.
    4. Step: Call env.step(key, state, action, env_params) to transition the environment.
    import jax
    import gymnax
    
    key = jax.random.key(0)
    key, key_reset, key_act, key_step = jax.random.split(key, 4)
    
    # Instantiate the environment & its settings.
    env, env_params = gymnax.make("Pendulum-v1")
    
    # Reset the environment.
    obs, state = env.reset(key_reset, env_params)
    
    # Sample a random action.
    action = env.action_space(env_params).sample(key_act)
    
    # Perform the step transition.
    n_obs, n_state, reward, done, _ = env.step(key_step, state, action, env_params)
  5. Install gymnax

    main

    You can install the latest stable release of gymnax from PyPI, or install the latest commit directly from the GitHub repository.

    To use JAX on hardware accelerators (like GPUs or TPUs), ensure you follow the JAX installation guide.

    # Install from PyPI
    pip install gymnax
    
    # Install latest commit from GitHub
    pip install git+https://github.com/RobertTLange/gymnax.git@main
  6. Scan through entire episode rollouts with lax.scan

    main

    For maximum performance and fast compilation, you can use jax.lax.scan to loop through an entire episode (reset and multiple steps) within a single compiled function. This avoids the overhead of Python loops and allows the entire rollout to be treated as a single JAX operation.

    def rollout(key_input, policy_params, env_params, steps_in_episode):
        """Rollout a jitted gymnax episode with lax.scan."""
        # Reset the environment
        key_reset, key_episode = jax.random.split(key_input)
        obs, state = env.reset(key_reset, env_params)
    
        def policy_step(state_input, tmp):
            """lax.scan compatible step transition in jax env."""
            obs, state, policy_params, key = state_input
            key, key_step, key_net = jax.random.split(key, 3)
            action = model.apply(policy_params, obs)
            next_obs, next_state, reward, done, _ = env.step(
                key_step, state, action, env_params
            )
            carry = [next_obs, next_state, policy_params, key]
            return carry, [obs, action, reward, next_obs, done]
    
        # Scan over episode step loop
        _, scan_out = jax.lax.scan(
            policy_step,
            [obs, state, policy_params, key_episode],
            (),
            steps_in_episode
        )
        # Return masked sum of rewards accumulated by agent in episode
        obs, action, reward, next_obs, done = scan_out
        return obs, action, reward, next_obs, done
  7. Perform efficient episode rollouts with jax.lax.scan

    main

    For high-performance RL training, you should implement episode rollouts using jax.lax.scan instead of Python loops. This allows the entire rollout loop to be compiled into a single XLA kernel via jax.jit.

    In a scan loop, you define a policy_step function that takes the current state (including the policy parameters and random key) and returns the updated state and the data to be recorded (observations, actions, rewards, etc.).

    def rollout(key_input, policy_params, env_params, steps_in_episode):
        key_reset, key_episode = jax.random.split(key_input)
        obs, state = env.reset(key_reset, env_params)
    
        def policy_step(state_input, _):
            obs, state, policy_params, key = state_input
            key, key_step, key_net = jax.random.split(key, 3)
            
            # Apply policy (e.g., from a Flax model)
            action = model.apply(policy_params, obs, key_net)
            
            next_obs, next_state, reward, done, _ = env.step(
                key_step, state, action, env_params
            )
            
            carry = [next_obs, next_state, policy_params, key]
            return carry, [obs, action, reward, next_obs, done]
    
        _, scan_out = jax.lax.scan(
            policy_step, [obs, state, policy_params, key_episode], (), steps_in_episode
        )
        
        obs, action, reward, next_obs, done = scan_out
        return obs, action, reward, next_obs, done
    
    # Compile the rollout for maximum speed
    jit_rollout = jax.jit(rollout, static_argnums=3)
    obs, action, reward, next_obs, done = jit_rollout(key, policy_params, env_params, 200)
  8. Use the basic gymnax API: make, reset, and step

    main

    The core workflow in gymnax involves creating an environment, resetting it to an initial state, and stepping through transitions. Unlike standard Gym, gymnax requires explicit handling of JAX random keys and environment parameters to remain pure and JIT-compatible.

    1. gymnax.make(env_name): Creates the environment object and its default env_params.
    2. env.reset(key, env_params): Returns the initial observation and state.
    3. env.action_space(env_params).sample(key): Samples a random action.
    4. env.step(key, state, action, env_params): Returns next_obs, next_state, reward, done, and info.
    import gymnax
    import jax
    import jax.numpy as jnp
    
    key = jax.random.key(0)
    key, key_reset, key_policy, key_step = jax.random.split(key, 4)
    
    # 1. Instantiate environment
    env, env_params = gymnax.make("Pendulum-v1")
    
    # 2. Reset
    obs, state = env.reset(key_reset, env_params)
    
    # 3. Sample action
    action = env.action_space(env_params).sample(key_policy)
    
    # 4. Step
    obs, state, reward, done, _ = env.step(key_step, state, action, env_params)
  9. Install dependencies for Meta-RL examples

    main

    To run the Meta-RL examples, you need gymnax along with several JAX-compatible libraries for neural networks, optimization, and probability distributions:

    !pip install --quiet dm-haiku optax distrax chex gymnax
  10. Meta-learning with varying environment parameters

    main

    To train a policy (like a Meta-LSTM) that is robust to variations in environment dynamics (e.g., different link lengths in Acrobot-v1), incorporate parameter sampling into your rollout function.

    Instead of using a fixed env_params, sample them from a distribution at the start of every episode within your rollout loop. This ensures the policy is evaluated against a diverse set of environment configurations during a single population evaluation.

    Key Steps:

    1. Define a sampling function for EnvParams (e.g., sample_link_params).
    2. In your rollout function, use jax.random.split to generate a unique key for environment parameter sampling.
    3. Use jax.lax.scan to iterate through time steps, passing the sampled env_params into env.step.
    4. Use jax.vmap to parallelize these rollouts across the population.
    # 1. Sample parameters
    def sample_link_params(key, min_link=0.1, max_link=1.9):
        link_length_1 = jax.random.uniform(key, (), minval=min_link, maxval=max_link)
        link_length_2 = 2 - link_length_1
        return EnvParams(link_length_1=link_length_1, link_length_2=link_length_2)
    
    # 2. Incorporate into rollout
    def rollout(key_input, policy_params, steps_in_episode):
        key_reset, key_episode, key_link = jax.random.split(key_input, 3)
        env_params = sample_link_params(key_link) # Sample per episode
        obs, state = env.reset(key_reset, env_params)
        
        # ... use lax.scan with env.step(..., env_params) ...
    
    # 3. Parallelize
    key_rollout = jax.vmap(rollout, in_axes=(0, None, None))
    pop_rollout = jax.jit(jax.vmap(key_rollout, in_axes=(None, 0, None)), static_argnums=2)
  11. Parallelize fitness rollouts with RolloutWrapper

    main

    To parallelize fitness rollouts across population members and initial conditions, use the RolloutWrapper from gymnax.experimental. This wrapper manages the interaction between a policy (e.g., a model from evosax.NetworkMapper) and a gymnax environment.

    1. Define your model using NetworkMapper.
    2. Initialize policy parameters using model.init.
    3. Instantiate RolloutWrapper with the model's apply function and the target env_name.
    4. Use manager.single_rollout for a single episode or manager.population_rollout for parallelized evaluations.
    import jax
    import jax.numpy as jnp
    import gymnax
    from evosax import NetworkMapper
    from gymnax.experimental import RolloutWrapper
    
    # 1. Define model
    model = NetworkMapper["MLP"](
        num_hidden_units=64,
        num_hidden_layers=2,
        num_output_units=3,
        hidden_activation="relu",
        output_activation="categorical",
    )
    
    # 2. Initialize params
    key = jax.random.key(0)
    env, env_params = gymnax.make("Acrobot-v1")
    pholder = jnp.zeros(env.observation_space(env_params).shape)
    policy_params = model.init(key, x=pholder, key=key)
    
    # 3. Setup RolloutWrapper
    manager = RolloutWrapper(model.apply, env_name="Acrobot-v1")
    
    # 4. Run rollout
    obs, action, reward, next_obs, done, cum_ret = manager.single_rollout(
        key, policy_params
    )