| import pytest |
|
|
| torch = pytest.importorskip("torch") |
| from torch import nn |
|
|
| from mitointeract_recovery.model import MitoInteractHead, TargetScaler |
|
|
|
|
| def test_head_output_shape(): |
| model = MitoInteractHead(16, 12, projection_dim=8, hidden_dim=16, dropout=0.0) |
| output = model(torch.randn(4, 16), torch.randn(4, 12)) |
| assert output.shape == (4,) |
|
|
|
|
| def test_head_uses_no_batch_norm(): |
| model = MitoInteractHead(16, 12, projection_dim=8, hidden_dim=16) |
| assert not any( |
| isinstance(module, nn.modules.batchnorm._BatchNorm) |
| for module in model.modules() |
| ) |
|
|
|
|
| def test_target_scaler_round_trip(): |
| values = torch.tensor([3.0, 5.0, 7.0, 9.0]) |
| scaler = TargetScaler.fit(values) |
| assert torch.allclose(scaler.decode(scaler.encode(values)), values) |
|
|
|
|
| def test_mismatched_batches_rejected(): |
| model = MitoInteractHead(16, 12, projection_dim=8, hidden_dim=16) |
| with pytest.raises(ValueError): |
| model(torch.randn(3, 16), torch.randn(4, 12)) |
|
|