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