"""Tests for SimplePopulationResponseModel.""" import sys import os sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "src")) import pytest import torch from gidflow.models import SimplePopulationResponseModel class TestSimplePopulationResponseModel: @pytest.fixture def model(self): return SimplePopulationResponseModel(num_genes=32, hidden_dim=64, perturbation_dim=32) def test_output_shape(self, model): B, N, G = 4, 10, 32 src = torch.randn(B, N, G) pert = torch.zeros(B, G) pert[:, :3] = 1.0 out = model(src, pert) assert out.shape == (B, N, G), f"Expected ({B},{N},{G}), got {out.shape}" def test_different_n_per_call(self, model): """Model should handle varying N between calls.""" for N in [1, 5, 32, 64]: src = torch.randn(2, N, 32) pert = torch.zeros(2, 32) out = model(src, pert) assert out.shape == (2, N, 32) def test_no_nan(self, model): src = torch.randn(3, 8, 32) pert = torch.zeros(3, 32) pert[:, 0] = 1.0 out = model(src, pert) assert not out.isnan().any(), "NaN in output" assert not out.isinf().any(), "Inf in output" def test_gradient_flows(self, model): src = torch.randn(2, 6, 32, requires_grad=True) pert = torch.randn(2, 32, requires_grad=True) out = model(src, pert) loss = out.mean() loss.backward() assert src.grad is not None assert pert.grad is not None def test_mask_accepted(self, model): """Forward accepts source_mask without error.""" B, N, G = 2, 8, 32 src = torch.randn(B, N, G) pert = torch.zeros(B, G) mask = torch.ones(B, N, dtype=torch.bool) out = model(src, pert, source_mask=mask) assert out.shape == (B, N, G) def test_gpu_if_available(self, model): if not torch.cuda.is_available(): pytest.skip("CUDA not available") device = torch.device("cuda") model = model.to(device) B, N, G = 2, 5, 32 src = torch.randn(B, N, G, device=device) pert = torch.zeros(B, G, device=device) out = model(src, pert) assert out.device.type == "cuda" def test_perturbation_changes_output(self, model): """Different perturbations should give different predictions.""" model.eval() src = torch.randn(1, 8, 32) pert1 = torch.zeros(1, 32); pert1[0, 0] = 1.0 pert2 = torch.zeros(1, 32); pert2[0, 5] = 1.0 with torch.no_grad(): out1 = model(src, pert1) out2 = model(src, pert2) assert not torch.allclose(out1, out2), "Different perturbations gave identical output" def test_batch_size_1(self, model): out = model(torch.randn(1, 4, 32), torch.zeros(1, 32)) assert out.shape == (1, 4, 32)