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()