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)