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()
|