From dd0e5b2d09c0af670a1058e8b0dc23d0b77ca686 Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Thu, 14 May 2026 12:13:34 +0200 Subject: [PATCH 01/16] feat: added vids/ folder to gitignore --- .gitignore | 1 + 1 file changed, 1 insertion(+) diff --git a/.gitignore b/.gitignore index dd76b50..1095a1e 100644 --- a/.gitignore +++ b/.gitignore @@ -6,6 +6,7 @@ outputs/ multirun/ metrics/ adjacency_debug.txt +vids/ # Python-generated files __pycache__/ From b84e8800ad0a16f4c373ee020f6d6d610caea588 Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Thu, 14 May 2026 13:39:47 +0200 Subject: [PATCH 02/16] feat: added video configs fields --- configs/simulation/default.yaml | 4 ++++ scripts/simulate.py | 3 +++ src/brittle_star_project/configs/config_simulation.py | 3 +++ 3 files changed, 10 insertions(+) diff --git a/configs/simulation/default.yaml b/configs/simulation/default.yaml index 61599a4..5bd448a 100644 --- a/configs/simulation/default.yaml +++ b/configs/simulation/default.yaml @@ -21,6 +21,10 @@ video_output_path: null # Camera ID to use for video recording (1 is usually the close-up camera) camera_id: 1 +video_width: int = 640 +video_height: int = 480 +video_fps: int = 60 + # Optional override for the metadata YAML file path. # If null, the script looks for `_metadata.yaml` alongside the model_path. metadata_path: null diff --git a/scripts/simulate.py b/scripts/simulate.py index 9f3e550..db8f70d 100644 --- a/scripts/simulate.py +++ b/scripts/simulate.py @@ -107,6 +107,9 @@ def main(dict_cfg: DictConfig) -> None: action_mask=action_mask, output_path=output_path, camera_id=sim_cfg.camera_id, + width=sim_cfg.video_width, + height=sim_cfg.video_height, + fps=sim_cfg.video_fps, ) save_evaluation_metadata( diff --git a/src/brittle_star_project/configs/config_simulation.py b/src/brittle_star_project/configs/config_simulation.py index 9fd1507..0a75b8d 100644 --- a/src/brittle_star_project/configs/config_simulation.py +++ b/src/brittle_star_project/configs/config_simulation.py @@ -26,6 +26,9 @@ class SimulationSettings: video_output_path: Optional[str] = None # Camera ID to use for video recording (1 is usually the close-up camera) camera_id: int = 1 + video_width: int = 640 + video_height: int = 480 + video_fps: int = 60 # Optional override for the sidecar metadata YAML file. # If None, it defaults to the model_path with a `_metadata.yaml` suffix. From 0defc9a7c462d0fb4a4ec023d7c2c2fe8c9934be Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Thu, 14 May 2026 13:44:53 +0200 Subject: [PATCH 03/16] feat: render seperate topdown and follow camera video --- .../render_poster_videos.py | 103 ++++++++++++++++++ .../render_poster_videos.sh | 11 ++ src/brittle_star_project/evaluation/video.py | 94 ++++++++++++++++ 3 files changed, 208 insertions(+) create mode 100644 scripts/poster_visualisations/render_poster_videos.py create mode 100755 scripts/poster_visualisations/render_poster_videos.sh diff --git a/scripts/poster_visualisations/render_poster_videos.py b/scripts/poster_visualisations/render_poster_videos.py new file mode 100644 index 0000000..b221687 --- /dev/null +++ b/scripts/poster_visualisations/render_poster_videos.py @@ -0,0 +1,103 @@ +from __future__ import annotations + +import argparse +from pathlib import Path + +from brittle_star_project.evaluation.checkpoint import load_metadata, metadata_to_configs +from brittle_star_project.evaluation.eval_env_builder import build_eval_env +from brittle_star_project.evaluation.video import record_episode_multi_camera + +_ARCH_DIR_MAP = { + "CENTRALIZED": "centralized", + "FULLY_CONNECTED": "fully-connected", + "RING": "ring", + "SEGMENT": "segment", +} + + +def _arch_dir(name: str) -> str: + return _ARCH_DIR_MAP.get(name, name.lower()) + + +def _resolve_overrides(overrides: list[str], count: int) -> list[str | None]: + if not overrides: + return [None] * count + if len(overrides) == 1 and count > 1: + return overrides * count + if len(overrides) != count: + raise ValueError("morphology overrides must match the number of models") + return overrides + + +def main() -> None: + parser = argparse.ArgumentParser(description="Render top-down and follow videos for poster.") + parser.add_argument("models", nargs="+", help="Paths to .flax checkpoints") + parser.add_argument( + "--morphology-override", + action="append", + default=[], + help="Override morphology YAML path (repeat to match models)", + ) + parser.add_argument("--output-root", default="vids/poster") + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--max-steps", type=int, required=True) + parser.add_argument("--topdown-camera", type=int, default=0) + parser.add_argument("--follow-camera", type=int, default=1) + parser.add_argument("--width", type=int, default=640) + parser.add_argument("--height", type=int, default=480) + parser.add_argument("--fps", type=int, default=60) + args = parser.parse_args() + + overrides = _resolve_overrides(args.morphology_override, len(args.models)) + output_root = Path(args.output_root) + + for model_path_str, override in zip(args.models, overrides): + model_path = Path(model_path_str) + metadata = load_metadata(model_path, None) + training = metadata_to_configs(metadata) + + bundle = build_eval_env( + model_path=model_path, + training=training, + metadata=metadata, + morphology_override_path=override, + ) + + arch_dir = _arch_dir(bundle.architecture) + arms_dir = f"{bundle.num_active_arms}arms" + out_dir = output_root / arms_dir / arch_dir + out_dir.mkdir(parents=True, exist_ok=True) + + output_paths = { + args.topdown_camera: out_dir / "topdown.mp4", + args.follow_camera: out_dir / "follow.mp4", + } + + result = record_episode_multi_camera( + env=bundle.env, + policy=bundle.policy, + seed=args.seed, + max_steps=args.max_steps, + action_low=bundle.action_low, + action_high=bundle.action_high, + action_mask=bundle.action_mask, + output_paths=output_paths, + camera_ids=[args.topdown_camera, args.follow_camera], + width=args.width, + height=args.height, + fps=args.fps, + ) + + final_dist = ( + "n/a" if result.final_xy_dist is None else f"{result.final_xy_dist:.3f}" + ) + print( + f"{arms_dir}/{arch_dir}: return={result.return_:.6f}, len={result.length}, " + f"target_reached={result.reached_target}, final_xy_dist={final_dist}" + ) + + bundle.env.close() + + +if __name__ == "__main__": + main() diff --git a/scripts/poster_visualisations/render_poster_videos.sh b/scripts/poster_visualisations/render_poster_videos.sh new file mode 100755 index 0000000..a6f5b59 --- /dev/null +++ b/scripts/poster_visualisations/render_poster_videos.sh @@ -0,0 +1,11 @@ +#!/usr/bin/env bash + +# Multi-camera renders per model -> vids/poster/{arms}arms/{arch}/topdown.mp4 + follow.mp4 + +path=$1 + +uv run scripts/poster_visualisations/render_poster_videos.py \ + "$path"/centralized.flax \ + "$path"/fully-connected.flax \ + "$path"/ring.flax \ + --max-steps 10000 --width 640 --height 480 --fps 60 \ No newline at end of file diff --git a/src/brittle_star_project/evaluation/video.py b/src/brittle_star_project/evaluation/video.py index b4b44db..0ba3ce7 100644 --- a/src/brittle_star_project/evaluation/video.py +++ b/src/brittle_star_project/evaluation/video.py @@ -147,3 +147,97 @@ def record_episode( final_xy_dist=final_dist, initial_target_distance=initial_dist, ) + + +def record_episode_multi_camera( + *, + env: BrittleStarEnv, + policy: ControlPolicy, + seed: int, + max_steps: int, + action_low: np.ndarray | None, + action_high: np.ndarray | None, + output_paths: dict[int, Path], + action_mask: np.ndarray | None = None, + camera_ids: list[int] | None = None, + fps: int = 60, + width: int = 640, + height: int = 480, +) -> EpisodeResult: + """Run one episode and render multiple camera views to separate files.""" + try: + import imageio + import mujoco + except ImportError as e: + raise ImportError( + "Video recording requires 'imageio' and 'mujoco'. " + "Please install the evaluation dependencies: `uv pip install .[evaluation]`" + ) from e + + if camera_ids is None: + camera_ids = list(output_paths.keys()) + + for cam_id in camera_ids: + if cam_id not in output_paths: + raise ValueError(f"Missing output path for camera {cam_id}") + + output_paths = {cam_id: output_paths[cam_id] for cam_id in camera_ids} + + for path in output_paths.values(): + path.parent.mkdir(parents=True, exist_ok=True) + + state = env.reset(seed=seed) + model = state.mj_model + data = state.mj_data + + renderer = mujoco.Renderer(model, width=width, height=height) + frames = {cam_id: [] for cam_id in camera_ids} + + ep_return = 0.0 + observations = _get_observations(state) + prev_dist = _get_xy_distance_to_target(observations) if observations else None + initial_dist = prev_dist + reached_target = _target_reached(state=state) + + steps = 0 + for _ in range(int(max_steps)): + for cam_id in camera_ids: + renderer.update_scene(data, camera=cam_id) + frames[cam_id].append(renderer.render()) + + 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 + + for cam_id in camera_ids: + renderer.update_scene(data, camera=cam_id) + frames[cam_id].append(renderer.render()) + + renderer.close() + + for cam_id, path in output_paths.items(): + imageio.mimsave(str(path), frames[cam_id], fps=fps) + + 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, + initial_target_distance=initial_dist, + ) From ecc476433228ff100505e8f967b42e10cfee17ef Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Thu, 14 May 2026 13:54:19 +0200 Subject: [PATCH 04/16] docs: updated simulation docs for render_poster_videos Co-authored-by: Copilot --- docs/api/simulation.md | 13 +++++++++++++ 1 file changed, 13 insertions(+) diff --git a/docs/api/simulation.md b/docs/api/simulation.md index 1e288b9..45d7878 100644 --- a/docs/api/simulation.md +++ b/docs/api/simulation.md @@ -38,4 +38,17 @@ uv run scripts/simulate.py \ Videos and evaluation metadata are stored in timestamped folders alongside the model: `runs/your_run/final_model_evaluations/eval_/simulation.mp4` +### Top-Down and Follow Cameras + +Using the following script, you can render a top-down and follow camera view for multiple models at once: + +```bash +uv run scripts/poster_visualisations/render_poster_videos.py \ + runs/final-models/centralized/.../final_model.flax \ + runs/final-models/fully-connected/.../final_model.flax \ + runs/final-models/ring/.../final_model.flax \ + --max-steps 10000 --width 640 --height 480 --fps 60 \ + --output-root vids/poster/ +``` + For batch evaluation and cross-model comparison, see the **[Evaluation Guide](./evaluation.md)**. From a125a0948fe539d92fcaa7a1a5f934411255201f Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Thu, 14 May 2026 18:19:48 +0200 Subject: [PATCH 05/16] feat: customizable camera & target --- .../render_poster_videos.py | 40 ++++++++++++-- .../render_poster_videos.sh | 3 +- .../environment/BrittleStarJaxEnvWrapper.py | 14 ++++- .../environment/env_wrapper.py | 7 ++- src/brittle_star_project/evaluation/video.py | 52 ++++++++++++++++++- 5 files changed, 105 insertions(+), 11 deletions(-) diff --git a/scripts/poster_visualisations/render_poster_videos.py b/scripts/poster_visualisations/render_poster_videos.py index b221687..d3cf697 100644 --- a/scripts/poster_visualisations/render_poster_videos.py +++ b/scripts/poster_visualisations/render_poster_videos.py @@ -40,14 +40,45 @@ def main() -> None: ) parser.add_argument("--output-root", default="vids/poster") parser.add_argument("--seed", type=int, default=0) - parser.add_argument("--max-steps", type=int, required=True) + parser.add_argument("--max-steps", type=int, default=2000) parser.add_argument("--topdown-camera", type=int, default=0) parser.add_argument("--follow-camera", type=int, default=1) - parser.add_argument("--width", type=int, default=640) - parser.add_argument("--height", type=int, default=480) + parser.add_argument("--topdown-camera-x", type=float, default=-3.0) + parser.add_argument("--topdown-camera-y", type=float, default=0.0) + parser.add_argument("--topdown-camera-z", type=float, default=None) + parser.add_argument("--topdown-camera-fovy", type=float, default=50.0) + parser.add_argument("--target-x", type=float, default=-6.0) + parser.add_argument("--target-y", type=float, default=0.0) + parser.add_argument("--width", type=int, default=1920) + parser.add_argument("--height", type=int, default=1088) parser.add_argument("--fps", type=int, default=60) args = parser.parse_args() + if (args.target_x is None) != (args.target_y is None): + raise ValueError("target-x and target-y must be provided together") + + target_xy = None + if args.target_x is not None: + target_xy = (float(args.target_x), float(args.target_y)) + + camera_fovy = None + if args.topdown_camera_fovy is not None: + camera_fovy = {args.topdown_camera: float(args.topdown_camera_fovy)} + + camera_x = None + if args.topdown_camera_x is not None: + camera_x = {args.topdown_camera: float(args.topdown_camera_x)} + + camera_y = None + if args.topdown_camera_y is not None: + camera_y = {args.topdown_camera: float(args.topdown_camera_y)} + + camera_z = None + if args.topdown_camera_z is not None: + camera_z = {args.topdown_camera: float(args.topdown_camera_z)} + + camera_xyz = (camera_x, camera_y, camera_z) + overrides = _resolve_overrides(args.morphology_override, len(args.models)) output_root = Path(args.output_root) @@ -83,6 +114,9 @@ def main() -> None: action_mask=bundle.action_mask, output_paths=output_paths, camera_ids=[args.topdown_camera, args.follow_camera], + camera_fovy=camera_fovy, + camera_xyz=camera_xyz, + target_xy=target_xy, width=args.width, height=args.height, fps=args.fps, diff --git a/scripts/poster_visualisations/render_poster_videos.sh b/scripts/poster_visualisations/render_poster_videos.sh index a6f5b59..10648f3 100755 --- a/scripts/poster_visualisations/render_poster_videos.sh +++ b/scripts/poster_visualisations/render_poster_videos.sh @@ -7,5 +7,4 @@ path=$1 uv run scripts/poster_visualisations/render_poster_videos.py \ "$path"/centralized.flax \ "$path"/fully-connected.flax \ - "$path"/ring.flax \ - --max-steps 10000 --width 640 --height 480 --fps 60 \ No newline at end of file + "$path"/ring.flax \ \ No newline at end of file diff --git a/src/brittle_star_project/environment/BrittleStarJaxEnvWrapper.py b/src/brittle_star_project/environment/BrittleStarJaxEnvWrapper.py index b514c11..cfb60cd 100644 --- a/src/brittle_star_project/environment/BrittleStarJaxEnvWrapper.py +++ b/src/brittle_star_project/environment/BrittleStarJaxEnvWrapper.py @@ -65,11 +65,21 @@ class BrittleStarJaxEnvWrapper: def single_observation_space(self): return self._env.observation_space - def reset(self, seed: int = 0): + def reset(self, seed: int = 0, target_position: tuple[float, float] | None = None): self.logger.info(f"Resetting vectorized environment environments with seed {seed}") self._action_rng, env_rng = jax.random.split(jax.random.PRNGKey(seed), 2) env_rngs = jnp.array(jax.random.split(env_rng, self._num_envs)) - state = self._vectorized_reset(rng=env_rngs) + + # If a target_position is provided, pass it through to the underlying env.reset + if target_position is None: + state = jax.jit(jax.vmap(lambda rng: self._env.reset(rng=rng)))(env_rngs) + else: + tp = jnp.asarray(target_position, dtype=jnp.float32) + tp_batched = jnp.tile(tp[None, :], (self._num_envs, 1)) + state = jax.jit(jax.vmap(lambda rng, t: self._env.reset(rng=rng, target_position=t)))( + env_rngs, tp_batched + ) + return state def sample_actions(self): diff --git a/src/brittle_star_project/environment/env_wrapper.py b/src/brittle_star_project/environment/env_wrapper.py index 8b3e51f..0d2463c 100644 --- a/src/brittle_star_project/environment/env_wrapper.py +++ b/src/brittle_star_project/environment/env_wrapper.py @@ -62,9 +62,12 @@ class BrittleStarEnv: return jax.random.PRNGKey(seed) - def reset(self, *, seed: int = 0): + def reset(self, *, seed: int = 0, target_position: tuple[float, float, float] | None = None): rng = self.make_rng(seed) - state = self._env.reset(rng=rng) + if target_position is not None: + state = self._env.reset(rng=rng, target_position=target_position) + else: + state = self._env.reset(rng=rng) return state def render(self, *, state: Any): diff --git a/src/brittle_star_project/evaluation/video.py b/src/brittle_star_project/evaluation/video.py index 0ba3ce7..a7a629d 100644 --- a/src/brittle_star_project/evaluation/video.py +++ b/src/brittle_star_project/evaluation/video.py @@ -25,6 +25,44 @@ def create_evaluation_dir(model_path: Path) -> Path: return eval_dir +def _ensure_offscreen_size(model, width: int, height: int) -> None: + vis_global = getattr(getattr(model, "vis", None), "global_", None) + if vis_global is None: + return + vis_global.offwidth = int(max(width, vis_global.offwidth)) + vis_global.offheight = int(max(height, vis_global.offheight)) + + +def _apply_camera_overrides( + model, + *, + camera_fovy: dict[int, float] | None = None, + camera_xyz: tuple[dict[int, float] | None, dict[int, float] | None, dict[int, float] | None] = (None, None, None), +) -> None: + if not camera_fovy and not (camera_xyz[0] or camera_xyz[1] or camera_xyz[2]): + return + + ncam = int(getattr(model, "ncam", 0)) + for cam_id, fovy in (camera_fovy or {}).items(): + if cam_id < 0 or cam_id >= ncam: + raise ValueError(f"Camera id {cam_id} is out of range") + model.cam_fovy[cam_id] = float(fovy) + + for cam_id, x in (camera_xyz[0] or {}).items(): + if cam_id < 0 or cam_id >= ncam: + raise ValueError(f"Camera id {cam_id} is out of range") + model.cam_pos[cam_id][0] = float(x) + + for cam_id, y in (camera_xyz[1] or {}).items(): + if cam_id < 0 or cam_id >= ncam: + raise ValueError(f"Camera id {cam_id} is out of range") + model.cam_pos[cam_id][1] = float(y) + + for cam_id, z in (camera_xyz[2] or {}).items(): + if cam_id < 0 or cam_id >= ncam: + raise ValueError(f"Camera id {cam_id} is out of range") + model.cam_pos[cam_id][2] = float(z) + def save_evaluation_metadata( eval_dir: Path, *, @@ -66,6 +104,7 @@ def record_episode( fps: int = 60, width: int = 640, height: int = 480, + target_xy: tuple[float, float] | None = None, ) -> EpisodeResult: """Run an episode headlessly and record a video using MuJoCo's Renderer and imageio. @@ -92,10 +131,12 @@ def record_episode( "Please install the evaluation dependencies: `uv pip install .[evaluation]`" ) from e - state = env.reset(seed=seed) + state = env.reset(seed=seed, target_position=target_xy) model = state.mj_model data = state.mj_data + _ensure_offscreen_size(model, width, height) + renderer = mujoco.Renderer(model, width=width, height=height) ep_return = 0.0 observations = _get_observations(state) @@ -160,6 +201,9 @@ def record_episode_multi_camera( output_paths: dict[int, Path], action_mask: np.ndarray | None = None, camera_ids: list[int] | None = None, + camera_fovy: dict[int, float] | None = None, + camera_xyz: tuple[dict[int, float] | None, dict[int, float] | None, dict[int, float] | None] = (None, None, None), + target_xy: tuple[float, float] | None = None, fps: int = 60, width: int = 640, height: int = 480, @@ -186,10 +230,14 @@ def record_episode_multi_camera( for path in output_paths.values(): path.parent.mkdir(parents=True, exist_ok=True) - state = env.reset(seed=seed) + state = env.reset(seed=seed, target_position=(target_xy[0], target_xy[1], 0.0)) model = state.mj_model data = state.mj_data + _apply_camera_overrides(model, camera_fovy=camera_fovy, camera_xyz=camera_xyz) + + _ensure_offscreen_size(model, width, height) + renderer = mujoco.Renderer(model, width=width, height=height) frames = {cam_id: [] for cam_id in camera_ids} From 748479c046211308c3886c5b5a15becd65615b6c Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Fri, 15 May 2026 01:11:22 +0200 Subject: [PATCH 06/16] fix: updated default params --- scripts/poster_visualisations/render_poster_videos.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/scripts/poster_visualisations/render_poster_videos.py b/scripts/poster_visualisations/render_poster_videos.py index d3cf697..3c418ff 100644 --- a/scripts/poster_visualisations/render_poster_videos.py +++ b/scripts/poster_visualisations/render_poster_videos.py @@ -45,12 +45,12 @@ def main() -> None: parser.add_argument("--follow-camera", type=int, default=1) parser.add_argument("--topdown-camera-x", type=float, default=-3.0) parser.add_argument("--topdown-camera-y", type=float, default=0.0) - parser.add_argument("--topdown-camera-z", type=float, default=None) - parser.add_argument("--topdown-camera-fovy", type=float, default=50.0) + parser.add_argument("--topdown-camera-z", type=float, default=5.0) + parser.add_argument("--topdown-camera-fovy", type=float, default=None) parser.add_argument("--target-x", type=float, default=-6.0) parser.add_argument("--target-y", type=float, default=0.0) - parser.add_argument("--width", type=int, default=1920) - parser.add_argument("--height", type=int, default=1088) + parser.add_argument("--width", type=int, default=2160) + parser.add_argument("--height", type=int, default=960) parser.add_argument("--fps", type=int, default=60) args = parser.parse_args() From 5af820b539b1c7d2a979b56857d5e4b045fae795 Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Fri, 15 May 2026 12:51:12 +0200 Subject: [PATCH 07/16] feat: static path image visualisation --- .../render_static_path_image.py | 268 ++++++++++++++++++ .../render_static_path_image.sh | 10 + 2 files changed, 278 insertions(+) create mode 100644 scripts/poster_visualisations/render_static_path_image.py create mode 100755 scripts/poster_visualisations/render_static_path_image.sh diff --git a/scripts/poster_visualisations/render_static_path_image.py b/scripts/poster_visualisations/render_static_path_image.py new file mode 100644 index 0000000..841a34a --- /dev/null +++ b/scripts/poster_visualisations/render_static_path_image.py @@ -0,0 +1,268 @@ +from __future__ import annotations + +import argparse +from pathlib import Path + +import numpy as np + +from brittle_star_project.evaluation.checkpoint import load_metadata, metadata_to_configs +from brittle_star_project.evaluation.eval_env_builder import build_eval_env +from brittle_star_project.evaluation.rollout import _get_observations, _maybe_clip_action, _target_reached +from brittle_star_project.evaluation.video import _apply_camera_overrides, _ensure_offscreen_size + + +def _enum_value(enum_obj, *names: str) -> int: + for name in names: + if hasattr(enum_obj, name): + return int(getattr(enum_obj, name)) + raise AttributeError(f"Could not find any of {names!r} on {enum_obj!r}") + + +def _hex_to_rgba(hex_color: str, alpha: float) -> np.ndarray: + color = hex_color.lstrip("#") + if len(color) != 6: + raise ValueError(f"Expected a 6-digit hex color, got {hex_color!r}") + red = int(color[0:2], 16) / 255.0 + green = int(color[2:4], 16) / 255.0 + blue = int(color[4:6], 16) / 255.0 + return np.asarray([red, green, blue, float(alpha)], dtype=np.float32) + + +def _cumulative_arc_length(points: np.ndarray) -> np.ndarray: + if len(points) == 0: + return np.zeros((0,), dtype=np.float32) + + deltas = np.diff(points, axis=0) + segment_lengths = np.linalg.norm(deltas, axis=1) + return np.concatenate(([0.0], np.cumsum(segment_lengths))).astype(np.float32) + + +def _sample_along_path(points: np.ndarray, count: int) -> np.ndarray: + if len(points) == 0: + return points + if count <= 1 or len(points) == 1: + return points[[0]] + + arc = _cumulative_arc_length(points) + total = float(arc[-1]) + if total <= 0.0: + return points[[0] * count] + + targets = np.linspace(0.0, total, num=count, dtype=np.float32) + sampled = np.empty((count, points.shape[1]), dtype=np.float32) + for idx, target in enumerate(targets): + upper = int(np.searchsorted(arc, target, side="right")) + lower = max(upper - 1, 0) + if upper >= len(points): + sampled[idx] = points[-1] + continue + + start = points[lower] + end = points[upper] + span = float(arc[upper] - arc[lower]) + if span <= 1e-8: + sampled[idx] = start + continue + + weight = (float(target) - float(arc[lower])) / span + sampled[idx] = start + weight * (end - start) + + return sampled + + +def _append_line(scene, mujoco, start: np.ndarray, end: np.ndarray, rgba: np.ndarray, width: float) -> None: + geom = scene.geoms[scene.ngeom] + mujoco.mjv_connector(geom, mujoco.mjtGeom.mjGEOM_LINE, float(width), start, end) + geom.rgba[:] = rgba + scene.ngeom += 1 + + +def _append_sphere(scene, mujoco, center: np.ndarray, radius: float, rgba: np.ndarray) -> None: + geom = scene.geoms[scene.ngeom] + mujoco.mjv_initGeom( + geom, + mujoco.mjtGeom.mjGEOM_SPHERE, + np.asarray([radius, 0.0, 0.0], dtype=np.float32), + center, + np.eye(3, dtype=np.float32).reshape(-1), + rgba, + ) + scene.ngeom += 1 + + +def _resolve_body_id(model, body_name: str, mujoco) -> int: + body_id = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_BODY, body_name) + if body_id >= 0: + return body_id + + for fallback_name in ("BrittleStarMorphology/central_disk", "central_disk"): + body_id = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_BODY, fallback_name) + if body_id >= 0: + return body_id + + for candidate_body_id in range(1, int(model.nbody)): + name = mujoco.mj_id2name(model, mujoco.mjtObj.mjOBJ_BODY, candidate_body_id) + if name: + return candidate_body_id + + raise ValueError(f"Body '{body_name}' not found in the model") + + +def main() -> None: + parser = argparse.ArgumentParser(description="Render a static path image from a rollout.") + parser.add_argument("model", help="Path to .flax checkpoint") + parser.add_argument("--morphology-override", default=None) + parser.add_argument("--output-path", required=True) + parser.add_argument("--ghost-overlay", action="store_true") + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--max-steps", type=int, default=2000) + parser.add_argument("--body-name", default="BrittleStarMorphology/central_disk") + parser.add_argument("--camera-id", type=int, default=0) + parser.add_argument("--camera-x", type=float, default=-3.0) + parser.add_argument("--camera-y", type=float, default=0.0) + parser.add_argument("--camera-z", type=float, default=5.0) + parser.add_argument("--camera-fovy", type=float, default=None) + parser.add_argument("--target-x", type=float, default=-6.0) + parser.add_argument("--target-y", type=float, default=0.0) + parser.add_argument("--width", type=int, default=2160) + parser.add_argument("--height", type=int, default=960) + parser.add_argument("--frame-stride", type=int, default=5) + parser.add_argument("--line-color", default="#50c4ba") + parser.add_argument("--line-width", type=float, default=2.0) + args = parser.parse_args() + + model_path = Path(args.model) + metadata = load_metadata(model_path, None) + training = metadata_to_configs(metadata) + + bundle = build_eval_env( + model_path=model_path, + training=training, + metadata=metadata, + morphology_override_path=args.morphology_override, + ) + + if (args.target_x is None) != (args.target_y is None): + raise ValueError("target-x and target-y must be provided together") + + target_xy = None + if args.target_x is not None: + target_xy = (float(args.target_x), float(args.target_y)) + + try: + import imageio + import mujoco + except ImportError as e: + raise ImportError( + "Static image rendering requires 'mujoco' and 'imageio'. " + "Please install the evaluation dependencies: `uv pip install .[evaluation]`" + ) from e + + reset_kwargs = {} + if target_xy is not None: + reset_kwargs["target_position"] = (target_xy[0], target_xy[1], 0.0) + + state = bundle.env.reset(seed=args.seed, **reset_kwargs) + model = state.mj_model + data = state.mj_data + + _apply_camera_overrides( + model, + camera_fovy={args.camera_id: float(args.camera_fovy)} if args.camera_fovy is not None else None, + camera_xyz=( + {args.camera_id: float(args.camera_x)} if args.camera_x is not None else None, + {args.camera_id: float(args.camera_y)} if args.camera_y is not None else None, + {args.camera_id: float(args.camera_z)} if args.camera_z is not None else None, + ), + ) + _ensure_offscreen_size(model, args.width, args.height) + + body_id = _resolve_body_id(model, args.body_name, mujoco) + + positions = [] + observations = _get_observations(state) + + for _ in range(int(args.max_steps)): + positions.append(np.asarray(data.xpos[body_id], dtype=np.float32)) + + obs_dict = observations or {} + action = bundle.policy.act(observations=obs_dict) + if bundle.action_mask is not None: + action = action[bundle.action_mask] + action = _maybe_clip_action(action, bundle.action_low, bundle.action_high) + + state = bundle.env.step(state=state, action=action) + data = state.mj_data + observations = _get_observations(state) + + if _target_reached(state=state): + break + + positions_arr = np.vstack(positions) + if len(positions_arr) < 2: + raise ValueError("Need at least two rollout positions to render a path") + + path_points = positions_arr.copy() + path_points[:, 2] += 0.02 + + if target_xy is not None: + target_point = np.asarray([target_xy[0], target_xy[1], path_points[:, 2].min()], dtype=np.float32) + else: + target_point = None + + ghost_count = max(1, len(path_points) // max(int(args.frame_stride), 1)) // 2 + ghost_points = _sample_along_path(path_points, ghost_count) + + line_rgba = _hex_to_rgba(args.line_color, 0.92) + ghost_base_rgba = _hex_to_rgba(args.line_color, 0.10) + target_rgba = np.asarray([0.90, 0.12, 0.12, 0.95], dtype=np.float32) + + ctx = mujoco.GLContext(args.width, args.height) + ctx.make_current() + try: + catmask = _enum_value(mujoco.mjtCatBit, "mjCAT_ALL") + camera_type = _enum_value(mujoco.mjtCamera, "mjCAMERA_FIXED") + font_scale = _enum_value(mujoco.mjtFontScale, "mjFONTSCALE_100") + + maxgeom = int(model.ngeom + len(path_points) + len(ghost_points) + 8) + scene = mujoco.MjvScene(model, maxgeom=maxgeom) + option = mujoco.MjvOption() + perturb = mujoco.MjvPerturb() + camera = mujoco.MjvCamera() + mujoco.mjv_defaultOption(option) + mujoco.mjv_defaultPerturb(perturb) + mujoco.mjv_defaultCamera(camera) + camera.type = camera_type + camera.fixedcamid = int(args.camera_id) + if hasattr(camera, "trackbodyid"): + camera.trackbodyid = -1 + + context = mujoco.MjrContext(model, font_scale) + viewport = mujoco.MjrRect(0, 0, args.width, args.height) + + mujoco.mjv_updateScene(model, data, option, perturb, camera, catmask, scene) + + for idx, ghost_point in enumerate(ghost_points): + alpha = 0.10 + 0.70 * (idx / max(len(ghost_points) - 1, 1)) + ghost_color = ghost_base_rgba.copy() + ghost_color[3] = float(alpha) + _append_sphere(scene, mujoco, ghost_point, 0.03, ghost_color) + + if target_point is not None: + _append_sphere(scene, mujoco, target_point, 0.025, target_rgba) + + rgb = np.empty((args.height, args.width, 3), dtype=np.uint8) + depth = np.empty((args.height, args.width), dtype=np.float32) + mujoco.mjr_render(viewport, scene, context) + mujoco.mjr_readPixels(rgb, depth, viewport, context) + imageio.imwrite(str(output_path := Path(args.output_path)), np.flipud(rgb)) + + context.free() + finally: + ctx.free() + + bundle.env.close() + + +if __name__ == "__main__": + main() diff --git a/scripts/poster_visualisations/render_static_path_image.sh b/scripts/poster_visualisations/render_static_path_image.sh new file mode 100755 index 0000000..ab21ca0 --- /dev/null +++ b/scripts/poster_visualisations/render_static_path_image.sh @@ -0,0 +1,10 @@ +#!/usr/bin/env bash + +# Static path image + optional ghost render + +path=$1 # path to .flax model with metadata.yaml alongside it + +uv run scripts/poster_visualisations/render_static_path_image.py \ + "$path" \ + --output-path vids/poster/5arms/centralized/path.png \ + --ghost-overlay From 91cfa92834195a75f31c59145a688202a30f2a2d Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Fri, 15 May 2026 15:33:01 +0200 Subject: [PATCH 08/16] feat: implemented frame streaming to prevent OOM errors --- src/brittle_star_project/evaluation/video.py | 62 ++++++++++---------- 1 file changed, 31 insertions(+), 31 deletions(-) diff --git a/src/brittle_star_project/evaluation/video.py b/src/brittle_star_project/evaluation/video.py index a7a629d..21017a3 100644 --- a/src/brittle_star_project/evaluation/video.py +++ b/src/brittle_star_project/evaluation/video.py @@ -239,7 +239,7 @@ def record_episode_multi_camera( _ensure_offscreen_size(model, width, height) renderer = mujoco.Renderer(model, width=width, height=height) - frames = {cam_id: [] for cam_id in camera_ids} + writers = {cam_id: imageio.get_writer(str(path), fps=fps) for cam_id, path in output_paths.items()} ep_return = 0.0 observations = _get_observations(state) @@ -248,38 +248,38 @@ def record_episode_multi_camera( reached_target = _target_reached(state=state) steps = 0 - for _ in range(int(max_steps)): + try: + for _ in range(int(max_steps)): + for cam_id in camera_ids: + renderer.update_scene(data, camera=cam_id) + writers[cam_id].append_data(renderer.render()) + + 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 + for cam_id in camera_ids: renderer.update_scene(data, camera=cam_id) - frames[cam_id].append(renderer.render()) - - 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 - - for cam_id in camera_ids: - renderer.update_scene(data, camera=cam_id) - frames[cam_id].append(renderer.render()) - - renderer.close() - - for cam_id, path in output_paths.items(): - imageio.mimsave(str(path), frames[cam_id], fps=fps) + writers[cam_id].append_data(renderer.render()) + finally: + renderer.close() + for writer in writers.values(): + writer.close() final_dist = _get_xy_distance_to_target(observations) if observations else None return EpisodeResult( From beb64ca5428e40b767e92c4a3b6ae875a1ab4a09 Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Fri, 15 May 2026 15:49:40 +0200 Subject: [PATCH 09/16] fix: cleanup of ghost brittle stars --- .../render_static_path_image.py | 96 ++++--------------- 1 file changed, 19 insertions(+), 77 deletions(-) diff --git a/scripts/poster_visualisations/render_static_path_image.py b/scripts/poster_visualisations/render_static_path_image.py index 841a34a..374f31f 100644 --- a/scripts/poster_visualisations/render_static_path_image.py +++ b/scripts/poster_visualisations/render_static_path_image.py @@ -7,7 +7,11 @@ import numpy as np from brittle_star_project.evaluation.checkpoint import load_metadata, metadata_to_configs from brittle_star_project.evaluation.eval_env_builder import build_eval_env -from brittle_star_project.evaluation.rollout import _get_observations, _maybe_clip_action, _target_reached +from brittle_star_project.evaluation.rollout import ( + _get_observations, + _maybe_clip_action, + _target_reached, +) from brittle_star_project.evaluation.video import _apply_camera_overrides, _ensure_offscreen_size @@ -28,55 +32,6 @@ def _hex_to_rgba(hex_color: str, alpha: float) -> np.ndarray: return np.asarray([red, green, blue, float(alpha)], dtype=np.float32) -def _cumulative_arc_length(points: np.ndarray) -> np.ndarray: - if len(points) == 0: - return np.zeros((0,), dtype=np.float32) - - deltas = np.diff(points, axis=0) - segment_lengths = np.linalg.norm(deltas, axis=1) - return np.concatenate(([0.0], np.cumsum(segment_lengths))).astype(np.float32) - - -def _sample_along_path(points: np.ndarray, count: int) -> np.ndarray: - if len(points) == 0: - return points - if count <= 1 or len(points) == 1: - return points[[0]] - - arc = _cumulative_arc_length(points) - total = float(arc[-1]) - if total <= 0.0: - return points[[0] * count] - - targets = np.linspace(0.0, total, num=count, dtype=np.float32) - sampled = np.empty((count, points.shape[1]), dtype=np.float32) - for idx, target in enumerate(targets): - upper = int(np.searchsorted(arc, target, side="right")) - lower = max(upper - 1, 0) - if upper >= len(points): - sampled[idx] = points[-1] - continue - - start = points[lower] - end = points[upper] - span = float(arc[upper] - arc[lower]) - if span <= 1e-8: - sampled[idx] = start - continue - - weight = (float(target) - float(arc[lower])) / span - sampled[idx] = start + weight * (end - start) - - return sampled - - -def _append_line(scene, mujoco, start: np.ndarray, end: np.ndarray, rgba: np.ndarray, width: float) -> None: - geom = scene.geoms[scene.ngeom] - mujoco.mjv_connector(geom, mujoco.mjtGeom.mjGEOM_LINE, float(width), start, end) - geom.rgba[:] = rgba - scene.ngeom += 1 - - def _append_sphere(scene, mujoco, center: np.ndarray, radius: float, rgba: np.ndarray) -> None: geom = scene.geoms[scene.ngeom] mujoco.mjv_initGeom( @@ -113,9 +68,8 @@ def main() -> None: parser.add_argument("model", help="Path to .flax checkpoint") parser.add_argument("--morphology-override", default=None) parser.add_argument("--output-path", required=True) - parser.add_argument("--ghost-overlay", action="store_true") parser.add_argument("--seed", type=int, default=0) - parser.add_argument("--max-steps", type=int, default=2000) + parser.add_argument("--max-steps", type=int, default=5000) parser.add_argument("--body-name", default="BrittleStarMorphology/central_disk") parser.add_argument("--camera-id", type=int, default=0) parser.add_argument("--camera-x", type=float, default=-3.0) @@ -127,8 +81,7 @@ def main() -> None: parser.add_argument("--width", type=int, default=2160) parser.add_argument("--height", type=int, default=960) parser.add_argument("--frame-stride", type=int, default=5) - parser.add_argument("--line-color", default="#50c4ba") - parser.add_argument("--line-width", type=float, default=2.0) + parser.add_argument("--path-color", default="#50c4ba") args = parser.parse_args() model_path = Path(args.model) @@ -168,7 +121,9 @@ def main() -> None: _apply_camera_overrides( model, - camera_fovy={args.camera_id: float(args.camera_fovy)} if args.camera_fovy is not None else None, + camera_fovy={args.camera_id: float(args.camera_fovy)} + if args.camera_fovy is not None + else None, camera_xyz=( {args.camera_id: float(args.camera_x)} if args.camera_x is not None else None, {args.camera_id: float(args.camera_y)} if args.camera_y is not None else None, @@ -203,19 +158,11 @@ def main() -> None: raise ValueError("Need at least two rollout positions to render a path") path_points = positions_arr.copy() - path_points[:, 2] += 0.02 + path_points[:, 2] -= 0.02 - if target_xy is not None: - target_point = np.asarray([target_xy[0], target_xy[1], path_points[:, 2].min()], dtype=np.float32) - else: - target_point = None - - ghost_count = max(1, len(path_points) // max(int(args.frame_stride), 1)) // 2 - ghost_points = _sample_along_path(path_points, ghost_count) - - line_rgba = _hex_to_rgba(args.line_color, 0.92) - ghost_base_rgba = _hex_to_rgba(args.line_color, 0.10) - target_rgba = np.asarray([0.90, 0.12, 0.12, 0.95], dtype=np.float32) + path_step = max(1, int(args.frame_stride)) * 4 + path_points_visible = path_points[::path_step] + path_rgba = _hex_to_rgba(args.path_color, 0.92) ctx = mujoco.GLContext(args.width, args.height) ctx.make_current() @@ -224,7 +171,7 @@ def main() -> None: camera_type = _enum_value(mujoco.mjtCamera, "mjCAMERA_FIXED") font_scale = _enum_value(mujoco.mjtFontScale, "mjFONTSCALE_100") - maxgeom = int(model.ngeom + len(path_points) + len(ghost_points) + 8) + maxgeom = int(model.ngeom + len(path_points_visible) + 8) scene = mujoco.MjvScene(model, maxgeom=maxgeom) option = mujoco.MjvOption() perturb = mujoco.MjvPerturb() @@ -242,20 +189,15 @@ def main() -> None: mujoco.mjv_updateScene(model, data, option, perturb, camera, catmask, scene) - for idx, ghost_point in enumerate(ghost_points): - alpha = 0.10 + 0.70 * (idx / max(len(ghost_points) - 1, 1)) - ghost_color = ghost_base_rgba.copy() - ghost_color[3] = float(alpha) - _append_sphere(scene, mujoco, ghost_point, 0.03, ghost_color) - - if target_point is not None: - _append_sphere(scene, mujoco, target_point, 0.025, target_rgba) + for idx, path_point in enumerate(path_points_visible): + path_rgba[3] = 0.10 + 0.70 * (idx / max(len(path_points_visible) - 1, 1)) + _append_sphere(scene, mujoco, path_point, 0.03, path_rgba) rgb = np.empty((args.height, args.width, 3), dtype=np.uint8) depth = np.empty((args.height, args.width), dtype=np.float32) mujoco.mjr_render(viewport, scene, context) mujoco.mjr_readPixels(rgb, depth, viewport, context) - imageio.imwrite(str(output_path := Path(args.output_path)), np.flipud(rgb)) + imageio.imwrite(args.output_path, np.flipud(rgb)) context.free() finally: From 3968bef66c896bdc68a47ad44ce46fda6faa5ac6 Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Fri, 15 May 2026 17:32:04 +0200 Subject: [PATCH 10/16] feat: customizable path & robot color --- .../render_static_path_image.py | 27 ++++++++++++++++++- 1 file changed, 26 insertions(+), 1 deletion(-) diff --git a/scripts/poster_visualisations/render_static_path_image.py b/scripts/poster_visualisations/render_static_path_image.py index 374f31f..c9a4416 100644 --- a/scripts/poster_visualisations/render_static_path_image.py +++ b/scripts/poster_visualisations/render_static_path_image.py @@ -81,7 +81,12 @@ def main() -> None: parser.add_argument("--width", type=int, default=2160) parser.add_argument("--height", type=int, default=960) parser.add_argument("--frame-stride", type=int, default=5) - parser.add_argument("--path-color", default="#50c4ba") + parser.add_argument("--path-color", default="#444444") + parser.add_argument( + "--robot-color", + default="#444444", + help="Hex color for brittle star robot (e.g. #ff0000)", + ) args = parser.parse_args() model_path = Path(args.model) @@ -134,6 +139,26 @@ def main() -> None: body_id = _resolve_body_id(model, args.body_name, mujoco) + # Optionally override robot color by recoloring geoms belonging to the robot's body subtree. + if args.robot_color is not None: + robot_rgba = _hex_to_rgba(args.robot_color, 1.0) + # Collect body IDs in the subtree rooted at `body_id` by walking parent links. + nbody = int(model.nbody) + body_parent = model.body_parentid + robot_body_ids = set([int(body_id)]) + for i in range(1, nbody): + cur = int(i) + # walk up until root (0) or until we hit the robot root + while cur not in (-1, 0, int(body_id)): + cur = int(body_parent[cur]) + if cur == int(body_id): + robot_body_ids.add(i) + + # Recolor geoms whose body id is in the robot subtree + for g in range(int(model.ngeom)): + if int(model.geom_bodyid[g]) in robot_body_ids: + model.geom_rgba[g][:] = robot_rgba + positions = [] observations = _get_observations(state) From ec4e069e2f4b4bc55ec6a03a6bf5ea0fb21b6d3d Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Fri, 15 May 2026 17:32:04 +0200 Subject: [PATCH 11/16] feat: customizable path & robot color --- .../render_static_path_image.py | 33 +++++++++++++++++-- 1 file changed, 30 insertions(+), 3 deletions(-) diff --git a/scripts/poster_visualisations/render_static_path_image.py b/scripts/poster_visualisations/render_static_path_image.py index 374f31f..5db1472 100644 --- a/scripts/poster_visualisations/render_static_path_image.py +++ b/scripts/poster_visualisations/render_static_path_image.py @@ -74,14 +74,21 @@ def main() -> None: parser.add_argument("--camera-id", type=int, default=0) parser.add_argument("--camera-x", type=float, default=-3.0) parser.add_argument("--camera-y", type=float, default=0.0) - parser.add_argument("--camera-z", type=float, default=5.0) + parser.add_argument("--camera-z", type=float, default=4.5) parser.add_argument("--camera-fovy", type=float, default=None) parser.add_argument("--target-x", type=float, default=-6.0) parser.add_argument("--target-y", type=float, default=0.0) parser.add_argument("--width", type=int, default=2160) parser.add_argument("--height", type=int, default=960) parser.add_argument("--frame-stride", type=int, default=5) - parser.add_argument("--path-color", default="#50c4ba") + parser.add_argument( + "--path-color", default="#FA9F42" + ) # ring = #888888, centralized = #2B4162, fully connected = FA9F42 + parser.add_argument( + "--robot-color", + default="#FA9F42", + help="Hex color for brittle star robot (e.g. #ff0000)", + ) args = parser.parse_args() model_path = Path(args.model) @@ -134,6 +141,26 @@ def main() -> None: body_id = _resolve_body_id(model, args.body_name, mujoco) + # Optionally override robot color by recoloring geoms belonging to the robot's body subtree. + if args.robot_color is not None: + robot_rgba = _hex_to_rgba(args.robot_color, 1.0) + # Collect body IDs in the subtree rooted at `body_id` by walking parent links. + nbody = int(model.nbody) + body_parent = model.body_parentid + robot_body_ids = set([int(body_id)]) + for i in range(1, nbody): + cur = int(i) + # walk up until root (0) or until we hit the robot root + while cur not in (-1, 0, int(body_id)): + cur = int(body_parent[cur]) + if cur == int(body_id): + robot_body_ids.add(i) + + # Recolor geoms whose body id is in the robot subtree + for g in range(int(model.ngeom)): + if int(model.geom_bodyid[g]) in robot_body_ids: + model.geom_rgba[g][:] = robot_rgba + positions = [] observations = _get_observations(state) @@ -160,7 +187,7 @@ def main() -> None: path_points = positions_arr.copy() path_points[:, 2] -= 0.02 - path_step = max(1, int(args.frame_stride)) * 4 + path_step = max(1, int(args.frame_stride)) * 3 path_points_visible = path_points[::path_step] path_rgba = _hex_to_rgba(args.path_color, 0.92) From c68a7ddf26b8cb87337416b9c2574bd31eedc186 Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Fri, 15 May 2026 19:50:15 +0200 Subject: [PATCH 12/16] cleanup: set default params & code deduplication --- .../render_poster_videos.py | 14 ++++-- .../render_static_path_image.py | 42 ++++------------ src/brittle_star_project/evaluation/video.py | 50 +++++++++++++++++-- 3 files changed, 66 insertions(+), 40 deletions(-) diff --git a/scripts/poster_visualisations/render_poster_videos.py b/scripts/poster_visualisations/render_poster_videos.py index 3c418ff..8a73638 100644 --- a/scripts/poster_visualisations/render_poster_videos.py +++ b/scripts/poster_visualisations/render_poster_videos.py @@ -40,18 +40,23 @@ def main() -> None: ) parser.add_argument("--output-root", default="vids/poster") parser.add_argument("--seed", type=int, default=0) - parser.add_argument("--max-steps", type=int, default=2000) + parser.add_argument("--max-steps", type=int, default=5000) parser.add_argument("--topdown-camera", type=int, default=0) parser.add_argument("--follow-camera", type=int, default=1) parser.add_argument("--topdown-camera-x", type=float, default=-3.0) parser.add_argument("--topdown-camera-y", type=float, default=0.0) - parser.add_argument("--topdown-camera-z", type=float, default=5.0) + parser.add_argument("--topdown-camera-z", type=float, default=4.5) parser.add_argument("--topdown-camera-fovy", type=float, default=None) parser.add_argument("--target-x", type=float, default=-6.0) parser.add_argument("--target-y", type=float, default=0.0) parser.add_argument("--width", type=int, default=2160) parser.add_argument("--height", type=int, default=960) parser.add_argument("--fps", type=int, default=60) + parser.add_argument( + "--robot-color", + default="#2B4162", # ring = #888888, centralized = #2B4162, fully connected = FA9F42 + help="Hex color for the brittle star robot", + ) args = parser.parse_args() if (args.target_x is None) != (args.target_y is None): @@ -117,14 +122,13 @@ def main() -> None: camera_fovy=camera_fovy, camera_xyz=camera_xyz, target_xy=target_xy, + robot_color=args.robot_color, width=args.width, height=args.height, fps=args.fps, ) - final_dist = ( - "n/a" if result.final_xy_dist is None else f"{result.final_xy_dist:.3f}" - ) + final_dist = "n/a" if result.final_xy_dist is None else f"{result.final_xy_dist:.3f}" print( f"{arms_dir}/{arch_dir}: return={result.return_:.6f}, len={result.length}, " f"target_reached={result.reached_target}, final_xy_dist={final_dist}" diff --git a/scripts/poster_visualisations/render_static_path_image.py b/scripts/poster_visualisations/render_static_path_image.py index 5db1472..24e0b53 100644 --- a/scripts/poster_visualisations/render_static_path_image.py +++ b/scripts/poster_visualisations/render_static_path_image.py @@ -12,7 +12,11 @@ from brittle_star_project.evaluation.rollout import ( _maybe_clip_action, _target_reached, ) -from brittle_star_project.evaluation.video import _apply_camera_overrides, _ensure_offscreen_size +from brittle_star_project.evaluation.video import ( + _apply_camera_overrides, + _ensure_offscreen_size, + hex_to_rgba, +) def _enum_value(enum_obj, *names: str) -> int: @@ -22,16 +26,6 @@ def _enum_value(enum_obj, *names: str) -> int: raise AttributeError(f"Could not find any of {names!r} on {enum_obj!r}") -def _hex_to_rgba(hex_color: str, alpha: float) -> np.ndarray: - color = hex_color.lstrip("#") - if len(color) != 6: - raise ValueError(f"Expected a 6-digit hex color, got {hex_color!r}") - red = int(color[0:2], 16) / 255.0 - green = int(color[2:4], 16) / 255.0 - blue = int(color[4:6], 16) / 255.0 - return np.asarray([red, green, blue, float(alpha)], dtype=np.float32) - - def _append_sphere(scene, mujoco, center: np.ndarray, radius: float, rgba: np.ndarray) -> None: geom = scene.geoms[scene.ngeom] mujoco.mjv_initGeom( @@ -45,24 +39,6 @@ def _append_sphere(scene, mujoco, center: np.ndarray, radius: float, rgba: np.nd scene.ngeom += 1 -def _resolve_body_id(model, body_name: str, mujoco) -> int: - body_id = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_BODY, body_name) - if body_id >= 0: - return body_id - - for fallback_name in ("BrittleStarMorphology/central_disk", "central_disk"): - body_id = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_BODY, fallback_name) - if body_id >= 0: - return body_id - - for candidate_body_id in range(1, int(model.nbody)): - name = mujoco.mj_id2name(model, mujoco.mjtObj.mjOBJ_BODY, candidate_body_id) - if name: - return candidate_body_id - - raise ValueError(f"Body '{body_name}' not found in the model") - - def main() -> None: parser = argparse.ArgumentParser(description="Render a static path image from a rollout.") parser.add_argument("model", help="Path to .flax checkpoint") @@ -139,11 +115,13 @@ def main() -> None: ) _ensure_offscreen_size(model, args.width, args.height) - body_id = _resolve_body_id(model, args.body_name, mujoco) + body_id = mujoco.mj_name2id( + model, mujoco.mjtObj.mjOBJ_BODY, "BrittleStarMorphology/central_disk" + ) # Optionally override robot color by recoloring geoms belonging to the robot's body subtree. if args.robot_color is not None: - robot_rgba = _hex_to_rgba(args.robot_color, 1.0) + robot_rgba = hex_to_rgba(args.robot_color, 1.0) # Collect body IDs in the subtree rooted at `body_id` by walking parent links. nbody = int(model.nbody) body_parent = model.body_parentid @@ -189,7 +167,7 @@ def main() -> None: path_step = max(1, int(args.frame_stride)) * 3 path_points_visible = path_points[::path_step] - path_rgba = _hex_to_rgba(args.path_color, 0.92) + path_rgba = hex_to_rgba(args.path_color, 0.92) ctx = mujoco.GLContext(args.width, args.height) ctx.make_current() diff --git a/src/brittle_star_project/evaluation/video.py b/src/brittle_star_project/evaluation/video.py index 21017a3..cf2e5c2 100644 --- a/src/brittle_star_project/evaluation/video.py +++ b/src/brittle_star_project/evaluation/video.py @@ -37,7 +37,11 @@ def _apply_camera_overrides( model, *, camera_fovy: dict[int, float] | None = None, - camera_xyz: tuple[dict[int, float] | None, dict[int, float] | None, dict[int, float] | None] = (None, None, None), + camera_xyz: tuple[ + dict[int, float] | None, + dict[int, float] | None, + dict[int, float] | None, + ] = (None, None, None), ) -> None: if not camera_fovy and not (camera_xyz[0] or camera_xyz[1] or camera_xyz[2]): return @@ -63,6 +67,17 @@ def _apply_camera_overrides( raise ValueError(f"Camera id {cam_id} is out of range") model.cam_pos[cam_id][2] = float(z) + +def hex_to_rgba(hex_color: str, alpha: float) -> np.ndarray: + color = hex_color.lstrip("#") + if len(color) != 6: + raise ValueError(f"Expected a 6-digit hex color, got {hex_color!r}") + red = int(color[0:2], 16) / 255.0 + green = int(color[2:4], 16) / 255.0 + blue = int(color[4:6], 16) / 255.0 + return np.asarray([red, green, blue, float(alpha)], dtype=np.float32) + + def save_evaluation_metadata( eval_dir: Path, *, @@ -202,8 +217,13 @@ def record_episode_multi_camera( action_mask: np.ndarray | None = None, camera_ids: list[int] | None = None, camera_fovy: dict[int, float] | None = None, - camera_xyz: tuple[dict[int, float] | None, dict[int, float] | None, dict[int, float] | None] = (None, None, None), + camera_xyz: tuple[ + dict[int, float] | None, + dict[int, float] | None, + dict[int, float] | None, + ] = (None, None, None), target_xy: tuple[float, float] | None = None, + robot_color: str | None = None, fps: int = 60, width: int = 640, height: int = 480, @@ -236,10 +256,34 @@ def record_episode_multi_camera( _apply_camera_overrides(model, camera_fovy=camera_fovy, camera_xyz=camera_xyz) + if robot_color is not None: + robot_body_id = mujoco.mj_name2id( + model, mujoco.mjtObj.mjOBJ_BODY, "BrittleStarMorphology/central_disk" + ) + if robot_body_id < 0: + raise ValueError("Body 'BrittleStarMorphology/central_disk' not found in the model") + + robot_rgba = hex_to_rgba(robot_color, 1.0) + body_parent = model.body_parentid + robot_body_ids = {int(robot_body_id)} + + for body_id in range(1, int(model.nbody)): + current_body_id = int(body_id) + while current_body_id not in (-1, 0, int(robot_body_id)): + current_body_id = int(body_parent[current_body_id]) + if current_body_id == int(robot_body_id): + robot_body_ids.add(body_id) + + for geom_id in range(int(model.ngeom)): + if int(model.geom_bodyid[geom_id]) in robot_body_ids: + model.geom_rgba[geom_id][:] = robot_rgba + _ensure_offscreen_size(model, width, height) renderer = mujoco.Renderer(model, width=width, height=height) - writers = {cam_id: imageio.get_writer(str(path), fps=fps) for cam_id, path in output_paths.items()} + writers = { + cam_id: imageio.get_writer(str(path), fps=fps) for cam_id, path in output_paths.items() + } ep_return = 0.0 observations = _get_observations(state) From 3eb6e0fb457c8b97c32080d0a7207088464e413a Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Fri, 15 May 2026 19:52:06 +0200 Subject: [PATCH 13/16] fix: function name error --- scripts/poster_visualisations/render_static_path_image.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/scripts/poster_visualisations/render_static_path_image.py b/scripts/poster_visualisations/render_static_path_image.py index 2a992aa..2057133 100644 --- a/scripts/poster_visualisations/render_static_path_image.py +++ b/scripts/poster_visualisations/render_static_path_image.py @@ -141,7 +141,7 @@ def main() -> None: # Optionally override robot color by recoloring geoms belonging to the robot's body subtree. if args.robot_color is not None: - robot_rgba = _hex_to_rgba(args.robot_color, 1.0) + robot_rgba = hex_to_rgba(args.robot_color, 1.0) # Collect body IDs in the subtree rooted at `body_id` by walking parent links. nbody = int(model.nbody) body_parent = model.body_parentid From 3489448a3db95312d9856add2e2ed4976633e5e5 Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Sat, 16 May 2026 01:50:35 +0200 Subject: [PATCH 14/16] feat: color map for archs --- .../poster_visualisations/render_poster_videos.py | 13 +++++++++++-- .../render_static_path_image.py | 12 +++++++++--- 2 files changed, 20 insertions(+), 5 deletions(-) diff --git a/scripts/poster_visualisations/render_poster_videos.py b/scripts/poster_visualisations/render_poster_videos.py index 8a73638..eca60fc 100644 --- a/scripts/poster_visualisations/render_poster_videos.py +++ b/scripts/poster_visualisations/render_poster_videos.py @@ -14,6 +14,12 @@ _ARCH_DIR_MAP = { "SEGMENT": "segment", } +ROBOT_COLOR_MAP = { + "CENTRALIZED": "#0D567C", # Blue + "FULLY_CONNECTED": "#8C0E0F", # Reddish + "RING": "#FCB304", # Pale Yellow +} + def _arch_dir(name: str) -> str: return _ARCH_DIR_MAP.get(name, name.lower()) @@ -54,7 +60,7 @@ def main() -> None: parser.add_argument("--fps", type=int, default=60) parser.add_argument( "--robot-color", - default="#2B4162", # ring = #888888, centralized = #2B4162, fully connected = FA9F42 + default="#2B4162", help="Hex color for the brittle star robot", ) args = parser.parse_args() @@ -109,6 +115,9 @@ def main() -> None: args.follow_camera: out_dir / "follow.mp4", } + print(bundle.architecture) + color = ROBOT_COLOR_MAP.get(bundle.architecture, args.robot_color) + result = record_episode_multi_camera( env=bundle.env, policy=bundle.policy, @@ -122,7 +131,7 @@ def main() -> None: camera_fovy=camera_fovy, camera_xyz=camera_xyz, target_xy=target_xy, - robot_color=args.robot_color, + robot_color=color, width=args.width, height=args.height, fps=args.fps, diff --git a/scripts/poster_visualisations/render_static_path_image.py b/scripts/poster_visualisations/render_static_path_image.py index 2057133..17fb4b5 100644 --- a/scripts/poster_visualisations/render_static_path_image.py +++ b/scripts/poster_visualisations/render_static_path_image.py @@ -18,6 +18,12 @@ from brittle_star_project.evaluation.video import ( hex_to_rgba, ) +ROBOT_COLOR_MAP = { + "CENTRALIZED": "#0D567C", # Blue + "FULLY_CONNECTED": "#8C0E0F", # Reddish + "RING": "#FCB304", # Pale Yellow +} + def _enum_value(enum_obj, *names: str) -> int: for name in names: @@ -119,9 +125,10 @@ def main() -> None: model, mujoco.mjtObj.mjOBJ_BODY, "BrittleStarMorphology/central_disk" ) + robot_rgba = hex_to_rgba(ROBOT_COLOR_MAP.get(bundle.architecture, args.robot_color), 1.0) + # Optionally override robot color by recoloring geoms belonging to the robot's body subtree. if args.robot_color is not None: - robot_rgba = hex_to_rgba(args.robot_color, 1.0) # Collect body IDs in the subtree rooted at `body_id` by walking parent links. nbody = int(model.nbody) body_parent = model.body_parentid @@ -141,7 +148,6 @@ def main() -> None: # Optionally override robot color by recoloring geoms belonging to the robot's body subtree. if args.robot_color is not None: - robot_rgba = hex_to_rgba(args.robot_color, 1.0) # Collect body IDs in the subtree rooted at `body_id` by walking parent links. nbody = int(model.nbody) body_parent = model.body_parentid @@ -187,7 +193,7 @@ def main() -> None: path_step = max(1, int(args.frame_stride)) * 3 path_points_visible = path_points[::path_step] - path_rgba = hex_to_rgba(args.path_color, 0.92) + path_rgba = hex_to_rgba(ROBOT_COLOR_MAP.get(bundle.architecture, args.path_color), 0.92) ctx = mujoco.GLContext(args.width, args.height) ctx.make_current() From 4b25e1c57027bea99d1e3db197a7bde4b6d328ce Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Sat, 16 May 2026 18:29:14 +0200 Subject: [PATCH 15/16] fix: adapted default frame stride --- scripts/poster_visualisations/render_static_path_image.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/scripts/poster_visualisations/render_static_path_image.py b/scripts/poster_visualisations/render_static_path_image.py index 17fb4b5..8229f52 100644 --- a/scripts/poster_visualisations/render_static_path_image.py +++ b/scripts/poster_visualisations/render_static_path_image.py @@ -62,7 +62,7 @@ def main() -> None: parser.add_argument("--target-y", type=float, default=0.0) parser.add_argument("--width", type=int, default=2160) parser.add_argument("--height", type=int, default=960) - parser.add_argument("--frame-stride", type=int, default=5) + parser.add_argument("--frame-stride", type=int, default=15) parser.add_argument( "--path-color", default="#FA9F42" ) # ring = #888888, centralized = #2B4162, fully connected = FA9F42 @@ -191,7 +191,7 @@ def main() -> None: path_points = positions_arr.copy() path_points[:, 2] -= 0.02 - path_step = max(1, int(args.frame_stride)) * 3 + path_step = max(1, int(args.frame_stride)) path_points_visible = path_points[::path_step] path_rgba = hex_to_rgba(ROBOT_COLOR_MAP.get(bundle.architecture, args.path_color), 0.92) From 1e7f8c1880976491bccc5a9b749203c755e266cd Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Tue, 19 May 2026 17:57:03 +0200 Subject: [PATCH 16/16] fix: removed unnecessary 'int' type before value --- configs/simulation/default.yaml | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/configs/simulation/default.yaml b/configs/simulation/default.yaml index 5bd448a..1ebc263 100644 --- a/configs/simulation/default.yaml +++ b/configs/simulation/default.yaml @@ -21,9 +21,9 @@ video_output_path: null # Camera ID to use for video recording (1 is usually the close-up camera) camera_id: 1 -video_width: int = 640 -video_height: int = 480 -video_fps: int = 60 +video_width: 640 +video_height: 80 +video_fps: 60 # Optional override for the metadata YAML file path. # If null, the script looks for `_metadata.yaml` alongside the model_path.