mishig/so101-block-sorting / scripts /lego_simulator.py
mishig's picture
download
raw
10.1 kB
"""SO-101 physics server with a fixed RGB camera and joint-position commands."""
import argparse
import json
import os
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from queue import Empty, Queue
import threading
import time
import sys
import numpy as np
from PIL import Image
from camera_recorder import CameraRecorder
ROOT = Path(__file__).resolve().parent.parent
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--port", type=int, default=int(os.environ.get("SO101_PORT", "8877")))
parser.add_argument("--headless", action="store_true", help="Render the front camera without opening a viewer.")
parser.add_argument("--fast", action="store_true", help="Skip real-time pacing; motor trajectories and physics steps are unchanged.")
args = parser.parse_args()
if not 1 <= args.port <= 65535:
parser.error("--port must be between 1 and 65535")
# EGL supports offscreen rendering on Linux without an X display. An
# explicit MUJOCO_GL setting (e.g. osmesa for software rendering) wins.
if args.headless and sys.platform.startswith("linux"):
os.environ.setdefault("MUJOCO_GL", "egl")
import mujoco
run = ROOT / "work/lego-run"
run.mkdir(parents=True, exist_ok=True)
results = ROOT / "results"
results.mkdir(parents=True, exist_ok=True)
model = mujoco.MjModel.from_xml_path(str(ROOT / "scene/scene.xml"))
data = mujoco.MjData(model)
if model.nkey:
mujoco.mj_resetDataKeyframe(model, data, 0)
# Start with the arm raised and the jaw open, using motor targets only.
data.ctrl[:6] = [0, -.2, .2, 1.4, 0, .6]
mujoco.mj_step(model, data, nstep=round(3 / model.opt.timestep))
requests = Queue()
keys = Queue()
renderer = mujoco.Renderer(model, height=720, width=960)
image_count = 0
state = {}
recorder = None
def robot_state():
return {"scene": "lego", "time": float(data.time), "qpos": data.qpos[:6].tolist(),
"ctrl": data.ctrl[:6].tolist(), "warning_counts": data.warning.number.tolist()}
def capture(label="camera"):
nonlocal image_count
renderer.update_scene(data, camera="front")
rgb = renderer.render()
safe_label = "".join(char if char.isalnum() or char in "-_." else "_" for char in str(label))[:120]
path = run / f"{image_count:04d}-{safe_label}.png"
Image.fromarray(rgb).save(path)
Image.fromarray(rgb).save(results / "lego-camera.png")
image_count += 1
if recorder:
recorder.event(data.time, "camera", label)
return str(path)
class Handler(BaseHTTPRequestHandler):
def log_message(self, *args):
pass
def do_GET(self):
if self.path != "/state":
self.send_error(404); return
payload = json.dumps(state).encode()
self.send_response(200); self.end_headers(); self.wfile.write(payload)
def do_POST(self):
try:
payload = json.loads(self.rfile.read(int(self.headers.get("Content-Length", "0"))))
response = Queue(maxsize=1)
requests.put((self.path, payload, response))
result = response.get(timeout=90)
self.send_response(200); self.end_headers()
self.wfile.write(json.dumps(result).encode())
except Exception as exc:
self.send_error(500, str(exc))
server = ThreadingHTTPServer(("127.0.0.1", args.port), Handler)
threading.Thread(target=server.serve_forever, daemon=True).start()
viewer = None
if not args.headless:
import mujoco.viewer
viewer = mujoco.viewer.launch_passive(
model, data, key_callback=keys.put, show_left_ui=False, show_right_ui=True
)
if viewer:
viewer.cam.type = mujoco.mjtCamera.mjCAMERA_FIXED
viewer.cam.fixedcamid = model.camera("front").id
viewer.sync()
state.update(robot_state())
state["camera"] = capture("initial")
print(f"Ready at http://127.0.0.1:{args.port}; camera={state['camera']}", flush=True)
active = None
paused = False
selected = 0
frame_dt = .016
steps = max(1, round(frame_dt / model.opt.timestep))
running = True
try:
while running and (viewer is None or viewer.is_running()):
started = time.monotonic()
if active is None:
try:
endpoint, command, response = requests.get_nowait()
except Empty:
endpoint = None
if endpoint:
if endpoint == "/move":
target = np.asarray(command["target"], dtype=float)
if target.shape != (6,) or not np.isfinite(target).all():
response.put({"error": "target must be six finite joint angles"})
else:
target = np.clip(target, model.actuator_ctrlrange[:6, 0], model.actuator_ctrlrange[:6, 1])
active = {"initial": data.ctrl[:6].copy(), "target": target,
"frames": max(1, round(float(command.get("seconds", 1.5)) / frame_dt)),
"i": 0, "response": response, "label": command.get("label", "move")}
if recorder:
recorder.event(data.time, "move", active["label"],
target=target.tolist(), seconds=float(command.get("seconds", 1.5)))
elif endpoint == "/camera":
response.put({**robot_state(), "image": capture(command.get("label", "camera"))})
elif endpoint == "/record/start":
if recorder:
response.put({"error": "Recording is already active"})
else:
try:
recorder = CameraRecorder(command.get("path", results / "front-camera.mp4"),
data.time, command.get("fps", 30))
response.put({"recording": True, "video": str(recorder.path), **robot_state()})
except Exception as exc:
response.put({"error": str(exc)})
elif endpoint == "/record/mark":
if recorder:
recorder.event(data.time, command.get("type", "phase"), command["label"])
response.put({"recording": True, **robot_state()})
else:
response.put({"error": "No active recording"})
elif endpoint == "/record/stop":
if recorder:
stopped = recorder
recorder = None
response.put(stopped.close())
else:
response.put({"error": "No active recording"})
elif endpoint == "/audit":
# Diagnostic ground truth, separate from camera observations.
blocks = {name: data.body(f"block_{name}").xpos.tolist() for name in ("red", "green", "blue")}
response.put({"blocks": blocks, "contacts": int(data.ncon), **robot_state()})
elif endpoint == "/shutdown":
response.put({"shutdown": True, **robot_state()})
running = False
else:
response.put({"error": "unknown endpoint"})
while not keys.empty():
key = keys.get_nowait()
if key == 32: paused = not paused
elif ord("1") <= key <= ord("6"): selected = key - ord("1")
elif key in (264, 265) and active is None:
data.ctrl[selected] = np.clip(data.ctrl[selected] + (.0873 if key == 265 else -.0873), *model.actuator_ctrlrange[selected])
# Fast idle mode waits between commands, unless a recording needs
# simulation time to advance through its observation/ending holds.
advance = not args.fast or active is not None or recorder is not None
if running and not paused and advance:
if active:
active["i"] += 1
t = min(1., active["i"] / active["frames"])
blend = t * t * (3 - 2 * t)
data.ctrl[:6] = active["initial"] * (1 - blend) + active["target"] * blend
mujoco.mj_step(model, data, nstep=steps)
if viewer:
viewer.sync()
if recorder and recorder.due(data.time):
renderer.update_scene(data, camera="front")
recorder.write(renderer.render(), data.time)
state.update(robot_state())
if active and active["i"] >= active["frames"] + 20:
image = capture(active["label"])
result = {**robot_state(), "image": image}
with (run / "actions.jsonl").open("a") as log:
log.write(json.dumps({"target": active["target"].tolist(), "label": active["label"], **result}) + "\n")
active["response"].put(result)
active = None
if not args.fast:
remaining = frame_dt - (time.monotonic() - started)
if remaining > 0:
time.sleep(remaining)
elif active is None and recorder is None:
# Keep HTTP and viewer input responsive without spinning.
time.sleep(.002)
finally:
if recorder: recorder.close()
if viewer: viewer.close()
renderer.close()
server.shutdown()
server.server_close()
if __name__ == "__main__":
main()

Xet Storage Details

Size:
10.1 kB
·
Xet hash:
2f4c4adae709f0579c31b4d1dbebd9245a3c773e2845c649e214b933ee2f5870

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.