import unittest import torch from transformers import Qwen2Config, Qwen2ForCausalLM from tree_decode import FieldPlan, decode_fields, prefill class TreeTests(unittest.TestCase): def setUp(self): torch.manual_seed(7) torch.set_num_threads(2) config = Qwen2Config(vocab_size=97, hidden_size=64, intermediate_size=128, num_hidden_layers=2, num_attention_heads=4, num_key_value_heads=2, attention_dropout=0) config._attn_implementation = 'sdpa' self.model = Qwen2ForCausalLM(config).eval() self.ids = torch.tensor([[1, 2, 3, 4, 5]]) self.cache = prefill(self.model, self.ids) def test_matches_independent_and_batch_after_continuation(self): suffixes = [[8, 9], [10, 11, 12, 13], [14]] plan = FieldPlan(suffixes, 5, 'cpu', torch.float32) ref = decode_fields(self.model, self.cache, plan, 'batch', 4) tree = decode_fields(self.model, self.cache, plan, 'tree', 4) torch.testing.assert_close(tree, ref, atol=1e-6, rtol=1e-5) for i, suffix in enumerate(suffixes): sequence = self.ids[0].tolist() + suffix for step in range(4): with torch.inference_mode(): logits = self.model(torch.tensor([sequence])).logits[0, -1] torch.testing.assert_close(tree[step, i], logits, atol=1e-6, rtol=1e-5) sequence.append(logits.argmax().item()) self.assertEqual(self.cache.get_seq_length(), 5) def test_no_sibling_leakage(self): original = FieldPlan([[8, 9], [10, 11, 12]], 5, 'cpu', torch.float32) changed = FieldPlan([[8, 9], [66, 67, 68]], 5, 'cpu', torch.float32) a = decode_fields(self.model, self.cache, original, 'tree', 3) b = decode_fields(self.model, self.cache, changed, 'tree', 3) torch.testing.assert_close(a[:, 0], b[:, 0], atol=1e-6, rtol=1e-5) self.assertFalse(torch.allclose(a[:, 1], b[:, 1])) def test_one_branch(self): plan = FieldPlan([[7, 8, 9]], 5, 'cpu', torch.float32) torch.testing.assert_close(decode_fields(self.model, self.cache, plan, 'tree', 2), decode_fields(self.model, self.cache, plan, 'batch', 2)) if __name__ == '__main__': unittest.main()