Instructions to use Synthyra/ESMFold2-Fast with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use Synthyra/ESMFold2-Fast with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("feature-extraction", model="Synthyra/ESMFold2-Fast", trust_remote_code=True)# pip install -U transformers accelerate # Load model directly from transformers import AutoTokenizer, AutoModel tokenizer = AutoTokenizer.from_pretrained("Synthyra/ESMFold2-Fast", trust_remote_code=True) model = AutoModel.from_pretrained("Synthyra/ESMFold2-Fast", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
Support verified live v2 confidence heads
Browse files
fastplms/models/esmfold2/confidence_checkpoint.py
ADDED
|
@@ -0,0 +1,171 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Load explicit head-only confidence checkpoints without changing base weights."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import hashlib
|
| 6 |
+
import json
|
| 7 |
+
import re
|
| 8 |
+
import torch
|
| 9 |
+
|
| 10 |
+
from pathlib import Path, PurePosixPath
|
| 11 |
+
from typing import TYPE_CHECKING, Any
|
| 12 |
+
from huggingface_hub import hf_hub_download
|
| 13 |
+
from safetensors.torch import load_file
|
| 14 |
+
from transformers.utils.hub import extract_commit_hash
|
| 15 |
+
|
| 16 |
+
if TYPE_CHECKING:
|
| 17 |
+
from .modeling_esmfold2_experimental import ESMFold2ExperimentalModel
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
_DOWNLOAD_OPTIONS = (
|
| 21 |
+
"cache_dir",
|
| 22 |
+
"token",
|
| 23 |
+
"local_files_only",
|
| 24 |
+
"force_download",
|
| 25 |
+
)
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
def _checkpoint_metadata(path: Path, source: dict[str, str]) -> dict[str, Any]:
|
| 29 |
+
if path.stat().st_size > 1_000_000:
|
| 30 |
+
raise ValueError("Confidence checkpoint metadata exceeds one megabyte.")
|
| 31 |
+
metadata = json.loads(path.read_text(encoding="utf-8"))
|
| 32 |
+
if not isinstance(metadata, dict) or metadata.get("schema_version") != 1:
|
| 33 |
+
raise ValueError("Unsupported confidence checkpoint metadata schema.")
|
| 34 |
+
for key in ("repo_id", "repo_type", "model_id", "base_weight_sha256"):
|
| 35 |
+
if metadata.get(key) != source[key]:
|
| 36 |
+
raise ValueError(
|
| 37 |
+
f"Confidence checkpoint {key} does not match its configured source."
|
| 38 |
+
)
|
| 39 |
+
update = metadata.get("update")
|
| 40 |
+
if type(update) is not int or update <= 0:
|
| 41 |
+
raise ValueError("A published training head must have a positive update count.")
|
| 42 |
+
if metadata.get("evaluation_status") != "pending":
|
| 43 |
+
raise ValueError(
|
| 44 |
+
"Rolling confidence heads must explicitly declare pending evaluation."
|
| 45 |
+
)
|
| 46 |
+
if (
|
| 47 |
+
metadata.get("head_state_format") != "native_confidence_head"
|
| 48 |
+
or metadata.get("checkpoint_kind") != "ema"
|
| 49 |
+
):
|
| 50 |
+
raise ValueError("Expected a native EMA confidence-head state dictionary.")
|
| 51 |
+
head_path = metadata.get("head_path")
|
| 52 |
+
if not isinstance(head_path, str) or not head_path:
|
| 53 |
+
raise ValueError("Confidence checkpoint metadata requires head_path.")
|
| 54 |
+
relative_path = PurePosixPath(head_path)
|
| 55 |
+
if (
|
| 56 |
+
relative_path.is_absolute()
|
| 57 |
+
or ".." in relative_path.parts
|
| 58 |
+
or "\\" in head_path
|
| 59 |
+
or relative_path.suffix != ".safetensors"
|
| 60 |
+
):
|
| 61 |
+
raise ValueError(
|
| 62 |
+
"Confidence checkpoint head_path must be a relative safetensors path."
|
| 63 |
+
)
|
| 64 |
+
digest = metadata.get("head_sha256")
|
| 65 |
+
if not isinstance(digest, str) or re.fullmatch(r"[0-9a-f]{64}", digest) is None:
|
| 66 |
+
raise ValueError(
|
| 67 |
+
"Confidence checkpoint metadata requires a SHA256 head digest."
|
| 68 |
+
)
|
| 69 |
+
if type(metadata.get("head_size")) is not int or metadata["head_size"] <= 0:
|
| 70 |
+
raise ValueError(
|
| 71 |
+
"Confidence checkpoint metadata requires a positive head_size."
|
| 72 |
+
)
|
| 73 |
+
return metadata
|
| 74 |
+
|
| 75 |
+
|
| 76 |
+
def install_confidence_checkpoint(
|
| 77 |
+
model: ESMFold2ExperimentalModel, *, download_options: dict[str, Any]
|
| 78 |
+
) -> None:
|
| 79 |
+
"""Resolve a rolling pointer once, verify its immutable head, and embed it."""
|
| 80 |
+
|
| 81 |
+
from .modeling_esmfold2_experimental import ConfidenceHead
|
| 82 |
+
|
| 83 |
+
source = model.config.confidence_head_source
|
| 84 |
+
if source is None:
|
| 85 |
+
return
|
| 86 |
+
if model.confidence_head is not None:
|
| 87 |
+
raise ValueError(
|
| 88 |
+
"An external confidence checkpoint cannot replace an embedded head."
|
| 89 |
+
)
|
| 90 |
+
devices = {parameter.device for parameter in model.parameters()}
|
| 91 |
+
device_map = getattr(model, "hf_device_map", {})
|
| 92 |
+
if (
|
| 93 |
+
len(devices) != 1
|
| 94 |
+
or any(device.type == "meta" for device in devices)
|
| 95 |
+
or "disk" in device_map.values()
|
| 96 |
+
):
|
| 97 |
+
raise ValueError(
|
| 98 |
+
"External confidence heads require a single resident model device; offload is unsupported."
|
| 99 |
+
)
|
| 100 |
+
if any(
|
| 101 |
+
getattr(getattr(module, "_hf_hook", None), "offload", False)
|
| 102 |
+
for module in model.modules()
|
| 103 |
+
):
|
| 104 |
+
raise ValueError("External confidence heads do not support an offloaded model.")
|
| 105 |
+
options = {
|
| 106 |
+
key: download_options[key]
|
| 107 |
+
for key in _DOWNLOAD_OPTIONS
|
| 108 |
+
if key in download_options
|
| 109 |
+
}
|
| 110 |
+
latest = Path(
|
| 111 |
+
hf_hub_download(
|
| 112 |
+
repo_id=source["repo_id"],
|
| 113 |
+
repo_type="dataset",
|
| 114 |
+
filename=source["latest_path"],
|
| 115 |
+
revision=source["revision"],
|
| 116 |
+
**options,
|
| 117 |
+
)
|
| 118 |
+
)
|
| 119 |
+
revision = extract_commit_hash(str(latest), None)
|
| 120 |
+
if revision is None or re.fullmatch(r"[0-9a-f]{40}", revision) is None:
|
| 121 |
+
raise ValueError("Could not resolve an immutable confidence dataset revision.")
|
| 122 |
+
metadata = _checkpoint_metadata(latest, source)
|
| 123 |
+
checkpoint = Path(
|
| 124 |
+
hf_hub_download(
|
| 125 |
+
repo_id=source["repo_id"],
|
| 126 |
+
repo_type="dataset",
|
| 127 |
+
filename=metadata["head_path"],
|
| 128 |
+
revision=revision,
|
| 129 |
+
**options,
|
| 130 |
+
)
|
| 131 |
+
)
|
| 132 |
+
if checkpoint.stat().st_size != metadata["head_size"]:
|
| 133 |
+
raise ValueError(
|
| 134 |
+
"Confidence checkpoint size does not match its publication metadata."
|
| 135 |
+
)
|
| 136 |
+
with checkpoint.open("rb") as handle:
|
| 137 |
+
digest = hashlib.file_digest(handle, "sha256").hexdigest()
|
| 138 |
+
if digest != metadata["head_sha256"]:
|
| 139 |
+
raise ValueError(
|
| 140 |
+
"Confidence checkpoint SHA256 does not match its publication metadata."
|
| 141 |
+
)
|
| 142 |
+
state = load_file(str(checkpoint), device="cpu")
|
| 143 |
+
# Head construction must not consume the caller's folding random-number stream.
|
| 144 |
+
with torch.random.fork_rng(devices=[]), torch.device("cpu"):
|
| 145 |
+
head = ConfidenceHead(model.config)
|
| 146 |
+
expected = head.state_dict()
|
| 147 |
+
if set(state) != set(expected):
|
| 148 |
+
raise ValueError("Confidence checkpoint keys do not match the native head.")
|
| 149 |
+
for key, tensor in state.items():
|
| 150 |
+
# Every parameter/buffer must retain the native architecture's exact shape.
|
| 151 |
+
if (
|
| 152 |
+
tensor.shape != expected[key].shape
|
| 153 |
+
or tensor.is_floating_point() != expected[key].is_floating_point()
|
| 154 |
+
):
|
| 155 |
+
raise ValueError(
|
| 156 |
+
f"Confidence checkpoint tensor {key!r} has an incompatible shape or dtype."
|
| 157 |
+
)
|
| 158 |
+
if not torch.isfinite(tensor).all().item():
|
| 159 |
+
raise ValueError(
|
| 160 |
+
f"Confidence checkpoint tensor {key!r} contains nonfinite values."
|
| 161 |
+
)
|
| 162 |
+
head.load_state_dict(state, strict=True)
|
| 163 |
+
parameter = next(model.parameters())
|
| 164 |
+
head.to(device=parameter.device, dtype=parameter.dtype)
|
| 165 |
+
head.train(model.training)
|
| 166 |
+
head.set_kernel_backend(model._kernel_backend)
|
| 167 |
+
model.confidence_head = head
|
| 168 |
+
model.config.confidence_head.enabled = True
|
| 169 |
+
model.config.confidence_head_resolved = {**metadata, "dataset_revision": revision}
|
| 170 |
+
# save_pretrained now persists a self-contained head and its exact provenance.
|
| 171 |
+
model.config.confidence_head_source = None
|
fastplms/models/esmfold2/configuration_esmfold2.py
CHANGED
|
@@ -16,7 +16,10 @@
|
|
| 16 |
|
| 17 |
from __future__ import annotations
|
| 18 |
|
|
|
|
|
|
|
| 19 |
from dataclasses import asdict, dataclass, field
|
|
|
|
| 20 |
from typing import Any, TypeVar, cast
|
| 21 |
from transformers.configuration_utils import PretrainedConfig
|
| 22 |
|
|
@@ -27,6 +30,35 @@ _ESMC_ATTENTION_IMPLEMENTATIONS = frozenset({"eager", "flex_attention", "sdpa"})
|
|
| 27 |
_ESMC_PRECISIONS = frozenset({"auto", "bf16", "fp32", "fp8"})
|
| 28 |
|
| 29 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 30 |
def _esmc_backbone_checkpoint_ids() -> tuple[str, str]:
|
| 31 |
"""Return the manifest-pinned official and FastPLMs ESMC repositories."""
|
| 32 |
|
|
@@ -284,6 +316,25 @@ class ESMFold2Config(PretrainedConfig):
|
|
| 284 |
|
| 285 |
for name, config_type in _NESTED_CONFIGS:
|
| 286 |
setattr(self, name, _nested_config(kwargs.get(name), config_type))
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 287 |
if not isinstance(self.msa_encoder.enabled, bool):
|
| 288 |
raise TypeError("msa_encoder.enabled must be a boolean.")
|
| 289 |
declared_msa_conditioning = kwargs.get("msa_conditioning")
|
|
@@ -332,4 +383,5 @@ __all__ = [
|
|
| 332 |
"ParcaeConfig",
|
| 333 |
"normalize_esmc_attention_implementation",
|
| 334 |
"normalize_esmc_id",
|
|
|
|
| 335 |
]
|
|
|
|
| 16 |
|
| 17 |
from __future__ import annotations
|
| 18 |
|
| 19 |
+
import re
|
| 20 |
+
|
| 21 |
from dataclasses import asdict, dataclass, field
|
| 22 |
+
from pathlib import PurePosixPath
|
| 23 |
from typing import Any, TypeVar, cast
|
| 24 |
from transformers.configuration_utils import PretrainedConfig
|
| 25 |
|
|
|
|
| 30 |
_ESMC_PRECISIONS = frozenset({"auto", "bf16", "fp32", "fp8"})
|
| 31 |
|
| 32 |
|
| 33 |
+
def validate_confidence_head_source(value: Any) -> dict[str, str] | None:
|
| 34 |
+
"""Validate an explicit, independently versioned confidence-head binding."""
|
| 35 |
+
|
| 36 |
+
if value is None:
|
| 37 |
+
return None
|
| 38 |
+
fields = {
|
| 39 |
+
"repo_id", "repo_type", "latest_path", "revision", "model_id", "base_weight_sha256"
|
| 40 |
+
}
|
| 41 |
+
if not isinstance(value, dict) or set(value) != fields:
|
| 42 |
+
raise ValueError(f"confidence_head_source requires exactly {sorted(fields)}.")
|
| 43 |
+
if any(not isinstance(item, str) or not item for item in value.values()):
|
| 44 |
+
raise ValueError("confidence_head_source values must be nonempty strings.")
|
| 45 |
+
if value["repo_type"] != "dataset":
|
| 46 |
+
raise ValueError("External confidence heads must come from a dataset repository.")
|
| 47 |
+
if value["model_id"] not in {"esmfold2_300", "esmfold2_600"}:
|
| 48 |
+
raise ValueError("External confidence heads support esmfold2_300 and esmfold2_600.")
|
| 49 |
+
path = PurePosixPath(value["latest_path"])
|
| 50 |
+
if (
|
| 51 |
+
path.is_absolute()
|
| 52 |
+
or ".." in path.parts
|
| 53 |
+
or "\\" in value["latest_path"]
|
| 54 |
+
or path.suffix != ".json"
|
| 55 |
+
):
|
| 56 |
+
raise ValueError("confidence_head_source.latest_path must be a relative JSON path.")
|
| 57 |
+
if re.fullmatch(r"[0-9a-f]{64}", value["base_weight_sha256"]) is None:
|
| 58 |
+
raise ValueError("confidence_head_source.base_weight_sha256 must be a SHA256 digest.")
|
| 59 |
+
return dict(value)
|
| 60 |
+
|
| 61 |
+
|
| 62 |
def _esmc_backbone_checkpoint_ids() -> tuple[str, str]:
|
| 63 |
"""Return the manifest-pinned official and FastPLMs ESMC repositories."""
|
| 64 |
|
|
|
|
| 316 |
|
| 317 |
for name, config_type in _NESTED_CONFIGS:
|
| 318 |
setattr(self, name, _nested_config(kwargs.get(name), config_type))
|
| 319 |
+
self.confidence_head_source = validate_confidence_head_source(
|
| 320 |
+
kwargs.get("confidence_head_source")
|
| 321 |
+
)
|
| 322 |
+
self.confidence_head_resolved = kwargs.get("confidence_head_resolved")
|
| 323 |
+
if self.confidence_head_source is not None:
|
| 324 |
+
if self.type != "experimental" or self.confidence_head.enabled:
|
| 325 |
+
raise ValueError(
|
| 326 |
+
"External confidence heads require a disabled experimental base head."
|
| 327 |
+
)
|
| 328 |
+
if self.confidence_head_resolved is not None:
|
| 329 |
+
raise ValueError("A confidence head cannot be both external and embedded.")
|
| 330 |
+
if self.confidence_head_resolved is not None:
|
| 331 |
+
if (
|
| 332 |
+
not isinstance(self.confidence_head_resolved, dict)
|
| 333 |
+
or not self.confidence_head.enabled
|
| 334 |
+
):
|
| 335 |
+
raise ValueError(
|
| 336 |
+
"Resolved confidence-head provenance requires an enabled embedded head."
|
| 337 |
+
)
|
| 338 |
if not isinstance(self.msa_encoder.enabled, bool):
|
| 339 |
raise TypeError("msa_encoder.enabled must be a boolean.")
|
| 340 |
declared_msa_conditioning = kwargs.get("msa_conditioning")
|
|
|
|
| 383 |
"ParcaeConfig",
|
| 384 |
"normalize_esmc_attention_implementation",
|
| 385 |
"normalize_esmc_id",
|
| 386 |
+
"validate_confidence_head_source",
|
| 387 |
]
|
fastplms/models/esmfold2/modeling_esmfold2_experimental.py
CHANGED
|
@@ -22,6 +22,7 @@ from transformers.modeling_utils import PreTrainedModel
|
|
| 22 |
|
| 23 |
from .attention import ESMFold2AttentionMixin
|
| 24 |
from .configuration_esmfold2 import ESMFold2Config
|
|
|
|
| 25 |
from .embedding import ESMFold2EmbeddingMixin
|
| 26 |
from .modeling_esmfold2 import (
|
| 27 |
ESMCPrecision,
|
|
@@ -606,6 +607,7 @@ class ESMFold2ExperimentalModel(ESMFold2EmbeddingMixin, ESMFold2AttentionMixin,
|
|
| 606 |
pretrained_model_name_or_path,
|
| 607 |
*model_args,
|
| 608 |
load_esmc: bool = True,
|
|
|
|
| 609 |
**kwargs,
|
| 610 |
):
|
| 611 |
if "config" not in kwargs:
|
|
@@ -620,6 +622,8 @@ class ESMFold2ExperimentalModel(ESMFold2EmbeddingMixin, ESMFold2AttentionMixin,
|
|
| 620 |
model, loading_info = loaded
|
| 621 |
else:
|
| 622 |
model = loaded
|
|
|
|
|
|
|
| 623 |
if load_esmc:
|
| 624 |
model.load_esmc(
|
| 625 |
model.config.esmc_id,
|
|
|
|
| 22 |
|
| 23 |
from .attention import ESMFold2AttentionMixin
|
| 24 |
from .configuration_esmfold2 import ESMFold2Config
|
| 25 |
+
from .confidence_checkpoint import install_confidence_checkpoint
|
| 26 |
from .embedding import ESMFold2EmbeddingMixin
|
| 27 |
from .modeling_esmfold2 import (
|
| 28 |
ESMCPrecision,
|
|
|
|
| 607 |
pretrained_model_name_or_path,
|
| 608 |
*model_args,
|
| 609 |
load_esmc: bool = True,
|
| 610 |
+
load_confidence_head: bool = True,
|
| 611 |
**kwargs,
|
| 612 |
):
|
| 613 |
if "config" not in kwargs:
|
|
|
|
| 622 |
model, loading_info = loaded
|
| 623 |
else:
|
| 624 |
model = loaded
|
| 625 |
+
if load_confidence_head and model.config.confidence_head_source is not None:
|
| 626 |
+
install_confidence_checkpoint(model, download_options=kwargs)
|
| 627 |
if load_esmc:
|
| 628 |
model.load_esmc(
|
| 629 |
model.config.esmc_id,
|
fastplms_bundle.py
CHANGED
|
The diff for this file is too large to render.
See raw diff
|
|
|
modeling_fastplms.py
CHANGED
|
@@ -13,7 +13,7 @@ from zipfile import ZIP_DEFLATED, ZipFile
|
|
| 13 |
|
| 14 |
from .fastplms_bundle import RUNTIME_DATA, RUNTIME_HASH
|
| 15 |
|
| 16 |
-
if RUNTIME_HASH != "
|
| 17 |
raise RuntimeError("FastPLMs runtime identity differs from the bridge.")
|
| 18 |
|
| 19 |
_RUNTIME_TEMPORARIES = []
|
|
|
|
| 13 |
|
| 14 |
from .fastplms_bundle import RUNTIME_DATA, RUNTIME_HASH
|
| 15 |
|
| 16 |
+
if RUNTIME_HASH != "94b14721cefe8f43a5b08716ad38dd2b2030a6c69fd63e8e9edacdc0e5515b23":
|
| 17 |
raise RuntimeError("FastPLMs runtime identity differs from the bridge.")
|
| 18 |
|
| 19 |
_RUNTIME_TEMPORARIES = []
|