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