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