Spaces:
Running
Running
| """Tests for the glass-box per-prediction confidence (dee.core.confidence). | |
| The signal is ESM-2's own posterior sharpness at a position, reconstructed | |
| from the ΞLL table (softmax shift-invariant) β a peaked distribution β high | |
| confidence, a flat one β low. Weakest-link aggregation across a variant's | |
| mutated positions. | |
| """ | |
| import pandas as pd | |
| import pytest | |
| from dee.core.confidence import ( | |
| attach_confidence, | |
| position_confidence_from_scores, | |
| variant_confidence, | |
| ) | |
| def _scores(rows): | |
| """rows: list of (position, mut_aa, delta_ll) β wt_aa fixed 'A'.""" | |
| return pd.DataFrame( | |
| [{"position": p, "wt_aa": "A", "mut_aa": m, "delta_ll": float(d)} for (p, m, d) in rows] | |
| ) | |
| _MUTS = "CDEFGHIKLMNPQRSTVWY" # 19 non-Ala AAs | |
| def test_peaked_position_is_high_confidence(): | |
| # One substitution dominates, the rest strongly disfavored β model is sure. | |
| rows = [(0, m, (5.0 if m == "C" else -10.0)) for m in _MUTS] | |
| conf = position_confidence_from_scores(_scores(rows)) | |
| assert conf[0] > 0.7 | |
| def test_flat_position_is_low_confidence(): | |
| # All substitutions equal to WT β uniform posterior over 20 β confidence ~0. | |
| rows = [(0, m, 0.0) for m in _MUTS] | |
| conf = position_confidence_from_scores(_scores(rows)) | |
| assert conf[0] < 0.05 | |
| def test_confidence_is_bounded(): | |
| rows = [(0, m, 3.0) for m in _MUTS] + [(1, m, -3.0) for m in _MUTS] | |
| conf = position_confidence_from_scores(_scores(rows)) | |
| for c in conf.values(): | |
| assert 0.0 <= c <= 1.0 | |
| def test_variant_confidence_is_weakest_link(): | |
| pos_conf = {0: 0.9, 5: 0.3, 9: 0.7} | |
| score, weakest = variant_confidence("A1C,A6D,A10E", pos_conf) | |
| assert score == pytest.approx(0.3) # min, not mean | |
| assert weakest == "A6D" | |
| def test_variant_confidence_none_when_no_scorable_site(): | |
| score, weakest = variant_confidence("A99Z", {0: 0.9}) # position not in map | |
| assert score is None | |
| assert weakest == "" | |
| def test_attach_confidence_sets_fields_and_skips_wt(): | |
| pos_conf = {0: 0.8, 1: 0.4} | |
| rows = [ | |
| {"Variant_ID": "WT", "Mutations_AA": ""}, | |
| {"Variant_ID": "V1", "Mutations_AA": "A1C,A2D"}, | |
| ] | |
| attach_confidence(rows, pos_conf) | |
| assert rows[0]["Confidence"] is None | |
| assert rows[1]["Confidence"] == pytest.approx(0.4) # weakest of {0.8, 0.4} | |
| assert rows[1]["Confidence_Weakest"] == "A2D" | |