From dd0e5b2d09c0af670a1058e8b0dc23d0b77ca686 Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Thu, 14 May 2026 12:13:34 +0200 Subject: [PATCH 01/25] 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/25] 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/25] 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/25] 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 becd6798e948ca6f08630786fdbdcce7f6b95ef0 Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Thu, 14 May 2026 14:30:24 +0200 Subject: [PATCH 05/25] feat: added script to download all runs in a given project --- scripts/download_wandb_project.py | 87 +++++++++++++++++++++++++++++++ 1 file changed, 87 insertions(+) create mode 100644 scripts/download_wandb_project.py diff --git a/scripts/download_wandb_project.py b/scripts/download_wandb_project.py new file mode 100644 index 0000000..7c04ad1 --- /dev/null +++ b/scripts/download_wandb_project.py @@ -0,0 +1,87 @@ +from pathlib import Path +from concurrent.futures import ThreadPoolExecutor, as_completed +import threading +import wandb +import argparse + +# tune these depending on network / W&B limits +MAX_RUN_WORKERS = 8 +MAX_FILE_WORKERS = 16 + +api = wandb.Api() + +print_lock = threading.Lock() + + +def safe_print(*args, **kwargs): + with print_lock: + print(*args, **kwargs) + + +def download_file(file, run_dir): + target = run_dir / file.name + try: + # skip existing files + if target.exists(): + return f"SKIP {target}" + + target.parent.mkdir(parents=True, exist_ok=True) + + file.download(root=run_dir, replace=False) + + return f"DONE {target}" + + except Exception as e: + return f"FAIL {target}: {e}" + + +def download_run(run, root): + run_dir = root / f"{run.name}" + run_dir.mkdir(parents=True, exist_ok=True) + + safe_print(f"\n=== {run.name} ({run.id}) ===") + + files = list(run.files()) + + with ThreadPoolExecutor(max_workers=MAX_FILE_WORKERS) as executor: + futures = [executor.submit(download_file, file, run_dir) for file in files] + + for future in as_completed(futures): + safe_print(future.result()) + + # OPTIONAL: download artifacts too + # for artifact in run.logged_artifacts(): + # artifact_dir = run_dir / "artifacts" / artifact.name + # artifact.download(root=artifact_dir) + + safe_print(f"Finished {run.name}") + + +def main(entity: str, project: str, root: Path): + root.mkdir(exist_ok=True) + + runs = list(api.runs(f"{entity}/{project}")) + + safe_print(f"Found {len(runs)} runs") + + with ThreadPoolExecutor(max_workers=MAX_RUN_WORKERS) as executor: + futures = [executor.submit(download_run, run, root) for run in runs] + + for future in as_completed(futures): + try: + future.result() + except Exception as e: + safe_print("RUN FAILED:", e) + + safe_print("\nAll downloads complete.") + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument("--entity", type=str, default="SEL3-2026-Groep-4") + parser.add_argument("--project", type=str, required=True) + parser.add_argument("--root", type=str, default="runs") + args = parser.parse_args() + + root = Path(args.root) + main(entity=args.entity, project=args.project, root=root) From a125a0948fe539d92fcaa7a1a5f934411255201f Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Thu, 14 May 2026 18:19:48 +0200 Subject: [PATCH 06/25] 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 03b1b3397490038281f3aa04afd92ca64a7eefdd Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Thu, 14 May 2026 21:42:50 +0200 Subject: [PATCH 07/25] feat(downloader): added downloading of artifacts + moved script to tool directory --- scripts/plots/analyze_convergence.py | 30 +++++--- scripts/{ => tools}/download_wandb_project.py | 69 +++++++++++++++++-- 2 files changed, 84 insertions(+), 15 deletions(-) rename scripts/{ => tools}/download_wandb_project.py (51%) diff --git a/scripts/plots/analyze_convergence.py b/scripts/plots/analyze_convergence.py index 2612bb9..08f905d 100644 --- a/scripts/plots/analyze_convergence.py +++ b/scripts/plots/analyze_convergence.py @@ -45,19 +45,22 @@ class Columns(str, Enum): """Column names expected in every evaluation CSV.""" ARCH = "architecture" - TIMESTEPS = "total_trained_timesteps" - REWARD = "accumulated_reward" + TIMESTEPS = "trained_timesteps" + REWARD = "eval_return" VELOCITY = "velocity" + EVAL_STEPS = "eval_steps" + FINAL_XY_DIST = "final_xy_dist" + INITIAL_XY_DIST = "initial_xy_dist" + REACHED_TARGET = "reached_target" # Maps architecture display names to the path of their evaluation CSV. # Update these paths once real evaluation data is available. FILE_MAPPING: dict[str, str] = { - "centralized 2 arms": "runs/dummy/dummy_centralized_2_arms.csv", - "centralized 5 arms": "runs/dummy/dummy_centralized_5_arms.csv", - "decentralized fully connected": "runs/dummy/dummy_decentralized_fully_connected.csv", - "decentralized ring-level": "runs/dummy/dummy_decentralized_ring-level.csv", - "decentralized segment-level": "runs/dummy/dummy_decentralized_segment-level.csv", + # "centralized 2 arms": "runs/dummy/dummy_centralized_2_arms.csv", + "centralized 5 arms": "runs/final-v2-centralized/checkpoint_evaluation.csv", + "decentralized fully connected": "runs/final-v2-fully-conn/checkpoint_evaluation.csv", + "decentralized ring-level": "runs/final-v2-ring/checkpoint_evaluation.csv", } # Architecture profiles for dummy data generation: (max_reward, max_velocity, sigmoid_speed) @@ -108,7 +111,13 @@ def load_metrics(file_mapping: dict[str, str]) -> pd.DataFrame: Loads one CSV per architecture, injects the architecture name as a column, and returns the combined DataFrame with only the required columns. """ - required = [Columns.TIMESTEPS, Columns.REWARD, Columns.VELOCITY] + required = [ + Columns.TIMESTEPS, + Columns.REWARD, + Columns.INITIAL_XY_DIST, + Columns.FINAL_XY_DIST, + Columns.EVAL_STEPS, + ] dfs = [] for arch_name, filepath in file_mapping.items(): @@ -125,6 +134,10 @@ def load_metrics(file_mapping: dict[str, str]) -> pd.DataFrame: df = df[required].copy() df[Columns.ARCH] = arch_name + df[Columns.VELOCITY] = (df[Columns.INITIAL_XY_DIST] - df[Columns.FINAL_XY_DIST]) / df[ + Columns.EVAL_STEPS + ] + dfs.append(df) return pd.concat(dfs, ignore_index=True) if dfs else pd.DataFrame() @@ -303,6 +316,7 @@ def plot_results(df: pd.DataFrame, results: pd.DataFrame, output_dir: str, **kwa def obtain_data() -> pd.DataFrame: """Resolves the file mapping, falling back to generated dummy CSVs if needed.""" global USING_DUMMY_DATA + if not any(os.path.exists(p) for p in FILE_MAPPING.values()): logger.info("No real evaluation files found. Generating dummy CSVs at expected locations.") generate_dummy_csvs(FILE_MAPPING) diff --git a/scripts/download_wandb_project.py b/scripts/tools/download_wandb_project.py similarity index 51% rename from scripts/download_wandb_project.py rename to scripts/tools/download_wandb_project.py index 7c04ad1..bf2fb11 100644 --- a/scripts/download_wandb_project.py +++ b/scripts/tools/download_wandb_project.py @@ -7,6 +7,7 @@ import argparse # tune these depending on network / W&B limits MAX_RUN_WORKERS = 8 MAX_FILE_WORKERS = 16 +MAX_ARTIFACT_WORKERS = 8 api = wandb.Api() @@ -20,19 +21,42 @@ def safe_print(*args, **kwargs): def download_file(file, run_dir): target = run_dir / file.name + try: # skip existing files if target.exists(): - return f"SKIP {target}" + return f"SKIP FILE {target}" target.parent.mkdir(parents=True, exist_ok=True) file.download(root=run_dir, replace=False) - return f"DONE {target}" + return f"DONE FILE {target}" except Exception as e: - return f"FAIL {target}: {e}" + return f"FAIL FILE {target}: {e}" + + +def sanitize_artifact_name(name: str): + return name.replace(":", "_") + + +def download_artifact(artifact, artifact_root): + try: + artifact_name = sanitize_artifact_name(artifact.name) + artifact_dir = artifact_root / artifact_name + + if artifact_dir.exists() and any(artifact_dir.iterdir()): + return f"SKIP ARTIFACT {artifact.name}" + + artifact_dir.mkdir(parents=True, exist_ok=True) + + artifact.download(root=artifact_dir) + + return f"DONE ARTIFACT {artifact.name}" + + except Exception as e: + return f"FAIL ARTIFACT {artifact.name}: {e}" def download_run(run, root): @@ -41,6 +65,9 @@ def download_run(run, root): safe_print(f"\n=== {run.name} ({run.id}) ===") + # ------------------------- + # Download regular run files + # ------------------------- files = list(run.files()) with ThreadPoolExecutor(max_workers=MAX_FILE_WORKERS) as executor: @@ -49,10 +76,38 @@ def download_run(run, root): for future in as_completed(futures): safe_print(future.result()) - # OPTIONAL: download artifacts too - # for artifact in run.logged_artifacts(): - # artifact_dir = run_dir / "artifacts" / artifact.name - # artifact.download(root=artifact_dir) + # ------------------------- + # Download logged artifacts + # ------------------------- + artifact_root = run_dir / "artifacts" + + try: + artifacts = list(run.logged_artifacts()) + safe_print(f"Found {len(artifacts)} artifacts for {run.name}") + + with ThreadPoolExecutor(max_workers=MAX_ARTIFACT_WORKERS) as executor: + futures = [ + executor.submit(download_artifact, artifact, artifact_root) + for artifact in artifacts + ] + + for future in as_completed(futures): + safe_print(future.result()) + + except Exception as e: + safe_print(f"Artifact download failed for {run.name}: {e}") + + # ------------------------- + # OPTIONAL: download used/input artifacts + # ------------------------- + # try: + # used_artifacts = list(run.used_artifacts()) + # used_root = run_dir / "used_artifacts" + # + # for artifact in used_artifacts: + # download_artifact(artifact, used_root) + # except Exception as e: + # safe_print(f"Used artifact download failed: {e}") safe_print(f"Finished {run.name}") From 748479c046211308c3886c5b5a15becd65615b6c Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Fri, 15 May 2026 01:11:22 +0200 Subject: [PATCH 08/25] 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 09/25] 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 10/25] 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 11/25] 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 12/25] 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 13/25] 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 14/25] 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 15/25] 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 7334565d69d3b0ea8314ef524013055135ac5c37 Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Fri, 15 May 2026 20:52:06 +0200 Subject: [PATCH 16/25] feat(convergence analysis): updated convergence analysis script to also print out index of found item --- .gitignore | 3 ++ configs/centralized-final.yaml | 6 ++-- configs/evaluation/poster.yaml | 6 ++-- configs/fully-connected-final.yaml | 6 ++-- configs/ring-final.yaml | 6 ++-- scripts/plots/analyze_convergence.py | 52 +++++++++++++++++++++++----- 6 files changed, 60 insertions(+), 19 deletions(-) diff --git a/.gitignore b/.gitignore index dd76b50..b40f5c8 100644 --- a/.gitignore +++ b/.gitignore @@ -523,3 +523,6 @@ Network Trash Folder Temporary Items .apdisk *.pdf + +# plot directory +poster_plots/ \ No newline at end of file diff --git a/configs/centralized-final.yaml b/configs/centralized-final.yaml index f40570d..9c50437 100644 --- a/configs/centralized-final.yaml +++ b/configs/centralized-final.yaml @@ -23,7 +23,7 @@ morphology: morph_mode: CENTRALIZED experiment: - exp_name: "final-models/centralized/" + exp_name: "final-models-v2/centralized/" seed: 42 torch_deterministic: true cuda: true @@ -34,8 +34,8 @@ logging: save_checkpoints: true upload_final_model: true upload_checkpoints: true - checkpoint_frequency: 20 - wandb_project_name: "final-models" + checkpoint_frequency: 10 + wandb_project_name: "final-models-v2" evaluation: evaluate_checkpoints: true diff --git a/configs/evaluation/poster.yaml b/configs/evaluation/poster.yaml index 0777cd1..948e4c7 100644 --- a/configs/evaluation/poster.yaml +++ b/configs/evaluation/poster.yaml @@ -9,12 +9,14 @@ eval_seed: 0 # Cross-model comparison settings # We use 10 episodes to get a more robust average for the final poster results. comparison_base_seed: 0 -comparison_num_episodes: 2 +comparison_num_episodes: 10 comparison_output_csv: "runs/evaluation/comparison.csv" # Paths to the .cleanrl_model files to be compared (relative to workspace root). comparison_models: - - "runs/input-space-2-arms/2026-05-02/08-14-58/final_model.flax" + - "runs/final-v2-centralized/artifacts/12-19-01_checkpoint_v22/checkpoint_step_230.flax" + - "runs/final-v2-fully-conn/artifacts/14-02-00_checkpoint_v17/checkpoint_step_180.flax" + - "runs/final-v2-ring/artifacts/15-27-03_checkpoint_v21/checkpoint_step_220.flax" # Path to the morphologies to evaluate against. comparison_morphologies: diff --git a/configs/fully-connected-final.yaml b/configs/fully-connected-final.yaml index 8d60451..29983e3 100644 --- a/configs/fully-connected-final.yaml +++ b/configs/fully-connected-final.yaml @@ -26,7 +26,7 @@ morphology: morph_mode: FULLY_CONNECTED experiment: - exp_name: "final-models/fully-connected/" + exp_name: "final-models-v2/fully-connected/" seed: 42 torch_deterministic: true cuda: true @@ -37,8 +37,8 @@ logging: save_checkpoints: true upload_final_model: true upload_checkpoints: true - checkpoint_frequency: 20 - wandb_project_name: "final-models" + checkpoint_frequency: 10 + wandb_project_name: "final-models-v2" evaluation: evaluate_checkpoints: true diff --git a/configs/ring-final.yaml b/configs/ring-final.yaml index ffba64f..a0d852a 100644 --- a/configs/ring-final.yaml +++ b/configs/ring-final.yaml @@ -26,7 +26,7 @@ morphology: morph_mode: RING experiment: - exp_name: "final-models/ring/" + exp_name: "final-models-v2/ring/" seed: 42 torch_deterministic: true cuda: true @@ -37,8 +37,8 @@ logging: save_checkpoints: true upload_final_model: true upload_checkpoints: true - checkpoint_frequency: 20 - wandb_project_name: "final-models" + checkpoint_frequency: 10 + wandb_project_name: "final-models-v2" evaluation: evaluate_checkpoints: true diff --git a/scripts/plots/analyze_convergence.py b/scripts/plots/analyze_convergence.py index 2612bb9..ec21a9a 100644 --- a/scripts/plots/analyze_convergence.py +++ b/scripts/plots/analyze_convergence.py @@ -44,6 +44,7 @@ class Columns(str, Enum): # ... (rest of the file remains same, just need to update plotting functions and obtain_data) """Column names expected in every evaluation CSV.""" + CHECKPOINT = "checkpoint" ARCH = "architecture" TIMESTEPS = "total_trained_timesteps" REWARD = "accumulated_reward" @@ -108,7 +109,18 @@ def load_metrics(file_mapping: dict[str, str]) -> pd.DataFrame: Loads one CSV per architecture, injects the architecture name as a column, and returns the combined DataFrame with only the required columns. """ +<<<<<<< Updated upstream required = [Columns.TIMESTEPS, Columns.REWARD, Columns.VELOCITY] +======= + required = [ + Columns.CHECKPOINT, + Columns.TIMESTEPS, + Columns.REWARD, + Columns.INITIAL_XY_DIST, + Columns.FINAL_XY_DIST, + Columns.EVAL_STEPS, + ] +>>>>>>> Stashed changes dfs = [] for arch_name, filepath in file_mapping.items(): @@ -130,11 +142,17 @@ def load_metrics(file_mapping: dict[str, str]) -> pd.DataFrame: return pd.concat(dfs, ignore_index=True) if dfs else pd.DataFrame() -def _convergence_timestep(series: pd.Series, timesteps: pd.Series) -> float: +def _convergence_timestep( + series: pd.Series, timesteps: pd.Series, checkpoints: pd.Series +) -> tuple[float, int, int]: """Returns the first timestep where the smoothed series reaches 95% of its peak.""" smoothed = series.rolling(window=SMOOTHING_WINDOW, min_periods=1).mean() threshold = smoothed.max() * CONVERGENCE_THRESHOLD - return timesteps[smoothed >= threshold].iloc[0] + + mask = smoothed >= threshold + first_idx = mask.idxmax() + + return timesteps.loc[first_idx], first_idx, checkpoints.loc[first_idx] def analyze_convergence(df: pd.DataFrame) -> pd.DataFrame: @@ -147,15 +165,27 @@ def analyze_convergence(df: pd.DataFrame) -> pd.DataFrame: for arch in df[Columns.ARCH].unique(): arch_data = df[df[Columns.ARCH] == arch].sort_values(Columns.TIMESTEPS) + reward_timestep, reward_checkpoint_idx, reward_checkpoint = _convergence_timestep( + arch_data[Columns.REWARD], + arch_data[Columns.TIMESTEPS], + arch_data[Columns.CHECKPOINT], + ) + + velocity_timestep, velocity_checkpoint_idx, velocity_checkpoint = _convergence_timestep( + arch_data[Columns.VELOCITY], + arch_data[Columns.TIMESTEPS], + arch_data[Columns.CHECKPOINT], + ) + results.append( { "Architecture": arch, - "Reward_Convergence_Timestep": _convergence_timestep( - arch_data[Columns.REWARD], arch_data[Columns.TIMESTEPS] - ), - "Velocity_Convergence_Timestep": _convergence_timestep( - arch_data[Columns.VELOCITY], arch_data[Columns.TIMESTEPS] - ), + "Reward_Convergence_Timestep": reward_timestep, + "Reward_Convergence_Checkpoint_Idx": reward_checkpoint_idx, + "Reward_Convergence_Checkpoint": reward_checkpoint, + "Velocity_Convergence_Timestep": velocity_timestep, + "Velocity_Convergence_Checkpoint_Idx": velocity_checkpoint_idx, + "Velocity_Convergence_Checkpoint": velocity_checkpoint, } ) @@ -319,6 +349,12 @@ def run_analysis(output_dir: str, **kwargs): return results = analyze_convergence(df) + print( + results[ + ["Architecture", "Reward_Convergence_Checkpoint_Idx", "Reward_Convergence_Checkpoint"] + ] + ) + plot_results(df, results, output_dir, **kwargs) logger.info("Analysis complete. Plots saved to disk.") From a2173f64074154e9abdc2ae9206649292daeec16 Mon Sep 17 00:00:00 2001 From: Tibo De Peuter Date: Fri, 15 May 2026 21:33:55 +0200 Subject: [PATCH 17/25] chore: streamline visualisations --- scripts/plots/analyze_comparisons.py | 4 ++-- scripts/plots/plot_config.py | 20 +++++++++----------- 2 files changed, 11 insertions(+), 13 deletions(-) diff --git a/scripts/plots/analyze_comparisons.py b/scripts/plots/analyze_comparisons.py index 3ac6e8a..0a036f8 100644 --- a/scripts/plots/analyze_comparisons.py +++ b/scripts/plots/analyze_comparisons.py @@ -182,7 +182,7 @@ def plot_grouped_bar( color="w", markerfacecolor=BEST_PERFORMER_COLOR, markersize=15, - label="Best Performance", + # label="Best Performance", ls="", ) ax.legend(**LEGEND_KWARGS, ncol=len(architectures) + 1) @@ -317,7 +317,7 @@ def plot_grouped_bar_alt( color="w", markerfacecolor=BEST_PERFORMER_COLOR, markersize=15, - label="Best Performance", + # label="Best Performance", ls="", ) ax.legend(**LEGEND_KWARGS, ncol=len(morphologies) + 1) diff --git a/scripts/plots/plot_config.py b/scripts/plots/plot_config.py index 5fd6ffd..e8cd78f 100644 --- a/scripts/plots/plot_config.py +++ b/scripts/plots/plot_config.py @@ -4,26 +4,24 @@ import matplotlib.pyplot as plt # Shared Color Palette (Colorblind friendly, high contrast) # Matches poster design COLORS = { - "CENTRALIZED": "#2B4162", # Deep Slate Blue - "FULLY_CONNECTED": "#FA9F42", # Vibrant Orange - "RING_LEVEL": "#4E937A", # Muted Teal - "SEGMENT_LEVEL": "#B4436C", # Soft Red - "DECENTRALIZED": "#4E937A", # Default decentralized fallback + "CENTRALIZED": "#0D567C", # Blue + "FULLY_CONNECTED": "#8C0E0F", # Reddish + "RING_LEVEL": "#E1BA6D", # Pale Yellow } -def apply_style(font_size=28): +def apply_style(font_size=36): """ Applies the shared typography and aesthetic settings to Matplotlib. """ plt.rcParams.update( { "font.size": font_size, - "axes.labelsize": font_size + 4, - "axes.titlesize": font_size + 8, - "xtick.labelsize": font_size - 4, - "ytick.labelsize": font_size - 4, - "legend.fontsize": font_size - 6, + "axes.labelsize": font_size, + "axes.titlesize": font_size, + "xtick.labelsize": font_size, + "ytick.labelsize": font_size, + "legend.fontsize": font_size, "axes.linewidth": 2, "axes.spines.top": False, "axes.spines.right": False, From 3489448a3db95312d9856add2e2ed4976633e5e5 Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Sat, 16 May 2026 01:50:35 +0200 Subject: [PATCH 18/25] 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 31a70480fb17b07f80d874e12814b42bf4c4808f Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Sat, 16 May 2026 12:35:07 +0200 Subject: [PATCH 19/25] other: prep for merge conflicts with other branch --- scripts/plots/analyze_comparisons.py | 2 +- scripts/plots/analyze_convergence.py | 13 ++++++------- 2 files changed, 7 insertions(+), 8 deletions(-) diff --git a/scripts/plots/analyze_comparisons.py b/scripts/plots/analyze_comparisons.py index 3ac6e8a..fa5c05c 100644 --- a/scripts/plots/analyze_comparisons.py +++ b/scripts/plots/analyze_comparisons.py @@ -172,7 +172,7 @@ def plot_grouped_bar( plt.FuncFormatter(lambda x, _: f"{x:.2f}" if abs(x) < 10 else f"{x:.0f}") ) - _add_square_placeholders(ax, x_ticks_pos, [f"{m} Arms" for m in morphologies]) + # _add_square_placeholders(ax, x_ticks_pos, [f"{m} Arms" for m in morphologies]) # Add custom legend entry for best performer ax.plot( diff --git a/scripts/plots/analyze_convergence.py b/scripts/plots/analyze_convergence.py index ec21a9a..cf2e2ce 100644 --- a/scripts/plots/analyze_convergence.py +++ b/scripts/plots/analyze_convergence.py @@ -44,11 +44,14 @@ class Columns(str, Enum): # ... (rest of the file remains same, just need to update plotting functions and obtain_data) """Column names expected in every evaluation CSV.""" - CHECKPOINT = "checkpoint" ARCH = "architecture" - TIMESTEPS = "total_trained_timesteps" - REWARD = "accumulated_reward" + TIMESTEPS = "trained_timesteps" + REWARD = "eval_return" VELOCITY = "velocity" + EVAL_STEPS = "eval_steps" + FINAL_XY_DIST = "final_xy_dist" + INITIAL_XY_DIST = "initial_xy_dist" + REACHED_TARGET = "reached_target" # Maps architecture display names to the path of their evaluation CSV. @@ -109,9 +112,6 @@ def load_metrics(file_mapping: dict[str, str]) -> pd.DataFrame: Loads one CSV per architecture, injects the architecture name as a column, and returns the combined DataFrame with only the required columns. """ -<<<<<<< Updated upstream - required = [Columns.TIMESTEPS, Columns.REWARD, Columns.VELOCITY] -======= required = [ Columns.CHECKPOINT, Columns.TIMESTEPS, @@ -120,7 +120,6 @@ def load_metrics(file_mapping: dict[str, str]) -> pd.DataFrame: Columns.FINAL_XY_DIST, Columns.EVAL_STEPS, ] ->>>>>>> Stashed changes dfs = [] for arch_name, filepath in file_mapping.items(): From fe7aecb62a3fa2293fd4acecc0ebf8d3a6c75f38 Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Sat, 16 May 2026 12:36:58 +0200 Subject: [PATCH 20/25] feat(plot colors): copied over the config from the other (chaotic) branch where config was wrongfully pushed to --- scripts/plots/plot_config.py | 20 +++++++++----------- 1 file changed, 9 insertions(+), 11 deletions(-) diff --git a/scripts/plots/plot_config.py b/scripts/plots/plot_config.py index 5fd6ffd..15f0e5c 100644 --- a/scripts/plots/plot_config.py +++ b/scripts/plots/plot_config.py @@ -4,26 +4,24 @@ import matplotlib.pyplot as plt # Shared Color Palette (Colorblind friendly, high contrast) # Matches poster design COLORS = { - "CENTRALIZED": "#2B4162", # Deep Slate Blue - "FULLY_CONNECTED": "#FA9F42", # Vibrant Orange - "RING_LEVEL": "#4E937A", # Muted Teal - "SEGMENT_LEVEL": "#B4436C", # Soft Red - "DECENTRALIZED": "#4E937A", # Default decentralized fallback + "CENTRALIZED": "#0D567C", # Blue + "FULLY_CONNECTED": "#8C0E0F", # Reddish + "RING_LEVEL": "#FCB305", # Pale Yellow } -def apply_style(font_size=28): +def apply_style(font_size=36): """ Applies the shared typography and aesthetic settings to Matplotlib. """ plt.rcParams.update( { "font.size": font_size, - "axes.labelsize": font_size + 4, - "axes.titlesize": font_size + 8, - "xtick.labelsize": font_size - 4, - "ytick.labelsize": font_size - 4, - "legend.fontsize": font_size - 6, + "axes.labelsize": font_size, + "axes.titlesize": font_size, + "xtick.labelsize": font_size, + "ytick.labelsize": font_size, + "legend.fontsize": font_size, "axes.linewidth": 2, "axes.spines.top": False, "axes.spines.right": False, From 708b06157729b3f383d967805f34c35575de7f16 Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Sat, 16 May 2026 15:07:15 +0200 Subject: [PATCH 21/25] feat(plotting): updated plots to not contain placeholder squares + updated colors --- scripts/plots/analyze_comparisons.py | 60 +++++++++++++++------------- scripts/plots/analyze_convergence.py | 1 + scripts/plots/plot_config.py | 4 +- 3 files changed, 35 insertions(+), 30 deletions(-) diff --git a/scripts/plots/analyze_comparisons.py b/scripts/plots/analyze_comparisons.py index fa5c05c..24e28fc 100644 --- a/scripts/plots/analyze_comparisons.py +++ b/scripts/plots/analyze_comparisons.py @@ -84,7 +84,8 @@ def plot_grouped_bar( fig, ax = plt.subplots(figsize=figsize) bar_width = 0.35 - x_indices = np.arange(len(morphologies)) + group_spacing = 1.3 + x_indices = np.arange(len(morphologies)) * group_spacing all_bars = {} all_means = [] @@ -118,7 +119,7 @@ def plot_grouped_bar( ) all_bars[arch] = (x_pos, means, stds, bars) - for m_idx, m in enumerate(morphologies): + for m_idx, _ in enumerate(morphologies): m_means = {arch: all_bars[arch][1][m_idx] for arch in architectures} best_arch = ( max(m_means, key=m_means.get) if higher_is_better else min(m_means, key=m_means.get) @@ -144,12 +145,13 @@ def plot_grouped_bar( x_ticks_pos = ( x_indices + + bar_width # center the label in the 3 bars + (bar_width / 2 if len(architectures) % 2 == 0 else 0) - (bar_width / 2 if len(architectures) == 2 else 0) ) ax.set_xticks(x_ticks_pos) ax.set_xticklabels([f"{m} Arms" for m in morphologies]) - ax.tick_params(axis="x", pad=25) # More padding for the squares + ax.tick_params(axis="x") # More padding for the squares # X-axis at zero ax.axhline(0, color="black", linewidth=1.5) @@ -175,17 +177,18 @@ def plot_grouped_bar( # _add_square_placeholders(ax, x_ticks_pos, [f"{m} Arms" for m in morphologies]) # Add custom legend entry for best performer - ax.plot( - [], - [], - marker=BEST_PERFORMER_MARKER, - color="w", - markerfacecolor=BEST_PERFORMER_COLOR, - markersize=15, - label="Best Performance", - ls="", - ) - ax.legend(**LEGEND_KWARGS, ncol=len(architectures) + 1) + # ax.plot( + # [], + # [], + # marker=BEST_PERFORMER_MARKER, + # color="w", + # markerfacecolor=BEST_PERFORMER_COLOR, + # markersize=15, + # label="Best Performance", + # ls="", + # ) + + ax.legend(**LEGEND_KWARGS, ncol=len(architectures)) ax.set_facecolor("white") fig.patch.set_facecolor("white") @@ -306,21 +309,22 @@ def plot_grouped_bar_alt( ) # In this alt plot, placeholders might be per architecture - _add_square_placeholders( - ax, x_indices, [arch.replace("_", "\n").title() for arch in architectures] - ) + # _add_square_placeholders( + # ax, x_indices, [arch.replace("_", "\n").title() for arch in architectures] + # ) - ax.plot( - [], - [], - marker=BEST_PERFORMER_MARKER, - color="w", - markerfacecolor=BEST_PERFORMER_COLOR, - markersize=15, - label="Best Performance", - ls="", - ) - ax.legend(**LEGEND_KWARGS, ncol=len(morphologies) + 1) + # ax.plot( + # [], + # [], + # marker=BEST_PERFORMER_MARKER, + # color="w", + # markerfacecolor=BEST_PERFORMER_COLOR, + # markersize=15, + # label="Best Performance", + # ls="", + # ) + + ax.legend(**LEGEND_KWARGS, ncol=len(morphologies)) ax.set_facecolor("white") fig.patch.set_facecolor("white") diff --git a/scripts/plots/analyze_convergence.py b/scripts/plots/analyze_convergence.py index cf2e2ce..df9826c 100644 --- a/scripts/plots/analyze_convergence.py +++ b/scripts/plots/analyze_convergence.py @@ -44,6 +44,7 @@ class Columns(str, Enum): # ... (rest of the file remains same, just need to update plotting functions and obtain_data) """Column names expected in every evaluation CSV.""" + CHECKPOINT = "checkpoint" ARCH = "architecture" TIMESTEPS = "trained_timesteps" REWARD = "eval_return" diff --git a/scripts/plots/plot_config.py b/scripts/plots/plot_config.py index 15f0e5c..48f9106 100644 --- a/scripts/plots/plot_config.py +++ b/scripts/plots/plot_config.py @@ -6,7 +6,7 @@ import matplotlib.pyplot as plt COLORS = { "CENTRALIZED": "#0D567C", # Blue "FULLY_CONNECTED": "#8C0E0F", # Reddish - "RING_LEVEL": "#FCB305", # Pale Yellow + "RING": "#FCB305", # Pale Yellow } @@ -42,7 +42,7 @@ BEST_PERFORMER_COLOR = "#D4AF37" # Gold # Centralized Legend Configuration LEGEND_KWARGS = { "loc": "upper center", - "bbox_to_anchor": (0.5, -0.5), + "bbox_to_anchor": (0.5, -0.12), "frameon": False, } From 4b25e1c57027bea99d1e3db197a7bde4b6d328ce Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Sat, 16 May 2026 18:29:14 +0200 Subject: [PATCH 22/25] 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 2493c8d2b306fa55bd1dee966706b4e2bfaf4283 Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Tue, 19 May 2026 10:16:18 +0200 Subject: [PATCH 23/25] feat(plotting): removed best marker legend entry --- scripts/plots/analyze_comparisons.py | 19 +++++++++---------- scripts/plots/analyze_convergence.py | 10 ++++++++++ 2 files changed, 19 insertions(+), 10 deletions(-) diff --git a/scripts/plots/analyze_comparisons.py b/scripts/plots/analyze_comparisons.py index 24e28fc..d001b03 100644 --- a/scripts/plots/analyze_comparisons.py +++ b/scripts/plots/analyze_comparisons.py @@ -6,18 +6,17 @@ Rate, Distance Remaining). """ import os -import pandas as pd -import numpy as np -import matplotlib.pyplot as plt +import matplotlib.pyplot as plt +import numpy as np +import pandas as pd from plot_config import ( - COLORS, - apply_style, - BEST_PERFORMER_MARKER, - BEST_PERFORMER_TEXT, BEST_PERFORMER_COLOR, - create_common_parser, + BEST_PERFORMER_TEXT, + COLORS, LEGEND_KWARGS, + apply_style, + create_common_parser, ) @@ -362,8 +361,8 @@ if __name__ == "__main__": plot_grouped_bar( df=df, metric_col="approx_max_velocity", - ylabel="Max Forward Velocity (cm/s)", - title="Graceful Degradation: Velocity Across Morphologies", + ylabel="", + title="Maximal forward velocity (in cm/s)", output_filename="poster_plot_velocity.png", output_dir=OUTPUT_DIR, higher_is_better=True, diff --git a/scripts/plots/analyze_convergence.py b/scripts/plots/analyze_convergence.py index 3bb078d..bc12c79 100644 --- a/scripts/plots/analyze_convergence.py +++ b/scripts/plots/analyze_convergence.py @@ -135,6 +135,9 @@ def load_metrics(file_mapping: dict[str, str]) -> pd.DataFrame: continue df = df[required].copy() + df[Columns.VELOCITY] = (df[Columns.INITIAL_XY_DIST] - df[Columns.FINAL_XY_DIST]) / df[ + Columns.EVAL_STEPS + ] df[Columns.ARCH] = arch_name df[Columns.VELOCITY] = (df[Columns.INITIAL_XY_DIST] - df[Columns.FINAL_XY_DIST]) / df[ Columns.EVAL_STEPS @@ -165,6 +168,8 @@ def analyze_convergence(df: pd.DataFrame) -> pd.DataFrame: """ results = [] + centralized_base = 0 + for arch in df[Columns.ARCH].unique(): arch_data = df[df[Columns.ARCH] == arch].sort_values(Columns.TIMESTEPS) @@ -192,6 +197,11 @@ def analyze_convergence(df: pd.DataFrame) -> pd.DataFrame: } ) + if arch == "centralized 5 arms": + centralized_base = reward_checkpoint + else: + print(arch, "speedup:", 1 - reward_checkpoint / centralized_base) + return pd.DataFrame(results) From 1ddd065ef5b77a08a18577345faf718770f2eb2c Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Tue, 19 May 2026 10:28:26 +0200 Subject: [PATCH 24/25] cleanup(plotting): removed commented code --- scripts/plots/analyze_comparisons.py | 30 ---------------------------- 1 file changed, 30 deletions(-) diff --git a/scripts/plots/analyze_comparisons.py b/scripts/plots/analyze_comparisons.py index d001b03..8a66b4c 100644 --- a/scripts/plots/analyze_comparisons.py +++ b/scripts/plots/analyze_comparisons.py @@ -173,20 +173,6 @@ def plot_grouped_bar( plt.FuncFormatter(lambda x, _: f"{x:.2f}" if abs(x) < 10 else f"{x:.0f}") ) - # _add_square_placeholders(ax, x_ticks_pos, [f"{m} Arms" for m in morphologies]) - - # Add custom legend entry for best performer - # ax.plot( - # [], - # [], - # marker=BEST_PERFORMER_MARKER, - # color="w", - # markerfacecolor=BEST_PERFORMER_COLOR, - # markersize=15, - # label="Best Performance", - # ls="", - # ) - ax.legend(**LEGEND_KWARGS, ncol=len(architectures)) ax.set_facecolor("white") fig.patch.set_facecolor("white") @@ -307,22 +293,6 @@ def plot_grouped_bar_alt( plt.FuncFormatter(lambda x, _: f"{x:.2f}" if abs(x) < 10 else f"{x:.0f}") ) - # In this alt plot, placeholders might be per architecture - # _add_square_placeholders( - # ax, x_indices, [arch.replace("_", "\n").title() for arch in architectures] - # ) - - # ax.plot( - # [], - # [], - # marker=BEST_PERFORMER_MARKER, - # color="w", - # markerfacecolor=BEST_PERFORMER_COLOR, - # markersize=15, - # label="Best Performance", - # ls="", - # ) - ax.legend(**LEGEND_KWARGS, ncol=len(morphologies)) ax.set_facecolor("white") fig.patch.set_facecolor("white") From 1e7f8c1880976491bccc5a9b749203c755e266cd Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Tue, 19 May 2026 17:57:03 +0200 Subject: [PATCH 25/25] 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.