devildasdf commited on
Commit
c7d6933
·
verified ·
1 Parent(s): 795f737

Optimize CPU feature encoding and add bounded batch inference with measured parity

Browse files
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
- # Training padding uses batch size; inference uses only the surviving candidates.
39
- maps = [candidates(row['goal'],row['elements'],max_candidates) for row in rows]
40
- count = max(1, max(map(len,maps),default=0))
41
- goal = torch.zeros((len(rows),max_tokens),dtype=torch.long)
42
- element = torch.zeros((len(rows),count,max_tokens),dtype=torch.long)
43
- features = torch.zeros((len(rows),count,5))
44
- mask = torch.zeros((len(rows),count),dtype=torch.bool)
45
- targets = torch.full((len(rows),),-100,dtype=torch.long)
46
- actions = torch.full((len(rows),),-100,dtype=torch.long)
47
- def tokens(text):
48
- return [vocab.get(token,1) for token in words(text)[:max_tokens]]
49
- for batch,row in enumerate(rows):
50
- ids = tokens(row['goal'])
51
- goal[batch,:len(ids)] = torch.tensor(ids,dtype=torch.long)
52
- if row.get('action') in KINDS:
53
- actions[batch] = KINDS.index(row['action'])
54
- for j,index in enumerate(maps[batch]):
55
- e = row['elements'][index]
56
- ids = tokens(e['role'] + ' ' + e['name'])
57
- element[batch,j,:len(ids)] = torch.tensor(ids,dtype=torch.long)
58
- features[batch,j] = torch.tensor(lexical(row['goal'],e))
59
- mask[batch,j] = True
60
- if row.get('target') == index:
61
- targets[batch] = j
62
- return (goal,element,features,mask),actions,targets,maps
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
- row = dict(goal=goal,elements=[asdict(e) for e in state.elements])
25
- inputs,_,_,maps = encode([row],self.vocab)
26
- if not maps[0]:
27
- return Decision(ticket,Action(Kind.ASK_USER,('No eligible DOM target; a fallback is required.',)))
 
 
 
 
 
 
 
 
 
28
  a,t = self.model(*inputs)
29
- ap = (a/self.temperatures[0]).softmax(-1)[0]
30
- tp = (t/self.temperatures[1]).softmax(-1)[0]
 
 
 
 
 
 
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[maps[0][ti]]
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.1"
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()