From 9d1e2c9bfff74c291a4f5fd5b50d567c8f4ada6e Mon Sep 17 00:00:00 2001 From: Tibo De Peuter Date: Tue, 28 Apr 2026 13:39:01 +0200 Subject: [PATCH] refactor: evaluation subpackage --- scripts/simulate.py | 320 ++---------------- src/brittle_star_project/__init__.py | 16 +- .../evaluation/__init__.py | 17 + .../evaluation/checkpoint.py | 105 ++++++ src/brittle_star_project/evaluation/policy.py | 89 +++++ .../evaluation/rollout.py | 163 +++++++++ src/brittle_star_project/render/__init__.py | 3 - src/brittle_star_project/render/renderer.py | 78 ----- 8 files changed, 409 insertions(+), 382 deletions(-) create mode 100644 src/brittle_star_project/evaluation/__init__.py create mode 100644 src/brittle_star_project/evaluation/checkpoint.py create mode 100644 src/brittle_star_project/evaluation/policy.py create mode 100644 src/brittle_star_project/evaluation/rollout.py delete mode 100644 src/brittle_star_project/render/__init__.py delete mode 100644 src/brittle_star_project/render/renderer.py diff --git a/scripts/simulate.py b/scripts/simulate.py index bee6214..7bb2800 100644 --- a/scripts/simulate.py +++ b/scripts/simulate.py @@ -4,280 +4,29 @@ Automatically extracts the training configuration (morphology, environment, etc. from the sidecar metadata YAML file to ensure simulation perfectly matches training. Override simulation settings via CLI, e.g.: uv run scripts/simulate.py \ - simulation.morphology_override=config/morphology/3_arms.yaml \ + simulation.morphology_override=configs/morphology/3_arms.yaml \ simulation.model_path=runs/.../final_model.flax """ from __future__ import annotations -import itertools -import time from pathlib import Path -from typing import Any -import flax import hydra -import jax -import jax.numpy as jnp import numpy as np -import yaml from omegaconf import DictConfig, OmegaConf +import yaml 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 from brittle_star_project.environment.obs_processing import create_obs_processor -from brittle_star_project.environment.env_config import ( - MorphologyConfig, - ArenaConfig, - EnvConfig, - ObservationBoundsConfig, -) +from brittle_star_project.environment.env_config import MorphologyConfig - -class PolicyAgent: - """Wraps a trained Flax actor for deterministic inference.""" - - def __init__( - self, - *, - sensor_params: Any, - actor_params: Any, - action_dim: int, - obs_processor: Any, - ) -> None: - from brittle_star_project.MLPs.mlps import Actor, GenericDenseLayersWithActivation - - # Infer layer sizes from params - try: - dense_params = ( - sensor_params.get("params", {}) - if isinstance(sensor_params, dict) - else sensor_params["params"] - ) - except Exception: - dense_params = sensor_params - - layer_sizes = [] - idx = 0 - while True: - key = f"Dense_{idx}" - if key not in dense_params: - break - layer_sizes.append(int(np.asarray(dense_params[key]["kernel"]).shape[1])) - idx += 1 - - if not layer_sizes: - raise ValueError("Could not infer Dense_* layers from 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) - self._params = { - "sensor_params": sensor_params, - "actor_params": actor_params, - } - self._obs_processor = obs_processor - - @staticmethod - def load( - path: Path, - *, - action_dim: int, - obs_processor: Any, - ) -> "PolicyAgent": - payload = path.read_bytes() - restored = flax.serialization.msgpack_restore(payload) - - sensor_params = None - actor_params = None - - # Extract params from restored checkpoint - if isinstance(restored, dict): - params_sub = restored.get("params", {}) - sensor_params = restored.get("sensor_params") or params_sub.get("sensor_params") - actor_params = restored.get("actor_params") or params_sub.get("actor_params") - elif isinstance(restored, (list, tuple)) and len(restored) >= 2: - params_part = restored[1] - if isinstance(params_part, dict): - sensor_params = params_part.get("0", params_part.get(0)) - actor_params = params_part.get("1", params_part.get(1)) - elif isinstance(params_part, (list, tuple)) and len(params_part) >= 2: - sensor_params = params_part[0] - actor_params = params_part[1] - - if sensor_params is None or actor_params is None: - raise ValueError(f"Could not extract sensor and actor params from checkpoint: {path}") - - return PolicyAgent( - sensor_params=sensor_params, - actor_params=actor_params, - action_dim=action_dim, - obs_processor=obs_processor, - ) - - def act(self, *, observations: dict[str, Any]) -> np.ndarray: - batched_obs = jax.tree.map(lambda x: jnp.asarray(x)[None, ...], observations) - obs = self._obs_processor(batched_obs)[0] - 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. - # (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] | None: - 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) or getattr(state, "truncated", False)) - - -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 _rollout_headless( - *, - env: BrittleStarEnv, - policy: PolicyAgent, - seed: int, - max_steps: int, - action_low: np.ndarray | None, - action_high: np.ndarray | None, - action_mask: np.ndarray | None = None, -) -> tuple[float, int, bool, float | None]: - 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) - - steps = 0 - for _ in range(int(max_steps)): - obs_dict = observations or {} - - action = policy.act(observations=obs_dict) - if action_mask is not None: - action = action[action_mask] - action = _maybe_clip_action(action, action_low, action_high) - - 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 _rollout_viewer( - *, - env: BrittleStarEnv, - policy: PolicyAgent, - seed: int, - state: Any, - control_dt: float, - max_steps: int | None, - action_low: np.ndarray | None, - action_high: np.ndarray | None, - action_mask: np.ndarray | None = None, -) -> None: - import mujoco.viewer - - model = state.mj_model - data = state.mj_data - - episode_return = 0.0 - observations = _get_observations(state) - prev_dist = _get_xy_distance_to_target(observations) - reached_target = _target_reached(state=state) - - steps = 0 - # Use the viewer as a context manager to avoid GLX teardown races. - 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() - - obs_dict = observations or {} - - action = policy.act(observations=obs_dict) - if action_mask is not None: - action = action[action_mask] - action = _maybe_clip_action(action, action_low, action_high) - - # 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 - - remaining = control_dt - (time.time() - step_start) - if remaining > 0: - time.sleep(remaining) - - 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 _load_metadata_yaml(model_path: Path) -> dict: - """Discover and load the sidecar metadata YAML file.""" - metadata_path = model_path.with_name(model_path.stem + "_metadata.yaml") - if not metadata_path.exists(): - raise FileNotFoundError( - f"Could not find metadata YAML for {model_path.name}. Expected it at {metadata_path}" - ) - with open(metadata_path, "r") as f: - return yaml.safe_load(f) +from brittle_star_project.evaluation.checkpoint import load_metadata, metadata_to_configs +from brittle_star_project.evaluation.policy import PolicyAgent +from brittle_star_project.evaluation.rollout import rollout_headless, rollout_viewer @hydra.main(config_path="../configs", config_name="main_config", version_base="1.3") @@ -297,36 +46,10 @@ def main(dict_cfg: DictConfig) -> None: raise ValueError(f"Expected a '.flax' checkpoint, got '{model_path.name}'.") # 2. Discover + load sidecar metadata YAML - metadata = _load_metadata_yaml(model_path) + metadata = load_metadata(model_path) # 3. Reconstruct typed configs from metadata - trained_morphology = OmegaConf.to_object( - OmegaConf.merge(OmegaConf.structured(MorphologyConfig), metadata.get("morphology", {})) - ) - trained_arena = OmegaConf.to_object( - OmegaConf.merge(OmegaConf.structured(ArenaConfig), metadata.get("arena", {})) - ) - - env_dict = metadata.get("environment", {}) - if isinstance(env_dict.get("task"), str): - from brittle_star_project.environment.env_types import Task - - try: - env_dict["task"] = Task[env_dict["task"]].name - except Exception: - try: - env_dict["task"] = Task(env_dict["task"]).name - except Exception: - pass - - trained_environment = OmegaConf.to_object( - OmegaConf.merge(OmegaConf.structured(EnvConfig), env_dict) - ) - trained_obs_bounds = OmegaConf.to_object( - OmegaConf.merge( - OmegaConf.structured(ObservationBoundsConfig), metadata.get("obs_bounds", {}) - ) - ) + training = metadata_to_configs(metadata) # 4. Determine environment morphology if sim_cfg.morphology_override is not None: @@ -339,14 +62,15 @@ def main(dict_cfg: DictConfig) -> None: OmegaConf.merge(OmegaConf.structured(MorphologyConfig), override_dict) ) else: - env_morphology = trained_morphology + env_morphology = training.morphology # 5. Build obs_processor with TRAINING morphology padding masks always padding_masks = compute_padding_masks( segments_per_arm=env_morphology.segments_per_arm, + reference_segments_per_arm=training.morphology.segments_per_arm, ) obs_processor = create_obs_processor( - bounds_dict=trained_obs_bounds.to_bounds_dict(), + bounds_dict=training.obs_bounds.to_bounds_dict(), padding_masks=padding_masks, ) @@ -358,23 +82,23 @@ def main(dict_cfg: DictConfig) -> None: raw_env = factory.create_environment( backend, env_morphology, - trained_arena, - trained_environment, + training.arena, + training.environment, ) env = BrittleStarEnv( raw_env, backend=backend, - config=trained_environment, + config=training.environment, morphology_config=env_morphology, ) state0 = env.reset(seed=seed) # Calculate the action dimension the model was trained with - trained_action_dim = sum(trained_morphology.segments_per_arm) * 2 + trained_action_dim = sum(training.morphology.segments_per_arm) * 2 # 7. Load policy - policy = PolicyAgent.load( + policy = PolicyAgent.from_checkpoint( model_path, action_dim=trained_action_dim, obs_processor=obs_processor ) @@ -402,7 +126,7 @@ def main(dict_cfg: DictConfig) -> None: if max_steps_i <= 0: raise ValueError("simulation.max_steps must be > 0") - ep_return, ep_len, reached_target, final_dist = _rollout_headless( + result = rollout_headless( env=env, policy=policy, seed=seed, @@ -411,11 +135,11 @@ def main(dict_cfg: DictConfig) -> None: action_high=action_high, action_mask=action_mask, ) - final_dist_str = "n/a" if final_dist is None else f"{final_dist:.3f}" + final_dist_str = "n/a" if result.final_xy_dist is None else f"{result.final_xy_dist:.3f}" print( "episode done: " - f"return={ep_return:.6f}, len={ep_len}, " - f"target_reached={reached_target}, final_xy_dist={final_dist_str}" + f"return={result.return_:.6f}, len={result.length}, " + f"target_reached={result.reached_target}, final_xy_dist={final_dist_str}" ) else: max_steps_val = None @@ -426,9 +150,9 @@ def main(dict_cfg: DictConfig) -> None: max_steps_val = max_steps_i model_dt = float(state0.mj_model.opt.timestep) - control_dt = model_dt * float(trained_environment.num_physics_steps_per_control_step) + control_dt = model_dt * float(training.environment.num_physics_steps_per_control_step) - _rollout_viewer( + rollout_viewer( env=env, policy=policy, seed=seed, diff --git a/src/brittle_star_project/__init__.py b/src/brittle_star_project/__init__.py index 3902cbb..4eec766 100644 --- a/src/brittle_star_project/__init__.py +++ b/src/brittle_star_project/__init__.py @@ -2,7 +2,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 -from .render import simulate_policy, SimulationConfig, ControlPolicy +from .evaluation import ( + PolicyAgent, + ControlPolicy, + load_metadata, + rollout_headless, + rollout_viewer, + EpisodeResult, +) __all__ = [ "ArenaConfig", @@ -12,7 +19,10 @@ __all__ = [ "EnvConfig", "MorphologyConfig", "Task", - "simulate_policy", - "SimulationConfig", + "PolicyAgent", "ControlPolicy", + "load_metadata", + "rollout_headless", + "rollout_viewer", + "EpisodeResult", ] diff --git a/src/brittle_star_project/evaluation/__init__.py b/src/brittle_star_project/evaluation/__init__.py new file mode 100644 index 0000000..73719e6 --- /dev/null +++ b/src/brittle_star_project/evaluation/__init__.py @@ -0,0 +1,17 @@ +from __future__ import annotations + +from .checkpoint import load_metadata, load_params, metadata_to_configs, TrainingConfig +from .policy import PolicyAgent, ControlPolicy +from .rollout import rollout_headless, rollout_viewer, EpisodeResult + +__all__ = [ + "load_metadata", + "load_params", + "metadata_to_configs", + "TrainingConfig", + "PolicyAgent", + "ControlPolicy", + "rollout_headless", + "rollout_viewer", + "EpisodeResult", +] diff --git a/src/brittle_star_project/evaluation/checkpoint.py b/src/brittle_star_project/evaluation/checkpoint.py new file mode 100644 index 0000000..ffbc27a --- /dev/null +++ b/src/brittle_star_project/evaluation/checkpoint.py @@ -0,0 +1,105 @@ +from __future__ import annotations + +import yaml +from dataclasses import dataclass +from pathlib import Path + +import flax +from omegaconf import OmegaConf + +from brittle_star_project.environment.env_config import ( + MorphologyConfig, + ArenaConfig, + EnvConfig, + ObservationBoundsConfig, +) + + +@dataclass +class TrainingConfig: + """Holds typed configurations extracted from a training run's metadata.""" + + morphology: MorphologyConfig + arena: ArenaConfig + environment: EnvConfig + obs_bounds: ObservationBoundsConfig + + +def load_params(path: Path) -> dict: + """Load model parameters from a .flax checkpoint file.""" + payload = path.read_bytes() + restored = flax.serialization.msgpack_restore(payload) + + sensor_params = None + actor_params = None + + # Extract params from restored checkpoint + if isinstance(restored, dict): + params_sub = restored.get("params", {}) + sensor_params = restored.get("sensor_params") or params_sub.get("sensor_params") + actor_params = restored.get("actor_params") or params_sub.get("actor_params") + elif isinstance(restored, (list, tuple)) and len(restored) >= 2: + params_part = restored[1] + if isinstance(params_part, dict): + sensor_params = params_part.get("0", params_part.get(0)) + actor_params = params_part.get("1", params_part.get(1)) + elif isinstance(params_part, (list, tuple)) and len(params_part) >= 2: + sensor_params = params_part[0] + actor_params = params_part[1] + + if sensor_params is None or actor_params is None: + raise ValueError(f"Could not extract sensor and actor params from checkpoint: {path}") + + return { + "sensor_params": sensor_params, + "actor_params": actor_params, + } + + +def load_metadata(model_path: Path) -> dict: + """Discover and load the sidecar metadata YAML file.""" + metadata_path = model_path.with_name(model_path.stem + "_metadata.yaml") + if not metadata_path.exists(): + raise FileNotFoundError( + f"Could not find metadata YAML for {model_path.name}. Expected it at {metadata_path}" + ) + with open(metadata_path, "r") as f: + return yaml.safe_load(f) + + +def metadata_to_configs(metadata: dict) -> TrainingConfig: + """Reconstruct typed configuration objects from a metadata dictionary.""" + trained_morphology = OmegaConf.to_object( + OmegaConf.merge(OmegaConf.structured(MorphologyConfig), metadata.get("morphology", {})) + ) + trained_arena = OmegaConf.to_object( + OmegaConf.merge(OmegaConf.structured(ArenaConfig), metadata.get("arena", {})) + ) + + env_dict = metadata.get("environment", {}) + if isinstance(env_dict.get("task"), str): + from brittle_star_project.environment.env_types import Task + + try: + env_dict["task"] = Task[env_dict["task"]].name + except Exception: + try: + env_dict["task"] = Task(env_dict["task"]).name + except Exception: + pass + + trained_environment = OmegaConf.to_object( + OmegaConf.merge(OmegaConf.structured(EnvConfig), env_dict) + ) + trained_obs_bounds = OmegaConf.to_object( + OmegaConf.merge( + OmegaConf.structured(ObservationBoundsConfig), metadata.get("obs_bounds", {}) + ) + ) + + return TrainingConfig( + morphology=trained_morphology, + arena=trained_arena, + environment=trained_environment, + obs_bounds=trained_obs_bounds, + ) diff --git a/src/brittle_star_project/evaluation/policy.py b/src/brittle_star_project/evaluation/policy.py new file mode 100644 index 0000000..5aa2d6e --- /dev/null +++ b/src/brittle_star_project/evaluation/policy.py @@ -0,0 +1,89 @@ +from __future__ import annotations + +from pathlib import Path +from typing import Any, Protocol + +import jax +import jax.numpy as jnp +import numpy as np + +from brittle_star_project.evaluation.checkpoint import load_params + + +class ControlPolicy(Protocol): + """Protocol for any policy that can produce actions from observations.""" + + def act(self, *, observations: dict[str, Any]) -> np.ndarray: ... + + +class PolicyAgent: + """Wraps a trained Flax actor for deterministic inference.""" + + def __init__( + self, + *, + sensor_params: Any, + actor_params: Any, + action_dim: int, + obs_processor: Any, + ) -> None: + from brittle_star_project.MLPs.mlps import Actor, GenericDenseLayersWithActivation + + # Infer layer sizes from params + try: + dense_params = ( + sensor_params.get("params", {}) + if isinstance(sensor_params, dict) + else sensor_params["params"] + ) + except Exception: + dense_params = sensor_params + + layer_sizes = [] + idx = 0 + while True: + key = f"Dense_{idx}" + if key not in dense_params: + break + layer_sizes.append(int(np.asarray(dense_params[key]["kernel"]).shape[1])) + idx += 1 + + if not layer_sizes: + raise ValueError("Could not infer Dense_* layers from 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) + self._params = { + "sensor_params": sensor_params, + "actor_params": actor_params, + } + self._obs_processor = obs_processor + + @classmethod + def from_checkpoint( + cls, + model_path: Path, + *, + action_dim: int, + obs_processor: Any, + ) -> "PolicyAgent": + """Load params from .flax and construct the agent.""" + params = load_params(model_path) + + return cls( + sensor_params=params["sensor_params"], + actor_params=params["actor_params"], + action_dim=action_dim, + obs_processor=obs_processor, + ) + + def act(self, *, observations: dict[str, Any]) -> np.ndarray: + """Return deterministic action (actor mean, no exploration noise).""" + batched_obs = jax.tree.map(lambda x: jnp.asarray(x)[None, ...], observations) + obs = self._obs_processor(batched_obs)[0] + hidden = self._sensor_apply(self._params["sensor_params"], obs) + mean, _log_std = self._actor_apply(self._params["actor_params"], hidden) + + return np.asarray(mean, dtype=np.float32).ravel() diff --git a/src/brittle_star_project/evaluation/rollout.py b/src/brittle_star_project/evaluation/rollout.py new file mode 100644 index 0000000..79c86ed --- /dev/null +++ b/src/brittle_star_project/evaluation/rollout.py @@ -0,0 +1,163 @@ +from __future__ import annotations + +import itertools +import time +from dataclasses import dataclass +from typing import Any + +import numpy as np + +from brittle_star_project import BrittleStarEnv +from brittle_star_project.evaluation.policy import ControlPolicy + + +@dataclass +class EpisodeResult: + return_: float + length: int + reached_target: bool + final_xy_dist: float | None + + +def _get_observations(state: Any) -> dict[str, Any] | None: + 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) or getattr(state, "truncated", False)) + + +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 rollout_headless( + *, + env: BrittleStarEnv, + policy: ControlPolicy, + seed: int, + max_steps: int, + action_low: np.ndarray | None, + action_high: np.ndarray | None, + action_mask: np.ndarray | None = None, +) -> EpisodeResult: + """Run an episode headlessly and return the result.""" + state = env.reset(seed=seed) + + ep_return = 0.0 + observations = _get_observations(state) + prev_dist = _get_xy_distance_to_target(observations) if observations else None + reached_target = _target_reached(state=state) + + steps = 0 + for _ in range(int(max_steps)): + obs_dict = observations or {} + + action = policy.act(observations=obs_dict) + if action_mask is not None: + action = action[action_mask] + action = _maybe_clip_action(action, action_low, action_high) + + state = env.step(state=state, action=action) + steps += 1 + + observations = _get_observations(state) + cur_dist = _get_xy_distance_to_target(observations) if observations else None + 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) if observations else None + return EpisodeResult( + return_=ep_return, + length=steps, + reached_target=reached_target, + final_xy_dist=final_dist, + ) + + +def rollout_viewer( + *, + env: BrittleStarEnv, + policy: ControlPolicy, + seed: int, + state: Any, + control_dt: float, + max_steps: int | None, + action_low: np.ndarray | None, + action_high: np.ndarray | None, + action_mask: np.ndarray | None = None, +) -> None: + """Run an episode using the interactive MuJoCo viewer.""" + import mujoco.viewer + + model = state.mj_model + data = state.mj_data + + episode_return = 0.0 + observations = _get_observations(state) + prev_dist = _get_xy_distance_to_target(observations) if observations else None + reached_target = _target_reached(state=state) + + steps = 0 + 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() + + obs_dict = observations or {} + action = policy.act(observations=obs_dict) + if action_mask is not None: + action = action[action_mask] + action = _maybe_clip_action(action, action_low, action_high) + + 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 observations else None + 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 + + remaining = control_dt - (time.time() - step_start) + if remaining > 0: + time.sleep(remaining) + + dist = _get_xy_distance_to_target(observations) if observations else None + 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}" + ) diff --git a/src/brittle_star_project/render/__init__.py b/src/brittle_star_project/render/__init__.py deleted file mode 100644 index 51e0aa1..0000000 --- a/src/brittle_star_project/render/__init__.py +++ /dev/null @@ -1,3 +0,0 @@ -from .renderer import simulate_policy, SimulationConfig, ControlPolicy - -__all__ = ["simulate_policy", "SimulationConfig", "ControlPolicy"] diff --git a/src/brittle_star_project/render/renderer.py b/src/brittle_star_project/render/renderer.py deleted file mode 100644 index 91e669c..0000000 --- a/src/brittle_star_project/render/renderer.py +++ /dev/null @@ -1,78 +0,0 @@ -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 - - if state is None: - raise ValueError("A valid environment state must be provided.") - - 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)