"""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"