From c68a7ddf26b8cb87337416b9c2574bd31eedc186 Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Fri, 15 May 2026 19:50:15 +0200 Subject: [PATCH] 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)