Spaces:
Running
Running
| """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) | |
| 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), | |
| } | |