rosdiff / worker /sim_runner.py
Chandra Kiran
Export runs as LeRobot datasets
b357d7d unverified
Raw History Blame Contribute Delete
21.4 kB
"""Run a scene in MuJoCo and record it: Foxglove MCAP, MP4 video, JSON summary.
Pure local code (no RunPod, no network) so it can be tested anywhere:
python -m worker.sim_runner SCENE_DIR/scene.xml OUT_DIR --duration 5
Outputs in OUT_DIR:
run.mcap /scene (foxglove.SceneUpdate: every visible geom, meshes included),
/tf (foxglove.FrameTransforms: every body's pose at `fps`),
/state (JSON: time, root height, fallen flag, contacts)
video.mp4 H.264 render from a camera tracking the robot (skipped if rendering is unavailable)
video.webm the same video as VP9, for browsers without H.264
trajectory.npz at `fps`: the robot's joint positions/velocities and the controls (dataset export), and
every body's world pose (replay renders)
summary.json what happened: duration, real-time factor, fall detection, warnings
"""
from __future__ import annotations
import argparse
import json
import math
import os
import struct
import sys
import time
from dataclasses import dataclass
from pathlib import Path
import mujoco
import numpy as np
from api.robots import reset_to_robot_keyframe
NS = 1_000_000_000
VISIBLE_GROUPS = (0, 1, 2) # MuJoCo's viewer shows groups 0-2 by default; Menagerie puts collision meshes in 3
@dataclass
class RunConfig:
duration_s: float = 5.0
controller: str = "hold" # "hold": position actuators track the first keyframe; "passive": zero control
fps: int = 30
width: int = 1280
height: int = 720
render: bool = True
root_body: str | None = "pelvis"
fall_height: float | None = 0.35 # root body below this (m) counts as fallen; None: fixed-base robot
home_qpos: tuple | None = None # the robot's starting pose when its model has no keyframe
camera_distance: float = 4.5
kp: float = 0.0 # PD gains for torque-controlled joints under "hold" (0: derive from the torque limits)
kd: float = 0.0
# --------------------------------------------------------------------------- geometry -> Foxglove
def _stl_bytes(model: mujoco.MjModel, mesh_id: int) -> bytes:
"""Binary STL of a compiled mesh (vertices are in the mesh frame, as MuJoCo stores them)."""
va, vn = model.mesh_vertadr[mesh_id], model.mesh_vertnum[mesh_id]
fa, fn = model.mesh_faceadr[mesh_id], model.mesh_facenum[mesh_id]
tris = model.mesh_vert[va : va + vn][model.mesh_face[fa : fa + fn]].astype(np.float32) # (fn, 3, 3)
normals = np.cross(tris[:, 1] - tris[:, 0], tris[:, 2] - tris[:, 0])
lengths = np.linalg.norm(normals, axis=1, keepdims=True)
normals = np.divide(normals, lengths, out=np.zeros_like(normals), where=lengths > 0)
record = np.dtype([("normal", "<f4", 3), ("verts", "<f4", (3, 3)), ("attr", "<u2")])
body = np.empty(fn, dtype=record)
body["normal"], body["verts"], body["attr"] = normals, tris, 0
return b"\0" * 80 + struct.pack("<I", fn) + body.tobytes()
def _scene_update(model: mujoco.MjModel):
"""One SceneEntity per body, frame-locked to that body's TF frame."""
from foxglove import messages as fm
def pose(pos, quat):
return fm.Pose(
position=fm.Vector3(x=float(pos[0]), y=float(pos[1]), z=float(pos[2])),
orientation=fm.Quaternion(w=float(quat[0]), x=float(quat[1]), y=float(quat[2]), z=float(quat[3])),
)
def color(rgba):
return fm.Color(r=float(rgba[0]), g=float(rgba[1]), b=float(rgba[2]), a=float(rgba[3]))
meshes: dict[int, bytes] = {}
entities = []
for body in range(model.nbody):
name = mujoco.mj_id2name(model, mujoco.mjtObj.mjOBJ_BODY, body) or f"body_{body}"
cubes, spheres, cylinders, models = [], [], [], []
for g in range(model.ngeom):
if model.geom_bodyid[g] != body or model.geom_group[g] not in VISIBLE_GROUPS:
continue
gtype, size = model.geom_type[g], model.geom_size[g]
rgba = model.geom_rgba[g]
if model.geom_matid[g] >= 0:
rgba = model.mat_rgba[model.geom_matid[g]]
p, c = pose(model.geom_pos[g], model.geom_quat[g]), color(rgba)
T = mujoco.mjtGeom
if gtype == T.mjGEOM_PLANE:
sx = size[0] * 2 if size[0] > 0 else 40.0
sy = size[1] * 2 if size[1] > 0 else 40.0
cubes.append(fm.CubePrimitive(pose=p, size=fm.Vector3(x=sx, y=sy, z=0.002), color=c))
elif gtype == T.mjGEOM_BOX:
cubes.append(
fm.CubePrimitive(pose=p, size=fm.Vector3(x=2 * size[0], y=2 * size[1], z=2 * size[2]), color=c)
)
elif gtype == T.mjGEOM_SPHERE:
d = 2 * size[0]
spheres.append(fm.SpherePrimitive(pose=p, size=fm.Vector3(x=d, y=d, z=d), color=c))
elif gtype == T.mjGEOM_ELLIPSOID:
spheres.append(
fm.SpherePrimitive(pose=p, size=fm.Vector3(x=2 * size[0], y=2 * size[1], z=2 * size[2]), color=c)
)
elif gtype in (T.mjGEOM_CYLINDER, T.mjGEOM_CAPSULE):
r, half = size[0], size[1]
cylinders.append(
fm.CylinderPrimitive(
pose=p, size=fm.Vector3(x=2 * r, y=2 * r, z=2 * half), bottom_scale=1.0, top_scale=1.0, color=c
)
)
if gtype == T.mjGEOM_CAPSULE: # end caps
rot = np.zeros(9)
mujoco.mju_quat2Mat(rot, model.geom_quat[g])
axis = rot.reshape(3, 3)[:, 2]
for sign in (-1, 1):
cp = model.geom_pos[g] + sign * half * axis
spheres.append(
fm.SpherePrimitive(
pose=pose(cp, [1, 0, 0, 0]), size=fm.Vector3(x=2 * r, y=2 * r, z=2 * r), color=c
)
)
elif gtype == T.mjGEOM_MESH:
mesh_id = int(model.geom_dataid[g])
if mesh_id not in meshes:
meshes[mesh_id] = _stl_bytes(model, mesh_id)
models.append(
fm.ModelPrimitive(
pose=p,
scale=fm.Vector3(x=1, y=1, z=1),
color=c,
override_color=True,
media_type="model/stl",
data=meshes[mesh_id],
)
)
# hfield / sdf geoms are not drawn in the 3D view (they are in the video)
if cubes or spheres or cylinders or models:
entities.append(
fm.SceneEntity(
frame_id=name,
id=name,
frame_locked=True,
cubes=cubes,
spheres=spheres,
cylinders=cylinders,
models=models,
)
)
return fm.SceneUpdate(entities=entities)
def _transforms(model: mujoco.MjModel, data: mujoco.MjData, t_ns: int):
from foxglove import messages as fm
ts = fm.Timestamp(sec=t_ns // NS, nsec=t_ns % NS)
out = []
for body in range(1, model.nbody):
name = mujoco.mj_id2name(model, mujoco.mjtObj.mjOBJ_BODY, body) or f"body_{body}"
p, q = data.xpos[body], data.xquat[body]
out.append(
fm.FrameTransform(
timestamp=ts,
parent_frame_id="world",
child_frame_id=name,
translation=fm.Vector3(x=float(p[0]), y=float(p[1]), z=float(p[2])),
rotation=fm.Quaternion(w=float(q[0]), x=float(q[1]), y=float(q[2]), z=float(q[3])),
)
)
return fm.FrameTransforms(transforms=out)
# --------------------------------------------------------------------------- simulation
CAMERA_DISTANCE = {"humanoid": 4.5, "quadruped": 2.5, "arm": 2.2}
def config_for(robot: str | None, duration_s: float, controller: str, render: bool = True) -> RunConfig:
"""Run settings for a robot from api.robots.ROBOTS (root body, fall threshold, starting pose, camera)."""
from api.robots import ROBOTS
if not robot:
return RunConfig(duration_s, controller, render=render, root_body=None, fall_height=None)
spec = ROBOTS[robot]
return RunConfig(
duration_s,
controller,
render=render,
root_body=spec.root_body,
fall_height=spec.fall_height,
home_qpos=spec.home_qpos,
camera_distance=CAMERA_DISTANCE.get(spec.kind, 4.5),
kp=spec.hold_kp,
kd=spec.hold_kd,
)
@dataclass
class HoldController:
"""Holds the robot's starting pose (its keyframe, or RobotSpec.home_qpos).
Position servos get the target joint angle as their control. Torque motors get a PD law,
tau = kp (q* - q) - kd dq, clipped to the motor's range; gains default to reaching the torque
limit at 0.25 rad of error, with kd = kp / 20.
"""
ctrl: np.ndarray # constant part (servo targets)
motors: np.ndarray # actuator indices driven by the PD law
qadr: np.ndarray # their joints' qpos addresses
vadr: np.ndarray # their joints' dof addresses
target: np.ndarray
kp: np.ndarray
kd: np.ndarray
lo: np.ndarray
hi: np.ndarray
def apply(self, data: mujoco.MjData) -> None:
data.ctrl[:] = self.ctrl
if len(self.motors):
tau = self.kp * (self.target - data.qpos[self.qadr]) - self.kd * data.qvel[self.vadr]
data.ctrl[self.motors] = np.clip(tau, self.lo, self.hi)
def _hold_controller(model: mujoco.MjModel, data: mujoco.MjData, config: RunConfig) -> HoldController | None:
"""Hold the pose the run starts in (`data` right after the reset)."""
if model.nu == 0:
return None
qpos = data.qpos.copy()
ctrl = model.key_ctrl[0].copy() if model.nkey and np.any(model.key_ctrl[0]) else np.zeros(model.nu)
motors, qadr, vadr, target, kp, kd, lo, hi = [], [], [], [], [], [], [], []
for a in range(model.nu):
if model.actuator_trntype[a] != mujoco.mjtTrn.mjTRN_JOINT:
continue
joint = model.actuator_trnid[a, 0]
q = qpos[model.jnt_qposadr[joint]]
if model.actuator_biastype[a] == mujoco.mjtBias.mjBIAS_AFFINE: # position servo
if not (model.nkey and np.any(model.key_ctrl[0])):
ctrl[a] = q
continue
# torque motor
limit = model.actuator_ctrlrange[a] if model.actuator_ctrllimited[a] else np.array([-100.0, 100.0])
gain = config.kp or float(max(abs(limit[0]), abs(limit[1]))) / 0.25
motors.append(a)
qadr.append(model.jnt_qposadr[joint])
vadr.append(model.jnt_dofadr[joint])
target.append(q)
kp.append(gain)
kd.append(config.kd or gain / 20)
lo.append(limit[0])
hi.append(limit[1])
arr = np.asarray
return HoldController(
ctrl, arr(motors, int), arr(qadr, int), arr(vadr, int), arr(target), arr(kp), arr(kd), arr(lo), arr(hi)
)
def _apply_controller(model, data, config: RunConfig, hold: HoldController | None) -> None:
if config.controller == "passive" or hold is None:
data.ctrl[:] = 0.0
return
hold.apply(data)
def _robot_state_index(model: mujoco.MjModel, root: int) -> tuple[np.ndarray, np.ndarray, list[str]]:
"""qpos/qvel indices of the robot's own joints (every joint when there is no robot), for trajectory.npz."""
qidx, vidx, names = [], [], []
for j in range(model.njnt):
b = model.jnt_bodyid[j]
while root > 0 and b > 0 and b != root:
b = model.body_parentid[b]
if root > 0 and b != root:
continue
nq = {0: 7, 1: 4}.get(int(model.jnt_type[j]), 1)
nv = {0: 6, 1: 3}.get(int(model.jnt_type[j]), 1)
qidx += range(model.jnt_qposadr[j], model.jnt_qposadr[j] + nq)
vidx += range(model.jnt_dofadr[j], model.jnt_dofadr[j] + nv)
name = mujoco.mj_id2name(model, mujoco.mjtObj.mjOBJ_JOINT, j) or f"joint_{j}"
names += [name] if nq == 1 else [f"{name}[{i}]" for i in range(nq)]
return np.asarray(qidx, int), np.asarray(vidx, int), names
def _open_videos(out_dir: Path, config: RunConfig) -> list:
"""H.264 MP4 (plays in Chrome, Safari, Edge, Firefox) and VP9 WebM (for browsers built without H.264)."""
import imageio_ffmpeg
common = dict(
size=(config.width, config.height), fps=config.fps, quality=None, pix_fmt_out="yuv420p", macro_block_size=8
)
specs = [
("video.mp4", "libx264", ["-crf", "23", "-movflags", "+faststart"]),
(
"video.webm",
"libvpx-vp9",
["-crf", "35", "-b:v", "0", "-deadline", "realtime", "-cpu-used", "8", "-row-mt", "1"],
),
]
writers = []
try:
for name, codec, params in specs:
w = imageio_ffmpeg.write_frames(str(out_dir / name), codec=codec, output_params=params, **common)
w.send(None)
writers.append(w)
except Exception:
_close_videos(writers)
raise
return writers
def _close_videos(writers: list) -> None:
for w in writers:
try:
w.close()
except Exception:
pass
def simulate(scene: Path, out_dir: Path, config: RunConfig) -> dict:
out_dir.mkdir(parents=True, exist_ok=True)
model = mujoco.MjModel.from_xml_path(str(scene))
data = mujoco.MjData(model)
reset_to_robot_keyframe(model, data, config.root_body, config.home_qpos)
target = _hold_controller(model, data, config)
if config.controller == "passive": # motors off (for position servos, ctrl = 0 would still hold a pose)
model.opt.disableflags |= mujoco.mjtDisableBit.mjDSBL_ACTUATION
root = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_BODY, config.root_body) if config.root_body else -1
import foxglove
mcap_path = out_dir / "run.mcap"
writer = foxglove.open_mcap(str(mcap_path), allow_overwrite=True)
from foxglove.channels import FrameTransformsChannel, SceneUpdateChannel
scene_ch = SceneUpdateChannel("/scene")
tf_ch = FrameTransformsChannel("/tf")
state_ch = foxglove.Channel("/state", message_encoding="json")
renderer = cam = None
videos: list = [] # ffmpeg frame writers: MP4 (H.264) and WebM (VP9) from the same frames
render_error = None
if config.render:
try:
model.vis.global_.offwidth = max(model.vis.global_.offwidth, config.width)
model.vis.global_.offheight = max(model.vis.global_.offheight, config.height)
renderer = mujoco.Renderer(model, config.height, config.width)
cam = mujoco.MjvCamera()
if root >= 0:
cam.type = mujoco.mjtCamera.mjCAMERA_TRACKING
cam.trackbodyid = root
else:
cam.type = mujoco.mjtCamera.mjCAMERA_FREE
cam.lookat[:] = model.stat.center
# Behind and slightly left of the robot, looking along +x (where scenes put their content).
cam.distance = 1.5 * model.stat.extent if root < 0 else config.camera_distance
cam.azimuth, cam.elevation = 20, -22
videos = _open_videos(out_dir, config)
except Exception as e: # no GL context, no ffmpeg, ...
_close_videos(videos)
renderer, videos = None, []
render_error = f"{type(e).__name__}: {e}"
scene_ch.log(_scene_update(model), log_time=0)
traj_q, traj_v, traj_names = _robot_state_index(model, root)
traj: dict[str, list] = {"time": [], "qpos": [], "qvel": [], "ctrl": [], "body_pos": [], "body_quat": []}
steps = max(1, int(math.ceil(config.duration_s / model.opt.timestep)))
frame_dt = 1.0 / config.fps
next_frame = 0.0
min_root_z, fell_at, max_contacts, frames, unstable = math.inf, None, 0, 0, None
wall_start = time.monotonic()
for step in range(steps + 1):
if data.time >= next_frame - 0.5 * model.opt.timestep: # capture on simulation time
next_frame += frame_dt
t_ns = int(round(data.time * NS))
tf_ch.log(_transforms(model, data, t_ns), log_time=t_ns)
state = {"time": round(float(data.time), 4), "contacts": int(data.ncon)}
if root >= 0:
z = float(data.xpos[root][2])
state["root_height"] = round(z, 4)
if config.fall_height is not None:
state["fallen"] = z < config.fall_height
state_ch.log(json.dumps(state).encode(), log_time=t_ns)
traj["time"].append(float(data.time))
traj["qpos"].append(data.qpos[traj_q].copy())
traj["qvel"].append(data.qvel[traj_v].copy())
traj["ctrl"].append(data.ctrl.copy())
traj["body_pos"].append(data.xpos.astype(np.float32))
traj["body_quat"].append(data.xquat.astype(np.float32))
if videos:
try:
renderer.update_scene(data, cam)
frame = np.ascontiguousarray(renderer.render())
for v in videos:
v.send(frame)
except Exception as e: # keep simulating and recording the MCAP without video
render_error = f"rendering stopped at t={data.time:.2f}s: {type(e).__name__}: {e}"
_close_videos(videos)
videos = []
frames += 1
if step == steps:
break
_apply_controller(model, data, config, target)
mujoco.mj_step(model, data)
if not (np.all(np.isfinite(data.qpos)) and np.all(np.isfinite(data.qvel))):
unstable = round(float(data.time), 4)
break
max_contacts = max(max_contacts, int(data.ncon))
if root >= 0:
z = float(data.xpos[root][2])
min_root_z = min(min_root_z, z)
if fell_at is None and config.fall_height is not None and z < config.fall_height:
fell_at = round(float(data.time), 3)
wall = time.monotonic() - wall_start
writer.close()
np.savez_compressed(
out_dir / "trajectory.npz",
time=np.asarray(traj["time"], np.float64),
qpos=np.asarray(traj["qpos"], np.float32).reshape(len(traj["time"]), -1),
qvel=np.asarray(traj["qvel"], np.float32).reshape(len(traj["time"]), -1),
ctrl=np.asarray(traj["ctrl"], np.float32).reshape(len(traj["time"]), -1),
joint_names=np.asarray(traj_names),
actuator_names=np.asarray(
[mujoco.mj_id2name(model, mujoco.mjtObj.mjOBJ_ACTUATOR, a) or f"actuator_{a}" for a in range(model.nu)]
),
body_pos=np.asarray(traj["body_pos"], np.float32), # world pose of every body (for replay renders)
body_quat=np.asarray(traj["body_quat"], np.float32),
body_names=np.asarray(
[mujoco.mj_id2name(model, mujoco.mjtObj.mjOBJ_BODY, b) or f"body_{b}" for b in range(model.nbody)]
),
body_parents=np.asarray(model.body_parentid, np.int32),
fps=np.asarray(config.fps),
)
_close_videos(videos)
if renderer is not None:
renderer.close()
warnings = [
f"{mujoco.mjtWarning(i).name} x{int(data.warning[i].number)}"
for i in range(len(data.warning))
if data.warning[i].number
]
summary = {
"simulated_seconds": round(float(data.time), 3),
"requested_seconds": config.duration_s,
"wall_seconds": round(wall, 2),
"realtime_factor": round(float(data.time) / wall, 2) if wall > 0 else None,
"controller": config.controller if target is not None else "passive (no actuators)",
"frames": frames,
"unstable_at": unstable,
"warnings": warnings,
"max_contacts": max_contacts,
"video": render_error is None and (out_dir / "video.mp4").is_file(),
"render_error": render_error,
"model": {"nbody": int(model.nbody), "ngeom": int(model.ngeom), "nq": int(model.nq), "nu": int(model.nu)},
}
if root >= 0:
summary.update(root_min_height=round(min_root_z, 3), root_final_height=round(float(data.xpos[root][2]), 3))
if config.fall_height is not None:
summary.update(fell=fell_at is not None, fell_at=fell_at)
(out_dir / "summary.json").write_text(json.dumps(summary, indent=2))
return summary
def main(argv: list[str] | None = None) -> int:
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
p.add_argument("scene", type=Path)
p.add_argument("out", type=Path)
p.add_argument("--duration", type=float, default=5.0)
p.add_argument("--controller", choices=["hold", "passive"], default="hold")
p.add_argument("--no-render", action="store_true")
p.add_argument("--robot", default="unitree_g1", help="robot the scene includes ('' for none)")
args = p.parse_args(argv)
os.environ.setdefault("MUJOCO_GL", "egl")
config = config_for(args.robot or None, args.duration, args.controller, render=not args.no_render)
summary = simulate(args.scene, args.out, config)
print(json.dumps(summary, indent=2))
return 0
if __name__ == "__main__":
sys.exit(main())