Optimize CPU feature encoding and add bounded batch inference with measured parity
Browse files- README.md +2 -0
- baim/features.py +84 -62
- baim/policy.py +22 -8
- docs/RUNTIME_V002.md +28 -0
- pyproject.toml +1 -1
- reports/runtime-optimization.json +34 -0
- tests/test_batch.py +43 -0
README.md
CHANGED
|
@@ -11,6 +11,8 @@ library_name: baim
|
|
| 11 |
|
| 12 |
# Devils Agent / BAIM — experimental research checkpoints
|
| 13 |
|
|
|
|
|
|
|
| 14 |
**Research prototype, not a production-ready general browser agent.** Four custom
|
| 15 |
checkpoints are stored under `models/`. They require the accompanying Python code;
|
| 16 |
this repository is not a standard Transformers `AutoModel` or hosted-inference
|
|
|
|
| 11 |
|
| 12 |
# Devils Agent / BAIM — experimental research checkpoints
|
| 13 |
|
| 14 |
+
Runtime v0.0.2: **2.38x faster median synthetic CPU prediction**, bounded microbatch inference, and unchanged checkpoint weights. See [measured results and limits](docs/RUNTIME_V002.md).
|
| 15 |
+
|
| 16 |
**Research prototype, not a production-ready general browser agent.** Four custom
|
| 17 |
checkpoints are stored under `models/`. They require the accompanying Python code;
|
| 18 |
this repository is not a standard Transformers `AutoModel` or hosted-inference
|
baim/features.py
CHANGED
|
@@ -1,62 +1,84 @@
|
|
| 1 |
-
"""Small vocabulary and deterministic candidate retrieval, without site rules."""
|
| 2 |
-
from collections import Counter
|
| 3 |
-
import re
|
| 4 |
-
import torch
|
| 5 |
-
|
| 6 |
-
KINDS = ('C', 'T', 'O')
|
| 7 |
-
ROLES = {'button', 'link', 'textbox', 'searchbox', 'combobox', 'checkbox', 'radio', 'listbox', 'spinbutton'}
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
def words(text):
|
| 11 |
-
return re.findall(r'\w+', text.casefold(), flags=re.UNICODE)
|
| 12 |
-
|
| 13 |
-
|
| 14 |
-
def fit_vocab(rows, limit=4096):
|
| 15 |
-
counts = Counter()
|
| 16 |
-
for row in rows:
|
| 17 |
-
counts.update(words(row['goal']))
|
| 18 |
-
for element in row['elements']:
|
| 19 |
-
counts.update(words(element['role'] + ' ' + element['name']))
|
| 20 |
-
return {'<pad>':0, '<unk>':1, **{token:i+2 for i, (token, _) in enumerate(counts.most_common(limit-2))}}
|
| 21 |
-
|
| 22 |
-
|
| 23 |
-
def lexical(goal, element):
|
| 24 |
-
goal_words, name_words = set(words(goal)), set(words(element['name']))
|
| 25 |
-
overlap = len(goal_words & name_words)
|
| 26 |
-
return [overlap / max(1,len(name_words)), overlap / max(1,len(goal_words)),
|
| 27 |
-
float(element['name'].casefold() in goal.casefold()),
|
| 28 |
-
float(element.get('visible', True)), float(element.get('enabled', True))]
|
| 29 |
-
|
| 30 |
-
|
| 31 |
-
def candidates(goal, elements, limit=40):
|
| 32 |
-
eligible = [i for i, e in enumerate(elements) if e.get('visible', True) and e.get('enabled', True)
|
| 33 |
-
and not e.get('sensitive', False) and e['role'] in ROLES]
|
| 34 |
-
return sorted(eligible, key=lambda i: (-lexical(goal,elements[i])[0], i))[:limit]
|
| 35 |
-
|
| 36 |
-
|
| 37 |
-
def encode(rows, vocab, max_tokens=24, max_candidates=40):
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Small vocabulary and deterministic candidate retrieval, without site rules."""
|
| 2 |
+
from collections import Counter
|
| 3 |
+
import re
|
| 4 |
+
import torch
|
| 5 |
+
|
| 6 |
+
KINDS = ('C', 'T', 'O')
|
| 7 |
+
ROLES = {'button', 'link', 'textbox', 'searchbox', 'combobox', 'checkbox', 'radio', 'listbox', 'spinbutton'}
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
def words(text):
|
| 11 |
+
return re.findall(r'\w+', text.casefold(), flags=re.UNICODE)
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
def fit_vocab(rows, limit=4096):
|
| 15 |
+
counts = Counter()
|
| 16 |
+
for row in rows:
|
| 17 |
+
counts.update(words(row['goal']))
|
| 18 |
+
for element in row['elements']:
|
| 19 |
+
counts.update(words(element['role'] + ' ' + element['name']))
|
| 20 |
+
return {'<pad>':0, '<unk>':1, **{token:i+2 for i, (token, _) in enumerate(counts.most_common(limit-2))}}
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
def lexical(goal, element):
|
| 24 |
+
goal_words, name_words = set(words(goal)), set(words(element['name']))
|
| 25 |
+
overlap = len(goal_words & name_words)
|
| 26 |
+
return [overlap / max(1,len(name_words)), overlap / max(1,len(goal_words)),
|
| 27 |
+
float(element['name'].casefold() in goal.casefold()),
|
| 28 |
+
float(element.get('visible', True)), float(element.get('enabled', True))]
|
| 29 |
+
|
| 30 |
+
|
| 31 |
+
def candidates(goal, elements, limit=40):
|
| 32 |
+
eligible = [i for i, e in enumerate(elements) if e.get('visible', True) and e.get('enabled', True)
|
| 33 |
+
and not e.get('sensitive', False) and e['role'] in ROLES]
|
| 34 |
+
return sorted(eligible, key=lambda i: (-lexical(goal,elements[i])[0], i))[:limit]
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
def encode(rows, vocab, max_tokens=24, max_candidates=40):
|
| 38 |
+
# Tokenize each goal once and construct six tensors per batch, not per element.
|
| 39 |
+
if max_tokens < 1 or max_candidates < 1:
|
| 40 |
+
raise ValueError('Token and candidate limits must be positive')
|
| 41 |
+
maps, prepared = [], []
|
| 42 |
+
for row in rows:
|
| 43 |
+
goal_words = words(row['goal'])
|
| 44 |
+
goal_set, folded = set(goal_words), row['goal'].casefold()
|
| 45 |
+
eligible = []
|
| 46 |
+
for i, e in enumerate(row['elements']):
|
| 47 |
+
if not (e.get('visible', True) and e.get('enabled', True)) or e.get('sensitive', False) or e['role'] not in ROLES:
|
| 48 |
+
continue
|
| 49 |
+
name_set = set(words(e['name']))
|
| 50 |
+
overlap = len(goal_set & name_set)
|
| 51 |
+
feature = [overlap/max(1,len(name_set)), overlap/max(1,len(goal_set)),
|
| 52 |
+
float(e['name'].casefold() in folded), float(e.get('visible',True)), float(e.get('enabled',True))]
|
| 53 |
+
eligible.append((i, feature))
|
| 54 |
+
eligible.sort(key=lambda item: (-item[1][0], item[0]))
|
| 55 |
+
selected = eligible[:max_candidates]
|
| 56 |
+
maps.append([i for i, _ in selected])
|
| 57 |
+
prepared.append((goal_words, selected))
|
| 58 |
+
count = max(1, max(map(len, maps), default=0))
|
| 59 |
+
goals, elements, features, masks, actions, targets = [], [], [], [], [], []
|
| 60 |
+
def tokens(tokens):
|
| 61 |
+
ids = [vocab.get(token, 1) for token in tokens[:max_tokens]]
|
| 62 |
+
return ids + [0]*(max_tokens-len(ids))
|
| 63 |
+
for row, (goal_words, selected) in zip(rows, prepared):
|
| 64 |
+
goals.append(tokens(goal_words))
|
| 65 |
+
ids, feats = [], []
|
| 66 |
+
target = -100
|
| 67 |
+
for j, (i, feature) in enumerate(selected):
|
| 68 |
+
e = row['elements'][i]
|
| 69 |
+
ids.append(tokens(words(e['role']+' '+e['name'])))
|
| 70 |
+
feats.append(feature)
|
| 71 |
+
if row.get('target') == i:
|
| 72 |
+
target = j
|
| 73 |
+
padding = count-len(selected)
|
| 74 |
+
elements.append(ids + [[0]*max_tokens for _ in range(padding)])
|
| 75 |
+
features.append(feats + [[0.0]*5 for _ in range(padding)])
|
| 76 |
+
masks.append([True]*len(selected) + [False]*padding)
|
| 77 |
+
actions.append(KINDS.index(row['action']) if row.get('action') in KINDS else -100)
|
| 78 |
+
targets.append(target)
|
| 79 |
+
batch = len(rows)
|
| 80 |
+
inputs = (torch.tensor(goals,dtype=torch.long).reshape(batch,max_tokens),
|
| 81 |
+
torch.tensor(elements,dtype=torch.long).reshape(batch,count,max_tokens),
|
| 82 |
+
torch.tensor(features,dtype=torch.float32).reshape(batch,count,5),
|
| 83 |
+
torch.tensor(masks,dtype=torch.bool).reshape(batch,count))
|
| 84 |
+
return inputs, torch.tensor(actions,dtype=torch.long), torch.tensor(targets,dtype=torch.long), maps
|
baim/policy.py
CHANGED
|
@@ -1,5 +1,4 @@
|
|
| 1 |
"""Learned three-action baseline with explicit uncertainty abstention."""
|
| 2 |
-
from dataclasses import asdict
|
| 3 |
import json
|
| 4 |
from pathlib import Path
|
| 5 |
import re
|
|
@@ -21,19 +20,34 @@ class LearnedPolicy:
|
|
| 21 |
|
| 22 |
@torch.inference_mode()
|
| 23 |
def predict(self, goal, state, ticket):
|
| 24 |
-
|
| 25 |
-
|
| 26 |
-
|
| 27 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 28 |
a,t = self.model(*inputs)
|
| 29 |
-
ap = (a/self.temperatures[0]).softmax(-1)
|
| 30 |
-
tp = (t/self.temperatures[1]).softmax(-1)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 31 |
ai,ti = int(ap.argmax()),int(tp.argmax())
|
| 32 |
ac,tc = float(ap[ai]),float(tp[ti])
|
| 33 |
if min(ac,tc) < self.threshold:
|
| 34 |
return Decision(ticket,Action(Kind.ASK_USER,('The policy is uncertain about this action or target.',)),ac,tc)
|
| 35 |
kind = Kind(KINDS[ai])
|
| 36 |
-
element = state.elements[
|
| 37 |
args = (element.ref,)
|
| 38 |
if kind in {Kind.TYPE,Kind.SELECT}:
|
| 39 |
# Literal copying, not inferred website/workflow logic. General span prediction is pending.
|
|
|
|
| 1 |
"""Learned three-action baseline with explicit uncertainty abstention."""
|
|
|
|
| 2 |
import json
|
| 3 |
from pathlib import Path
|
| 4 |
import re
|
|
|
|
| 20 |
|
| 21 |
@torch.inference_mode()
|
| 22 |
def predict(self, goal, state, ticket):
|
| 23 |
+
return self.predict_batch([(goal, state, ticket)])[0]
|
| 24 |
+
|
| 25 |
+
@torch.inference_mode()
|
| 26 |
+
def predict_batch(self, requests):
|
| 27 |
+
"""Bounded CPU microbatch. Every decision retains its own authority ticket."""
|
| 28 |
+
if len(requests) > 64:
|
| 29 |
+
raise ValueError('A microbatch may contain at most 64 requests')
|
| 30 |
+
if not requests:
|
| 31 |
+
return []
|
| 32 |
+
rows = [dict(goal=goal,elements=[dict(role=e.role,name=e.name,visible=e.visible,
|
| 33 |
+
enabled=e.enabled,sensitive=e.sensitive) for e in state.elements])
|
| 34 |
+
for goal,state,_ in requests]
|
| 35 |
+
inputs,_,_,maps = encode(rows,self.vocab)
|
| 36 |
a,t = self.model(*inputs)
|
| 37 |
+
ap = (a/self.temperatures[0]).softmax(-1)
|
| 38 |
+
tp = (t/self.temperatures[1]).softmax(-1)
|
| 39 |
+
return [self._decision(goal,state,ticket,mapping,ap[i],tp[i])
|
| 40 |
+
for i,((goal,state,ticket),mapping) in enumerate(zip(requests,maps))]
|
| 41 |
+
|
| 42 |
+
def _decision(self, goal, state, ticket, mapping, ap, tp):
|
| 43 |
+
if not mapping:
|
| 44 |
+
return Decision(ticket,Action(Kind.ASK_USER,('No eligible DOM target; a fallback is required.',)))
|
| 45 |
ai,ti = int(ap.argmax()),int(tp.argmax())
|
| 46 |
ac,tc = float(ap[ai]),float(tp[ti])
|
| 47 |
if min(ac,tc) < self.threshold:
|
| 48 |
return Decision(ticket,Action(Kind.ASK_USER,('The policy is uncertain about this action or target.',)),ac,tc)
|
| 49 |
kind = Kind(KINDS[ai])
|
| 50 |
+
element = state.elements[mapping[ti]]
|
| 51 |
args = (element.ref,)
|
| 52 |
if kind in {Kind.TYPE,Kind.SELECT}:
|
| 53 |
# Literal copying, not inferred website/workflow logic. General span prediction is pending.
|
docs/RUNTIME_V002.md
ADDED
|
@@ -0,0 +1,28 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Runtime v0.0.2
|
| 2 |
+
|
| 3 |
+
This release optimizes inference code; checkpoint weights and model intelligence are unchanged.
|
| 4 |
+
|
| 5 |
+
Feature encoding tokenizes each goal once, reuses lexical scores for retrieval, and
|
| 6 |
+
constructs tensors once per batch instead of repeatedly assigning tiny tensors.
|
| 7 |
+
Policy conversion no longer recursively copies unused DOM fields.
|
| 8 |
+
|
| 9 |
+
`LearnedPolicy.predict_batch([(goal, state, ticket), ...])` supports at most 64
|
| 10 |
+
independent requests. Each returned decision retains its own ticket and observed
|
| 11 |
+
target references. Confidence abstention, sensitive-target exclusion, and value-copy
|
| 12 |
+
restrictions remain in effect. Empty batches return an empty list.
|
| 13 |
+
|
| 14 |
+
Measured on the development Windows CPU with two Torch threads and the mean
|
| 15 |
+
checkpoint: median single-request prediction 2.103 ms before / 0.884 ms after
|
| 16 |
+
(2.38x); p95 3.304 / 1.327 ms. A 16-request batch took 6.253 ms median, about
|
| 17 |
+
0.391 ms per request. These are synthetic CPU policy timings, not browser-task
|
| 18 |
+
latency or target Linux VPS results. Loading, network, and browser time are excluded.
|
| 19 |
+
|
| 20 |
+
All encoded inputs, targets, and candidate maps were bit-identical across 1,440
|
| 21 |
+
validation/test/novel-wording examples. Individual action/ticket parity was checked
|
| 22 |
+
on 160 requests; batch actions matched individual actions on those requests.
|
| 23 |
+
Novel-wording raw step match remains 55.83%; this release does not fix generalization.
|
| 24 |
+
|
| 25 |
+
Evidence: `reports/runtime-optimization.json`. Run `python -m unittest discover -s tests`
|
| 26 |
+
for runtime regression tests. The existing `baim.bench_policy` command can measure
|
| 27 |
+
the current feature encoder independently. Compare against Hub revision
|
| 28 |
+
`795f73703a1e33b82636b2167d7d9507581efd4d` for the prior implementation.
|
pyproject.toml
CHANGED
|
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
|
| 4 |
|
| 5 |
[project]
|
| 6 |
name = "baim"
|
| 7 |
-
version = "0.0.
|
| 8 |
description = "Experimental CPU-first browser action intelligence"
|
| 9 |
requires-python = ">=3.12"
|
| 10 |
dependencies = []
|
|
|
|
| 4 |
|
| 5 |
[project]
|
| 6 |
name = "baim"
|
| 7 |
+
version = "0.0.2"
|
| 8 |
description = "Experimental CPU-first browser action intelligence"
|
| 9 |
requires-python = ">=3.12"
|
| 10 |
dependencies = []
|
reports/runtime-optimization.json
ADDED
|
@@ -0,0 +1,34 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"scope": "Windows CPU, two Torch threads, synthetic data; unchanged model weights",
|
| 3 |
+
"splits": {
|
| 4 |
+
"validation": {
|
| 5 |
+
"rows": 480,
|
| 6 |
+
"encoding_bit_exact": true,
|
| 7 |
+
"raw_step_match": 1.0
|
| 8 |
+
},
|
| 9 |
+
"test": {
|
| 10 |
+
"rows": 480,
|
| 11 |
+
"encoding_bit_exact": true,
|
| 12 |
+
"raw_step_match": 1.0
|
| 13 |
+
},
|
| 14 |
+
"novel_wording": {
|
| 15 |
+
"rows": 480,
|
| 16 |
+
"encoding_bit_exact": true,
|
| 17 |
+
"raw_step_match": 0.5583333333333333
|
| 18 |
+
}
|
| 19 |
+
},
|
| 20 |
+
"decision_parity_rows": 160,
|
| 21 |
+
"old_predict": {
|
| 22 |
+
"median_ms": 2.1032000076957047,
|
| 23 |
+
"p95_ms": 3.303700010292232
|
| 24 |
+
},
|
| 25 |
+
"new_predict": {
|
| 26 |
+
"median_ms": 0.8837499772198498,
|
| 27 |
+
"p95_ms": 1.3269999762997031
|
| 28 |
+
},
|
| 29 |
+
"batch16": {
|
| 30 |
+
"median_ms": 6.252999999560416,
|
| 31 |
+
"p95_ms": 9.634000016376376
|
| 32 |
+
},
|
| 33 |
+
"speedup": 2.37985862733719
|
| 34 |
+
}
|
tests/test_batch.py
ADDED
|
@@ -0,0 +1,43 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import unittest
|
| 2 |
+
from pathlib import Path
|
| 3 |
+
import torch
|
| 4 |
+
from baim.policy import LearnedPolicy
|
| 5 |
+
from baim.features import encode, candidates, lexical
|
| 6 |
+
from baim.state import State, Element
|
| 7 |
+
from baim.actions import Ticket, Kind
|
| 8 |
+
|
| 9 |
+
class BatchTests(unittest.TestCase):
|
| 10 |
+
@classmethod
|
| 11 |
+
def setUpClass(cls):
|
| 12 |
+
torch.set_num_threads(2)
|
| 13 |
+
cls.policy=LearnedPolicy(Path(__file__).resolve().parents[1]/'models/v000-mean')
|
| 14 |
+
|
| 15 |
+
def test_batch_isolates_tickets_refs_and_sensitive_targets(self):
|
| 16 |
+
requests=[]
|
| 17 |
+
for i in range(4):
|
| 18 |
+
elements=() if i==3 else (Element(f'{i}-safe','button','amber birch','0'),
|
| 19 |
+
Element(f'{i}-secret','textbox','password','1',sensitive=True))
|
| 20 |
+
state=State(str(i),i,'fixture',elements)
|
| 21 |
+
ticket=Ticket('session',str(i),0,i,str(i),i,state.state_hash)
|
| 22 |
+
requests.append(('Click amber birch.',state,ticket))
|
| 23 |
+
batch=self.policy.predict_batch(requests)
|
| 24 |
+
for request,decision in zip(requests,batch):
|
| 25 |
+
self.assertEqual(decision.action,self.policy.predict(*request).action)
|
| 26 |
+
self.assertEqual(decision.ticket,request[2])
|
| 27 |
+
if decision.action.kind!=Kind.ASK_USER:
|
| 28 |
+
self.assertEqual(decision.action.args[0],request[1].elements[0].ref)
|
| 29 |
+
self.assertEqual(batch[-1].action.kind,Kind.ASK_USER)
|
| 30 |
+
self.assertEqual(self.policy.predict_batch([]),[])
|
| 31 |
+
with self.assertRaises(ValueError):self.policy.predict_batch(requests*17)
|
| 32 |
+
|
| 33 |
+
def test_encoding_preserves_retrieval_and_lexical_features(self):
|
| 34 |
+
row={'goal':'Click café', 'elements':[{'role':'button','name':'café'},
|
| 35 |
+
{'role':'button','name':'café','sensitive':True},
|
| 36 |
+
{'role':'button','name':'other'}, {'role':'button','name':'café','visible':False}]}
|
| 37 |
+
inputs,_,_,maps=encode([row],self.policy.vocab)
|
| 38 |
+
self.assertEqual(maps[0],candidates(row['goal'],row['elements']))
|
| 39 |
+
for i,index in enumerate(maps[0]):
|
| 40 |
+
self.assertTrue(torch.equal(inputs[2][0,i],torch.tensor(lexical(row['goal'],row['elements'][index]))))
|
| 41 |
+
self.assertEqual(encode([],self.policy.vocab)[0][0].shape[0],0)
|
| 42 |
+
|
| 43 |
+
if __name__=='__main__':unittest.main()
|