1
Fork 0

feat: customizable path & robot color

This commit is contained in:
Jona Reynaert 2026-05-15 17:32:04 +02:00
parent beb64ca542
commit ec4e069e2f

View file

@ -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)