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