DeskForge / deskforge.py
asnassar's picture
Run GPU calls from the web agent and screenshot views as Gradio events (fixes ZeroGPU 'runs limit')
ec54602 verified
Raw History Blame Contribute Delete
13.9 kB
"""DeskForge models (and their base models), prompt, and action handling.
The model replies with one line of pyautogui code. We never exec() it: the line
is parsed with `ast` into an `Action`, and the action is replayed on the X11
display with xdotool.
"""
import spaces # MUST come before torch / transformers (ZeroGPU)
import ast
import re
import time
from dataclasses import dataclass
import torch
from PIL import Image, ImageDraw, ImageFont
from transformers import AutoModelForImageTextToText, AutoProcessor
MAX_PIXELS = 2_097_152 # training-time cap on screenshot area
SYSTEM_PROMPT = """You are a computer-use agent operating a desktop graphical interface. At each step you see the user's task, a screenshot of the current screen, and the actions you have already taken. Reply with the next action as pyautogui code and nothing else -- no explanation, no code fence, no commentary.
Coordinates are fractions of the screen, not pixels: x runs from 0.0 at the left edge to 1.0 at the right edge, y from 0.0 at the top to 1.0 at the bottom. Write both with four decimals.
These are the only actions available:
pyautogui.click(x=0.0000, y=0.0000)
pyautogui.doubleClick(x=0.0000, y=0.0000)
pyautogui.rightClick(x=0.0000, y=0.0000)
pyautogui.middleClick(x=0.0000, y=0.0000)
computer.tripleClick(x=0.0000, y=0.0000)
pyautogui.moveTo(x=0.0000, y=0.0000)
pyautogui.dragTo(x=0.0000, y=0.0000, button='left')
pyautogui.scroll(-4)
pyautogui.hscroll(4)
pyautogui.write(message='text to type')
pyautogui.press('enter')
pyautogui.hotkey(['ctrl', 'c'])
computer.wait()
computer.terminate(status='success')
To scroll at a particular place, move there first and then scroll. When the task is finished, or cannot be finished, end with computer.terminate."""
def fit_qwen(image, max_pixels=MAX_PIXELS):
"""Downscale to at most max_pixels (even sides, Lanczos), as in Qwen training."""
w, h = image.size
if w * h <= max_pixels:
return image
s = (max_pixels / (w * h)) ** 0.5
return image.resize(
(max(2, int(w * s) // 2 * 2), max(2, int(h * s) // 2 * 2)), Image.LANCZOS
)
def fit_gemma(image, max_pixels=MAX_PIXELS):
"""Downscale to at most max_pixels (bilinear, sides rounded down), as in Gemma training."""
w, h = image.size
if w * h <= max_pixels:
return image
ratio = w / h
height = (max_pixels / ratio) ** 0.5
return image.resize((max(1, int(height * ratio)), max(1, int(height))), Image.BILINEAR)
@dataclass(frozen=True)
class Family:
deskforge: str # fine-tuned model
base: str # what it was fine-tuned from, for the duel
base_label: str
fit: object
FAMILIES = {
"Qwen3.5-4B": Family("docling-project/DeskForge-Qwen3.5-4B", "Qwen/Qwen3.5-4B",
"Qwen3.5-4B (base)", fit_qwen),
"Gemma4-E4B": Family("docling-project/DeskForge-Gemma4-E4B", "google/gemma-4-E4B-it",
"Gemma4-E4B (base)", fit_gemma),
}
DEFAULT_FAMILY = "Gemma4-E4B" # its base model loses clearly, which makes the clearest duel
def user_text(instruction):
# The model was trained on single-target steps, so every instruction is
# presented as the first step of its own task.
return f"Task: {instruction}\n\nActions already taken:\n(none -- this is the first step)\n\nNext action:"
def chat(instruction):
return [
{"role": "system", "content": SYSTEM_PROMPT},
{
"role": "user",
"content": [{"type": "image"}, {"type": "text", "text": user_text(instruction)}],
},
]
def _load(name):
return AutoModelForImageTextToText.from_pretrained(name, dtype=torch.bfloat16).eval().to("cuda")
# Each base model uses its fine-tune's processor, so both sides of a duel get
# the screenshot encoded exactly as in DeskForge training (e.g. Gemma's 1,120
# image tokens), and only the weights differ.
PROCESSORS = {f: AutoProcessor.from_pretrained(fam.deskforge) for f, fam in FAMILIES.items()}
MODELS = {(f, which): _load(getattr(fam, which))
for f, fam in FAMILIES.items() for which in ("deskforge", "base")}
def _generate(family, which, instruction, screenshot, max_new_tokens):
processor, model = PROCESSORS[family], MODELS[(family, which)]
prompt = processor.apply_chat_template(
chat(instruction.strip()), tokenize=False, add_generation_prompt=True,
enable_thinking=False,
)
inputs = processor(
text=[prompt], images=[FAMILIES[family].fit(screenshot)], return_tensors="pt"
).to(model.device)
t0 = time.perf_counter()
with torch.inference_mode():
output = model.generate(**inputs, max_new_tokens=int(max_new_tokens), do_sample=False)
elapsed = time.perf_counter() - t0
action = processor.decode(
output[0, inputs["input_ids"].shape[1]:], skip_special_tokens=True
).strip()
return action, elapsed
@spaces.GPU(duration=30, size="xlarge") # four models resident (~50 GB)
def predict(instruction: str, screenshot, max_new_tokens: int = 128,
family: str = DEFAULT_FAMILY, which: str = "deskforge"):
"""Run one model once. Returns (pyautogui action code, inference seconds)."""
return _generate(family, which, instruction, screenshot.convert("RGB"), max_new_tokens)
@spaces.GPU(duration=30, size="xlarge") # a pair takes ~10 s; quota is checked against this
def predict_duel(instruction: str, screenshot, max_new_tokens: int = 128, family: str = DEFAULT_FAMILY):
"""DeskForge and its base on the same screenshot in one GPU call: {"deskforge": (code, s), "base": (code, s)}."""
screenshot = screenshot.convert("RGB")
return {which: _generate(family, which, instruction, screenshot, max_new_tokens)
for which in ("deskforge", "base")}
# --------------------------------------------------------------------------
# Parsing the model's pyautogui line into a structured action
# --------------------------------------------------------------------------
POINT_KINDS = {"click", "doubleClick", "rightClick", "middleClick", "tripleClick", "moveTo", "dragTo"}
KNOWN_KINDS = POINT_KINDS | {"scroll", "hscroll", "write", "press", "hotkey", "wait", "terminate"}
@dataclass
class Action:
kind: str # one of KNOWN_KINDS, or "unknown"
raw: str
x: float | None = None # fractional screen coordinates
y: float | None = None
amount: int = 0 # scroll clicks
text: str = "" # write / terminate status
keys: tuple = () # press / hotkey
@property
def has_point(self):
return self.x is not None and self.y is not None
def parse_action(code: str) -> Action:
"""Parse e.g. `pyautogui.click(x=0.1234, y=0.5678)` without executing it."""
# Base models often wrap the call in a code fence (```py ... ```); take the call itself.
m = re.search(r"(?:pyautogui|computer)\.\w+\([^\n]*\)", code)
line = m.group(0) if m else next((l.strip() for l in code.splitlines() if l.strip()), "")
try:
call = ast.parse(line, mode="eval").body
if not (isinstance(call, ast.Call) and isinstance(call.func, ast.Attribute)):
raise ValueError
kind = call.func.attr
args = [ast.literal_eval(a) for a in call.args]
kw = {k.arg: ast.literal_eval(k.value) for k in call.keywords}
except (SyntaxError, ValueError):
return Action("unknown", line)
if kind not in KNOWN_KINDS:
return Action("unknown", line)
a = Action(kind, line)
if kind in POINT_KINDS:
x = kw.get("x", args[0] if len(args) > 0 else None)
y = kw.get("y", args[1] if len(args) > 1 else None)
if isinstance(x, (int, float)) and isinstance(y, (int, float)):
a.x, a.y = min(max(float(x), 0.0), 1.0), min(max(float(y), 0.0), 1.0)
elif kind != "moveTo":
return Action("unknown", line)
elif kind in ("scroll", "hscroll"):
a.amount = int(kw.get("clicks", args[0] if args else 0))
elif kind == "write":
a.text = str(kw.get("message", args[0] if args else ""))
elif kind in ("press", "hotkey"):
keys = kw.get("keys", args[0] if len(args) == 1 else args)
a.keys = tuple(keys) if isinstance(keys, (list, tuple)) else (keys,)
a.keys = tuple(str(k) for k in a.keys)
elif kind == "terminate":
a.text = str(kw.get("status", args[0] if args else ""))
return a
# --------------------------------------------------------------------------
# Replaying an action with xdotool
# --------------------------------------------------------------------------
KEYSYMS = {
"enter": "Return", "return": "Return", "esc": "Escape", "escape": "Escape",
"tab": "Tab", "backspace": "BackSpace", "delete": "Delete", "del": "Delete",
"insert": "Insert", "space": "space", "up": "Up", "down": "Down",
"left": "Left", "right": "Right", "home": "Home", "end": "End",
"pageup": "Prior", "pgup": "Prior", "pagedown": "Next", "pgdn": "Next",
"ctrl": "ctrl", "control": "ctrl", "alt": "alt", "shift": "shift",
"win": "super", "super": "super", "cmd": "super", "command": "super",
"+": "plus", "-": "minus",
}
BUTTONS = {"click": ("1", 1), "doubleClick": ("1", 2), "tripleClick": ("1", 3),
"rightClick": ("3", 1), "middleClick": ("2", 1)}
def keysym(key: str) -> str:
k = key.strip().lower()
if re.fullmatch(r"f\d{1,2}", k):
return k.upper()
return KEYSYMS.get(k, key.strip())
def to_pixels(action: Action, width: int, height: int):
if not action.has_point:
return None
return min(round(action.x * width), width - 1), min(round(action.y * height), height - 1)
def xdotool_commands(action: Action, width: int, height: int) -> list[list[str]]:
"""The xdotool invocations (argv lists) that perform `action` on a width x height screen."""
point = to_pixels(action, width, height)
move = ["mousemove", str(point[0]), str(point[1])] if point else []
k = action.kind
if k in BUTTONS:
button, repeat = BUTTONS[k]
return [["xdotool", *move, "click", "--repeat", str(repeat), "--delay", "80", button]]
if k == "moveTo":
return [["xdotool", *move]] if move else []
if k == "dragTo":
return [["xdotool", "mousedown", "1", "sleep", "0.1", *move, "sleep", "0.1", "mouseup", "1"]]
if k in ("scroll", "hscroll") and action.amount:
if k == "scroll":
button = "4" if action.amount > 0 else "5"
else:
button = "7" if action.amount > 0 else "6"
return [["xdotool", "click", "--repeat", str(min(abs(action.amount), 25)), "--delay", "30", button]]
if k == "write" and action.text:
return [["xdotool", "type", "--delay", "25", "--", action.text[:500]]]
if k == "press":
return [["xdotool", "key", "--", keysym(key)] for key in action.keys]
if k == "hotkey" and action.keys:
return [["xdotool", "key", "--", "+".join(keysym(key) for key in action.keys)]]
return []
# --------------------------------------------------------------------------
# Drawing
# --------------------------------------------------------------------------
def draw_marker(image, point):
"""Draw a crosshair + circle at the predicted click point.
A white underlay keeps the marker visible on both light and dark UIs.
"""
out = image.copy()
d = ImageDraw.Draw(out)
x, y = point
r = max(10, min(out.width, out.height) // 60)
d.ellipse([x - r, y - r, x + r, y + r], outline=(255, 255, 255), width=6)
d.ellipse([x - r, y - r, x + r, y + r], outline=(255, 64, 0), width=3)
for (x0, y0, x1, y1) in [
(x - r * 1.8, y, x + r * 1.8, y), # horizontal
(x, y - r * 1.8, x, y + r * 1.8), # vertical
]:
d.line([x0, y0, x1, y1], fill=(255, 255, 255), width=9)
d.line([x0, y0, x1, y1], fill=(255, 64, 0), width=4)
return out
def annotate(screenshot, action: Action):
point = to_pixels(action, screenshot.width, screenshot.height)
return draw_marker(screenshot, point) if point else screenshot.copy()
DESKFORGE_COLOR = (210, 96, 58) # terracotta, the UI accent
BASE_COLOR = (59, 110, 190) # blue
TARGET_COLOR = (74, 178, 64) # green
def _font(size):
try:
return ImageFont.truetype("DejaVuSans-Bold.ttf", size)
except OSError:
return ImageFont.load_default()
def draw_duel(screenshot, actions: dict, target_box=None):
"""Both models' predicted points (orange = DeskForge, blue = base) and the true target box."""
out = screenshot.copy()
d = ImageDraw.Draw(out)
font = _font(13)
if target_box:
x0, y0, x1, y1 = target_box
d.rectangle([x0 - 2, y0 - 2, x1 + 2, y1 + 2], outline=(255, 255, 255), width=5)
d.rectangle([x0 - 2, y0 - 2, x1 + 2, y1 + 2], outline=TARGET_COLOR, width=3)
r = max(9, min(out.width, out.height) // 70)
for which, color, label in (("base", BASE_COLOR, "Base"), ("deskforge", DESKFORGE_COLOR, "DeskForge")):
action = actions.get(which)
point = to_pixels(action, out.width, out.height) if action else None
if not point:
continue
x, y = point
d.ellipse([x - r, y - r, x + r, y + r], outline=(255, 255, 255), width=6)
d.ellipse([x - r, y - r, x + r, y + r], outline=color, width=3)
d.line([x - r * 1.7, y, x + r * 1.7, y], fill=(255, 255, 255), width=7)
d.line([x, y - r * 1.7, x, y + r * 1.7], fill=(255, 255, 255), width=7)
d.line([x - r * 1.7, y, x + r * 1.7, y], fill=color, width=3)
d.line([x, y - r * 1.7, x, y + r * 1.7], fill=color, width=3)
tw = d.textlength(label, font=font)
lx = min(x + r + 4, out.width - tw - 8)
ly = y - r - 18 if which == "deskforge" else y + r + 2
d.rounded_rectangle([lx, ly, lx + tw + 8, ly + 17], radius=4, fill=color)
d.text((lx + 4, ly + 1), label, fill=(255, 255, 255), font=font)
return out