devils-agent / tests /test_model.py
devildasdf's picture
Upload experimental BAIM code, research checkpoints and measured evaluations
795f737 verified
Raw
History Blame Contribute Delete
1.83 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
@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())