vl/demos/computer_use.py
9.7 KB · 179 lines · python Raw
1 """Computer-use click loop with JEV-9B-VL System 1 (POST /v1/decide), in a real headless Chromium (Playwright).
2
3 Every step: screenshot -> every interactive element gets a numbered red box (Set-of-Mark) -> one System 1 `choice`
4 over "click element k" for each mark plus "the task is complete" -> the browser really clicks the chosen element.
5 Variants: "marks" (the options are only the mark numbers: the model must read the screenshot to know what each mark
6 is) and "marks+text" (each option also carries the element's visible text, like an accessibility tree).
7 Tasks are randomised per seed (which email / product / colour / size / setting; element order and contents shuffled).
8
9 python3 computer_use.py run # 3 apps x N seeds x 2 variants -> cu_results.json
10 python3 computer_use.py record APP SEED VARIANT OUT.mp4
11 """
12 import base64, io, json, os, sys, time, subprocess, tempfile
13 import requests
14 from PIL import Image, ImageDraw, ImageFont
15
16 HERE = os.path.dirname(os.path.abspath(__file__))
17 URL = "file://" + os.path.join(HERE, "webapp.html")
18 API = os.environ.get("JEV_URL", "http://localhost:8000") + "/v1/decide"
19 VW, VH = 1100, 640
20 FONT = "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf"
21 FONTB = "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf"
22 MAX_STEPS = 10
23
24
25 def task_for(app, world, seed):
26 import random
27 r = random.Random(seed * 7919 + len(app))
28 if app == "mail":
29 m = r.choice(world["MAILS"])
30 return f"Archive the email from {m['from']} about the {m['subject']}.", {"mail": m["id"]}
31 if app == "shop":
32 p, c, s = r.choice(world["PRODUCTS"]), r.choice(world["COLORS"]), r.choice(world["SIZES"])
33 return f"Buy a {c} {p['name']} in size {s}: add it to the cart and check out.", {"name": p["name"], "color": c, "size": s}
34 sec = r.choice(list(world["SECTIONS"])); t = r.choice(world["SECTIONS"][sec])
35 return f"Turn on \"{t}\" in the {sec} settings and save the change.", {"toggle": t}
36
37
38 def goal_met(app, S, g):
39 if app == "mail":
40 return S["archived"] == [g["mail"]] and not S["deleted"]
41 if app == "shop":
42 return bool(S["ordered"]) and S["cart"] == [{"name": g["name"], "color": g["color"], "size": g["size"]}]
43 on = [t for t, v in (S["saved"] or {}).items() if v]
44 return on == [g["toggle"]]
45
46
47 ELEMS_JS = """() => [...document.querySelectorAll('[data-click]')].map(e => { const r = e.getBoundingClientRect();
48 return {x: r.x, y: r.y, w: r.width, h: r.height, text: (e.innerText || '').trim().replace(/\\s+/g, ' ').slice(0, 60),
49 kind: e.classList.contains('swatch') ? 'colour swatch' : e.classList.contains('sw') ? 'switch' : e.classList.contains('row') ? 'list row' : e.classList.contains('card') ? 'product card' : 'button'}; })
50 .filter(e => e.w > 0 && e.h > 0 && e.y < innerHeight)"""
51
52
53 def mark(png, elems):
54 im = Image.open(io.BytesIO(png)).convert("RGB"); d = ImageDraw.Draw(im); f = ImageFont.truetype(FONTB, 15)
55 for i, e in enumerate(elems):
56 d.rectangle((e["x"] - 2, e["y"] - 2, e["x"] + e["w"] + 2, e["y"] + e["h"] + 2), outline=(230, 0, 0), width=2)
57 lab = str(i + 1); tw = d.textlength(lab, font=f)
58 ly = e["y"] - 20 if e["y"] >= 22 else e["y"] + e["h"] + 2 # label above the box, or below it at the top edge
59 d.rectangle((e["x"] - 2, ly, e["x"] + tw + 6, ly + 18), fill=(230, 0, 0))
60 d.text((e["x"] + 2, ly + 1), lab, fill="white", font=f)
61 return im
62
63
64 def to_url(im):
65 b = io.BytesIO(); im.save(b, "PNG"); return "data:image/png;base64," + base64.b64encode(b.getvalue()).decode()
66
67
68 def episode(page, app, seed, variant, frames=None):
69 page.goto(f"{URL}?app={app}&seed={seed}"); page.wait_for_timeout(50)
70 world = page.evaluate("() => window.WORLD")
71 task, goal = task_for(app, world, seed)
72 hist, steps, stopped, lat = [], [], False, []
73 for step in range(MAX_STEPS):
74 elems = page.evaluate(ELEMS_JS)
75 im = mark(page.screenshot(), elems)
76 if variant == "marks":
77 opts = [f"click element {i + 1}" for i in range(len(elems))]
78 else:
79 opts = [f"click element {i + 1}: {e['text'] or e['kind']}" + (f" ({e['kind']})" if e["text"] else "") for i, e in enumerate(elems)]
80 opts.append("the task is complete")
81 state = ["Browser screenshot. Every clickable element has a numbered red box.\n", {"image": to_url(im)},
82 f"\nTask: {task}\nActions taken so far: {'; '.join(hist) if hist else 'none'}"]
83 t0 = time.time()
84 r = requests.post(API, json={"kind": "choice", "state": state, "question": "What should be done next to complete the task?",
85 "options": opts, "thinking": "off"}, timeout=120).json()
86 lat.append(time.time() - t0)
87 probs = r["probabilities"]; k = max(range(len(probs)), key=probs.__getitem__)
88 if frames is not None:
89 frames.append({"img": im, "task": task, "step": step + 1, "opts": opts, "probs": probs, "k": k, "ms": 1000 * lat[-1], "variant": variant})
90 steps.append({"options": len(opts), "choice": opts[k], "p": probs[k]})
91 if k == len(elems):
92 stopped = True
93 break
94 e = elems[k]
95 hist.append(f"clicked element {k + 1}" + (f" ({e['text'] or e['kind']})" if variant != "marks" else ""))
96 page.mouse.click(e["x"] + e["w"] / 2, e["y"] + e["h"] / 2); page.wait_for_timeout(30)
97 S = page.evaluate("() => window.S")
98 ok = goal_met(app, S, goal)
99 return {"app": app, "seed": seed, "variant": variant, "task": task, "success": bool(ok and stopped), "goal_met": bool(ok),
100 "stopped": stopped, "steps": len(steps), "trace": steps, "latency_s": lat}
101
102
103 def worker(jobs):
104 from playwright.sync_api import sync_playwright
105 out = []
106 with sync_playwright() as p:
107 b = p.chromium.launch(); page = b.new_page(viewport={"width": VW, "height": VH})
108 for app, seed, variant in jobs:
109 out.append(episode(page, app, seed, variant))
110 b.close()
111 return out
112
113
114 def run(n=10, procs=6):
115 from multiprocessing import Pool
116 jobs = [(a, s, v) for v in ("marks", "marks+text") for a in ("mail", "shop", "settings") for s in range(1, n + 1)]
117 chunks = [jobs[i::procs] for i in range(procs)]
118 t0 = time.time()
119 with Pool(procs) as pool:
120 res = [r for part in pool.map(worker, chunks) for r in part]
121 summ = {}
122 for v in ("marks", "marks+text"):
123 for a in ("mail", "shop", "settings", "all"):
124 rs = [r for r in res if r["variant"] == v and (a == "all" or r["app"] == a)]
125 lat = sorted(x for r in rs for x in r["latency_s"])
126 summ[f"{v}/{a}"] = {"episodes": len(rs), "success": sum(r["success"] for r in rs) / len(rs),
127 "mean_steps": sum(r["steps"] for r in rs) / len(rs), "median_ms": 1000 * lat[len(lat) // 2]}
128 json.dump({"summary": summ, "seconds": time.time() - t0, "episodes": res}, open(os.path.join(HERE, "cu_results.json"), "w"), indent=1)
129 for k, v in summ.items():
130 print(f"{k:22s} success {v['success']:.0%} steps {v['mean_steps']:.1f} median {v['median_ms']:.0f} ms/decision (n={v['episodes']})")
131
132
133 def panel(fr, W=1280, H=720):
134 """Video frame: marked screenshot on the left, task + option probabilities on the right."""
135 can = Image.new("RGB", (W, H), (250, 250, 252)); d = ImageDraw.Draw(can)
136 f, fb, fs = ImageFont.truetype(FONT, 17), ImageFont.truetype(FONTB, 19), ImageFont.truetype(FONT, 15)
137 shot = fr["img"].resize((880, int(880 * VH / VW))); can.paste(shot, (10, 80))
138 d.text((14, 12), f"{os.environ.get('JEV_LABEL', 'JEV-9B')} System 1 · computer use: screenshot → which element to click", fill=(20, 30, 50), font=fb)
139 d.text((14, 44), f"Task: {fr['task']}", fill=(40, 40, 40), font=f)
140 x0 = 905; d.text((x0, 86), f"Step {fr['step']} · one forward pass · {fr['ms']:.0f} ms", fill=(20, 30, 50), font=fb)
141 order = sorted(range(len(fr["probs"])), key=lambda i: -fr["probs"][i])[:6]
142 y = 126
143 for i in order:
144 p = fr["probs"][i]; lab = fr["opts"][i].replace("click element", "click").replace("the task is complete", "task complete")
145 lab = lab if len(lab) <= 34 else lab[:33] + "…"
146 d.text((x0, y), lab, fill=(30, 30, 30) if i != fr["k"] else (200, 0, 0), font=fs)
147 d.rectangle((x0, y + 21, x0 + 340, y + 33), fill=(225, 228, 235))
148 d.rectangle((x0, y + 21, x0 + int(340 * p), y + 33), fill=(200, 0, 0) if i == fr["k"] else (120, 140, 170))
149 d.text((x0 + 300, y), f"{p:.2f}", fill=(30, 30, 30), font=fs)
150 y += 50
151 note = ("options = numbered marks only\n(the model must read the screenshot)" if fr.get("variant") == "marks" else
152 "options = mark number + visible text;\ncolour swatches and switches have no\ntext: the screenshot decides")
153 d.text((x0, H - 80), note, fill=(90, 90, 90), font=fs)
154 return can
155
156
157 def record(app, seed, variant, out):
158 from playwright.sync_api import sync_playwright
159 frames = []
160 with sync_playwright() as p:
161 b = p.chromium.launch(); page = b.new_page(viewport={"width": VW, "height": VH})
162 res = episode(page, app, seed, variant, frames); b.close()
163 tmp = tempfile.mkdtemp()
164 n = 0
165 for fr in frames:
166 im = panel(fr)
167 for _ in range(3): # 1.5 s per decision at 2 fps
168 im.save(f"{tmp}/{n:04d}.png"); n += 1
169 subprocess.run(["ffmpeg", "-y", "-loglevel", "error", "-framerate", "2", "-i", f"{tmp}/%04d.png", "-c:v", "libx264", "-pix_fmt", "yuv420p",
170 "-vf", "fps=10", out], check=True)
171 print(out, "success" if res["success"] else "FAILED", res["steps"], "steps", [round(1000 * x) for x in res["latency_s"]], "ms")
172
173
174 if __name__ == "__main__":
175 if sys.argv[1] == "run":
176 run(int(sys.argv[2]) if len(sys.argv) > 2 else 10)
177 else:
178 record(sys.argv[2], int(sys.argv[3]), sys.argv[4], sys.argv[5])
179