File size: 2,340 Bytes
f36843f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
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()