devils-agent / tests /test_batch.py
devildasdf's picture
Optimize CPU feature encoding and add bounded batch inference with measured parity
c7d6933 verified
Raw History Blame Contribute Delete
2.2 kB
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()