devils-agent / baim /features.py
devildasdf's picture
Optimize CPU feature encoding and add bounded batch inference with measured parity
c7d6933 verified
Raw History Blame Contribute Delete
4.13 kB
"""Small vocabulary and deterministic candidate retrieval, without site rules."""
from collections import Counter
import re
import torch
KINDS = ('C', 'T', 'O')
ROLES = {'button', 'link', 'textbox', 'searchbox', 'combobox', 'checkbox', 'radio', 'listbox', 'spinbutton'}
def words(text):
return re.findall(r'\w+', text.casefold(), flags=re.UNICODE)
def fit_vocab(rows, limit=4096):
counts = Counter()
for row in rows:
counts.update(words(row['goal']))
for element in row['elements']:
counts.update(words(element['role'] + ' ' + element['name']))
return {'<pad>':0, '<unk>':1, **{token:i+2 for i, (token, _) in enumerate(counts.most_common(limit-2))}}
def lexical(goal, element):
goal_words, name_words = set(words(goal)), set(words(element['name']))
overlap = len(goal_words & name_words)
return [overlap / max(1,len(name_words)), overlap / max(1,len(goal_words)),
float(element['name'].casefold() in goal.casefold()),
float(element.get('visible', True)), float(element.get('enabled', True))]
def candidates(goal, elements, limit=40):
eligible = [i for i, e in enumerate(elements) if e.get('visible', True) and e.get('enabled', True)
and not e.get('sensitive', False) and e['role'] in ROLES]
return sorted(eligible, key=lambda i: (-lexical(goal,elements[i])[0], i))[:limit]
def encode(rows, vocab, max_tokens=24, max_candidates=40):
# Tokenize each goal once and construct six tensors per batch, not per element.
if max_tokens < 1 or max_candidates < 1:
raise ValueError('Token and candidate limits must be positive')
maps, prepared = [], []
for row in rows:
goal_words = words(row['goal'])
goal_set, folded = set(goal_words), row['goal'].casefold()
eligible = []
for i, e in enumerate(row['elements']):
if not (e.get('visible', True) and e.get('enabled', True)) or e.get('sensitive', False) or e['role'] not in ROLES:
continue
name_set = set(words(e['name']))
overlap = len(goal_set & name_set)
feature = [overlap/max(1,len(name_set)), overlap/max(1,len(goal_set)),
float(e['name'].casefold() in folded), float(e.get('visible',True)), float(e.get('enabled',True))]
eligible.append((i, feature))
eligible.sort(key=lambda item: (-item[1][0], item[0]))
selected = eligible[:max_candidates]
maps.append([i for i, _ in selected])
prepared.append((goal_words, selected))
count = max(1, max(map(len, maps), default=0))
goals, elements, features, masks, actions, targets = [], [], [], [], [], []
def tokens(tokens):
ids = [vocab.get(token, 1) for token in tokens[:max_tokens]]
return ids + [0]*(max_tokens-len(ids))
for row, (goal_words, selected) in zip(rows, prepared):
goals.append(tokens(goal_words))
ids, feats = [], []
target = -100
for j, (i, feature) in enumerate(selected):
e = row['elements'][i]
ids.append(tokens(words(e['role']+' '+e['name'])))
feats.append(feature)
if row.get('target') == i:
target = j
padding = count-len(selected)
elements.append(ids + [[0]*max_tokens for _ in range(padding)])
features.append(feats + [[0.0]*5 for _ in range(padding)])
masks.append([True]*len(selected) + [False]*padding)
actions.append(KINDS.index(row['action']) if row.get('action') in KINDS else -100)
targets.append(target)
batch = len(rows)
inputs = (torch.tensor(goals,dtype=torch.long).reshape(batch,max_tokens),
torch.tensor(elements,dtype=torch.long).reshape(batch,count,max_tokens),
torch.tensor(features,dtype=torch.float32).reshape(batch,count,5),
torch.tensor(masks,dtype=torch.bool).reshape(batch,count))
return inputs, torch.tensor(actions,dtype=torch.long), torch.tensor(targets,dtype=torch.long), maps