| 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 |
|
|
|
|
| @unittest.skipIf(torch is None,'training extras not installed') |
| 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()) |
|
|