syntheogenesis / tests /test_confidence.py
Tengo Gzirishvili
Glass-box confidence + epistasis-in-the-loop + the Field Atlas
ed32186
Raw
History Blame Contribute Delete
2.38 kB
"""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"