File size: 4,127 Bytes
c7d6933
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
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
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
"""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