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