vl/demos/robot_arm.py
19.4 KB · 324 lines · python Raw
1 """Robot-arm pick-and-place with JEV-9B-VL System 1 (POST /v1/decide), simulated in MuJoCo.
2
3 A gantry arm with a two-finger gripper above a table with three coloured cubes and a yellow tray. Every step the model
4 sees two camera images (front and top; the top view shows a black dot under the gripper) and the gripper state, and
5 picks one of 8 actions in one forward pass. Grasping is simplified: closing the gripper low over a cube attaches the
6 cube; opening releases it and it falls under MuJoCo physics.
7
8 MUJOCO_GL=glfw DISPLAY=:99 python3 robot_arm.py run 20
9 MUJOCO_GL=glfw DISPLAY=:99 python3 robot_arm.py record SEED OUT.mp4
10 """
11 import base64, io, json, os, random, subprocess, sys, tempfile, time
12 os.environ.setdefault("MUJOCO_GL", "glfw"); os.environ.setdefault("DISPLAY", ":99")
13 import numpy as np, mujoco, requests
14 from PIL import Image, ImageDraw, ImageFont
15
16 API = os.environ.get("JEV_URL", "http://localhost:8000") + "/v1/decide"
17 HERE = os.path.dirname(os.path.abspath(__file__))
18 FONT, FONTB = "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf"
19 ACTIONS = ["move left", "move right", "move forward (away from the front camera)", "move backward (toward the front camera)",
20 "lower the gripper", "raise the gripper", "close the gripper (grasp)", "open the gripper (release)"]
21 COLORS = {"red": "0.85 0.15 0.15 1", "green": "0.15 0.65 0.25 1", "blue": "0.15 0.3 0.85 1"}
22 STEP, ZSTEP, ZTOP, ZLOW, HALF = 0.04, 0.05, 0.25, 0.05, 0.02
23 MAX_STEPS = 45
24
25
26 def xml(cubes, tray):
27 cube_xml = "".join(f'<body name="{n}" pos="{x} {y} {HALF + 0.001}"><freejoint name="j_{n}"/>'
28 f'<geom type="box" size="{HALF} {HALF} {HALF}" rgba="{COLORS[n]}" mass="0.05" friction="1 0.01 0.001"/></body>'
29 for n, (x, y) in cubes.items())
30 tx, ty = tray
31 return f"""<mujoco><option timestep="0.002"/>
32 <visual><headlight ambient=".3 .3 .3" diffuse=".3 .3 .3"/><global offwidth="640" offheight="480"/><quality shadowsize="2048"/></visual>
33 <asset><texture name="g" type="2d" builtin="checker" rgb1=".72 .70 .66" rgb2=".62 .60 .56" width="256" height="256"/>
34 <material name="g" texture="g" texrepeat="6 6"/></asset>
35 <worldbody>
36 <light pos="0 0 1.6" dir="0 0 -1" diffuse=".45 .45 .45" castshadow="false"/>
37 <geom type="plane" size=".45 .45 .01" material="g"/>
38 <geom type="box" size=".012 .012 .22" pos="-.42 .42 .22" rgba=".35 .35 .4 1" group="2"/><geom type="box" size=".012 .012 .22" pos=".42 .42 .22" rgba=".35 .35 .4 1" group="2"/>
39 <geom type="box" size=".43 .012 .012" pos="0 .42 .44" rgba=".35 .35 .4 1" group="2"/>
40 <body name="tray" pos="{tx} {ty} 0">
41 <geom type="box" size=".065 .065 .003" pos="0 0 .003" rgba=".95 .8 .1 1"/>
42 <geom type="box" size=".065 .004 .02" pos="0 .061 .02" rgba=".95 .8 .1 1"/><geom type="box" size=".065 .004 .02" pos="0 -.061 .02" rgba=".95 .8 .1 1"/>
43 <geom type="box" size=".004 .065 .02" pos=".061 0 .02" rgba=".95 .8 .1 1"/><geom type="box" size=".004 .065 .02" pos="-.061 0 .02" rgba=".95 .8 .1 1"/>
44 </body>
45 {cube_xml}
46 <body name="dot" mocap="true" pos="0 0 .0015"><geom type="cylinder" size=".012 .001" rgba="0 0 0 0.85" contype="0" conaffinity="0"/></body>
47 <body name="arm" mocap="true" pos="0 0 {ZTOP}">
48 <geom type="cylinder" size=".012 .2" pos="0 0 .23" rgba=".55 .55 .6 1" contype="0" conaffinity="0" group="2"/>
49 <geom type="box" size=".035 .02 .015" rgba=".25 .25 .3 1" contype="0" conaffinity="0" group="2"/>
50 <camera name="wrist" pos="0 0 -0.02" xyaxes="1 0 0 0 1 0" fovy="70"/>
51 </body>
52 <body name="fl" mocap="true" pos="-.03 0 {ZTOP - .035}"><geom type="box" size=".005 .015 .03" rgba=".15 .15 .15 1" contype="0" conaffinity="0" group="2"/></body>
53 <body name="fr" mocap="true" pos=".03 0 {ZTOP - .035}"><geom type="box" size=".005 .015 .03" rgba=".15 .15 .15 1" contype="0" conaffinity="0" group="2"/></body>
54 <camera name="front" pos="0 -0.78 0.62" xyaxes="1 0 0 0 0.62 0.78"/>
55 <camera name="top" pos="0 0 1.05" xyaxes="1 0 0 0 1 0"/>
56 </worldbody></mujoco>"""
57
58
59 class Env:
60 def __init__(self, seed):
61 r = random.Random(seed)
62 while True: # cubes and tray apart from each other, inside reach
63 pts = [(round(r.uniform(-.3, .3), 3), round(r.uniform(-.25, .25), 3)) for _ in range(4)]
64 if all(np.hypot(a[0] - b[0], a[1] - b[1]) > .12 for i, a in enumerate(pts) for b in pts[i + 1:]):
65 break
66 self.cubes = dict(zip(COLORS, pts[:3])); self.tray = pts[3]
67 self.target = r.choice(list(COLORS))
68 self.m = mujoco.MjModel.from_xml_string(xml(self.cubes, self.tray)); self.d = mujoco.MjData(self.m)
69 self.g = np.array([round(r.uniform(-.2, .2), 2), round(r.uniform(-.2, .2), 2), ZTOP]); self.closed = False; self.held = None
70 self.ren = mujoco.Renderer(self.m, 300, 400); self.ren_top = mujoco.Renderer(self.m, 300, 300); self.ren_w = mujoco.Renderer(self.m, 300, 300)
71 self._pose(); self.settle(200)
72
73 def mocap(self, name):
74 return self.m.body_mocapid[mujoco.mj_name2id(self.m, mujoco.mjtObj.mjOBJ_BODY, name)]
75
76 def cube_q(self, n):
77 return self.m.jnt_qposadr[mujoco.mj_name2id(self.m, mujoco.mjtObj.mjOBJ_JOINT, f"j_{n}")]
78
79 def cube_pos(self, n):
80 a = self.cube_q(n); return self.d.qpos[a:a + 3].copy()
81
82 def _pose(self):
83 x, y, z = self.g; w = 0.022 if self.closed else 0.03
84 self.d.mocap_pos[self.mocap("arm")] = [x, y, z]; self.d.mocap_pos[self.mocap("dot")] = [x, y, .0015]
85 self.d.mocap_pos[self.mocap("fl")] = [x - w, y, z - .035]; self.d.mocap_pos[self.mocap("fr")] = [x + w, y, z - .035]
86 if self.held:
87 a = self.cube_q(self.held); self.d.qpos[a:a + 3] = [x, y, z - .045]; self.d.qpos[a + 3:a + 7] = [1, 0, 0, 0]
88 va = self.m.jnt_dofadr[mujoco.mj_name2id(self.m, mujoco.mjtObj.mjOBJ_JOINT, f"j_{self.held}")]; self.d.qvel[va:va + 6] = 0
89
90 def settle(self, n=300):
91 for _ in range(n):
92 self._pose(); mujoco.mj_step(self.m, self.d)
93
94 def step(self, a):
95 x, y, z = self.g
96 if a == 0: x -= STEP
97 elif a == 1: x += STEP
98 elif a == 2: y += STEP
99 elif a == 3: y -= STEP
100 elif a == 4: z = max(ZLOW, z - ZSTEP)
101 elif a == 5: z = min(ZTOP, z + ZSTEP)
102 elif a == 6:
103 self.closed = True
104 if self.held is None and z <= ZLOW + 1e-6:
105 near = [n for n in COLORS if np.hypot(*(self.cube_pos(n)[:2] - [x, y])) < 0.03 and self.cube_pos(n)[2] < .06]
106 self.held = near[0] if near else None
107 elif a == 7:
108 self.closed = False; self.held = None
109 self.g = np.clip([x, y, z], [-.38, -.32, ZLOW], [.38, .36, ZTOP]); self.settle()
110
111 def done(self):
112 p = self.cube_pos(self.target)
113 return self.held is None and abs(p[0] - self.tray[0]) < .06 and abs(p[1] - self.tray[1]) < .06 and p[2] < .06
114
115 def wrist(self):
116 opt = mujoco.MjvOption(); opt.geomgroup[2] = 0
117 self.ren_w.update_scene(self.d, camera="wrist", scene_option=opt); im = Image.fromarray(self.ren_w.render())
118 d = ImageDraw.Draw(im); c = im.width // 2
119 d.line((c - 18, c, c + 18, c), fill=(255, 255, 255), width=2); d.line((c, c - 18, c, c + 18), fill=(255, 255, 255), width=2)
120 return im
121
122 def images(self):
123 self.ren.update_scene(self.d, camera="front"); f = Image.fromarray(self.ren.render())
124 opt = mujoco.MjvOption(); opt.geomgroup[2] = 0 # hide gantry and arm column from above
125 self.ren_top.update_scene(self.d, camera="top", scene_option=opt); t = Image.fromarray(self.ren_top.render())
126 return f, t
127
128
129 def url(im):
130 b = io.BytesIO(); im.save(b, "PNG"); return "data:image/png;base64," + base64.b64encode(b.getvalue()).decode()
131
132
133 def episode(seed, frames=None):
134 env = Env(seed); hist, lat, ok = [], [], False
135 for step in range(MAX_STEPS):
136 f, t = env.images()
137 grip = ("closed, holding the " + env.held + " cube") if env.held else ("closed, holding nothing" if env.closed else "open")
138 state = ["Front camera:", {"image": url(f)}, "\nTop camera (the black dot is directly under the gripper; up in this image is forward):",
139 {"image": url(t)}, f"\nTask: pick up the {env.target} cube and put it in the yellow tray.\nGripper: {grip}.\n"
140 f"Gripper height: {'raised' if env.g[2] >= ZTOP - 1e-6 else ('lowest (at cube height)' if env.g[2] <= ZLOW + 1e-6 else 'in between')}.\n"
141 f"Last actions: {'; '.join(hist[-3:]) if hist else 'none'}"]
142 t0 = time.time()
143 r = requests.post(API, json={"kind": "choice", "state": state, "question": "Which action should the robot take next?",
144 "options": ACTIONS, "thinking": "off"}, timeout=120).json()
145 lat.append(time.time() - t0)
146 p = r["probabilities"]; a = int(np.argmax(p))
147 if frames is not None:
148 frames.append({"front": f, "top": t, "probs": p, "a": a, "step": step + 1, "ms": 1000 * lat[-1], "task": env.target, "grip": grip})
149 env.step(a); hist.append(ACTIONS[a].split(" (")[0])
150 if env.done():
151 ok = True
152 if frames is not None:
153 f, t = env.images(); frames.append({"front": f, "top": t, "probs": None, "a": None, "step": step + 1, "ms": 0, "task": env.target, "grip": "open"})
154 break
155 held_ever = any(h == "close the gripper" for h in hist)
156 return {"seed": seed, "target": env.target, "success": ok, "steps": len(hist), "actions": hist, "latency_s": lat}
157
158
159 def run(n=20):
160 res = [episode(s) for s in range(1, n + 1)]
161 lat = sorted(x for r in res for x in r["latency_s"])
162 summ = {"episodes": n, "success": sum(r["success"] for r in res) / n, "mean_steps_success": float(np.mean([r["steps"] for r in res if r["success"]] or [0])),
163 "median_ms": 1000 * lat[len(lat) // 2]}
164 json.dump({"summary": summ, "episodes": res}, open(os.path.join(HERE, "arm_results.json"), "w"), indent=1)
165 print(json.dumps(summ)); [print(r["seed"], r["target"], r["success"], r["steps"], " ".join(a.split()[0][:2] + a.split()[-1][:2] for a in r["actions"])) for r in res]
166
167
168 def frame(fr, W=1280, H=720):
169 can = Image.new("RGB", (W, H), (248, 248, 250)); d = ImageDraw.Draw(can)
170 f, fb, fs = ImageFont.truetype(FONT, 17), ImageFont.truetype(FONTB, 19), ImageFont.truetype(FONT, 15)
171 d.text((14, 12), f"{os.environ.get('JEV_LABEL', 'JEV-9B')} System 1 · robot arm: camera images → action probabilities", fill=(20, 30, 50), font=fb)
172 d.text((14, 44), f"Task: pick up the {fr['task']} cube and put it in the yellow tray · gripper: {fr['grip']}", fill=(40, 40, 40), font=f)
173 can.paste(fr["front"].resize((560, 420)), (14, 84)); can.paste(fr["top"].resize((420, 420)), (584, 84))
174 d.text((14, 510), "front camera", fill=(90, 90, 90), font=fs); d.text((584, 510), "top camera", fill=(90, 90, 90), font=fs)
175 x0 = 1018
176 if fr["probs"] is None:
177 d.text((x0, 100), "Task complete", fill=(0, 140, 60), font=fb)
178 else:
179 d.text((x0, 90), f"Step {fr['step']} · {fr['ms']:.0f} ms", fill=(20, 30, 50), font=fb)
180 y = 130
181 for i, a in enumerate(ACTIONS):
182 p = fr["probs"][i]; c = (200, 0, 0) if i == fr["a"] else (120, 140, 170)
183 d.text((x0, y), a.split(" (")[0], fill=(200, 0, 0) if i == fr["a"] else (30, 30, 30), font=fs)
184 d.rectangle((x0, y + 21, x0 + 240, y + 31), fill=(225, 228, 235)); d.rectangle((x0, y + 21, x0 + int(240 * p), y + 31), fill=c)
185 d.text((x0 + 200, y), f"{p:.2f}", fill=(30, 30, 30), font=fs); y += 46
186 return can
187
188
189 def record(seed, out):
190 frames = []; res = episode(seed, frames); tmp = tempfile.mkdtemp(); n = 0
191 for fr in frames:
192 im = frame(fr)
193 for _ in range(2 if fr["probs"] is not None else 6):
194 im.save(f"{tmp}/{n:04d}.png"); n += 1
195 subprocess.run(["ffmpeg", "-y", "-loglevel", "error", "-framerate", "4", "-i", f"{tmp}/%04d.png", "-c:v", "libx264", "-pix_fmt", "yuv420p",
196 "-vf", "fps=12", out], check=True)
197 print(out, "success" if res["success"] else "FAILED", res["steps"], "steps")
198
199
200 if __name__ == "__main__" and sys.argv[1] in ("run", "view", "record"):
201 if sys.argv[1] == "run":
202 run(int(sys.argv[2]) if len(sys.argv) > 2 else 20)
203 elif sys.argv[1] == "view":
204 e = Env(int(sys.argv[2])); f, t = e.images(); f.save("/tmp/opencode/arm_front.png"); t.save("/tmp/opencode/arm_top.png"); print(e.target, e.cubes, e.tray, e.g)
205 else:
206 record(int(sys.argv[2]), sys.argv[3])
207
208
209 # ---------------------------------------------------------------- v2: System 1 visual servoing (binary decisions)
210 ZCARRY = 0.15
211
212
213 def ask(imgs, q, opts):
214 state = []
215 for title, im in imgs:
216 state += [title, {"image": url(im)}, "\n"]
217 t0 = time.time()
218 p = requests.post(API, json={"kind": "choice", "state": state, "question": q, "options": opts, "thinking": "off"},
219 timeout=120).json()["probabilities"]
220 return p, time.time() - t0
221
222
223 def zoom(env, t, half=60):
224 """Crop of the top image centred on the black dot (projected position), upscaled 2.5x."""
225 f = (t.height / 2) / np.tan(np.radians(45 / 2)); u = t.width / 2 + f * env.g[0] / 1.05; v = t.height / 2 - f * env.g[1] / 1.05
226 u, v = int(np.clip(u, half, t.width - half)), int(np.clip(v, half, t.height - half))
227 return t.crop((u - half, v - half, u + half, v + half)).resize((300, 300), Image.BICUBIC)
228
229
230 def servo(env, what, ref, log, frames, phase, max_iter=36):
231 """Move the gripper over `what` by an adaptive bisection on each axis, driven only by two System 1 decisions per
232 step: is `what` left or right of `ref`, and above or below it, in the top camera image (plus a zoomed view of the
233 area around the gripper once the steps are small). A flipped answer halves that axis' step; three equal answers in a
234 row double it."""
235 step, hist = {"x": 0.08, "y": 0.08}, {"x": [], "y": []}
236 for it in range(max_iter):
237 f, t = env.images()
238 imgs = [("Top camera image:", t)]
239 if max(step.values()) <= 0.04:
240 imgs.append(("Zoomed top view around the black dot (2.5x):", zoom(env, t)))
241 plr, l1 = ask(imgs, f"In the top camera view, is the {what} to the left or to the right of the {ref}?", ["left", "right"])
242 pud, l2 = ask(imgs, f"In the top camera view, is the {what} above or below the {ref}?", ["above", "below"])
243 sx = -1 if plr[0] > plr[1] else 1; sy = 1 if pud[0] > pud[1] else -1
244 for ax, sgn in (("x", sx), ("y", sy)):
245 h = hist[ax]
246 if h and sgn != h[-1]:
247 step[ax] /= 2
248 elif len(h) >= 2 and h[-1] == h[-2] == sgn:
249 step[ax] = min(0.08, step[ax] * 2)
250 h.append(sgn)
251 log["decisions"] += 2; log["lat"] += [l1, l2]
252 if frames is not None:
253 frames.append({"front": f, "top": t, "phase": phase, "plr": plr, "pud": pud, "ms": 1000 * (l1 + l2), "task": env.target,
254 "grip": "holding the " + env.held + " cube" if env.held else "open"})
255 if step["x"] < 0.012 and step["y"] < 0.012:
256 return
257 env.g = np.clip(env.g + [sx * step["x"], sy * step["y"], 0], [-.38, -.32, ZLOW], [.38, .36, ZTOP]); env.settle(150)
258
259
260 def episode2(seed, frames=None):
261 env = Env(seed); log = {"decisions": 0, "lat": []}
262 servo(env, f"{env.target} cube", "black dot", log, frames, "1 · move over the cube")
263 c = env.cube_pos(env.target); log["err_cube_cm"] = round(100 * float(np.hypot(c[0] - env.g[0], c[1] - env.g[1])), 1)
264 env.g[2] = ZLOW; env.settle(); env.step(6) # lower and grasp
265 grasped = env.held == env.target
266 env.g[2] = ZCARRY; env.settle()
267 servo(env, "yellow tray", f"{env.target} cube held by the gripper", log, frames, "2 · carry it over the tray")
268 log["err_tray_cm"] = round(100 * float(np.hypot(env.tray[0] - env.g[0], env.tray[1] - env.g[1])), 1)
269 env.step(7); env.settle(600) # release
270 ok = env.done()
271 if frames is not None:
272 f, t = env.images(); frames.append({"front": f, "top": t, "phase": "done" if ok else "missed", "plr": None, "pud": None, "ms": 0,
273 "task": env.target, "grip": "open"})
274 return {"seed": seed, "target": env.target, "grasped": bool(grasped), "success": bool(ok), "decisions": log["decisions"], "latency_s": log["lat"],
275 "err_cube_cm": log.get("err_cube_cm"), "err_tray_cm": log.get("err_tray_cm")}
276
277
278 def run2(n=20):
279 t0 = time.time(); res = [episode2(s) for s in range(1, n + 1)]
280 lat = sorted(x for r in res for x in r["latency_s"])
281 summ = {"episodes": n, "success": sum(r["success"] for r in res) / n, "grasped": sum(r["grasped"] for r in res) / n,
282 "mean_decisions": float(np.mean([r["decisions"] for r in res])), "median_ms": 1000 * lat[len(lat) // 2], "seconds": time.time() - t0}
283 json.dump({"summary": summ, "episodes": res}, open(os.path.join(HERE, "arm2_results.json"), "w"), indent=1)
284 print(json.dumps(summ)); [print(r["seed"], r["target"], "grasp", r["grasped"], "success", r["success"], r["decisions"], "err cube/tray cm", r["err_cube_cm"], r["err_tray_cm"]) for r in res]
285
286
287 def frame2(fr, W=1280, H=720):
288 can = Image.new("RGB", (W, H), (248, 248, 250)); d = ImageDraw.Draw(can)
289 f, fb, fs = ImageFont.truetype(FONT, 17), ImageFont.truetype(FONTB, 19), ImageFont.truetype(FONT, 16)
290 d.text((14, 12), f"{os.environ.get('JEV_LABEL', 'JEV-9B')} System 1 · robot arm: camera image → two yes/no-style decisions per step", fill=(20, 30, 50), font=fb)
291 d.text((14, 44), f"Task: pick up the {fr['task']} cube and put it in the yellow tray · gripper: {fr['grip']}", fill=(40, 40, 40), font=f)
292 can.paste(fr["front"].resize((560, 420)), (14, 84)); can.paste(fr["top"].resize((420, 420)), (584, 84))
293 d.text((14, 510), "front camera", fill=(90, 90, 90), font=fs); d.text((584, 510), "top camera (what System 1 reads)", fill=(90, 90, 90), font=fs)
294 x0 = 1016; d.text((x0, 90), fr["phase"].replace(" · ", ". "), fill=(20, 30, 50), font=ImageFont.truetype(FONTB, 16))
295 if fr["plr"] is None:
296 d.text((x0, 130), "Task complete" if fr["phase"] == "done" else "Missed", fill=(0, 140, 60) if fr["phase"] == "done" else (200, 0, 0), font=fb)
297 else:
298 d.text((x0, 125), f"{fr['ms']:.0f} ms for 2 decisions", fill=(60, 60, 60), font=fs)
299 y = 165
300 for title, labs, p in (("target left or right?", ["left", "right"], fr["plr"]), ("target above or below?", ["above", "below"], fr["pud"])):
301 d.text((x0, y), title, fill=(30, 30, 30), font=fs); y += 26
302 for lab, v in zip(labs, p):
303 c = (200, 0, 0) if v == max(p) else (120, 140, 170)
304 d.text((x0, y), lab, fill=(30, 30, 30), font=fs); d.rectangle((x0 + 62, y + 4, x0 + 62 + int(130 * v), y + 16), fill=c)
305 d.text((x0 + 200, y), f"{v:.2f}", fill=(30, 30, 30), font=fs); y += 26
306 y += 18
307 d.text((x0, y + 10), "step size halves each time\na decision flips (bisection)", fill=(90, 90, 90), font=fs)
308 return can
309
310
311 def record2(seed, out):
312 frames = []; res = episode2(seed, frames); tmp = tempfile.mkdtemp(); n = 0
313 for fr in frames:
314 im = frame2(fr)
315 for _ in range(2 if fr["plr"] is not None else 8):
316 im.save(f"{tmp}/{n:04d}.png"); n += 1
317 subprocess.run(["ffmpeg", "-y", "-loglevel", "error", "-framerate", "4", "-i", f"{tmp}/%04d.png", "-c:v", "libx264", "-pix_fmt", "yuv420p",
318 "-vf", "fps=12", out], check=True)
319 print(out, "success" if res["success"] else "FAILED", res["decisions"], "decisions")
320
321
322 if __name__ == "__main__" and sys.argv[1] in ("run2", "record2"):
323 run2(int(sys.argv[2])) if sys.argv[1] == "run2" else record2(int(sys.argv[2]), sys.argv[3])
324