Download tests/test_data_and_metrics.py from ChatterjeeLab/ReMEDi: direct link, hf CLI and curl.
- Browser
- Download file 2.46 kB
-
https://huggingface.co/ChatterjeeLab/ReMEDi/resolve/main/tests/test_data_and_metrics.py
- Command line
-
hf download hf://ChatterjeeLab/ReMEDi/tests/test_data_and_metrics.py
-
curl -L -o test_data_and_metrics.py https://huggingface.co/ChatterjeeLab/ReMEDi/resolve/main/tests/test_data_and_metrics.py
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) | |