1
Fork 0
This repository has been archived on 2026-08-15. You can view files and clone it, but you cannot make any changes to it's state, such as pushing and creating new issues, pull requests or comments.
2026SEL3-project-Brittle_St.../src/brittle_star_project/environment/BrittleStarJaxEnvWrapper.py

97 lines
3 KiB
Python

import jax
import jax.numpy as jnp
from brittle_star_project import (
EnvConfig,
BrittleStarEnvFactory,
MorphologyConfig,
ArenaConfig,
Backend,
)
from brittle_star_project.environment import from_file
class BrittleStarJaxEnvWrapper:
def __init__(
self,
morphology: MorphologyConfig,
arena: ArenaConfig,
env_config: EnvConfig,
num_envs: int,
backend: Backend = Backend.MJX,
):
self._morphology = morphology
self._arena = arena
self._env_config = env_config
self._backend = backend
self._num_envs = num_envs
self._env = BrittleStarEnvFactory.create_environment(
self._backend, self._morphology, self._arena, self._env_config
)
self._vectorized_reset = jax.jit(jax.vmap(self._env.reset))
self._vectorized_step = jax.jit(jax.vmap(self._env.step))
self._vectorized_action_sample = jax.jit(jax.vmap(self._env.action_space.sample))
self._action_rng = None
@property
def backend(self):
return self._backend
@property
def raw(self):
return self._env
@property
def single_action_space(self):
return self._env.action_space
@property
def single_observation_space(self):
return self._env.observation_space
def reset(self, seed: int = 0):
self._action_rng, env_rng = jax.random.split(jax.random.PRNGKey(seed), 2)
env_rngs = jnp.array(jax.random.split(env_rng, self._num_envs))
return self._vectorized_reset(rng=env_rngs)
def sample_actions(self):
assert self._action_rng is not None, "Call reset() before sample_actions()"
self._action_rng, *sub_rngs = jnp.array(
jax.random.split(self._action_rng, self._num_envs + 1)
)
return self._vectorized_action_sample(rng=jnp.array(sub_rngs))
def step(self, state, action):
return self._vectorized_step(state=state, action=action)
def close(self):
self._env.close()
@staticmethod
def default(num_envs: int, backend: Backend = Backend.MJX) -> "BrittleStarJaxEnvWrapper":
morphology = MorphologyConfig()
arena = ArenaConfig()
env_config = EnvConfig()
return BrittleStarJaxEnvWrapper(
morphology, arena, env_config, num_envs=num_envs, backend=backend
)
@staticmethod
def from_config(
config_path: str, num_envs: int, backend: Backend = Backend.MJX
) -> "BrittleStarJaxEnvWrapper":
morphology_cfg, arena_cfg, env_cfg = from_file(config_path)
return BrittleStarJaxEnvWrapper(
morphology_cfg, arena_cfg, env_cfg, num_envs=num_envs, backend=backend
)
def __str__(self):
morphology_str = str(self._morphology)
arena_str = str(self._arena)
env_config_str = str(self._env_config)
return (
f"BrittleStarJaxEnvWrapper(backend={self._backend}, num_envs={self._num_envs}, "
+ f"morphology={morphology_str}, arena={arena_str}, env_config={env_config_str})"
)