File size: 3,421 Bytes
9d901ad | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 | """Tests for DeepPTR data loading utilities."""
import numpy as np
import pytest
import torch
from anndata import AnnData
from scptr.deep._data import setup_dataloaders
from scptr.deep._utils import get_library_sizes
@pytest.fixture
def simple_adata():
"""AnnData with spliced/unspliced counts."""
rng = np.random.RandomState(42)
n, g = 100, 20
s = rng.poisson(5, size=(n, g)).astype(np.float32)
u = rng.poisson(2, size=(n, g)).astype(np.float32)
adata = AnnData(X=s)
adata.layers["spliced"] = s
adata.layers["unspliced"] = u
adata.obs["cell_type"] = [f"type_{i % 3}" for i in range(n)]
adata.obs["cell_type"] = adata.obs["cell_type"].astype("category")
return adata
class TestGetLibrarySizes:
def test_shapes(self, simple_adata):
l_u, l_s = get_library_sizes(simple_adata)
assert l_u.shape == (simple_adata.n_obs,)
assert l_s.shape == (simple_adata.n_obs,)
def test_positive(self, simple_adata):
l_u, l_s = get_library_sizes(simple_adata)
assert (l_u >= 1.0).all()
assert (l_s >= 1.0).all()
def test_correct_sums(self, simple_adata):
l_u, l_s = get_library_sizes(simple_adata)
expected_s = simple_adata.layers["spliced"].sum(axis=1)
expected_u = simple_adata.layers["unspliced"].sum(axis=1)
np.testing.assert_allclose(l_s, np.clip(expected_s, 1.0, None), rtol=1e-5)
np.testing.assert_allclose(l_u, np.clip(expected_u, 1.0, None), rtol=1e-5)
def test_missing_layer_raises(self):
adata = AnnData(X=np.zeros((5, 3)))
with pytest.raises(KeyError, match="Missing required layer"):
get_library_sizes(adata)
class TestSetupDataloaders:
def test_returns_four(self, simple_adata):
train_dl, val_dl, train_idx, val_idx = setup_dataloaders(
simple_adata, batch_size=16, val_frac=0.2, seed=0
)
assert isinstance(train_dl, torch.utils.data.DataLoader)
assert isinstance(val_dl, torch.utils.data.DataLoader)
assert len(train_idx) + len(val_idx) == simple_adata.n_obs
def test_no_overlap(self, simple_adata):
_, _, train_idx, val_idx = setup_dataloaders(
simple_adata, batch_size=16, val_frac=0.2, seed=0
)
assert len(set(train_idx) & set(val_idx)) == 0
def test_batch_contents(self, simple_adata):
train_dl, _, _, _ = setup_dataloaders(
simple_adata, batch_size=16, val_frac=0.1, seed=0
)
s, u, l_s, l_u = next(iter(train_dl))
assert s.ndim == 2
assert u.ndim == 2
assert l_s.ndim == 1
assert l_u.ndim == 1
assert s.shape[1] == simple_adata.n_vars
def test_stratified_split(self, simple_adata):
_, _, train_idx, val_idx = setup_dataloaders(
simple_adata,
batch_size=16,
val_frac=0.2,
stratify_key="cell_type",
seed=0,
)
# All cell types should appear in both splits
train_types = set(simple_adata.obs["cell_type"].values[train_idx])
val_types = set(simple_adata.obs["cell_type"].values[val_idx])
assert train_types == val_types
def test_reproducible(self, simple_adata):
_, _, idx1, _ = setup_dataloaders(simple_adata, seed=42)
_, _, idx2, _ = setup_dataloaders(simple_adata, seed=42)
np.testing.assert_array_equal(idx1, idx2)
|