vl/demos/computer_use.py
| 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 | |