1
Fork 0

refactor: top level imports

Co-authored-by: Robin Meersman <echteenrobin@gmail.com>
This commit is contained in:
Tibo De Peuter 2026-04-16 10:26:26 +02:00
parent a8298702c8
commit a33b8a2e8e
Signed by: tdpeuter
SSH key fingerprint: SHA256:u/h/LVoqKF1Iz02uOyxe6hcjmoZASCGV2HM0TG9ZMoU

View file

@ -1,13 +1,11 @@
import jax
import jax.numpy as jnp
from brittle_star_project import (
EnvConfig,
BrittleStarEnvFactory,
MorphologyConfig,
ArenaConfig,
Backend,
)
from experiment_logger import get_logger
from .env_config import EnvConfig, MorphologyConfig, ArenaConfig
from .env_types import Backend
from .factory import BrittleStarEnvFactory
from .padded_obs_wrapper import compute_padding_masks, pad_observations_batched
class BrittleStarJaxEnvWrapper:
@ -28,8 +26,7 @@ class BrittleStarJaxEnvWrapper:
self._backend, self._morphology, self._arena, self._env_config
)
from brittle_star_project.environment.padded_obs_wrapper import compute_padding_masks
# Pre-compute masks for observation padding
self._padding_masks = compute_padding_masks(self._morphology.segments_per_arm)
self._vectorized_reset = jax.jit(jax.vmap(self._env.reset))
@ -38,8 +35,6 @@ class BrittleStarJaxEnvWrapper:
self._action_rng = None
from experiment_logger import get_logger
self.logger = get_logger()
self.logger.info(
f"Initialized BrittleStarJaxEnvWrapper with {num_envs} envs on {backend.value}"
@ -66,7 +61,6 @@ class BrittleStarJaxEnvWrapper:
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))
state = self._vectorized_reset(rng=env_rngs)
from brittle_star_project.environment.padded_obs_wrapper import pad_observations_batched
state = state.replace(
observations=pad_observations_batched(state.observations, self._padding_masks)
@ -82,7 +76,6 @@ class BrittleStarJaxEnvWrapper:
def step(self, state, action):
next_state = self._vectorized_step(state=state, action=action)
from brittle_star_project.environment.padded_obs_wrapper import pad_observations_batched
next_state = next_state.replace(
observations=pad_observations_batched(next_state.observations, self._padding_masks)