← Files Biohub ESMARCHIVED FILE
scripts/biohub_esm_lib/esmc.py
12.4 KB · Sep 30, 2026 · 23:14 UTC
"""Deterministic ESMC single-mask mutation scoring."""
from __future__ import annotations
import math
import re
from typing import Any
from .errors import SchemaDriftError, ValidationError
from .validation import validate_esmc_sequence
SUBSTITUTION_PATTERN = re.compile(r"^([A-Z])([1-9][0-9]*)([A-Z])$")
def _finite_score_number(scoring: dict[str, Any], field: str) -> float:
value = scoring.get(field)
if isinstance(value, bool) or not isinstance(value, (int, float)):
raise ValidationError(f"mutation score {field} must be a finite number")
try:
number = float(value)
except (OverflowError, ValueError) as exc:
raise ValidationError(f"mutation score {field} must be a finite number") from exc
if not math.isfinite(number):
raise ValidationError(f"mutation score {field} must be a finite number")
return number
def _format_score_number(value: float) -> str:
if value == 0:
value = 0.0
return format(value, ".4g")
def render_mutation_score_svg(score: dict[str, Any]) -> bytes:
"""Render a deterministic, self-contained mutation-score summary card."""
mutation = score.get("mutation")
scoring = score.get("scoring")
if not isinstance(mutation, dict) or not isinstance(scoring, dict):
raise ValidationError("mutation score is missing mutation or scoring data")
label = mutation.get("label")
match = SUBSTITUTION_PATTERN.fullmatch(label) if isinstance(label, str) else None
if match is None:
raise ValidationError("mutation score label must be one uppercase substitution")
wild_type, _, alternate = match.groups()
llr = _finite_score_number(scoring, "log_likelihood_ratio")
wild_type_log_probability = _finite_score_number(scoring, "wild_type_log_probability")
alternate_log_probability = _finite_score_number(scoring, "alternate_log_probability")
if llr > 0:
accent = "#2563eb"
comparison = "Alternate has higher model log probability"
elif llr < 0:
accent = "#d97706"
comparison = "Alternate has lower model log probability"
else:
accent = "#64748b"
comparison = "Model log probabilities are equal"
llr_text = _format_score_number(llr)
wild_type_text = _format_score_number(wild_type_log_probability)
alternate_text = _format_score_number(alternate_log_probability)
svg = f"""<svg xmlns="http://www.w3.org/2000/svg" width="720" height="360" viewBox="0 0 720 360" role="img" aria-labelledby="title description">
<title id="title">ESMC mutation score for {label}</title>
<desc id="description">{comparison}. This is a model hypothesis, not experimental evidence.</desc>
<rect width="720" height="360" rx="24" fill="#f8fafc"/>
<rect x="1" y="1" width="718" height="358" rx="23" fill="none" stroke="#cbd5e1" stroke-width="2"/>
<rect x="36" y="30" width="174" height="28" rx="14" fill="#e2e8f0"/>
<text x="123" y="49" text-anchor="middle" font-family="ui-sans-serif, system-ui, sans-serif" font-size="12" font-weight="700" letter-spacing="1.2" fill="#334155">MODEL HYPOTHESIS</text>
<text x="36" y="94" font-family="ui-sans-serif, system-ui, sans-serif" font-size="26" font-weight="700" fill="#0f172a">{label} mutation score</text>
<text x="36" y="120" font-family="ui-sans-serif, system-ui, sans-serif" font-size="15" fill="#475569">One masked-residue context · natural-log probabilities</text>
<text x="36" y="176" font-family="ui-sans-serif, system-ui, sans-serif" font-size="54" font-weight="750" fill="{accent}">{llr_text}</text>
<text x="36" y="204" font-family="ui-sans-serif, system-ui, sans-serif" font-size="14" font-weight="650" fill="#334155">log-likelihood ratio (alternate − wild type)</text>
<text x="36" y="230" font-family="ui-sans-serif, system-ui, sans-serif" font-size="14" fill="{accent}">{comparison}</text>
<rect x="400" y="78" width="132" height="130" rx="16" fill="#ffffff" stroke="#cbd5e1"/>
<text x="466" y="108" text-anchor="middle" font-family="ui-sans-serif, system-ui, sans-serif" font-size="12" font-weight="700" letter-spacing="1" fill="#64748b">WILD TYPE</text>
<text x="466" y="157" text-anchor="middle" font-family="ui-sans-serif, system-ui, sans-serif" font-size="38" font-weight="750" fill="#0f172a">{wild_type}</text>
<text x="466" y="188" text-anchor="middle" font-family="ui-monospace, SFMono-Regular, monospace" font-size="14" fill="#475569">ln P = {wild_type_text}</text>
<rect x="552" y="78" width="132" height="130" rx="16" fill="#ffffff" stroke="{accent}" stroke-width="2"/>
<text x="618" y="108" text-anchor="middle" font-family="ui-sans-serif, system-ui, sans-serif" font-size="12" font-weight="700" letter-spacing="1" fill="#64748b">ALTERNATE</text>
<text x="618" y="157" text-anchor="middle" font-family="ui-sans-serif, system-ui, sans-serif" font-size="38" font-weight="750" fill="{accent}">{alternate}</text>
<text x="618" y="188" text-anchor="middle" font-family="ui-monospace, SFMono-Regular, monospace" font-size="14" fill="#475569">ln P = {alternate_text}</text>
<line x1="36" y1="274" x2="684" y2="274" stroke="#cbd5e1"/>
<text x="36" y="304" font-family="ui-sans-serif, system-ui, sans-serif" font-size="14" font-weight="650" fill="#334155">Interpret with care</text>
<text x="36" y="329" font-family="ui-sans-serif, system-ui, sans-serif" font-size="13" fill="#475569">Model hypothesis — not experimental fitness, stability, activity, binding, or function.</text>
</svg>
"""
return svg.encode("utf-8")
def validate_single_substitution(sequence: str, mutation: str) -> dict[str, Any]:
"""Validate one one-based substitution against an exact ESMC sequence."""
normalized = validate_esmc_sequence(sequence)
if not isinstance(mutation, str) or mutation != mutation.strip():
raise ValidationError("mutation must be one uppercase substitution such as W43F")
match = SUBSTITUTION_PATTERN.fullmatch(mutation)
if match is None:
raise ValidationError("mutation must be one uppercase substitution such as W43F")
wild_type, position_text, alternate = match.groups()
position = int(position_text)
if position > len(normalized):
raise ValidationError("mutation position is outside the normalized sequence")
if normalized[position - 1] != wild_type:
raise ValidationError("mutation wild-type residue does not match the normalized sequence")
if alternate == wild_type:
raise ValidationError("mutation alternate residue must differ from wild type")
if alternate not in "ACDEFGHIKLMNPQRSTVWY":
raise ValidationError("mutation alternate residue must be a canonical amino acid")
return {
"label": mutation,
"wild_type": wild_type,
"alternate": alternate,
"position_one_based": position,
"sequence_index_zero_based": position - 1,
# The pinned ESMC tokenizer prepends exactly one BOS/CLS token.
"logits_index_zero_based": position,
"bos_offset": 1,
}
def validate_sequence_logits(
value: Any,
*,
expected_positions: int,
minimum_width: int,
) -> list[list[float]]:
"""Validate the pinned single-sequence managed logits tensor as JSON data."""
if not isinstance(value, list) or len(value) != expected_positions:
raise SchemaDriftError(
"managed ESMC sequence logits must be a two-dimensional [L+2, V] tensor",
raw={"observed_type": type(value).__name__},
)
width: int | None = None
result: list[list[float]] = []
for row in value:
if not isinstance(row, list) or not row:
raise SchemaDriftError(
"managed ESMC sequence logits contain an invalid row",
raw={"positions": len(value)},
)
if width is None:
width = len(row)
if width < minimum_width:
raise SchemaDriftError(
"managed ESMC sequence logits vocabulary is too small",
raw={"positions": len(value), "vocabulary_width": width},
)
elif len(row) != width:
raise SchemaDriftError(
"managed ESMC sequence logits contain jagged rows",
raw={"positions": len(value)},
)
normalized_row: list[float] = []
for item in row:
if isinstance(item, bool) or not isinstance(item, (int, float)):
raise SchemaDriftError(
"managed ESMC sequence logits must contain numeric values",
raw={"positions": len(value), "vocabulary_width": width},
)
try:
number = float(item)
except (OverflowError, ValueError) as exc:
raise SchemaDriftError(
"managed ESMC sequence logits must contain finite numeric values",
raw={"positions": len(value), "vocabulary_width": width},
) from exc
if not math.isfinite(number):
raise SchemaDriftError(
"managed ESMC sequence logits must contain finite values",
raw={"positions": len(value), "vocabulary_width": width},
)
normalized_row.append(number)
result.append(normalized_row)
return result
def derive_single_mask_llr(
sequence_logits: list[list[float]],
mutation: dict[str, Any],
*,
wild_type_token_id: int,
alternate_token_id: int,
mask_token_id: int,
) -> dict[str, Any]:
"""Derive an LLR from one full-vocabulary logit row using natural logs."""
for label, token_id in (
("wild-type", wild_type_token_id),
("alternate", alternate_token_id),
("mask", mask_token_id),
):
if isinstance(token_id, bool) or not isinstance(token_id, int) or token_id < 0:
raise ValidationError(f"{label} token id is invalid")
row = sequence_logits[mutation["logits_index_zero_based"]]
if max(wild_type_token_id, alternate_token_id, mask_token_id) >= len(row):
raise SchemaDriftError(
"managed ESMC sequence logits do not cover the pinned tokenizer ids",
raw={"vocabulary_width": len(row)},
)
wild_type_logit = row[wild_type_token_id]
alternate_logit = row[alternate_token_id]
llr = alternate_logit - wild_type_logit
maximum = max(row)
shifted_log_normalizer = math.log(math.fsum(math.exp(value - maximum) for value in row))
wild_type_log_probability = (wild_type_logit - maximum) - shifted_log_normalizer
alternate_log_probability = (alternate_logit - maximum) - shifted_log_normalizer
if not all(
math.isfinite(value)
for value in (
shifted_log_normalizer,
wild_type_log_probability,
alternate_log_probability,
llr,
)
):
raise SchemaDriftError(
"managed ESMC logits produce non-finite log probabilities or LLR",
raw={
"wild_type_logit": wild_type_logit,
"alternate_logit": alternate_logit,
},
)
return {
"schema_version": "1.0",
"analysis": "single-masked-residue-log-likelihood-ratio",
"evidence_class": "model-generated-hypothesis",
"mutation": dict(mutation),
"masking": {
"context_count": 1,
"masked_positions_one_based": [mutation["position_one_based"]],
"mask_token_id": mask_token_id,
},
"scoring": {
"definition": (
f"ln P({mutation['alternate']} | {mutation['wild_type']}"
f"{mutation['position_one_based']} masked) - "
f"ln P({mutation['wild_type']} | {mutation['wild_type']}"
f"{mutation['position_one_based']} masked)"
),
"log_base": "e",
"normalization": ("log_softmax over the full returned sequence-logit vocabulary"),
"llr_computation": "alternate_logit - wild_type_logit in one masked context",
"wild_type_vocabulary_index": wild_type_token_id,
"alternate_vocabulary_index": alternate_token_id,
"wild_type_logit": wild_type_logit,
"alternate_logit": alternate_logit,
"wild_type_log_probability": wild_type_log_probability,
"alternate_log_probability": alternate_log_probability,
"log_likelihood_ratio": llr,
},
"interpretation": (
"A model hypothesis, not experimental fitness, stability, activity, "
"binding, or function."
),
}
SHA-256: 96cbaa293ba29ecc91ba2b20cfda039d7b42fcb6511b0aa2a694d454e0a977b4