syntheogenesis / dee /core /benchmark.py
github-actions[bot]
Deploy c0f4b8a
6c0aef4
Raw
History Blame Contribute Delete
6.59 kB
"""Validation harness — does the engine actually predict? Receipts, not claims.
The glass-box thesis lives or dies on this: every score we surface is framed
as a prediction to verify, so we owe the user proof the predictions track
reality. This harness correlates the engine's zero-shot ESM-2 ΔLL ranking
against measured deep-mutational-scanning (DMS) fitness from published studies
(e.g. ProteinGym) and reports the two numbers a protein engineer actually
cares about:
* Spearman ρ — does the ranking order match measured fitness order?
* top-decile precision — of the variants we rank in the top 10%, what
fraction are genuinely high-fitness (top quartile measured)? i.e. "if I
only make the picks you put on top, how often are they real?"
Pure numpy, predictions injected — so it's unit-testable without ESM. The
reproducible CLI (scripts/run_benchmarks.py) runs it against real DMS with the
live model where the weights exist; the app only ever serves results a real
run produced. No number is fabricated here.
"""
from __future__ import annotations
import re
from dataclasses import dataclass
from typing import List, Optional, Sequence
import numpy as np
_LABEL_RE = re.compile(r"^([A-Za-z])(\d+)([A-Za-z*])$")
def predict_additive(scores_df, labels: Sequence[str]) -> List[float]:
"""Sum-of-single-site ESM-2 ΔLL for each (possibly multi-site) DMS label
('A1C', 'A1C:D5E', or 'A1C,D5E') — the SAME additive scoring the design
engine itself uses (dee.optimizer.search), so this validates exactly what
we ship rather than a different metric. ``scores_df`` is the long-format
output of ``ESM2Scorer.score_all_substitutions`` (columns: position
0-indexed, mut_aa, delta_ll).
Single source of truth for both the reproducible CLI
(scripts/run_benchmarks.py) and the admin in-Space endpoint — never
duplicated, so they can't silently drift apart.
A label with any unparseable or out-of-table site becomes NaN (never
silently coerced to 0) — spearman()/top_decile_precision() mask NaNs out
via ``np.isfinite``, so a bad label reduces the effective N rather than
corrupting the correlation."""
lut = {(int(r.position), str(r.mut_aa)): float(r.delta_ll)
for r in scores_df.itertuples(index=False)}
out: List[float] = []
for lab in labels:
total, ok = 0.0, True
for tok in re.split(r"[:,]", str(lab).strip()):
m = _LABEL_RE.match(tok.strip())
if not m:
ok = False
break
pos, mut = int(m.group(2)) - 1, m.group(3).upper()
key = (pos, mut)
if key not in lut:
ok = False
break
total += lut[key]
out.append(total if ok else float("nan"))
return out
def _average_ranks(x: np.ndarray) -> np.ndarray:
"""Ranks with ties resolved to the average rank (proper Spearman ties)."""
order = np.argsort(x, kind="mergesort")
ranks = np.empty(len(x), dtype=np.float64)
sx = x[order]
i = 0
n = len(x)
while i < n:
j = i
while j + 1 < n and sx[j + 1] == sx[i]:
j += 1
avg = (i + j) / 2.0 + 1.0 # 1-indexed average rank over the tie block
ranks[order[i:j + 1]] = avg
i = j + 1
return ranks
def spearman(a: Sequence[float], b: Sequence[float]) -> Optional[float]:
"""Spearman rank correlation. None if <3 points or no variance."""
a = np.asarray(a, dtype=np.float64)
b = np.asarray(b, dtype=np.float64)
mask = np.isfinite(a) & np.isfinite(b)
a, b = a[mask], b[mask]
if len(a) < 3:
return None
ra, rb = _average_ranks(a), _average_ranks(b)
if ra.std() < 1e-9 or rb.std() < 1e-9:
return None
return float(np.corrcoef(ra, rb)[0, 1])
def top_decile_precision(
predicted: Sequence[float], measured: Sequence[float],
*, pred_frac: float = 0.10, true_frac: float = 0.25,
) -> Optional[float]:
"""Of the top ``pred_frac`` by prediction, the fraction that land in the
top ``true_frac`` by measured fitness. The 'are your top picks real?' number.
None if too few points to form a meaningful top set."""
p = np.asarray(predicted, dtype=np.float64)
m = np.asarray(measured, dtype=np.float64)
mask = np.isfinite(p) & np.isfinite(m)
p, m = p[mask], m[mask]
n = len(p)
k = int(round(n * pred_frac))
if n < 10 or k < 1:
return None
top_pred_idx = np.argsort(-p)[:k]
true_cut = np.quantile(m, 1.0 - true_frac)
hits = int(np.sum(m[top_pred_idx] >= true_cut))
return float(hits / k)
@dataclass
class DatasetResult:
"""One DMS assay's validation result."""
name: str
protein: str # e.g. UniProt/DMS id
n: int # variants scored
spearman: Optional[float]
top_decile_precision: Optional[float]
source: str = "" # citation / DOI / dataset id (provenance)
def as_dict(self) -> dict:
return {
"name": self.name, "protein": self.protein, "n": self.n,
"spearman": None if self.spearman is None else round(self.spearman, 4),
"top_decile_precision": None if self.top_decile_precision is None
else round(self.top_decile_precision, 4),
"source": self.source,
}
def evaluate_dataset(
name: str, protein: str,
predicted: Sequence[float], measured: Sequence[float],
*, source: str = "",
) -> DatasetResult:
"""Score one aligned (predicted, measured) DMS assay."""
p = np.asarray(predicted, dtype=np.float64)
m = np.asarray(measured, dtype=np.float64)
mask = np.isfinite(p) & np.isfinite(m)
return DatasetResult(
name=name, protein=protein, n=int(mask.sum()),
spearman=spearman(p, m),
top_decile_precision=top_decile_precision(p, m),
source=source,
)
def summarize(results: Sequence[DatasetResult]) -> dict:
"""Headline across datasets: median Spearman, median top-decile precision,
dataset + variant counts. Medians (robust to a couple of hard assays)."""
rhos = [r.spearman for r in results if r.spearman is not None]
precs = [r.top_decile_precision for r in results if r.top_decile_precision is not None]
return {
"n_datasets": len(results),
"n_variants": int(sum(r.n for r in results)),
"median_spearman": None if not rhos else round(float(np.median(rhos)), 4),
"median_top_decile_precision": None if not precs
else round(float(np.median(precs)), 4),
}