Add trained contextual action candidate with browser evaluation and explicit real-web limits
c8e5620 verified Download tests/test_model.py from devildasdf/devils-agent: direct link, hf CLI and curl.
- Browser
- Download file 2.55 kB
-
https://huggingface.co/devildasdf/devils-agent/resolve/main/tests/test_model.py
- Command line
-
hf download hf://devildasdf/devils-agent/tests/test_model.py
-
curl -L -o test_model.py https://huggingface.co/devildasdf/devils-agent/resolve/main/tests/test_model.py
2.55 kB
| import unittest | |
| try: | |
| import torch | |
| from baim.features import encode,fit_vocab,candidates | |
| from baim.model import PointerPolicy | |
| except ImportError: | |
| torch = None | |
| from baim.synthetic import generate | |
| class ModelTests(unittest.TestCase): | |
| def test_candidate_filter_rejects_hidden_disabled_and_sensitive(self): | |
| sample = next(generate('train',1,7)) | |
| for index,key in enumerate(['visible','enabled','sensitive']): | |
| sample['elements'][index][key] = key == 'sensitive' | |
| result = candidates(sample['goal'],sample['elements']) | |
| self.assertFalse(set(result)&{0,1,2}) | |
| def test_pointer_is_equivariant_to_candidate_order(self): | |
| torch.manual_seed(1) | |
| torch.set_num_threads(2) | |
| rows = list(generate('train',2,11)) | |
| vocab = fit_vocab(rows) | |
| x,*_ = encode(rows,vocab) | |
| model = PointerPolicy(vocab_size=len(vocab)).eval() | |
| with torch.inference_mode(): | |
| a,t = model(*x) | |
| permutation = torch.arange(x[1].shape[1]-1,-1,-1) | |
| changed = (x[0],x[1][:,permutation],x[2][:,permutation],x[3][:,permutation]) | |
| other_a,other_t = model(*changed) | |
| torch.testing.assert_close(a,other_a) | |
| torch.testing.assert_close(t[:,permutation],other_t) | |
| def test_all_encoders_handle_padded_candidates(self): | |
| rows = list(generate('train',3,19)) | |
| vocab = fit_vocab(rows) | |
| x,*_ = encode(rows,vocab) | |
| for architecture in ['mean','gru','transformer']: | |
| model = PointerPolicy(vocab_size=len(vocab),encoder=architecture).eval() | |
| with torch.inference_mode(): | |
| a,t = model(*x) | |
| self.assertTrue(torch.isfinite(a).all()) | |
| self.assertTrue(torch.isfinite(t).all()) | |
| def test_contextual_action_is_order_equivariant_and_trainable(self): | |
| rows = list(generate('train',3,29)) | |
| vocab = fit_vocab(rows) | |
| x,*_ = encode(rows,vocab) | |
| model = PointerPolicy(vocab_size=len(vocab),contextual_action=True).eval() | |
| a,t = model(*x) | |
| permutation = torch.arange(x[1].shape[1]-1,-1,-1) | |
| other_a,other_t = model(x[0],x[1][:,permutation],x[2][:,permutation],x[3][:,permutation]) | |
| torch.testing.assert_close(a,other_a) | |
| torch.testing.assert_close(t[:,permutation],other_t) | |
| a.sum().backward() | |
| self.assertIsNotNone(model.pointer[0].weight.grad) | |
| self.assertTrue(torch.isfinite(model.pointer[0].weight.grad).all()) | |