|
|
| from __future__ import annotations |
|
|
| from pathlib import Path |
| from typing import Any |
|
|
| import torch |
| import torch.nn as nn |
|
|
|
|
| MAX_CTX = 96 |
| MAX_ELEMS = 72 |
| MAX_ETXT = 28 |
|
|
|
|
| def _tok(text: str) -> list[str]: |
| text = (text or "").lower() |
| for c in "\n\t\r.,:;!?()[]{}<>/\\|\"'`~@#$%^&*+=_-": |
| text = text.replace(c, " ") |
| return [x for x in text.split() if x] |
|
|
|
|
| def _enc(text: str, n: int, vocab: dict[str, int]) -> list[int]: |
| ids = [vocab.get(t, 1) for t in _tok(text)[:n]] |
| return ids + [0] * (n - len(ids)) |
|
|
|
|
| def _state_for(element_id: str, history: list[dict[str, Any]], states: dict[str, int]) -> int: |
| clicked = typed = selected = False |
| for h in history: |
| if h.get("target_element_id") == element_id: |
| clicked = clicked or h.get("action", h.get("type")) == "click" |
| typed = typed or h.get("action", h.get("type")) == "type" |
| selected = selected or h.get("action", h.get("type")) == "select" |
| key = "_".join( |
| name for name, flag in [ |
| ("clicked", clicked), |
| ("typed", typed), |
| ("selected", selected), |
| ] if flag |
| ) or "none" |
| return states.get(key, 0) |
|
|
|
|
| class _TorchPolicy(nn.Module): |
| def __init__(self, vocab_size: int, role_size: int, state_size: int, d: int, layers: int, heads: int, action_size: int) -> None: |
| super().__init__() |
| self.emb = nn.Embedding(vocab_size, d, padding_idx=0) |
| self.role = nn.Embedding(role_size, d) |
| self.state = nn.Embedding(state_size, d) |
| self.eproj = nn.Linear(d * 3, d) |
| layer = nn.TransformerEncoderLayer(d, heads, d * 4, 0.1, batch_first=True, activation="gelu") |
| self.tr = nn.TransformerEncoder(layer, layers) |
| self.action = nn.Linear(d, action_size) |
| self.elem = nn.Linear(d, 1) |
|
|
| def _mean(self, ids: torch.Tensor) -> torch.Tensor: |
| x = self.emb(ids) |
| m = (ids != 0).float().unsqueeze(-1) |
| return (x * m).sum(1) / m.sum(1).clamp_min(1) |
|
|
| def forward(self, ctx: torch.Tensor, et: torch.Tensor, er: torch.Tensor, es: torch.Tensor, em: torch.Tensor): |
| b, e, t = et.shape |
| cv = self._mean(ctx) |
| ev = self._mean(et.reshape(b * e, t)).reshape(b, e, -1) |
| ev = self.eproj(torch.cat([ev, self.role(er), self.state(es)], -1)) |
| seq = torch.cat([cv[:, None, :], ev], 1) |
| pad = torch.cat([torch.zeros(b, 1, dtype=torch.bool, device=em.device), ~em], 1) |
| z = self.tr(seq, src_key_padding_mask=pad) |
| return self.action(z[:, 0]), self.elem(z[:, 1:]).squeeze(-1).masked_fill(~em, -1e9) |
|
|
|
|
| class LoadedUIActionPolicy: |
| def __init__(self, checkpoint_path: str | Path, device: str = "cpu") -> None: |
| self.checkpoint_path = Path(checkpoint_path) |
| self.device = torch.device(device) |
| self.ckpt = torch.load(self.checkpoint_path, map_location=self.device, weights_only=False) |
| self.vocab = self.ckpt["vocab"] |
| self.roles = self.ckpt["roles"] |
| self.states = self.ckpt["states"] |
| self.actions = self.ckpt["actions"] |
| cfg = self.ckpt["config"] |
|
|
| self.model = _TorchPolicy( |
| vocab_size=len(self.vocab), |
| role_size=len(self.roles), |
| state_size=len(self.states), |
| d=cfg["d"], |
| layers=cfg["layers"], |
| heads=cfg["heads"], |
| action_size=len(self.actions), |
| ).to(self.device) |
| self.model.load_state_dict(self.ckpt["state_dict"]) |
| self.model.eval() |
|
|
| def predict_raw( |
| self, |
| goal: str, |
| elements: list[dict[str, Any]], |
| history: list[dict[str, Any]] | None = None, |
| step_index: int = 0, |
| ) -> dict[str, Any]: |
| history = history or [] |
| elems = elements[:MAX_ELEMS] |
|
|
| ctx_text = " ".join( |
| [f"step {step_index}", goal] |
| + [ |
| f"{h.get('action', h.get('type', ''))} {h.get('target_element_id', '')} {h.get('text', h.get('value', ''))}" |
| for h in history |
| ] |
| ) |
|
|
| et, er, es, em = [], [], [], [] |
| for i in range(MAX_ELEMS): |
| if i < len(elems): |
| e = elems[i] |
| et.append(_enc(" ".join([ |
| str(e.get("role", "")), |
| str(e.get("name", "")), |
| str(e.get("text", "")), |
| str(e.get("value", "")), |
| str(e.get("section", "")), |
| ]), MAX_ETXT, self.vocab)) |
| er.append(self.roles.get(str(e.get("role", "")), 0)) |
| es.append(_state_for(str(e.get("element_id", "")), history, self.states)) |
| em.append(bool(e.get("visible", True) and e.get("enabled", True))) |
| else: |
| et.append([0] * MAX_ETXT) |
| er.append(0) |
| es.append(0) |
| em.append(False) |
|
|
| with torch.no_grad(): |
| al, el = self.model( |
| torch.tensor([_enc(ctx_text, MAX_CTX, self.vocab)], device=self.device), |
| torch.tensor([et], device=self.device), |
| torch.tensor([er], device=self.device), |
| torch.tensor([es], device=self.device), |
| torch.tensor([em], device=self.device).bool(), |
| ) |
| action_probs = torch.softmax(al, dim=-1)[0].detach().cpu() |
| elem_probs = torch.softmax(el, dim=-1)[0][:len(elems)].detach().cpu() |
|
|
| action_order = [self.actions[i] for i in torch.argsort(action_probs, descending=True).tolist()] |
| elem_order = torch.argsort(elem_probs, descending=True).tolist() |
|
|
| compatible = { |
| "select": {"combobox"}, |
| "type": {"textbox"}, |
| "click": {"button", "link", "checkbox", "row", "tab", "menuitem"}, |
| } |
|
|
| selected_action = "done" |
| selected_idx = None |
| for action in action_order: |
| if action in {"done", "wait", "press"}: |
| selected_action = action |
| selected_idx = None |
| break |
| for idx in elem_order: |
| if idx < len(elems): |
| e = elems[idx] |
| if ( |
| e.get("visible", True) |
| and e.get("enabled", True) |
| and str(e.get("role", "")) in compatible.get(action, set()) |
| ): |
| selected_action = action |
| selected_idx = idx |
| break |
| if selected_idx is not None: |
| break |
|
|
| target_id = elems[selected_idx]["element_id"] if selected_idx is not None else None |
| confidence = float(action_probs[self.actions.index(selected_action)]) |
| if selected_idx is not None: |
| confidence *= float(elem_probs[selected_idx]) |
|
|
| return { |
| "action": selected_action, |
| "target_element_id": target_id, |
| "confidence": confidence, |
| "top_actions": [ |
| {"action": self.actions[i], "prob": float(action_probs[i])} |
| for i in torch.argsort(action_probs, descending=True).tolist()[:6] |
| ], |
| "top_elements": [ |
| { |
| "element_id": elems[i]["element_id"], |
| "role": elems[i].get("role", ""), |
| "name": elems[i].get("name", ""), |
| "text": elems[i].get("text", ""), |
| "prob": float(elem_probs[i]), |
| } |
| for i in elem_order[: min(8, len(elems))] |
| ], |
| } |
|
|