lhallee commited on
Commit
2dbf793
·
verified ·
1 Parent(s): 6577ccb

Update FastPLMs runtime and model cards

Browse files
README.md CHANGED
@@ -40,9 +40,8 @@ This model requires Python 3.11-3.14, PyTorch 2.13, and Transformers 5.13.
40
 
41
  The artifact requirements include the structure dependencies.
42
 
43
- The release contract requires a CUDA device. The current validated target is
44
- the exact NVIDIA GH200 on Linux aarch64. Linux x86-64, CPU-only, Windows, and
45
- macOS structure runs are not release evidence.
46
 
47
  The Hub quick start needs network access for the first download. For an
48
  air-gapped run, build the manifest-pinned local artifact first and use the
@@ -153,6 +152,7 @@ with torch.inference_mode():
153
  output = model.infer(
154
  "MKTLLILAVVAAALA",
155
  num_recycles=4,
 
156
  )
157
 
158
  print(output["mean_plddt"])
@@ -185,7 +185,7 @@ folding requests raise.
185
  - Redistributable: `true`
186
  - Complete weight publication required: `false`
187
 
188
- ## Validation and provenance
189
 
190
  FastPLMs pins the checkpoint, upstream source revisions, state transformation,
191
  and required files in `models.toml`. Built artifacts record exact source
 
40
 
41
  The artifact requirements include the structure dependencies.
42
 
43
+ Validation runs in Docker on any compatible CUDA device. Record the container,
44
+ hardware, precision, and inputs; no GPU product or workstation is required.
 
45
 
46
  The Hub quick start needs network access for the first download. For an
47
  air-gapped run, build the manifest-pinned local artifact first and use the
 
152
  output = model.infer(
153
  "MKTLLILAVVAAALA",
154
  num_recycles=4,
155
+ verbose=False,
156
  )
157
 
158
  print(output["mean_plddt"])
 
185
  - Redistributable: `true`
186
  - Complete weight publication required: `false`
187
 
188
+ ## Validation and sources
189
 
190
  FastPLMs pins the checkpoint, upstream source revisions, state transformation,
191
  and required files in `models.toml`. Built artifacts record exact source
fastplms/models.toml CHANGED
@@ -392,7 +392,7 @@ checkpoint_license = "MIT"
392
  hub_license = "mit"
393
  weights_publication_allowed = true
394
  state_transform = "identity"
395
- conversion_provenance = "Input: each pinned Biohub ESMFold2 checkpoint and its separately pinned ESMC checkpoint. Transformation: apply identity to preserve the folding checkpoint exactly, load its parameters in FP32 for CUDA BF16-autocast execution, retain canonical BF16 ESMC weights, and optionally rebuild exactly 80 ESMC attention output projections as transient Transformer Engine linears. Output: the corresponding pinned Synthyra ESMFold2 checkpoint plus its declared ESMC precision policy. Validation: release parity covers exact canonical state, learned projection, prepared features, and seeded BF16 folding; experimental FP8 validation covers strict unavailable-device behavior, all four variants, and three BF16-to-FP8 reload cycles on the standard variant. Limitation: only the four manifest-listed ESMFold2 variants are supported; FP8 is experimental, applies only to inference-time ESMC execution, and requires direct CUDA loading with Transformer Engine availability."
396
  representative = "esmfold2"
397
  documentation = "docs/esmfold2.md"
398
  test_tiers = ["check", "compliance", "structure", "feature", "artifact", "benchmark"]
@@ -1237,3 +1237,39 @@ official_files = [
1237
  "model.safetensors=sha256:4e903b740ad6ad704ec60881bfd593e0d6c874a630ffa0f0838276e0b665088f",
1238
  ]
1239
  auto_map = { AutoConfig = "fastplms.models.esmfold2.configuration_esmfold2.ESMFold2Config", AutoModel = "fastplms.models.esmfold2.modeling_esmfold2_experimental.ESMFold2ExperimentalModel", AutoModelForSequenceClassification = "fastplms.models.esmfold2.modeling_esmfold2_classification.ESMFold2ExperimentalForSequenceClassification", AutoModelForTokenClassification = "fastplms.models.esmfold2.modeling_esmfold2_classification.ESMFold2ExperimentalForTokenClassification" }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
392
  hub_license = "mit"
393
  weights_publication_allowed = true
394
  state_transform = "identity"
395
+ conversion_provenance = "Input: each pinned Biohub ESMFold2 checkpoint and its separately pinned ESMC checkpoint. Transformation: apply identity to preserve the folding checkpoint exactly, load its parameters in FP32 for CUDA BF16-autocast execution, retain canonical BF16 ESMC weights, and optionally rebuild exactly 80 ESMC attention output projections as transient Transformer Engine linears. Output: the corresponding pinned Synthyra ESMFold2 checkpoint plus its declared ESMC precision policy. Validation: release parity covers exact canonical state, learned projection, prepared features, and seeded BF16 folding; experimental FP8 validation covers strict unavailable-device behavior, the four 6B-backbone variants, and three BF16-to-FP8 reload cycles on the standard variant. Limitation: only the six manifest-listed ESMFold2 variants are supported; FP8 is experimental, applies only to inference-time ESMC execution, and requires direct CUDA loading with Transformer Engine availability."
396
  representative = "esmfold2"
397
  documentation = "docs/esmfold2.md"
398
  test_tiers = ["check", "compliance", "structure", "feature", "artifact", "benchmark"]
 
1237
  "model.safetensors=sha256:4e903b740ad6ad704ec60881bfd593e0d6c874a630ffa0f0838276e0b665088f",
1238
  ]
1239
  auto_map = { AutoConfig = "fastplms.models.esmfold2.configuration_esmfold2.ESMFold2Config", AutoModel = "fastplms.models.esmfold2.modeling_esmfold2_experimental.ESMFold2ExperimentalModel", AutoModelForSequenceClassification = "fastplms.models.esmfold2.modeling_esmfold2_classification.ESMFold2ExperimentalForSequenceClassification", AutoModelForTokenClassification = "fastplms.models.esmfold2.modeling_esmfold2_classification.ESMFold2ExperimentalForTokenClassification" }
1240
+
1241
+ [[models]]
1242
+ id = "esmfold2_300"
1243
+ family = "esmfold2"
1244
+ size_category = "structure"
1245
+ generation_contract = "not_applicable"
1246
+ msa_conditioning = false
1247
+ publication_status = "published"
1248
+ fast_repo = "Synthyra/ESMFold2-300"
1249
+ fast_revision = "a38a62ae930d157484b331c2bf4241684573adba"
1250
+ fast_files = ["config.json=git-sha1:47ec20cf8b234c3b41d6f3ae1bdfe95d4eb4849e", "model.safetensors=sha256:44d6797c5efebf24753d502b40950e0874871c96ceea14f2d7f7e39cebac67fd"]
1251
+ official_repo = "biohub/ESMFold2-Experimental-Fast-base300M-step1500k"
1252
+ official_revision = "21531e59002c9205284715e28ee802dafb430637"
1253
+ official_files = ["config.json=git-sha1:8c9a04fe22b0e5fca77bc4e2861a12c9494ef4d4", "model.safetensors=sha256:44d6797c5efebf24753d502b40950e0874871c96ceea14f2d7f7e39cebac67fd"]
1254
+ notes = "Experimental Fast checkpoint with a frozen 300M ESM++ backbone, tensor-exact in BF16 with the pinned step-1500000 source, 24 folding blocks, no MSA conditioning, and no confidence head. BF16 execution uses FP32 folding parameters with CUDA autocast. FP8 is unsupported. Docker BF16 inference validation passed on the compact Protein G case."
1255
+ auto_map = { AutoConfig = "fastplms.models.esmfold2.configuration_esmfold2.ESMFold2Config", AutoModel = "fastplms.models.esmfold2.modeling_esmfold2_experimental.ESMFold2ExperimentalModel", AutoModelForSequenceClassification = "fastplms.models.esmfold2.modeling_esmfold2_classification.ESMFold2ExperimentalForSequenceClassification", AutoModelForTokenClassification = "fastplms.models.esmfold2.modeling_esmfold2_classification.ESMFold2ExperimentalForTokenClassification" }
1256
+ backbone_model = "esmc_small"
1257
+ backbone = { repo = "biohub/ESMC-300M-1500000", revision = "56803b6378b82e16c3b24aac49d1fce4445540b7", files = ["config.json=git-sha1:7fe728a0eb3fb81b24491d6cc1de816bf7797c27", "model.safetensors=sha256:8bd6cacf9b5a92d51954b64b20407f1f9f564a7e4849b8663470784d2a8b7ed2", "tokenizer.json=git-sha1:81c797f56768b22dec0301fa771f018b7e43e98c", "tokenizer_config.json=git-sha1:f49f57b24a1c93bd544974811e8ecbd61b7fae89", "special_tokens_map.json=git-sha1:c907ee1dc19b24241749b32d665c291c7e6e8e4b"] }
1258
+
1259
+ [[models]]
1260
+ id = "esmfold2_600"
1261
+ family = "esmfold2"
1262
+ size_category = "structure"
1263
+ generation_contract = "not_applicable"
1264
+ msa_conditioning = false
1265
+ publication_status = "published"
1266
+ fast_repo = "Synthyra/ESMFold2-600"
1267
+ fast_revision = "71c67d0b2b73dc245ea7c3cc0d0476439a882d08"
1268
+ fast_files = ["config.json=git-sha1:8e271837cbdada96c4974c8e543f84065e0f06f1", "model.safetensors=sha256:11a53c1b4700b6c62a5a464fc3ec7076160c19e8584157f88e13275116cfd602"]
1269
+ official_repo = "biohub/ESMFold2-Experimental-Fast-base600M-step1500k"
1270
+ official_revision = "15cf2d6648692f6c17cee1297d8a285476fffa9b"
1271
+ official_files = ["config.json=git-sha1:95517e555f17a1eca4b68866c89033a1f6916a5d", "model.safetensors=sha256:11a53c1b4700b6c62a5a464fc3ec7076160c19e8584157f88e13275116cfd602"]
1272
+ notes = "Experimental Fast checkpoint with a frozen 600M ESM++ backbone, tensor-exact in BF16 with the pinned step-1500000 source, 24 folding blocks, no MSA conditioning, and no confidence head. BF16 execution uses FP32 folding parameters with CUDA autocast. FP8 is unsupported. Configuration, weight identities, and artifact reload are verified; this model is not inference-validated."
1273
+ auto_map = { AutoConfig = "fastplms.models.esmfold2.configuration_esmfold2.ESMFold2Config", AutoModel = "fastplms.models.esmfold2.modeling_esmfold2_experimental.ESMFold2ExperimentalModel", AutoModelForSequenceClassification = "fastplms.models.esmfold2.modeling_esmfold2_classification.ESMFold2ExperimentalForSequenceClassification", AutoModelForTokenClassification = "fastplms.models.esmfold2.modeling_esmfold2_classification.ESMFold2ExperimentalForTokenClassification" }
1274
+ backbone_model = "esmc_large"
1275
+ backbone = { repo = "biohub/ESMC-600M-1500000", revision = "21af9cc429af76ebda6c48074fb624db4735aaaf", files = ["config.json=git-sha1:ec29f6009b21d710f64bf1c058f3a9710833d692", "model.safetensors=sha256:d6869f5ae0f11e5dc829b195e062e87cfcc2f851a08a5edbaf5d1083ae7f76cc", "tokenizer.json=git-sha1:81c797f56768b22dec0301fa771f018b7e43e98c", "tokenizer_config.json=git-sha1:f49f57b24a1c93bd544974811e8ecbd61b7fae89", "special_tokens_map.json=git-sha1:c907ee1dc19b24241749b32d665c291c7e6e8e4b"] }
fastplms/models/esmfold/modeling_fast_esmfold.py CHANGED
@@ -16,10 +16,13 @@ from __future__ import annotations
16
 
17
  import torch
18
  import torch.nn as nn
 
19
  from contextvars import ContextVar
20
  from dataclasses import dataclass
 
21
  from typing import Any
22
  from einops import rearrange
 
23
  from torch.nn import functional as F
24
  from transformers.modeling_outputs import ModelOutput
25
  from transformers.models.esm.configuration_esm import EsmConfig
@@ -142,6 +145,78 @@ _ESMFOLD_CAPTURED_ATTENTIONS: ContextVar[
142
  )
143
 
144
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
145
  def _align_internal_esm_attentions(
146
  attentions: tuple[torch.Tensor, ...],
147
  residue_mask: torch.Tensor,
@@ -674,8 +749,13 @@ class FastEsmForProteinFolding(FastPLMsAttentionMixin, EsmForProteinFolding):
674
  output_attentions: bool | None = None,
675
  output_hidden_states: bool | None = None,
676
  return_dict: bool | None = None,
 
677
  ) -> FastEsmForProteinFoldingOutput | tuple[Any, ...]:
678
- """Run folding with Meta ESMFold's 0-to-100 pLDDT convention."""
 
 
 
 
679
 
680
  config = getattr(self, "config", None)
681
  resolved_attentions = (
@@ -698,14 +778,19 @@ class FastEsmForProteinFolding(FastPLMsAttentionMixin, EsmForProteinFolding):
698
  capture_token = _ESMFOLD_CAPTURED_ATTENTIONS.set(None)
699
  captured_attentions: tuple[torch.Tensor, ...] | None = None
700
  try:
701
- output = super().forward(
702
- input_ids,
703
- attention_mask=attention_mask,
704
- position_ids=position_ids,
705
- masking_pattern=masking_pattern,
706
  num_recycles=num_recycles,
707
- output_hidden_states=resolved_hidden_states,
708
- )
 
 
 
 
 
 
 
 
709
  captured_attentions = _ESMFOLD_CAPTURED_ATTENTIONS.get()
710
  finally:
711
  _ESMFOLD_CAPTURED_ATTENTIONS.reset(capture_token)
@@ -743,12 +828,14 @@ class FastEsmForProteinFolding(FastPLMsAttentionMixin, EsmForProteinFolding):
743
  num_recycles: int | None = None,
744
  residue_index_offset: int | None = 512,
745
  chain_linker: str | None = "G" * 25,
 
746
  ):
747
  """Fold raw sequences through Meta ESMFold's public input contract.
748
 
749
  Transformers v5 narrows ``infer`` even though ``forward`` retains the
750
  required controls. This adapter restores recycle selection, explicit
751
  residue indices, masking, and colon-delimited multimer preparation.
 
752
  """
753
 
754
  sequence_batch = [sequences] if isinstance(sequences, str) else sequences
@@ -816,6 +903,7 @@ class FastEsmForProteinFolding(FastPLMsAttentionMixin, EsmForProteinFolding):
816
  position_ids=residx,
817
  masking_pattern=masking_pattern,
818
  num_recycles=num_recycles,
 
819
  )
820
  output["atom37_atom_exists"] = output["atom37_atom_exists"] * linker_mask.unsqueeze(2)
821
  output["mean_plddt"] = (output["plddt"] * output["atom37_atom_exists"]).sum(
 
16
 
17
  import torch
18
  import torch.nn as nn
19
+ from collections.abc import Iterator
20
  from contextvars import ContextVar
21
  from dataclasses import dataclass
22
+ from contextlib import contextmanager
23
  from typing import Any
24
  from einops import rearrange
25
+ from tqdm.auto import tqdm
26
  from torch.nn import functional as F
27
  from transformers.modeling_outputs import ModelOutput
28
  from transformers.models.esm.configuration_esm import EsmConfig
 
145
  )
146
 
147
 
148
+ @contextmanager
149
+ def _folding_progress(
150
+ model: nn.Module,
151
+ *,
152
+ num_recycles: int | None,
153
+ verbose: bool,
154
+ ) -> Iterator[None]:
155
+ """Report actual embedding, recycling, and confidence stages."""
156
+
157
+ if not verbose:
158
+ yield
159
+ return
160
+
161
+ trunk = getattr(model, "trunk", None)
162
+ blocks = getattr(trunk, "blocks", None)
163
+ structure_module = getattr(trunk, "structure_module", None)
164
+ embedding = getattr(model, "esm", None)
165
+ confidence_heads = tuple(
166
+ getattr(model, name, None)
167
+ for name in ("distogram_head", "lm_head", "lddt_head", "ptm_head")
168
+ )
169
+ if not isinstance(blocks, nn.ModuleList) or not isinstance(structure_module, nn.Module):
170
+ raise RuntimeError("ESMFold progress requires the standard folding trunk modules.")
171
+
172
+ if num_recycles is None:
173
+ recycle_passes = getattr(getattr(trunk, "config", None), "max_recycles", 0)
174
+ else:
175
+ recycle_passes = num_recycles + 1
176
+ confidence_modules = tuple(
177
+ head for head in confidence_heads if isinstance(head, nn.Module)
178
+ )
179
+ total = (
180
+ int(isinstance(embedding, nn.Module))
181
+ + max(int(recycle_passes), 0) * (len(blocks) + 1)
182
+ + len(confidence_modules)
183
+ )
184
+ progress = tqdm(total=total, desc="ESMFold embeddings", unit="stage")
185
+
186
+ def update_progress(stage: str, *_args: Any) -> None:
187
+ progress.set_description(f"ESMFold {stage}")
188
+ progress.update(1)
189
+
190
+ handles = []
191
+ try:
192
+ if isinstance(embedding, nn.Module):
193
+ handles.append(
194
+ embedding.register_forward_hook(
195
+ lambda *_args: update_progress("embeddings", *_args)
196
+ )
197
+ )
198
+ for block in blocks:
199
+ handles.append(
200
+ block.register_forward_hook(
201
+ lambda *_args: update_progress("recycling", *_args)
202
+ )
203
+ )
204
+ handles.append(
205
+ structure_module.register_forward_hook(
206
+ lambda *_args: update_progress("recycling", *_args)
207
+ )
208
+ )
209
+ handles.extend(
210
+ head.register_forward_hook(lambda *_args: update_progress("confidence", *_args))
211
+ for head in confidence_modules
212
+ )
213
+ yield
214
+ finally:
215
+ for handle in handles:
216
+ handle.remove()
217
+ progress.close()
218
+
219
+
220
  def _align_internal_esm_attentions(
221
  attentions: tuple[torch.Tensor, ...],
222
  residue_mask: torch.Tensor,
 
749
  output_attentions: bool | None = None,
750
  output_hidden_states: bool | None = None,
751
  return_dict: bool | None = None,
752
+ verbose: bool = False,
753
  ) -> FastEsmForProteinFoldingOutput | tuple[Any, ...]:
754
+ """Run folding with Meta ESMFold's 0-to-100 pLDDT convention.
755
+
756
+ Set ``verbose=True`` to display progress for embeddings, recycling, and
757
+ confidence heads. The default keeps the call silent.
758
+ """
759
 
760
  config = getattr(self, "config", None)
761
  resolved_attentions = (
 
778
  capture_token = _ESMFOLD_CAPTURED_ATTENTIONS.set(None)
779
  captured_attentions: tuple[torch.Tensor, ...] | None = None
780
  try:
781
+ with _folding_progress(
782
+ self,
 
 
 
783
  num_recycles=num_recycles,
784
+ verbose=verbose,
785
+ ):
786
+ output = super().forward(
787
+ input_ids,
788
+ attention_mask=attention_mask,
789
+ position_ids=position_ids,
790
+ masking_pattern=masking_pattern,
791
+ num_recycles=num_recycles,
792
+ output_hidden_states=resolved_hidden_states,
793
+ )
794
  captured_attentions = _ESMFOLD_CAPTURED_ATTENTIONS.get()
795
  finally:
796
  _ESMFOLD_CAPTURED_ATTENTIONS.reset(capture_token)
 
828
  num_recycles: int | None = None,
829
  residue_index_offset: int | None = 512,
830
  chain_linker: str | None = "G" * 25,
831
+ verbose: bool = False,
832
  ):
833
  """Fold raw sequences through Meta ESMFold's public input contract.
834
 
835
  Transformers v5 narrows ``infer`` even though ``forward`` retains the
836
  required controls. This adapter restores recycle selection, explicit
837
  residue indices, masking, and colon-delimited multimer preparation.
838
+ Set ``verbose=True`` to display folding progress.
839
  """
840
 
841
  sequence_batch = [sequences] if isinstance(sequences, str) else sequences
 
903
  position_ids=residx,
904
  masking_pattern=masking_pattern,
905
  num_recycles=num_recycles,
906
+ verbose=verbose,
907
  )
908
  output["atom37_atom_exists"] = output["atom37_atom_exists"] * linker_mask.unsqueeze(2)
909
  output["mean_plddt"] = (output["plddt"] * output["atom37_atom_exists"]).sum(
fastplms/registry.py CHANGED
The diff for this file is too large to render. See raw diff
 
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 != "1526d8452bea4d4b6015a989393b467d1d8c4a98553a1299586a6d6eb28b83a0":
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 != "21bd51688943411febc393b21b61f6efb210ea084dc4614ef0ee6f71f877c0c5":
17
  raise RuntimeError("FastPLMs runtime identity differs from the bridge.")
18
 
19
  _RUNTIME_TEMPORARIES = []