ReMEDi / tests /test_data_and_metrics.py
pranamanam's picture
Upload 62 files
3f98d52 verified
Raw History Blame Contribute Delete
2.46 kB
import numpy as np
import pandas as pd
import pytest
from scipy import sparse
from remedi.data import normalize_counts, standardize_obs
from remedi.splits import make_splits
from remedi.metrics import intervention_metrics
from remedi.demo import synthetic_data
from remedi.io import ResponseData
from remedi.uncertainty import cross_block_losses, calibration_radius
def test_sparse_and_dense_count_normalization_agree():
x=np.array([[0,2,3],[0,0,0],[1,0,2]],float)
np.testing.assert_allclose(normalize_counts(x),normalize_counts(sparse.csr_matrix(x)).toarray())
def test_missing_dose_units_fail_closed():
obs=pd.DataFrame({"drug":["x"],"dose":[1],"context":["c"],"block":["r"],"plate":["p"]})
mapping={"drug":"drug","dose":"dose","context":"context","block":"block","control_group":"plate"}
with pytest.raises(ValueError,match="dose_unit"):
standardize_obs(obs,mapping)
def test_response_roundtrip_and_split_leakage(tmp_path):
data=synthetic_data(molecules=24)
data.save(tmp_path)
restored=ResponseData.load(tmp_path)
np.testing.assert_allclose(data.response,restored.response)
restored.obs.loc[0,"split"]="new_split"
with pytest.raises(ValueError,match="multiple splits"):restored.validate()
def test_molecule_groups_never_cross_splits():
smiles=["C"*k+"c1ccccc1" for k in range(1,20)]
result=make_splits(pd.DataFrame({"smiles":smiles+smiles}),seed=2)
assert len(result)==19
assert result.groupby("molecule_id").split.nunique().max()==1
assert set(result.split)=={"train","tune","calibration","test"}
def test_mixture_of_bad_endpoints_is_not_scored_as_a_good_average():
# Responses -1 and +1 both have squared loss 1 from target 0.
result=intervention_metrics([.5,.5],[1.,1.],tolerance=.1)
assert result["endpoint_loss"]==1.
assert result["success"]==0.
def test_cross_block_shared_control_is_rejected():
data=synthetic_data(molecules=24)
group=data.obs.groupby(["molecule_id","dose_um","context"])
indices=[g.index.to_list() for _,g in group][:2]
data.obs.loc[indices[0][1],"control_id"]=data.obs.loc[indices[0][0],"control_id"]
with pytest.raises(ValueError,match="reuse a control"):
cross_block_losses(data,[indices[1]],indices[0],np.ones(data.response.shape[1]))
def test_calibration_is_an_explicit_empirical_quantile():
assert calibration_radius([1.,2.,3.,4.],.8)==4.
with pytest.raises(ValueError): calibration_radius([], .9)