← Files Biohub ESMARCHIVED FILE
scripts/biohub_esm_lib/esmc_landscape.py
12.1 KB · Sep 30, 2026 · 23:14 UTC
"""Deterministic ESMC single-mask mutation-landscape analysis."""
from __future__ import annotations
import csv
import io
import math
from typing import Any
from .errors import SchemaDriftError, ValidationError
from .validation import validate_esmc_sequence
CANONICAL_AMINO_ACIDS = tuple("ACDEFGHIKLMNPQRSTVWY")
# A complete landscape retains one full [L+2, V] response per residue, so its
# in-memory and raw-artifact footprint grows quadratically with sequence length.
ESMC_LANDSCAPE_MAX_RESIDUES = 512
def validate_landscape_sequence(sequence: str) -> str:
"""Require canonical residues for a complete substitution landscape."""
normalized = validate_esmc_sequence(sequence)
if len(normalized) > ESMC_LANDSCAPE_MAX_RESIDUES:
raise ValidationError(
"ESMC mutation landscapes are limited to "
f"{ESMC_LANDSCAPE_MAX_RESIDUES} residues to bound raw-logit memory and artifact size"
)
if any(residue not in CANONICAL_AMINO_ACIDS for residue in normalized):
raise ValidationError(
"ESMC mutation landscapes require a sequence of 20 canonical amino acids"
)
return normalized
def _token_id(tokenizer: Any, token: str, label: str) -> int:
value = tokenizer.convert_tokens_to_ids(token)
if isinstance(value, bool) or not isinstance(value, int) or value < 0:
raise ValidationError(f"pinned ESMC tokenizer returned an invalid {label} token id")
return value
def encode_esmc_sequence(tokenizer: Any, sequence: str) -> dict[str, Any]:
"""Encode one validated sequence with exactly one BOS and one EOS token."""
normalized = validate_esmc_sequence(sequence)
bos_token_id = tokenizer.cls_token_id
eos_token_id = tokenizer.eos_token_id
mask_token_id = tokenizer.mask_token_id
for label, token_id in (
("BOS", bos_token_id),
("EOS", eos_token_id),
("mask", mask_token_id),
):
if isinstance(token_id, bool) or not isinstance(token_id, int) or token_id < 0:
raise ValidationError(f"pinned ESMC tokenizer returned an invalid {label} token id")
residue_token_ids = [
_token_id(tokenizer, residue, f"residue {residue}") for residue in normalized
]
encoded = [bos_token_id, *residue_token_ids, eos_token_id]
if len(encoded) != len(normalized) + 2:
raise ValidationError("pinned ESMC tokenizer must add exactly one BOS and one EOS token")
if any(
token_id in {bos_token_id, eos_token_id, mask_token_id} for token_id in residue_token_ids
):
raise ValidationError("pinned ESMC tokenizer mapped a residue to a control token")
return {
"sequence": normalized,
"tokens": encoded,
"bos_token_id": bos_token_id,
"eos_token_id": eos_token_id,
"mask_token_id": mask_token_id,
"bos_offset": 1,
}
def canonical_token_ids(tokenizer: Any) -> dict[str, int]:
"""Resolve the 20 canonical amino-acid IDs from the pinned tokenizer."""
result = {
residue: _token_id(tokenizer, residue, f"canonical residue {residue}")
for residue in CANONICAL_AMINO_ACIDS
}
if len(set(result.values())) != len(result):
raise ValidationError("pinned ESMC tokenizer canonical residue IDs must be unique")
return result
def mask_esmc_position(
encoded_tokens: list[int],
*,
position_one_based: int,
mask_token_id: int,
) -> list[int]:
"""Replace exactly one biological residue while preserving BOS/EOS."""
residue_count = len(encoded_tokens) - 2
if (
isinstance(position_one_based, bool)
or not isinstance(position_one_based, int)
or not 1 <= position_one_based <= residue_count
):
raise ValidationError("ESMC landscape position is outside the sequence")
if isinstance(mask_token_id, bool) or not isinstance(mask_token_id, int) or mask_token_id < 0:
raise ValidationError("ESMC landscape mask token id is invalid")
masked = list(encoded_tokens)
masked[position_one_based] = mask_token_id
changed = [
index
for index, pair in enumerate(zip(encoded_tokens, masked, strict=True))
if pair[0] != pair[1]
]
if changed != [position_one_based]:
raise ValidationError("ESMC landscape must replace exactly one encoded residue")
return masked
def _full_vocabulary_entropy_bits(row: list[float]) -> float:
maximum = max(row)
shifted_total = math.fsum(math.exp(value - maximum) for value in row)
log_normalizer = maximum + math.log(shifted_total)
entropy = -math.fsum(
probability * ((value - log_normalizer) / math.log(2.0))
for value in row
if (probability := math.exp(value - log_normalizer)) > 0.0
)
if not math.isfinite(entropy):
raise SchemaDriftError("managed ESMC logits produce non-finite entropy")
return entropy
def derive_mutation_landscape(
sequence: str,
masked_position_rows: list[list[float]],
*,
canonical_ids: dict[str, int],
mask_token_id: int,
summary_count: int = 5,
) -> dict[str, Any]:
"""Derive tutorial entropy and canonical LLRs from one row per mask."""
normalized = validate_landscape_sequence(sequence)
if len(masked_position_rows) != len(normalized):
raise SchemaDriftError(
"managed ESMC landscape must return one masked-position row per residue"
)
if set(canonical_ids) != set(CANONICAL_AMINO_ACIDS):
raise ValidationError("ESMC landscape canonical tokenizer mapping is incomplete")
if isinstance(summary_count, bool) or not isinstance(summary_count, int) or summary_count < 1:
raise ValidationError("ESMC landscape summary count must be positive")
minimum_width = max(*canonical_ids.values(), mask_token_id) + 1
positions: list[dict[str, Any]] = []
vocabulary_width: int | None = None
for position, (wild_type, row) in enumerate(
zip(normalized, masked_position_rows, strict=True), start=1
):
if not isinstance(row, list) or len(row) < minimum_width:
raise SchemaDriftError(
"managed ESMC landscape row does not cover the pinned tokenizer IDs"
)
if vocabulary_width is None:
vocabulary_width = len(row)
elif len(row) != vocabulary_width:
raise SchemaDriftError("managed ESMC landscape rows have inconsistent vocabularies")
if any(
isinstance(value, bool)
or not isinstance(value, (int, float))
or not math.isfinite(float(value))
for value in row
):
raise SchemaDriftError("managed ESMC landscape row must contain finite numbers")
normalized_row = [float(value) for value in row]
wild_type_logit = normalized_row[canonical_ids[wild_type]]
canonical_llr = {
residue: normalized_row[token_id] - wild_type_logit
for residue, token_id in canonical_ids.items()
}
negative_count = sum(
canonical_llr[residue] < 0.0
for residue in CANONICAL_AMINO_ACIDS
if residue != wild_type
)
positions.append(
{
"position_one_based": position,
"wild_type": wild_type,
"tutorial_full_vocabulary_entropy_bits": _full_vocabulary_entropy_bits(
normalized_row
),
"negative_substitution_fraction": negative_count / 19,
"canonical_llr": canonical_llr,
}
)
count = min(summary_count, len(positions))
def summary_item(position: dict[str, Any]) -> dict[str, Any]:
return {
"position_one_based": position["position_one_based"],
"wild_type": position["wild_type"],
"label": f"{position['wild_type']}{position['position_one_based']}",
"tutorial_full_vocabulary_entropy_bits": position[
"tutorial_full_vocabulary_entropy_bits"
],
"negative_substitution_fraction": position["negative_substitution_fraction"],
}
constrained = sorted(
positions,
key=lambda item: (
item["tutorial_full_vocabulary_entropy_bits"],
item["position_one_based"],
),
)[:count]
tolerant = sorted(
positions,
key=lambda item: (
-item["tutorial_full_vocabulary_entropy_bits"],
item["position_one_based"],
),
)[:count]
return {
"schema_version": "1.0",
"analysis": "single-mask-canonical-mutation-landscape",
"evidence_class": "model-generated-hypothesis",
"metrics": {
"entropy": {
"name": "tutorial_full_vocabulary_entropy_bits",
"log_base": 2,
"units": "bits",
"token_set": "complete returned sequence-logit vocabulary",
"vocabulary_width": vocabulary_width,
},
"negative_substitution_fraction": {
"definition": "fraction of 19 non-wild-type canonical substitutions with LLR < 0",
"denominator": 19,
},
"canonical_llr": {
"definition": "alternate logit minus wild-type logit in the same masked context",
"log_base": "e",
"amino_acid_order": list(CANONICAL_AMINO_ACIDS),
},
},
"positions": positions,
"summary": {
"most_constrained": [summary_item(item) for item in constrained],
"most_tolerant": [summary_item(item) for item in tolerant],
},
"interpretation": (
"A model-generated prioritization signal, not an experimental measurement of "
"fitness, stability, activity, binding, or function."
),
}
def render_mutation_landscape_csv(landscape: dict[str, Any]) -> bytes:
"""Render the per-position landscape in a deterministic tabular format."""
positions = landscape.get("positions")
if not isinstance(positions, list) or not positions:
raise ValidationError("ESMC mutation landscape has no positions")
fieldnames = [
"position_one_based",
"wild_type",
"tutorial_full_vocabulary_entropy_bits",
"negative_substitution_fraction",
*(f"llr_{residue}" for residue in CANONICAL_AMINO_ACIDS),
]
output = io.StringIO(newline="")
writer = csv.DictWriter(output, fieldnames=fieldnames, lineterminator="\n")
writer.writeheader()
for position in positions:
if not isinstance(position, dict) or not isinstance(position.get("canonical_llr"), dict):
raise ValidationError("ESMC mutation landscape position is malformed")
writer.writerow(
{
"position_one_based": position.get("position_one_based"),
"wild_type": position.get("wild_type"),
"tutorial_full_vocabulary_entropy_bits": position.get(
"tutorial_full_vocabulary_entropy_bits"
),
"negative_substitution_fraction": position.get("negative_substitution_fraction"),
**{
f"llr_{residue}": position["canonical_llr"].get(residue)
for residue in CANONICAL_AMINO_ACIDS
},
}
)
return output.getvalue().encode("utf-8")
def summarize_reported_usage(responses: list[dict[str, Any]]) -> dict[str, Any]:
"""Aggregate only usage fields explicitly returned by the managed service."""
result: dict[str, Any] = {"managed_request_count": len(responses)}
for field in ("tokens_used", "credits_used"):
values: list[float] = []
for response in responses:
container = response.get("data") if isinstance(response.get("data"), dict) else response
value = container.get(field)
if (
isinstance(value, bool)
or not isinstance(value, (int, float))
or not math.isfinite(float(value))
or value < 0
):
continue
values.append(float(value))
result[f"responses_reporting_{field}"] = len(values)
result[field] = math.fsum(values) if len(values) == len(responses) else None
return result
SHA-256: f8114e9620d73e4a524d63202c0bed35596e4c0a49a850389cb1372b0c051cdd