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')
|