From 6db1129ca7de6f3739a32bdad76e19f919ba5a0f Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Thu, 2 Apr 2026 10:26:52 +0200 Subject: [PATCH 01/10] feat: implemented simulator --- experiments/simulate.py | 396 ++++++++++++++++++++++++++++++++++++---- 1 file changed, 361 insertions(+), 35 deletions(-) diff --git a/experiments/simulate.py b/experiments/simulate.py index 28062fa..3b79082 100644 --- a/experiments/simulate.py +++ b/experiments/simulate.py @@ -1,47 +1,346 @@ from __future__ import annotations import argparse +import time from pathlib import Path +from typing import Any +import flax +import jax +import jax.numpy as jnp +import numpy as np from brittle_star_project import ( Backend, - BrittleStarEnv, - BrittleStarEnvFactory, - SimulationConfig, - simulate_policy, ) from brittle_star_project.environment import from_json -from brittle_star_project.rl import RLModel # imports concrete models via rl.__init__ -from brittle_star_project.rl.base import get_rl_model_registry -MODEL_BY_NAME = get_rl_model_registry() -MODEL_OPTIONS = sorted(MODEL_BY_NAME) +def _flatten_obs_dict(obs_dict: dict[str, Any]) -> jnp.ndarray: + """Flatten the env's observation dict into a 1D vector. + + concatenates values in the dict's iteration order and skips empty arrays. + """ + + parts: list[jnp.ndarray] = [] + for v in obs_dict.values(): + arr = jnp.asarray(v) + if arr.size == 0: + continue + parts.append(arr.reshape((-1,))) + + if not parts: + return jnp.zeros((0,), dtype=jnp.float32) + return jnp.concatenate(parts, axis=0) + + +# A minimal policy class to load a CleanRL/Flax checkpoint and run inference. +class CleanRLPPOPolicy: + + def __init__( + self, + *, + network_params: Any, + actor_params: Any, + action_dim: int, + deterministic: bool = True, + seed: int = 0, + ) -> None: + from brittle_star_project.rl import Actor, Network + + self._network = Network() + self._actor = Actor(action_dim=action_dim) + self._network_apply = jax.jit(self._network.apply) + self._actor_apply = jax.jit(self._actor.apply) + self._params = { + "network_params": network_params, + "actor_params": actor_params, + } + self._deterministic = deterministic + self._rng = jax.random.PRNGKey(int(seed)) + + @staticmethod + def load( + path: Path, + *, + action_dim: int, + deterministic: bool, + seed: int, + ) -> "CleanRLPPOPolicy": + def _get_index(container: Any, idx: int) -> Any: + if isinstance(container, (list, tuple)): + return container[idx] + if isinstance(container, dict): + return container.get(idx, container.get(str(idx))) + raise KeyError(idx) + + def _looks_like_indexed_dict(container: Any) -> bool: + return ( + isinstance(container, dict) + and container + and all(str(k).isdigit() for k in container.keys()) + ) + + def _parse_checkpoint(restored_obj: Any) -> tuple[Any, Any, Any, Any]: + """Extract (args_dict, network_params, actor_params, critic_params). + + `src/train.py` saves: + flax.serialization.to_bytes([vars(args), [net, actor, critic]]) + + `msgpack_restore()` occasionally restores lists as dicts keyed by + string indices ("0", "1", ...), so we accept both shapes. + """ + + args_part: Any | None = None + params_part: Any = restored_obj + + if isinstance(restored_obj, (list, tuple)) and len(restored_obj) >= 2: + args_part = restored_obj[0] + params_part = restored_obj[1] + elif _looks_like_indexed_dict(restored_obj) and ( + "0" in restored_obj or "1" in restored_obj + ): + args_part = restored_obj.get("0", restored_obj.get(0)) + params_part = restored_obj.get("1", restored_obj.get(1)) + + if _looks_like_indexed_dict(params_part): + network_params = _get_index(params_part, 0) + actor_params = _get_index(params_part, 1) + critic_params = _get_index(params_part, 2) + if network_params is None or actor_params is None: + raise ValueError("Missing required params in checkpoint") + return args_part, network_params, actor_params, critic_params + + if isinstance(params_part, (list, tuple)) and len(params_part) >= 2: + network_params = params_part[0] + actor_params = params_part[1] + critic_params = params_part[2] if len(params_part) >= 3 else None + return args_part, network_params, actor_params, critic_params + + raise ValueError( + f"Unexpected .cleanrl_model structure in {path}. " + "Expected [args_dict, [network_params, actor_params, critic_params]] " + "or an equivalent dict-indexed variant." + ) + + payload = path.read_bytes() + restored = flax.serialization.msgpack_restore(payload) + _args_dict, network_params, actor_params, _critic_params = _parse_checkpoint(restored) + + return CleanRLPPOPolicy( + network_params=network_params, + actor_params=actor_params, + action_dim=action_dim, + deterministic=deterministic, + seed=seed, + ) + + def reset(self, seed: int) -> None: + self._rng = jax.random.PRNGKey(int(seed)) + + def act(self, *, observations: dict[str, Any]) -> np.ndarray: + obs = _flatten_obs_dict(observations) + hidden = self._network_apply(self._params["network_params"], obs) + mean, log_std = self._actor_apply(self._params["actor_params"], hidden) + + if self._deterministic: + action = mean + else: + self._rng, sub = jax.random.split(self._rng) + noise = jax.random.normal(sub, shape=mean.shape) + action = mean + noise * jnp.exp(log_std) + + return np.asarray(action, dtype=np.float32).ravel() + + +def _get_observations(state: Any) -> dict[str, Any]: + return getattr(state, "observations", None) + + +def _get_xy_distance_to_target(observations: dict[str, Any]) -> float | None: + return float(np.asarray(observations["xy_distance_to_target"]).reshape(-1)[0]) + + +def _target_reached(*, state: Any) -> bool: + return bool(getattr(state, "terminated", False)) + + +def _rollout_one_episode_headless( + *, + env: Any, + policy: CleanRLPPOPolicy, + seed: int, + max_steps: int, +) -> tuple[float, int, bool, float | None]: + """Run one rollout up to `max_steps`. + + Returns (return, length, reached_target, final_xy_dist). + """ + + policy.reset(seed) + state = env.reset(seed=seed) + + ep_return = 0.0 + + observations = _get_observations(state) + prev_dist = _get_xy_distance_to_target(observations) + reached_target = _target_reached(state=state) + + # NOTE: In the MJC backend, `state.reward` is always 0.0. + # To get a meaningful return, we compute a simple progress reward: + # r_t = d_{t-1} - d_t + # where d is `xy_distance_to_target`. + steps = 0 + for _ in range(int(max_steps)): + action = policy.act(observations=observations) + + nu = int(state.mj_model.nu) + if nu > 0 and action.shape != (nu,): + raise ValueError(f"Policy returned action shape {action.shape}, expected ({nu},)") + + state = env.step(state=state, action=action) + steps += 1 + + observations = _get_observations(state) + cur_dist = _get_xy_distance_to_target(observations) + if prev_dist is not None and cur_dist is not None: + ep_return += prev_dist - cur_dist + prev_dist = cur_dist + + reached_target = _target_reached(state=state) + if reached_target: + break + + final_dist = _get_xy_distance_to_target(observations) + return ep_return, steps, reached_target, final_dist + + +def _run_one_episode_viewer( + *, + env: Any, + policy: CleanRLPPOPolicy, + seed: int, + state: Any, + control_dt: float, + max_steps: int, +) -> None: + import mujoco.viewer + + model = state.mj_model + data = state.mj_data + + seed = int(seed) + episode_return = 0.0 + observations = _get_observations(state) + prev_dist = _get_xy_distance_to_target(observations) + reached_target = _target_reached(state=state) + + viewer = mujoco.viewer.launch_passive(model, data) + try: + steps = 0 + for _step_idx in range(int(max_steps)): + if not viewer.is_running(): + break + step_start = time.time() + + # One control step. We do the env step under the viewer lock. + action = policy.act(observations=observations) + if model.nu > 0 and action.shape != (int(model.nu),): + raise ValueError( + f"Policy returned action shape {action.shape}, expected ({int(model.nu)},)" + ) + # The passive viewer runs a GUI thread; protect MuJoCo state mutation. + with viewer.lock(): + state = env.step(state=state, action=action) + if not viewer.is_running(): + break + viewer.sync() + + steps += 1 + + observations = _get_observations(state) + cur_dist = _get_xy_distance_to_target(observations) + if prev_dist is not None and cur_dist is not None: + episode_return += prev_dist - cur_dist + prev_dist = cur_dist + + reached_target = _target_reached(state=state) + if reached_target: + break + + # Real-time pacing so the viewer doesn't run as fast as possible. + remaining = control_dt - (time.time() - step_start) + if remaining > 0: + time.sleep(remaining) + + # Done: target reached, fixed horizon reached, or window closed. + if viewer.is_running(): + dist = _get_xy_distance_to_target(observations) + dist_str = "n/a" if dist is None else f"{dist:.3f}" + print( + "episode done: " + f"return={episode_return:.6f}, len={steps}, " + f"target_reached={reached_target}, final_xy_dist={dist_str}" + ) + viewer.close() + finally: + # Ensure the GUI thread stops before the env/model/data are torn down. + try: + viewer.close() + except Exception: + pass + for _ in range(200): + if not viewer.is_running(): + break + time.sleep(0.01) def parse_args() -> argparse.Namespace: - p = argparse.ArgumentParser(description="Simulate a trained policy in the MuJoCo viewer.") + from brittle_star_project.environment import Task + + p = argparse.ArgumentParser( + description="Run a trained policy for exactly one episode (viewer or headless)." + ) p.add_argument( "--model", type=str, - default=None, - help="Path to a saved model artifact. If omitted, a model is created from --model-type.", + required=True, + help=("Path to a CleanRL/Flax '.cleanrl_model' checkpoint (saved by src/train.py)."), ) p.add_argument( - "--model-type", - choices=MODEL_OPTIONS, - default="random", - help="Which model class to instantiate when --model is omitted.", + "--deterministic", + action=argparse.BooleanOptionalAction, + default=True, + help="Use mean action (deterministic) or sample actions (stochastic).", + ) + p.add_argument( + "--headless", + action="store_true", + help="Run without the MuJoCo viewer (still exactly one episode).", + ) + p.add_argument( + "--max-steps", + type=int, + required=True, + help=( + "Number of control steps to run (fixed horizon). " + "This script stops when this many steps are reached, or earlier if " + "the target is reached (directed locomotion)." + ), ) p.add_argument( "--backend", choices=[b for b in Backend], - default=Backend.MJX, + default=Backend.MJC, ) - p.add_argument("--seed", type=int, default=None) + p.add_argument("--seed", type=int, default=0) return p.parse_args() def main() -> None: + from brittle_star_project.environment import ( + BrittleStarEnv, + BrittleStarEnvFactory, + ) + args = parse_args() morphology_cfg, arena_cfg, env_cfg = from_json("../configs/test.json") @@ -63,30 +362,57 @@ def main() -> None: # policy/model. nu = int(state.mj_model.nu) - if args.model is not None: - model_path = Path(args.model) - policy = RLModel.load(model_path) - if hasattr(policy, "nu"): - policy.nu = nu - else: - model_cls = MODEL_BY_NAME[str(args.model_type)] - policy = model_cls(seed=seed_for_env) - if hasattr(policy, "nu"): - policy.nu = nu + model_path = Path(args.model) + if model_path.suffix != ".cleanrl_model": + raise ValueError(f"Expected a '.cleanrl_model' checkpoint, got '{model_path.name}'.") - # If the policy/model has a `seed` attribute, use the provided seed (or default) to reset it. - default_seed = int(getattr(policy, "seed", seed_for_env)) - if args.seed is not None and hasattr(policy, "reset"): + policy = CleanRLPPOPolicy.load( + model_path, + action_dim=nu, + deterministic=bool(args.deterministic), + seed=seed_for_env, + ) + + # Reset policy RNG if a seed was provided. + default_seed = seed_for_env + if args.seed is not None: policy.reset(int(args.seed)) # ======= SIMULATION ======= - rollout_cfg = SimulationConfig( - realtime=True, - seed=int(args.seed) if args.seed is not None else default_seed, - ) + if args.headless: + max_steps = int(args.max_steps) + if max_steps <= 0: + raise ValueError("--max-steps must be > 0") - simulate_policy(policy, rollout_cfg, state) + ep_seed = int(args.seed) if args.seed is not None else default_seed + ep_return, ep_len, reached_target, final_dist = _rollout_one_episode_headless( + env=env, + policy=policy, + seed=ep_seed, + max_steps=max_steps, + ) + final_dist_str = "n/a" if final_dist is None else f"{final_dist:.3f}" + print( + "episode done: " + f"return={ep_return:.6f}, len={ep_len}, " + f"target_reached={reached_target}, final_xy_dist={final_dist_str}" + ) + else: + max_steps = int(args.max_steps) + if max_steps <= 0: + raise ValueError("--max-steps must be > 0") + + model_dt = float(state.mj_model.opt.timestep) + control_dt = model_dt * float(env_cfg.num_physics_steps_per_control_step) + _run_one_episode_viewer( + env=env, + policy=policy, + seed=int(args.seed) if args.seed is not None else default_seed, + state=state, + control_dt=control_dt, + max_steps=max_steps, + ) env.close() From 5329ae18ba809d1d25b0ec4f97cfb67f1334eb05 Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Thu, 2 Apr 2026 10:36:18 +0200 Subject: [PATCH 02/10] fix: removed deterministic flag --- experiments/simulate.py | 39 ++++----------------------------------- 1 file changed, 4 insertions(+), 35 deletions(-) diff --git a/experiments/simulate.py b/experiments/simulate.py index 3b79082..3dca83c 100644 --- a/experiments/simulate.py +++ b/experiments/simulate.py @@ -41,8 +41,6 @@ class CleanRLPPOPolicy: network_params: Any, actor_params: Any, action_dim: int, - deterministic: bool = True, - seed: int = 0, ) -> None: from brittle_star_project.rl import Actor, Network @@ -54,16 +52,12 @@ class CleanRLPPOPolicy: "network_params": network_params, "actor_params": actor_params, } - self._deterministic = deterministic - self._rng = jax.random.PRNGKey(int(seed)) @staticmethod def load( path: Path, *, action_dim: int, - deterministic: bool, - seed: int, ) -> "CleanRLPPOPolicy": def _get_index(container: Any, idx: int) -> Any: if isinstance(container, (list, tuple)): @@ -129,26 +123,16 @@ class CleanRLPPOPolicy: network_params=network_params, actor_params=actor_params, action_dim=action_dim, - deterministic=deterministic, - seed=seed, ) - def reset(self, seed: int) -> None: - self._rng = jax.random.PRNGKey(int(seed)) - def act(self, *, observations: dict[str, Any]) -> np.ndarray: obs = _flatten_obs_dict(observations) hidden = self._network_apply(self._params["network_params"], obs) - mean, log_std = self._actor_apply(self._params["actor_params"], hidden) + mean, _log_std = self._actor_apply(self._params["actor_params"], hidden) - if self._deterministic: - action = mean - else: - self._rng, sub = jax.random.split(self._rng) - noise = jax.random.normal(sub, shape=mean.shape) - action = mean + noise * jnp.exp(log_std) - - return np.asarray(action, dtype=np.float32).ravel() + # Always evaluate with the actor mean. + # (Sampling adds exploration noise, which is useful for training but not for evaluation.) + return np.asarray(mean, dtype=np.float32).ravel() def _get_observations(state: Any) -> dict[str, Any]: @@ -174,8 +158,6 @@ def _rollout_one_episode_headless( Returns (return, length, reached_target, final_xy_dist). """ - - policy.reset(seed) state = env.reset(seed=seed) ep_return = 0.0 @@ -294,8 +276,6 @@ def _run_one_episode_viewer( def parse_args() -> argparse.Namespace: - from brittle_star_project.environment import Task - p = argparse.ArgumentParser( description="Run a trained policy for exactly one episode (viewer or headless)." ) @@ -305,12 +285,6 @@ def parse_args() -> argparse.Namespace: required=True, help=("Path to a CleanRL/Flax '.cleanrl_model' checkpoint (saved by src/train.py)."), ) - p.add_argument( - "--deterministic", - action=argparse.BooleanOptionalAction, - default=True, - help="Use mean action (deterministic) or sample actions (stochastic).", - ) p.add_argument( "--headless", action="store_true", @@ -369,14 +343,9 @@ def main() -> None: policy = CleanRLPPOPolicy.load( model_path, action_dim=nu, - deterministic=bool(args.deterministic), - seed=seed_for_env, ) - # Reset policy RNG if a seed was provided. default_seed = seed_for_env - if args.seed is not None: - policy.reset(int(args.seed)) # ======= SIMULATION ======= From 76ba835ee64b32e44f239fe15eef161af3931181 Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Wed, 15 Apr 2026 14:37:08 +0200 Subject: [PATCH 03/10] feat: added cli config path to simulate script --- scripts/simulate.py | 30 +++++++++++++++++-- .../environment/env_wrapper.py | 16 ++++++++-- 2 files changed, 41 insertions(+), 5 deletions(-) diff --git a/scripts/simulate.py b/scripts/simulate.py index 90dc65f..b38d9a4 100644 --- a/scripts/simulate.py +++ b/scripts/simulate.py @@ -12,7 +12,7 @@ import numpy as np from brittle_star_project import ( Backend, ) -from brittle_star_project.environment import from_file +from brittle_star_project.environment import ArenaConfig, EnvConfig, MorphologyConfig, from_file def _flatten_obs_dict(obs_dict: dict[str, Any]) -> jnp.ndarray: """Flatten the env's observation dict into a 1D vector. @@ -279,6 +279,16 @@ def parse_args() -> argparse.Namespace: p = argparse.ArgumentParser( description="Run a trained policy for exactly one episode (viewer or headless)." ) + p.add_argument( + "--config-path", + type=str, + default=None, + help=( + "Path to an environment JSON config (morphology/arena/env). " + "If omitted, uses the environment defaults. " + "Relative paths are resolved from the repository root." + ), + ) p.add_argument( "--model", type=str, @@ -317,7 +327,16 @@ def main() -> None: args = parse_args() - morphology_cfg, arena_cfg, env_cfg = from_file("../configs/test.yaml") + if args.config_path is None: + morphology_cfg = MorphologyConfig() + arena_cfg = ArenaConfig() + env_cfg = EnvConfig() + else: + repo_root = Path(__file__).resolve().parents[1] + config_path = Path(args.config_path) + if not config_path.is_absolute(): + config_path = repo_root / config_path + morphology_cfg, arena_cfg, env_cfg = from_file(str(config_path)) # ======= ENVIRONMENT SETUP ======= @@ -325,7 +344,12 @@ def main() -> None: factory = BrittleStarEnvFactory() raw_env = factory.create_environment(backend, morphology_cfg, arena_cfg, env_cfg) - env = BrittleStarEnv(raw_env, backend=backend, config=env_cfg) + env = BrittleStarEnv( + raw_env, + backend=backend, + config=env_cfg, + morphology_config=morphology_cfg, + ) seed_for_env = int(args.seed) if args.seed is not None else 0 state = env.reset(seed=seed_for_env) diff --git a/src/brittle_star_project/environment/env_wrapper.py b/src/brittle_star_project/environment/env_wrapper.py index f74a1ca..8b3e51f 100644 --- a/src/brittle_star_project/environment/env_wrapper.py +++ b/src/brittle_star_project/environment/env_wrapper.py @@ -6,7 +6,7 @@ from typing import Any import numpy as np -from .env_config import EnvConfig +from .env_config import EnvConfig, MorphologyConfig from .env_types import Backend @@ -25,10 +25,18 @@ class BrittleStarEnv: 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: + def __init__( + self, + env: Any, + *, + backend: Backend, + config: EnvConfig, + morphology_config: MorphologyConfig | None = None, + ) -> None: self._env = env self._backend = backend self._config = config + self._morphology_config = morphology_config @property def raw(self) -> Any: @@ -42,6 +50,10 @@ class BrittleStarEnv: def config(self) -> EnvConfig: return self._config + @property + def morphology_config(self) -> MorphologyConfig | None: + return self._morphology_config + def make_rng(self, seed: int): if self._backend == Backend.MJC: return np.random.RandomState(seed) From ea9ed21c9e60193cab4ca7700a9727b98125ce59 Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Wed, 15 Apr 2026 17:31:41 +0200 Subject: [PATCH 04/10] fix: further adapted simulate script to the new pipeline --- scripts/simulate.py | 149 ++++++++++++++++++++++++++------------------ 1 file changed, 87 insertions(+), 62 deletions(-) diff --git a/scripts/simulate.py b/scripts/simulate.py index b38d9a4..98fd18e 100644 --- a/scripts/simulate.py +++ b/scripts/simulate.py @@ -1,6 +1,7 @@ from __future__ import annotations import argparse +import itertools import time from pathlib import Path from typing import Any @@ -38,18 +39,18 @@ class CleanRLPPOPolicy: def __init__( self, *, - network_params: Any, + sensor_params: Any, actor_params: Any, action_dim: int, ) -> None: - from brittle_star_project.rl import Actor, Network + from brittle_star_project.MLPs.mlps import Actor, GenericDenseLayersWithActivation - self._network = Network() + self._sensor = GenericDenseLayersWithActivation() self._actor = Actor(action_dim=action_dim) - self._network_apply = jax.jit(self._network.apply) + self._sensor_apply = jax.jit(self._sensor.apply) self._actor_apply = jax.jit(self._actor.apply) self._params = { - "network_params": network_params, + "sensor_params": sensor_params, "actor_params": actor_params, } @@ -73,14 +74,19 @@ class CleanRLPPOPolicy: and all(str(k).isdigit() for k in container.keys()) ) - def _parse_checkpoint(restored_obj: Any) -> tuple[Any, Any, Any, Any]: - """Extract (args_dict, network_params, actor_params, critic_params). + def _parse_checkpoint(restored_obj: Any) -> tuple[Any, Any, Any, Any, Any]: + """Extract checkpoint parts. - `src/train.py` saves: - flax.serialization.to_bytes([vars(args), [net, actor, critic]]) + Returns (args_dict, sensor_params, actor_params, critic_params, + feature_extractor_params). - `msgpack_restore()` occasionally restores lists as dicts keyed by - string indices ("0", "1", ...), so we accept both shapes. + `PPOTrainer` saves: + flax.serialization.to_bytes( + [vars(args), [sensor, actor, critic, feature_extractor]] + ) + + `msgpack_restore()` may restore lists as dicts keyed by string + indices ("0", "1", ...), so we accept both shapes. """ args_part: Any | None = None @@ -96,38 +102,55 @@ class CleanRLPPOPolicy: params_part = restored_obj.get("1", restored_obj.get(1)) if _looks_like_indexed_dict(params_part): - network_params = _get_index(params_part, 0) + sensor_params = _get_index(params_part, 0) actor_params = _get_index(params_part, 1) critic_params = _get_index(params_part, 2) - if network_params is None or actor_params is None: + feature_extractor_params = _get_index(params_part, 3) + if sensor_params is None or actor_params is None: raise ValueError("Missing required params in checkpoint") - return args_part, network_params, actor_params, critic_params + return ( + args_part, + sensor_params, + actor_params, + critic_params, + feature_extractor_params, + ) if isinstance(params_part, (list, tuple)) and len(params_part) >= 2: - network_params = params_part[0] + sensor_params = params_part[0] actor_params = params_part[1] critic_params = params_part[2] if len(params_part) >= 3 else None - return args_part, network_params, actor_params, critic_params + feature_extractor_params = params_part[3] if len(params_part) >= 4 else None + return ( + args_part, + sensor_params, + actor_params, + critic_params, + feature_extractor_params, + ) raise ValueError( - f"Unexpected .cleanrl_model structure in {path}. " - "Expected [args_dict, [network_params, actor_params, critic_params]] " + f"Unexpected checkpoint structure in {path}. " + "Expected [args_dict, [sensor_params, actor_params, critic_params, " + "feature_extractor_params]] " "or an equivalent dict-indexed variant." ) payload = path.read_bytes() restored = flax.serialization.msgpack_restore(payload) - _args_dict, network_params, actor_params, _critic_params = _parse_checkpoint(restored) + _args_dict, sensor_params, actor_params, _critic_params, _feature_extractor_params = ( + _parse_checkpoint(restored) + ) return CleanRLPPOPolicy( - network_params=network_params, + sensor_params=sensor_params, actor_params=actor_params, action_dim=action_dim, ) def act(self, *, observations: dict[str, Any]) -> np.ndarray: obs = _flatten_obs_dict(observations) - hidden = self._network_apply(self._params["network_params"], obs) + hidden = self._sensor_apply(self._params["sensor_params"], obs) mean, _log_std = self._actor_apply(self._params["actor_params"], hidden) # Always evaluate with the actor mean. @@ -135,7 +158,7 @@ class CleanRLPPOPolicy: return np.asarray(mean, dtype=np.float32).ravel() -def _get_observations(state: Any) -> dict[str, Any]: +def _get_observations(state: Any) -> dict[str, Any] | None: return getattr(state, "observations", None) @@ -202,7 +225,7 @@ def _run_one_episode_viewer( seed: int, state: Any, control_dt: float, - max_steps: int, + max_steps: int | None, ) -> None: import mujoco.viewer @@ -215,23 +238,28 @@ def _run_one_episode_viewer( prev_dist = _get_xy_distance_to_target(observations) reached_target = _target_reached(state=state) - viewer = mujoco.viewer.launch_passive(model, data) - try: - steps = 0 - for _step_idx in range(int(max_steps)): + steps = 0 + # Use the viewer as a context manager to avoid GLX teardown races + # (e.g. GLXBadDrawable from X_GLXSwapBuffers after a window is destroyed). + with mujoco.viewer.launch_passive(model, data) as viewer: + step_iter = ( + range(int(max_steps)) if max_steps is not None else itertools.count() + ) + for _step_idx in step_iter: if not viewer.is_running(): break step_start = time.time() - # One control step. We do the env step under the viewer lock. action = policy.act(observations=observations) if model.nu > 0 and action.shape != (int(model.nu),): raise ValueError( f"Policy returned action shape {action.shape}, expected ({int(model.nu)},)" ) + # The passive viewer runs a GUI thread; protect MuJoCo state mutation. with viewer.lock(): state = env.step(state=state, action=action) + if not viewer.is_running(): break viewer.sync() @@ -248,31 +276,17 @@ def _run_one_episode_viewer( if reached_target: break - # Real-time pacing so the viewer doesn't run as fast as possible. remaining = control_dt - (time.time() - step_start) if remaining > 0: time.sleep(remaining) - # Done: target reached, fixed horizon reached, or window closed. - if viewer.is_running(): - dist = _get_xy_distance_to_target(observations) - dist_str = "n/a" if dist is None else f"{dist:.3f}" - print( - "episode done: " - f"return={episode_return:.6f}, len={steps}, " - f"target_reached={reached_target}, final_xy_dist={dist_str}" - ) - viewer.close() - finally: - # Ensure the GUI thread stops before the env/model/data are torn down. - try: - viewer.close() - except Exception: - pass - for _ in range(200): - if not viewer.is_running(): - break - time.sleep(0.01) + dist = _get_xy_distance_to_target(observations) + dist_str = "n/a" if dist is None else f"{dist:.3f}" + print( + "episode done: " + f"return={episode_return:.6f}, len={steps}, " + f"target_reached={reached_target}, final_xy_dist={dist_str}" + ) def parse_args() -> argparse.Namespace: @@ -293,7 +307,9 @@ def parse_args() -> argparse.Namespace: "--model", type=str, required=True, - help=("Path to a CleanRL/Flax '.cleanrl_model' checkpoint (saved by src/train.py)."), + help=( + "Path to the Flax checkpoint saved by scripts/train.py (final_model.flax)." + ), ) p.add_argument( "--headless", @@ -303,17 +319,17 @@ def parse_args() -> argparse.Namespace: p.add_argument( "--max-steps", type=int, - required=True, + default=None, help=( - "Number of control steps to run (fixed horizon). " - "This script stops when this many steps are reached, or earlier if " - "the target is reached (directed locomotion)." + "Number of control steps to run. " + "In --headless mode this is required and acts as a fixed horizon. " + "In viewer mode the default is infinite (run until window closed or target reached)." ), ) p.add_argument( "--backend", - choices=[b for b in Backend], - default=Backend.MJC, + choices=[b.value for b in Backend], + default=Backend.MJC.value, ) p.add_argument("--seed", type=int, default=0) return p.parse_args() @@ -340,7 +356,7 @@ def main() -> None: # ======= ENVIRONMENT SETUP ======= - backend = args.backend + backend = Backend(args.backend) factory = BrittleStarEnvFactory() raw_env = factory.create_environment(backend, morphology_cfg, arena_cfg, env_cfg) @@ -361,8 +377,11 @@ def main() -> None: nu = int(state.mj_model.nu) model_path = Path(args.model) - if model_path.suffix != ".cleanrl_model": - raise ValueError(f"Expected a '.cleanrl_model' checkpoint, got '{model_path.name}'.") + if model_path.name != "final_model.flax" or model_path.suffix != ".flax": + raise ValueError( + "Expected the training artifact 'final_model.flax', " + f"got '{model_path.name}'." + ) policy = CleanRLPPOPolicy.load( model_path, @@ -374,6 +393,8 @@ def main() -> None: # ======= SIMULATION ======= if args.headless: + if args.max_steps is None: + raise ValueError("--max-steps is required in --headless mode") max_steps = int(args.max_steps) if max_steps <= 0: raise ValueError("--max-steps must be > 0") @@ -392,9 +413,13 @@ def main() -> None: f"target_reached={reached_target}, final_xy_dist={final_dist_str}" ) else: - max_steps = int(args.max_steps) - if max_steps <= 0: - raise ValueError("--max-steps must be > 0") + max_steps: int | None + if args.max_steps is None: + max_steps = None + else: + max_steps = int(args.max_steps) + if max_steps <= 0: + raise ValueError("--max-steps must be > 0") model_dt = float(state.mj_model.opt.timestep) control_dt = model_dt * float(env_cfg.num_physics_steps_per_control_step) From f6022cc9130982fe4e386fbc3f03500b485c5b03 Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Wed, 15 Apr 2026 18:57:14 +0200 Subject: [PATCH 05/10] fix: formatted simulate script --- scripts/simulate.py | 13 ++++--------- 1 file changed, 4 insertions(+), 9 deletions(-) diff --git a/scripts/simulate.py b/scripts/simulate.py index 98fd18e..4afdc24 100644 --- a/scripts/simulate.py +++ b/scripts/simulate.py @@ -15,6 +15,7 @@ from brittle_star_project import ( ) from brittle_star_project.environment import ArenaConfig, EnvConfig, MorphologyConfig, from_file + def _flatten_obs_dict(obs_dict: dict[str, Any]) -> jnp.ndarray: """Flatten the env's observation dict into a 1D vector. @@ -35,7 +36,6 @@ def _flatten_obs_dict(obs_dict: dict[str, Any]) -> jnp.ndarray: # A minimal policy class to load a CleanRL/Flax checkpoint and run inference. class CleanRLPPOPolicy: - def __init__( self, *, @@ -242,9 +242,7 @@ def _run_one_episode_viewer( # Use the viewer as a context manager to avoid GLX teardown races # (e.g. GLXBadDrawable from X_GLXSwapBuffers after a window is destroyed). with mujoco.viewer.launch_passive(model, data) as viewer: - step_iter = ( - range(int(max_steps)) if max_steps is not None else itertools.count() - ) + step_iter = range(int(max_steps)) if max_steps is not None else itertools.count() for _step_idx in step_iter: if not viewer.is_running(): break @@ -307,9 +305,7 @@ def parse_args() -> argparse.Namespace: "--model", type=str, required=True, - help=( - "Path to the Flax checkpoint saved by scripts/train.py (final_model.flax)." - ), + help=("Path to the Flax checkpoint saved by scripts/train.py (final_model.flax)."), ) p.add_argument( "--headless", @@ -379,8 +375,7 @@ def main() -> None: model_path = Path(args.model) if model_path.name != "final_model.flax" or model_path.suffix != ".flax": raise ValueError( - "Expected the training artifact 'final_model.flax', " - f"got '{model_path.name}'." + f"Expected the training artifact 'final_model.flax', got '{model_path.name}'." ) policy = CleanRLPPOPolicy.load( From 248683eb0c84c65568e0d5bf266377cce955c3c8 Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Wed, 15 Apr 2026 19:02:46 +0200 Subject: [PATCH 06/10] fix: comment mismatch Co-authored-by: Copilot Autofix powered by AI <175728472+Copilot@users.noreply.github.com> --- scripts/simulate.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/scripts/simulate.py b/scripts/simulate.py index 4afdc24..1e1253b 100644 --- a/scripts/simulate.py +++ b/scripts/simulate.py @@ -296,7 +296,7 @@ def parse_args() -> argparse.Namespace: type=str, default=None, help=( - "Path to an environment JSON config (morphology/arena/env). " + "Path to an environment YAML config (morphology/arena/env). " "If omitted, uses the environment defaults. " "Relative paths are resolved from the repository root." ), From cb6e3d97142fa3a8dea816835f96385774ae4874 Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Thu, 16 Apr 2026 10:48:37 +0200 Subject: [PATCH 07/10] fix: universal .flax files --- scripts/simulate.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/scripts/simulate.py b/scripts/simulate.py index 4afdc24..38e4ae5 100644 --- a/scripts/simulate.py +++ b/scripts/simulate.py @@ -373,9 +373,9 @@ def main() -> None: nu = int(state.mj_model.nu) model_path = Path(args.model) - if model_path.name != "final_model.flax" or model_path.suffix != ".flax": + if model_path.suffix != ".flax": raise ValueError( - f"Expected the training artifact 'final_model.flax', got '{model_path.name}'." + f"Expected the training artifact '.flax', got '{model_path.name}'." ) policy = CleanRLPPOPolicy.load( From 3fb8ffc9d0820c3a93e7fbc2d018c61995ee3a2f Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Thu, 16 Apr 2026 15:18:08 +0200 Subject: [PATCH 08/10] fix: adapted simulation config --- configs/simulation/default.yaml | 7 ++++++- src/brittle_star_project/configs/config_simulation.py | 8 +++++++- 2 files changed, 13 insertions(+), 2 deletions(-) diff --git a/configs/simulation/default.yaml b/configs/simulation/default.yaml index 00c4e44..8645f00 100644 --- a/configs/simulation/default.yaml +++ b/configs/simulation/default.yaml @@ -8,4 +8,9 @@ model_path: null model_type: "random" # Execution backend (MJX or BRAX) -backend: "MJX" +backend: "MJC" + +# Script behavior +headless: false +# In headless mode this is required; in viewer mode null means "infinite". +max_steps: null diff --git a/src/brittle_star_project/configs/config_simulation.py b/src/brittle_star_project/configs/config_simulation.py index 56747ae..00488fc 100644 --- a/src/brittle_star_project/configs/config_simulation.py +++ b/src/brittle_star_project/configs/config_simulation.py @@ -1,5 +1,6 @@ from dataclasses import dataclass from typing import Optional + from brittle_star_project.environment.env_types import Backend @@ -9,4 +10,9 @@ class SimulationSettings: model_path: Optional[str] = None model_type: str = "random" - backend: Backend = Backend.MJX + backend: Backend = Backend.MJC + + # Script behavior + headless: bool = False + # If None, viewer mode runs until window closed or target reached. + max_steps: Optional[int] = None From 395b04d9a85d27a561927cea5695470e64948269 Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Wed, 22 Apr 2026 18:47:35 +0200 Subject: [PATCH 09/10] feat: adapted simulate to trained config --- configs/simulation/default.yaml | 11 +- scripts/simulate.py | 285 +++++++++++++++++- .../configs/config_simulation.py | 9 +- src/experiment_logger/unified_logger.py | 59 +++- 4 files changed, 334 insertions(+), 30 deletions(-) diff --git a/configs/simulation/default.yaml b/configs/simulation/default.yaml index 8645f00..84694e0 100644 --- a/configs/simulation/default.yaml +++ b/configs/simulation/default.yaml @@ -4,13 +4,12 @@ # Path to the trained model (optional) model_path: null -# Type of model to use if no path is provided (e.g., random) -model_type: "random" - -# Execution backend (MJX or BRAX) -backend: "MJC" - # Script behavior headless: false # In headless mode this is required; in viewer mode null means "infinite". max_steps: null + +# Optional: path to the Hydra config.yaml used during training. +# When set, scripts/simulate.py will use it to default morphology/arena/environment/architecture +# to match training (unless you explicitly override those keys via CLI). +trained_config_path: null diff --git a/scripts/simulate.py b/scripts/simulate.py index 28660b9..a9a51f4 100644 --- a/scripts/simulate.py +++ b/scripts/simulate.py @@ -3,6 +3,10 @@ Uses Hydra to load the same BrittleStarConfig that was used during training. Override settings via CLI, e.g.: python scripts/simulate.py morphology=3_arms + +To replay a run using the *exact* Hydra config used during training, pass: + python scripts/simulate.py simulation.trained_config_path=runs/.../.hydra/config.yaml \ + simulation.model_path=runs/.../final_model.flax """ from __future__ import annotations @@ -17,11 +21,16 @@ import hydra import jax import jax.numpy as jnp import numpy as np -from omegaconf import DictConfig, OmegaConf +import yaml +from omegaconf import DictConfig, OmegaConf, open_dict -from brittle_star_project import BrittleStarEnv, BrittleStarEnvFactory +from brittle_star_project import Backend, BrittleStarEnv, BrittleStarEnvFactory from brittle_star_project.configs.main_config import BrittleStarConfig from brittle_star_project.configs.register_configs import register_configs +from brittle_star_project.environment.padded_obs_wrapper import ( + compute_padding_masks, + pad_observation, +) _ALLOWED_OBS_KEYS = { "joint_position", @@ -36,6 +45,74 @@ _ALLOWED_OBS_KEYS = { "xy_distance_to_target", } + +def _dense_layer_sizes_from_params(params: Any) -> list[int]: + """Infer GenericDenseLayersWithActivation.layer_sizes from a Flax params tree.""" + + try: + dense_params = params["params"] + except Exception as exc: + raise ValueError("Unexpected sensor params structure (missing 'params')") from exc + + layer_sizes: list[int] = [] + idx = 0 + while True: + key = f"Dense_{idx}" + if key not in dense_params: + break + kernel = dense_params[key]["kernel"] + layer_sizes.append(int(np.asarray(kernel).shape[1])) + idx += 1 + + if not layer_sizes: + raise ValueError("Could not infer Dense_* layers from sensor params") + return layer_sizes + + +def _infer_action_dim_from_actor_params(params: Any) -> int | None: + """Best-effort infer action_dim from a Flax Actor params tree.""" + + try: + dense0 = params["params"]["Dense_0"] + bias = dense0.get("bias") + kernel = dense0.get("kernel") + except Exception: + return None + + if bias is not None: + try: + return int(np.asarray(bias).shape[0]) + except Exception: + return None + + if kernel is not None: + try: + return int(np.asarray(kernel).shape[1]) + except Exception: + return None + + return None + + +def _has_cli_override(overrides: list[str], key: str) -> bool: + prefixes = (f"{key}=", f"{key}.", f"+{key}=", f"+{key}.") + return any(str(o).startswith(prefixes) for o in overrides) + + +def _maybe_clip_action( + action: np.ndarray, + low: np.ndarray | None, + high: np.ndarray | None, +) -> np.ndarray: + if low is None or high is None: + return action + low = np.asarray(low, dtype=np.float32).ravel() + high = np.asarray(high, dtype=np.float32).ravel() + if low.shape != action.shape or high.shape != action.shape: + return action + return np.clip(action, low, high) + + def _transform_obs_dict(obs_dict: dict[str, Any]) -> jnp.ndarray: """Flatten the env's observation dict into a 1D vector. @@ -69,9 +146,8 @@ class CleanRLPPOPolicy: ) -> None: from brittle_star_project.MLPs.mlps import Actor, GenericDenseLayersWithActivation - hidden_dim = int(sensor_params["params"]["Dense_0"]["kernel"].shape[1]) - - self._sensor = GenericDenseLayersWithActivation(layer_sizes=[hidden_dim, hidden_dim]) + layer_sizes = _dense_layer_sizes_from_params(sensor_params) + self._sensor = GenericDenseLayersWithActivation(layer_sizes=layer_sizes) self._actor = Actor(action_dim=action_dim) self._sensor_apply = jax.jit(self._sensor.apply) self._actor_apply = jax.jit(self._actor.apply) @@ -156,6 +232,28 @@ class CleanRLPPOPolicy: feature_extractor_params, ) + # Accept a plain dict-shaped Flax params mapping commonly produced + # by saving `agent_state.params` directly. Typical keys are + # 'sensor_params' and 'actor_params', or sometimes nested under 'params'. + if isinstance(restored_obj, dict): + # Top-level params dict + params_sub = restored_obj.get("params", {}) + sensor_params = restored_obj.get("sensor_params") or params_sub.get("sensor_params") + actor_params = restored_obj.get("actor_params") or params_sub.get("actor_params") + critic_params = restored_obj.get("critic_params") or params_sub.get("critic_params") + feature_extractor_params = restored_obj.get( + "feature_extractor_params" + ) or params_sub.get("feature_extractor_params") + # Some checkpoints only save actor+sensor as top-level + if sensor_params is not None and actor_params is not None: + return ( + cfg_part, + sensor_params, + actor_params, + critic_params, + feature_extractor_params, + ) + raise ValueError( f"Unexpected checkpoint structure in {path}. " "Expected [config_dict, [sensor_params, actor_params, critic_params, " @@ -168,6 +266,16 @@ class CleanRLPPOPolicy: _parse_checkpoint(restored) ) + ckpt_action_dim = _infer_action_dim_from_actor_params(actor_params) + if ckpt_action_dim is not None and ckpt_action_dim != action_dim: + raise ValueError( + "Checkpoint/env mismatch: " + f"checkpoint expects action_dim={ckpt_action_dim}, " + f"env provides action_dim={action_dim}. " + "Use the same Hydra config (morphology/arena/environment) " + "that was used during training." + ) + return CleanRLPPOPolicy( sensor_params=sensor_params, actor_params=actor_params, @@ -193,7 +301,7 @@ def _get_xy_distance_to_target(observations: dict[str, Any]) -> float | None: def _target_reached(*, state: Any) -> bool: - return bool(getattr(state, "terminated", False)) + return bool(getattr(state, "terminated", False) or getattr(state, "truncated", False)) def _rollout_one_episode_headless( @@ -202,6 +310,9 @@ def _rollout_one_episode_headless( policy: CleanRLPPOPolicy, seed: int, max_steps: int, + action_low: np.ndarray | None, + action_high: np.ndarray | None, + padding_masks: dict[str, Any] | None, ) -> tuple[float, int, bool, float | None]: """Run one rollout up to max_steps. @@ -220,7 +331,12 @@ def _rollout_one_episode_headless( steps = 0 for _ in range(int(max_steps)): - action = policy.act(observations=observations) + obs_dict = observations or {} + if padding_masks is not None: + obs_dict = pad_observation(obs_dict, padding_masks) + + action = policy.act(observations=obs_dict) + action = _maybe_clip_action(action, action_low, action_high) nu = int(state.mj_model.nu) if nu > 0 and action.shape != (nu,): @@ -251,6 +367,9 @@ def _run_one_episode_viewer( state: Any, control_dt: float, max_steps: int | None, + action_low: np.ndarray | None, + action_high: np.ndarray | None, + padding_masks: dict[str, Any] | None, ) -> None: import mujoco.viewer @@ -272,7 +391,12 @@ def _run_one_episode_viewer( break step_start = time.time() - action = policy.act(observations=observations or {}) + obs_dict = observations or {} + if padding_masks is not None: + obs_dict = pad_observation(obs_dict, padding_masks) + + action = policy.act(observations=obs_dict) + action = _maybe_clip_action(action, action_low, action_high) if model.nu > 0 and action.shape != (int(model.nu),): raise ValueError( f"Policy returned action shape {action.shape}, expected ({int(model.nu)},)" @@ -321,16 +445,126 @@ def _infer_checkpoint_obs_dim(policy: CleanRLPPOPolicy) -> int | None: return None +def _load_trained_config(path: Path) -> DictConfig: + """Load a trained config YAML. + + Supports both: + - Hydra's run config (e.g. runs/.../.hydra/config.yaml) + - This project's logger metadata YAMLs, which may contain + ``!!python/object/apply:...`` tags for Enums. + + For safety, we *do not* execute Python constructors from YAML; we only + treat these tags as data and extract their scalar arguments. + """ + + if not path.exists(): + raise FileNotFoundError(f"trained_config_path does not exist: '{path}'.") + if not path.is_file(): + raise ValueError(f"trained_config_path must be a file, got: '{path}'.") + + try: + return OmegaConf.load(path) + except Exception as exc: + python_apply_prefix = "tag:yaml.org,2002:python/object/apply:" + + class _SafeLoaderWithPythonApply(yaml.SafeLoader): + pass + + def _construct_python_apply( + loader: yaml.SafeLoader, + _tag_suffix: str, + node: yaml.Node, + ) -> Any: + if isinstance(node, yaml.SequenceNode): + seq = loader.construct_sequence(node) + if len(seq) == 1: + return seq[0] + return seq + if isinstance(node, yaml.MappingNode): + return loader.construct_mapping(node) + return loader.construct_scalar(node) + + _SafeLoaderWithPythonApply.add_multi_constructor( + python_apply_prefix, _construct_python_apply + ) + + try: + data = yaml.load(path.read_text(encoding="utf-8"), Loader=_SafeLoaderWithPythonApply) + except Exception as yaml_exc: + raise ValueError( + "Failed to load trained_config_path as YAML. " + "If this is a Hydra run, pass the run's '.hydra/config.yaml' file. " + f"Got: '{path}'." + ) from yaml_exc + + if not isinstance(data, dict): + raise ValueError( + "trained_config_path must contain a YAML mapping (dict-like) at the root. " + f"Got type={type(data).__name__} from '{path}'." + ) from exc + + # Normalize known enum-like strings to their Enum *names* so OmegaConf's + # structured config merge behaves like the normal Hydra config. + from brittle_star_project.environment.env_types import Task + + env_cfg = data.get("environment") + if isinstance(env_cfg, dict) and isinstance(env_cfg.get("task"), str): + task_str = str(env_cfg["task"]) + try: + env_cfg["task"] = Task[task_str].name + except Exception: + try: + env_cfg["task"] = Task(task_str).name + except Exception: + pass + + return OmegaConf.create(data) + + @hydra.main(config_path="../configs", config_name="main_config", version_base="1.3") def main(dict_cfg: DictConfig) -> None: - # Convert DictConfig to structured dataclass, ensuring the root schema is applied. - config: BrittleStarConfig = OmegaConf.to_object( - OmegaConf.merge(OmegaConf.structured(BrittleStarConfig), dict_cfg) - ) + # Compose against the structured schema first, so missing keys are validated. + cfg = OmegaConf.merge(OmegaConf.structured(BrittleStarConfig), dict_cfg) - backend = config.simulation.backend + # Optional: override env-defining sections (morphology/arena/environment/architecture) + # using the exact Hydra config that was used for training. + trained_cfg_path = cfg.simulation.trained_config_path + if trained_cfg_path: + overrides_raw = OmegaConf.select(cfg, "hydra.overrides.task") or [] + overrides = [str(o) for o in overrides_raw] + + trained_cfg_path_abs = Path(hydra.utils.to_absolute_path(trained_cfg_path)) + trained_cfg = _load_trained_config(trained_cfg_path_abs) + if "hydra" in trained_cfg: + with open_dict(trained_cfg): + del trained_cfg["hydra"] + + with open_dict(cfg): + for key in ("morphology", "arena", "environment", "architecture"): + if key in trained_cfg and not _has_cli_override(overrides, key): + base_node = OmegaConf.select(cfg, key) + override_node = OmegaConf.select(trained_cfg, key) + try: + cfg[key] = OmegaConf.merge(base_node, override_node) + except Exception as exc: + raise ValueError( + "Failed to merge trained config into the active Hydra config. " + f"Key={key!r}, trained_config_path='{trained_cfg_path_abs}'." + ) from exc + + # Convert DictConfig to structured dataclass. + config: BrittleStarConfig = OmegaConf.to_object(cfg) + + backend = Backend.MJC seed = int(config.experiment.seed) + if getattr(config.architecture, "name", None) != "centralized": + raise ValueError( + "simulate.py currently only supports architecture=centralized. " + f"Got architecture.name={getattr(config.architecture, 'name', None)!r}. " + "(Training supports decentralized, but simulation wiring for it isn't implemented.)" + ) + model_path_str = config.simulation.model_path if model_path_str is None: raise ValueError( @@ -359,15 +593,28 @@ def main(dict_cfg: DictConfig) -> None: state0 = env.reset(seed=seed) + # Match training's padded observation layout for amputated morphologies. + padding_masks = compute_padding_masks(config.morphology.segments_per_arm) + + # Match training's action clipping behavior. + action_space = getattr(raw_env, "action_space", None) + action_low = ( + None if action_space is None else np.asarray(action_space.low, dtype=np.float32).ravel() + ) + action_high = ( + None if action_space is None else np.asarray(action_space.high, dtype=np.float32).ravel() + ) + # ======= MODEL SETUP ======= nu = int(state0.mj_model.nu) policy = CleanRLPPOPolicy.load(model_path, action_dim=nu) # Helpful early failure when configs don't match the checkpoint. observations0 = _get_observations(state0) - env_obs_dim = int(_transform_obs_dict(observations0 or {}).shape[0]) + obs0_dict = pad_observation(observations0 or {}, padding_masks) + env_obs_dim = int(_transform_obs_dict(obs0_dict).shape[0]) ckpt_obs_dim = _infer_checkpoint_obs_dim(policy) - + if ckpt_obs_dim is not None and ckpt_obs_dim != env_obs_dim: raise ValueError( "Checkpoint/env mismatch: " @@ -375,7 +622,7 @@ def main(dict_cfg: DictConfig) -> None: "Use the same Hydra config (morphology/arena/environment) " "that was used during training." ) - + # ======= SIMULATION ======= headless = bool(config.simulation.headless) max_steps = config.simulation.max_steps @@ -392,6 +639,9 @@ def main(dict_cfg: DictConfig) -> None: policy=policy, seed=seed, max_steps=max_steps_i, + action_low=action_low, + action_high=action_high, + padding_masks=padding_masks, ) final_dist_str = "n/a" if final_dist is None else f"{final_dist:.3f}" print( @@ -418,6 +668,9 @@ def main(dict_cfg: DictConfig) -> None: state=state0, control_dt=control_dt, max_steps=max_steps_val, + action_low=action_low, + action_high=action_high, + padding_masks=padding_masks, ) env.close() diff --git a/src/brittle_star_project/configs/config_simulation.py b/src/brittle_star_project/configs/config_simulation.py index 00488fc..872ca20 100644 --- a/src/brittle_star_project/configs/config_simulation.py +++ b/src/brittle_star_project/configs/config_simulation.py @@ -1,18 +1,19 @@ from dataclasses import dataclass from typing import Optional -from brittle_star_project.environment.env_types import Backend - @dataclass class SimulationSettings: """Settings for the simulation script.""" model_path: Optional[str] = None - model_type: str = "random" - backend: Backend = Backend.MJC # Script behavior headless: bool = False # If None, viewer mode runs until window closed or target reached. max_steps: Optional[int] = None + + # Optional: point to a Hydra config.yaml from a training run (e.g. runs/.../.hydra/config.yaml). + # When set, the simulation script can override + # morphology/arena/environment/architecture to match. + trained_config_path: Optional[str] = None diff --git a/src/experiment_logger/unified_logger.py b/src/experiment_logger/unified_logger.py index 634638d..f37136b 100644 --- a/src/experiment_logger/unified_logger.py +++ b/src/experiment_logger/unified_logger.py @@ -6,6 +6,7 @@ This logger ensures all experimental data is preserved by writing to: 3. stdout (for real-time monitoring) """ +from enum import Enum import logging import yaml import sys @@ -25,6 +26,33 @@ _active_logger: Optional[Any] = None _proxy_instance: Optional["LoggerProxy"] = None +def _sanitize_for_yaml(obj: Any) -> Any: + """Convert non-primitive values into YAML-safe structures. + + In particular, avoids PyYAML serializing Enums as + ``!!python/object/apply:...`` which OmegaConf will not load. + """ + + if isinstance(obj, Enum): + return obj.name + if isinstance(obj, Path): + return str(obj) + if isinstance(obj, (np.generic, jnp.ndarray)): + try: + return obj.item() + except Exception: + pass + if isinstance(obj, np.ndarray): + return obj.tolist() + if isinstance(obj, dict): + return {str(k): _sanitize_for_yaml(v) for k, v in obj.items()} + if isinstance(obj, list): + return [_sanitize_for_yaml(v) for v in obj] + if isinstance(obj, tuple): + return [_sanitize_for_yaml(v) for v in obj] + return obj + + def get_logger() -> "LoggerProxy": """Retrieve the global LoggerProxy. @@ -216,7 +244,13 @@ class UnifiedLogger: """Save configuration to disk.""" try: with open(self.config_file, "w") as f: - yaml.dump(self.full_config, f, default_flow_style=False, indent=2, sort_keys=False) + yaml.safe_dump( + _sanitize_for_yaml(self.full_config), + f, + default_flow_style=False, + indent=2, + sort_keys=False, + ) self.info(f"Config saved to {self.config_file}") except Exception as e: self.error(f"Error saving config: {e}") @@ -296,7 +330,12 @@ class UnifiedLogger: else: serializable_metric[k] = v f.write("---\n") - yaml.dump(serializable_metric, f, default_flow_style=False) + yaml.safe_dump( + _sanitize_for_yaml(serializable_metric), + f, + default_flow_style=False, + sort_keys=False, + ) self.metrics_buffer.clear() except Exception as e: self.error(f"Error flushing metrics: {e}") @@ -321,7 +360,13 @@ class UnifiedLogger: if metadata: metadata_path = self.checkpoints_dir / f"{prefix}_step_{step}_metadata.yaml" with open(metadata_path, "w") as f: - yaml.dump(metadata, f, default_flow_style=False) + yaml.safe_dump( + _sanitize_for_yaml(metadata), + f, + default_flow_style=False, + indent=2, + sort_keys=False, + ) self.info(f"Checkpoint saved: {checkpoint_path}") @@ -357,7 +402,13 @@ class UnifiedLogger: if metadata: metadata_path = self.run_dir / "final_model_metadata.yaml" with open(metadata_path, "w") as f: - yaml.dump(metadata, f, default_flow_style=False) + yaml.safe_dump( + _sanitize_for_yaml(metadata), + f, + default_flow_style=False, + indent=2, + sort_keys=False, + ) self.info(f"Final model saved: {final_model_path}") From 94e50f3d83c74ee372ef07b401a801800c1c67b3 Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Wed, 22 Apr 2026 18:51:08 +0200 Subject: [PATCH 10/10] fix: used jax type to remove warning --- .../environment/padded_obs_wrapper.py | 30 ++++++++++++++++--- 1 file changed, 26 insertions(+), 4 deletions(-) diff --git a/src/brittle_star_project/environment/padded_obs_wrapper.py b/src/brittle_star_project/environment/padded_obs_wrapper.py index 4886284..ae64de4 100644 --- a/src/brittle_star_project/environment/padded_obs_wrapper.py +++ b/src/brittle_star_project/environment/padded_obs_wrapper.py @@ -9,6 +9,8 @@ flattened observation maintains the correct physical mapping to the neural netwo from __future__ import annotations from typing import Any, Sequence + +import jax import jax.numpy as jnp # Observation keys whose size scales with the number of joints (2 per segment). @@ -78,11 +80,12 @@ def pad_observation( """Pad an observation dict using spatial insertion.""" padded = {} for key, value in obs.items(): + padded_dtype = _padding_dtype(value) if key in _JOINT_SCALED_KEYS: - out = jnp.zeros(masks["target_size_2x"], dtype=value.dtype) + out = jnp.zeros(masks["target_size_2x"], dtype=padded_dtype) padded[key] = out.at[masks["mask_2x"]].set(value) elif key in _SEGMENT_SCALED_KEYS: - out = jnp.zeros(masks["target_size_1x"], dtype=value.dtype) + out = jnp.zeros(masks["target_size_1x"], dtype=padded_dtype) padded[key] = out.at[masks["mask_1x"]].set(value) else: padded[key] = value @@ -97,12 +100,31 @@ def pad_observations_batched( padded = {} for key, value in obs.items(): batch_size = value.shape[0] + padded_dtype = _padding_dtype(value) if key in _JOINT_SCALED_KEYS: - out = jnp.zeros((batch_size, masks["target_size_2x"]), dtype=value.dtype) + out = jnp.zeros((batch_size, masks["target_size_2x"]), dtype=padded_dtype) padded[key] = out.at[:, masks["mask_2x"]].set(value) elif key in _SEGMENT_SCALED_KEYS: - out = jnp.zeros((batch_size, masks["target_size_1x"]), dtype=value.dtype) + out = jnp.zeros((batch_size, masks["target_size_1x"]), dtype=padded_dtype) padded[key] = out.at[:, masks["mask_1x"]].set(value) else: padded[key] = value return padded + + +def _padding_dtype(value: Any) -> jnp.dtype: + """Choose a JAX-safe dtype for padding arrays. + + When JAX x64 is disabled, allocating float64 zeros emits a warning. We + preserve the original dtype whenever it is supported, and otherwise fall + back to float32 for padding buffers. + """ + + dtype = getattr(value, "dtype", None) + if dtype is None: + dtype = jnp.asarray(value).dtype + else: + dtype = jnp.dtype(dtype) + if dtype == jnp.float64 and not jax.config.read("jax_enable_x64"): + return jnp.float32 + return dtype