ReMEDi / tests /test_preparation.py
pranamanam's picture
Upload 62 files
3f98d52 verified
Raw History Blame Contribute Delete
2.11 kB
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')