File size: 2,111 Bytes
ebad435
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
53
54
55
56
57
58
59
60
61
import unittest

import torch

from predictor_training.confidence import ConfidenceTokenHead


class ConfidenceTokenHeadTest(unittest.TestCase):
    def setUp(self) -> None:
        torch.manual_seed(0)
        self.head = ConfidenceTokenHead(
            dim=16,
            token_dim=8,
            context_dim=4,
            num_heads=2,
            ffn_dim=32,
            dropout=0.0,
            num_steps=3,
        )

    def test_output_shape_and_predictor_features_are_detached(self) -> None:
        transformed = torch.randn(2, 7, 16, requires_grad=True)
        predicted = torch.randn(2, 7, 16, requires_grad=True)
        anchor = torch.randn(2, 7, 16, requires_grad=True)
        output = self.head(
            transformed_hidden=transformed,
            pred_hidden=predicted,
            anchor_hidden=anchor,
            chunk_position=torch.tensor([0.0, 1.0]),
            step_id=torch.tensor([1, 3]),
        )
        self.assertEqual(output.shape, (2,))
        output.sum().backward()
        self.assertIsNone(transformed.grad)
        self.assertIsNone(predicted.grad)
        self.assertIsNone(anchor.grad)
        self.assertTrue(any(parameter.grad is not None for parameter in self.head.parameters()))

    def test_rejects_unsupported_step(self) -> None:
        value = torch.randn(1, 7, 16)
        with self.assertRaisesRegex(ValueError, "supports step_id"):
            self.head(
                transformed_hidden=value,
                pred_hidden=value,
                anchor_hidden=value,
                chunk_position=torch.tensor([0.0]),
                step_id=torch.tensor([4]),
            )

    def test_production_dimensions_and_parameter_count(self) -> None:
        head = ConfidenceTokenHead()
        self.assertEqual(head.token_projection[0].in_features, 1536)
        self.assertEqual(head.token_projection[0].out_features, 512)
        self.assertEqual(head.output[0].in_features, 1156)
        count = sum(parameter.numel() for parameter in head.parameters())
        self.assertEqual(count, 4_872_065)


if __name__ == "__main__":
    unittest.main()