Download tests/test_disca_input_variant.py from Cccccz/Self-Forcing: direct link, hf CLI and curl.
- Browser
- Download file 1.99 kB
-
https://huggingface.co/Cccccz/Self-Forcing/resolve/main/tests/test_disca_input_variant.py
- Command line
-
hf download hf://Cccccz/Self-Forcing/tests/test_disca_input_variant.py
-
curl -L -o test_disca_input_variant.py https://huggingface.co/Cccccz/Self-Forcing/resolve/main/tests/test_disca_input_variant.py
1.99 kB
| import unittest | |
| import torch | |
| from torch import nn | |
| from predictor_training.single_block import SingleBlockPredictor | |
| class EchoBlock(nn.Module): | |
| def forward(self, value: torch.Tensor, **kwargs) -> torch.Tensor: | |
| return value | |
| def build_model(input_variant: str) -> SingleBlockPredictor: | |
| model = SingleBlockPredictor( | |
| EchoBlock(), | |
| dim=8, | |
| gradient_checkpointing=False, | |
| input_variant=input_variant, | |
| ).eval() | |
| with torch.no_grad(): | |
| model.residual_out.weight.copy_(torch.eye(8)) | |
| return model | |
| def inputs() -> dict[str, object]: | |
| return { | |
| "current_tokens": torch.randn(2, 3, 8), | |
| "anchor_hidden": torch.randn(2, 3, 8), | |
| "timestep_modulation": torch.randn(2, 6, 8), | |
| "grid_sizes": torch.tensor([[1, 1, 3], [1, 1, 3]]), | |
| "freqs": torch.empty(0), | |
| "history_k": torch.randn(2, 2, 1, 8), | |
| "history_v": torch.randn(2, 2, 1, 8), | |
| "cross_k": torch.randn(2, 1, 1, 8), | |
| "cross_v": torch.randn(2, 1, 1, 8), | |
| "current_start": 2, | |
| } | |
| class DiscaInputVariantTest(unittest.TestCase): | |
| def test_disca_output_ignores_previous_chunk_hidden(self) -> None: | |
| torch.manual_seed(0) | |
| model = build_model(input_variant="disca") | |
| common = inputs() | |
| first = model(previous_hidden=torch.randn(2, 3, 8), **common) | |
| second = model(previous_hidden=torch.randn(2, 3, 8), **common) | |
| self.assertTrue(torch.equal(first, second)) | |
| def test_disca_physically_removes_previous_chunk_channel(self) -> None: | |
| baseline = build_model(input_variant="self_forcing") | |
| disca = build_model(input_variant="disca") | |
| self.assertEqual(baseline.fusion.proj_in.weight.shape, (16, 24)) | |
| self.assertEqual(disca.fusion.proj_in.weight.shape, (16, 16)) | |
| self.assertTrue(hasattr(baseline.fusion, "previous_norm")) | |
| self.assertFalse(hasattr(disca.fusion, "previous_norm")) | |
| if __name__ == "__main__": | |
| unittest.main() | |