1
Fork 0

Environment setup + train loop (#5)

Mujoco environment setup (vectorized on GPU) + training loop + simulate script
This commit is contained in:
RobinMeersman 2026-03-26 09:53:54 +01:00 committed by GitHub
parent af9cf1fdc5
commit a85a7b8d89
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
27 changed files with 2155 additions and 234 deletions

View file

@ -0,0 +1,14 @@
from .environment.env_types import Backend, Task
from .environment.env_config import ArenaConfig, EnvConfig, MorphologyConfig
from .environment.factory import BrittleStarEnvFactory
from .environment.env_wrapper import BrittleStarEnv
__all__ = [
"ArenaConfig",
"Backend",
"BrittleStarEnv",
"BrittleStarEnvFactory",
"EnvConfig",
"MorphologyConfig",
"Task",
]

View file

@ -0,0 +1,10 @@
import flax.struct
import jax.numpy as jnp
@flax.struct.dataclass
class EpisodeStatistics:
episode_returns: jnp.array
episode_lengths: jnp.array
returned_episode_returns: jnp.array
returned_episode_lengths: jnp.array

View file

@ -0,0 +1,103 @@
from dataclasses import dataclass
@dataclass
class PPOArgs:
"""
source: https://github.com/vwxyzjn/cleanrl/blob/master/cleanrl/ppo_atari_envpool_xla_jax_scan.py
"""
# the name of this experiment
exp_name: str = "brittle_star_ppo"
# seed of the experiment
seed: int = 1
# if toggled, `torch.backends.cudnn.deterministic=False`
torch_deterministic: bool = True
# if toggled, cuda will be enabled by default
cuda: bool = True
# if toggled, this experiment will be tracked with Weights and Biases
track: bool = False
# the wandb's project name
wandb_project_name: str = "PPO-Modularity"
# the entity (team) of wandb's project
wandb_entity: str | None = None
# whether to capture videos of the agent performances (check out `videos` folder)
capture_video: bool = False
# whether to save model into the `runs/{run_name}` folder
save_model: bool = True
# whether to upload the saved model to huggingface
upload_model: bool = False
# the user or org name of the model repository from the Hugging Face Hub
hf_entity: str = ""
# ==== Algorithm specific dataclasses ====
# the id of the environment
env_id: str = "" # todo
# total timesteps of the experiments
total_timesteps: int = 10000000
# the learning rate of the optimizer
learning_rate: float = 2.5e-4
# the number of parallel game environments
num_envs: int = 16
# the number of steps to run in each environment per policy rollout
num_steps: int = 128
# Toggle learning rate annealing for policy and value networks
anneal_lr: bool = True
# the discount factor gamma
gamma: float = 0.99
# the lambda for the general advantage estimation
gae_lambda: float = 0.95
# the number of mini-batches
num_minibatches: int = 4
# the K epochs to update the policy
update_epochs: int = 4
# Toggles advantages normalization
norm_adv: bool = True
# the surrogate clipping coefficient
clip_coef: float = 0.1
# Toggles whether or not to use a clipped loss for the value function, as per the paper.
clip_vloss: bool = True
# coefficient of the entropy
ent_coef: float = 0.01
# coefficient of the value function
vf_coef: float = 0.5
# the maximum norm for the gradient clipping
max_grad_norm: float = 0.5
# the target KL divergence threshold
target_kl: float | None = None
# ==== to be filled in runtime ====
# the batch size (computed in runtime)
batch_size: int = 0
# the mini-batch size (computed in runtime)
minibatch_size: int = 0
# the number of iterations (computed in runtime)
num_iterations: int = 0

View file

@ -0,0 +1,8 @@
from .PPOArgs import PPOArgs
from .EpisodeStatistics import EpisodeStatistics
__all__ = [
"PPOArgs",
"EpisodeStatistics",
]

View file

@ -0,0 +1,78 @@
import jax
import jax.numpy as jnp
from brittle_star_project import (
EnvConfig,
BrittleStarEnvFactory,
MorphologyConfig,
ArenaConfig,
Backend,
)
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
)

View file

@ -0,0 +1,15 @@
from .env_config import ArenaConfig, EnvConfig, MorphologyConfig
from .env_types import Backend, Task
from .env_wrapper import BrittleStarEnv, StepResult
from .factory import BrittleStarEnvFactory
__all__ = [
"ArenaConfig",
"EnvConfig",
"MorphologyConfig",
"Backend",
"Task",
"BrittleStarEnv",
"StepResult",
"BrittleStarEnvFactory",
]

View file

@ -0,0 +1,53 @@
from __future__ import annotations
from dataclasses import dataclass, field
from .env_types import Task
@dataclass(frozen=True, slots=True)
class MorphologyConfig:
num_arms: int = 5
num_segments_per_arm: int = 4
use_p_control: bool = True
use_torque_control: bool = False
@dataclass(frozen=True, slots=True)
class ArenaConfig:
size: tuple[float, float] = (10.0, 5.0)
sand_ground_color: bool = True
attach_target: bool = True
wall_height: float = 1.5
wall_thickness: float = 0.1
@dataclass(frozen=True, slots=True)
class EnvConfig:
"""Shared environment settings.
Note: Some tasks have additional parameters (see fields below).
"""
task: Task = Task.DIRECTED_LOCOMOTION
simulation_time: float = 5.0
num_physics_steps_per_control_step: int = 10
time_scale: int = 2
camera_ids: list[int] = field(default_factory=lambda: [0, 1])
# (height, width)
render_size: tuple[int, int] = (480, 640)
joint_randomization_noise_scale: float = 0.0
# Directed locomotion
target_distance: float = 3.0
# Light escape
# Per docs in upstream env config: integer factors of 200.
light_perlin_noise_scale: int = 0
@staticmethod
def from_json(path: str) -> EnvConfig:
pass

View file

@ -0,0 +1,21 @@
from __future__ import annotations
from enum import Enum
class Backend(str, Enum):
"""Physics backend.
- MJC: MuJoCo C engine
- MJX: MuJoCo XLA (JAX) engine
"""
MJC = "MJC"
MJX = "MJX"
class Task(str, Enum):
"""Which brittle-star task/environment to instantiate."""
DIRECTED_LOCOMOTION = "directed_locomotion"
LIGHT_ESCAPE = "light_escape"

View file

@ -0,0 +1,87 @@
from __future__ import annotations
import inspect
from dataclasses import dataclass
from typing import Any
import numpy as np
from .env_config import EnvConfig
from .env_types import Backend
@dataclass(slots=True)
class StepResult:
state: Any
reward: float | None = None
terminated: bool | None = None
truncated: bool | None = None
info: dict[str, Any] | None = None
class BrittleStarEnv:
"""Thin wrapper around the underlying DualMuJoCoEnvironment.
Goal: hide backend-specific RNG setup and provide a stable place to plug in RL.
"""
def __init__(self, env: Any, *, backend: Backend, config: EnvConfig) -> None:
self._env = env
self._backend = backend
self._config = config
@property
def raw(self) -> Any:
return self._env
@property
def backend(self) -> Backend:
return self._backend
@property
def config(self) -> EnvConfig:
return self._config
def make_rng(self, seed: int):
if self._backend == Backend.MJC:
return np.random.RandomState(seed)
import jax
return jax.random.PRNGKey(seed)
def reset(self, *, seed: int = 0):
rng = self.make_rng(seed)
state = self._env.reset(rng=rng)
return state
def render(self, *, state: Any):
return self._env.render(state=state)
def close(self) -> None:
self._env.close()
def step(self, *, state: Any, action: Any, rng: Any | None = None) -> StepResult:
"""Best-effort step wrapper.
Different env libraries return different tuples; we normalize common cases.
"""
if not hasattr(self._env, "step"):
raise AttributeError("Underlying env has no step() method")
step_fn = self._env.step
sig = inspect.signature(step_fn)
params = list(sig.parameters)
# Common patterns:
# - step(state, action)
# - step(state, action, rng)
# - step(state, action, key)
# We pass rng only if the callable accepts a 3rd arg.
if len(params) >= 3 and rng is not None:
out = step_fn(state, action, rng)
else:
out = step_fn(state, action)
return out

View file

@ -0,0 +1,106 @@
from __future__ import annotations
from dataclasses import asdict
from moojoco.environment.dual import DualMuJoCoEnvironment
from .env_config import ArenaConfig, EnvConfig, MorphologyConfig
from .env_types import Backend, Task
class BrittleStarEnvFactory:
"""Creates brittle-star morphology, arena, and task environment instances."""
@staticmethod
def create_morphology(config: MorphologyConfig):
from biorobot.brittle_star.mjcf.morphology.morphology import (
MJCFBrittleStarMorphology,
)
from biorobot.brittle_star.mjcf.morphology.specification.default import (
default_brittle_star_morphology_specification,
)
spec = default_brittle_star_morphology_specification(
num_arms=config.num_arms,
num_segments_per_arm=config.num_segments_per_arm,
use_p_control=config.use_p_control,
use_torque_control=config.use_torque_control,
)
return MJCFBrittleStarMorphology(specification=spec)
@staticmethod
def create_arena(config: ArenaConfig):
from biorobot.brittle_star.mjcf.arena.aquarium import (
AquariumArenaConfiguration,
MJCFAquariumArena,
)
arena_config = AquariumArenaConfiguration(**asdict(config))
return MJCFAquariumArena(configuration=arena_config)
@staticmethod
def create_environment_configuration(config: EnvConfig):
# Import locally so the project can still be imported without these deps.
from biorobot.brittle_star.environment.directed_locomotion.shared import (
BrittleStarDirectedLocomotionEnvironmentConfiguration,
)
from biorobot.brittle_star.environment.light_escape.shared import (
BrittleStarLightEscapeEnvironmentConfiguration,
)
common = dict(
joint_randomization_noise_scale=config.joint_randomization_noise_scale,
render_mode="human",
simulation_time=config.simulation_time,
num_physics_steps_per_control_step=config.num_physics_steps_per_control_step,
time_scale=config.time_scale,
camera_ids=config.camera_ids,
render_size=config.render_size,
)
match config.task:
case Task.DIRECTED_LOCOMOTION:
return BrittleStarDirectedLocomotionEnvironmentConfiguration(
target_distance=config.target_distance,
**common,
)
case Task.LIGHT_ESCAPE:
return BrittleStarLightEscapeEnvironmentConfiguration(
light_perlin_noise_scale=config.light_perlin_noise_scale,
**common,
)
case _:
raise ValueError(f"Unsupported task: {config.task}")
@staticmethod
def create_environment(
backend: Backend,
morphology_config: MorphologyConfig,
arena_config: ArenaConfig,
env_config: EnvConfig,
) -> DualMuJoCoEnvironment:
from biorobot.brittle_star.environment.directed_locomotion.dual import (
BrittleStarDirectedLocomotionEnvironment,
)
from biorobot.brittle_star.environment.light_escape.dual import (
BrittleStarLightEscapeEnvironment,
)
morphology = BrittleStarEnvFactory.create_morphology(morphology_config)
arena = BrittleStarEnvFactory.create_arena(arena_config)
env_configuration = BrittleStarEnvFactory.create_environment_configuration(env_config)
match env_config.task:
case Task.DIRECTED_LOCOMOTION:
env_class = BrittleStarDirectedLocomotionEnvironment
case Task.LIGHT_ESCAPE:
env_class = BrittleStarLightEscapeEnvironment
case _:
raise ValueError(f"Unsupported task: {env_config.task}")
return env_class.from_morphology_and_arena(
morphology=morphology,
arena=arena,
configuration=env_configuration,
backend=backend.value,
)

View file

@ -0,0 +1,75 @@
from __future__ import annotations
import time
from dataclasses import dataclass
from typing import Any, Protocol
import numpy as np
@dataclass
class SimulationConfig:
realtime: bool = True
seed: int = 0
class ControlPolicy(Protocol):
def act(self, *, obs: np.ndarray | None = None, t: float = 0.0) -> np.ndarray: ...
def _default_observations(data: Any) -> np.ndarray:
qpos = np.asarray(data.qpos, dtype=np.float32).ravel()
qvel = np.asarray(data.qvel, dtype=np.float32).ravel()
return np.concatenate([qpos, qvel], axis=0)
def simulate_policy(
policy: ControlPolicy,
config: SimulationConfig,
state: Any | None = None,
) -> None:
"""Open MuJoCo's native viewer and step using actions from a policy.
This path drives MuJoCo physics directly (mj_step) and uses the policy output
as `data.ctrl`.
"""
import mujoco.viewer
model = state.mj_model
data = state.mj_data
start = time.time()
with mujoco.viewer.launch_passive(model, data) as viewer:
while viewer.is_running():
step_start = time.time()
t = time.time() - start
# Input vector for the policy
# TODO: custom input
obs = _default_observations(data)
# Policy action
ctrl = policy.act(obs=obs, t=t)
# Check if the policy output vector give an input for each actuator (nu)
# TODO: what if model trained on full morphology but we want to test on a damaged one?
# (nu mismatch)
if model.nu > 0:
ctrl = np.asarray(ctrl, dtype=np.float32).ravel()
if ctrl.shape != (model.nu,):
raise ValueError(
f"Policy returned ctrl shape {ctrl.shape}, expected ({model.nu},)"
)
data.ctrl[:] = ctrl
# Step the simulation and update the viewer
mujoco.mj_step(model, data)
viewer.sync()
# If we're running in realtime mode, sleep to maintain real-time pacing.
if config.realtime:
remaining = model.opt.timestep - (time.time() - step_start)
if remaining > 0:
time.sleep(remaining)

View file

@ -0,0 +1,71 @@
from dataclasses import dataclass, fields
import flax
import flax.linen as nn
import jax.numpy as jnp
import jax.tree_util
import numpy as np
from flax.linen.initializers import constant, orthogonal
class Network(nn.Module):
"""
Dummy model only used for testing purposes
inspired by: https://github.com/vwxyzjn/cleanrl/blob/master/cleanrl/ppo_atari_envpool_xla_jax_scan.py
"""
hidden_dim: int = 195
@nn.compact
def __call__(self, x):
x = nn.Dense(self.hidden_dim, kernel_init=orthogonal(np.sqrt(2)), bias_init=constant(0.0))(
x
)
x = nn.relu(x)
x = nn.Dense(self.hidden_dim, kernel_init=orthogonal(np.sqrt(2)), bias_init=constant(0.0))(
x
)
x = nn.relu(x)
return x
class Critic(nn.Module):
@nn.compact
def __call__(self, x):
return nn.Dense(1, kernel_init=orthogonal(1), bias_init=constant(0.0))(x)
class Actor(nn.Module):
action_dim: int
@nn.compact
def __call__(self, x):
mean = nn.Dense(self.action_dim, kernel_init=orthogonal(0.01), bias_init=constant(0.0))(x)
log_std = self.param("log_std", nn.initializers.zeros, (self.action_dim,))
return mean, log_std
@jax.tree_util.register_dataclass
@dataclass
class AgentParams:
network_params: flax.core.FrozenDict
actor_params: flax.core.FrozenDict
critic_params: flax.core.FrozenDict
@jax.tree_util.register_dataclass
@dataclass
class Storage:
obs: jnp.array
actions: jnp.array
logprobs: jnp.array
dones: jnp.array
values: jnp.array
advantages: jnp.array
returns: jnp.array
rewards: jnp.array
def replace(self, **kwargs) -> "Storage":
fs = fields(self)
return Storage(**{f.name: kwargs.get(f.name, getattr(self, f.name)) for f in fs})

View file

@ -0,0 +1,25 @@
from .DummyAgent import Network, Critic, Actor, AgentParams, Storage
from .base import (
RLAlgorithm,
RLModel,
Transition,
create_model,
register_rl_model,
registered_model_types,
)
from .random_policy_model import RandomPolicyModel
__all__ = [
"RLAlgorithm",
"RLModel",
"RandomPolicyModel",
"Transition",
"create_model",
"register_rl_model",
"registered_model_types",
"Network",
"Critic",
"Actor",
"AgentParams",
"Storage",
]

View file

@ -0,0 +1,162 @@
from __future__ import annotations
from abc import ABC, abstractmethod
from dataclasses import dataclass
import json
from pathlib import Path
from typing import Any
@dataclass(frozen=True, slots=True)
class Transition:
"""A minimal transition container for RL.
This is intentionally generic because the underlying env state type may be a
JAX pytree, a numpy struct, or something library-specific.
"""
obs: Any
action: Any
reward: float
next_obs: Any
terminated: bool
truncated: bool
info: dict[str, Any] | None = None
class RLAlgorithm(ABC):
"""Insertable RL algorithm interface."""
@abstractmethod
def select_action(self, *, obs: Any, rng: Any | None = None) -> Any:
raise NotImplementedError
def observe(self, transition: Transition) -> None:
"""Optional hook to store transitions."""
def update(self, *, rng: Any | None = None) -> dict[str, float]:
"""Optional hook to run one training update."""
return {}
def save(self, path: str) -> None:
raise NotImplementedError("Save not implemented")
def load(self, path: str) -> None:
raise NotImplementedError("Load not implemented")
_RL_MODEL_REGISTRY: dict[str, type["RLModel"]] = {}
def registered_model_types() -> list[str]:
return sorted(_RL_MODEL_REGISTRY)
def create_model(type_name: str, *, payload: dict[str, Any]) -> "RLModel":
model_cls = _RL_MODEL_REGISTRY.get(type_name)
if model_cls is None:
known = ", ".join(sorted(_RL_MODEL_REGISTRY)) or "<none>"
raise ValueError(f"Unknown RLModel type '{type_name}'. Known: {known}")
return model_cls.from_payload(payload)
def get_rl_model_registry() -> dict[str, type["RLModel"]]:
"""Return a copy of the current RLModel registry.
The registry is populated by importing concrete model modules that use the
`@register_rl_model(...)` decorator.
"""
return dict(_RL_MODEL_REGISTRY)
def register_rl_model(*type_names: str):
"""Decorator to register an `RLModel` for generic loading.
Concrete model modules should apply this decorator, so `base.py` never needs
to import concrete models (avoids circular imports).
"""
if not type_names:
raise TypeError("register_rl_model() requires at least one type name")
primary = type_names[0]
def _decorator(cls: type[RLModel]):
for name in type_names:
_RL_MODEL_REGISTRY[name] = cls
cls.type_name = primary
return cls
return _decorator
class RLModel(ABC):
"""Serializable policy/model interface.
This is the artifact that `train.py` writes and `simulate.py` loads.
"""
# Overwritten by the `@register_rl_model(...)` decorator.
type_name: str = "RLModel"
def reset(self, seed: int | None = None) -> None:
"""Optional hook for RNG/stateful models."""
@abstractmethod
def act(self, *, obs: Any | None = None, t: float = 0.0) -> Any:
raise NotImplementedError
def train(self, *, env: Any, num_epochs: int = 1) -> None:
"""Optional training hook.
Many models won't learn; for those this can be a no-op.
"""
_ = (env, num_epochs)
def to_payload(self) -> dict[str, Any]:
"""Return JSON-serializable model parameters."""
return {}
@classmethod
def from_payload(cls, payload: dict[str, Any]) -> "RLModel":
"""Reconstruct a model from `to_payload()` output."""
return cls(**payload) # type: ignore[arg-type]
def save(self, path: str | Path) -> Path:
out = Path(path)
out.parent.mkdir(parents=True, exist_ok=True)
doc = {
"type": self.type_name,
"version": 1,
"payload": self.to_payload(),
}
out.write_text(json.dumps(doc, indent=2, sort_keys=True) + "\n")
return out
@classmethod
def load(cls, path: str | Path) -> "RLModel":
p = Path(path)
doc = json.loads(p.read_text())
type_name = doc.get("type")
if not isinstance(type_name, str):
raise ValueError("Model artifact missing string field 'type'")
model_cls = _RL_MODEL_REGISTRY.get(type_name)
if model_cls is None:
known = ", ".join(sorted(_RL_MODEL_REGISTRY)) or "<none>"
raise ValueError(f"Unknown RLModel type '{type_name}'. Known: {known}")
payload = doc.get("payload")
# Backward compatibility: older artifacts stored fields at top-level.
if payload is None:
payload = {k: v for k, v in doc.items() if k not in ("type", "version")}
if not isinstance(payload, dict):
raise ValueError("Model artifact field 'payload' must be an object")
return model_cls.from_payload(payload)

View file

@ -0,0 +1,52 @@
from __future__ import annotations
from dataclasses import dataclass, field
import numpy as np
from .base import RLModel, register_rl_model
@register_rl_model("random")
@dataclass(slots=True)
class RandomPolicyModel(RLModel):
"""A minimal, serializable policy model that outputs random controls.
This is intentionally *not* a learning algorithm yet. It exists so we can:
- produce a stable model artifact from `train.py`
- load that artifact in `simulate.py`
- drive the MuJoCo viewer with the model's actions
"""
nu: int = 0
seed: int = 0
ctrl_noise_scale: float = 0.5
_rng: np.random.RandomState = field(init=False, repr=False)
def __post_init__(self) -> None:
self.reset(self.seed)
def reset(self, seed: int | None = None) -> None:
if seed is not None:
self.seed = int(seed)
self._rng = np.random.RandomState(self.seed)
def act(self, *, obs: np.ndarray | None = None, t: float = 0.0) -> np.ndarray:
if self.nu <= 0:
return np.zeros((0,), dtype=np.float32)
ctrl = self.ctrl_noise_scale * self._rng.randn(self.nu)
return ctrl.astype(np.float32)
def to_payload(self) -> dict[str, object]:
return {
"seed": int(self.seed),
"ctrl_noise_scale": float(self.ctrl_noise_scale),
}
@classmethod
def from_payload(cls, payload: dict[str, object]) -> RandomPolicyModel:
return cls(
seed=int(payload.get("seed", 0)),
ctrl_noise_scale=float(payload.get("ctrl_noise_scale", 0.5)),
)