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