File size: 2,113 Bytes
3f98d52
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import numpy as np
import pandas as pd
import anndata as ad
import pytest
from remedi.data import prepare_h5ad


def fixture(path, order=None):
    rng=np.random.default_rng(3)
    obs=[]
    for block in ['r1','r2']:
        for drug,smiles in [('DMSO','CS(C)=O'),('a','CCO'),('b','CCN')]:
            for cell in range(4):
                obs.append(dict(drug=drug,smiles=smiles,dose=100.,context='c',block=block))
    x=ad.AnnData(rng.poisson(3,size=(len(obs),10)).astype(float),obs=pd.DataFrame(obs,index=[str(i) for i in range(len(obs))]),var=pd.DataFrame(index=[f'g{i}' for i in range(10)]))
    if order is not None:x=x[:,order].copy()
    x.write_h5ad(path)


def test_external_gene_reordering_is_invariant(tmp_path):
    first=tmp_path/'a.h5ad';second=tmp_path/'b.h5ad'
    fixture(first);fixture(second,np.arange(10)[::-1])
    mapping=dict(drug='drug',dose='dose',context='context',block='block',control_group='block',dose_unit='nM',smiles='smiles',control_names=['DMSO'])
    splits=pd.DataFrame({'smiles':['CCO','CCN'],'split':['train','test']})
    a=prepare_h5ad(first,tmp_path/'fit',mapping,splits,min_cells=2,dimensions=3)
    b=prepare_h5ad(second,tmp_path/'transfer',mapping,splits,min_cells=2,dimensions=3,feature_model=tmp_path/'fit/cell_feature_model.joblib')
    np.testing.assert_allclose(a.response,b.response)
    assert a.metadata['feature_space_id']==b.metadata['feature_space_id']
    np.testing.assert_allclose(a.obs.dose_um,.1)


def test_external_missing_genes_rejected(tmp_path):
    first=tmp_path/'a.h5ad';second=tmp_path/'b.h5ad'
    fixture(first);fixture(second,np.arange(9))
    mapping=dict(drug='drug',dose='dose',context='context',block='block',control_group='block',dose_unit='nM',smiles='smiles',control_names=['DMSO'])
    splits=pd.DataFrame({'smiles':['CCO','CCN'],'split':['train','test']})
    prepare_h5ad(first,tmp_path/'fit',mapping,splits,min_cells=2,dimensions=3)
    with pytest.raises(ValueError,match='misses fitted genes'):
        prepare_h5ad(second,tmp_path/'transfer',mapping,splits,min_cells=2,feature_model=tmp_path/'fit/cell_feature_model.joblib')