File size: 6,182 Bytes
8387666 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 | # /// script
# dependencies = [
# "transformers>=5.14.0",
# "peft>=0.19.0",
# "datasets>=4.0",
# "torch>=2.5",
# "torchvision>=0.20",
# "Pillow>=10.0",
# "accelerate>=1.0",
# "num2words",
# ]
# ///
"""ScreenSpot-v2 grounding accuracy for a (base or LoRA-adapted) Gemma 4.
A prediction is correct when the model's predicted click point lands inside the
ground-truth bounding box. We report overall accuracy plus a breakdown by
data_type (text vs icon) and data_source (web/mobile/desktop) — the axes that
matter for a browser agent.
# base model
python gemma4/eval_screenspot.py --model google/gemma-4-E4B-it
# our fine-tune (adapter on top of the base)
python gemma4/eval_screenspot.py --model google/gemma-4-E4B-it \
--adapter khalidFlex/gemma4-gui-agent --limit 300
"""
import argparse
import re
import torch
from datasets import load_dataset
from PIL import Image
from transformers import AutoModelForImageTextToText, AutoProcessor
# MUST byte-match gemma4/prep_data.py's SYSTEM_PROMPT (train == eval framing).
SYSTEM_PROMPT = """You are a GUI agent. You are given a task, a screenshot of the screen, and your previous actions. Complete the task by calling one or more of these Python functions:
click(x, y) left-click at coordinates
double_click(x, y) double-click at coordinates
move_mouse(x, y) move cursor without clicking
type(text) type text at the cursor
press(keys) press a key or key combo (e.g. "enter")
scroll(direction, amount) scroll "up" or "down" by amount
drag(from_coord, to_coord) drag from [x1, y1] to [x2, y2]
navigate_back() browser back
wait(seconds) wait for the screen to settle
final_answer(answer) task complete: report the result
All coordinates are normalized floats in [0, 1] — x runs left to right, y runs top to bottom.
For each step: first reason briefly inside <think></think>, then act with function calls inside <code></code>. When the task is fully complete, call final_answer with a short summary."""
USER_TEMPLATE = (
"Please generate the next move according to the UI screenshot, instruction "
"and previous actions.\n\nInstruction: {instruction}\n\nPrevious actions:\nNone"
)
NUM = re.compile(r"-?\d+\.?\d*")
def parse_point(text: str):
"""Pull the first click(x, y) (or first two numbers) as a normalized point."""
m = re.search(r"click\(\s*x\s*=\s*([-\d.]+)\s*,\s*y\s*=\s*([-\d.]+)", text)
if m:
return float(m.group(1)), float(m.group(2))
m = re.search(r"\b(?:click|double_click)\(\s*([-\d.]+)\s*,\s*([-\d.]+)", text)
if m:
return float(m.group(1)), float(m.group(2))
nums = [float(x) for x in NUM.findall(text)]
if len(nums) >= 2:
return nums[0], nums[1]
return None
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--model", default="google/gemma-4-E4B-it")
ap.add_argument("--adapter", default=None, help="LoRA adapter repo (optional)")
ap.add_argument("--dataset", default="lmms-lab/ScreenSpot-v2")
ap.add_argument("--split", default="train")
ap.add_argument("--limit", type=int, default=0, help="0 = full set")
ap.add_argument("--max-new-tokens", type=int, default=64)
args = ap.parse_args()
print(f"[eval] loading {args.model}" + (f" + adapter {args.adapter}" if args.adapter else ""))
processor = AutoProcessor.from_pretrained(args.model)
model = AutoModelForImageTextToText.from_pretrained(
args.model, dtype=torch.bfloat16, device_map="auto", attn_implementation="eager",
)
if args.adapter:
from peft import PeftModel
model = PeftModel.from_pretrained(model, args.adapter)
model = model.merge_and_unload()
model.eval()
ds = load_dataset(args.dataset, split=args.split)
if args.limit:
ds = ds.select(range(min(args.limit, len(ds))))
print(f"[eval] {len(ds)} samples")
hits = 0
by_type, by_source = {}, {}
for i, ex in enumerate(ds):
img = ex["image"].convert("RGB")
W, H = img.size
bx, by, bw, bh = ex["bbox"] # absolute pixels [x, y, w, h]
# Same single-user-turn shape as training: [image] + "SYSTEM\n\nuser".
user_text = f"{SYSTEM_PROMPT}\n\n{USER_TEMPLATE.format(instruction=ex['instruction'])}"
messages = [
{"role": "user", "content": [
{"type": "image", "image": img},
{"type": "text", "text": user_text},
]},
]
inputs = processor.apply_chat_template(
messages, add_generation_prompt=True,
tokenize=True, return_dict=True, return_tensors="pt",
).to(model.device)
with torch.no_grad():
out = model.generate(**inputs, max_new_tokens=args.max_new_tokens, do_sample=False)
gen = processor.decode(out[0][inputs["input_ids"].shape[1]:], skip_special_tokens=True)
pt = parse_point(gen)
ok = False
if pt:
px, py = pt
# Model emits normalized 0..1; if it emitted pixels, fall back to raw.
ax = px * W if px <= 1.0 else px
ay = py * H if py <= 1.0 else py
ok = (bx <= ax <= bx + bw) and (by <= ay <= by + bh)
hits += int(ok)
dt, dsrc = ex.get("data_type", "?"), ex.get("data_source", "?")
by_type.setdefault(dt, [0, 0]); by_source.setdefault(dsrc, [0, 0])
by_type[dt][0] += int(ok); by_type[dt][1] += 1
by_source[dsrc][0] += int(ok); by_source[dsrc][1] += 1
if (i + 1) % 25 == 0:
print(f" [{i+1}/{len(ds)}] running acc={hits/(i+1):.1%}")
n = len(ds)
print(f"\n[eval] ==== ScreenSpot-v2 ({args.adapter or args.model}) ====")
print(f"[eval] OVERALL: {hits}/{n} = {hits/n:.1%}\n")
print("[eval] by data_type:")
for k, (h, t) in sorted(by_type.items()):
print(f" {k:8s}: {h}/{t} = {h/t:.1%}")
print("[eval] by data_source:")
for k, (h, t) in sorted(by_source.items()):
print(f" {k:12s}: {h}/{t} = {h/t:.1%}")
if __name__ == "__main__":
main()
|