Download tests/test_batch.py from devildasdf/devils-agent: direct link, hf CLI and curl.
- Browser
- Download file 2.2 kB
-
https://huggingface.co/devildasdf/devils-agent/resolve/main/tests/test_batch.py
- Command line
-
hf download hf://devildasdf/devils-agent/tests/test_batch.py
-
curl -L -o test_batch.py https://huggingface.co/devildasdf/devils-agent/resolve/main/tests/test_batch.py
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): | |
| 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() | |