lhallee commited on
Commit
121842a
·
verified ·
1 Parent(s): 7b9689e

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 != "ab074f9ea4b20dcfaf9e2ca9082480ee0f9dba02c92d14da5dd3879aa86bf7ed":
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 = []