File size: 6,029 Bytes
07fcdfe
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
"""Tests for real data interfaces (scPerturb + PDGrapher).

Norman2019 and PDGrapher real_lognorm are loaded from disk; tests skip
gracefully when the files are not present.
"""
import sys
import os

sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "src"))

import pytest
import torch
from torch.utils.data import DataLoader

from gidflow.data import (
    ScPerturbPopulationDataset,
    PDGrapherPseudobulkDataset,
    population_collate_fn,
    select_hvg_by_variance,
    normalize_counts,
)
import numpy as np

NORMAN_PATH = (
    "/data/boom/Protein/regulatory_field/data/raw/"
    "scPerturb/rna_protein/NormanWeissman2019_filtered.h5ad"
)
PDG_DIR = "/data/boom/Protein/PDGrapher/data/processed/torch_data/real_lognorm"

norman_available   = pytest.mark.skipif(not os.path.exists(NORMAN_PATH),   reason="Norman2019 h5ad not found")
pdgrapher_available = pytest.mark.skipif(not os.path.isdir(PDG_DIR),        reason="PDGrapher real_lognorm not found")


# ---------------------------------------------------------------------------
# gene_selection utils
# ---------------------------------------------------------------------------
class TestGeneSelection:

    def test_hvg_shape(self):
        X = np.random.randn(100, 500).astype(np.float32)
        idx = select_hvg_by_variance(X, n_genes=50)
        assert idx.shape == (50,)
        assert len(set(idx)) == 50

    def test_hvg_fewer_than_requested(self):
        X = np.random.randn(10, 20).astype(np.float32)
        idx = select_hvg_by_variance(X, n_genes=100)
        assert len(idx) == 20   # capped at n_genes

    def test_normalize_no_log(self):
        X = np.array([[1.0, 3.0], [2.0, 2.0]], dtype=np.float32)
        out = normalize_counts(X, target_sum=4.0, log1p=False)
        assert np.allclose(out.sum(axis=1), [4.0, 4.0])

    def test_normalize_log1p(self):
        X = np.ones((5, 10), dtype=np.float32) * 100
        out = normalize_counts(X, target_sum=1e4, log1p=True)
        assert (out > 0).all()
        assert out.dtype == np.float32


# ---------------------------------------------------------------------------
# ScPerturbPopulationDataset (Norman2019)
# ---------------------------------------------------------------------------
@norman_available
class TestScPerturbPopulationDataset:

    @pytest.fixture(scope="class")
    def ds(self):
        import warnings; warnings.filterwarnings("ignore")
        return ScPerturbPopulationDataset(
            NORMAN_PATH,
            n_hvg=500,
            min_cells_per_cond=30,
            max_source_cells=16,
            max_target_cells=16,
            use_single_pert_only=True,
            force_include_pert_genes=True,
            seed=0,
        )

    def test_num_conditions(self, ds):
        assert len(ds) >= 50, f"Expected ≥50 conditions, got {len(ds)}"

    def test_item_shapes(self, ds):
        item = ds[0]
        G = ds.num_genes
        assert item["source_cells"].shape == (16, G)
        assert item["target_cells"].shape == (16, G)
        assert item["perturbation"].shape == (G,)

    def test_perturbation_is_multihot(self, ds):
        item = ds[0]
        pert = item["perturbation"]
        assert ((pert == 0) | (pert == 1)).all()
        assert pert.sum() >= 1

    def test_most_perts_in_gene_space(self, ds):
        covered = ds.n_genes_in_hvg_that_are_targets
        assert covered / len(ds) >= 0.8, f"Only {covered}/{len(ds)} pert genes covered"

    def test_source_target_differ(self, ds):
        item = ds[5]
        src_mean = item["source_cells"].mean()
        tgt_mean = item["target_cells"].mean()
        # They come from different cell populations — means will differ slightly
        assert item["source_cells"].shape == item["target_cells"].shape

    def test_collate_fn_integration(self, ds):
        loader = DataLoader(ds, batch_size=4, collate_fn=population_collate_fn, shuffle=False)
        batch = next(iter(loader))
        assert batch.source_cells.shape[0] == 4
        assert batch.source_cells.shape[2] == ds.num_genes
        assert not batch.source_cells.isnan().any()

    def test_reproducibility(self):
        import warnings; warnings.filterwarnings("ignore")
        ds1 = ScPerturbPopulationDataset(NORMAN_PATH, n_hvg=200, min_cells_per_cond=30,
                                          max_source_cells=8, seed=7)
        ds2 = ScPerturbPopulationDataset(NORMAN_PATH, n_hvg=200, min_cells_per_cond=30,
                                          max_source_cells=8, seed=7)
        item1 = ds1[0]
        item2 = ds2[0]
        assert torch.allclose(item1["perturbation"], item2["perturbation"])


# ---------------------------------------------------------------------------
# PDGrapherPseudobulkDataset
# ---------------------------------------------------------------------------
@pdgrapher_available
class TestPDGrapherPseudobulkDataset:

    @pytest.fixture(scope="class")
    def ds(self):
        return PDGrapherPseudobulkDataset(PDG_DIR, cell_lines=["A549"], max_items=200)

    def test_length(self, ds):
        assert len(ds) == 200

    def test_item_shapes(self, ds):
        item = ds[0]
        G = ds.num_genes
        assert item["source_cells"].shape == (1, G)
        assert item["target_cells"].shape == (1, G)
        assert item["perturbation"].shape == (G,)

    def test_num_genes(self, ds):
        assert ds.num_genes == 10716

    def test_perturbation_nonzero(self, ds):
        for i in range(10):
            item = ds[i]
            assert item["perturbation"].sum() >= 1

    def test_collate_fn_integration(self, ds):
        loader = DataLoader(ds, batch_size=4, collate_fn=population_collate_fn)
        batch = next(iter(loader))
        assert batch.source_cells.shape == (4, 1, ds.num_genes)
        assert batch.target_cells.shape == (4, 1, ds.num_genes)

    def test_no_nan(self, ds):
        for i in range(5):
            item = ds[i]
            assert not item["source_cells"].isnan().any()
            assert not item["target_cells"].isnan().any()