syntheogenesis / dee /core /confidence.py
Tengo Gzirishvili
Glass-box confidence + epistasis-in-the-loop + the Field Atlas
ed32186
Raw
History Blame Contribute Delete
4.4 kB
"""Per-prediction model certainty — the glass-box "how sure is ESM-2?" signal.
Every ΔLL score answers "is this substitution favorable?" — but not "how much
should I trust that call?". This module adds the second axis, honestly: it
measures how SHARP ESM-2's amino-acid distribution is at each position. A
well-conserved position where the model overwhelmingly prefers one residue
(low entropy) is one it's confident about; a flexible position where many
residues are near-equiprobable (high entropy) is one where any single call is
a weaker signal — and we say so rather than hide it.
Crucially this is NOT calibrated against wet-lab outcomes (we don't claim it
is): it's the model's own posterior confidence, surfaced instead of buried.
The bench-calibrated axis is the separate active-learning / calibration work.
Cost is zero extra compute: softmax is shift-invariant, so the full
per-position posterior over the 20 amino acids is recoverable from the ΔLL
table we already computed — ``p(a) ∝ exp(ΔLL_a)`` with ``ΔLL_wt = 0`` — no
extra ESM-2 forward pass.
"""
from __future__ import annotations
import math
import re
from typing import Dict, Optional, Sequence, Tuple
import numpy as np
# 20 canonical AAs; H_max is the entropy of a uniform posterior over them.
_N_AA = 20
_H_MAX = math.log(_N_AA)
_LABEL_RE = re.compile(r"^([A-Za-z])(\d+)([A-Za-z*])$") # "F76Y"
def _entropy_confidence(delta_lls: Sequence[float]) -> float:
"""Confidence in [0,1] for one position from its 19 mutant ΔLLs.
Reconstruct the posterior p(a) ∝ exp(ΔLL_a) (wt's ΔLL is 0 by definition),
then confidence = 1 - H(p)/H_max. 1.0 = the model is certain (one residue
dominates); 0.0 = maximally uncertain (flat over all 20)."""
vals = np.asarray(list(delta_lls) + [0.0], dtype=np.float64) # + implicit WT
vals = vals - vals.max() # numerically-stable softmax
p = np.exp(vals)
p = p / p.sum()
nz = p[p > 0]
h = float(-(nz * np.log(nz)).sum())
conf = 1.0 - h / _H_MAX
return float(min(1.0, max(0.0, conf)))
def position_confidence_from_scores(scores_df) -> Dict[int, float]:
"""position (0-indexed) -> confidence in [0,1], from the long-format ΔLL
DataFrame (columns: position, delta_ll, ...). Positions with fewer than a
near-complete substitution set are still scored on whatever is present."""
out: Dict[int, float] = {}
if scores_df is None or len(scores_df) == 0:
return out
for pos, group in scores_df.groupby("position"):
out[int(pos)] = _entropy_confidence([float(x) for x in group["delta_ll"].values])
return out
def _parse_positions(mutations: str) -> list:
"""'V60L,F76Y' -> [('V60L', 59), ('F76Y', 75)] (label, 0-indexed pos)."""
out = []
for tok in re.split(r"[,\s;]+", (mutations or "").strip()):
m = _LABEL_RE.match(tok.strip())
if not m:
continue
pos = int(m.group(2)) - 1
if pos >= 0:
out.append((tok.strip(), pos))
return out
def variant_confidence(
mutations: str, pos_conf: Dict[int, float],
) -> Tuple[Optional[float], str]:
"""Overall confidence for a variant = the WEAKEST-LINK of its mutated
positions (one uncertain position honestly undermines the whole call).
Returns (confidence or None if no scorable position, weakest-site label)."""
sites = _parse_positions(mutations)
scored = [(lbl, pos_conf[pos]) for (lbl, pos) in sites if pos in pos_conf]
if not scored:
return (None, "")
weakest_label, weakest = min(scored, key=lambda t: t[1])
return (round(weakest, 4), weakest_label)
def attach_confidence(rows: list, pos_conf: Dict[int, float]) -> list:
"""In-place: add ``Confidence`` (0-1, weakest-link) and
``Confidence_Weakest`` (the label of the least-certain changed residue) to
each variant row dict that has a ``Mutations_AA``. WT / mutation-less rows
get Confidence = None. Returns the same list for chaining."""
for r in rows:
if not isinstance(r, dict):
continue
muts = str(r.get("Mutations_AA", "") or "")
if not muts:
r["Confidence"] = None
r["Confidence_Weakest"] = ""
continue
conf, weakest = variant_confidence(muts, pos_conf)
r["Confidence"] = conf
r["Confidence_Weakest"] = weakest
return rows