Self-Forcing / tests /test_confidence_token.py
Cccccz's picture
Upload tests
ebad435 verified
Raw History Blame Contribute Delete
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()