Snapshot: full project state
This commit is contained in:
205
bench.py
Normal file
205
bench.py
Normal file
@@ -0,0 +1,205 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Vision model benchmark: Ollama vision models vs Android screenshots.
|
||||
Ground truth = uiautomator XML dumps. Scores: identify, describe, read, ground, agent JSON.
|
||||
"""
|
||||
import base64, json, time, os, sys, re
|
||||
import xml.etree.ElementTree as ET
|
||||
import urllib.request
|
||||
|
||||
HOST = "http://10.30.20.128:11434"
|
||||
BENCH = os.path.expanduser("~/vision_bench")
|
||||
|
||||
MODELS = [
|
||||
"moondream:1.8b",
|
||||
"granite3.2-vision:2b",
|
||||
"qwen2.5vl:3b",
|
||||
"llava-phi3:3.8b",
|
||||
"gemma3:4b",
|
||||
"llava:7b",
|
||||
"qwen2.5vl:7b",
|
||||
"minicpm-v4.5:latest",
|
||||
]
|
||||
|
||||
# name, image, xml (None = no ground truth, manual grade)
|
||||
SITUATIONS = [
|
||||
("s1_dialog", "s1_dialog.png", "ui1.xml"),
|
||||
("s3_chrome", "s3_browser.png", "ui3.xml"),
|
||||
("s_home", "s_home.png", "s_home.xml"),
|
||||
("s_settings", "s_settings.png", "s_settings.xml"),
|
||||
("s_article", "s_article.png", "s_article.xml"),
|
||||
("s_calc", "s_calc.png", "s_calc.xml"),
|
||||
("s_keyboard", "s_keyboard.png", "s_keyboard.xml"),
|
||||
("s_drawer", "s_drawer.png", "s_drawer.xml"),
|
||||
]
|
||||
|
||||
FRIENDLY = {
|
||||
"android": "Android system dialog (Quickstep launcher chooser)",
|
||||
"com.android.chrome": "Chrome browser",
|
||||
"com.android.launcher3": "Home screen launcher",
|
||||
"com.android.settings": "Settings app",
|
||||
}
|
||||
|
||||
def parse_xml(path):
|
||||
tree = ET.parse(path)
|
||||
root = tree.getroot()
|
||||
nodes = [n for n in root.iter("node")]
|
||||
pkg = nodes[0].get("package") if nodes else ""
|
||||
texts = [n.get("text") for n in nodes if n.get("text")]
|
||||
descs = [n.get("content-desc") for n in nodes if n.get("content-desc")]
|
||||
clickable = []
|
||||
for n in nodes:
|
||||
if n.get("clickable") == "true":
|
||||
label = n.get("text") or n.get("content-desc")
|
||||
if label and n.get("bounds"):
|
||||
b = n.get("bounds")
|
||||
m = re.match(r"\[(\d+),(\d+)\]\[(\d+),(\d+)\]", b)
|
||||
if m:
|
||||
x1, y1, x2, y2 = map(int, m.groups())
|
||||
clickable.append({"label": label, "center": [(x1+x2)//2, (y1+y2)//2]})
|
||||
return {"package": pkg, "texts": texts, "descs": descs, "clickable": clickable}
|
||||
|
||||
def call(model, prompt, img_path, keep_alive="10m"):
|
||||
with open(img_path, "rb") as f:
|
||||
b64 = base64.b64encode(f.read()).decode()
|
||||
payload = {
|
||||
"model": model, "stream": False, "keep_alive": keep_alive,
|
||||
"options": {"temperature": 0, "num_predict": 220},
|
||||
"messages": [{"role": "user", "content": prompt, "images": [b64]}],
|
||||
}
|
||||
req = urllib.request.Request(HOST + "/api/chat", data=json.dumps(payload).encode(),
|
||||
headers={"Content-Type": "application/json"})
|
||||
opener = urllib.request.build_opener(urllib.request.ProxyHandler({})) # LAN: bypass any system proxy
|
||||
t0 = time.time()
|
||||
try:
|
||||
with opener.open(req, timeout=240) as r:
|
||||
d = json.loads(r.read())
|
||||
except Exception as e:
|
||||
return {"error": str(e), "wall_s": round(time.time()-t0, 2)}
|
||||
msg = d.get("message", {})
|
||||
content = msg.get("content", "")
|
||||
if not content:
|
||||
content = msg.get("thinking", "")
|
||||
return {
|
||||
"wall_s": round(time.time()-t0, 2),
|
||||
"load_s": round(d.get("load_duration", 0)/1e9, 2),
|
||||
"prompt_eval_s": round(d.get("prompt_eval_duration", 0)/1e9, 2),
|
||||
"eval_s": round(d.get("eval_duration", 0)/1e9, 2),
|
||||
"eval_count": d.get("eval_count", 0),
|
||||
"tok_s": round(d.get("eval_count", 0) / max(d.get("eval_duration", 1)/1e9, 0.01), 1),
|
||||
"content": content.strip(),
|
||||
"done": d.get("done_reason", "?"),
|
||||
}
|
||||
|
||||
def score_identify(gt, answer):
|
||||
if not gt: return None
|
||||
target = gt["package"]
|
||||
keys = []
|
||||
if target == "android": keys = ["dialog", "quickstep", "launcher", "home"]
|
||||
elif "chrome" in target: keys = ["chrome", "browser", "welcome", "terms"]
|
||||
elif "settings" in target: keys = ["settings"]
|
||||
hit = any(k in answer.lower() for k in keys)
|
||||
return 1.0 if hit else 0.0
|
||||
|
||||
def score_read(gt, answer):
|
||||
if not gt: return None
|
||||
labels = [c["label"] for c in gt["clickable"]]
|
||||
if not labels: return None
|
||||
hits = sum(1 for l in labels if l.lower() in answer.lower())
|
||||
return hits / len(labels)
|
||||
|
||||
def score_ground(gt, answer, W=1024, H=720):
|
||||
if not gt or not gt["clickable"]: return None
|
||||
# parse JSON with coords from answer
|
||||
m = re.search(r"\{[^{}]*\"x\"\s*:\s*(\d+)[^{}]*\"y\"\s*:\s*(\d+)[^{}]*\}", answer.replace("'", '"'))
|
||||
if not m:
|
||||
m = re.search(r"\(?(\d{2,4})\s*[,x]\s*(\d{2,4})\)?", answer)
|
||||
if not m: return 0.0
|
||||
x, y = int(m.group(1)), int(m.group(2))
|
||||
tx, ty = gt["clickable"][0]["center"]
|
||||
dist = ((x-tx)**2 + (y-ty)**2) ** 0.5
|
||||
# score 1.0 within 5% of diagonal, linear decay to 30%
|
||||
diag = (W*W + H*H) ** 0.5
|
||||
return max(0.0, round(1.0 - (dist / (0.3*diag)), 2))
|
||||
|
||||
def score_agent(gt, answer):
|
||||
try:
|
||||
j = json.loads(answer.strip().strip("`"))
|
||||
if not isinstance(j, dict) or "elements" not in j: return {"parse": 0.0, "cover": None}
|
||||
elems = j["elements"]
|
||||
cover = None
|
||||
if gt and gt["clickable"]:
|
||||
labels = [c["label"].lower() for c in gt["clickable"]]
|
||||
hits = sum(1 for l in labels if any(l in str(e).lower() for e in elems))
|
||||
cover = hits / len(labels)
|
||||
return {"parse": 1.0, "cover": cover}
|
||||
except Exception:
|
||||
return {"parse": 0.0, "cover": None}
|
||||
|
||||
def vram(model):
|
||||
try:
|
||||
opener = urllib.request.build_opener(urllib.request.ProxyHandler({}))
|
||||
with opener.open(HOST + "/api/ps", timeout=10) as r:
|
||||
ps = json.loads(r.read())
|
||||
for m in ps.get("models", []):
|
||||
if m.get("name") == model:
|
||||
return {"size_gb": round(m.get("size", 0)/1e9, 2),
|
||||
"vram_gb": round(m.get("size_vram", 0)/1e9, 2)}
|
||||
except Exception:
|
||||
pass
|
||||
return None
|
||||
|
||||
def unload(model):
|
||||
try:
|
||||
opener = urllib.request.build_opener(urllib.request.ProxyHandler({}))
|
||||
req = urllib.request.Request(HOST + "/api/generate",
|
||||
data=json.dumps({"model": model, "keep_alive": 0}).encode(),
|
||||
headers={"Content-Type": "application/json"})
|
||||
opener.open(req, timeout=30).read()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def main():
|
||||
out_path = os.path.join(BENCH, "bench_results.jsonl")
|
||||
gts = {}
|
||||
for name, img, xml in SITUATIONS:
|
||||
gts[name] = parse_xml(os.path.join(BENCH, xml)) if xml else None
|
||||
|
||||
for mi, model in enumerate(MODELS):
|
||||
print(f"\n===== MODEL {mi+1}/{len(MODELS)}: {model} =====", flush=True)
|
||||
for name, img, _ in SITUATIONS:
|
||||
gt = gts[name]
|
||||
img_path = os.path.join(BENCH, img)
|
||||
tasks = {
|
||||
"identify": "What application or system screen is shown in this Android screenshot? Answer in one short line.",
|
||||
"describe": "Describe this Android screenshot in detail. List every visible UI element and all text you can read.",
|
||||
"ground": f'The button with label "{gt["clickable"][0]["label"] if gt and gt["clickable"] else "the main button"}" is on this 1024x720 screen. Output ONLY valid JSON: {{"x": <int>, "y": <int>}} giving the pixel center of that button.',
|
||||
"agent": 'You are a UI automation agent looking at an Android screenshot (1024x720). Output ONLY valid JSON, no commentary: {"summary": "<one line>", "elements": [{"label": "<text>", "role": "button|text", "center": [x, y]}], "next_action": "<which button a bot should tap>"}',
|
||||
}
|
||||
for tname, prompt in tasks.items():
|
||||
r = call(model, prompt, img_path)
|
||||
rec = {"model": model, "situation": name, "task": tname,
|
||||
"wall_s": r.get("wall_s"), "load_s": r.get("load_s"),
|
||||
"eval_s": r.get("eval_s"), "tok_s": r.get("tok_s"),
|
||||
"eval_count": r.get("eval_count"), "done": r.get("done"),
|
||||
"content": (r.get("content") or "")[:600],
|
||||
"error": r.get("error")}
|
||||
if tname == "identify": rec["score"] = score_identify(gt, r.get("content", ""))
|
||||
elif tname == "read": rec["score"] = score_read(gt, r.get("content", ""))
|
||||
elif tname == "ground": rec["score"] = score_ground(gt, r.get("content", ""))
|
||||
elif tname == "agent":
|
||||
s = score_agent(gt, r.get("content", ""))
|
||||
rec["score"] = s["parse"] if s["cover"] is None else round(0.5*s["parse"] + 0.5*s["cover"], 2)
|
||||
elif tname == "describe": rec["score"] = score_read(gt, r.get("content", ""))
|
||||
with open(out_path, "a") as f:
|
||||
f.write(json.dumps(rec) + "\n")
|
||||
print(f" {name}/{tname}: wall={r.get('wall_s')}s eval={r.get('eval_s')}s tok/s={r.get('tok_s')} score={rec.get('score')} err={r.get('error')}", flush=True)
|
||||
vr = vram(model)
|
||||
print(f" VRAM: {vr}", flush=True)
|
||||
with open(out_path, "a") as f:
|
||||
f.write(json.dumps({"model": model, "vram": vr, "meta": True}) + "\n")
|
||||
unload(model)
|
||||
time.sleep(2)
|
||||
print("\nBENCH_DONE", flush=True)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user