File size: 2,197 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
import unittest
from pathlib import Path
import torch
from baim.policy import LearnedPolicy
from baim.features import encode, candidates, lexical
from baim.state import State, Element
from baim.actions import Ticket, Kind

class BatchTests(unittest.TestCase):
    @classmethod
    def setUpClass(cls):
        torch.set_num_threads(2)
        cls.policy=LearnedPolicy(Path(__file__).resolve().parents[1]/'models/v000-mean')

    def test_batch_isolates_tickets_refs_and_sensitive_targets(self):
        requests=[]
        for i in range(4):
            elements=() if i==3 else (Element(f'{i}-safe','button','amber birch','0'),
                Element(f'{i}-secret','textbox','password','1',sensitive=True))
            state=State(str(i),i,'fixture',elements)
            ticket=Ticket('session',str(i),0,i,str(i),i,state.state_hash)
            requests.append(('Click amber birch.',state,ticket))
        batch=self.policy.predict_batch(requests)
        for request,decision in zip(requests,batch):
            self.assertEqual(decision.action,self.policy.predict(*request).action)
            self.assertEqual(decision.ticket,request[2])
            if decision.action.kind!=Kind.ASK_USER:
                self.assertEqual(decision.action.args[0],request[1].elements[0].ref)
        self.assertEqual(batch[-1].action.kind,Kind.ASK_USER)
        self.assertEqual(self.policy.predict_batch([]),[])
        with self.assertRaises(ValueError):self.policy.predict_batch(requests*17)

    def test_encoding_preserves_retrieval_and_lexical_features(self):
        row={'goal':'Click café', 'elements':[{'role':'button','name':'café'},
            {'role':'button','name':'café','sensitive':True},
            {'role':'button','name':'other'}, {'role':'button','name':'café','visible':False}]}
        inputs,_,_,maps=encode([row],self.policy.vocab)
        self.assertEqual(maps[0],candidates(row['goal'],row['elements']))
        for i,index in enumerate(maps[0]):
            self.assertTrue(torch.equal(inputs[2][0,i],torch.tensor(lexical(row['goal'],row['elements'][index]))))
        self.assertEqual(encode([],self.policy.vocab)[0][0].shape[0],0)

if __name__=='__main__':unittest.main()