Download tests/test_confidence_token.py from Cccccz/Self-Forcing: direct link, hf CLI and curl.
- Browser
- Download file 2.11 kB
-
https://huggingface.co/Cccccz/Self-Forcing/resolve/main/tests/test_confidence_token.py
- Command line
-
hf download hf://Cccccz/Self-Forcing/tests/test_confidence_token.py
-
curl -L -o test_confidence_token.py https://huggingface.co/Cccccz/Self-Forcing/resolve/main/tests/test_confidence_token.py
2.11 kB
| 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() | |