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