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
Apply coding standards from 1cb5747 (files only)
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- README.md +10 -6
- fastplms/attention/_core.py +1 -0
- fastplms/attention/_kernel_lock.py +1 -0
- fastplms/attention/interfaces.py +1 -0
- fastplms/embeddings/batches.py +379 -378
- fastplms/embeddings/identity.py +508 -507
- fastplms/embeddings/inputs.py +264 -263
- fastplms/embeddings/output.py +215 -215
- fastplms/embeddings/pooling.py +1 -0
- fastplms/embeddings/runner.py +420 -419
- fastplms/embeddings/storage.py +1 -0
- fastplms/models/_esm_rotary.py +1 -0
- fastplms/models/classification_probe.py +21 -20
- fastplms/models/esm_plusplus/modeling_esm_plusplus.py +25 -25
- fastplms/models/esmfold2/__init__.py +1 -0
- fastplms/models/esmfold2/configuration_esmfold2.py +0 -1
- fastplms/models/esmfold2/embedding.py +1 -0
- fastplms/models/esmfold2/esmfold2_affine3d.py +134 -112
- fastplms/models/esmfold2/esmfold2_aligner.py +19 -18
- fastplms/models/esmfold2/esmfold2_atom_indexer.py +3 -3
- fastplms/models/esmfold2/esmfold2_conformers.py +3 -3
- fastplms/models/esmfold2/esmfold2_input_builder.py +3 -2
- fastplms/models/esmfold2/esmfold2_metrics.py +69 -54
- fastplms/models/esmfold2/esmfold2_misc.py +73 -58
- fastplms/models/esmfold2/esmfold2_mmcif_parsing.py +5 -4
- fastplms/models/esmfold2/esmfold2_molecular_complex.py +6 -6
- fastplms/models/esmfold2/esmfold2_msa.py +3 -2
- fastplms/models/esmfold2/esmfold2_msa_filter_sequences.py +15 -13
- fastplms/models/esmfold2/esmfold2_normalize_coordinates.py +24 -21
- fastplms/models/esmfold2/esmfold2_output.py +11 -10
- fastplms/models/esmfold2/esmfold2_paired_msa.py +3 -2
- fastplms/models/esmfold2/esmfold2_parsing.py +4 -3
- fastplms/models/esmfold2/esmfold2_predicted_aligned_error.py +45 -39
- fastplms/models/esmfold2/esmfold2_prepare_input.py +4 -3
- fastplms/models/esmfold2/esmfold2_processor.py +4 -3
- fastplms/models/esmfold2/esmfold2_protein_chain.py +200 -187
- fastplms/models/esmfold2/esmfold2_protein_complex.py +104 -97
- fastplms/models/esmfold2/esmfold2_protein_structure.py +65 -58
- fastplms/models/esmfold2/esmfold2_residue_constants.py +2 -1
- fastplms/models/esmfold2/esmfold2_sequential_dataclass.py +3 -2
- fastplms/models/esmfold2/esmfold2_system.py +2 -0
- fastplms/models/esmfold2/esmfold2_types.py +1 -0
- fastplms/models/esmfold2/esmfold2_utils_types.py +2 -0
- fastplms/models/esmfold2/modeling_esmfold2.py +219 -202
- fastplms/models/esmfold2/modeling_esmfold2_classification.py +21 -20
- fastplms/models/esmfold2/modeling_esmfold2_common.py +494 -450
- fastplms/models/esmfold2/modeling_esmfold2_experimental.py +194 -185
- fastplms/models/esmfold2/protein_utils.py +44 -43
- fastplms/models/esmfold2/reproducibility.py +3 -3
- fastplms/models/ttt.py +1 -0
README.md
CHANGED
|
@@ -16,11 +16,12 @@ Load the published model, fold two protein chains together, and write an mmCIF
|
|
| 16 |
file. The example omits `num_sampling_steps` and uses the model default.
|
| 17 |
|
| 18 |
```python
|
| 19 |
-
from pathlib import Path
|
| 20 |
-
|
| 21 |
import torch
|
|
|
|
|
|
|
| 22 |
from transformers import AutoModel
|
| 23 |
|
|
|
|
| 24 |
model = AutoModel.from_pretrained(
|
| 25 |
"Synthyra/ESMFold2-Fast",
|
| 26 |
trust_remote_code=True,
|
|
@@ -109,11 +110,13 @@ state mixture and projection, followed by one trainable transformer probe.
|
|
| 109 |
|
| 110 |
```python
|
| 111 |
import torch
|
|
|
|
| 112 |
from transformers import (
|
| 113 |
AutoModelForSequenceClassification,
|
| 114 |
AutoModelForTokenClassification,
|
| 115 |
)
|
| 116 |
|
|
|
|
| 117 |
model_id = "Synthyra/ESMFold2-Fast"
|
| 118 |
sequence_model = AutoModelForSequenceClassification.from_pretrained(
|
| 119 |
model_id, num_labels=2, trust_remote_code=True
|
|
@@ -123,11 +126,11 @@ token_model = AutoModelForTokenClassification.from_pretrained(
|
|
| 123 |
).eval()
|
| 124 |
sequences = ["MSTNPKPQRKTKRNT", "MKTIIALSYIFCLVFA"]
|
| 125 |
batch = sequence_model.prepare_classifier_inputs(sequences)
|
| 126 |
-
biological = batch["attention_mask"].bool()
|
| 127 |
|
| 128 |
-
sequence_labels = torch.zeros(len(sequences), dtype=torch.long)
|
| 129 |
-
token_labels = torch.full_like(batch["input_ids"], -100)
|
| 130 |
-
token_labels[biological] = 0
|
| 131 |
|
| 132 |
with torch.inference_mode():
|
| 133 |
sequence_output = sequence_model(**batch, labels=sequence_labels)
|
|
@@ -147,6 +150,7 @@ python -m pip install "datasets>=4.8,<5" "peft>=0.19,<0.20"
|
|
| 147 |
```python
|
| 148 |
from peft import LoraConfig, TaskType, get_peft_model
|
| 149 |
|
|
|
|
| 150 |
peft_model = get_peft_model(
|
| 151 |
sequence_model,
|
| 152 |
LoraConfig(
|
|
|
|
| 16 |
file. The example omits `num_sampling_steps` and uses the model default.
|
| 17 |
|
| 18 |
```python
|
|
|
|
|
|
|
| 19 |
import torch
|
| 20 |
+
|
| 21 |
+
from pathlib import Path
|
| 22 |
from transformers import AutoModel
|
| 23 |
|
| 24 |
+
|
| 25 |
model = AutoModel.from_pretrained(
|
| 26 |
"Synthyra/ESMFold2-Fast",
|
| 27 |
trust_remote_code=True,
|
|
|
|
| 110 |
|
| 111 |
```python
|
| 112 |
import torch
|
| 113 |
+
|
| 114 |
from transformers import (
|
| 115 |
AutoModelForSequenceClassification,
|
| 116 |
AutoModelForTokenClassification,
|
| 117 |
)
|
| 118 |
|
| 119 |
+
|
| 120 |
model_id = "Synthyra/ESMFold2-Fast"
|
| 121 |
sequence_model = AutoModelForSequenceClassification.from_pretrained(
|
| 122 |
model_id, num_labels=2, trust_remote_code=True
|
|
|
|
| 126 |
).eval()
|
| 127 |
sequences = ["MSTNPKPQRKTKRNT", "MKTIIALSYIFCLVFA"]
|
| 128 |
batch = sequence_model.prepare_classifier_inputs(sequences)
|
| 129 |
+
biological = batch["attention_mask"].bool() # (b, l)
|
| 130 |
|
| 131 |
+
sequence_labels = torch.zeros(len(sequences), dtype=torch.long) # (b,)
|
| 132 |
+
token_labels = torch.full_like(batch["input_ids"], -100) # (b, l)
|
| 133 |
+
token_labels[biological] = 0 # selected biological positions; labels stay (b, l)
|
| 134 |
|
| 135 |
with torch.inference_mode():
|
| 136 |
sequence_output = sequence_model(**batch, labels=sequence_labels)
|
|
|
|
| 150 |
```python
|
| 151 |
from peft import LoraConfig, TaskType, get_peft_model
|
| 152 |
|
| 153 |
+
|
| 154 |
peft_model = get_peft_model(
|
| 155 |
sequence_model,
|
| 156 |
LoraConfig(
|
fastplms/attention/_core.py
CHANGED
|
@@ -9,6 +9,7 @@ from __future__ import annotations
|
|
| 9 |
|
| 10 |
import warnings
|
| 11 |
import torch
|
|
|
|
| 12 |
from collections import OrderedDict
|
| 13 |
from collections.abc import Callable
|
| 14 |
from dataclasses import dataclass
|
|
|
|
| 9 |
|
| 10 |
import warnings
|
| 11 |
import torch
|
| 12 |
+
|
| 13 |
from collections import OrderedDict
|
| 14 |
from collections.abc import Callable
|
| 15 |
from dataclasses import dataclass
|
fastplms/attention/_kernel_lock.py
CHANGED
|
@@ -4,6 +4,7 @@ from __future__ import annotations
|
|
| 4 |
|
| 5 |
import json
|
| 6 |
import os
|
|
|
|
| 7 |
from pathlib import Path
|
| 8 |
from typing import Any
|
| 9 |
|
|
|
|
| 4 |
|
| 5 |
import json
|
| 6 |
import os
|
| 7 |
+
|
| 8 |
from pathlib import Path
|
| 9 |
from typing import Any
|
| 10 |
|
fastplms/attention/interfaces.py
CHANGED
|
@@ -3,6 +3,7 @@
|
|
| 3 |
from __future__ import annotations
|
| 4 |
|
| 5 |
import torch
|
|
|
|
| 6 |
from collections.abc import Mapping
|
| 7 |
from functools import partial
|
| 8 |
from typing import Any
|
|
|
|
| 3 |
from __future__ import annotations
|
| 4 |
|
| 5 |
import torch
|
| 6 |
+
|
| 7 |
from collections.abc import Mapping
|
| 8 |
from functools import partial
|
| 9 |
from typing import Any
|
fastplms/embeddings/batches.py
CHANGED
|
@@ -1,378 +1,379 @@
|
|
| 1 |
-
"""Execute model-specific batches and return ordered residue-aware CPU tensors."""
|
| 2 |
-
|
| 3 |
-
from __future__ import annotations
|
| 4 |
-
|
| 5 |
-
import torch
|
| 6 |
-
|
| 7 |
-
from
|
| 8 |
-
from
|
| 9 |
-
from
|
| 10 |
-
from
|
| 11 |
-
|
| 12 |
-
|
| 13 |
-
from .
|
| 14 |
-
from .
|
| 15 |
-
from .
|
| 16 |
-
|
| 17 |
-
|
| 18 |
-
|
| 19 |
-
|
| 20 |
-
|
| 21 |
-
|
| 22 |
-
|
| 23 |
-
|
| 24 |
-
|
| 25 |
-
|
| 26 |
-
|
| 27 |
-
|
| 28 |
-
|
| 29 |
-
|
| 30 |
-
|
| 31 |
-
|
| 32 |
-
|
| 33 |
-
|
| 34 |
-
|
| 35 |
-
|
| 36 |
-
|
| 37 |
-
|
| 38 |
-
|
| 39 |
-
|
| 40 |
-
|
| 41 |
-
|
| 42 |
-
|
| 43 |
-
|
| 44 |
-
|
| 45 |
-
|
| 46 |
-
|
| 47 |
-
|
| 48 |
-
|
| 49 |
-
|
| 50 |
-
|
| 51 |
-
|
| 52 |
-
|
| 53 |
-
|
| 54 |
-
|
| 55 |
-
|
| 56 |
-
|
| 57 |
-
|
| 58 |
-
|
| 59 |
-
|
| 60 |
-
|
| 61 |
-
|
| 62 |
-
|
| 63 |
-
|
| 64 |
-
|
| 65 |
-
|
| 66 |
-
|
| 67 |
-
|
| 68 |
-
|
| 69 |
-
|
| 70 |
-
|
| 71 |
-
|
| 72 |
-
or not callable(
|
| 73 |
-
|
| 74 |
-
|
| 75 |
-
|
| 76 |
-
|
| 77 |
-
|
| 78 |
-
|
| 79 |
-
|
| 80 |
-
|
| 81 |
-
|
| 82 |
-
|
| 83 |
-
|
| 84 |
-
|
| 85 |
-
|
| 86 |
-
|
| 87 |
-
|
| 88 |
-
|
| 89 |
-
|
| 90 |
-
|
| 91 |
-
|
| 92 |
-
|
| 93 |
-
|
| 94 |
-
|
| 95 |
-
|
| 96 |
-
|
| 97 |
-
|
| 98 |
-
|
| 99 |
-
|
| 100 |
-
|
| 101 |
-
|
| 102 |
-
|
| 103 |
-
|
| 104 |
-
|
| 105 |
-
|
| 106 |
-
|
| 107 |
-
|
| 108 |
-
|
| 109 |
-
|
| 110 |
-
|
| 111 |
-
|
| 112 |
-
|
| 113 |
-
|
| 114 |
-
|
| 115 |
-
|
| 116 |
-
|
| 117 |
-
|
| 118 |
-
|
| 119 |
-
|
| 120 |
-
|
| 121 |
-
|
| 122 |
-
|
| 123 |
-
|
| 124 |
-
|
| 125 |
-
|
| 126 |
-
|
| 127 |
-
|
| 128 |
-
|
| 129 |
-
|
| 130 |
-
|
| 131 |
-
#
|
| 132 |
-
|
| 133 |
-
|
| 134 |
-
|
| 135 |
-
|
| 136 |
-
|
| 137 |
-
|
| 138 |
-
|
| 139 |
-
|
| 140 |
-
|
| 141 |
-
|
| 142 |
-
|
| 143 |
-
|
| 144 |
-
"
|
| 145 |
-
"
|
| 146 |
-
|
| 147 |
-
|
| 148 |
-
|
| 149 |
-
#
|
| 150 |
-
#
|
| 151 |
-
|
| 152 |
-
|
| 153 |
-
|
| 154 |
-
|
| 155 |
-
|
| 156 |
-
|
| 157 |
-
|
| 158 |
-
|
| 159 |
-
|
| 160 |
-
|
| 161 |
-
|
| 162 |
-
|
| 163 |
-
|
| 164 |
-
|
| 165 |
-
|
| 166 |
-
|
| 167 |
-
|
| 168 |
-
|
| 169 |
-
|
| 170 |
-
#
|
| 171 |
-
|
| 172 |
-
|
| 173 |
-
|
| 174 |
-
|
| 175 |
-
|
| 176 |
-
|
| 177 |
-
|
| 178 |
-
|
| 179 |
-
|
| 180 |
-
|
| 181 |
-
|
| 182 |
-
|
| 183 |
-
|
| 184 |
-
|
| 185 |
-
|
| 186 |
-
|
| 187 |
-
|
| 188 |
-
|
| 189 |
-
|
| 190 |
-
|
| 191 |
-
|
| 192 |
-
|
| 193 |
-
|
| 194 |
-
|
| 195 |
-
|
| 196 |
-
|
| 197 |
-
|
| 198 |
-
|
| 199 |
-
|
| 200 |
-
|
| 201 |
-
|
| 202 |
-
|
| 203 |
-
|
| 204 |
-
|
| 205 |
-
|
| 206 |
-
|
| 207 |
-
|
| 208 |
-
|
| 209 |
-
|
| 210 |
-
|
| 211 |
-
|
| 212 |
-
|
| 213 |
-
|
| 214 |
-
|
| 215 |
-
|
| 216 |
-
|
| 217 |
-
|
| 218 |
-
|
| 219 |
-
self.
|
| 220 |
-
|
| 221 |
-
|
| 222 |
-
|
| 223 |
-
|
| 224 |
-
|
| 225 |
-
|
| 226 |
-
|
| 227 |
-
|
| 228 |
-
|
| 229 |
-
|
| 230 |
-
|
| 231 |
-
|
| 232 |
-
|
| 233 |
-
|
| 234 |
-
|
| 235 |
-
|
| 236 |
-
|
| 237 |
-
|
| 238 |
-
|
| 239 |
-
|
| 240 |
-
|
| 241 |
-
|
| 242 |
-
|
| 243 |
-
|
| 244 |
-
|
| 245 |
-
|
| 246 |
-
|
| 247 |
-
|
| 248 |
-
|
| 249 |
-
|
| 250 |
-
|
| 251 |
-
|
| 252 |
-
|
| 253 |
-
|
| 254 |
-
|
| 255 |
-
|
| 256 |
-
|
| 257 |
-
|
| 258 |
-
|
| 259 |
-
|
| 260 |
-
|
| 261 |
-
|
| 262 |
-
|
| 263 |
-
|
| 264 |
-
|
| 265 |
-
|
| 266 |
-
|
| 267 |
-
|
| 268 |
-
|
| 269 |
-
|
| 270 |
-
|
| 271 |
-
|
| 272 |
-
|
| 273 |
-
|
| 274 |
-
|
| 275 |
-
|
| 276 |
-
|
| 277 |
-
|
| 278 |
-
|
| 279 |
-
|
| 280 |
-
|
| 281 |
-
|
| 282 |
-
|
| 283 |
-
|
| 284 |
-
|
| 285 |
-
|
| 286 |
-
|
| 287 |
-
|
| 288 |
-
|
| 289 |
-
|
| 290 |
-
|
| 291 |
-
|
| 292 |
-
|
| 293 |
-
|
| 294 |
-
|
| 295 |
-
|
| 296 |
-
|
| 297 |
-
|
| 298 |
-
|
| 299 |
-
|
| 300 |
-
|
| 301 |
-
|
| 302 |
-
|
| 303 |
-
|
| 304 |
-
|
| 305 |
-
|
| 306 |
-
|
| 307 |
-
|
| 308 |
-
|
| 309 |
-
|
| 310 |
-
|
| 311 |
-
|
| 312 |
-
|
| 313 |
-
|
| 314 |
-
|
| 315 |
-
|
| 316 |
-
|
| 317 |
-
and X.shape[
|
| 318 |
-
and
|
| 319 |
-
|
| 320 |
-
|
| 321 |
-
|
| 322 |
-
|
| 323 |
-
and self.
|
| 324 |
-
and
|
| 325 |
-
and X.shape[
|
| 326 |
-
and X.shape[
|
| 327 |
-
and
|
| 328 |
-
|
| 329 |
-
|
| 330 |
-
|
| 331 |
-
|
| 332 |
-
"(b,
|
| 333 |
-
"
|
| 334 |
-
|
| 335 |
-
|
| 336 |
-
|
| 337 |
-
|
| 338 |
-
|
| 339 |
-
|
| 340 |
-
|
| 341 |
-
|
| 342 |
-
|
| 343 |
-
|
| 344 |
-
|
| 345 |
-
|
| 346 |
-
|
| 347 |
-
|
| 348 |
-
|
| 349 |
-
|
| 350 |
-
|
| 351 |
-
|
| 352 |
-
|
| 353 |
-
|
| 354 |
-
|
| 355 |
-
|
| 356 |
-
|
| 357 |
-
|
| 358 |
-
|
| 359 |
-
|
| 360 |
-
|
| 361 |
-
|
| 362 |
-
|
| 363 |
-
|
| 364 |
-
|
| 365 |
-
|
| 366 |
-
|
| 367 |
-
|
| 368 |
-
|
| 369 |
-
|
| 370 |
-
|
| 371 |
-
|
| 372 |
-
|
| 373 |
-
|
| 374 |
-
|
| 375 |
-
|
| 376 |
-
|
| 377 |
-
|
| 378 |
-
|
|
|
|
|
|
| 1 |
+
"""Execute model-specific batches and return ordered residue-aware CPU tensors."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import torch
|
| 6 |
+
|
| 7 |
+
from collections.abc import Callable, Iterator, Sequence
|
| 8 |
+
from contextlib import contextmanager
|
| 9 |
+
from dataclasses import dataclass, field
|
| 10 |
+
from typing import Any
|
| 11 |
+
from torch import Tensor
|
| 12 |
+
|
| 13 |
+
from .identity import _model_device
|
| 14 |
+
from .inputs import _planned_batches
|
| 15 |
+
from .pooling import Pooler
|
| 16 |
+
from .types import EmbeddingBatch, EmbeddingInput, EmbeddingRecord
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
_MAX_PARTI_RESIDUES = 2_048
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
def _validate_parti_length(M: Tensor) -> None:
|
| 23 |
+
"""Reject an oversized attention graph before model inference."""
|
| 24 |
+
|
| 25 |
+
# M: (b, l)
|
| 26 |
+
n_residues = int(M.to(dtype=torch.int64).sum(dim=1).max().item())
|
| 27 |
+
if n_residues > _MAX_PARTI_RESIDUES:
|
| 28 |
+
raise ValueError(f"parti supports at most {_MAX_PARTI_RESIDUES:,} biological residues.")
|
| 29 |
+
|
| 30 |
+
|
| 31 |
+
def select_hidden_state_embeddings(
|
| 32 |
+
last_hidden_state: Tensor,
|
| 33 |
+
hidden_states: tuple[Tensor, ...] | None,
|
| 34 |
+
*,
|
| 35 |
+
hidden_state_index: int = -1,
|
| 36 |
+
store_all_hidden_states: bool = False,
|
| 37 |
+
) -> Tensor:
|
| 38 |
+
"""Select one hidden state or stack every state without changing values."""
|
| 39 |
+
# last_hidden_state and each hidden_states entry: (b, l, d)
|
| 40 |
+
if store_all_hidden_states:
|
| 41 |
+
if not hidden_states:
|
| 42 |
+
raise ValueError("store_all_hidden_states requires model hidden states.")
|
| 43 |
+
# H has shape (b, n, l, d), where n follows the model's output order.
|
| 44 |
+
return torch.stack(hidden_states, dim=1) # (b, n, l, d)
|
| 45 |
+
if hidden_state_index == -1:
|
| 46 |
+
return last_hidden_state # (b, l, d)
|
| 47 |
+
if not hidden_states:
|
| 48 |
+
raise ValueError("hidden_state_index requires model hidden states.")
|
| 49 |
+
return hidden_states[hidden_state_index] # (b, l, d)
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
def _residue_embeddings(X: Tensor, M: Tensor) -> list[Tensor]:
|
| 53 |
+
"""Copy every sample's biological residues to the host in one transfer.
|
| 54 |
+
|
| 55 |
+
Boolean indexing packs the selected rows in batch order, so splitting the
|
| 56 |
+
packed rows by residue count gives the values that indexing each sample
|
| 57 |
+
would. Each returned tensor owns its storage, as a per-sample copy does.
|
| 58 |
+
"""
|
| 59 |
+
# X: (b, l, d); M: (b, l)
|
| 60 |
+
residue_counts = M.sum(dim=1).tolist() # b counts r_i
|
| 61 |
+
packed = X[M].detach().cpu() # (sum of r_i, d)
|
| 62 |
+
return [sample.clone() for sample in torch.split(packed, residue_counts)] # each: (r_i, d)
|
| 63 |
+
|
| 64 |
+
|
| 65 |
+
@contextmanager
|
| 66 |
+
def _temporary_eval(model: Any) -> Iterator[None]:
|
| 67 |
+
was_training = getattr(model, "training", None)
|
| 68 |
+
eval_method = getattr(model, "eval", None)
|
| 69 |
+
train_method = getattr(model, "train", None)
|
| 70 |
+
if (
|
| 71 |
+
not isinstance(was_training, bool)
|
| 72 |
+
or not callable(eval_method)
|
| 73 |
+
or not callable(train_method)
|
| 74 |
+
):
|
| 75 |
+
yield
|
| 76 |
+
return
|
| 77 |
+
eval_method()
|
| 78 |
+
try:
|
| 79 |
+
yield
|
| 80 |
+
finally:
|
| 81 |
+
train_method(was_training)
|
| 82 |
+
|
| 83 |
+
|
| 84 |
+
def _biological_residue_mask(
|
| 85 |
+
input_ids: Tensor,
|
| 86 |
+
attention_mask: Tensor,
|
| 87 |
+
tokenizer: Any,
|
| 88 |
+
) -> Tensor:
|
| 89 |
+
"""Remove padding and tokenizer-declared special tokens from M."""
|
| 90 |
+
|
| 91 |
+
# input_ids, attention_mask: (b, l)
|
| 92 |
+
M = attention_mask.to(dtype=torch.bool) # (b, l)
|
| 93 |
+
special_ids = tuple(int(token_id) for token_id in getattr(tokenizer, "all_special_ids", ()))
|
| 94 |
+
if special_ids:
|
| 95 |
+
specials = torch.tensor( # (n_special,)
|
| 96 |
+
special_ids,
|
| 97 |
+
device=input_ids.device,
|
| 98 |
+
dtype=input_ids.dtype,
|
| 99 |
+
)
|
| 100 |
+
M = M & ~torch.isin(input_ids, specials) # (b, l)
|
| 101 |
+
return M # (b, l)
|
| 102 |
+
|
| 103 |
+
|
| 104 |
+
def _generic_embedding_batch(
|
| 105 |
+
model: Any,
|
| 106 |
+
sequences: list[str],
|
| 107 |
+
*,
|
| 108 |
+
tokenizer: Any | None,
|
| 109 |
+
max_length: int | None,
|
| 110 |
+
truncate: bool,
|
| 111 |
+
need_attentions: bool,
|
| 112 |
+
model_kwargs: dict[str, Any],
|
| 113 |
+
) -> EmbeddingBatch:
|
| 114 |
+
config = getattr(model, "config", None)
|
| 115 |
+
model_type = str(getattr(config, "model_type", "")).lower()
|
| 116 |
+
if tokenizer is None:
|
| 117 |
+
tokenizer = getattr(model, "tokenizer", None)
|
| 118 |
+
|
| 119 |
+
if tokenizer is None and model_type == "e1":
|
| 120 |
+
output = model._embed(sequences, return_attention_mask=True, **model_kwargs)
|
| 121 |
+
if not isinstance(output, tuple) or len(output) != 2:
|
| 122 |
+
raise TypeError("E1 _embed must return (X, residue_mask).")
|
| 123 |
+
X, M = output # (b, l, d), (b, l)
|
| 124 |
+
preparer = getattr(model, "prep_tokens", None)
|
| 125 |
+
if preparer is not None and hasattr(preparer, "get_batch_kwargs"):
|
| 126 |
+
prepared = preparer.get_batch_kwargs(sequences, device=X.device)
|
| 127 |
+
input_ids = prepared["input_ids"] # (b, l)
|
| 128 |
+
boundary_ids = preparer.boundary_token_ids.to( # (n_boundary,)
|
| 129 |
+
device=input_ids.device, dtype=input_ids.dtype
|
| 130 |
+
)
|
| 131 |
+
# E1 wraps each raw sequence in BOS, context-label, terminal-label,
|
| 132 |
+
# and EOS tokens. Only amino-acid rows are biological residues.
|
| 133 |
+
M = M.to(dtype=torch.bool) & ~torch.isin(input_ids, boundary_ids) # (b, l)
|
| 134 |
+
if need_attentions:
|
| 135 |
+
raise ValueError("parti is not available for tokenizer-free E1 embedding.")
|
| 136 |
+
return EmbeddingBatch( # X: (b, l, d); residue_mask: (b, l)
|
| 137 |
+
X=X,
|
| 138 |
+
residue_mask=M.to(dtype=torch.bool),
|
| 139 |
+
)
|
| 140 |
+
if tokenizer is None:
|
| 141 |
+
raise ValueError("A tokenizer is required for this model's embedding path.")
|
| 142 |
+
|
| 143 |
+
tokenize_kwargs: dict[str, Any] = {
|
| 144 |
+
"return_tensors": "pt",
|
| 145 |
+
"padding": True,
|
| 146 |
+
"truncation": truncate,
|
| 147 |
+
}
|
| 148 |
+
if max_length is not None and truncate:
|
| 149 |
+
# ``max_length`` is a biological-residue limit. Tokenizer limits include
|
| 150 |
+
# boundary tokens, so reserve their declared width instead of dropping
|
| 151 |
+
# residues at the exact boundary.
|
| 152 |
+
special_token_count = 0
|
| 153 |
+
num_special_tokens_to_add = getattr(tokenizer, "num_special_tokens_to_add", None)
|
| 154 |
+
if callable(num_special_tokens_to_add):
|
| 155 |
+
special_token_count = int(num_special_tokens_to_add(pair=False))
|
| 156 |
+
tokenize_kwargs["max_length"] = max_length + special_token_count
|
| 157 |
+
sequence_tokenizer = getattr(model, "_tokenize_sequence_batch", None)
|
| 158 |
+
if callable(sequence_tokenizer):
|
| 159 |
+
encoded = sequence_tokenizer(sequences, tokenizer=tokenizer, **tokenize_kwargs)
|
| 160 |
+
else:
|
| 161 |
+
encoded = tokenizer(sequences, **tokenize_kwargs)
|
| 162 |
+
device = _model_device(model)
|
| 163 |
+
input_ids = encoded["input_ids"].to(device) # (b, l)
|
| 164 |
+
attention_mask = encoded.get( # (b, l)
|
| 165 |
+
"attention_mask",
|
| 166 |
+
input_ids.new_ones(input_ids.shape),
|
| 167 |
+
).to(device)
|
| 168 |
+
M = _biological_residue_mask(input_ids, attention_mask, tokenizer) # (b, l)
|
| 169 |
+
if need_attentions:
|
| 170 |
+
# Validate l before either the backbone or its quadratic attention graph
|
| 171 |
+
# is materialized. M has shape (b, l).
|
| 172 |
+
_validate_parti_length(M)
|
| 173 |
+
X = model._embed(input_ids, attention_mask, **model_kwargs) # (b, l, d)
|
| 174 |
+
attentions = None
|
| 175 |
+
if need_attentions:
|
| 176 |
+
output = model(
|
| 177 |
+
input_ids=input_ids,
|
| 178 |
+
attention_mask=attention_mask,
|
| 179 |
+
output_attentions=True,
|
| 180 |
+
return_dict=True,
|
| 181 |
+
)
|
| 182 |
+
attentions = getattr(output, "attentions", None) # each: (b, h, l, l)
|
| 183 |
+
if attentions is None:
|
| 184 |
+
raise ValueError("The model did not return attentions required by parti.")
|
| 185 |
+
return EmbeddingBatch( # X: (b, l, d); M: (b, l)
|
| 186 |
+
X=X,
|
| 187 |
+
residue_mask=M,
|
| 188 |
+
attentions=attentions,
|
| 189 |
+
)
|
| 190 |
+
|
| 191 |
+
|
| 192 |
+
@dataclass(eq=False)
|
| 193 |
+
class BatchExecutor:
|
| 194 |
+
"""Model and batch policy for one bounded embedding window at a time."""
|
| 195 |
+
|
| 196 |
+
model: Any
|
| 197 |
+
batch_size: int
|
| 198 |
+
max_tokens_per_batch: int | None
|
| 199 |
+
max_length: int | None
|
| 200 |
+
truncate: bool
|
| 201 |
+
model_kwargs: dict[str, Any]
|
| 202 |
+
hidden_state_source: str
|
| 203 |
+
normalized_decoder_inputs: tuple[str, ...] | None
|
| 204 |
+
decoder_input_ids: Tensor | None
|
| 205 |
+
decoder_attention_mask: Tensor | None
|
| 206 |
+
_embedding_batch_fn: Callable[..., EmbeddingBatch] | None
|
| 207 |
+
tokenizer: Any | None
|
| 208 |
+
store_all_hidden_states: bool
|
| 209 |
+
full_embeddings: bool
|
| 210 |
+
dtype: torch.dtype | None
|
| 211 |
+
pooler: Pooler | None
|
| 212 |
+
attention_backend: str | None
|
| 213 |
+
need_attentions: bool
|
| 214 |
+
model_type: str = field(init=False)
|
| 215 |
+
resolved_tokenizer: Any = field(init=False)
|
| 216 |
+
|
| 217 |
+
def __post_init__(self) -> None:
|
| 218 |
+
config = getattr(self.model, "config", None)
|
| 219 |
+
self.model_type = str(getattr(config, "model_type", "")).lower()
|
| 220 |
+
self.resolved_tokenizer = (
|
| 221 |
+
self.tokenizer if self.tokenizer is not None else getattr(self.model, "tokenizer", None)
|
| 222 |
+
)
|
| 223 |
+
|
| 224 |
+
def run_window(
|
| 225 |
+
self,
|
| 226 |
+
window_records: Sequence[EmbeddingInput],
|
| 227 |
+
*,
|
| 228 |
+
window_start: int,
|
| 229 |
+
) -> tuple[list[EmbeddingRecord], dict[str, tuple[int, int]]]:
|
| 230 |
+
"""Restore source order after length-bucketed inference and pooling."""
|
| 231 |
+
|
| 232 |
+
pool_slices: dict[str, tuple[int, int]] = {}
|
| 233 |
+
window_results: dict[int, EmbeddingRecord] = {}
|
| 234 |
+
for local_positions in _planned_batches(
|
| 235 |
+
window_records,
|
| 236 |
+
range(len(window_records)),
|
| 237 |
+
batch_size=self.batch_size,
|
| 238 |
+
max_tokens_per_batch=self.max_tokens_per_batch,
|
| 239 |
+
max_length=self.max_length,
|
| 240 |
+
truncate=self.truncate,
|
| 241 |
+
):
|
| 242 |
+
batch_positions = [window_start + position for position in local_positions]
|
| 243 |
+
batch_records = [window_records[position] for position in local_positions]
|
| 244 |
+
sequences = [
|
| 245 |
+
record.sequence[: self.max_length]
|
| 246 |
+
if self.truncate and self.max_length is not None
|
| 247 |
+
else record.sequence
|
| 248 |
+
for record in batch_records
|
| 249 |
+
]
|
| 250 |
+
batch_model_kwargs = dict(self.model_kwargs)
|
| 251 |
+
if self.model_type == "fast_ankh" or self.hidden_state_source == "decoder":
|
| 252 |
+
batch_model_kwargs["hidden_state_source"] = self.hidden_state_source
|
| 253 |
+
if self.normalized_decoder_inputs is not None:
|
| 254 |
+
batch_model_kwargs["decoder_inputs"] = [
|
| 255 |
+
self.normalized_decoder_inputs[position] for position in batch_positions
|
| 256 |
+
]
|
| 257 |
+
if self.decoder_input_ids is not None:
|
| 258 |
+
# decoder_input_ids: (n_records, l_decoder)
|
| 259 |
+
indices = torch.tensor( # (b,)
|
| 260 |
+
batch_positions,
|
| 261 |
+
device=self.decoder_input_ids.device,
|
| 262 |
+
dtype=torch.long,
|
| 263 |
+
)
|
| 264 |
+
batch_model_kwargs["decoder_input_ids"] = ( # (b, l_decoder)
|
| 265 |
+
self.decoder_input_ids.index_select(0, indices)
|
| 266 |
+
)
|
| 267 |
+
if self.decoder_attention_mask is not None:
|
| 268 |
+
# decoder_attention_mask: (n_records, l_decoder)
|
| 269 |
+
indices = torch.tensor( # (b,)
|
| 270 |
+
batch_positions,
|
| 271 |
+
device=self.decoder_attention_mask.device,
|
| 272 |
+
dtype=torch.long,
|
| 273 |
+
)
|
| 274 |
+
batch_model_kwargs["decoder_attention_mask"] = (
|
| 275 |
+
self.decoder_attention_mask.index_select(0, indices) # (b, l_decoder)
|
| 276 |
+
)
|
| 277 |
+
custom_batch = self._embedding_batch_fn or getattr(self.model, "_embedding_batch", None)
|
| 278 |
+
if custom_batch is not None:
|
| 279 |
+
if self.model_type == "fast_ankh":
|
| 280 |
+
batch = custom_batch(
|
| 281 |
+
sequences,
|
| 282 |
+
tokenizer=self.resolved_tokenizer,
|
| 283 |
+
max_length=self.max_length,
|
| 284 |
+
truncate=self.truncate,
|
| 285 |
+
need_attentions=self.need_attentions,
|
| 286 |
+
**batch_model_kwargs,
|
| 287 |
+
)
|
| 288 |
+
else:
|
| 289 |
+
batch = custom_batch(sequences, **batch_model_kwargs)
|
| 290 |
+
if not isinstance(batch, EmbeddingBatch):
|
| 291 |
+
raise TypeError("_embedding_batch must return EmbeddingBatch.")
|
| 292 |
+
else:
|
| 293 |
+
batch = _generic_embedding_batch(
|
| 294 |
+
self.model,
|
| 295 |
+
sequences,
|
| 296 |
+
tokenizer=self.tokenizer,
|
| 297 |
+
max_length=self.max_length,
|
| 298 |
+
truncate=self.truncate,
|
| 299 |
+
need_attentions=self.need_attentions,
|
| 300 |
+
model_kwargs=batch_model_kwargs,
|
| 301 |
+
)
|
| 302 |
+
X = batch.X # (b, l, d) or (b, n_states, l, d)
|
| 303 |
+
raw_mask = batch.residue_mask # (b, l)
|
| 304 |
+
if not isinstance(X, Tensor) or not isinstance(raw_mask, Tensor):
|
| 305 |
+
raise TypeError("Embedding batches must provide Tensor X and residue_mask.")
|
| 306 |
+
if X.is_meta or raw_mask.is_meta:
|
| 307 |
+
raise ValueError("Embedding batches cannot contain meta tensors.")
|
| 308 |
+
if not X.is_floating_point():
|
| 309 |
+
raise TypeError("Embedding batches must use a floating-point X dtype.")
|
| 310 |
+
if raw_mask.is_complex() or not bool(torch.isfinite(raw_mask).all()):
|
| 311 |
+
raise ValueError("Embedding residue_mask must contain finite binary values.")
|
| 312 |
+
if not bool(((raw_mask == 0) | (raw_mask == 1)).all()):
|
| 313 |
+
raise ValueError("Embedding residue_mask must contain finite binary values.")
|
| 314 |
+
M = raw_mask.to(device=X.device, dtype=torch.bool) # (b, l)
|
| 315 |
+
valid_X_shape = (
|
| 316 |
+
X.ndim == 3
|
| 317 |
+
and X.shape[0] == len(batch_records)
|
| 318 |
+
and X.shape[-1] > 0
|
| 319 |
+
and M.shape == X.shape[:2]
|
| 320 |
+
)
|
| 321 |
+
valid_all_states_shape = (
|
| 322 |
+
X.ndim == 4
|
| 323 |
+
and self.store_all_hidden_states
|
| 324 |
+
and self.full_embeddings
|
| 325 |
+
and X.shape[0] == len(batch_records)
|
| 326 |
+
and X.shape[1] > 0
|
| 327 |
+
and X.shape[-1] > 0
|
| 328 |
+
and M.shape == (X.shape[0], X.shape[2])
|
| 329 |
+
)
|
| 330 |
+
if not (valid_X_shape or valid_all_states_shape):
|
| 331 |
+
raise ValueError(
|
| 332 |
+
"Embedding batches must provide X with shape (b, l, d), or "
|
| 333 |
+
"(b, states, l, d) when storing all hidden states, and "
|
| 334 |
+
"residue_mask with shape (b, l)."
|
| 335 |
+
)
|
| 336 |
+
if not bool(M.any(dim=1).all()):
|
| 337 |
+
raise ValueError("Every embedding sample must contain a biological residue.")
|
| 338 |
+
finite_selected = ( # X.shape
|
| 339 |
+
torch.isfinite(X) | ~M.unsqueeze(-1)
|
| 340 |
+
if X.ndim == 3
|
| 341 |
+
else torch.isfinite(X) | ~M[:, None, :, None]
|
| 342 |
+
)
|
| 343 |
+
if not bool(finite_selected.all()):
|
| 344 |
+
raise ValueError("Biological residue embeddings produced non-finite output.")
|
| 345 |
+
if self.need_attentions:
|
| 346 |
+
# Validate the biological graph only after mask integrity is established.
|
| 347 |
+
_validate_parti_length(M)
|
| 348 |
+
if self.dtype is not None:
|
| 349 |
+
X = X.to(dtype=self.dtype) # unchanged shape
|
| 350 |
+
|
| 351 |
+
if self.full_embeddings:
|
| 352 |
+
if X.ndim == 4:
|
| 353 |
+
values = [
|
| 354 |
+
X_i[:, M_i, :].detach().cpu() # (n_states, r_i, d)
|
| 355 |
+
for X_i, M_i in zip(X, M, strict=True)
|
| 356 |
+
]
|
| 357 |
+
else:
|
| 358 |
+
values = _residue_embeddings(X, M) # each: (r_i, d)
|
| 359 |
+
else:
|
| 360 |
+
if self.pooler is None:
|
| 361 |
+
raise RuntimeError(
|
| 362 |
+
"Pooled embedding output was requested without an initialized pooler."
|
| 363 |
+
)
|
| 364 |
+
Y = self.pooler( # (b, n_poolers * d)
|
| 365 |
+
X,
|
| 366 |
+
M,
|
| 367 |
+
attentions=batch.attentions,
|
| 368 |
+
attention_backend=self.attention_backend,
|
| 369 |
+
)
|
| 370 |
+
pool_slices = self.pooler.output_slices(X.shape[-1])
|
| 371 |
+
values = list(Y.detach().cpu().unbind(0)) # each: (n_poolers * d,)
|
| 372 |
+
for position, record, value in zip(batch_positions, batch_records, values, strict=True):
|
| 373 |
+
window_results[position] = EmbeddingRecord(record.id, record.sequence, value)
|
| 374 |
+
|
| 375 |
+
new_records = [
|
| 376 |
+
window_results[position]
|
| 377 |
+
for position in range(window_start, window_start + len(window_records))
|
| 378 |
+
]
|
| 379 |
+
return new_records, pool_slices
|
fastplms/embeddings/identity.py
CHANGED
|
@@ -1,510 +1,511 @@
|
|
| 1 |
-
"""Deterministic identity for embedding inputs, models, tokenizers, and execution."""
|
| 2 |
-
|
| 3 |
-
from __future__ import annotations
|
| 4 |
-
|
| 5 |
-
import hashlib
|
| 6 |
-
import json
|
| 7 |
-
import platform
|
| 8 |
-
import torch
|
| 9 |
-
|
| 10 |
-
from
|
| 11 |
-
from
|
| 12 |
-
from
|
| 13 |
-
|
|
|
|
| 14 |
from .inputs import _InputSpool
|
| 15 |
from .storage import tensor_sha256
|
| 16 |
from .types import EmbeddingInput
|
| 17 |
-
|
| 18 |
-
|
| 19 |
-
_RUN_FINGERPRINT_SCHEMA_VERSION = 3
|
| 20 |
-
_MODEL_STATE_HASH_CHUNK_BYTES = 16 * 1024**2
|
| 21 |
-
|
| 22 |
-
|
| 23 |
-
def _model_device(model: Any) -> torch.device:
|
| 24 |
-
try:
|
| 25 |
-
return torch.device(next(model.parameters()).device)
|
| 26 |
-
except (AttributeError, StopIteration):
|
| 27 |
-
return torch.device("cpu")
|
| 28 |
-
|
| 29 |
-
|
| 30 |
-
def _attention_backend(model: Any) -> str | None:
|
| 31 |
-
config = getattr(model, "config", None)
|
| 32 |
-
for name in ("_attn_implementation", "attn_implementation", "attn_backend"):
|
| 33 |
-
value = getattr(config, name, None)
|
| 34 |
-
if value:
|
| 35 |
-
return str(value)
|
| 36 |
-
return None
|
| 37 |
-
|
| 38 |
-
|
| 39 |
-
def _attention_kernel_metadata(backend: str | None) -> dict[str, Any] | None:
|
| 40 |
-
if backend not in {"flash_attention_2", "flash_attention_3"}:
|
| 41 |
-
return None
|
| 42 |
-
from fastplms.registry import get_model_registry
|
| 43 |
-
|
| 44 |
-
spec = get_model_registry().attention_kernels[backend]
|
| 45 |
-
return {
|
| 46 |
-
"repository": spec.repository,
|
| 47 |
-
"revision": spec.revision,
|
| 48 |
-
"version": spec.version,
|
| 49 |
-
"expected_variant": spec.expected_variant,
|
| 50 |
-
"dtypes": list(spec.dtypes),
|
| 51 |
-
}
|
| 52 |
-
|
| 53 |
-
|
| 54 |
-
def _fingerprint_jsonable(value: Any) -> Any:
|
| 55 |
-
if isinstance(value, Mapping):
|
| 56 |
-
return {str(key): _fingerprint_jsonable(item) for key, item in value.items()}
|
| 57 |
-
if isinstance(value, (list, tuple)):
|
| 58 |
-
return [_fingerprint_jsonable(item) for item in value]
|
| 59 |
-
if isinstance(value, (set, frozenset)):
|
| 60 |
-
return sorted((_fingerprint_jsonable(item) for item in value), key=repr)
|
| 61 |
-
if isinstance(value, Path):
|
| 62 |
-
return str(value)
|
| 63 |
-
if isinstance(value, Tensor):
|
| 64 |
-
return {
|
| 65 |
-
"dtype": str(value.dtype).removeprefix("torch."),
|
| 66 |
-
"shape": list(value.shape),
|
| 67 |
-
"sha256": tensor_sha256(value),
|
| 68 |
-
}
|
| 69 |
-
if isinstance(value, torch.dtype):
|
| 70 |
-
return str(value).removeprefix("torch.")
|
| 71 |
-
if isinstance(value, torch.device):
|
| 72 |
-
return str(value)
|
| 73 |
-
if value is None or isinstance(value, (str, int, float, bool)):
|
| 74 |
-
return value
|
| 75 |
-
return {
|
| 76 |
-
"class": f"{value.__class__.__module__}.{value.__class__.__qualname__}",
|
| 77 |
-
"value": str(value),
|
| 78 |
-
}
|
| 79 |
-
|
| 80 |
-
|
| 81 |
-
def _tokenizer_content_sha256(tokenizer: Any) -> str:
|
| 82 |
-
content: dict[str, Any] = {
|
| 83 |
-
"init_kwargs": getattr(tokenizer, "init_kwargs", None),
|
| 84 |
-
"special_tokens_map": getattr(tokenizer, "special_tokens_map", None),
|
| 85 |
-
"model_max_length": getattr(tokenizer, "model_max_length", None),
|
| 86 |
-
"padding_side": getattr(tokenizer, "padding_side", None),
|
| 87 |
-
"truncation_side": getattr(tokenizer, "truncation_side", None),
|
| 88 |
-
}
|
| 89 |
-
get_vocab = getattr(tokenizer, "get_vocab", None)
|
| 90 |
-
if callable(get_vocab):
|
| 91 |
-
content["vocabulary"] = get_vocab()
|
| 92 |
-
get_added_vocab = getattr(tokenizer, "get_added_vocab", None)
|
| 93 |
-
if callable(get_added_vocab):
|
| 94 |
-
content["added_vocabulary"] = get_added_vocab()
|
| 95 |
-
backend = getattr(tokenizer, "backend_tokenizer", None)
|
| 96 |
-
backend_to_str = getattr(backend, "to_str", None)
|
| 97 |
-
if callable(backend_to_str):
|
| 98 |
-
content["backend"] = backend_to_str()
|
| 99 |
-
serialized = json.dumps(
|
| 100 |
-
_fingerprint_jsonable(content),
|
| 101 |
-
sort_keys=True,
|
| 102 |
-
separators=(",", ":"),
|
| 103 |
-
ensure_ascii=False,
|
| 104 |
-
).encode()
|
| 105 |
-
return hashlib.sha256(serialized).hexdigest()
|
| 106 |
-
|
| 107 |
-
|
| 108 |
-
def _tokenizer_metadata(model: Any, tokenizer: Any | None) -> dict[str, Any]:
|
| 109 |
-
resolved = tokenizer if tokenizer is not None else getattr(model, "tokenizer", None)
|
| 110 |
-
if resolved is None:
|
| 111 |
-
# Raw-sequence families such as E1 retain their loader context on the
|
| 112 |
-
# model/encoder rather than exposing a Transformers tokenizer. Bind the
|
| 113 |
-
# non-secret source policy to resume identity without serializing a Hub
|
| 114 |
-
# token or forcing lazy tokenizer initialization.
|
| 115 |
-
for candidate in (model, getattr(model, "model", None)):
|
| 116 |
-
settings = getattr(candidate, "__dict__", {}).get("_fastplms_tokenizer_kwargs")
|
| 117 |
-
if isinstance(settings, Mapping):
|
| 118 |
-
token_value = settings.get("token")
|
| 119 |
-
return {
|
| 120 |
-
"mode": "native-sequence",
|
| 121 |
-
"source": (
|
| 122 |
-
str(settings.get("tokenizer_source"))
|
| 123 |
-
if settings.get("tokenizer_source") is not None
|
| 124 |
-
else None
|
| 125 |
-
),
|
| 126 |
-
"revision": settings.get("revision"),
|
| 127 |
-
"cache_dir": (
|
| 128 |
-
str(settings.get("cache_dir"))
|
| 129 |
-
if settings.get("cache_dir") is not None
|
| 130 |
-
else None
|
| 131 |
-
),
|
| 132 |
-
"local_files_only": bool(settings.get("local_files_only", False)),
|
| 133 |
-
"token_policy": (
|
| 134 |
-
"disabled"
|
| 135 |
-
if token_value is False
|
| 136 |
-
else "provided"
|
| 137 |
-
if token_value is not None
|
| 138 |
-
else "default"
|
| 139 |
-
),
|
| 140 |
-
}
|
| 141 |
-
return {"mode": "native-sequence"}
|
| 142 |
-
return {
|
| 143 |
-
"mode": "tokenizer",
|
| 144 |
-
"class": f"{resolved.__class__.__module__}.{resolved.__class__.__qualname__}",
|
| 145 |
-
"name_or_path": getattr(resolved, "name_or_path", None),
|
| 146 |
-
"vocab_size": getattr(resolved, "vocab_size", None),
|
| 147 |
-
"special_token_ids": list(getattr(resolved, "all_special_ids", ())),
|
| 148 |
-
"content_sha256": _tokenizer_content_sha256(resolved),
|
| 149 |
-
}
|
| 150 |
-
|
| 151 |
-
|
| 152 |
-
def _software_versions() -> dict[str, str | None]:
|
| 153 |
-
try:
|
| 154 |
-
import fastplms
|
| 155 |
-
|
| 156 |
-
fastplms_version = fastplms.__version__
|
| 157 |
-
except (AttributeError, ImportError):
|
| 158 |
-
fastplms_version = None
|
| 159 |
-
try:
|
| 160 |
-
import safetensors
|
| 161 |
-
|
| 162 |
-
safetensors_version = safetensors.__version__
|
| 163 |
-
except ImportError:
|
| 164 |
-
safetensors_version = None
|
| 165 |
-
try:
|
| 166 |
-
import transformers
|
| 167 |
-
|
| 168 |
-
transformers_version = transformers.__version__
|
| 169 |
-
except ImportError:
|
| 170 |
-
transformers_version = None
|
| 171 |
-
return {
|
| 172 |
-
"fastplms": fastplms_version,
|
| 173 |
-
"python": platform.python_version(),
|
| 174 |
-
"safetensors": safetensors_version,
|
| 175 |
-
"torch": torch.__version__,
|
| 176 |
-
"torch_cuda": torch.version.cuda,
|
| 177 |
-
"transformers": transformers_version,
|
| 178 |
-
}
|
| 179 |
-
|
| 180 |
-
|
| 181 |
-
def _adapter_identity_metadata(model: Any) -> dict[str, Any] | None:
|
| 182 |
-
"""Return deterministic PEFT/adapter identity without tensor payloads."""
|
| 183 |
-
|
| 184 |
-
peft_config = getattr(model, "peft_config", None)
|
| 185 |
-
if not isinstance(peft_config, Mapping) or not peft_config:
|
| 186 |
-
return None
|
| 187 |
-
configurations: dict[str, Any] = {}
|
| 188 |
-
for name, config in sorted(peft_config.items(), key=lambda item: str(item[0])):
|
| 189 |
-
to_dict = getattr(config, "to_dict", None)
|
| 190 |
-
if callable(to_dict):
|
| 191 |
-
value = to_dict()
|
| 192 |
-
else:
|
| 193 |
-
try:
|
| 194 |
-
value = vars(config)
|
| 195 |
-
except TypeError:
|
| 196 |
-
value = config
|
| 197 |
-
configurations[str(name)] = _fingerprint_jsonable(value)
|
| 198 |
-
active_adapters = getattr(model, "active_adapters", None)
|
| 199 |
-
if callable(active_adapters):
|
| 200 |
-
active_adapters = active_adapters()
|
| 201 |
-
return {
|
| 202 |
-
"active": _fingerprint_jsonable(active_adapters),
|
| 203 |
-
"configurations": configurations,
|
| 204 |
-
}
|
| 205 |
-
|
| 206 |
-
|
| 207 |
-
def _execution_identity_metadata(model: Any) -> dict[str, Any]:
|
| 208 |
-
"""Capture runtime policy that can change persisted numerical results."""
|
| 209 |
-
|
| 210 |
-
parameter_dtypes = sorted(
|
| 211 |
-
{
|
| 212 |
-
str(parameter.dtype).removeprefix("torch.")
|
| 213 |
-
for parameter in getattr(model, "parameters", lambda: ())()
|
| 214 |
-
}
|
| 215 |
-
)
|
| 216 |
-
return {
|
| 217 |
-
"device": _model_device(model).type,
|
| 218 |
-
"hf_device_map": _fingerprint_jsonable(getattr(model, "hf_device_map", None)),
|
| 219 |
-
"parameter_dtypes": parameter_dtypes,
|
| 220 |
-
"software": _software_versions(),
|
| 221 |
-
}
|
| 222 |
-
|
| 223 |
-
|
| 224 |
-
def _first_metadata_value(*values: Any) -> Any:
|
| 225 |
-
for value in values:
|
| 226 |
-
if isinstance(value, str):
|
| 227 |
-
if value.strip():
|
| 228 |
-
return value
|
| 229 |
-
elif value is not None:
|
| 230 |
-
return value
|
| 231 |
-
return None
|
| 232 |
-
|
| 233 |
-
|
| 234 |
-
def _model_identity_metadata(model: Any) -> dict[str, Any]:
|
| 235 |
-
"""Resolve model and checkpoint identity, including local artifact fallbacks."""
|
| 236 |
-
|
| 237 |
-
config = getattr(model, "config", None)
|
| 238 |
-
checkpoint_revision = _first_metadata_value(
|
| 239 |
-
getattr(config, "fastplms_checkpoint_revision", None),
|
| 240 |
-
getattr(config, "_commit_hash", None),
|
| 241 |
-
)
|
| 242 |
-
return {
|
| 243 |
-
"model_id": _first_metadata_value(
|
| 244 |
-
getattr(config, "fastplms_model_id", None),
|
| 245 |
-
getattr(config, "_name_or_path", None),
|
| 246 |
-
),
|
| 247 |
-
"model_revision": _first_metadata_value(
|
| 248 |
-
getattr(config, "_commit_hash", None),
|
| 249 |
-
checkpoint_revision,
|
| 250 |
-
),
|
| 251 |
-
"checkpoint_repo_id": getattr(config, "fastplms_checkpoint_repo_id", None),
|
| 252 |
-
"checkpoint_revision": checkpoint_revision,
|
| 253 |
-
"checkpoint_hash": _first_metadata_value(
|
| 254 |
-
getattr(model, "checkpoint_hash", None),
|
| 255 |
-
getattr(config, "checkpoint_hash", None),
|
| 256 |
-
getattr(config, "fastplms_checkpoint_hash", None),
|
| 257 |
-
),
|
| 258 |
-
"weights_revision": getattr(config, "fastplms_weights_revision", None),
|
| 259 |
-
"runtime_revision": getattr(config, "fastplms_runtime_revision", None),
|
| 260 |
-
"source_tree_sha256": getattr(config, "fastplms_source_tree_sha256", None),
|
| 261 |
-
"runtime_bundle_sha256": getattr(config, "fastplms_runtime_bundle_sha256", None),
|
| 262 |
-
}
|
| 263 |
-
|
| 264 |
-
|
| 265 |
-
def _bounded_tensor_chunks(X: Tensor, max_elements: int) -> Iterable[Tensor]:
|
| 266 |
-
"""Yield X in logical row-major order without materializing a full copy."""
|
| 267 |
-
|
| 268 |
-
# X: (...)
|
| 269 |
-
if X.numel() == 0:
|
| 270 |
-
return
|
| 271 |
-
if X.ndim == 0:
|
| 272 |
-
yield X
|
| 273 |
-
return
|
| 274 |
-
trailing_elements = 1
|
| 275 |
-
for size in X.shape[1:]:
|
| 276 |
-
trailing_elements *= int(size)
|
| 277 |
-
if trailing_elements <= max_elements:
|
| 278 |
-
rows_per_chunk = max(1, max_elements // trailing_elements)
|
| 279 |
-
for start in range(0, X.shape[0], rows_per_chunk):
|
| 280 |
-
yield X[start : start + rows_per_chunk] # (chunk_rows, ...)
|
| 281 |
-
return
|
| 282 |
-
for row in X:
|
| 283 |
-
yield from _bounded_tensor_chunks(row, max_elements)
|
| 284 |
-
|
| 285 |
-
|
| 286 |
-
def _model_state_sha256(model: Any) -> str:
|
| 287 |
-
"""Hash named parameters and persistent buffers using bounded CPU copies."""
|
| 288 |
-
|
| 289 |
-
# Never cache this digest from tensor identity or ``Tensor._version``.
|
| 290 |
-
# ``Parameter.data`` and independent tensor aliases can mutate shared storage
|
| 291 |
-
# without changing either signal, while persisted resume identity must bind
|
| 292 |
-
# the authoritative bytes visible at the start of this run.
|
| 293 |
-
state = model.state_dict(keep_vars=True)
|
| 294 |
-
digest = hashlib.sha256()
|
| 295 |
-
for name, value in sorted(state.items()):
|
| 296 |
-
if not isinstance(value, Tensor):
|
| 297 |
-
raise TypeError(f"Model state entry {name!r} is not a tensor.")
|
| 298 |
-
if value.is_meta:
|
| 299 |
-
raise ValueError(
|
| 300 |
-
f"Cannot fingerprint meta-device model state entry {name!r}; pass "
|
| 301 |
-
"model_state_fingerprint with a caller-owned state identity."
|
| 302 |
-
)
|
| 303 |
-
header = json.dumps(
|
| 304 |
-
{
|
| 305 |
-
"name": name,
|
| 306 |
-
"dtype": str(value.dtype).removeprefix("torch."),
|
| 307 |
-
"shape": list(value.shape),
|
| 308 |
-
},
|
| 309 |
-
sort_keys=True,
|
| 310 |
-
separators=(",", ":"),
|
| 311 |
-
).encode()
|
| 312 |
-
digest.update(len(header).to_bytes(8, "big"))
|
| 313 |
-
digest.update(header)
|
| 314 |
-
max_elements = max(1, _MODEL_STATE_HASH_CHUNK_BYTES // value.element_size())
|
| 315 |
-
for chunk in _bounded_tensor_chunks(value.detach(), max_elements):
|
| 316 |
-
cpu_chunk = chunk.to(device="cpu").contiguous() # chunk.shape
|
| 317 |
-
digest.update(cpu_chunk.reshape(-1).view(torch.uint8).numpy().tobytes())
|
| 318 |
-
return digest.hexdigest()
|
| 319 |
-
|
| 320 |
-
|
| 321 |
-
def _input_sha256(records: Iterable[EmbeddingInput]) -> str:
|
| 322 |
-
"""Hash an ordered input stream without constructing a duplicate JSON payload."""
|
| 323 |
-
|
| 324 |
-
precomputed = getattr(records, "input_fingerprint", None)
|
| 325 |
-
if isinstance(precomputed, str):
|
| 326 |
-
return precomputed
|
| 327 |
-
digest = hashlib.sha256()
|
| 328 |
-
count = 0
|
| 329 |
-
for record in records:
|
| 330 |
-
count += 1
|
| 331 |
-
for value in (record.id, record.sequence):
|
| 332 |
-
encoded = value.encode("utf-8")
|
| 333 |
-
digest.update(len(encoded).to_bytes(8, "big"))
|
| 334 |
-
digest.update(encoded)
|
| 335 |
-
digest.update(count.to_bytes(8, "big"))
|
| 336 |
-
return digest.hexdigest()
|
| 337 |
-
|
| 338 |
-
|
| 339 |
-
def _run_fingerprint(
|
| 340 |
-
model: Any,
|
| 341 |
-
records: Sequence[EmbeddingInput],
|
| 342 |
-
*,
|
| 343 |
-
pooling: Sequence[str],
|
| 344 |
-
full_embeddings: bool,
|
| 345 |
-
max_length: int | None,
|
| 346 |
-
truncate: bool,
|
| 347 |
-
dtype: torch.dtype | None,
|
| 348 |
-
model_kwargs: dict[str, Any],
|
| 349 |
-
tokenizer_metadata: dict[str, Any],
|
| 350 |
-
model_state_fingerprint: str | None,
|
| 351 |
-
persist_output: bool,
|
| 352 |
-
embedding_context: Mapping[str, Any],
|
| 353 |
-
batch_size: int,
|
| 354 |
-
batch_window_size: int,
|
| 355 |
-
max_tokens_per_batch: int | None,
|
| 356 |
-
) -> tuple[str, str, str | None, str]:
|
| 357 |
-
input_fingerprint = _input_sha256(records)
|
| 358 |
-
attention_backend = _attention_backend(model)
|
| 359 |
-
model_identity = _model_identity_metadata(model)
|
| 360 |
-
if model_state_fingerprint is None and persist_output:
|
| 361 |
-
resolved_model_state_fingerprint = _model_state_sha256(model)
|
| 362 |
-
model_state_fingerprint_source = "computed"
|
| 363 |
-
elif model_state_fingerprint is not None:
|
| 364 |
-
resolved_model_state_fingerprint = model_state_fingerprint.strip()
|
| 365 |
-
if not resolved_model_state_fingerprint:
|
| 366 |
-
raise ValueError("model_state_fingerprint must not be empty.")
|
| 367 |
-
model_state_fingerprint_source = "caller"
|
| 368 |
-
else:
|
| 369 |
-
resolved_model_state_fingerprint = None
|
| 370 |
-
model_state_fingerprint_source = "not-computed"
|
| 371 |
-
payload = {
|
| 372 |
-
"fingerprint_schema_version": _RUN_FINGERPRINT_SCHEMA_VERSION,
|
| 373 |
-
"input_fingerprint": input_fingerprint,
|
| 374 |
-
"model_state_fingerprint": resolved_model_state_fingerprint,
|
| 375 |
-
"model_state_fingerprint_source": model_state_fingerprint_source,
|
| 376 |
-
"model_class": f"{model.__class__.__module__}.{model.__class__.__qualname__}",
|
| 377 |
-
**model_identity,
|
| 378 |
-
"attention_backend": attention_backend,
|
| 379 |
-
"attention_kernel": _attention_kernel_metadata(attention_backend),
|
| 380 |
-
"layer": repr(
|
| 381 |
-
getattr(model, "embedding_layer", model_kwargs.get("hidden_state_index", -1))
|
| 382 |
-
),
|
| 383 |
-
"projection": getattr(model, "embedding_projection", None),
|
| 384 |
-
"esmc_source": getattr(model, "_esmc_source", None),
|
| 385 |
-
"esmc_revision": getattr(model, "_esmc_source_revision", None),
|
| 386 |
-
"esmc_files": getattr(model, "_esmc_source_files", None),
|
| 387 |
-
"token_policy": getattr(model, "embedding_token_policy", None),
|
| 388 |
-
"tokenizer": tokenizer_metadata,
|
| 389 |
-
"adapter": _adapter_identity_metadata(model),
|
| 390 |
-
"execution": _execution_identity_metadata(model),
|
| 391 |
-
"embedding_context": _fingerprint_jsonable(embedding_context),
|
| 392 |
-
"pooling": list(pooling),
|
| 393 |
-
"full_embeddings": full_embeddings,
|
| 394 |
-
"max_length": max_length,
|
| 395 |
-
"truncate": truncate,
|
| 396 |
-
"dtype": str(dtype) if dtype is not None else None,
|
| 397 |
-
"batching": {
|
| 398 |
-
"batch_size": batch_size,
|
| 399 |
-
"batch_window_size": batch_window_size,
|
| 400 |
-
"max_tokens_per_batch": max_tokens_per_batch,
|
| 401 |
-
"input_storage": ("disk-spool" if isinstance(records, _InputSpool) else "memory"),
|
| 402 |
-
},
|
| 403 |
-
"model_kwargs": {
|
| 404 |
-
key: _fingerprint_jsonable(value) for key, value in sorted(model_kwargs.items())
|
| 405 |
-
},
|
| 406 |
-
"residue_mask_policy": "attention-mask-minus-special-tokens",
|
| 407 |
-
}
|
| 408 |
-
run_fingerprint = hashlib.sha256(
|
| 409 |
-
json.dumps(payload, sort_keys=True, separators=(",", ":")).encode()
|
| 410 |
-
).hexdigest()
|
| 411 |
-
return (
|
| 412 |
-
input_fingerprint,
|
| 413 |
-
run_fingerprint,
|
| 414 |
-
resolved_model_state_fingerprint,
|
| 415 |
-
model_state_fingerprint_source,
|
| 416 |
-
)
|
| 417 |
-
|
| 418 |
-
|
| 419 |
-
def _ordered_string_sha256(values: Sequence[str]) -> str:
|
| 420 |
-
digest = hashlib.sha256()
|
| 421 |
-
for value in values:
|
| 422 |
-
encoded = value.encode("utf-8")
|
| 423 |
-
digest.update(len(encoded).to_bytes(8, "big"))
|
| 424 |
-
digest.update(encoded)
|
| 425 |
-
digest.update(len(values).to_bytes(8, "big"))
|
| 426 |
-
return digest.hexdigest()
|
| 427 |
-
|
| 428 |
-
|
| 429 |
-
def _embedding_context(
|
| 430 |
-
model: Any,
|
| 431 |
-
records: Sequence[EmbeddingInput],
|
| 432 |
-
*,
|
| 433 |
-
hidden_state_source: str,
|
| 434 |
-
decoder_inputs: Sequence[str] | None,
|
| 435 |
-
decoder_input_ids: Tensor | None,
|
| 436 |
-
decoder_attention_mask: Tensor | None,
|
| 437 |
-
model_kwargs: Mapping[str, Any],
|
| 438 |
-
) -> tuple[dict[str, Any], tuple[str, ...] | None]:
|
| 439 |
-
if hidden_state_source not in {"encoder", "decoder"}:
|
| 440 |
-
raise ValueError("hidden_state_source must be 'encoder' or 'decoder'.")
|
| 441 |
-
hidden_state_index = model_kwargs.get("hidden_state_index", -1)
|
| 442 |
-
if not isinstance(hidden_state_index, int) or isinstance(hidden_state_index, bool):
|
| 443 |
-
raise TypeError("hidden_state_index must be an integer.")
|
| 444 |
-
store_all_hidden_states = model_kwargs.get("store_all_hidden_states", False)
|
| 445 |
-
if not isinstance(store_all_hidden_states, bool):
|
| 446 |
-
raise TypeError("store_all_hidden_states must be a boolean.")
|
| 447 |
-
normalized_decoder_inputs: tuple[str, ...] | None = None
|
| 448 |
-
has_decoder_inputs = decoder_inputs is not None
|
| 449 |
-
has_decoder_ids = decoder_input_ids is not None
|
| 450 |
-
if hidden_state_source == "encoder":
|
| 451 |
-
if has_decoder_inputs or has_decoder_ids or decoder_attention_mask is not None:
|
| 452 |
-
raise ValueError("Decoder inputs are only valid when hidden_state_source='decoder'.")
|
| 453 |
-
else:
|
| 454 |
-
if has_decoder_inputs == has_decoder_ids:
|
| 455 |
-
raise ValueError(
|
| 456 |
-
"Decoder embedding requires exactly one of decoder_inputs or decoder_input_ids."
|
| 457 |
-
)
|
| 458 |
-
decoder_input_fingerprint: str | None = None
|
| 459 |
-
if decoder_inputs is not None:
|
| 460 |
-
if isinstance(decoder_inputs, (str, bytes)) or not isinstance(decoder_inputs, Sequence):
|
| 461 |
-
raise TypeError("decoder_inputs must be an aligned sequence of strings.")
|
| 462 |
-
normalized_decoder_inputs = tuple(decoder_inputs)
|
| 463 |
-
if not all(isinstance(value, str) and value for value in normalized_decoder_inputs):
|
| 464 |
-
raise ValueError("decoder_inputs must contain non-empty strings.")
|
| 465 |
-
if len(normalized_decoder_inputs) != len(records):
|
| 466 |
-
raise ValueError("decoder_inputs must align one-to-one with embedding inputs.")
|
| 467 |
-
decoder_input_fingerprint = _ordered_string_sha256(normalized_decoder_inputs)
|
| 468 |
-
if decoder_attention_mask is not None:
|
| 469 |
-
raise ValueError("decoder_attention_mask requires decoder_input_ids.")
|
| 470 |
-
if decoder_input_ids is not None:
|
| 471 |
-
if not isinstance(decoder_input_ids, Tensor) or decoder_input_ids.ndim != 2:
|
| 472 |
-
raise ValueError("decoder_input_ids must have shape (batch, sequence).")
|
| 473 |
-
if decoder_input_ids.shape[0] != len(records):
|
| 474 |
-
raise ValueError("decoder_input_ids must align one-to-one with embedding inputs.")
|
| 475 |
-
if decoder_input_ids.dtype == torch.bool or decoder_input_ids.is_floating_point():
|
| 476 |
-
raise TypeError("decoder_input_ids must use an integer token dtype.")
|
| 477 |
-
decoder_input_fingerprint = tensor_sha256(decoder_input_ids)
|
| 478 |
-
decoder_mask_fingerprint: str | None = None
|
| 479 |
-
if decoder_attention_mask is not None:
|
| 480 |
-
if not isinstance(decoder_attention_mask, Tensor):
|
| 481 |
-
raise TypeError("decoder_attention_mask must be a tensor.")
|
| 482 |
-
if decoder_input_ids is None or decoder_attention_mask.shape != decoder_input_ids.shape:
|
| 483 |
-
raise ValueError("decoder_attention_mask must match decoder_input_ids shape.")
|
| 484 |
-
decoder_mask_fingerprint = tensor_sha256(decoder_attention_mask)
|
| 485 |
-
|
| 486 |
-
context: dict[str, Any] = {
|
| 487 |
-
"hidden_state_source": hidden_state_source,
|
| 488 |
-
"hidden_state_index": hidden_state_index,
|
| 489 |
-
"store_all_hidden_states": store_all_hidden_states,
|
| 490 |
-
"decoder_input_fingerprint": decoder_input_fingerprint,
|
| 491 |
-
"decoder_attention_mask_fingerprint": decoder_mask_fingerprint,
|
| 492 |
-
"decoder_alignment": "input-position" if hidden_state_source == "decoder" else None,
|
| 493 |
-
}
|
| 494 |
-
metadata_hook = getattr(model, "_embedding_metadata", None)
|
| 495 |
-
model_metadata: Mapping[str, Any] | None = None
|
| 496 |
-
if callable(metadata_hook):
|
| 497 |
-
model_metadata = metadata_hook(**context)
|
| 498 |
-
if not isinstance(model_metadata, Mapping):
|
| 499 |
-
raise TypeError("_embedding_metadata must return a mapping.")
|
| 500 |
-
context["model_embedding"] = _fingerprint_jsonable(model_metadata)
|
| 501 |
-
if hidden_state_source == "decoder":
|
| 502 |
-
has_decoder_batch = callable(getattr(model, "_embedding_batch", None))
|
| 503 |
-
declares_decoder_stack = (
|
| 504 |
-
model_metadata is not None and model_metadata.get("hidden_state_stack") == "decoder"
|
| 505 |
-
)
|
| 506 |
-
if not has_decoder_batch or not declares_decoder_stack:
|
| 507 |
-
raise ValueError(
|
| 508 |
-
f"{model.__class__.__name__} does not declare decoder embedding support."
|
| 509 |
-
)
|
| 510 |
-
return context, normalized_decoder_inputs
|
|
|
|
| 1 |
+
"""Deterministic identity for embedding inputs, models, tokenizers, and execution."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import hashlib
|
| 6 |
+
import json
|
| 7 |
+
import platform
|
| 8 |
+
import torch
|
| 9 |
+
|
| 10 |
+
from collections.abc import Iterable, Mapping, Sequence
|
| 11 |
+
from pathlib import Path
|
| 12 |
+
from typing import Any
|
| 13 |
+
from torch import Tensor
|
| 14 |
+
|
| 15 |
from .inputs import _InputSpool
|
| 16 |
from .storage import tensor_sha256
|
| 17 |
from .types import EmbeddingInput
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
_RUN_FINGERPRINT_SCHEMA_VERSION = 3
|
| 21 |
+
_MODEL_STATE_HASH_CHUNK_BYTES = 16 * 1024**2
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
def _model_device(model: Any) -> torch.device:
|
| 25 |
+
try:
|
| 26 |
+
return torch.device(next(model.parameters()).device)
|
| 27 |
+
except (AttributeError, StopIteration):
|
| 28 |
+
return torch.device("cpu")
|
| 29 |
+
|
| 30 |
+
|
| 31 |
+
def _attention_backend(model: Any) -> str | None:
|
| 32 |
+
config = getattr(model, "config", None)
|
| 33 |
+
for name in ("_attn_implementation", "attn_implementation", "attn_backend"):
|
| 34 |
+
value = getattr(config, name, None)
|
| 35 |
+
if value:
|
| 36 |
+
return str(value)
|
| 37 |
+
return None
|
| 38 |
+
|
| 39 |
+
|
| 40 |
+
def _attention_kernel_metadata(backend: str | None) -> dict[str, Any] | None:
|
| 41 |
+
if backend not in {"flash_attention_2", "flash_attention_3"}:
|
| 42 |
+
return None
|
| 43 |
+
from fastplms.registry import get_model_registry
|
| 44 |
+
|
| 45 |
+
spec = get_model_registry().attention_kernels[backend]
|
| 46 |
+
return {
|
| 47 |
+
"repository": spec.repository,
|
| 48 |
+
"revision": spec.revision,
|
| 49 |
+
"version": spec.version,
|
| 50 |
+
"expected_variant": spec.expected_variant,
|
| 51 |
+
"dtypes": list(spec.dtypes),
|
| 52 |
+
}
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
def _fingerprint_jsonable(value: Any) -> Any:
|
| 56 |
+
if isinstance(value, Mapping):
|
| 57 |
+
return {str(key): _fingerprint_jsonable(item) for key, item in value.items()}
|
| 58 |
+
if isinstance(value, (list, tuple)):
|
| 59 |
+
return [_fingerprint_jsonable(item) for item in value]
|
| 60 |
+
if isinstance(value, (set, frozenset)):
|
| 61 |
+
return sorted((_fingerprint_jsonable(item) for item in value), key=repr)
|
| 62 |
+
if isinstance(value, Path):
|
| 63 |
+
return str(value)
|
| 64 |
+
if isinstance(value, Tensor):
|
| 65 |
+
return {
|
| 66 |
+
"dtype": str(value.dtype).removeprefix("torch."),
|
| 67 |
+
"shape": list(value.shape),
|
| 68 |
+
"sha256": tensor_sha256(value),
|
| 69 |
+
}
|
| 70 |
+
if isinstance(value, torch.dtype):
|
| 71 |
+
return str(value).removeprefix("torch.")
|
| 72 |
+
if isinstance(value, torch.device):
|
| 73 |
+
return str(value)
|
| 74 |
+
if value is None or isinstance(value, (str, int, float, bool)):
|
| 75 |
+
return value
|
| 76 |
+
return {
|
| 77 |
+
"class": f"{value.__class__.__module__}.{value.__class__.__qualname__}",
|
| 78 |
+
"value": str(value),
|
| 79 |
+
}
|
| 80 |
+
|
| 81 |
+
|
| 82 |
+
def _tokenizer_content_sha256(tokenizer: Any) -> str:
|
| 83 |
+
content: dict[str, Any] = {
|
| 84 |
+
"init_kwargs": getattr(tokenizer, "init_kwargs", None),
|
| 85 |
+
"special_tokens_map": getattr(tokenizer, "special_tokens_map", None),
|
| 86 |
+
"model_max_length": getattr(tokenizer, "model_max_length", None),
|
| 87 |
+
"padding_side": getattr(tokenizer, "padding_side", None),
|
| 88 |
+
"truncation_side": getattr(tokenizer, "truncation_side", None),
|
| 89 |
+
}
|
| 90 |
+
get_vocab = getattr(tokenizer, "get_vocab", None)
|
| 91 |
+
if callable(get_vocab):
|
| 92 |
+
content["vocabulary"] = get_vocab()
|
| 93 |
+
get_added_vocab = getattr(tokenizer, "get_added_vocab", None)
|
| 94 |
+
if callable(get_added_vocab):
|
| 95 |
+
content["added_vocabulary"] = get_added_vocab()
|
| 96 |
+
backend = getattr(tokenizer, "backend_tokenizer", None)
|
| 97 |
+
backend_to_str = getattr(backend, "to_str", None)
|
| 98 |
+
if callable(backend_to_str):
|
| 99 |
+
content["backend"] = backend_to_str()
|
| 100 |
+
serialized = json.dumps(
|
| 101 |
+
_fingerprint_jsonable(content),
|
| 102 |
+
sort_keys=True,
|
| 103 |
+
separators=(",", ":"),
|
| 104 |
+
ensure_ascii=False,
|
| 105 |
+
).encode()
|
| 106 |
+
return hashlib.sha256(serialized).hexdigest()
|
| 107 |
+
|
| 108 |
+
|
| 109 |
+
def _tokenizer_metadata(model: Any, tokenizer: Any | None) -> dict[str, Any]:
|
| 110 |
+
resolved = tokenizer if tokenizer is not None else getattr(model, "tokenizer", None)
|
| 111 |
+
if resolved is None:
|
| 112 |
+
# Raw-sequence families such as E1 retain their loader context on the
|
| 113 |
+
# model/encoder rather than exposing a Transformers tokenizer. Bind the
|
| 114 |
+
# non-secret source policy to resume identity without serializing a Hub
|
| 115 |
+
# token or forcing lazy tokenizer initialization.
|
| 116 |
+
for candidate in (model, getattr(model, "model", None)):
|
| 117 |
+
settings = getattr(candidate, "__dict__", {}).get("_fastplms_tokenizer_kwargs")
|
| 118 |
+
if isinstance(settings, Mapping):
|
| 119 |
+
token_value = settings.get("token")
|
| 120 |
+
return {
|
| 121 |
+
"mode": "native-sequence",
|
| 122 |
+
"source": (
|
| 123 |
+
str(settings.get("tokenizer_source"))
|
| 124 |
+
if settings.get("tokenizer_source") is not None
|
| 125 |
+
else None
|
| 126 |
+
),
|
| 127 |
+
"revision": settings.get("revision"),
|
| 128 |
+
"cache_dir": (
|
| 129 |
+
str(settings.get("cache_dir"))
|
| 130 |
+
if settings.get("cache_dir") is not None
|
| 131 |
+
else None
|
| 132 |
+
),
|
| 133 |
+
"local_files_only": bool(settings.get("local_files_only", False)),
|
| 134 |
+
"token_policy": (
|
| 135 |
+
"disabled"
|
| 136 |
+
if token_value is False
|
| 137 |
+
else "provided"
|
| 138 |
+
if token_value is not None
|
| 139 |
+
else "default"
|
| 140 |
+
),
|
| 141 |
+
}
|
| 142 |
+
return {"mode": "native-sequence"}
|
| 143 |
+
return {
|
| 144 |
+
"mode": "tokenizer",
|
| 145 |
+
"class": f"{resolved.__class__.__module__}.{resolved.__class__.__qualname__}",
|
| 146 |
+
"name_or_path": getattr(resolved, "name_or_path", None),
|
| 147 |
+
"vocab_size": getattr(resolved, "vocab_size", None),
|
| 148 |
+
"special_token_ids": list(getattr(resolved, "all_special_ids", ())),
|
| 149 |
+
"content_sha256": _tokenizer_content_sha256(resolved),
|
| 150 |
+
}
|
| 151 |
+
|
| 152 |
+
|
| 153 |
+
def _software_versions() -> dict[str, str | None]:
|
| 154 |
+
try:
|
| 155 |
+
import fastplms
|
| 156 |
+
|
| 157 |
+
fastplms_version = fastplms.__version__
|
| 158 |
+
except (AttributeError, ImportError):
|
| 159 |
+
fastplms_version = None
|
| 160 |
+
try:
|
| 161 |
+
import safetensors
|
| 162 |
+
|
| 163 |
+
safetensors_version = safetensors.__version__
|
| 164 |
+
except ImportError:
|
| 165 |
+
safetensors_version = None
|
| 166 |
+
try:
|
| 167 |
+
import transformers
|
| 168 |
+
|
| 169 |
+
transformers_version = transformers.__version__
|
| 170 |
+
except ImportError:
|
| 171 |
+
transformers_version = None
|
| 172 |
+
return {
|
| 173 |
+
"fastplms": fastplms_version,
|
| 174 |
+
"python": platform.python_version(),
|
| 175 |
+
"safetensors": safetensors_version,
|
| 176 |
+
"torch": torch.__version__,
|
| 177 |
+
"torch_cuda": torch.version.cuda,
|
| 178 |
+
"transformers": transformers_version,
|
| 179 |
+
}
|
| 180 |
+
|
| 181 |
+
|
| 182 |
+
def _adapter_identity_metadata(model: Any) -> dict[str, Any] | None:
|
| 183 |
+
"""Return deterministic PEFT/adapter identity without tensor payloads."""
|
| 184 |
+
|
| 185 |
+
peft_config = getattr(model, "peft_config", None)
|
| 186 |
+
if not isinstance(peft_config, Mapping) or not peft_config:
|
| 187 |
+
return None
|
| 188 |
+
configurations: dict[str, Any] = {}
|
| 189 |
+
for name, config in sorted(peft_config.items(), key=lambda item: str(item[0])):
|
| 190 |
+
to_dict = getattr(config, "to_dict", None)
|
| 191 |
+
if callable(to_dict):
|
| 192 |
+
value = to_dict()
|
| 193 |
+
else:
|
| 194 |
+
try:
|
| 195 |
+
value = vars(config)
|
| 196 |
+
except TypeError:
|
| 197 |
+
value = config
|
| 198 |
+
configurations[str(name)] = _fingerprint_jsonable(value)
|
| 199 |
+
active_adapters = getattr(model, "active_adapters", None)
|
| 200 |
+
if callable(active_adapters):
|
| 201 |
+
active_adapters = active_adapters()
|
| 202 |
+
return {
|
| 203 |
+
"active": _fingerprint_jsonable(active_adapters),
|
| 204 |
+
"configurations": configurations,
|
| 205 |
+
}
|
| 206 |
+
|
| 207 |
+
|
| 208 |
+
def _execution_identity_metadata(model: Any) -> dict[str, Any]:
|
| 209 |
+
"""Capture runtime policy that can change persisted numerical results."""
|
| 210 |
+
|
| 211 |
+
parameter_dtypes = sorted(
|
| 212 |
+
{
|
| 213 |
+
str(parameter.dtype).removeprefix("torch.")
|
| 214 |
+
for parameter in getattr(model, "parameters", lambda: ())()
|
| 215 |
+
}
|
| 216 |
+
)
|
| 217 |
+
return {
|
| 218 |
+
"device": _model_device(model).type,
|
| 219 |
+
"hf_device_map": _fingerprint_jsonable(getattr(model, "hf_device_map", None)),
|
| 220 |
+
"parameter_dtypes": parameter_dtypes,
|
| 221 |
+
"software": _software_versions(),
|
| 222 |
+
}
|
| 223 |
+
|
| 224 |
+
|
| 225 |
+
def _first_metadata_value(*values: Any) -> Any:
|
| 226 |
+
for value in values:
|
| 227 |
+
if isinstance(value, str):
|
| 228 |
+
if value.strip():
|
| 229 |
+
return value
|
| 230 |
+
elif value is not None:
|
| 231 |
+
return value
|
| 232 |
+
return None
|
| 233 |
+
|
| 234 |
+
|
| 235 |
+
def _model_identity_metadata(model: Any) -> dict[str, Any]:
|
| 236 |
+
"""Resolve model and checkpoint identity, including local artifact fallbacks."""
|
| 237 |
+
|
| 238 |
+
config = getattr(model, "config", None)
|
| 239 |
+
checkpoint_revision = _first_metadata_value(
|
| 240 |
+
getattr(config, "fastplms_checkpoint_revision", None),
|
| 241 |
+
getattr(config, "_commit_hash", None),
|
| 242 |
+
)
|
| 243 |
+
return {
|
| 244 |
+
"model_id": _first_metadata_value(
|
| 245 |
+
getattr(config, "fastplms_model_id", None),
|
| 246 |
+
getattr(config, "_name_or_path", None),
|
| 247 |
+
),
|
| 248 |
+
"model_revision": _first_metadata_value(
|
| 249 |
+
getattr(config, "_commit_hash", None),
|
| 250 |
+
checkpoint_revision,
|
| 251 |
+
),
|
| 252 |
+
"checkpoint_repo_id": getattr(config, "fastplms_checkpoint_repo_id", None),
|
| 253 |
+
"checkpoint_revision": checkpoint_revision,
|
| 254 |
+
"checkpoint_hash": _first_metadata_value(
|
| 255 |
+
getattr(model, "checkpoint_hash", None),
|
| 256 |
+
getattr(config, "checkpoint_hash", None),
|
| 257 |
+
getattr(config, "fastplms_checkpoint_hash", None),
|
| 258 |
+
),
|
| 259 |
+
"weights_revision": getattr(config, "fastplms_weights_revision", None),
|
| 260 |
+
"runtime_revision": getattr(config, "fastplms_runtime_revision", None),
|
| 261 |
+
"source_tree_sha256": getattr(config, "fastplms_source_tree_sha256", None),
|
| 262 |
+
"runtime_bundle_sha256": getattr(config, "fastplms_runtime_bundle_sha256", None),
|
| 263 |
+
}
|
| 264 |
+
|
| 265 |
+
|
| 266 |
+
def _bounded_tensor_chunks(X: Tensor, max_elements: int) -> Iterable[Tensor]:
|
| 267 |
+
"""Yield X in logical row-major order without materializing a full copy."""
|
| 268 |
+
|
| 269 |
+
# X: (...)
|
| 270 |
+
if X.numel() == 0:
|
| 271 |
+
return
|
| 272 |
+
if X.ndim == 0:
|
| 273 |
+
yield X
|
| 274 |
+
return
|
| 275 |
+
trailing_elements = 1
|
| 276 |
+
for size in X.shape[1:]:
|
| 277 |
+
trailing_elements *= int(size)
|
| 278 |
+
if trailing_elements <= max_elements:
|
| 279 |
+
rows_per_chunk = max(1, max_elements // trailing_elements)
|
| 280 |
+
for start in range(0, X.shape[0], rows_per_chunk):
|
| 281 |
+
yield X[start : start + rows_per_chunk] # (chunk_rows, ...)
|
| 282 |
+
return
|
| 283 |
+
for row in X:
|
| 284 |
+
yield from _bounded_tensor_chunks(row, max_elements)
|
| 285 |
+
|
| 286 |
+
|
| 287 |
+
def _model_state_sha256(model: Any) -> str:
|
| 288 |
+
"""Hash named parameters and persistent buffers using bounded CPU copies."""
|
| 289 |
+
|
| 290 |
+
# Never cache this digest from tensor identity or ``Tensor._version``.
|
| 291 |
+
# ``Parameter.data`` and independent tensor aliases can mutate shared storage
|
| 292 |
+
# without changing either signal, while persisted resume identity must bind
|
| 293 |
+
# the authoritative bytes visible at the start of this run.
|
| 294 |
+
state = model.state_dict(keep_vars=True)
|
| 295 |
+
digest = hashlib.sha256()
|
| 296 |
+
for name, value in sorted(state.items()):
|
| 297 |
+
if not isinstance(value, Tensor):
|
| 298 |
+
raise TypeError(f"Model state entry {name!r} is not a tensor.")
|
| 299 |
+
if value.is_meta:
|
| 300 |
+
raise ValueError(
|
| 301 |
+
f"Cannot fingerprint meta-device model state entry {name!r}; pass "
|
| 302 |
+
"model_state_fingerprint with a caller-owned state identity."
|
| 303 |
+
)
|
| 304 |
+
header = json.dumps(
|
| 305 |
+
{
|
| 306 |
+
"name": name,
|
| 307 |
+
"dtype": str(value.dtype).removeprefix("torch."),
|
| 308 |
+
"shape": list(value.shape),
|
| 309 |
+
},
|
| 310 |
+
sort_keys=True,
|
| 311 |
+
separators=(",", ":"),
|
| 312 |
+
).encode()
|
| 313 |
+
digest.update(len(header).to_bytes(8, "big"))
|
| 314 |
+
digest.update(header)
|
| 315 |
+
max_elements = max(1, _MODEL_STATE_HASH_CHUNK_BYTES // value.element_size())
|
| 316 |
+
for chunk in _bounded_tensor_chunks(value.detach(), max_elements):
|
| 317 |
+
cpu_chunk = chunk.to(device="cpu").contiguous() # chunk.shape
|
| 318 |
+
digest.update(cpu_chunk.reshape(-1).view(torch.uint8).numpy().tobytes())
|
| 319 |
+
return digest.hexdigest()
|
| 320 |
+
|
| 321 |
+
|
| 322 |
+
def _input_sha256(records: Iterable[EmbeddingInput]) -> str:
|
| 323 |
+
"""Hash an ordered input stream without constructing a duplicate JSON payload."""
|
| 324 |
+
|
| 325 |
+
precomputed = getattr(records, "input_fingerprint", None)
|
| 326 |
+
if isinstance(precomputed, str):
|
| 327 |
+
return precomputed
|
| 328 |
+
digest = hashlib.sha256()
|
| 329 |
+
count = 0
|
| 330 |
+
for record in records:
|
| 331 |
+
count += 1
|
| 332 |
+
for value in (record.id, record.sequence):
|
| 333 |
+
encoded = value.encode("utf-8")
|
| 334 |
+
digest.update(len(encoded).to_bytes(8, "big"))
|
| 335 |
+
digest.update(encoded)
|
| 336 |
+
digest.update(count.to_bytes(8, "big"))
|
| 337 |
+
return digest.hexdigest()
|
| 338 |
+
|
| 339 |
+
|
| 340 |
+
def _run_fingerprint(
|
| 341 |
+
model: Any,
|
| 342 |
+
records: Sequence[EmbeddingInput],
|
| 343 |
+
*,
|
| 344 |
+
pooling: Sequence[str],
|
| 345 |
+
full_embeddings: bool,
|
| 346 |
+
max_length: int | None,
|
| 347 |
+
truncate: bool,
|
| 348 |
+
dtype: torch.dtype | None,
|
| 349 |
+
model_kwargs: dict[str, Any],
|
| 350 |
+
tokenizer_metadata: dict[str, Any],
|
| 351 |
+
model_state_fingerprint: str | None,
|
| 352 |
+
persist_output: bool,
|
| 353 |
+
embedding_context: Mapping[str, Any],
|
| 354 |
+
batch_size: int,
|
| 355 |
+
batch_window_size: int,
|
| 356 |
+
max_tokens_per_batch: int | None,
|
| 357 |
+
) -> tuple[str, str, str | None, str]:
|
| 358 |
+
input_fingerprint = _input_sha256(records)
|
| 359 |
+
attention_backend = _attention_backend(model)
|
| 360 |
+
model_identity = _model_identity_metadata(model)
|
| 361 |
+
if model_state_fingerprint is None and persist_output:
|
| 362 |
+
resolved_model_state_fingerprint = _model_state_sha256(model)
|
| 363 |
+
model_state_fingerprint_source = "computed"
|
| 364 |
+
elif model_state_fingerprint is not None:
|
| 365 |
+
resolved_model_state_fingerprint = model_state_fingerprint.strip()
|
| 366 |
+
if not resolved_model_state_fingerprint:
|
| 367 |
+
raise ValueError("model_state_fingerprint must not be empty.")
|
| 368 |
+
model_state_fingerprint_source = "caller"
|
| 369 |
+
else:
|
| 370 |
+
resolved_model_state_fingerprint = None
|
| 371 |
+
model_state_fingerprint_source = "not-computed"
|
| 372 |
+
payload = {
|
| 373 |
+
"fingerprint_schema_version": _RUN_FINGERPRINT_SCHEMA_VERSION,
|
| 374 |
+
"input_fingerprint": input_fingerprint,
|
| 375 |
+
"model_state_fingerprint": resolved_model_state_fingerprint,
|
| 376 |
+
"model_state_fingerprint_source": model_state_fingerprint_source,
|
| 377 |
+
"model_class": f"{model.__class__.__module__}.{model.__class__.__qualname__}",
|
| 378 |
+
**model_identity,
|
| 379 |
+
"attention_backend": attention_backend,
|
| 380 |
+
"attention_kernel": _attention_kernel_metadata(attention_backend),
|
| 381 |
+
"layer": repr(
|
| 382 |
+
getattr(model, "embedding_layer", model_kwargs.get("hidden_state_index", -1))
|
| 383 |
+
),
|
| 384 |
+
"projection": getattr(model, "embedding_projection", None),
|
| 385 |
+
"esmc_source": getattr(model, "_esmc_source", None),
|
| 386 |
+
"esmc_revision": getattr(model, "_esmc_source_revision", None),
|
| 387 |
+
"esmc_files": getattr(model, "_esmc_source_files", None),
|
| 388 |
+
"token_policy": getattr(model, "embedding_token_policy", None),
|
| 389 |
+
"tokenizer": tokenizer_metadata,
|
| 390 |
+
"adapter": _adapter_identity_metadata(model),
|
| 391 |
+
"execution": _execution_identity_metadata(model),
|
| 392 |
+
"embedding_context": _fingerprint_jsonable(embedding_context),
|
| 393 |
+
"pooling": list(pooling),
|
| 394 |
+
"full_embeddings": full_embeddings,
|
| 395 |
+
"max_length": max_length,
|
| 396 |
+
"truncate": truncate,
|
| 397 |
+
"dtype": str(dtype) if dtype is not None else None,
|
| 398 |
+
"batching": {
|
| 399 |
+
"batch_size": batch_size,
|
| 400 |
+
"batch_window_size": batch_window_size,
|
| 401 |
+
"max_tokens_per_batch": max_tokens_per_batch,
|
| 402 |
+
"input_storage": ("disk-spool" if isinstance(records, _InputSpool) else "memory"),
|
| 403 |
+
},
|
| 404 |
+
"model_kwargs": {
|
| 405 |
+
key: _fingerprint_jsonable(value) for key, value in sorted(model_kwargs.items())
|
| 406 |
+
},
|
| 407 |
+
"residue_mask_policy": "attention-mask-minus-special-tokens",
|
| 408 |
+
}
|
| 409 |
+
run_fingerprint = hashlib.sha256(
|
| 410 |
+
json.dumps(payload, sort_keys=True, separators=(",", ":")).encode()
|
| 411 |
+
).hexdigest()
|
| 412 |
+
return (
|
| 413 |
+
input_fingerprint,
|
| 414 |
+
run_fingerprint,
|
| 415 |
+
resolved_model_state_fingerprint,
|
| 416 |
+
model_state_fingerprint_source,
|
| 417 |
+
)
|
| 418 |
+
|
| 419 |
+
|
| 420 |
+
def _ordered_string_sha256(values: Sequence[str]) -> str:
|
| 421 |
+
digest = hashlib.sha256()
|
| 422 |
+
for value in values:
|
| 423 |
+
encoded = value.encode("utf-8")
|
| 424 |
+
digest.update(len(encoded).to_bytes(8, "big"))
|
| 425 |
+
digest.update(encoded)
|
| 426 |
+
digest.update(len(values).to_bytes(8, "big"))
|
| 427 |
+
return digest.hexdigest()
|
| 428 |
+
|
| 429 |
+
|
| 430 |
+
def _embedding_context(
|
| 431 |
+
model: Any,
|
| 432 |
+
records: Sequence[EmbeddingInput],
|
| 433 |
+
*,
|
| 434 |
+
hidden_state_source: str,
|
| 435 |
+
decoder_inputs: Sequence[str] | None,
|
| 436 |
+
decoder_input_ids: Tensor | None,
|
| 437 |
+
decoder_attention_mask: Tensor | None,
|
| 438 |
+
model_kwargs: Mapping[str, Any],
|
| 439 |
+
) -> tuple[dict[str, Any], tuple[str, ...] | None]:
|
| 440 |
+
if hidden_state_source not in {"encoder", "decoder"}:
|
| 441 |
+
raise ValueError("hidden_state_source must be 'encoder' or 'decoder'.")
|
| 442 |
+
hidden_state_index = model_kwargs.get("hidden_state_index", -1)
|
| 443 |
+
if not isinstance(hidden_state_index, int) or isinstance(hidden_state_index, bool):
|
| 444 |
+
raise TypeError("hidden_state_index must be an integer.")
|
| 445 |
+
store_all_hidden_states = model_kwargs.get("store_all_hidden_states", False)
|
| 446 |
+
if not isinstance(store_all_hidden_states, bool):
|
| 447 |
+
raise TypeError("store_all_hidden_states must be a boolean.")
|
| 448 |
+
normalized_decoder_inputs: tuple[str, ...] | None = None
|
| 449 |
+
has_decoder_inputs = decoder_inputs is not None
|
| 450 |
+
has_decoder_ids = decoder_input_ids is not None
|
| 451 |
+
if hidden_state_source == "encoder":
|
| 452 |
+
if has_decoder_inputs or has_decoder_ids or decoder_attention_mask is not None:
|
| 453 |
+
raise ValueError("Decoder inputs are only valid when hidden_state_source='decoder'.")
|
| 454 |
+
else:
|
| 455 |
+
if has_decoder_inputs == has_decoder_ids:
|
| 456 |
+
raise ValueError(
|
| 457 |
+
"Decoder embedding requires exactly one of decoder_inputs or decoder_input_ids."
|
| 458 |
+
)
|
| 459 |
+
decoder_input_fingerprint: str | None = None
|
| 460 |
+
if decoder_inputs is not None:
|
| 461 |
+
if isinstance(decoder_inputs, (str, bytes)) or not isinstance(decoder_inputs, Sequence):
|
| 462 |
+
raise TypeError("decoder_inputs must be an aligned sequence of strings.")
|
| 463 |
+
normalized_decoder_inputs = tuple(decoder_inputs)
|
| 464 |
+
if not all(isinstance(value, str) and value for value in normalized_decoder_inputs):
|
| 465 |
+
raise ValueError("decoder_inputs must contain non-empty strings.")
|
| 466 |
+
if len(normalized_decoder_inputs) != len(records):
|
| 467 |
+
raise ValueError("decoder_inputs must align one-to-one with embedding inputs.")
|
| 468 |
+
decoder_input_fingerprint = _ordered_string_sha256(normalized_decoder_inputs)
|
| 469 |
+
if decoder_attention_mask is not None:
|
| 470 |
+
raise ValueError("decoder_attention_mask requires decoder_input_ids.")
|
| 471 |
+
if decoder_input_ids is not None:
|
| 472 |
+
if not isinstance(decoder_input_ids, Tensor) or decoder_input_ids.ndim != 2:
|
| 473 |
+
raise ValueError("decoder_input_ids must have shape (batch, sequence).")
|
| 474 |
+
if decoder_input_ids.shape[0] != len(records):
|
| 475 |
+
raise ValueError("decoder_input_ids must align one-to-one with embedding inputs.")
|
| 476 |
+
if decoder_input_ids.dtype == torch.bool or decoder_input_ids.is_floating_point():
|
| 477 |
+
raise TypeError("decoder_input_ids must use an integer token dtype.")
|
| 478 |
+
decoder_input_fingerprint = tensor_sha256(decoder_input_ids)
|
| 479 |
+
decoder_mask_fingerprint: str | None = None
|
| 480 |
+
if decoder_attention_mask is not None:
|
| 481 |
+
if not isinstance(decoder_attention_mask, Tensor):
|
| 482 |
+
raise TypeError("decoder_attention_mask must be a tensor.")
|
| 483 |
+
if decoder_input_ids is None or decoder_attention_mask.shape != decoder_input_ids.shape:
|
| 484 |
+
raise ValueError("decoder_attention_mask must match decoder_input_ids shape.")
|
| 485 |
+
decoder_mask_fingerprint = tensor_sha256(decoder_attention_mask)
|
| 486 |
+
|
| 487 |
+
context: dict[str, Any] = {
|
| 488 |
+
"hidden_state_source": hidden_state_source,
|
| 489 |
+
"hidden_state_index": hidden_state_index,
|
| 490 |
+
"store_all_hidden_states": store_all_hidden_states,
|
| 491 |
+
"decoder_input_fingerprint": decoder_input_fingerprint,
|
| 492 |
+
"decoder_attention_mask_fingerprint": decoder_mask_fingerprint,
|
| 493 |
+
"decoder_alignment": "input-position" if hidden_state_source == "decoder" else None,
|
| 494 |
+
}
|
| 495 |
+
metadata_hook = getattr(model, "_embedding_metadata", None)
|
| 496 |
+
model_metadata: Mapping[str, Any] | None = None
|
| 497 |
+
if callable(metadata_hook):
|
| 498 |
+
model_metadata = metadata_hook(**context)
|
| 499 |
+
if not isinstance(model_metadata, Mapping):
|
| 500 |
+
raise TypeError("_embedding_metadata must return a mapping.")
|
| 501 |
+
context["model_embedding"] = _fingerprint_jsonable(model_metadata)
|
| 502 |
+
if hidden_state_source == "decoder":
|
| 503 |
+
has_decoder_batch = callable(getattr(model, "_embedding_batch", None))
|
| 504 |
+
declares_decoder_stack = (
|
| 505 |
+
model_metadata is not None and model_metadata.get("hidden_state_stack") == "decoder"
|
| 506 |
+
)
|
| 507 |
+
if not has_decoder_batch or not declares_decoder_stack:
|
| 508 |
+
raise ValueError(
|
| 509 |
+
f"{model.__class__.__name__} does not declare decoder embedding support."
|
| 510 |
+
)
|
| 511 |
+
return context, normalized_decoder_inputs
|
fastplms/embeddings/inputs.py
CHANGED
|
@@ -1,263 +1,264 @@
|
|
| 1 |
-
"""Normalize ordered inputs and plan bounded windows without retaining a full stream."""
|
| 2 |
-
|
| 3 |
-
from __future__ import annotations
|
| 4 |
-
|
| 5 |
-
import hashlib
|
| 6 |
-
import sqlite3
|
| 7 |
-
import tempfile
|
| 8 |
-
|
| 9 |
-
from
|
| 10 |
-
from
|
| 11 |
-
|
| 12 |
-
|
| 13 |
-
|
| 14 |
-
|
| 15 |
-
|
| 16 |
-
|
| 17 |
-
|
| 18 |
-
|
| 19 |
-
|
| 20 |
-
|
| 21 |
-
|
| 22 |
-
|
| 23 |
-
|
| 24 |
-
|
| 25 |
-
|
| 26 |
-
|
| 27 |
-
|
| 28 |
-
|
| 29 |
-
|
| 30 |
-
|
| 31 |
-
|
| 32 |
-
|
| 33 |
-
|
| 34 |
-
|
| 35 |
-
|
| 36 |
-
|
| 37 |
-
|
| 38 |
-
|
| 39 |
-
|
| 40 |
-
|
| 41 |
-
|
| 42 |
-
|
| 43 |
-
|
| 44 |
-
|
| 45 |
-
|
| 46 |
-
|
| 47 |
-
|
| 48 |
-
|
| 49 |
-
|
| 50 |
-
|
| 51 |
-
|
| 52 |
-
|
| 53 |
-
|
| 54 |
-
|
| 55 |
-
|
| 56 |
-
|
| 57 |
-
|
| 58 |
-
|
| 59 |
-
|
| 60 |
-
|
| 61 |
-
|
| 62 |
-
|
| 63 |
-
|
| 64 |
-
|
| 65 |
-
|
| 66 |
-
|
| 67 |
-
|
| 68 |
-
|
| 69 |
-
|
| 70 |
-
|
| 71 |
-
|
| 72 |
-
|
| 73 |
-
|
| 74 |
-
|
| 75 |
-
|
| 76 |
-
|
| 77 |
-
|
| 78 |
-
|
| 79 |
-
self.
|
| 80 |
-
self._connection.
|
| 81 |
-
|
| 82 |
-
"
|
| 83 |
-
|
| 84 |
-
|
| 85 |
-
|
| 86 |
-
|
| 87 |
-
|
| 88 |
-
|
| 89 |
-
|
| 90 |
-
|
| 91 |
-
|
| 92 |
-
|
| 93 |
-
digest.update(encoded)
|
| 94 |
-
|
| 95 |
-
|
| 96 |
-
|
| 97 |
-
|
| 98 |
-
|
| 99 |
-
|
| 100 |
-
|
| 101 |
-
|
| 102 |
-
|
| 103 |
-
|
| 104 |
-
self._connection.
|
| 105 |
-
self._connection
|
| 106 |
-
|
| 107 |
-
|
| 108 |
-
|
| 109 |
-
|
| 110 |
-
|
| 111 |
-
|
| 112 |
-
|
| 113 |
-
|
| 114 |
-
self.
|
| 115 |
-
|
| 116 |
-
|
| 117 |
-
|
| 118 |
-
|
| 119 |
-
|
| 120 |
-
|
| 121 |
-
|
| 122 |
-
|
| 123 |
-
|
| 124 |
-
|
| 125 |
-
|
| 126 |
-
|
| 127 |
-
|
| 128 |
-
|
| 129 |
-
|
| 130 |
-
|
| 131 |
-
|
| 132 |
-
|
| 133 |
-
|
| 134 |
-
|
| 135 |
-
|
| 136 |
-
|
| 137 |
-
|
| 138 |
-
|
| 139 |
-
|
| 140 |
-
|
| 141 |
-
|
| 142 |
-
|
| 143 |
-
|
| 144 |
-
|
| 145 |
-
|
| 146 |
-
|
| 147 |
-
"
|
| 148 |
-
|
| 149 |
-
|
| 150 |
-
|
| 151 |
-
|
| 152 |
-
|
| 153 |
-
|
| 154 |
-
|
| 155 |
-
|
| 156 |
-
|
| 157 |
-
|
| 158 |
-
|
| 159 |
-
|
| 160 |
-
|
| 161 |
-
|
| 162 |
-
|
| 163 |
-
|
| 164 |
-
|
| 165 |
-
|
| 166 |
-
|
| 167 |
-
|
| 168 |
-
|
| 169 |
-
|
| 170 |
-
|
| 171 |
-
|
| 172 |
-
|
| 173 |
-
|
| 174 |
-
|
| 175 |
-
|
| 176 |
-
|
| 177 |
-
|
| 178 |
-
|
| 179 |
-
|
| 180 |
-
|
| 181 |
-
|
| 182 |
-
|
| 183 |
-
|
| 184 |
-
|
| 185 |
-
|
| 186 |
-
|
| 187 |
-
|
| 188 |
-
|
| 189 |
-
|
| 190 |
-
|
| 191 |
-
|
| 192 |
-
|
| 193 |
-
|
| 194 |
-
|
| 195 |
-
|
| 196 |
-
|
| 197 |
-
|
| 198 |
-
|
| 199 |
-
|
| 200 |
-
|
| 201 |
-
|
| 202 |
-
|
| 203 |
-
|
| 204 |
-
|
| 205 |
-
|
| 206 |
-
|
| 207 |
-
|
| 208 |
-
|
| 209 |
-
|
| 210 |
-
|
| 211 |
-
|
| 212 |
-
|
| 213 |
-
|
| 214 |
-
|
| 215 |
-
|
| 216 |
-
|
| 217 |
-
|
| 218 |
-
|
| 219 |
-
|
| 220 |
-
|
| 221 |
-
f"{
|
| 222 |
-
"
|
| 223 |
-
|
| 224 |
-
|
| 225 |
-
|
| 226 |
-
|
| 227 |
-
|
| 228 |
-
|
| 229 |
-
|
| 230 |
-
|
| 231 |
-
|
| 232 |
-
|
| 233 |
-
|
| 234 |
-
|
| 235 |
-
|
| 236 |
-
|
| 237 |
-
|
| 238 |
-
|
| 239 |
-
|
| 240 |
-
|
| 241 |
-
|
| 242 |
-
|
| 243 |
-
|
| 244 |
-
|
| 245 |
-
|
| 246 |
-
|
| 247 |
-
|
| 248 |
-
|
| 249 |
-
f"
|
| 250 |
-
|
| 251 |
-
|
| 252 |
-
|
| 253 |
-
|
| 254 |
-
|
| 255 |
-
|
| 256 |
-
|
| 257 |
-
|
| 258 |
-
|
| 259 |
-
|
| 260 |
-
|
| 261 |
-
|
| 262 |
-
|
| 263 |
-
|
|
|
|
|
|
| 1 |
+
"""Normalize ordered inputs and plan bounded windows without retaining a full stream."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import hashlib
|
| 6 |
+
import sqlite3
|
| 7 |
+
import tempfile
|
| 8 |
+
|
| 9 |
+
from collections.abc import Iterable, Iterator, Mapping, Sequence
|
| 10 |
+
from pathlib import Path
|
| 11 |
+
from typing import overload
|
| 12 |
+
|
| 13 |
+
from .types import EmbeddingInput
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
def iter_fasta(path: str | Path) -> Iterator[EmbeddingInput]:
|
| 17 |
+
"""Yield FASTA records in source order without reading the file into memory."""
|
| 18 |
+
|
| 19 |
+
identifier: str | None = None
|
| 20 |
+
sequence_parts: list[str] = []
|
| 21 |
+
found_record = False
|
| 22 |
+
with Path(path).open("r", encoding="utf-8") as handle:
|
| 23 |
+
for line_number, raw_line in enumerate(handle, start=1):
|
| 24 |
+
line = raw_line.strip()
|
| 25 |
+
if not line:
|
| 26 |
+
continue
|
| 27 |
+
if line.startswith(">"):
|
| 28 |
+
if identifier is not None:
|
| 29 |
+
found_record = True
|
| 30 |
+
yield EmbeddingInput(identifier, "".join(sequence_parts))
|
| 31 |
+
identifier = line[1:].strip().split(maxsplit=1)[0]
|
| 32 |
+
if not identifier:
|
| 33 |
+
raise ValueError(f"Missing FASTA identifier on line {line_number}.")
|
| 34 |
+
sequence_parts = []
|
| 35 |
+
else:
|
| 36 |
+
if identifier is None:
|
| 37 |
+
raise ValueError(
|
| 38 |
+
f"Sequence data precedes the first FASTA header on line {line_number}."
|
| 39 |
+
)
|
| 40 |
+
sequence_parts.append("".join(line.split()))
|
| 41 |
+
if identifier is not None:
|
| 42 |
+
found_record = True
|
| 43 |
+
yield EmbeddingInput(identifier, "".join(sequence_parts))
|
| 44 |
+
if not found_record:
|
| 45 |
+
raise ValueError(f"No FASTA records found in {path}.")
|
| 46 |
+
|
| 47 |
+
|
| 48 |
+
def parse_fasta(path: str | Path) -> list[EmbeddingInput]:
|
| 49 |
+
"""Parse FASTA records while preserving identifiers, order, and duplicates."""
|
| 50 |
+
|
| 51 |
+
return list(iter_fasta(path))
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
def _normalize_input_item(
|
| 55 |
+
position: int,
|
| 56 |
+
item: str | EmbeddingInput | tuple[str, str],
|
| 57 |
+
) -> EmbeddingInput:
|
| 58 |
+
if isinstance(item, EmbeddingInput):
|
| 59 |
+
return item
|
| 60 |
+
if isinstance(item, str):
|
| 61 |
+
return EmbeddingInput(str(position), item)
|
| 62 |
+
if isinstance(item, tuple) and len(item) == 2:
|
| 63 |
+
return EmbeddingInput(str(item[0]), str(item[1]))
|
| 64 |
+
raise TypeError(
|
| 65 |
+
"inputs must contain sequences, EmbeddingInput values, or (id, sequence) tuples."
|
| 66 |
+
)
|
| 67 |
+
|
| 68 |
+
|
| 69 |
+
class _InputSpool(Sequence[EmbeddingInput]):
|
| 70 |
+
"""Immutable disk-backed normalized inputs with an incremental digest."""
|
| 71 |
+
|
| 72 |
+
def __init__(
|
| 73 |
+
self,
|
| 74 |
+
values: Iterable[str | EmbeddingInput | tuple[str, str]],
|
| 75 |
+
) -> None:
|
| 76 |
+
self._temporary: tempfile.TemporaryDirectory[str] | None = tempfile.TemporaryDirectory(
|
| 77 |
+
prefix="fastplms-inputs-"
|
| 78 |
+
)
|
| 79 |
+
self.path = Path(self._temporary.name) / "inputs.sqlite"
|
| 80 |
+
self._connection: sqlite3.Connection | None = sqlite3.connect(self.path)
|
| 81 |
+
self._connection.execute(
|
| 82 |
+
"CREATE TABLE inputs ("
|
| 83 |
+
"position INTEGER PRIMARY KEY, input_id TEXT NOT NULL, sequence TEXT NOT NULL)"
|
| 84 |
+
)
|
| 85 |
+
digest = hashlib.sha256()
|
| 86 |
+
count = 0
|
| 87 |
+
pending: list[tuple[int, str, str]] = []
|
| 88 |
+
try:
|
| 89 |
+
for position, item in enumerate(values):
|
| 90 |
+
record = _normalize_input_item(position, item)
|
| 91 |
+
for value in (record.id, record.sequence):
|
| 92 |
+
encoded = value.encode("utf-8")
|
| 93 |
+
digest.update(len(encoded).to_bytes(8, "big"))
|
| 94 |
+
digest.update(encoded)
|
| 95 |
+
pending.append((position, record.id, record.sequence))
|
| 96 |
+
count += 1
|
| 97 |
+
if len(pending) == 1_024:
|
| 98 |
+
self._connection.executemany("INSERT INTO inputs VALUES (?, ?, ?)", pending)
|
| 99 |
+
pending.clear()
|
| 100 |
+
if pending:
|
| 101 |
+
self._connection.executemany("INSERT INTO inputs VALUES (?, ?, ?)", pending)
|
| 102 |
+
if count == 0:
|
| 103 |
+
raise ValueError("inputs must contain at least one sequence.")
|
| 104 |
+
self._connection.commit()
|
| 105 |
+
self._connection.close()
|
| 106 |
+
self._connection = sqlite3.connect(
|
| 107 |
+
f"{self.path.resolve().as_uri()}?mode=ro",
|
| 108 |
+
uri=True,
|
| 109 |
+
)
|
| 110 |
+
except BaseException:
|
| 111 |
+
self.close()
|
| 112 |
+
raise
|
| 113 |
+
digest.update(count.to_bytes(8, "big"))
|
| 114 |
+
self.input_fingerprint = digest.hexdigest()
|
| 115 |
+
self._count = count
|
| 116 |
+
|
| 117 |
+
def _require_connection(self) -> sqlite3.Connection:
|
| 118 |
+
if self._connection is None:
|
| 119 |
+
raise RuntimeError("Input spool is closed.")
|
| 120 |
+
return self._connection
|
| 121 |
+
|
| 122 |
+
def __len__(self) -> int:
|
| 123 |
+
return self._count
|
| 124 |
+
|
| 125 |
+
def __iter__(self) -> Iterator[EmbeddingInput]:
|
| 126 |
+
cursor = self._require_connection().execute(
|
| 127 |
+
"SELECT input_id, sequence FROM inputs ORDER BY position"
|
| 128 |
+
)
|
| 129 |
+
while rows := cursor.fetchmany(1_024):
|
| 130 |
+
for input_id, sequence in rows:
|
| 131 |
+
yield EmbeddingInput(input_id, sequence)
|
| 132 |
+
|
| 133 |
+
@overload
|
| 134 |
+
def __getitem__(self, index: int, /) -> EmbeddingInput: ...
|
| 135 |
+
|
| 136 |
+
@overload
|
| 137 |
+
def __getitem__(self, index: slice, /) -> list[EmbeddingInput]: ...
|
| 138 |
+
|
| 139 |
+
def __getitem__(self, index: int | slice) -> EmbeddingInput | list[EmbeddingInput]:
|
| 140 |
+
connection = self._require_connection()
|
| 141 |
+
|
| 142 |
+
if isinstance(index, slice):
|
| 143 |
+
start, stop, step = index.indices(self._count)
|
| 144 |
+
if step != 1:
|
| 145 |
+
return [self[position] for position in range(start, stop, step)]
|
| 146 |
+
rows = connection.execute(
|
| 147 |
+
"SELECT input_id, sequence FROM inputs "
|
| 148 |
+
"WHERE position >= ? AND position < ? ORDER BY position",
|
| 149 |
+
(start, stop),
|
| 150 |
+
).fetchall()
|
| 151 |
+
return [EmbeddingInput(input_id, sequence) for input_id, sequence in rows]
|
| 152 |
+
position = index + self._count if index < 0 else index
|
| 153 |
+
if position < 0 or position >= self._count:
|
| 154 |
+
raise IndexError(index)
|
| 155 |
+
row = connection.execute(
|
| 156 |
+
"SELECT input_id, sequence FROM inputs WHERE position = ?", (position,)
|
| 157 |
+
).fetchone()
|
| 158 |
+
if row is None:
|
| 159 |
+
raise IndexError(index)
|
| 160 |
+
return EmbeddingInput(row[0], row[1])
|
| 161 |
+
|
| 162 |
+
def close(self) -> None:
|
| 163 |
+
connection = getattr(self, "_connection", None)
|
| 164 |
+
if connection is not None:
|
| 165 |
+
connection.close()
|
| 166 |
+
self._connection = None
|
| 167 |
+
temporary = getattr(self, "_temporary", None)
|
| 168 |
+
if temporary is not None:
|
| 169 |
+
temporary.cleanup()
|
| 170 |
+
self._temporary = None
|
| 171 |
+
|
| 172 |
+
def __del__(self) -> None:
|
| 173 |
+
self.close()
|
| 174 |
+
|
| 175 |
+
|
| 176 |
+
def _normalize_inputs(
|
| 177 |
+
inputs: (Iterable[str | EmbeddingInput | tuple[str, str]] | Mapping[str, str] | str | Path),
|
| 178 |
+
*,
|
| 179 |
+
disk_backed: bool,
|
| 180 |
+
) -> Sequence[EmbeddingInput]:
|
| 181 |
+
is_fasta_path = isinstance(inputs, Path)
|
| 182 |
+
if isinstance(inputs, str):
|
| 183 |
+
try:
|
| 184 |
+
is_fasta_path = Path(inputs).is_file()
|
| 185 |
+
except OSError:
|
| 186 |
+
is_fasta_path = False
|
| 187 |
+
should_spool = disk_backed or is_fasta_path or not isinstance(inputs, (str, Sequence, Mapping))
|
| 188 |
+
values: Iterable[str | EmbeddingInput | tuple[str, str]]
|
| 189 |
+
if isinstance(inputs, Path):
|
| 190 |
+
values = iter_fasta(inputs)
|
| 191 |
+
elif isinstance(inputs, str):
|
| 192 |
+
values = iter_fasta(inputs) if is_fasta_path else [inputs]
|
| 193 |
+
elif isinstance(inputs, Mapping):
|
| 194 |
+
values = inputs.items()
|
| 195 |
+
else:
|
| 196 |
+
values = inputs
|
| 197 |
+
if should_spool:
|
| 198 |
+
return _InputSpool(values)
|
| 199 |
+
records: list[EmbeddingInput] = []
|
| 200 |
+
for position, item in enumerate(values):
|
| 201 |
+
records.append(_normalize_input_item(position, item))
|
| 202 |
+
if not records:
|
| 203 |
+
raise ValueError("inputs must contain at least one sequence.")
|
| 204 |
+
return records
|
| 205 |
+
|
| 206 |
+
|
| 207 |
+
def _validate_untruncated_lengths(
|
| 208 |
+
records: Sequence[EmbeddingInput],
|
| 209 |
+
*,
|
| 210 |
+
max_length: int | None,
|
| 211 |
+
truncate: bool,
|
| 212 |
+
) -> None:
|
| 213 |
+
"""Fail before inference when a biological-residue limit would be exceeded."""
|
| 214 |
+
|
| 215 |
+
if max_length is None or truncate:
|
| 216 |
+
return
|
| 217 |
+
for position, record in enumerate(records):
|
| 218 |
+
residue_count = len(record.sequence)
|
| 219 |
+
if residue_count > max_length:
|
| 220 |
+
raise ValueError(
|
| 221 |
+
f"Input at position {position} with id {record.id!r} has "
|
| 222 |
+
f"{residue_count} biological residues, exceeding max_length={max_length} "
|
| 223 |
+
"while truncate=False."
|
| 224 |
+
)
|
| 225 |
+
|
| 226 |
+
|
| 227 |
+
def _planned_batches(
|
| 228 |
+
records: Sequence[EmbeddingInput],
|
| 229 |
+
positions: range,
|
| 230 |
+
*,
|
| 231 |
+
batch_size: int,
|
| 232 |
+
max_tokens_per_batch: int | None,
|
| 233 |
+
max_length: int | None,
|
| 234 |
+
truncate: bool,
|
| 235 |
+
) -> Iterator[list[int]]:
|
| 236 |
+
"""Length-bucket one bounded window while retaining stable output positions."""
|
| 237 |
+
|
| 238 |
+
def effective_length(position: int) -> int:
|
| 239 |
+
length = len(records[position].sequence)
|
| 240 |
+
return min(length, max_length) if truncate and max_length is not None else length
|
| 241 |
+
|
| 242 |
+
ordered = sorted(positions, key=lambda position: (-effective_length(position), position))
|
| 243 |
+
batch: list[int] = []
|
| 244 |
+
longest = 0
|
| 245 |
+
for position in ordered:
|
| 246 |
+
length = effective_length(position)
|
| 247 |
+
if max_tokens_per_batch is not None and length > max_tokens_per_batch:
|
| 248 |
+
raise ValueError(
|
| 249 |
+
f"Input at position {position} has {length} residues, exceeding "
|
| 250 |
+
f"max_tokens_per_batch={max_tokens_per_batch}."
|
| 251 |
+
)
|
| 252 |
+
candidate_longest = max(longest, length)
|
| 253 |
+
exceeds_tokens = (
|
| 254 |
+
max_tokens_per_batch is not None
|
| 255 |
+
and candidate_longest * (len(batch) + 1) > max_tokens_per_batch
|
| 256 |
+
)
|
| 257 |
+
if batch and (len(batch) >= batch_size or exceeds_tokens):
|
| 258 |
+
yield batch
|
| 259 |
+
batch = []
|
| 260 |
+
longest = 0
|
| 261 |
+
batch.append(position)
|
| 262 |
+
longest = max(longest, length)
|
| 263 |
+
if batch:
|
| 264 |
+
yield batch
|
fastplms/embeddings/output.py
CHANGED
|
@@ -1,215 +1,215 @@
|
|
| 1 |
-
"""Resume validation and transactional publication of ordered embedding windows."""
|
| 2 |
-
|
| 3 |
-
from __future__ import annotations
|
| 4 |
-
|
| 5 |
-
from collections.abc import Sequence
|
| 6 |
-
from pathlib import Path
|
| 7 |
-
from typing import Any
|
| 8 |
-
|
| 9 |
-
from .identity import _RUN_FINGERPRINT_SCHEMA_VERSION
|
| 10 |
-
from .pooling import Pooler
|
| 11 |
-
from .storage import (
|
| 12 |
-
SafetensorsStreamWriter,
|
| 13 |
-
append_sqlite_records,
|
| 14 |
-
initialize_sqlite_run,
|
| 15 |
-
load_result,
|
| 16 |
-
load_sqlite_result,
|
| 17 |
-
safetensors_result_exists,
|
| 18 |
-
save_result,
|
| 19 |
-
tensor_sha256,
|
| 20 |
-
update_sqlite_run_metadata,
|
| 21 |
-
)
|
| 22 |
-
from .types import EmbeddingInput, EmbeddingRecord, EmbeddingResult, LazyTensorReference
|
| 23 |
-
|
| 24 |
-
|
| 25 |
-
def _output_exists(path: str | Path, format: str) -> bool:
|
| 26 |
-
path = Path(path)
|
| 27 |
-
if format == "sqlite":
|
| 28 |
-
return path.is_file()
|
| 29 |
-
return safetensors_result_exists(path)
|
| 30 |
-
|
| 31 |
-
|
| 32 |
-
def _output_descriptor(position: int, record: EmbeddingRecord) -> dict[str, Any]:
|
| 33 |
-
tensor = record.tensor
|
| 34 |
-
if isinstance(tensor, LazyTensorReference):
|
| 35 |
-
dtype = tensor.dtype
|
| 36 |
-
shape = tensor.shape
|
| 37 |
-
digest = tensor.sha256
|
| 38 |
-
else:
|
| 39 |
-
dtype = str(tensor.dtype).removeprefix("torch.")
|
| 40 |
-
shape = tuple(tensor.shape)
|
| 41 |
-
digest = tensor_sha256(tensor)
|
| 42 |
-
return {
|
| 43 |
-
"position": position,
|
| 44 |
-
"id": record.id,
|
| 45 |
-
"dtype": dtype,
|
| 46 |
-
"shape": shape,
|
| 47 |
-
"sha256": digest,
|
| 48 |
-
}
|
| 49 |
-
|
| 50 |
-
|
| 51 |
-
class EmbeddingOutput:
|
| 52 |
-
"""Own the resumable prefix and the commit state of one output destination."""
|
| 53 |
-
|
| 54 |
-
def __init__(
|
| 55 |
-
self,
|
| 56 |
-
records: Sequence[EmbeddingInput],
|
| 57 |
-
*,
|
| 58 |
-
output: str | Path | None,
|
| 59 |
-
format: str,
|
| 60 |
-
resume: bool,
|
| 61 |
-
shard_size: int,
|
| 62 |
-
run_fingerprint: str,
|
| 63 |
-
input_fingerprint: str,
|
| 64 |
-
model_state_fingerprint: str | None,
|
| 65 |
-
model_state_fingerprint_source: str,
|
| 66 |
-
pooler: Pooler | None,
|
| 67 |
-
pooling_names: Sequence[str],
|
| 68 |
-
) -> None:
|
| 69 |
-
self.output = output
|
| 70 |
-
self.format = format
|
| 71 |
-
self.shard_size = shard_size
|
| 72 |
-
self.completed: EmbeddingResult | None = None
|
| 73 |
-
output_already_exists = output is not None and _output_exists(output, format)
|
| 74 |
-
existing: EmbeddingResult | None = None
|
| 75 |
-
self.start_position = 0
|
| 76 |
-
if output is not None and resume and output_already_exists:
|
| 77 |
-
if format == "sqlite":
|
| 78 |
-
try:
|
| 79 |
-
existing = load_sqlite_result(output, run_id=run_fingerprint)
|
| 80 |
-
except KeyError:
|
| 81 |
-
existing = load_result(output, format=format)
|
| 82 |
-
else:
|
| 83 |
-
existing = load_result(output, format=format)
|
| 84 |
-
if existing.metadata.get("fingerprint_schema_version") != (
|
| 85 |
-
_RUN_FINGERPRINT_SCHEMA_VERSION
|
| 86 |
-
):
|
| 87 |
-
raise ValueError(
|
| 88 |
-
"Existing embeddings use an incompatible run fingerprint schema; "
|
| 89 |
-
"choose another output or set resume=False."
|
| 90 |
-
)
|
| 91 |
-
if existing.metadata.get("run_fingerprint") != run_fingerprint:
|
| 92 |
-
raise ValueError(
|
| 93 |
-
"Existing embeddings were produced by a different run fingerprint; "
|
| 94 |
-
"choose another output or set resume=False."
|
| 95 |
-
)
|
| 96 |
-
if len(existing) > len(records):
|
| 97 |
-
raise ValueError(
|
| 98 |
-
"Existing embeddings are not an ordered prefix of the requested inputs."
|
| 99 |
-
)
|
| 100 |
-
prefix_matches = all(
|
| 101 |
-
(observed.id, observed.sequence) == (expected.id, expected.sequence)
|
| 102 |
-
for expected, observed in zip(records, existing, strict=False)
|
| 103 |
-
)
|
| 104 |
-
if not prefix_matches:
|
| 105 |
-
raise ValueError(
|
| 106 |
-
"Existing embeddings are not an ordered prefix of the requested inputs."
|
| 107 |
-
)
|
| 108 |
-
if len(existing) == len(records) and existing.metadata.get("complete", True):
|
| 109 |
-
self.completed = existing
|
| 110 |
-
return
|
| 111 |
-
self.start_position = len(existing)
|
| 112 |
-
|
| 113 |
-
self.sqlite_run_id: str | None = None
|
| 114 |
-
self.sqlite_replace_on_first_commit = False
|
| 115 |
-
self.sqlite_initial_metadata: dict[str, Any] | None = None
|
| 116 |
-
if output is not None and format == "sqlite":
|
| 117 |
-
self.sqlite_initial_metadata = {
|
| 118 |
-
"format_version": 1,
|
| 119 |
-
"fingerprint_schema_version": _RUN_FINGERPRINT_SCHEMA_VERSION,
|
| 120 |
-
"run_fingerprint": run_fingerprint,
|
| 121 |
-
"input_fingerprint": input_fingerprint,
|
| 122 |
-
"model_state_fingerprint": model_state_fingerprint,
|
| 123 |
-
"model_state_fingerprint_source": model_state_fingerprint_source,
|
| 124 |
-
"complete": False,
|
| 125 |
-
}
|
| 126 |
-
self.sqlite_run_id = run_fingerprint
|
| 127 |
-
if not resume and output_already_exists:
|
| 128 |
-
try:
|
| 129 |
-
load_sqlite_result(output, run_id=run_fingerprint)
|
| 130 |
-
except KeyError:
|
| 131 |
-
pass
|
| 132 |
-
else:
|
| 133 |
-
# Keep an exact prior run readable until replacement inference
|
| 134 |
-
# has produced the first complete commit window.
|
| 135 |
-
self.sqlite_replace_on_first_commit = True
|
| 136 |
-
if not self.sqlite_replace_on_first_commit:
|
| 137 |
-
initialize_sqlite_run(
|
| 138 |
-
output,
|
| 139 |
-
self.sqlite_initial_metadata,
|
| 140 |
-
resume=resume,
|
| 141 |
-
)
|
| 142 |
-
|
| 143 |
-
stream_safetensors = output is not None and format == "safetensors"
|
| 144 |
-
self.output_records: list[EmbeddingRecord] = (
|
| 145 |
-
[] if self.sqlite_run_id is not None or stream_safetensors else list(existing or ())
|
| 146 |
-
)
|
| 147 |
-
self.output_descriptors: list[dict[str, Any]] | None = [] if output is None else None
|
| 148 |
-
self.pool_slices: dict[str, tuple[int, int]] = {}
|
| 149 |
-
if existing and pooler is not None:
|
| 150 |
-
pooled_width = existing[0].load_tensor().shape[-1]
|
| 151 |
-
if pooled_width % len(pooling_names) != 0:
|
| 152 |
-
raise ValueError("Stored pooled width is inconsistent with pooling metadata.")
|
| 153 |
-
self.pool_slices = pooler.output_slices(pooled_width // len(pooling_names))
|
| 154 |
-
|
| 155 |
-
self.safetensors_writer: SafetensorsStreamWriter | None = None
|
| 156 |
-
if stream_safetensors:
|
| 157 |
-
if output is None:
|
| 158 |
-
raise RuntimeError(
|
| 159 |
-
"Safetensors streaming was enabled without an output destination."
|
| 160 |
-
)
|
| 161 |
-
transactional_overwrite = output_already_exists and not resume
|
| 162 |
-
self.safetensors_writer = SafetensorsStreamWriter(
|
| 163 |
-
output,
|
| 164 |
-
{
|
| 165 |
-
"format_version": 1,
|
| 166 |
-
"fingerprint_schema_version": _RUN_FINGERPRINT_SCHEMA_VERSION,
|
| 167 |
-
"run_fingerprint": run_fingerprint,
|
| 168 |
-
"input_fingerprint": input_fingerprint,
|
| 169 |
-
"model_state_fingerprint": model_state_fingerprint,
|
| 170 |
-
"model_state_fingerprint_source": model_state_fingerprint_source,
|
| 171 |
-
"complete": False,
|
| 172 |
-
},
|
| 173 |
-
shard_size=shard_size,
|
| 174 |
-
existing=existing or (),
|
| 175 |
-
reuse_existing=bool(resume and existing is not None),
|
| 176 |
-
publish_initial=not transactional_overwrite,
|
| 177 |
-
publish_incremental=not transactional_overwrite,
|
| 178 |
-
)
|
| 179 |
-
|
| 180 |
-
def append(self, window_start: int, new_records: list[EmbeddingRecord]) -> None:
|
| 181 |
-
"""Commit a complete ordered window at the storage format's granularity."""
|
| 182 |
-
|
| 183 |
-
if self.output_descriptors is not None:
|
| 184 |
-
self.output_descriptors.extend(
|
| 185 |
-
_output_descriptor(window_start + offset, record)
|
| 186 |
-
for offset, record in enumerate(new_records)
|
| 187 |
-
)
|
| 188 |
-
if self.output is not None and self.sqlite_run_id is not None:
|
| 189 |
-
append_sqlite_records(
|
| 190 |
-
self.output,
|
| 191 |
-
self.sqlite_run_id,
|
| 192 |
-
window_start,
|
| 193 |
-
new_records,
|
| 194 |
-
replace_metadata=(
|
| 195 |
-
self.sqlite_initial_metadata if self.sqlite_replace_on_first_commit else None
|
| 196 |
-
),
|
| 197 |
-
)
|
| 198 |
-
self.sqlite_replace_on_first_commit = False
|
| 199 |
-
elif self.safetensors_writer is not None:
|
| 200 |
-
self.safetensors_writer.append(new_records)
|
| 201 |
-
else:
|
| 202 |
-
self.output_records.extend(new_records)
|
| 203 |
-
|
| 204 |
-
def finish(self, metadata: dict[str, Any]) -> EmbeddingResult:
|
| 205 |
-
"""Publish completion only after every window has committed."""
|
| 206 |
-
|
| 207 |
-
if self.output is not None and self.sqlite_run_id is not None:
|
| 208 |
-
update_sqlite_run_metadata(self.output, self.sqlite_run_id, metadata)
|
| 209 |
-
return load_sqlite_result(self.output, run_id=self.sqlite_run_id)
|
| 210 |
-
if self.safetensors_writer is not None:
|
| 211 |
-
return self.safetensors_writer.publish(complete=True, metadata=metadata)
|
| 212 |
-
result = EmbeddingResult(self.output_records, metadata)
|
| 213 |
-
if self.output is not None:
|
| 214 |
-
return save_result(result, self.output, format=self.format, shard_size=self.shard_size)
|
| 215 |
-
return result
|
|
|
|
| 1 |
+
"""Resume validation and transactional publication of ordered embedding windows."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
from collections.abc import Sequence
|
| 6 |
+
from pathlib import Path
|
| 7 |
+
from typing import Any
|
| 8 |
+
|
| 9 |
+
from .identity import _RUN_FINGERPRINT_SCHEMA_VERSION
|
| 10 |
+
from .pooling import Pooler
|
| 11 |
+
from .storage import (
|
| 12 |
+
SafetensorsStreamWriter,
|
| 13 |
+
append_sqlite_records,
|
| 14 |
+
initialize_sqlite_run,
|
| 15 |
+
load_result,
|
| 16 |
+
load_sqlite_result,
|
| 17 |
+
safetensors_result_exists,
|
| 18 |
+
save_result,
|
| 19 |
+
tensor_sha256,
|
| 20 |
+
update_sqlite_run_metadata,
|
| 21 |
+
)
|
| 22 |
+
from .types import EmbeddingInput, EmbeddingRecord, EmbeddingResult, LazyTensorReference
|
| 23 |
+
|
| 24 |
+
|
| 25 |
+
def _output_exists(path: str | Path, format: str) -> bool:
|
| 26 |
+
path = Path(path)
|
| 27 |
+
if format == "sqlite":
|
| 28 |
+
return path.is_file()
|
| 29 |
+
return safetensors_result_exists(path)
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
def _output_descriptor(position: int, record: EmbeddingRecord) -> dict[str, Any]:
|
| 33 |
+
tensor = record.tensor
|
| 34 |
+
if isinstance(tensor, LazyTensorReference):
|
| 35 |
+
dtype = tensor.dtype
|
| 36 |
+
shape = tensor.shape
|
| 37 |
+
digest = tensor.sha256
|
| 38 |
+
else:
|
| 39 |
+
dtype = str(tensor.dtype).removeprefix("torch.")
|
| 40 |
+
shape = tuple(tensor.shape)
|
| 41 |
+
digest = tensor_sha256(tensor)
|
| 42 |
+
return {
|
| 43 |
+
"position": position,
|
| 44 |
+
"id": record.id,
|
| 45 |
+
"dtype": dtype,
|
| 46 |
+
"shape": shape,
|
| 47 |
+
"sha256": digest,
|
| 48 |
+
}
|
| 49 |
+
|
| 50 |
+
|
| 51 |
+
class EmbeddingOutput:
|
| 52 |
+
"""Own the resumable prefix and the commit state of one output destination."""
|
| 53 |
+
|
| 54 |
+
def __init__(
|
| 55 |
+
self,
|
| 56 |
+
records: Sequence[EmbeddingInput],
|
| 57 |
+
*,
|
| 58 |
+
output: str | Path | None,
|
| 59 |
+
format: str,
|
| 60 |
+
resume: bool,
|
| 61 |
+
shard_size: int,
|
| 62 |
+
run_fingerprint: str,
|
| 63 |
+
input_fingerprint: str,
|
| 64 |
+
model_state_fingerprint: str | None,
|
| 65 |
+
model_state_fingerprint_source: str,
|
| 66 |
+
pooler: Pooler | None,
|
| 67 |
+
pooling_names: Sequence[str],
|
| 68 |
+
) -> None:
|
| 69 |
+
self.output = output
|
| 70 |
+
self.format = format
|
| 71 |
+
self.shard_size = shard_size
|
| 72 |
+
self.completed: EmbeddingResult | None = None
|
| 73 |
+
output_already_exists = output is not None and _output_exists(output, format)
|
| 74 |
+
existing: EmbeddingResult | None = None
|
| 75 |
+
self.start_position = 0
|
| 76 |
+
if output is not None and resume and output_already_exists:
|
| 77 |
+
if format == "sqlite":
|
| 78 |
+
try:
|
| 79 |
+
existing = load_sqlite_result(output, run_id=run_fingerprint)
|
| 80 |
+
except KeyError:
|
| 81 |
+
existing = load_result(output, format=format)
|
| 82 |
+
else:
|
| 83 |
+
existing = load_result(output, format=format)
|
| 84 |
+
if existing.metadata.get("fingerprint_schema_version") != (
|
| 85 |
+
_RUN_FINGERPRINT_SCHEMA_VERSION
|
| 86 |
+
):
|
| 87 |
+
raise ValueError(
|
| 88 |
+
"Existing embeddings use an incompatible run fingerprint schema; "
|
| 89 |
+
"choose another output or set resume=False."
|
| 90 |
+
)
|
| 91 |
+
if existing.metadata.get("run_fingerprint") != run_fingerprint:
|
| 92 |
+
raise ValueError(
|
| 93 |
+
"Existing embeddings were produced by a different run fingerprint; "
|
| 94 |
+
"choose another output or set resume=False."
|
| 95 |
+
)
|
| 96 |
+
if len(existing) > len(records):
|
| 97 |
+
raise ValueError(
|
| 98 |
+
"Existing embeddings are not an ordered prefix of the requested inputs."
|
| 99 |
+
)
|
| 100 |
+
prefix_matches = all(
|
| 101 |
+
(observed.id, observed.sequence) == (expected.id, expected.sequence)
|
| 102 |
+
for expected, observed in zip(records, existing, strict=False)
|
| 103 |
+
)
|
| 104 |
+
if not prefix_matches:
|
| 105 |
+
raise ValueError(
|
| 106 |
+
"Existing embeddings are not an ordered prefix of the requested inputs."
|
| 107 |
+
)
|
| 108 |
+
if len(existing) == len(records) and existing.metadata.get("complete", True):
|
| 109 |
+
self.completed = existing
|
| 110 |
+
return
|
| 111 |
+
self.start_position = len(existing)
|
| 112 |
+
|
| 113 |
+
self.sqlite_run_id: str | None = None
|
| 114 |
+
self.sqlite_replace_on_first_commit = False
|
| 115 |
+
self.sqlite_initial_metadata: dict[str, Any] | None = None
|
| 116 |
+
if output is not None and format == "sqlite":
|
| 117 |
+
self.sqlite_initial_metadata = {
|
| 118 |
+
"format_version": 1,
|
| 119 |
+
"fingerprint_schema_version": _RUN_FINGERPRINT_SCHEMA_VERSION,
|
| 120 |
+
"run_fingerprint": run_fingerprint,
|
| 121 |
+
"input_fingerprint": input_fingerprint,
|
| 122 |
+
"model_state_fingerprint": model_state_fingerprint,
|
| 123 |
+
"model_state_fingerprint_source": model_state_fingerprint_source,
|
| 124 |
+
"complete": False,
|
| 125 |
+
}
|
| 126 |
+
self.sqlite_run_id = run_fingerprint
|
| 127 |
+
if not resume and output_already_exists:
|
| 128 |
+
try:
|
| 129 |
+
load_sqlite_result(output, run_id=run_fingerprint)
|
| 130 |
+
except KeyError:
|
| 131 |
+
pass
|
| 132 |
+
else:
|
| 133 |
+
# Keep an exact prior run readable until replacement inference
|
| 134 |
+
# has produced the first complete commit window.
|
| 135 |
+
self.sqlite_replace_on_first_commit = True
|
| 136 |
+
if not self.sqlite_replace_on_first_commit:
|
| 137 |
+
initialize_sqlite_run(
|
| 138 |
+
output,
|
| 139 |
+
self.sqlite_initial_metadata,
|
| 140 |
+
resume=resume,
|
| 141 |
+
)
|
| 142 |
+
|
| 143 |
+
stream_safetensors = output is not None and format == "safetensors"
|
| 144 |
+
self.output_records: list[EmbeddingRecord] = (
|
| 145 |
+
[] if self.sqlite_run_id is not None or stream_safetensors else list(existing or ())
|
| 146 |
+
)
|
| 147 |
+
self.output_descriptors: list[dict[str, Any]] | None = [] if output is None else None
|
| 148 |
+
self.pool_slices: dict[str, tuple[int, int]] = {}
|
| 149 |
+
if existing and pooler is not None:
|
| 150 |
+
pooled_width = existing[0].load_tensor().shape[-1]
|
| 151 |
+
if pooled_width % len(pooling_names) != 0:
|
| 152 |
+
raise ValueError("Stored pooled width is inconsistent with pooling metadata.")
|
| 153 |
+
self.pool_slices = pooler.output_slices(pooled_width // len(pooling_names))
|
| 154 |
+
|
| 155 |
+
self.safetensors_writer: SafetensorsStreamWriter | None = None
|
| 156 |
+
if stream_safetensors:
|
| 157 |
+
if output is None:
|
| 158 |
+
raise RuntimeError(
|
| 159 |
+
"Safetensors streaming was enabled without an output destination."
|
| 160 |
+
)
|
| 161 |
+
transactional_overwrite = output_already_exists and not resume
|
| 162 |
+
self.safetensors_writer = SafetensorsStreamWriter(
|
| 163 |
+
output,
|
| 164 |
+
{
|
| 165 |
+
"format_version": 1,
|
| 166 |
+
"fingerprint_schema_version": _RUN_FINGERPRINT_SCHEMA_VERSION,
|
| 167 |
+
"run_fingerprint": run_fingerprint,
|
| 168 |
+
"input_fingerprint": input_fingerprint,
|
| 169 |
+
"model_state_fingerprint": model_state_fingerprint,
|
| 170 |
+
"model_state_fingerprint_source": model_state_fingerprint_source,
|
| 171 |
+
"complete": False,
|
| 172 |
+
},
|
| 173 |
+
shard_size=shard_size,
|
| 174 |
+
existing=existing or (),
|
| 175 |
+
reuse_existing=bool(resume and existing is not None),
|
| 176 |
+
publish_initial=not transactional_overwrite,
|
| 177 |
+
publish_incremental=not transactional_overwrite,
|
| 178 |
+
)
|
| 179 |
+
|
| 180 |
+
def append(self, window_start: int, new_records: list[EmbeddingRecord]) -> None:
|
| 181 |
+
"""Commit a complete ordered window at the storage format's granularity."""
|
| 182 |
+
|
| 183 |
+
if self.output_descriptors is not None:
|
| 184 |
+
self.output_descriptors.extend(
|
| 185 |
+
_output_descriptor(window_start + offset, record)
|
| 186 |
+
for offset, record in enumerate(new_records)
|
| 187 |
+
)
|
| 188 |
+
if self.output is not None and self.sqlite_run_id is not None:
|
| 189 |
+
append_sqlite_records(
|
| 190 |
+
self.output,
|
| 191 |
+
self.sqlite_run_id,
|
| 192 |
+
window_start,
|
| 193 |
+
new_records,
|
| 194 |
+
replace_metadata=(
|
| 195 |
+
self.sqlite_initial_metadata if self.sqlite_replace_on_first_commit else None
|
| 196 |
+
),
|
| 197 |
+
)
|
| 198 |
+
self.sqlite_replace_on_first_commit = False
|
| 199 |
+
elif self.safetensors_writer is not None:
|
| 200 |
+
self.safetensors_writer.append(new_records)
|
| 201 |
+
else:
|
| 202 |
+
self.output_records.extend(new_records)
|
| 203 |
+
|
| 204 |
+
def finish(self, metadata: dict[str, Any]) -> EmbeddingResult:
|
| 205 |
+
"""Publish completion only after every window has committed."""
|
| 206 |
+
|
| 207 |
+
if self.output is not None and self.sqlite_run_id is not None:
|
| 208 |
+
update_sqlite_run_metadata(self.output, self.sqlite_run_id, metadata)
|
| 209 |
+
return load_sqlite_result(self.output, run_id=self.sqlite_run_id)
|
| 210 |
+
if self.safetensors_writer is not None:
|
| 211 |
+
return self.safetensors_writer.publish(complete=True, metadata=metadata)
|
| 212 |
+
result = EmbeddingResult(self.output_records, metadata)
|
| 213 |
+
if self.output is not None:
|
| 214 |
+
return save_result(result, self.output, format=self.format, shard_size=self.shard_size)
|
| 215 |
+
return result
|
fastplms/embeddings/pooling.py
CHANGED
|
@@ -4,6 +4,7 @@ from __future__ import annotations
|
|
| 4 |
|
| 5 |
import math
|
| 6 |
import torch
|
|
|
|
| 7 |
from collections.abc import Sequence
|
| 8 |
from torch import Tensor
|
| 9 |
|
|
|
|
| 4 |
|
| 5 |
import math
|
| 6 |
import torch
|
| 7 |
+
|
| 8 |
from collections.abc import Sequence
|
| 9 |
from torch import Tensor
|
| 10 |
|
fastplms/embeddings/runner.py
CHANGED
|
@@ -1,425 +1,426 @@
|
|
| 1 |
-
"""Coordinate input preparation, run identity, batch execution, and publication."""
|
| 2 |
-
|
| 3 |
-
from __future__ import annotations
|
| 4 |
-
|
| 5 |
-
import torch
|
| 6 |
-
|
| 7 |
-
from
|
| 8 |
-
from
|
|
|
|
| 9 |
from torch import Tensor
|
| 10 |
|
| 11 |
from . import identity
|
| 12 |
from .batches import (
|
| 13 |
-
BatchExecutor,
|
| 14 |
-
_residue_embeddings as _residue_embeddings,
|
| 15 |
-
_temporary_eval,
|
| 16 |
-
select_hidden_state_embeddings as select_hidden_state_embeddings,
|
| 17 |
-
)
|
| 18 |
-
from .identity import (
|
| 19 |
-
_RUN_FINGERPRINT_SCHEMA_VERSION,
|
| 20 |
-
_adapter_identity_metadata,
|
| 21 |
-
_attention_backend,
|
| 22 |
-
_attention_kernel_metadata,
|
| 23 |
-
_embedding_context,
|
| 24 |
-
_execution_identity_metadata,
|
| 25 |
-
_fingerprint_jsonable,
|
| 26 |
-
_model_identity_metadata,
|
| 27 |
-
_run_fingerprint,
|
| 28 |
_tokenizer_metadata,
|
| 29 |
-
)
|
| 30 |
-
from .inputs import (
|
| 31 |
-
_InputSpool,
|
| 32 |
-
_normalize_inputs,
|
| 33 |
-
_validate_untruncated_lengths,
|
| 34 |
-
iter_fasta as iter_fasta,
|
| 35 |
-
parse_fasta as parse_fasta,
|
| 36 |
-
)
|
| 37 |
-
from .output import EmbeddingOutput
|
| 38 |
-
from .pooling import Pooler
|
| 39 |
-
from .types import EmbeddingBatch, EmbeddingInput, EmbeddingResult
|
| 40 |
-
|
| 41 |
-
|
| 42 |
-
_DEFAULT_BATCH_WINDOW_MULTIPLIER = 16
|
| 43 |
-
_SUPPORTED_STORAGE_FORMATS = frozenset({"safetensors", "sqlite"})
|
| 44 |
-
|
| 45 |
-
|
| 46 |
-
def embed_dataset(
|
| 47 |
-
model: Any,
|
| 48 |
-
inputs: (Iterable[str | EmbeddingInput | tuple[str, str]] | Mapping[str, str] | str | Path),
|
| 49 |
-
*,
|
| 50 |
-
batch_size: int = 2,
|
| 51 |
-
pooling: str | Sequence[str] | None = None,
|
| 52 |
-
full_embeddings: bool = False,
|
| 53 |
-
output: str | Path | None = None,
|
| 54 |
-
format: str = "safetensors",
|
| 55 |
-
resume: bool = True,
|
| 56 |
-
tokenizer: Any | None = None,
|
| 57 |
-
max_length: int | None = None,
|
| 58 |
-
truncate: bool = True,
|
| 59 |
-
dtype: torch.dtype | None = torch.float32,
|
| 60 |
-
shard_size: int = 2 * 1024**3,
|
| 61 |
-
model_state_fingerprint: str | None = None,
|
| 62 |
-
batch_window_size: int | None = None,
|
| 63 |
-
max_tokens_per_batch: int | None = None,
|
| 64 |
-
hidden_state_source: str = "encoder",
|
| 65 |
-
decoder_inputs: Sequence[str] | None = None,
|
| 66 |
-
decoder_input_ids: Tensor | None = None,
|
| 67 |
-
decoder_attention_mask: Tensor | None = None,
|
| 68 |
-
_embedding_batch_fn: Callable[..., EmbeddingBatch] | None = None,
|
| 69 |
-
_embedding_batch_identity: Mapping[str, Any] | None = None,
|
| 70 |
-
_allowed_unsupported_pooling: Sequence[str] = (),
|
| 71 |
-
**model_kwargs: Any,
|
| 72 |
-
) -> EmbeddingResult:
|
| 73 |
-
"""Embed protein sequences with stable ordering and residue-only pooling."""
|
| 74 |
-
|
| 75 |
-
for name, value in (
|
| 76 |
-
("batch_size", batch_size),
|
| 77 |
-
("shard_size", shard_size),
|
| 78 |
-
):
|
| 79 |
-
if not isinstance(value, int) or isinstance(value, bool):
|
| 80 |
-
raise TypeError(f"{name} must be a positive integer.")
|
| 81 |
-
if value <= 0:
|
| 82 |
-
raise ValueError(f"{name} must be a positive integer.")
|
| 83 |
-
for optional_name, optional_value in (
|
| 84 |
-
("max_length", max_length),
|
| 85 |
-
("max_tokens_per_batch", max_tokens_per_batch),
|
| 86 |
-
("batch_window_size", batch_window_size),
|
| 87 |
-
):
|
| 88 |
-
if optional_value is not None and (
|
| 89 |
-
not isinstance(optional_value, int) or isinstance(optional_value, bool)
|
| 90 |
-
):
|
| 91 |
-
raise TypeError(f"{optional_name} must be a positive integer when provided.")
|
| 92 |
-
if optional_value is not None and optional_value <= 0:
|
| 93 |
-
raise ValueError(f"{optional_name} must be a positive integer when provided.")
|
| 94 |
-
for name, value in (
|
| 95 |
-
("full_embeddings", full_embeddings),
|
| 96 |
-
("resume", resume),
|
| 97 |
-
("truncate", truncate),
|
| 98 |
-
):
|
| 99 |
-
if not isinstance(value, bool):
|
| 100 |
-
raise TypeError(f"{name} must be a boolean.")
|
| 101 |
-
if not isinstance(format, str):
|
| 102 |
-
raise TypeError("format must be a string.")
|
| 103 |
-
if output is not None and not isinstance(output, (str, Path)):
|
| 104 |
-
raise TypeError("output must be a path or None.")
|
| 105 |
-
if model_state_fingerprint is not None and (
|
| 106 |
-
not isinstance(model_state_fingerprint, str) or not model_state_fingerprint
|
| 107 |
-
):
|
| 108 |
-
raise ValueError("model_state_fingerprint must be a non-empty string when provided.")
|
| 109 |
-
if hidden_state_source not in {"encoder", "decoder"}:
|
| 110 |
-
raise ValueError("hidden_state_source must be 'encoder' or 'decoder'.")
|
| 111 |
-
hidden_state_index = model_kwargs.get("hidden_state_index", -1)
|
| 112 |
-
if not isinstance(hidden_state_index, int) or isinstance(hidden_state_index, bool):
|
| 113 |
-
raise TypeError("hidden_state_index must be an integer.")
|
| 114 |
-
store_all_hidden_states = model_kwargs.get("store_all_hidden_states", False)
|
| 115 |
-
if not isinstance(store_all_hidden_states, bool):
|
| 116 |
-
raise TypeError("store_all_hidden_states must be a boolean.")
|
| 117 |
-
if decoder_input_ids is not None:
|
| 118 |
-
if not isinstance(decoder_input_ids, Tensor):
|
| 119 |
-
raise TypeError("decoder_input_ids must be a tensor.")
|
| 120 |
-
if decoder_input_ids.is_meta:
|
| 121 |
-
raise ValueError("decoder_input_ids cannot be a meta tensor.")
|
| 122 |
-
if decoder_input_ids.ndim != 2 or decoder_input_ids.shape[1] == 0:
|
| 123 |
-
raise ValueError("decoder_input_ids must have non-empty shape (batch, sequence).")
|
| 124 |
-
if decoder_input_ids.dtype not in {torch.int32, torch.int64}:
|
| 125 |
-
raise TypeError("decoder_input_ids must use torch.int32 or torch.int64.")
|
| 126 |
-
if decoder_attention_mask is not None:
|
| 127 |
-
if not isinstance(decoder_attention_mask, Tensor):
|
| 128 |
-
raise TypeError("decoder_attention_mask must be a tensor.")
|
| 129 |
-
if decoder_attention_mask.is_meta:
|
| 130 |
-
raise ValueError("decoder_attention_mask cannot be a meta tensor.")
|
| 131 |
-
if decoder_attention_mask.is_complex() or not bool(
|
| 132 |
-
torch.isfinite(decoder_attention_mask).all()
|
| 133 |
-
):
|
| 134 |
-
raise ValueError("decoder_attention_mask must contain finite binary values.")
|
| 135 |
-
if not bool(((decoder_attention_mask == 0) | (decoder_attention_mask == 1)).all()):
|
| 136 |
-
raise ValueError("decoder_attention_mask must contain finite binary values.")
|
| 137 |
-
pooling_names = (
|
| 138 |
-
(("mean",) if not full_embeddings else ())
|
| 139 |
-
if pooling is None
|
| 140 |
-
else ((pooling,) if isinstance(pooling, str) else tuple(pooling))
|
| 141 |
-
)
|
| 142 |
-
if full_embeddings and pooling is not None:
|
| 143 |
-
raise ValueError("full_embeddings=True cannot be combined with pooling.")
|
| 144 |
-
if not full_embeddings and not pooling_names:
|
| 145 |
-
raise ValueError("pooling is required unless full_embeddings=True.")
|
| 146 |
-
pooler = Pooler(pooling_names) if pooling_names else None
|
| 147 |
-
|
| 148 |
-
if batch_size <= 0:
|
| 149 |
-
raise ValueError("batch_size must be positive.")
|
| 150 |
-
if format == "pth" or (output is not None and Path(output).suffix.lower() == ".pth"):
|
| 151 |
-
raise ValueError("Writing pickle-based .pth embeddings is not supported.")
|
| 152 |
-
if format not in _SUPPORTED_STORAGE_FORMATS:
|
| 153 |
-
raise ValueError("format must be 'safetensors' or 'sqlite'.")
|
| 154 |
-
if max_length is not None and max_length <= 0:
|
| 155 |
-
raise ValueError("max_length must be positive when provided.")
|
| 156 |
-
if max_tokens_per_batch is not None and max_tokens_per_batch <= 0:
|
| 157 |
-
raise ValueError("max_tokens_per_batch must be positive when provided.")
|
| 158 |
-
if not isinstance(dtype, (torch.dtype, type(None))):
|
| 159 |
-
raise TypeError("dtype must be a torch.dtype or None.")
|
| 160 |
-
if batch_window_size is not None and batch_window_size <= 0:
|
| 161 |
-
raise ValueError("batch_window_size must be positive when provided.")
|
| 162 |
-
if _embedding_batch_fn is not None and not callable(_embedding_batch_fn):
|
| 163 |
-
raise TypeError("_embedding_batch_fn must be callable when provided.")
|
| 164 |
-
if _embedding_batch_fn is not None and _embedding_batch_identity is None:
|
| 165 |
-
raise ValueError(
|
| 166 |
-
"_embedding_batch_identity is required with _embedding_batch_fn so persisted "
|
| 167 |
-
"runs bind the family-specific embedding behavior."
|
| 168 |
-
)
|
| 169 |
-
if _embedding_batch_identity is not None and not isinstance(_embedding_batch_identity, Mapping):
|
| 170 |
-
raise TypeError("_embedding_batch_identity must be a mapping when provided.")
|
| 171 |
-
if isinstance(_allowed_unsupported_pooling, (str, bytes)) or not isinstance(
|
| 172 |
-
_allowed_unsupported_pooling, Sequence
|
| 173 |
-
):
|
| 174 |
-
raise TypeError("_allowed_unsupported_pooling must be a sequence of pooler names.")
|
| 175 |
-
if not all(isinstance(name, str) for name in _allowed_unsupported_pooling):
|
| 176 |
-
raise TypeError("_allowed_unsupported_pooling must contain only strings.")
|
| 177 |
-
allowed_unsupported_pooling = frozenset(_allowed_unsupported_pooling)
|
| 178 |
-
if allowed_unsupported_pooling and _embedding_batch_fn is None:
|
| 179 |
-
raise ValueError(
|
| 180 |
-
"_allowed_unsupported_pooling is only valid with a family-specific _embedding_batch_fn."
|
| 181 |
-
)
|
| 182 |
-
resolved_batch_window_size = (
|
| 183 |
-
batch_size * _DEFAULT_BATCH_WINDOW_MULTIPLIER
|
| 184 |
-
if batch_window_size is None
|
| 185 |
-
else batch_window_size
|
| 186 |
-
)
|
| 187 |
-
if resolved_batch_window_size < batch_size:
|
| 188 |
-
raise ValueError("batch_window_size must be at least batch_size.")
|
| 189 |
-
records = _normalize_inputs(inputs, disk_backed=output is not None)
|
| 190 |
-
_validate_untruncated_lengths(
|
| 191 |
-
records,
|
| 192 |
-
max_length=max_length,
|
| 193 |
-
truncate=truncate,
|
| 194 |
-
)
|
| 195 |
-
pooling_names = (
|
| 196 |
-
(("mean",) if not full_embeddings else ())
|
| 197 |
-
if pooling is None
|
| 198 |
-
else ((pooling,) if isinstance(pooling, str) else tuple(pooling))
|
| 199 |
-
)
|
| 200 |
-
if full_embeddings:
|
| 201 |
-
if pooling is not None:
|
| 202 |
-
raise ValueError("full_embeddings=True cannot be combined with pooling.")
|
| 203 |
-
elif not pooling_names:
|
| 204 |
-
raise ValueError("pooling is required unless full_embeddings=True.")
|
| 205 |
-
store_all_hidden_states = bool(model_kwargs.get("store_all_hidden_states", False))
|
| 206 |
-
if store_all_hidden_states and not full_embeddings:
|
| 207 |
-
raise ValueError("store_all_hidden_states=True requires full_embeddings=True.")
|
| 208 |
-
|
| 209 |
-
unsupported = set(getattr(model, "embedding_unsupported_pooling", ()))
|
| 210 |
-
unknown_pooling_overrides = allowed_unsupported_pooling.difference(unsupported)
|
| 211 |
-
if unknown_pooling_overrides:
|
| 212 |
-
raise ValueError(
|
| 213 |
-
"_allowed_unsupported_pooling may only override poolers declared unsupported "
|
| 214 |
-
f"by the model; unknown overrides: {sorted(unknown_pooling_overrides)}."
|
| 215 |
-
)
|
| 216 |
-
unsupported.difference_update(allowed_unsupported_pooling)
|
| 217 |
-
requested_unsupported = unsupported.intersection(pooling_names)
|
| 218 |
-
if requested_unsupported:
|
| 219 |
-
raise ValueError(
|
| 220 |
-
f"{model.__class__.__name__} does not support pooling operations "
|
| 221 |
-
f"{sorted(requested_unsupported)}."
|
| 222 |
-
)
|
| 223 |
-
|
| 224 |
-
# Constructing the pooler validates names and duplicate operations before
|
| 225 |
-
# any checkpoint hashing, tokenization, or inference occurs.
|
| 226 |
-
pooler = Pooler(pooling_names) if pooling_names else None
|
| 227 |
-
embedding_context, normalized_decoder_inputs = _embedding_context(
|
| 228 |
-
model,
|
| 229 |
-
records,
|
| 230 |
-
hidden_state_source=hidden_state_source,
|
| 231 |
-
decoder_inputs=decoder_inputs,
|
| 232 |
-
decoder_input_ids=decoder_input_ids,
|
| 233 |
-
decoder_attention_mask=decoder_attention_mask,
|
| 234 |
-
model_kwargs=model_kwargs,
|
| 235 |
-
)
|
| 236 |
-
if _embedding_batch_identity is not None:
|
| 237 |
-
embedding_context["family_adapter"] = _fingerprint_jsonable(_embedding_batch_identity)
|
| 238 |
-
if allowed_unsupported_pooling:
|
| 239 |
-
embedding_context["family_adapter_pooling_override"] = sorted(
|
| 240 |
-
allowed_unsupported_pooling
|
| 241 |
-
)
|
| 242 |
-
|
| 243 |
-
# A pending automatic attention request settles here, inside the caller's
|
| 244 |
-
# autocast context, so the fingerprint records the backend that executes.
|
| 245 |
-
attention_resolution = getattr(model, "attention_resolution", None)
|
| 246 |
-
if attention_resolution is not None and attention_resolution.deferred:
|
| 247 |
-
model.resolve_attn_implementation()
|
| 248 |
-
|
| 249 |
-
tokenizer_metadata = _tokenizer_metadata(model, tokenizer)
|
| 250 |
-
(
|
| 251 |
-
input_fingerprint,
|
| 252 |
-
run_fingerprint,
|
| 253 |
-
resolved_model_state_fingerprint,
|
| 254 |
-
model_state_fingerprint_source,
|
| 255 |
-
) = _run_fingerprint(
|
| 256 |
-
model,
|
| 257 |
-
records,
|
| 258 |
-
pooling=pooling_names,
|
| 259 |
-
full_embeddings=full_embeddings,
|
| 260 |
-
max_length=max_length,
|
| 261 |
-
truncate=truncate,
|
| 262 |
-
dtype=dtype,
|
| 263 |
-
model_kwargs=model_kwargs,
|
| 264 |
-
tokenizer_metadata=tokenizer_metadata,
|
| 265 |
-
model_state_fingerprint=model_state_fingerprint,
|
| 266 |
-
persist_output=output is not None,
|
| 267 |
-
embedding_context=embedding_context,
|
| 268 |
-
batch_size=batch_size,
|
| 269 |
-
batch_window_size=resolved_batch_window_size,
|
| 270 |
-
max_tokens_per_batch=max_tokens_per_batch,
|
| 271 |
-
)
|
| 272 |
-
destination = EmbeddingOutput(
|
| 273 |
-
records,
|
| 274 |
-
output=output,
|
| 275 |
-
format=format,
|
| 276 |
-
resume=resume,
|
| 277 |
-
shard_size=shard_size,
|
| 278 |
-
run_fingerprint=run_fingerprint,
|
| 279 |
-
input_fingerprint=input_fingerprint,
|
| 280 |
-
model_state_fingerprint=resolved_model_state_fingerprint,
|
| 281 |
-
model_state_fingerprint_source=model_state_fingerprint_source,
|
| 282 |
-
pooler=pooler,
|
| 283 |
-
pooling_names=pooling_names,
|
| 284 |
-
)
|
| 285 |
-
if destination.completed is not None:
|
| 286 |
-
return destination.completed
|
| 287 |
-
|
| 288 |
-
attention_backend = _attention_backend(model)
|
| 289 |
-
executor = BatchExecutor(
|
| 290 |
-
model=model,
|
| 291 |
-
batch_size=batch_size,
|
| 292 |
-
max_tokens_per_batch=max_tokens_per_batch,
|
| 293 |
-
max_length=max_length,
|
| 294 |
-
truncate=truncate,
|
| 295 |
-
model_kwargs=model_kwargs,
|
| 296 |
-
hidden_state_source=hidden_state_source,
|
| 297 |
-
normalized_decoder_inputs=normalized_decoder_inputs,
|
| 298 |
-
decoder_input_ids=decoder_input_ids,
|
| 299 |
-
decoder_attention_mask=decoder_attention_mask,
|
| 300 |
-
_embedding_batch_fn=_embedding_batch_fn,
|
| 301 |
-
tokenizer=tokenizer,
|
| 302 |
-
store_all_hidden_states=store_all_hidden_states,
|
| 303 |
-
full_embeddings=full_embeddings,
|
| 304 |
-
dtype=dtype,
|
| 305 |
-
pooler=pooler,
|
| 306 |
-
attention_backend=attention_backend,
|
| 307 |
-
need_attentions="parti" in pooling_names,
|
| 308 |
-
)
|
| 309 |
-
pool_slices = destination.pool_slices
|
| 310 |
-
with _temporary_eval(model), torch.inference_mode():
|
| 311 |
-
for window_start in range(
|
| 312 |
-
destination.start_position, len(records), resolved_batch_window_size
|
| 313 |
-
):
|
| 314 |
-
window_stop = min(window_start + resolved_batch_window_size, len(records))
|
| 315 |
-
window_records = records[window_start:window_stop]
|
| 316 |
-
if not isinstance(window_records, Sequence):
|
| 317 |
-
raise RuntimeError("The immutable embedding spool returned a non-sequence window.")
|
| 318 |
-
new_records, pool_slices = executor.run_window(
|
| 319 |
-
window_records, window_start=window_start
|
| 320 |
-
)
|
| 321 |
-
destination.append(window_start, new_records)
|
| 322 |
-
|
| 323 |
software_versions = identity._software_versions()
|
| 324 |
-
projection = getattr(model, "embedding_projection", None)
|
| 325 |
-
resolved_layer = getattr(
|
| 326 |
-
model,
|
| 327 |
-
"embedding_layer",
|
| 328 |
-
model_kwargs.get("hidden_state_index", -1),
|
| 329 |
-
)
|
| 330 |
-
token_policy = getattr(
|
| 331 |
-
model,
|
| 332 |
-
"embedding_token_policy",
|
| 333 |
-
{
|
| 334 |
-
"unit": "residue",
|
| 335 |
-
"include": ["biological residues"],
|
| 336 |
-
"exclude": [
|
| 337 |
-
"BOS",
|
| 338 |
-
"EOS",
|
| 339 |
-
"padding",
|
| 340 |
-
"chain delimiters",
|
| 341 |
-
"non-protein tokens",
|
| 342 |
-
],
|
| 343 |
-
},
|
| 344 |
-
)
|
| 345 |
-
model_identity = _model_identity_metadata(model)
|
| 346 |
-
metadata: dict[str, Any] = {
|
| 347 |
-
"format_version": 1,
|
| 348 |
-
"fingerprint_schema_version": _RUN_FINGERPRINT_SCHEMA_VERSION,
|
| 349 |
-
"run_fingerprint": run_fingerprint,
|
| 350 |
-
"input_fingerprint": input_fingerprint,
|
| 351 |
-
"model_state_fingerprint": resolved_model_state_fingerprint,
|
| 352 |
-
"model_state_fingerprint_source": model_state_fingerprint_source,
|
| 353 |
-
"model_class": f"{model.__class__.__module__}.{model.__class__.__qualname__}",
|
| 354 |
-
**model_identity,
|
| 355 |
-
"dtype": str(dtype).removeprefix("torch.") if dtype is not None else "model",
|
| 356 |
-
"attention_backend": attention_backend,
|
| 357 |
-
"attention_kernel": _attention_kernel_metadata(attention_backend),
|
| 358 |
-
"layer": resolved_layer,
|
| 359 |
-
"projection": projection,
|
| 360 |
-
"esmc_source": getattr(model, "_esmc_source", None),
|
| 361 |
-
"esmc_revision": getattr(model, "_esmc_source_revision", None),
|
| 362 |
-
"esmc_files": getattr(model, "_esmc_source_files", None),
|
| 363 |
-
"token_policy": token_policy,
|
| 364 |
-
"tokenizer": tokenizer_metadata,
|
| 365 |
-
**embedding_context,
|
| 366 |
-
"pooling": list(pooling_names),
|
| 367 |
-
"pool_slices": pool_slices,
|
| 368 |
-
"full_embeddings": full_embeddings,
|
| 369 |
-
"max_length": max_length,
|
| 370 |
-
"truncate": truncate,
|
| 371 |
-
"truncation": {"enabled": truncate, "max_length": max_length},
|
| 372 |
-
"batching": {
|
| 373 |
-
"batch_size": batch_size,
|
| 374 |
-
"batch_window_size": resolved_batch_window_size,
|
| 375 |
-
"max_tokens_per_batch": max_tokens_per_batch,
|
| 376 |
-
"input_storage": ("disk-spool" if isinstance(records, _InputSpool) else "memory"),
|
| 377 |
-
"ordering": "bounded-length-bucketed-stable-output",
|
| 378 |
-
"resume_commit_granularity": (
|
| 379 |
-
"not-applicable"
|
| 380 |
-
if output is None
|
| 381 |
-
else "batch-window"
|
| 382 |
-
if format == "sqlite"
|
| 383 |
-
else "shard-flush"
|
| 384 |
-
),
|
| 385 |
-
},
|
| 386 |
-
"residue_mask_policy": "biological-residues-only",
|
| 387 |
-
"record_count": len(records),
|
| 388 |
-
"descriptor_index": (
|
| 389 |
-
"memory-metadata"
|
| 390 |
-
if output is None
|
| 391 |
-
else "sqlite-records"
|
| 392 |
-
if format == "sqlite"
|
| 393 |
-
else "safetensors-generation-index"
|
| 394 |
-
),
|
| 395 |
-
"storage_format": format if output is not None else "memory",
|
| 396 |
-
"software": software_versions,
|
| 397 |
-
"execution": _execution_identity_metadata(model),
|
| 398 |
-
"adapter": _adapter_identity_metadata(model),
|
| 399 |
-
"torch_version": software_versions["torch"],
|
| 400 |
-
"transformers_version": software_versions["transformers"],
|
| 401 |
-
"complete": True,
|
| 402 |
-
}
|
| 403 |
-
if destination.output_descriptors is not None:
|
| 404 |
-
metadata["outputs"] = destination.output_descriptors
|
| 405 |
-
metadata["tensor_hashes"] = [item["sha256"] for item in destination.output_descriptors]
|
| 406 |
-
status = getattr(model, "esmc_precision_status", None)
|
| 407 |
-
if status is not None:
|
| 408 |
-
metadata["esmc_precision"] = status.as_dict() if hasattr(status, "as_dict") else status
|
| 409 |
-
return destination.finish(metadata)
|
| 410 |
-
|
| 411 |
-
|
| 412 |
-
class EmbeddingMixin:
|
| 413 |
-
"""Small delegation mixin shared by FastPLMs model classes."""
|
| 414 |
-
|
| 415 |
-
def embed_dataset(self, inputs: Any, **kwargs: Any) -> EmbeddingResult:
|
| 416 |
-
return embed_dataset(self, inputs, **kwargs)
|
| 417 |
-
|
| 418 |
-
|
| 419 |
-
__all__ = [
|
| 420 |
-
"EmbeddingMixin",
|
| 421 |
-
"embed_dataset",
|
| 422 |
-
"iter_fasta",
|
| 423 |
-
"parse_fasta",
|
| 424 |
-
"select_hidden_state_embeddings",
|
| 425 |
-
]
|
|
|
|
| 1 |
+
"""Coordinate input preparation, run identity, batch execution, and publication."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import torch
|
| 6 |
+
|
| 7 |
+
from collections.abc import Callable, Iterable, Mapping, Sequence
|
| 8 |
+
from pathlib import Path
|
| 9 |
+
from typing import Any
|
| 10 |
from torch import Tensor
|
| 11 |
|
| 12 |
from . import identity
|
| 13 |
from .batches import (
|
| 14 |
+
BatchExecutor,
|
| 15 |
+
_residue_embeddings as _residue_embeddings,
|
| 16 |
+
_temporary_eval,
|
| 17 |
+
select_hidden_state_embeddings as select_hidden_state_embeddings,
|
| 18 |
+
)
|
| 19 |
+
from .identity import (
|
| 20 |
+
_RUN_FINGERPRINT_SCHEMA_VERSION,
|
| 21 |
+
_adapter_identity_metadata,
|
| 22 |
+
_attention_backend,
|
| 23 |
+
_attention_kernel_metadata,
|
| 24 |
+
_embedding_context,
|
| 25 |
+
_execution_identity_metadata,
|
| 26 |
+
_fingerprint_jsonable,
|
| 27 |
+
_model_identity_metadata,
|
| 28 |
+
_run_fingerprint,
|
| 29 |
_tokenizer_metadata,
|
| 30 |
+
)
|
| 31 |
+
from .inputs import (
|
| 32 |
+
_InputSpool,
|
| 33 |
+
_normalize_inputs,
|
| 34 |
+
_validate_untruncated_lengths,
|
| 35 |
+
iter_fasta as iter_fasta,
|
| 36 |
+
parse_fasta as parse_fasta,
|
| 37 |
+
)
|
| 38 |
+
from .output import EmbeddingOutput
|
| 39 |
+
from .pooling import Pooler
|
| 40 |
+
from .types import EmbeddingBatch, EmbeddingInput, EmbeddingResult
|
| 41 |
+
|
| 42 |
+
|
| 43 |
+
_DEFAULT_BATCH_WINDOW_MULTIPLIER = 16
|
| 44 |
+
_SUPPORTED_STORAGE_FORMATS = frozenset({"safetensors", "sqlite"})
|
| 45 |
+
|
| 46 |
+
|
| 47 |
+
def embed_dataset(
|
| 48 |
+
model: Any,
|
| 49 |
+
inputs: (Iterable[str | EmbeddingInput | tuple[str, str]] | Mapping[str, str] | str | Path),
|
| 50 |
+
*,
|
| 51 |
+
batch_size: int = 2,
|
| 52 |
+
pooling: str | Sequence[str] | None = None,
|
| 53 |
+
full_embeddings: bool = False,
|
| 54 |
+
output: str | Path | None = None,
|
| 55 |
+
format: str = "safetensors",
|
| 56 |
+
resume: bool = True,
|
| 57 |
+
tokenizer: Any | None = None,
|
| 58 |
+
max_length: int | None = None,
|
| 59 |
+
truncate: bool = True,
|
| 60 |
+
dtype: torch.dtype | None = torch.float32,
|
| 61 |
+
shard_size: int = 2 * 1024**3,
|
| 62 |
+
model_state_fingerprint: str | None = None,
|
| 63 |
+
batch_window_size: int | None = None,
|
| 64 |
+
max_tokens_per_batch: int | None = None,
|
| 65 |
+
hidden_state_source: str = "encoder",
|
| 66 |
+
decoder_inputs: Sequence[str] | None = None,
|
| 67 |
+
decoder_input_ids: Tensor | None = None,
|
| 68 |
+
decoder_attention_mask: Tensor | None = None,
|
| 69 |
+
_embedding_batch_fn: Callable[..., EmbeddingBatch] | None = None,
|
| 70 |
+
_embedding_batch_identity: Mapping[str, Any] | None = None,
|
| 71 |
+
_allowed_unsupported_pooling: Sequence[str] = (),
|
| 72 |
+
**model_kwargs: Any,
|
| 73 |
+
) -> EmbeddingResult:
|
| 74 |
+
"""Embed protein sequences with stable ordering and residue-only pooling."""
|
| 75 |
+
|
| 76 |
+
for name, value in (
|
| 77 |
+
("batch_size", batch_size),
|
| 78 |
+
("shard_size", shard_size),
|
| 79 |
+
):
|
| 80 |
+
if not isinstance(value, int) or isinstance(value, bool):
|
| 81 |
+
raise TypeError(f"{name} must be a positive integer.")
|
| 82 |
+
if value <= 0:
|
| 83 |
+
raise ValueError(f"{name} must be a positive integer.")
|
| 84 |
+
for optional_name, optional_value in (
|
| 85 |
+
("max_length", max_length),
|
| 86 |
+
("max_tokens_per_batch", max_tokens_per_batch),
|
| 87 |
+
("batch_window_size", batch_window_size),
|
| 88 |
+
):
|
| 89 |
+
if optional_value is not None and (
|
| 90 |
+
not isinstance(optional_value, int) or isinstance(optional_value, bool)
|
| 91 |
+
):
|
| 92 |
+
raise TypeError(f"{optional_name} must be a positive integer when provided.")
|
| 93 |
+
if optional_value is not None and optional_value <= 0:
|
| 94 |
+
raise ValueError(f"{optional_name} must be a positive integer when provided.")
|
| 95 |
+
for name, value in (
|
| 96 |
+
("full_embeddings", full_embeddings),
|
| 97 |
+
("resume", resume),
|
| 98 |
+
("truncate", truncate),
|
| 99 |
+
):
|
| 100 |
+
if not isinstance(value, bool):
|
| 101 |
+
raise TypeError(f"{name} must be a boolean.")
|
| 102 |
+
if not isinstance(format, str):
|
| 103 |
+
raise TypeError("format must be a string.")
|
| 104 |
+
if output is not None and not isinstance(output, (str, Path)):
|
| 105 |
+
raise TypeError("output must be a path or None.")
|
| 106 |
+
if model_state_fingerprint is not None and (
|
| 107 |
+
not isinstance(model_state_fingerprint, str) or not model_state_fingerprint
|
| 108 |
+
):
|
| 109 |
+
raise ValueError("model_state_fingerprint must be a non-empty string when provided.")
|
| 110 |
+
if hidden_state_source not in {"encoder", "decoder"}:
|
| 111 |
+
raise ValueError("hidden_state_source must be 'encoder' or 'decoder'.")
|
| 112 |
+
hidden_state_index = model_kwargs.get("hidden_state_index", -1)
|
| 113 |
+
if not isinstance(hidden_state_index, int) or isinstance(hidden_state_index, bool):
|
| 114 |
+
raise TypeError("hidden_state_index must be an integer.")
|
| 115 |
+
store_all_hidden_states = model_kwargs.get("store_all_hidden_states", False)
|
| 116 |
+
if not isinstance(store_all_hidden_states, bool):
|
| 117 |
+
raise TypeError("store_all_hidden_states must be a boolean.")
|
| 118 |
+
if decoder_input_ids is not None:
|
| 119 |
+
if not isinstance(decoder_input_ids, Tensor):
|
| 120 |
+
raise TypeError("decoder_input_ids must be a tensor.")
|
| 121 |
+
if decoder_input_ids.is_meta:
|
| 122 |
+
raise ValueError("decoder_input_ids cannot be a meta tensor.")
|
| 123 |
+
if decoder_input_ids.ndim != 2 or decoder_input_ids.shape[1] == 0:
|
| 124 |
+
raise ValueError("decoder_input_ids must have non-empty shape (batch, sequence).")
|
| 125 |
+
if decoder_input_ids.dtype not in {torch.int32, torch.int64}:
|
| 126 |
+
raise TypeError("decoder_input_ids must use torch.int32 or torch.int64.")
|
| 127 |
+
if decoder_attention_mask is not None:
|
| 128 |
+
if not isinstance(decoder_attention_mask, Tensor):
|
| 129 |
+
raise TypeError("decoder_attention_mask must be a tensor.")
|
| 130 |
+
if decoder_attention_mask.is_meta:
|
| 131 |
+
raise ValueError("decoder_attention_mask cannot be a meta tensor.")
|
| 132 |
+
if decoder_attention_mask.is_complex() or not bool(
|
| 133 |
+
torch.isfinite(decoder_attention_mask).all()
|
| 134 |
+
):
|
| 135 |
+
raise ValueError("decoder_attention_mask must contain finite binary values.")
|
| 136 |
+
if not bool(((decoder_attention_mask == 0) | (decoder_attention_mask == 1)).all()):
|
| 137 |
+
raise ValueError("decoder_attention_mask must contain finite binary values.")
|
| 138 |
+
pooling_names = (
|
| 139 |
+
(("mean",) if not full_embeddings else ())
|
| 140 |
+
if pooling is None
|
| 141 |
+
else ((pooling,) if isinstance(pooling, str) else tuple(pooling))
|
| 142 |
+
)
|
| 143 |
+
if full_embeddings and pooling is not None:
|
| 144 |
+
raise ValueError("full_embeddings=True cannot be combined with pooling.")
|
| 145 |
+
if not full_embeddings and not pooling_names:
|
| 146 |
+
raise ValueError("pooling is required unless full_embeddings=True.")
|
| 147 |
+
pooler = Pooler(pooling_names) if pooling_names else None
|
| 148 |
+
|
| 149 |
+
if batch_size <= 0:
|
| 150 |
+
raise ValueError("batch_size must be positive.")
|
| 151 |
+
if format == "pth" or (output is not None and Path(output).suffix.lower() == ".pth"):
|
| 152 |
+
raise ValueError("Writing pickle-based .pth embeddings is not supported.")
|
| 153 |
+
if format not in _SUPPORTED_STORAGE_FORMATS:
|
| 154 |
+
raise ValueError("format must be 'safetensors' or 'sqlite'.")
|
| 155 |
+
if max_length is not None and max_length <= 0:
|
| 156 |
+
raise ValueError("max_length must be positive when provided.")
|
| 157 |
+
if max_tokens_per_batch is not None and max_tokens_per_batch <= 0:
|
| 158 |
+
raise ValueError("max_tokens_per_batch must be positive when provided.")
|
| 159 |
+
if not isinstance(dtype, (torch.dtype, type(None))):
|
| 160 |
+
raise TypeError("dtype must be a torch.dtype or None.")
|
| 161 |
+
if batch_window_size is not None and batch_window_size <= 0:
|
| 162 |
+
raise ValueError("batch_window_size must be positive when provided.")
|
| 163 |
+
if _embedding_batch_fn is not None and not callable(_embedding_batch_fn):
|
| 164 |
+
raise TypeError("_embedding_batch_fn must be callable when provided.")
|
| 165 |
+
if _embedding_batch_fn is not None and _embedding_batch_identity is None:
|
| 166 |
+
raise ValueError(
|
| 167 |
+
"_embedding_batch_identity is required with _embedding_batch_fn so persisted "
|
| 168 |
+
"runs bind the family-specific embedding behavior."
|
| 169 |
+
)
|
| 170 |
+
if _embedding_batch_identity is not None and not isinstance(_embedding_batch_identity, Mapping):
|
| 171 |
+
raise TypeError("_embedding_batch_identity must be a mapping when provided.")
|
| 172 |
+
if isinstance(_allowed_unsupported_pooling, (str, bytes)) or not isinstance(
|
| 173 |
+
_allowed_unsupported_pooling, Sequence
|
| 174 |
+
):
|
| 175 |
+
raise TypeError("_allowed_unsupported_pooling must be a sequence of pooler names.")
|
| 176 |
+
if not all(isinstance(name, str) for name in _allowed_unsupported_pooling):
|
| 177 |
+
raise TypeError("_allowed_unsupported_pooling must contain only strings.")
|
| 178 |
+
allowed_unsupported_pooling = frozenset(_allowed_unsupported_pooling)
|
| 179 |
+
if allowed_unsupported_pooling and _embedding_batch_fn is None:
|
| 180 |
+
raise ValueError(
|
| 181 |
+
"_allowed_unsupported_pooling is only valid with a family-specific _embedding_batch_fn."
|
| 182 |
+
)
|
| 183 |
+
resolved_batch_window_size = (
|
| 184 |
+
batch_size * _DEFAULT_BATCH_WINDOW_MULTIPLIER
|
| 185 |
+
if batch_window_size is None
|
| 186 |
+
else batch_window_size
|
| 187 |
+
)
|
| 188 |
+
if resolved_batch_window_size < batch_size:
|
| 189 |
+
raise ValueError("batch_window_size must be at least batch_size.")
|
| 190 |
+
records = _normalize_inputs(inputs, disk_backed=output is not None)
|
| 191 |
+
_validate_untruncated_lengths(
|
| 192 |
+
records,
|
| 193 |
+
max_length=max_length,
|
| 194 |
+
truncate=truncate,
|
| 195 |
+
)
|
| 196 |
+
pooling_names = (
|
| 197 |
+
(("mean",) if not full_embeddings else ())
|
| 198 |
+
if pooling is None
|
| 199 |
+
else ((pooling,) if isinstance(pooling, str) else tuple(pooling))
|
| 200 |
+
)
|
| 201 |
+
if full_embeddings:
|
| 202 |
+
if pooling is not None:
|
| 203 |
+
raise ValueError("full_embeddings=True cannot be combined with pooling.")
|
| 204 |
+
elif not pooling_names:
|
| 205 |
+
raise ValueError("pooling is required unless full_embeddings=True.")
|
| 206 |
+
store_all_hidden_states = bool(model_kwargs.get("store_all_hidden_states", False))
|
| 207 |
+
if store_all_hidden_states and not full_embeddings:
|
| 208 |
+
raise ValueError("store_all_hidden_states=True requires full_embeddings=True.")
|
| 209 |
+
|
| 210 |
+
unsupported = set(getattr(model, "embedding_unsupported_pooling", ()))
|
| 211 |
+
unknown_pooling_overrides = allowed_unsupported_pooling.difference(unsupported)
|
| 212 |
+
if unknown_pooling_overrides:
|
| 213 |
+
raise ValueError(
|
| 214 |
+
"_allowed_unsupported_pooling may only override poolers declared unsupported "
|
| 215 |
+
f"by the model; unknown overrides: {sorted(unknown_pooling_overrides)}."
|
| 216 |
+
)
|
| 217 |
+
unsupported.difference_update(allowed_unsupported_pooling)
|
| 218 |
+
requested_unsupported = unsupported.intersection(pooling_names)
|
| 219 |
+
if requested_unsupported:
|
| 220 |
+
raise ValueError(
|
| 221 |
+
f"{model.__class__.__name__} does not support pooling operations "
|
| 222 |
+
f"{sorted(requested_unsupported)}."
|
| 223 |
+
)
|
| 224 |
+
|
| 225 |
+
# Constructing the pooler validates names and duplicate operations before
|
| 226 |
+
# any checkpoint hashing, tokenization, or inference occurs.
|
| 227 |
+
pooler = Pooler(pooling_names) if pooling_names else None
|
| 228 |
+
embedding_context, normalized_decoder_inputs = _embedding_context(
|
| 229 |
+
model,
|
| 230 |
+
records,
|
| 231 |
+
hidden_state_source=hidden_state_source,
|
| 232 |
+
decoder_inputs=decoder_inputs,
|
| 233 |
+
decoder_input_ids=decoder_input_ids,
|
| 234 |
+
decoder_attention_mask=decoder_attention_mask,
|
| 235 |
+
model_kwargs=model_kwargs,
|
| 236 |
+
)
|
| 237 |
+
if _embedding_batch_identity is not None:
|
| 238 |
+
embedding_context["family_adapter"] = _fingerprint_jsonable(_embedding_batch_identity)
|
| 239 |
+
if allowed_unsupported_pooling:
|
| 240 |
+
embedding_context["family_adapter_pooling_override"] = sorted(
|
| 241 |
+
allowed_unsupported_pooling
|
| 242 |
+
)
|
| 243 |
+
|
| 244 |
+
# A pending automatic attention request settles here, inside the caller's
|
| 245 |
+
# autocast context, so the fingerprint records the backend that executes.
|
| 246 |
+
attention_resolution = getattr(model, "attention_resolution", None)
|
| 247 |
+
if attention_resolution is not None and attention_resolution.deferred:
|
| 248 |
+
model.resolve_attn_implementation()
|
| 249 |
+
|
| 250 |
+
tokenizer_metadata = _tokenizer_metadata(model, tokenizer)
|
| 251 |
+
(
|
| 252 |
+
input_fingerprint,
|
| 253 |
+
run_fingerprint,
|
| 254 |
+
resolved_model_state_fingerprint,
|
| 255 |
+
model_state_fingerprint_source,
|
| 256 |
+
) = _run_fingerprint(
|
| 257 |
+
model,
|
| 258 |
+
records,
|
| 259 |
+
pooling=pooling_names,
|
| 260 |
+
full_embeddings=full_embeddings,
|
| 261 |
+
max_length=max_length,
|
| 262 |
+
truncate=truncate,
|
| 263 |
+
dtype=dtype,
|
| 264 |
+
model_kwargs=model_kwargs,
|
| 265 |
+
tokenizer_metadata=tokenizer_metadata,
|
| 266 |
+
model_state_fingerprint=model_state_fingerprint,
|
| 267 |
+
persist_output=output is not None,
|
| 268 |
+
embedding_context=embedding_context,
|
| 269 |
+
batch_size=batch_size,
|
| 270 |
+
batch_window_size=resolved_batch_window_size,
|
| 271 |
+
max_tokens_per_batch=max_tokens_per_batch,
|
| 272 |
+
)
|
| 273 |
+
destination = EmbeddingOutput(
|
| 274 |
+
records,
|
| 275 |
+
output=output,
|
| 276 |
+
format=format,
|
| 277 |
+
resume=resume,
|
| 278 |
+
shard_size=shard_size,
|
| 279 |
+
run_fingerprint=run_fingerprint,
|
| 280 |
+
input_fingerprint=input_fingerprint,
|
| 281 |
+
model_state_fingerprint=resolved_model_state_fingerprint,
|
| 282 |
+
model_state_fingerprint_source=model_state_fingerprint_source,
|
| 283 |
+
pooler=pooler,
|
| 284 |
+
pooling_names=pooling_names,
|
| 285 |
+
)
|
| 286 |
+
if destination.completed is not None:
|
| 287 |
+
return destination.completed
|
| 288 |
+
|
| 289 |
+
attention_backend = _attention_backend(model)
|
| 290 |
+
executor = BatchExecutor(
|
| 291 |
+
model=model,
|
| 292 |
+
batch_size=batch_size,
|
| 293 |
+
max_tokens_per_batch=max_tokens_per_batch,
|
| 294 |
+
max_length=max_length,
|
| 295 |
+
truncate=truncate,
|
| 296 |
+
model_kwargs=model_kwargs,
|
| 297 |
+
hidden_state_source=hidden_state_source,
|
| 298 |
+
normalized_decoder_inputs=normalized_decoder_inputs,
|
| 299 |
+
decoder_input_ids=decoder_input_ids,
|
| 300 |
+
decoder_attention_mask=decoder_attention_mask,
|
| 301 |
+
_embedding_batch_fn=_embedding_batch_fn,
|
| 302 |
+
tokenizer=tokenizer,
|
| 303 |
+
store_all_hidden_states=store_all_hidden_states,
|
| 304 |
+
full_embeddings=full_embeddings,
|
| 305 |
+
dtype=dtype,
|
| 306 |
+
pooler=pooler,
|
| 307 |
+
attention_backend=attention_backend,
|
| 308 |
+
need_attentions="parti" in pooling_names,
|
| 309 |
+
)
|
| 310 |
+
pool_slices = destination.pool_slices
|
| 311 |
+
with _temporary_eval(model), torch.inference_mode():
|
| 312 |
+
for window_start in range(
|
| 313 |
+
destination.start_position, len(records), resolved_batch_window_size
|
| 314 |
+
):
|
| 315 |
+
window_stop = min(window_start + resolved_batch_window_size, len(records))
|
| 316 |
+
window_records = records[window_start:window_stop]
|
| 317 |
+
if not isinstance(window_records, Sequence):
|
| 318 |
+
raise RuntimeError("The immutable embedding spool returned a non-sequence window.")
|
| 319 |
+
new_records, pool_slices = executor.run_window(
|
| 320 |
+
window_records, window_start=window_start
|
| 321 |
+
)
|
| 322 |
+
destination.append(window_start, new_records)
|
| 323 |
+
|
| 324 |
software_versions = identity._software_versions()
|
| 325 |
+
projection = getattr(model, "embedding_projection", None)
|
| 326 |
+
resolved_layer = getattr(
|
| 327 |
+
model,
|
| 328 |
+
"embedding_layer",
|
| 329 |
+
model_kwargs.get("hidden_state_index", -1),
|
| 330 |
+
)
|
| 331 |
+
token_policy = getattr(
|
| 332 |
+
model,
|
| 333 |
+
"embedding_token_policy",
|
| 334 |
+
{
|
| 335 |
+
"unit": "residue",
|
| 336 |
+
"include": ["biological residues"],
|
| 337 |
+
"exclude": [
|
| 338 |
+
"BOS",
|
| 339 |
+
"EOS",
|
| 340 |
+
"padding",
|
| 341 |
+
"chain delimiters",
|
| 342 |
+
"non-protein tokens",
|
| 343 |
+
],
|
| 344 |
+
},
|
| 345 |
+
)
|
| 346 |
+
model_identity = _model_identity_metadata(model)
|
| 347 |
+
metadata: dict[str, Any] = {
|
| 348 |
+
"format_version": 1,
|
| 349 |
+
"fingerprint_schema_version": _RUN_FINGERPRINT_SCHEMA_VERSION,
|
| 350 |
+
"run_fingerprint": run_fingerprint,
|
| 351 |
+
"input_fingerprint": input_fingerprint,
|
| 352 |
+
"model_state_fingerprint": resolved_model_state_fingerprint,
|
| 353 |
+
"model_state_fingerprint_source": model_state_fingerprint_source,
|
| 354 |
+
"model_class": f"{model.__class__.__module__}.{model.__class__.__qualname__}",
|
| 355 |
+
**model_identity,
|
| 356 |
+
"dtype": str(dtype).removeprefix("torch.") if dtype is not None else "model",
|
| 357 |
+
"attention_backend": attention_backend,
|
| 358 |
+
"attention_kernel": _attention_kernel_metadata(attention_backend),
|
| 359 |
+
"layer": resolved_layer,
|
| 360 |
+
"projection": projection,
|
| 361 |
+
"esmc_source": getattr(model, "_esmc_source", None),
|
| 362 |
+
"esmc_revision": getattr(model, "_esmc_source_revision", None),
|
| 363 |
+
"esmc_files": getattr(model, "_esmc_source_files", None),
|
| 364 |
+
"token_policy": token_policy,
|
| 365 |
+
"tokenizer": tokenizer_metadata,
|
| 366 |
+
**embedding_context,
|
| 367 |
+
"pooling": list(pooling_names),
|
| 368 |
+
"pool_slices": pool_slices,
|
| 369 |
+
"full_embeddings": full_embeddings,
|
| 370 |
+
"max_length": max_length,
|
| 371 |
+
"truncate": truncate,
|
| 372 |
+
"truncation": {"enabled": truncate, "max_length": max_length},
|
| 373 |
+
"batching": {
|
| 374 |
+
"batch_size": batch_size,
|
| 375 |
+
"batch_window_size": resolved_batch_window_size,
|
| 376 |
+
"max_tokens_per_batch": max_tokens_per_batch,
|
| 377 |
+
"input_storage": ("disk-spool" if isinstance(records, _InputSpool) else "memory"),
|
| 378 |
+
"ordering": "bounded-length-bucketed-stable-output",
|
| 379 |
+
"resume_commit_granularity": (
|
| 380 |
+
"not-applicable"
|
| 381 |
+
if output is None
|
| 382 |
+
else "batch-window"
|
| 383 |
+
if format == "sqlite"
|
| 384 |
+
else "shard-flush"
|
| 385 |
+
),
|
| 386 |
+
},
|
| 387 |
+
"residue_mask_policy": "biological-residues-only",
|
| 388 |
+
"record_count": len(records),
|
| 389 |
+
"descriptor_index": (
|
| 390 |
+
"memory-metadata"
|
| 391 |
+
if output is None
|
| 392 |
+
else "sqlite-records"
|
| 393 |
+
if format == "sqlite"
|
| 394 |
+
else "safetensors-generation-index"
|
| 395 |
+
),
|
| 396 |
+
"storage_format": format if output is not None else "memory",
|
| 397 |
+
"software": software_versions,
|
| 398 |
+
"execution": _execution_identity_metadata(model),
|
| 399 |
+
"adapter": _adapter_identity_metadata(model),
|
| 400 |
+
"torch_version": software_versions["torch"],
|
| 401 |
+
"transformers_version": software_versions["transformers"],
|
| 402 |
+
"complete": True,
|
| 403 |
+
}
|
| 404 |
+
if destination.output_descriptors is not None:
|
| 405 |
+
metadata["outputs"] = destination.output_descriptors
|
| 406 |
+
metadata["tensor_hashes"] = [item["sha256"] for item in destination.output_descriptors]
|
| 407 |
+
status = getattr(model, "esmc_precision_status", None)
|
| 408 |
+
if status is not None:
|
| 409 |
+
metadata["esmc_precision"] = status.as_dict() if hasattr(status, "as_dict") else status
|
| 410 |
+
return destination.finish(metadata)
|
| 411 |
+
|
| 412 |
+
|
| 413 |
+
class EmbeddingMixin:
|
| 414 |
+
"""Small delegation mixin shared by FastPLMs model classes."""
|
| 415 |
+
|
| 416 |
+
def embed_dataset(self, inputs: Any, **kwargs: Any) -> EmbeddingResult:
|
| 417 |
+
return embed_dataset(self, inputs, **kwargs)
|
| 418 |
+
|
| 419 |
+
|
| 420 |
+
__all__ = [
|
| 421 |
+
"EmbeddingMixin",
|
| 422 |
+
"embed_dataset",
|
| 423 |
+
"iter_fasta",
|
| 424 |
+
"parse_fasta",
|
| 425 |
+
"select_hidden_state_embeddings",
|
| 426 |
+
]
|
fastplms/embeddings/storage.py
CHANGED
|
@@ -9,6 +9,7 @@ import sqlite3
|
|
| 9 |
import struct
|
| 10 |
import numpy as np
|
| 11 |
import torch
|
|
|
|
| 12 |
from bisect import bisect_right
|
| 13 |
from collections.abc import Iterable, Iterator, Sequence
|
| 14 |
from pathlib import Path
|
|
|
|
| 9 |
import struct
|
| 10 |
import numpy as np
|
| 11 |
import torch
|
| 12 |
+
|
| 13 |
from bisect import bisect_right
|
| 14 |
from collections.abc import Iterable, Iterator, Sequence
|
| 15 |
from pathlib import Path
|
fastplms/models/_esm_rotary.py
CHANGED
|
@@ -9,6 +9,7 @@ Transformers implementation.
|
|
| 9 |
from __future__ import annotations
|
| 10 |
|
| 11 |
import torch
|
|
|
|
| 12 |
from torch import nn
|
| 13 |
|
| 14 |
|
|
|
|
| 9 |
from __future__ import annotations
|
| 10 |
|
| 11 |
import torch
|
| 12 |
+
|
| 13 |
from torch import nn
|
| 14 |
|
| 15 |
|
fastplms/models/classification_probe.py
CHANGED
|
@@ -3,9 +3,9 @@
|
|
| 3 |
from __future__ import annotations
|
| 4 |
|
| 5 |
import math
|
| 6 |
-
from typing import Any
|
| 7 |
-
|
| 8 |
import torch
|
|
|
|
|
|
|
| 9 |
from torch import nn
|
| 10 |
from torch.nn import functional as F
|
| 11 |
from transformers.modeling_outputs import (
|
|
@@ -14,6 +14,7 @@ from transformers.modeling_outputs import (
|
|
| 14 |
TokenClassifierOutput,
|
| 15 |
)
|
| 16 |
|
|
|
|
| 17 |
try:
|
| 18 |
from fastplms.attention import (
|
| 19 |
AttentionBackend,
|
|
@@ -145,7 +146,7 @@ def token_classification_loss(
|
|
| 145 |
if problem_type == "regression":
|
| 146 |
targets = labels.to(logits.dtype)
|
| 147 |
if num_labels == 1 and targets.ndim == logits.ndim - 1:
|
| 148 |
-
targets = targets.unsqueeze(-1)
|
| 149 |
if targets.shape != logits.shape:
|
| 150 |
raise ValueError(
|
| 151 |
"Token regression labels must match logits, except that the final "
|
|
@@ -210,7 +211,7 @@ class ProbeSelfAttention(nn.Module):
|
|
| 210 |
sequence_length,
|
| 211 |
self.num_heads,
|
| 212 |
self.head_size,
|
| 213 |
-
).transpose(1, 2)
|
| 214 |
|
| 215 |
def forward(
|
| 216 |
self,
|
|
@@ -220,11 +221,11 @@ class ProbeSelfAttention(nn.Module):
|
|
| 220 |
output_attentions: bool,
|
| 221 |
) -> tuple[torch.Tensor, torch.Tensor | None]:
|
| 222 |
batch_size, sequence_length, _ = hidden_states.shape
|
| 223 |
-
query, key, value = self.qkv(hidden_states).chunk(3, dim=-1)
|
| 224 |
-
query = self._reshape(query)
|
| 225 |
-
key = self._reshape(key)
|
| 226 |
-
value = self._reshape(value)
|
| 227 |
-
query, key = self.rotary(query, key)
|
| 228 |
if output_attentions and self.backend != AttentionBackend.EAGER:
|
| 229 |
raise ValueError(
|
| 230 |
f"output_attentions=True is unavailable for {self.backend.value!r}; "
|
|
@@ -241,11 +242,11 @@ class ProbeSelfAttention(nn.Module):
|
|
| 241 |
dropout = self.dropout if self.training else 0.0
|
| 242 |
attention_weights = None
|
| 243 |
if self.backend == AttentionBackend.EAGER:
|
| 244 |
-
scores = query @ key.transpose(-2, -1) / math.sqrt(self.head_size)
|
| 245 |
if attention_mask_4d is not None:
|
| 246 |
-
scores = scores.masked_fill(~attention_mask_4d, float("-inf"))
|
| 247 |
-
attention_weights = scores.softmax(dim=-1)
|
| 248 |
-
context = F.dropout(attention_weights, p=dropout, training=self.training) @ value
|
| 249 |
elif self.backend == AttentionBackend.SDPA:
|
| 250 |
context = F.scaled_dot_product_attention(
|
| 251 |
query,
|
|
@@ -253,7 +254,7 @@ class ProbeSelfAttention(nn.Module):
|
|
| 253 |
value,
|
| 254 |
attn_mask=attention_mask_4d,
|
| 255 |
dropout_p=dropout,
|
| 256 |
-
)
|
| 257 |
elif self.backend == AttentionBackend.FLEX_ATTENTION:
|
| 258 |
if flex_attention is None:
|
| 259 |
raise RuntimeError("'flex_attention' was requested but is unavailable.")
|
|
@@ -272,15 +273,15 @@ class ProbeSelfAttention(nn.Module):
|
|
| 272 |
block_mask=flex_block_mask,
|
| 273 |
scale=1.0 / math.sqrt(self.head_size),
|
| 274 |
kernel_options={"PRESCALE_QK": True, "BLOCK_N": 32},
|
| 275 |
-
)
|
| 276 |
else:
|
| 277 |
raise AssertionError(f"Unhandled attention backend {self.backend.value!r}.")
|
| 278 |
context = context.transpose(1, 2).contiguous().view(
|
| 279 |
batch_size,
|
| 280 |
sequence_length,
|
| 281 |
self.hidden_size,
|
| 282 |
-
)
|
| 283 |
-
return self.output(context), attention_weights
|
| 284 |
|
| 285 |
|
| 286 |
class ProteinTransformerProbe(nn.Module):
|
|
@@ -451,7 +452,7 @@ class SequenceClassificationProbe(_ClassificationProbe):
|
|
| 451 |
embeddings.shape[:2],
|
| 452 |
device=embeddings.device,
|
| 453 |
dtype=torch.bool,
|
| 454 |
-
)
|
| 455 |
outputs = self._forward_transformer(
|
| 456 |
embeddings,
|
| 457 |
attention_mask,
|
|
@@ -460,8 +461,8 @@ class SequenceClassificationProbe(_ClassificationProbe):
|
|
| 460 |
)
|
| 461 |
if self.pooler is None:
|
| 462 |
raise AssertionError("Sequence classification requires a configured pooler.")
|
| 463 |
-
pooled = self.pooler(outputs.last_hidden_state, attention_mask)
|
| 464 |
-
logits = self.classifier(pooled)
|
| 465 |
loss = None
|
| 466 |
if labels is not None:
|
| 467 |
problem_type = resolve_problem_type(self.config, labels, num_labels=self.num_labels)
|
|
|
|
| 3 |
from __future__ import annotations
|
| 4 |
|
| 5 |
import math
|
|
|
|
|
|
|
| 6 |
import torch
|
| 7 |
+
|
| 8 |
+
from typing import Any
|
| 9 |
from torch import nn
|
| 10 |
from torch.nn import functional as F
|
| 11 |
from transformers.modeling_outputs import (
|
|
|
|
| 14 |
TokenClassifierOutput,
|
| 15 |
)
|
| 16 |
|
| 17 |
+
|
| 18 |
try:
|
| 19 |
from fastplms.attention import (
|
| 20 |
AttentionBackend,
|
|
|
|
| 146 |
if problem_type == "regression":
|
| 147 |
targets = labels.to(logits.dtype)
|
| 148 |
if num_labels == 1 and targets.ndim == logits.ndim - 1:
|
| 149 |
+
targets = targets.unsqueeze(-1) # (..., 1), matching single-target logits
|
| 150 |
if targets.shape != logits.shape:
|
| 151 |
raise ValueError(
|
| 152 |
"Token regression labels must match logits, except that the final "
|
|
|
|
| 211 |
sequence_length,
|
| 212 |
self.num_heads,
|
| 213 |
self.head_size,
|
| 214 |
+
).transpose(1, 2) # (b, h, l, d_h)
|
| 215 |
|
| 216 |
def forward(
|
| 217 |
self,
|
|
|
|
| 221 |
output_attentions: bool,
|
| 222 |
) -> tuple[torch.Tensor, torch.Tensor | None]:
|
| 223 |
batch_size, sequence_length, _ = hidden_states.shape
|
| 224 |
+
query, key, value = self.qkv(hidden_states).chunk(3, dim=-1) # each (b, l, d)
|
| 225 |
+
query = self._reshape(query) # (b, h, l, d_h)
|
| 226 |
+
key = self._reshape(key) # (b, h, l, d_h)
|
| 227 |
+
value = self._reshape(value) # (b, h, l, d_h)
|
| 228 |
+
query, key = self.rotary(query, key) # each (b, h, l, d_h)
|
| 229 |
if output_attentions and self.backend != AttentionBackend.EAGER:
|
| 230 |
raise ValueError(
|
| 231 |
f"output_attentions=True is unavailable for {self.backend.value!r}; "
|
|
|
|
| 242 |
dropout = self.dropout if self.training else 0.0
|
| 243 |
attention_weights = None
|
| 244 |
if self.backend == AttentionBackend.EAGER:
|
| 245 |
+
scores = query @ key.transpose(-2, -1) / math.sqrt(self.head_size) # (b, h, l, l)
|
| 246 |
if attention_mask_4d is not None:
|
| 247 |
+
scores = scores.masked_fill(~attention_mask_4d, float("-inf")) # (b, h, l, l)
|
| 248 |
+
attention_weights = scores.softmax(dim=-1) # (b, h, l, l)
|
| 249 |
+
context = F.dropout(attention_weights, p=dropout, training=self.training) @ value # (b, h, l, d_h)
|
| 250 |
elif self.backend == AttentionBackend.SDPA:
|
| 251 |
context = F.scaled_dot_product_attention(
|
| 252 |
query,
|
|
|
|
| 254 |
value,
|
| 255 |
attn_mask=attention_mask_4d,
|
| 256 |
dropout_p=dropout,
|
| 257 |
+
) # (b, h, l, d_h)
|
| 258 |
elif self.backend == AttentionBackend.FLEX_ATTENTION:
|
| 259 |
if flex_attention is None:
|
| 260 |
raise RuntimeError("'flex_attention' was requested but is unavailable.")
|
|
|
|
| 273 |
block_mask=flex_block_mask,
|
| 274 |
scale=1.0 / math.sqrt(self.head_size),
|
| 275 |
kernel_options={"PRESCALE_QK": True, "BLOCK_N": 32},
|
| 276 |
+
) # (b, h, l, d_h)
|
| 277 |
else:
|
| 278 |
raise AssertionError(f"Unhandled attention backend {self.backend.value!r}.")
|
| 279 |
context = context.transpose(1, 2).contiguous().view(
|
| 280 |
batch_size,
|
| 281 |
sequence_length,
|
| 282 |
self.hidden_size,
|
| 283 |
+
) # (b, l, d)
|
| 284 |
+
return self.output(context), attention_weights # (b, l, d), optional (b, h, l, l)
|
| 285 |
|
| 286 |
|
| 287 |
class ProteinTransformerProbe(nn.Module):
|
|
|
|
| 452 |
embeddings.shape[:2],
|
| 453 |
device=embeddings.device,
|
| 454 |
dtype=torch.bool,
|
| 455 |
+
) # (b, l)
|
| 456 |
outputs = self._forward_transformer(
|
| 457 |
embeddings,
|
| 458 |
attention_mask,
|
|
|
|
| 461 |
)
|
| 462 |
if self.pooler is None:
|
| 463 |
raise AssertionError("Sequence classification requires a configured pooler.")
|
| 464 |
+
pooled = self.pooler(outputs.last_hidden_state, attention_mask) # (b, d)
|
| 465 |
+
logits = self.classifier(pooled) # (b, num_labels)
|
| 466 |
loss = None
|
| 467 |
if labels is not None:
|
| 468 |
problem_type = resolve_problem_type(self.config, labels, num_labels=self.num_labels)
|
fastplms/models/esm_plusplus/modeling_esm_plusplus.py
CHANGED
|
@@ -6,15 +6,15 @@ import importlib
|
|
| 6 |
import importlib.metadata
|
| 7 |
import math
|
| 8 |
import os
|
|
|
|
|
|
|
|
|
|
|
|
|
| 9 |
from collections.abc import Sequence
|
| 10 |
from contextlib import contextmanager
|
| 11 |
from dataclasses import asdict, dataclass
|
| 12 |
from functools import partial
|
| 13 |
from typing import Any, ClassVar
|
| 14 |
-
|
| 15 |
-
import torch
|
| 16 |
-
import torch.nn as nn
|
| 17 |
-
import torch.nn.functional as F
|
| 18 |
from einops import rearrange
|
| 19 |
from tokenizers import Tokenizer
|
| 20 |
from tokenizers.models import BPE
|
|
@@ -378,7 +378,7 @@ class RotaryEmbedding(torch.nn.Module):
|
|
| 378 |
inv_freq = self._compute_inv_freq(buffer_device)
|
| 379 |
self._clear_cache()
|
| 380 |
self.register_buffer("inv_freq", inv_freq, persistent=False)
|
| 381 |
-
arange = torch.arange(0, self.dim, 2, device=buffer_device, dtype=torch.float32)
|
| 382 |
scale = (
|
| 383 |
(arange + 0.4 * self.dim) / (1.4 * self.dim) if self.scale_base is not None else None
|
| 384 |
)
|
|
@@ -446,18 +446,18 @@ class RotaryEmbedding(torch.nn.Module):
|
|
| 446 |
cos_angles = torch.cos(angles) # (l, d / 2)
|
| 447 |
sin_angles = torch.sin(angles) # (l, d / 2)
|
| 448 |
if self.scale is None:
|
| 449 |
-
self._cos_cached = cos_angles.to(dtype)
|
| 450 |
-
self._sin_cached = sin_angles.to(dtype)
|
| 451 |
-
self._cos_full_cached = torch.cat((self._cos_cached, self._cos_cached), dim=-1)
|
| 452 |
-
self._sin_full_cached = torch.cat((self._sin_cached, self._sin_cached), dim=-1)
|
| 453 |
return
|
| 454 |
|
| 455 |
centered_positions = (
|
| 456 |
torch.arange(seqlen, dtype=self.scale.dtype, device=self.scale.device) - seqlen // 2
|
| 457 |
-
) / self.scale_base
|
| 458 |
-
scale = self.scale ** centered_positions.unsqueeze(-1)
|
| 459 |
-
self._cos_cached = (cos_angles * scale).to(dtype)
|
| 460 |
-
self._sin_cached = (sin_angles * scale).to(dtype)
|
| 461 |
self._cos_k_cached = (cos_angles / scale).to(dtype)
|
| 462 |
self._sin_k_cached = (sin_angles / scale).to(dtype)
|
| 463 |
|
|
@@ -574,12 +574,12 @@ class MultiHeadAttention(nn.Module):
|
|
| 574 |
) -> tuple[torch.Tensor, torch.Tensor | None, list[torch.Tensor] | None]:
|
| 575 |
# x: (b, l, d)
|
| 576 |
qkv = self.layernorm_qkv(x) # (b, l, 3 * d)
|
| 577 |
-
query_sequence, key_sequence, value_sequence = torch.chunk(qkv, 3, dim=-1)
|
| 578 |
query_sequence, key_sequence = (
|
| 579 |
self.q_ln(query_sequence).to(query_sequence.dtype),
|
| 580 |
self.k_ln(key_sequence).to(query_sequence.dtype),
|
| 581 |
-
)
|
| 582 |
-
query_sequence, key_sequence = self._apply_rotary(query_sequence, key_sequence)
|
| 583 |
query_heads, key_heads, value_heads = map(
|
| 584 |
self.reshaper, (query_sequence, key_sequence, value_sequence)
|
| 585 |
) # each (b, h, l, d_h)
|
|
@@ -596,7 +596,7 @@ class MultiHeadAttention(nn.Module):
|
|
| 596 |
flash_padding_layout=flash_padding_layout,
|
| 597 |
)
|
| 598 |
|
| 599 |
-
output = self.out_proj(attn_output)
|
| 600 |
return output, attn_weights, s_max
|
| 601 |
|
| 602 |
def _attn(
|
|
@@ -682,9 +682,9 @@ class MultiHeadAttention(nn.Module):
|
|
| 682 |
attention_mask_2d: torch.Tensor | None = None,
|
| 683 |
flash_padding_layout: FlashPaddingLayout | None = None,
|
| 684 |
) -> tuple[torch.Tensor, None]:
|
| 685 |
-
query_tokens = query_heads.transpose(1, 2).contiguous()
|
| 686 |
-
key_tokens = key_heads.transpose(1, 2).contiguous()
|
| 687 |
-
value_tokens = value_heads.transpose(1, 2).contiguous()
|
| 688 |
attn_output = kernels_flash_attention_func(
|
| 689 |
query_states=query_tokens,
|
| 690 |
key_states=key_tokens,
|
|
@@ -693,8 +693,8 @@ class MultiHeadAttention(nn.Module):
|
|
| 693 |
causal=False,
|
| 694 |
implementation=self.attn_backend.value,
|
| 695 |
padding_layout=flash_padding_layout,
|
| 696 |
-
)
|
| 697 |
-
return rearrange(attn_output, "b s h d -> b s (h d)"), None
|
| 698 |
|
| 699 |
def _flex_attn(
|
| 700 |
self,
|
|
@@ -1015,11 +1015,11 @@ class TransformerStack(nn.Module):
|
|
| 1015 |
# finite without allowing their states to enter residue attention.
|
| 1016 |
attention_mask_4d = (
|
| 1017 |
mask_pattern[:, None, :, None] == mask_pattern[:, None, None, :]
|
| 1018 |
-
)
|
| 1019 |
else:
|
| 1020 |
attention_mask_4d = (
|
| 1021 |
mask_pattern.unsqueeze(-1) == mask_pattern.unsqueeze(-2)
|
| 1022 |
-
).unsqueeze(1)
|
| 1023 |
backend = (
|
| 1024 |
resolve_attention_backend_for_call(
|
| 1025 |
self.attention_backend,
|
|
@@ -1837,7 +1837,7 @@ class ESMplusplusForSequenceClassification(ESMplusplusForMaskedLM, EmbeddingMixi
|
|
| 1837 |
inputs_embeds.shape[:2],
|
| 1838 |
dtype=torch.bool,
|
| 1839 |
device=inputs_embeds.device,
|
| 1840 |
-
)
|
| 1841 |
|
| 1842 |
output = super().forward(
|
| 1843 |
input_ids=input_ids,
|
|
|
|
| 6 |
import importlib.metadata
|
| 7 |
import math
|
| 8 |
import os
|
| 9 |
+
import torch
|
| 10 |
+
import torch.nn as nn
|
| 11 |
+
import torch.nn.functional as F
|
| 12 |
+
|
| 13 |
from collections.abc import Sequence
|
| 14 |
from contextlib import contextmanager
|
| 15 |
from dataclasses import asdict, dataclass
|
| 16 |
from functools import partial
|
| 17 |
from typing import Any, ClassVar
|
|
|
|
|
|
|
|
|
|
|
|
|
| 18 |
from einops import rearrange
|
| 19 |
from tokenizers import Tokenizer
|
| 20 |
from tokenizers.models import BPE
|
|
|
|
| 378 |
inv_freq = self._compute_inv_freq(buffer_device)
|
| 379 |
self._clear_cache()
|
| 380 |
self.register_buffer("inv_freq", inv_freq, persistent=False)
|
| 381 |
+
arange = torch.arange(0, self.dim, 2, device=buffer_device, dtype=torch.float32) # (d / 2,)
|
| 382 |
scale = (
|
| 383 |
(arange + 0.4 * self.dim) / (1.4 * self.dim) if self.scale_base is not None else None
|
| 384 |
)
|
|
|
|
| 446 |
cos_angles = torch.cos(angles) # (l, d / 2)
|
| 447 |
sin_angles = torch.sin(angles) # (l, d / 2)
|
| 448 |
if self.scale is None:
|
| 449 |
+
self._cos_cached = cos_angles.to(dtype) # (l, d / 2)
|
| 450 |
+
self._sin_cached = sin_angles.to(dtype) # (l, d / 2)
|
| 451 |
+
self._cos_full_cached = torch.cat((self._cos_cached, self._cos_cached), dim=-1) # (l, d)
|
| 452 |
+
self._sin_full_cached = torch.cat((self._sin_cached, self._sin_cached), dim=-1) # (l, d)
|
| 453 |
return
|
| 454 |
|
| 455 |
centered_positions = (
|
| 456 |
torch.arange(seqlen, dtype=self.scale.dtype, device=self.scale.device) - seqlen // 2
|
| 457 |
+
) / self.scale_base # (l,)
|
| 458 |
+
scale = self.scale ** centered_positions.unsqueeze(-1) # (l, d / 2)
|
| 459 |
+
self._cos_cached = (cos_angles * scale).to(dtype) # (l, d / 2)
|
| 460 |
+
self._sin_cached = (sin_angles * scale).to(dtype) # (l, d / 2)
|
| 461 |
self._cos_k_cached = (cos_angles / scale).to(dtype)
|
| 462 |
self._sin_k_cached = (sin_angles / scale).to(dtype)
|
| 463 |
|
|
|
|
| 574 |
) -> tuple[torch.Tensor, torch.Tensor | None, list[torch.Tensor] | None]:
|
| 575 |
# x: (b, l, d)
|
| 576 |
qkv = self.layernorm_qkv(x) # (b, l, 3 * d)
|
| 577 |
+
query_sequence, key_sequence, value_sequence = torch.chunk(qkv, 3, dim=-1) # each (b, l, d)
|
| 578 |
query_sequence, key_sequence = (
|
| 579 |
self.q_ln(query_sequence).to(query_sequence.dtype),
|
| 580 |
self.k_ln(key_sequence).to(query_sequence.dtype),
|
| 581 |
+
) # each (b, l, d)
|
| 582 |
+
query_sequence, key_sequence = self._apply_rotary(query_sequence, key_sequence) # each (b, l, d)
|
| 583 |
query_heads, key_heads, value_heads = map(
|
| 584 |
self.reshaper, (query_sequence, key_sequence, value_sequence)
|
| 585 |
) # each (b, h, l, d_h)
|
|
|
|
| 596 |
flash_padding_layout=flash_padding_layout,
|
| 597 |
)
|
| 598 |
|
| 599 |
+
output = self.out_proj(attn_output) # (b, l, d)
|
| 600 |
return output, attn_weights, s_max
|
| 601 |
|
| 602 |
def _attn(
|
|
|
|
| 682 |
attention_mask_2d: torch.Tensor | None = None,
|
| 683 |
flash_padding_layout: FlashPaddingLayout | None = None,
|
| 684 |
) -> tuple[torch.Tensor, None]:
|
| 685 |
+
query_tokens = query_heads.transpose(1, 2).contiguous() # (b, l, h, d_h)
|
| 686 |
+
key_tokens = key_heads.transpose(1, 2).contiguous() # (b, l, h, d_h)
|
| 687 |
+
value_tokens = value_heads.transpose(1, 2).contiguous() # (b, l, h, d_h)
|
| 688 |
attn_output = kernels_flash_attention_func(
|
| 689 |
query_states=query_tokens,
|
| 690 |
key_states=key_tokens,
|
|
|
|
| 693 |
causal=False,
|
| 694 |
implementation=self.attn_backend.value,
|
| 695 |
padding_layout=flash_padding_layout,
|
| 696 |
+
) # (b, l, h, d_h)
|
| 697 |
+
return rearrange(attn_output, "b s h d -> b s (h d)"), None # (b, l, h * d_h), None
|
| 698 |
|
| 699 |
def _flex_attn(
|
| 700 |
self,
|
|
|
|
| 1015 |
# finite without allowing their states to enter residue attention.
|
| 1016 |
attention_mask_4d = (
|
| 1017 |
mask_pattern[:, None, :, None] == mask_pattern[:, None, None, :]
|
| 1018 |
+
) # (b, 1, l, l)
|
| 1019 |
else:
|
| 1020 |
attention_mask_4d = (
|
| 1021 |
mask_pattern.unsqueeze(-1) == mask_pattern.unsqueeze(-2)
|
| 1022 |
+
).unsqueeze(1) # (b, 1, l, l)
|
| 1023 |
backend = (
|
| 1024 |
resolve_attention_backend_for_call(
|
| 1025 |
self.attention_backend,
|
|
|
|
| 1837 |
inputs_embeds.shape[:2],
|
| 1838 |
dtype=torch.bool,
|
| 1839 |
device=inputs_embeds.device,
|
| 1840 |
+
) # (b, l)
|
| 1841 |
|
| 1842 |
output = super().forward(
|
| 1843 |
input_ids=input_ids,
|
fastplms/models/esmfold2/__init__.py
CHANGED
|
@@ -5,6 +5,7 @@ from __future__ import annotations
|
|
| 5 |
from importlib import import_module
|
| 6 |
from typing import TYPE_CHECKING, Any
|
| 7 |
|
|
|
|
| 8 |
if TYPE_CHECKING:
|
| 9 |
from .configuration_esmfold2 import ESMFold2Config as ESMFold2Config
|
| 10 |
from .modeling_esmfold2 import ESMFold2Model as ESMFold2Model
|
|
|
|
| 5 |
from importlib import import_module
|
| 6 |
from typing import TYPE_CHECKING, Any
|
| 7 |
|
| 8 |
+
|
| 9 |
if TYPE_CHECKING:
|
| 10 |
from .configuration_esmfold2 import ESMFold2Config as ESMFold2Config
|
| 11 |
from .modeling_esmfold2 import ESMFold2Model as ESMFold2Model
|
fastplms/models/esmfold2/configuration_esmfold2.py
CHANGED
|
@@ -18,7 +18,6 @@ from __future__ import annotations
|
|
| 18 |
|
| 19 |
from dataclasses import asdict, dataclass, field
|
| 20 |
from typing import Any, TypeVar, cast
|
| 21 |
-
|
| 22 |
from transformers.configuration_utils import PretrainedConfig
|
| 23 |
|
| 24 |
from fastplms.attention import canonical_checkpoint_attention_backend
|
|
|
|
| 18 |
|
| 19 |
from dataclasses import asdict, dataclass, field
|
| 20 |
from typing import Any, TypeVar, cast
|
|
|
|
| 21 |
from transformers.configuration_utils import PretrainedConfig
|
| 22 |
|
| 23 |
from fastplms.attention import canonical_checkpoint_attention_backend
|
fastplms/models/esmfold2/embedding.py
CHANGED
|
@@ -3,6 +3,7 @@
|
|
| 3 |
from __future__ import annotations
|
| 4 |
|
| 5 |
import torch
|
|
|
|
| 6 |
from typing import Any, ClassVar
|
| 7 |
from torch import Tensor
|
| 8 |
|
|
|
|
| 3 |
from __future__ import annotations
|
| 4 |
|
| 5 |
import torch
|
| 6 |
+
|
| 7 |
from typing import Any, ClassVar
|
| 8 |
from torch import Tensor
|
| 9 |
|
fastplms/models/esmfold2/esmfold2_affine3d.py
CHANGED
|
@@ -2,10 +2,11 @@
|
|
| 2 |
|
| 3 |
from __future__ import annotations
|
| 4 |
|
|
|
|
|
|
|
|
|
|
| 5 |
from dataclasses import dataclass
|
| 6 |
from typing import Any, Self
|
| 7 |
-
|
| 8 |
-
import torch
|
| 9 |
from torch.nn import functional as F
|
| 10 |
|
| 11 |
from .esmfold2_misc import fp32_autocast_context
|
|
@@ -19,23 +20,26 @@ def _index_tuple(index: Any) -> tuple[Any, ...]:
|
|
| 19 |
|
| 20 |
def _sqrt_subgradient(values: torch.Tensor) -> torch.Tensor:
|
| 21 |
"""Square root with a zero subgradient for non-positive inputs."""
|
|
|
|
| 22 |
|
| 23 |
-
result = torch.zeros_like(values)
|
| 24 |
-
positive = values > 0
|
| 25 |
-
result[positive] = torch.sqrt(values[positive])
|
| 26 |
-
return result
|
| 27 |
|
| 28 |
|
| 29 |
def _quat_invert(quaternion: torch.Tensor) -> torch.Tensor:
|
| 30 |
-
|
| 31 |
-
|
|
|
|
| 32 |
|
| 33 |
|
| 34 |
def _quat_mult(left: torch.Tensor, right: torch.Tensor) -> torch.Tensor:
|
| 35 |
"""Hamilton product for real-first quaternion tensors."""
|
|
|
|
| 36 |
|
| 37 |
-
aw, ax, ay, az = torch.unbind(left, -1)
|
| 38 |
-
bw, bx, by, bz = torch.unbind(right, -1)
|
| 39 |
return torch.stack(
|
| 40 |
(
|
| 41 |
aw * bw - ax * bx - ay * by - az * bz,
|
|
@@ -44,7 +48,7 @@ def _quat_mult(left: torch.Tensor, right: torch.Tensor) -> torch.Tensor:
|
|
| 44 |
aw * bz + ax * by - ay * bx + az * bw,
|
| 45 |
),
|
| 46 |
-1,
|
| 47 |
-
)
|
| 48 |
|
| 49 |
|
| 50 |
def _quat_rotation(
|
|
@@ -52,9 +56,10 @@ def _quat_rotation(
|
|
| 52 |
points: torch.Tensor,
|
| 53 |
) -> torch.Tensor:
|
| 54 |
"""Rotate points using normalized real-first quaternions."""
|
|
|
|
| 55 |
|
| 56 |
-
aw, ax, ay, az = torch.unbind(quaternion, -1)
|
| 57 |
-
bx, by, bz = torch.unbind(points, -1)
|
| 58 |
product = torch.stack(
|
| 59 |
(
|
| 60 |
-ax * bx - ay * by - az * bz,
|
|
@@ -63,8 +68,8 @@ def _quat_rotation(
|
|
| 63 |
aw * bz + ax * by - ay * bx,
|
| 64 |
),
|
| 65 |
-1,
|
| 66 |
-
)
|
| 67 |
-
return _quat_mult(product, _quat_invert(quaternion))[..., 1:]
|
| 68 |
|
| 69 |
|
| 70 |
def _graham_schmidt(
|
|
@@ -73,17 +78,18 @@ def _graham_schmidt(
|
|
| 73 |
eps: float = 1e-12,
|
| 74 |
) -> torch.Tensor:
|
| 75 |
"""Construct a right-handed orthonormal frame from two directions."""
|
|
|
|
| 76 |
|
| 77 |
with fp32_autocast_context(x_axis.device.type):
|
| 78 |
-
e1 = xy_plane
|
| 79 |
-
denominator = torch.sqrt((x_axis**2).sum(dim=-1, keepdim=True) + eps)
|
| 80 |
-
x_axis = x_axis / denominator
|
| 81 |
-
projection = (x_axis * e1).sum(dim=-1, keepdim=True)
|
| 82 |
-
e1 = e1 - x_axis * projection
|
| 83 |
-
denominator = torch.sqrt((e1**2).sum(dim=-1, keepdim=True) + eps)
|
| 84 |
-
e1 = e1 / denominator
|
| 85 |
-
e2 = torch.cross(x_axis, e1, dim=-1)
|
| 86 |
-
return torch.stack([x_axis, e1, e2], dim=-1)
|
| 87 |
|
| 88 |
|
| 89 |
class Rotation:
|
|
@@ -129,6 +135,7 @@ class Rotation:
|
|
| 129 |
|
| 130 |
@classmethod
|
| 131 |
def _from_tensor(cls, tensor: torch.Tensor) -> Self:
|
|
|
|
| 132 |
return cls(tensor) # type: ignore[call-arg]
|
| 133 |
|
| 134 |
def to(self, **kwargs) -> Self:
|
|
@@ -137,15 +144,17 @@ class Rotation:
|
|
| 137 |
def detach(self, *args, **kwargs) -> Self:
|
| 138 |
return self._from_tensor(self.tensor.detach(**kwargs))
|
| 139 |
|
| 140 |
-
def tensor_apply(self, func) -> Self:
|
| 141 |
-
|
| 142 |
-
|
|
|
|
| 143 |
|
| 144 |
|
| 145 |
class RotationQuat(Rotation):
|
| 146 |
"""A rotation represented by a real-first quaternion."""
|
| 147 |
|
| 148 |
-
def __init__(self, quats: torch.Tensor, normalized: bool = False):
|
|
|
|
| 149 |
if not isinstance(quats, torch.Tensor):
|
| 150 |
raise TypeError("quats must be a Torch tensor.")
|
| 151 |
if quats.ndim == 0 or quats.shape[-1] != 4:
|
|
@@ -156,32 +165,32 @@ class RotationQuat(Rotation):
|
|
| 156 |
raise TypeError("normalized must be a boolean.")
|
| 157 |
self._normalized = normalized
|
| 158 |
if normalized:
|
| 159 |
-
quats = F.normalize(quats.to(torch.float32), dim=-1)
|
| 160 |
-
self._quats = quats.where(quats[..., :1] >= 0, -quats)
|
| 161 |
else:
|
| 162 |
-
self._quats = quats.to(torch.float32)
|
| 163 |
|
| 164 |
@property
|
| 165 |
def tensor(self) -> torch.Tensor:
|
| 166 |
-
return self._quats
|
| 167 |
|
| 168 |
@property
|
| 169 |
def shape(self) -> torch.Size:
|
| 170 |
return self._quats.shape[:-1]
|
| 171 |
|
| 172 |
@classmethod
|
| 173 |
-
def identity(cls, shape, **tensor_kwargs) -> RotationQuat:
|
| 174 |
-
quaternions = torch.ones((*shape, 4), **tensor_kwargs)
|
| 175 |
-
selector = torch.tensor([1, 0, 0, 0], device=quaternions.device)
|
| 176 |
-
return cls(quaternions * selector)
|
| 177 |
|
| 178 |
@classmethod
|
| 179 |
-
def random(cls, shape, **tensor_kwargs) -> RotationQuat:
|
| 180 |
-
return cls(torch.randn((*shape, 4), **tensor_kwargs), normalized=True)
|
| 181 |
|
| 182 |
def __getitem__(self, idx: Any) -> RotationQuat:
|
| 183 |
indices = _index_tuple(idx)
|
| 184 |
-
return RotationQuat(self._quats[(*indices, slice(None))])
|
| 185 |
|
| 186 |
def normalized(self) -> RotationQuat:
|
| 187 |
if self._normalized:
|
|
@@ -192,9 +201,9 @@ class RotationQuat(Rotation):
|
|
| 192 |
return self
|
| 193 |
|
| 194 |
def as_matrix(self) -> RotationMatrix:
|
| 195 |
-
quaternion = self.normalized().tensor
|
| 196 |
-
r, i, j, k = torch.unbind(quaternion, -1)
|
| 197 |
-
scale = 2.0 / torch.linalg.norm(quaternion, dim=-1)
|
| 198 |
elements = torch.stack(
|
| 199 |
(
|
| 200 |
1 - scale * (j * j + k * k),
|
|
@@ -208,54 +217,56 @@ class RotationQuat(Rotation):
|
|
| 208 |
1 - scale * (i * i + j * j),
|
| 209 |
),
|
| 210 |
-1,
|
| 211 |
-
)
|
| 212 |
-
return RotationMatrix(elements.reshape((*quaternion.shape[:-1], 3, 3)))
|
| 213 |
|
| 214 |
def compose(self, other: RotationQuat) -> RotationQuat:
|
| 215 |
with fp32_autocast_context(self.device.type):
|
| 216 |
-
return RotationQuat(_quat_mult(self._quats, other._quats))
|
| 217 |
|
| 218 |
def convert_compose(self, other: Rotation) -> RotationQuat:
|
| 219 |
return self.compose(other.as_quat())
|
| 220 |
|
| 221 |
def apply(self, points: torch.Tensor) -> torch.Tensor:
|
| 222 |
-
|
|
|
|
| 223 |
|
| 224 |
def invert(self) -> RotationQuat:
|
| 225 |
-
return RotationQuat(_quat_invert(self._quats))
|
| 226 |
|
| 227 |
|
| 228 |
class RotationMatrix(Rotation):
|
| 229 |
"""A rotation represented by a dense FP32 matrix."""
|
| 230 |
|
| 231 |
-
def __init__(self, rots: torch.Tensor):
|
|
|
|
| 232 |
if not isinstance(rots, torch.Tensor):
|
| 233 |
raise TypeError("rots must be a Torch tensor.")
|
| 234 |
if rots.ndim > 0 and rots.shape[-1] == 9:
|
| 235 |
-
rots = rots.unflatten(-1, (3, 3))
|
| 236 |
if rots.ndim < 2 or rots.shape[-2:] != (3, 3):
|
| 237 |
raise ValueError(
|
| 238 |
"rots must have trailing shape (3, 3) or flattened width 9, got "
|
| 239 |
f"shape {tuple(rots.shape)}."
|
| 240 |
)
|
| 241 |
-
self._rots = rots.to(torch.float32)
|
| 242 |
|
| 243 |
@property
|
| 244 |
def tensor(self) -> torch.Tensor:
|
| 245 |
-
return self._rots.flatten(-2)
|
| 246 |
|
| 247 |
@property
|
| 248 |
def shape(self) -> torch.Size:
|
| 249 |
return self._rots.shape[:-2]
|
| 250 |
|
| 251 |
@classmethod
|
| 252 |
-
def identity(cls, shape, **tensor_kwargs) -> RotationMatrix:
|
| 253 |
-
matrix = torch.eye(3, **tensor_kwargs)
|
| 254 |
-
matrix = matrix.view(*(1 for _ in shape), 3, 3)
|
| 255 |
-
return cls(matrix.expand(*shape, -1, -1))
|
| 256 |
|
| 257 |
@classmethod
|
| 258 |
-
def random(cls, shape, **tensor_kwargs) -> RotationMatrix:
|
| 259 |
return RotationQuat.random(shape, **tensor_kwargs).as_matrix()
|
| 260 |
|
| 261 |
@staticmethod
|
|
@@ -264,23 +275,24 @@ class RotationMatrix(Rotation):
|
|
| 264 |
xy_plane: torch.Tensor,
|
| 265 |
eps: float = 1e-12,
|
| 266 |
) -> RotationMatrix:
|
| 267 |
-
|
|
|
|
| 268 |
|
| 269 |
def __getitem__(self, idx: Any) -> RotationMatrix:
|
| 270 |
indices = _index_tuple(idx)
|
| 271 |
-
return RotationMatrix(self._rots[(*indices, slice(None), slice(None))])
|
| 272 |
|
| 273 |
def as_matrix(self) -> RotationMatrix:
|
| 274 |
return self
|
| 275 |
|
| 276 |
def to_3x3(self) -> torch.Tensor:
|
| 277 |
-
return self._rots
|
| 278 |
|
| 279 |
def as_quat(self, normalize: bool = False) -> RotationQuat:
|
| 280 |
m00, m01, m02, m10, m11, m12, m20, m21, m22 = torch.unbind(
|
| 281 |
self._rots.flatten(-2),
|
| 282 |
dim=-1,
|
| 283 |
-
)
|
| 284 |
q_abs = _sqrt_subgradient(
|
| 285 |
torch.stack(
|
| 286 |
(
|
|
@@ -291,7 +303,7 @@ class RotationMatrix(Rotation):
|
|
| 291 |
),
|
| 292 |
dim=-1,
|
| 293 |
)
|
| 294 |
-
)
|
| 295 |
products = torch.stack(
|
| 296 |
(
|
| 297 |
q_abs[..., 0] ** 2,
|
|
@@ -312,29 +324,30 @@ class RotationMatrix(Rotation):
|
|
| 312 |
q_abs[..., 3] ** 2,
|
| 313 |
),
|
| 314 |
dim=-1,
|
| 315 |
-
).unflatten(-1, (4, 4))
|
| 316 |
-
floor = torch.tensor(0.1).to(dtype=q_abs.dtype, device=q_abs.device)
|
| 317 |
-
candidates = products / (2.0 * q_abs[..., None].max(floor))
|
| 318 |
-
best = torch.zeros_like(q_abs, dtype=torch.bool)
|
| 319 |
-
best.scatter_(-1, q_abs.argmax(dim=-1, keepdim=True), True)
|
| 320 |
-
quaternion = candidates[best, :].reshape(q_abs.shape)
|
| 321 |
-
return RotationQuat(quaternion)
|
| 322 |
|
| 323 |
def compose(self, other: RotationMatrix) -> RotationMatrix:
|
| 324 |
with fp32_autocast_context(self.device.type):
|
| 325 |
-
return RotationMatrix(self._rots @ other._rots)
|
| 326 |
|
| 327 |
def convert_compose(self, other: Rotation) -> RotationMatrix:
|
| 328 |
return self.compose(other.as_matrix())
|
| 329 |
|
| 330 |
def apply(self, points: torch.Tensor) -> torch.Tensor:
|
|
|
|
| 331 |
with fp32_autocast_context(self.device.type):
|
| 332 |
if self._rots.shape[-3] == 1:
|
| 333 |
-
return points @ self._rots.transpose(-1, -2).squeeze(-3)
|
| 334 |
-
return torch.einsum("...ij,...j", self._rots, points)
|
| 335 |
|
| 336 |
def invert(self) -> RotationMatrix:
|
| 337 |
-
return RotationMatrix(self._rots.transpose(-1, -2))
|
| 338 |
|
| 339 |
|
| 340 |
@dataclass(frozen=True)
|
|
@@ -378,7 +391,7 @@ class Affine3D:
|
|
| 378 |
|
| 379 |
@property
|
| 380 |
def tensor(self) -> torch.Tensor:
|
| 381 |
-
return torch.cat((self.rot.tensor, self.trans), dim=-1)
|
| 382 |
|
| 383 |
@staticmethod
|
| 384 |
def identity(
|
|
@@ -409,12 +422,13 @@ class Affine3D:
|
|
| 409 |
rotation_type: type[Rotation] = RotationMatrix,
|
| 410 |
**tensor_kwargs,
|
| 411 |
) -> Affine3D:
|
| 412 |
-
translation = torch.randn((*shape, 3), **tensor_kwargs).mul(std)
|
| 413 |
-
rotation = rotation_type.random(shape, **tensor_kwargs)
|
| 414 |
return Affine3D(trans=translation, rot=rotation)
|
| 415 |
|
| 416 |
@staticmethod
|
| 417 |
def from_tensor(tensor: torch.Tensor) -> Affine3D:
|
|
|
|
| 418 |
if not isinstance(tensor, torch.Tensor):
|
| 419 |
raise TypeError("tensor must be a Torch tensor.")
|
| 420 |
if tensor.ndim == 0:
|
|
@@ -426,17 +440,17 @@ class Affine3D:
|
|
| 426 |
"matrix-form affine tensors must have trailing shape (3, 4) or "
|
| 427 |
f"(4, 4), got {tuple(tensor.shape)}."
|
| 428 |
)
|
| 429 |
-
translation = tensor[..., :3, 3]
|
| 430 |
-
rotation: Rotation = RotationMatrix(tensor[..., :3, :3])
|
| 431 |
elif width == 6:
|
| 432 |
-
translation = tensor[..., -3:]
|
| 433 |
-
rotation = RotationQuat(F.pad(tensor[..., :3], (1, 0), value=1))
|
| 434 |
elif width == 7:
|
| 435 |
-
translation = tensor[..., -3:]
|
| 436 |
-
rotation = RotationQuat(tensor[..., :4])
|
| 437 |
elif width == 12:
|
| 438 |
-
translation = tensor[..., -3:]
|
| 439 |
-
rotation = RotationMatrix(tensor[..., :-3].unflatten(-1, (3, 3)))
|
| 440 |
else:
|
| 441 |
raise RuntimeError(
|
| 442 |
f"Cannot detect rotation format from {tensor.shape[-1] - 3}-d flat vector"
|
|
@@ -448,6 +462,7 @@ class Affine3D:
|
|
| 448 |
translation: torch.Tensor,
|
| 449 |
rotation: torch.Tensor,
|
| 450 |
) -> Affine3D:
|
|
|
|
| 451 |
return Affine3D(translation, RotationMatrix(rotation))
|
| 452 |
|
| 453 |
@staticmethod
|
|
@@ -457,13 +472,14 @@ class Affine3D:
|
|
| 457 |
xy_plane: torch.Tensor,
|
| 458 |
eps: float = 1e-10,
|
| 459 |
) -> Affine3D:
|
| 460 |
-
|
| 461 |
-
|
|
|
|
| 462 |
rotation = RotationMatrix.from_graham_schmidt(
|
| 463 |
x_axis,
|
| 464 |
plane_direction,
|
| 465 |
eps,
|
| 466 |
-
)
|
| 467 |
return Affine3D(trans=origin, rot=rotation)
|
| 468 |
|
| 469 |
@staticmethod
|
|
@@ -478,7 +494,7 @@ class Affine3D:
|
|
| 478 |
|
| 479 |
def __getitem__(self, idx: Any) -> Affine3D:
|
| 480 |
indices = _index_tuple(idx)
|
| 481 |
-
translation = self.trans[(*indices, slice(None))]
|
| 482 |
return Affine3D(trans=translation, rot=self.rot[idx])
|
| 483 |
|
| 484 |
def to(self, **kwargs) -> Affine3D:
|
|
@@ -490,9 +506,10 @@ class Affine3D:
|
|
| 490 |
self.rot.detach(**kwargs),
|
| 491 |
)
|
| 492 |
|
| 493 |
-
def tensor_apply(self, func) -> Affine3D:
|
| 494 |
-
|
| 495 |
-
|
|
|
|
| 496 |
|
| 497 |
def as_matrix(self) -> Affine3D:
|
| 498 |
return Affine3D(trans=self.trans, rot=self.rot.as_matrix())
|
|
@@ -509,8 +526,8 @@ class Affine3D:
|
|
| 509 |
autoconvert: bool = False,
|
| 510 |
) -> Affine3D:
|
| 511 |
compose_rotation = self.rot.convert_compose if autoconvert else self.rot.compose
|
| 512 |
-
rotation = compose_rotation(other.rot)
|
| 513 |
-
translation = self.rot.apply(other.trans) + self.trans
|
| 514 |
return Affine3D(trans=translation, rot=rotation)
|
| 515 |
|
| 516 |
def compose_rotation(
|
|
@@ -522,28 +539,31 @@ class Affine3D:
|
|
| 522 |
return Affine3D(trans=self.trans, rot=compose(other))
|
| 523 |
|
| 524 |
def scale(self, value: torch.Tensor | float) -> Affine3D:
|
|
|
|
| 525 |
return Affine3D(self.trans * value, self.rot)
|
| 526 |
|
| 527 |
def mask(self, mask: torch.Tensor, with_zero: bool = False) -> Affine3D:
|
|
|
|
| 528 |
if with_zero:
|
| 529 |
masked = torch.zeros_like(self.tensor).where(
|
| 530 |
mask[..., None],
|
| 531 |
self.tensor,
|
| 532 |
-
)
|
| 533 |
return Affine3D.from_tensor(masked)
|
| 534 |
identity = self.identity(
|
| 535 |
self.shape,
|
| 536 |
rotation_type=type(self.rot),
|
| 537 |
device=self.device,
|
| 538 |
dtype=self.dtype,
|
| 539 |
-
).tensor
|
| 540 |
-
return Affine3D.from_tensor(identity.where(mask[..., None], self.tensor))
|
| 541 |
|
| 542 |
def apply(self, points: torch.Tensor) -> torch.Tensor:
|
| 543 |
-
|
|
|
|
| 544 |
|
| 545 |
def invert(self) -> Affine3D:
|
| 546 |
-
rotation = self.rot.invert()
|
| 547 |
return Affine3D(trans=-rotation.apply(self.trans), rot=rotation)
|
| 548 |
|
| 549 |
|
|
@@ -551,6 +571,7 @@ def build_affine3d_from_coordinates(
|
|
| 551 |
coords: torch.Tensor,
|
| 552 |
) -> tuple[Affine3D, torch.Tensor]:
|
| 553 |
"""Build residue frames from X with shape (b, l, 3, 3)."""
|
|
|
|
| 554 |
|
| 555 |
if not isinstance(coords, torch.Tensor):
|
| 556 |
raise TypeError("coords must be a Torch tensor.")
|
|
@@ -567,39 +588,40 @@ def build_affine3d_from_coordinates(
|
|
| 567 |
dim=-1,
|
| 568 |
),
|
| 569 |
dim=-1,
|
| 570 |
-
)
|
| 571 |
|
| 572 |
def backbone_affine(positions: torch.Tensor) -> Affine3D:
|
| 573 |
-
|
| 574 |
-
|
|
|
|
| 575 |
|
| 576 |
-
coords = coords.clone().float()
|
| 577 |
-
coords[~coord_mask] = 0
|
| 578 |
average = coords.masked_fill(~coord_mask[..., None, None], 0).sum(1) / (
|
| 579 |
coord_mask.sum(-1)[..., None, None] + 1e-8
|
| 580 |
-
)
|
| 581 |
-
average_affine = backbone_affine(average.float()).as_matrix()
|
| 582 |
|
| 583 |
b, length, _, _ = coords.shape
|
| 584 |
-
rotation = average_affine.rot.tensor[..., None, :].expand(b, length, 9)
|
| 585 |
-
translation = average_affine.trans[..., None, :].expand(b, length, 3)
|
| 586 |
identity = RotationMatrix.identity(
|
| 587 |
(b, length),
|
| 588 |
dtype=torch.float32,
|
| 589 |
device=coords.device,
|
| 590 |
requires_grad=False,
|
| 591 |
-
)
|
| 592 |
rotation = rotation.where(
|
| 593 |
coord_mask.any(-1)[..., None, None],
|
| 594 |
identity.tensor,
|
| 595 |
-
)
|
| 596 |
-
missing_frame = Affine3D(translation, RotationMatrix(rotation))
|
| 597 |
|
| 598 |
-
residue_frame = backbone_affine(coords.float())
|
| 599 |
residue_frame = Affine3D.from_tensor(
|
| 600 |
residue_frame.tensor.where(
|
| 601 |
coord_mask[..., None],
|
| 602 |
missing_frame.tensor,
|
| 603 |
)
|
| 604 |
-
)
|
| 605 |
-
return residue_frame, coord_mask
|
|
|
|
| 2 |
|
| 3 |
from __future__ import annotations
|
| 4 |
|
| 5 |
+
import torch
|
| 6 |
+
|
| 7 |
+
from collections.abc import Callable
|
| 8 |
from dataclasses import dataclass
|
| 9 |
from typing import Any, Self
|
|
|
|
|
|
|
| 10 |
from torch.nn import functional as F
|
| 11 |
|
| 12 |
from .esmfold2_misc import fp32_autocast_context
|
|
|
|
| 20 |
|
| 21 |
def _sqrt_subgradient(values: torch.Tensor) -> torch.Tensor:
|
| 22 |
"""Square root with a zero subgradient for non-positive inputs."""
|
| 23 |
+
# values: arbitrary shape; every element keeps its position.
|
| 24 |
|
| 25 |
+
result = torch.zeros_like(values) # values.shape
|
| 26 |
+
positive = values > 0 # values.shape
|
| 27 |
+
result[positive] = torch.sqrt(values[positive]) # (n_positive,) selected elements
|
| 28 |
+
return result # values.shape
|
| 29 |
|
| 30 |
|
| 31 |
def _quat_invert(quaternion: torch.Tensor) -> torch.Tensor:
|
| 32 |
+
# quaternion: (..., 4), real component first.
|
| 33 |
+
conjugate_sign = torch.tensor([1, -1, -1, -1], device=quaternion.device) # (4,)
|
| 34 |
+
return quaternion * conjugate_sign # (..., 4)
|
| 35 |
|
| 36 |
|
| 37 |
def _quat_mult(left: torch.Tensor, right: torch.Tensor) -> torch.Tensor:
|
| 38 |
"""Hamilton product for real-first quaternion tensors."""
|
| 39 |
+
# left, right: (..., 4); their leading dimensions must broadcast.
|
| 40 |
|
| 41 |
+
aw, ax, ay, az = torch.unbind(left, -1) # each left.shape[:-1]
|
| 42 |
+
bw, bx, by, bz = torch.unbind(right, -1) # each right.shape[:-1]
|
| 43 |
return torch.stack(
|
| 44 |
(
|
| 45 |
aw * bw - ax * bx - ay * by - az * bz,
|
|
|
|
| 48 |
aw * bz + ax * by - ay * bx + az * bw,
|
| 49 |
),
|
| 50 |
-1,
|
| 51 |
+
) # (..., 4), broadcast leading dimensions
|
| 52 |
|
| 53 |
|
| 54 |
def _quat_rotation(
|
|
|
|
| 56 |
points: torch.Tensor,
|
| 57 |
) -> torch.Tensor:
|
| 58 |
"""Rotate points using normalized real-first quaternions."""
|
| 59 |
+
# quaternion: (..., 4); points: (..., 3); leading dimensions broadcast.
|
| 60 |
|
| 61 |
+
aw, ax, ay, az = torch.unbind(quaternion, -1) # each quaternion.shape[:-1]
|
| 62 |
+
bx, by, bz = torch.unbind(points, -1) # each points.shape[:-1]
|
| 63 |
product = torch.stack(
|
| 64 |
(
|
| 65 |
-ax * bx - ay * by - az * bz,
|
|
|
|
| 68 |
aw * bz + ax * by - ay * bx,
|
| 69 |
),
|
| 70 |
-1,
|
| 71 |
+
) # (..., 4), broadcast leading dimensions
|
| 72 |
+
return _quat_mult(product, _quat_invert(quaternion))[..., 1:] # (..., 3)
|
| 73 |
|
| 74 |
|
| 75 |
def _graham_schmidt(
|
|
|
|
| 78 |
eps: float = 1e-12,
|
| 79 |
) -> torch.Tensor:
|
| 80 |
"""Construct a right-handed orthonormal frame from two directions."""
|
| 81 |
+
# x_axis, xy_plane: (..., 3); ... contains arbitrary frame batch axes.
|
| 82 |
|
| 83 |
with fp32_autocast_context(x_axis.device.type):
|
| 84 |
+
e1 = xy_plane # (..., 3)
|
| 85 |
+
denominator = torch.sqrt((x_axis**2).sum(dim=-1, keepdim=True) + eps) # (..., 1)
|
| 86 |
+
x_axis = x_axis / denominator # (..., 3)
|
| 87 |
+
projection = (x_axis * e1).sum(dim=-1, keepdim=True) # (..., 1)
|
| 88 |
+
e1 = e1 - x_axis * projection # (..., 3)
|
| 89 |
+
denominator = torch.sqrt((e1**2).sum(dim=-1, keepdim=True) + eps) # (..., 1)
|
| 90 |
+
e1 = e1 / denominator # (..., 3)
|
| 91 |
+
e2 = torch.cross(x_axis, e1, dim=-1) # (..., 3)
|
| 92 |
+
return torch.stack([x_axis, e1, e2], dim=-1) # (..., 3, 3)
|
| 93 |
|
| 94 |
|
| 95 |
class Rotation:
|
|
|
|
| 135 |
|
| 136 |
@classmethod
|
| 137 |
def _from_tensor(cls, tensor: torch.Tensor) -> Self:
|
| 138 |
+
# tensor: (..., r), where r is the concrete rotation representation width.
|
| 139 |
return cls(tensor) # type: ignore[call-arg]
|
| 140 |
|
| 141 |
def to(self, **kwargs) -> Self:
|
|
|
|
| 144 |
def detach(self, *args, **kwargs) -> Self:
|
| 145 |
return self._from_tensor(self.tensor.detach(**kwargs))
|
| 146 |
|
| 147 |
+
def tensor_apply(self, func: Callable[[torch.Tensor], torch.Tensor]) -> Self:
|
| 148 |
+
# self.tensor: (..., r); func receives one (...) component and determines its output shape.
|
| 149 |
+
transformed = [func(component) for component in self.tensor.unbind(dim=-1)] # one func-shaped array per rotation component
|
| 150 |
+
return self._from_tensor(torch.stack(transformed, dim=-1)) # rotation tensor: (*func_output_shape, n_components)
|
| 151 |
|
| 152 |
|
| 153 |
class RotationQuat(Rotation):
|
| 154 |
"""A rotation represented by a real-first quaternion."""
|
| 155 |
|
| 156 |
+
def __init__(self, quats: torch.Tensor, normalized: bool = False) -> None:
|
| 157 |
+
# quats: (..., 4); ... is the rotation batch shape.
|
| 158 |
if not isinstance(quats, torch.Tensor):
|
| 159 |
raise TypeError("quats must be a Torch tensor.")
|
| 160 |
if quats.ndim == 0 or quats.shape[-1] != 4:
|
|
|
|
| 165 |
raise TypeError("normalized must be a boolean.")
|
| 166 |
self._normalized = normalized
|
| 167 |
if normalized:
|
| 168 |
+
quats = F.normalize(quats.to(torch.float32), dim=-1) # (..., 4)
|
| 169 |
+
self._quats = quats.where(quats[..., :1] >= 0, -quats) # (..., 4)
|
| 170 |
else:
|
| 171 |
+
self._quats = quats.to(torch.float32) # (..., 4)
|
| 172 |
|
| 173 |
@property
|
| 174 |
def tensor(self) -> torch.Tensor:
|
| 175 |
+
return self._quats # (..., 4)
|
| 176 |
|
| 177 |
@property
|
| 178 |
def shape(self) -> torch.Size:
|
| 179 |
return self._quats.shape[:-1]
|
| 180 |
|
| 181 |
@classmethod
|
| 182 |
+
def identity(cls, shape: tuple[int, ...], **tensor_kwargs) -> RotationQuat:
|
| 183 |
+
quaternions = torch.ones((*shape, 4), **tensor_kwargs) # (*shape, 4)
|
| 184 |
+
selector = torch.tensor([1, 0, 0, 0], device=quaternions.device) # (4,)
|
| 185 |
+
return cls(quaternions * selector) # quaternion tensor: (*shape, 4)
|
| 186 |
|
| 187 |
@classmethod
|
| 188 |
+
def random(cls, shape: tuple[int, ...], **tensor_kwargs) -> RotationQuat:
|
| 189 |
+
return cls(torch.randn((*shape, 4), **tensor_kwargs), normalized=True) # quaternion tensor: (*shape, 4)
|
| 190 |
|
| 191 |
def __getitem__(self, idx: Any) -> RotationQuat:
|
| 192 |
indices = _index_tuple(idx)
|
| 193 |
+
return RotationQuat(self._quats[(*indices, slice(None))]) # quaternion tensor: (*indexed_shape, 4)
|
| 194 |
|
| 195 |
def normalized(self) -> RotationQuat:
|
| 196 |
if self._normalized:
|
|
|
|
| 201 |
return self
|
| 202 |
|
| 203 |
def as_matrix(self) -> RotationMatrix:
|
| 204 |
+
quaternion = self.normalized().tensor # (..., 4)
|
| 205 |
+
r, i, j, k = torch.unbind(quaternion, -1) # each (...)
|
| 206 |
+
scale = 2.0 / torch.linalg.norm(quaternion, dim=-1) # (...)
|
| 207 |
elements = torch.stack(
|
| 208 |
(
|
| 209 |
1 - scale * (j * j + k * k),
|
|
|
|
| 217 |
1 - scale * (i * i + j * j),
|
| 218 |
),
|
| 219 |
-1,
|
| 220 |
+
) # (..., 9)
|
| 221 |
+
return RotationMatrix(elements.reshape((*quaternion.shape[:-1], 3, 3))) # rotation tensor: (..., 3, 3)
|
| 222 |
|
| 223 |
def compose(self, other: RotationQuat) -> RotationQuat:
|
| 224 |
with fp32_autocast_context(self.device.type):
|
| 225 |
+
return RotationQuat(_quat_mult(self._quats, other._quats)) # quaternion tensor: (..., 4), broadcast leading shapes
|
| 226 |
|
| 227 |
def convert_compose(self, other: Rotation) -> RotationQuat:
|
| 228 |
return self.compose(other.as_quat())
|
| 229 |
|
| 230 |
def apply(self, points: torch.Tensor) -> torch.Tensor:
|
| 231 |
+
# points: (..., 3); quaternion and point leading shapes must broadcast.
|
| 232 |
+
return _quat_rotation(self.normalized()._quats, points) # (..., 3), broadcast rotation/point leading shapes
|
| 233 |
|
| 234 |
def invert(self) -> RotationQuat:
|
| 235 |
+
return RotationQuat(_quat_invert(self._quats)) # quaternion tensor: (..., 4)
|
| 236 |
|
| 237 |
|
| 238 |
class RotationMatrix(Rotation):
|
| 239 |
"""A rotation represented by a dense FP32 matrix."""
|
| 240 |
|
| 241 |
+
def __init__(self, rots: torch.Tensor) -> None:
|
| 242 |
+
# rots: (..., 9) or (..., 3, 3); ... is the rotation batch shape.
|
| 243 |
if not isinstance(rots, torch.Tensor):
|
| 244 |
raise TypeError("rots must be a Torch tensor.")
|
| 245 |
if rots.ndim > 0 and rots.shape[-1] == 9:
|
| 246 |
+
rots = rots.unflatten(-1, (3, 3)) # (..., 3, 3)
|
| 247 |
if rots.ndim < 2 or rots.shape[-2:] != (3, 3):
|
| 248 |
raise ValueError(
|
| 249 |
"rots must have trailing shape (3, 3) or flattened width 9, got "
|
| 250 |
f"shape {tuple(rots.shape)}."
|
| 251 |
)
|
| 252 |
+
self._rots = rots.to(torch.float32) # (..., 3, 3)
|
| 253 |
|
| 254 |
@property
|
| 255 |
def tensor(self) -> torch.Tensor:
|
| 256 |
+
return self._rots.flatten(-2) # (..., 9)
|
| 257 |
|
| 258 |
@property
|
| 259 |
def shape(self) -> torch.Size:
|
| 260 |
return self._rots.shape[:-2]
|
| 261 |
|
| 262 |
@classmethod
|
| 263 |
+
def identity(cls, shape: tuple[int, ...], **tensor_kwargs) -> RotationMatrix:
|
| 264 |
+
matrix = torch.eye(3, **tensor_kwargs) # (3, 3)
|
| 265 |
+
matrix = matrix.view(*(1 for _ in shape), 3, 3) # (*singleton_shape, 3, 3); one singleton per requested axis
|
| 266 |
+
return cls(matrix.expand(*shape, -1, -1)) # rotation tensor: (*shape, 3, 3)
|
| 267 |
|
| 268 |
@classmethod
|
| 269 |
+
def random(cls, shape: tuple[int, ...], **tensor_kwargs) -> RotationMatrix:
|
| 270 |
return RotationQuat.random(shape, **tensor_kwargs).as_matrix()
|
| 271 |
|
| 272 |
@staticmethod
|
|
|
|
| 275 |
xy_plane: torch.Tensor,
|
| 276 |
eps: float = 1e-12,
|
| 277 |
) -> RotationMatrix:
|
| 278 |
+
# x_axis, xy_plane: (..., 3).
|
| 279 |
+
return RotationMatrix(_graham_schmidt(x_axis, xy_plane, eps)) # rotation tensor: (..., 3, 3)
|
| 280 |
|
| 281 |
def __getitem__(self, idx: Any) -> RotationMatrix:
|
| 282 |
indices = _index_tuple(idx)
|
| 283 |
+
return RotationMatrix(self._rots[(*indices, slice(None), slice(None))]) # rotation tensor: (*indexed_shape, 3, 3)
|
| 284 |
|
| 285 |
def as_matrix(self) -> RotationMatrix:
|
| 286 |
return self
|
| 287 |
|
| 288 |
def to_3x3(self) -> torch.Tensor:
|
| 289 |
+
return self._rots # (..., 3, 3)
|
| 290 |
|
| 291 |
def as_quat(self, normalize: bool = False) -> RotationQuat:
|
| 292 |
m00, m01, m02, m10, m11, m12, m20, m21, m22 = torch.unbind(
|
| 293 |
self._rots.flatten(-2),
|
| 294 |
dim=-1,
|
| 295 |
+
) # each (...)
|
| 296 |
q_abs = _sqrt_subgradient(
|
| 297 |
torch.stack(
|
| 298 |
(
|
|
|
|
| 303 |
),
|
| 304 |
dim=-1,
|
| 305 |
)
|
| 306 |
+
) # (..., 4)
|
| 307 |
products = torch.stack(
|
| 308 |
(
|
| 309 |
q_abs[..., 0] ** 2,
|
|
|
|
| 324 |
q_abs[..., 3] ** 2,
|
| 325 |
),
|
| 326 |
dim=-1,
|
| 327 |
+
).unflatten(-1, (4, 4)) # (..., 4, 4)
|
| 328 |
+
floor = torch.tensor(0.1).to(dtype=q_abs.dtype, device=q_abs.device) # ()
|
| 329 |
+
candidates = products / (2.0 * q_abs[..., None].max(floor)) # (..., 4, 4)
|
| 330 |
+
best = torch.zeros_like(q_abs, dtype=torch.bool) # (..., 4)
|
| 331 |
+
best.scatter_(-1, q_abs.argmax(dim=-1, keepdim=True), True) # (..., 4); selected index shape (..., 1)
|
| 332 |
+
quaternion = candidates[best, :].reshape(q_abs.shape) # (..., 4)
|
| 333 |
+
return RotationQuat(quaternion) # quaternion tensor: (..., 4)
|
| 334 |
|
| 335 |
def compose(self, other: RotationMatrix) -> RotationMatrix:
|
| 336 |
with fp32_autocast_context(self.device.type):
|
| 337 |
+
return RotationMatrix(self._rots @ other._rots) # rotation tensor: (..., 3, 3), broadcast leading shapes
|
| 338 |
|
| 339 |
def convert_compose(self, other: Rotation) -> RotationMatrix:
|
| 340 |
return self.compose(other.as_matrix())
|
| 341 |
|
| 342 |
def apply(self, points: torch.Tensor) -> torch.Tensor:
|
| 343 |
+
# points: (..., 3); matrix leading shapes broadcast with point axes.
|
| 344 |
with fp32_autocast_context(self.device.type):
|
| 345 |
if self._rots.shape[-3] == 1:
|
| 346 |
+
return points @ self._rots.transpose(-1, -2).squeeze(-3) # (..., 3), with leading axes set by matmul broadcasting
|
| 347 |
+
return torch.einsum("...ij,...j", self._rots, points) # (..., 3), broadcast leading shapes
|
| 348 |
|
| 349 |
def invert(self) -> RotationMatrix:
|
| 350 |
+
return RotationMatrix(self._rots.transpose(-1, -2)) # rotation tensor: (..., 3, 3)
|
| 351 |
|
| 352 |
|
| 353 |
@dataclass(frozen=True)
|
|
|
|
| 391 |
|
| 392 |
@property
|
| 393 |
def tensor(self) -> torch.Tensor:
|
| 394 |
+
return torch.cat((self.rot.tensor, self.trans), dim=-1) # (..., r + 3); r is 4 for quaternions or 9 for matrices
|
| 395 |
|
| 396 |
@staticmethod
|
| 397 |
def identity(
|
|
|
|
| 422 |
rotation_type: type[Rotation] = RotationMatrix,
|
| 423 |
**tensor_kwargs,
|
| 424 |
) -> Affine3D:
|
| 425 |
+
translation = torch.randn((*shape, 3), **tensor_kwargs).mul(std) # (*shape, 3)
|
| 426 |
+
rotation = rotation_type.random(shape, **tensor_kwargs) # rotation batch shape: shape
|
| 427 |
return Affine3D(trans=translation, rot=rotation)
|
| 428 |
|
| 429 |
@staticmethod
|
| 430 |
def from_tensor(tensor: torch.Tensor) -> Affine3D:
|
| 431 |
+
# tensor: (..., 6/7/12), (..., 3, 4), or (..., 4, 4); ... is frame batch shape.
|
| 432 |
if not isinstance(tensor, torch.Tensor):
|
| 433 |
raise TypeError("tensor must be a Torch tensor.")
|
| 434 |
if tensor.ndim == 0:
|
|
|
|
| 440 |
"matrix-form affine tensors must have trailing shape (3, 4) or "
|
| 441 |
f"(4, 4), got {tuple(tensor.shape)}."
|
| 442 |
)
|
| 443 |
+
translation = tensor[..., :3, 3] # (..., 3)
|
| 444 |
+
rotation: Rotation = RotationMatrix(tensor[..., :3, :3]) # rotation tensor: (..., 3, 3)
|
| 445 |
elif width == 6:
|
| 446 |
+
translation = tensor[..., -3:] # (..., 3)
|
| 447 |
+
rotation = RotationQuat(F.pad(tensor[..., :3], (1, 0), value=1)) # quaternion tensor: (..., 4)
|
| 448 |
elif width == 7:
|
| 449 |
+
translation = tensor[..., -3:] # (..., 3)
|
| 450 |
+
rotation = RotationQuat(tensor[..., :4]) # quaternion tensor: (..., 4)
|
| 451 |
elif width == 12:
|
| 452 |
+
translation = tensor[..., -3:] # (..., 3)
|
| 453 |
+
rotation = RotationMatrix(tensor[..., :-3].unflatten(-1, (3, 3))) # rotation tensor: (..., 3, 3)
|
| 454 |
else:
|
| 455 |
raise RuntimeError(
|
| 456 |
f"Cannot detect rotation format from {tensor.shape[-1] - 3}-d flat vector"
|
|
|
|
| 462 |
translation: torch.Tensor,
|
| 463 |
rotation: torch.Tensor,
|
| 464 |
) -> Affine3D:
|
| 465 |
+
# translation: (..., 3); rotation: (..., 3, 3) or (..., 9).
|
| 466 |
return Affine3D(translation, RotationMatrix(rotation))
|
| 467 |
|
| 468 |
@staticmethod
|
|
|
|
| 472 |
xy_plane: torch.Tensor,
|
| 473 |
eps: float = 1e-10,
|
| 474 |
) -> Affine3D:
|
| 475 |
+
# neg_x_axis, origin, xy_plane: (..., 3), with broadcast-compatible leading axes.
|
| 476 |
+
x_axis = origin - neg_x_axis # (..., 3)
|
| 477 |
+
plane_direction = xy_plane - origin # (..., 3)
|
| 478 |
rotation = RotationMatrix.from_graham_schmidt(
|
| 479 |
x_axis,
|
| 480 |
plane_direction,
|
| 481 |
eps,
|
| 482 |
+
) # rotation tensor: (..., 3, 3)
|
| 483 |
return Affine3D(trans=origin, rot=rotation)
|
| 484 |
|
| 485 |
@staticmethod
|
|
|
|
| 494 |
|
| 495 |
def __getitem__(self, idx: Any) -> Affine3D:
|
| 496 |
indices = _index_tuple(idx)
|
| 497 |
+
translation = self.trans[(*indices, slice(None))] # (*indexed_shape, 3)
|
| 498 |
return Affine3D(trans=translation, rot=self.rot[idx])
|
| 499 |
|
| 500 |
def to(self, **kwargs) -> Affine3D:
|
|
|
|
| 506 |
self.rot.detach(**kwargs),
|
| 507 |
)
|
| 508 |
|
| 509 |
+
def tensor_apply(self, func: Callable[[torch.Tensor], torch.Tensor]) -> Affine3D:
|
| 510 |
+
# self.tensor: (..., r + 3); func maps each (...) component to a common output shape.
|
| 511 |
+
components = [func(value) for value in self.tensor.unbind(dim=-1)] # one func-shaped array per affine component
|
| 512 |
+
return Affine3D.from_tensor(torch.stack(components, dim=-1)) # affine tensor: (*func_output_shape, r + 3)
|
| 513 |
|
| 514 |
def as_matrix(self) -> Affine3D:
|
| 515 |
return Affine3D(trans=self.trans, rot=self.rot.as_matrix())
|
|
|
|
| 526 |
autoconvert: bool = False,
|
| 527 |
) -> Affine3D:
|
| 528 |
compose_rotation = self.rot.convert_compose if autoconvert else self.rot.compose
|
| 529 |
+
rotation = compose_rotation(other.rot) # rotation batch shape: broadcast(self.shape, other.shape)
|
| 530 |
+
translation = self.rot.apply(other.trans) + self.trans # (..., 3), broadcast transform shapes
|
| 531 |
return Affine3D(trans=translation, rot=rotation)
|
| 532 |
|
| 533 |
def compose_rotation(
|
|
|
|
| 539 |
return Affine3D(trans=self.trans, rot=compose(other))
|
| 540 |
|
| 541 |
def scale(self, value: torch.Tensor | float) -> Affine3D:
|
| 542 |
+
# value must broadcast with translation (..., 3); rotation shape is retained.
|
| 543 |
return Affine3D(self.trans * value, self.rot)
|
| 544 |
|
| 545 |
def mask(self, mask: torch.Tensor, with_zero: bool = False) -> Affine3D:
|
| 546 |
+
# mask: self.shape; affine tensor: (*self.shape, r + 3).
|
| 547 |
if with_zero:
|
| 548 |
masked = torch.zeros_like(self.tensor).where(
|
| 549 |
mask[..., None],
|
| 550 |
self.tensor,
|
| 551 |
+
) # (..., r + 3)
|
| 552 |
return Affine3D.from_tensor(masked)
|
| 553 |
identity = self.identity(
|
| 554 |
self.shape,
|
| 555 |
rotation_type=type(self.rot),
|
| 556 |
device=self.device,
|
| 557 |
dtype=self.dtype,
|
| 558 |
+
).tensor # (..., r + 3)
|
| 559 |
+
return Affine3D.from_tensor(identity.where(mask[..., None], self.tensor)) # affine batch shape: self.shape
|
| 560 |
|
| 561 |
def apply(self, points: torch.Tensor) -> torch.Tensor:
|
| 562 |
+
# points: (..., 3); result uses broadcast transform/point leading axes.
|
| 563 |
+
return self.rot.apply(points) + self.trans # (..., 3), broadcast transform/point leading shapes
|
| 564 |
|
| 565 |
def invert(self) -> Affine3D:
|
| 566 |
+
rotation = self.rot.invert() # rotation batch shape: self.shape
|
| 567 |
return Affine3D(trans=-rotation.apply(self.trans), rot=rotation)
|
| 568 |
|
| 569 |
|
|
|
|
| 571 |
coords: torch.Tensor,
|
| 572 |
) -> tuple[Affine3D, torch.Tensor]:
|
| 573 |
"""Build residue frames from X with shape (b, l, 3, 3)."""
|
| 574 |
+
# coords: (b, l, 3, 3); final axes are N/CA/C atoms and xyz coordinates.
|
| 575 |
|
| 576 |
if not isinstance(coords, torch.Tensor):
|
| 577 |
raise TypeError("coords must be a Torch tensor.")
|
|
|
|
| 588 |
dim=-1,
|
| 589 |
),
|
| 590 |
dim=-1,
|
| 591 |
+
) # (b, l)
|
| 592 |
|
| 593 |
def backbone_affine(positions: torch.Tensor) -> Affine3D:
|
| 594 |
+
# positions: (..., 3, 3); final axes are N/CA/C atoms and xyz coordinates.
|
| 595 |
+
n, ca, c = positions.unbind(dim=-2) # each (..., 3)
|
| 596 |
+
return Affine3D.from_graham_schmidt(c, ca, n) # affine batch shape: positions.shape[:-2]
|
| 597 |
|
| 598 |
+
coords = coords.clone().float() # (b, l, 3, 3)
|
| 599 |
+
coords[~coord_mask] = 0 # (n_missing, 3, 3) selected residue coordinates
|
| 600 |
average = coords.masked_fill(~coord_mask[..., None, None], 0).sum(1) / (
|
| 601 |
coord_mask.sum(-1)[..., None, None] + 1e-8
|
| 602 |
+
) # (b, 3, 3)
|
| 603 |
+
average_affine = backbone_affine(average.float()).as_matrix() # affine batch shape: (b,)
|
| 604 |
|
| 605 |
b, length, _, _ = coords.shape
|
| 606 |
+
rotation = average_affine.rot.tensor[..., None, :].expand(b, length, 9) # (b, l, 9)
|
| 607 |
+
translation = average_affine.trans[..., None, :].expand(b, length, 3) # (b, l, 3)
|
| 608 |
identity = RotationMatrix.identity(
|
| 609 |
(b, length),
|
| 610 |
dtype=torch.float32,
|
| 611 |
device=coords.device,
|
| 612 |
requires_grad=False,
|
| 613 |
+
) # rotation batch shape: (b, l)
|
| 614 |
rotation = rotation.where(
|
| 615 |
coord_mask.any(-1)[..., None, None],
|
| 616 |
identity.tensor,
|
| 617 |
+
) # (b, l, 9)
|
| 618 |
+
missing_frame = Affine3D(translation, RotationMatrix(rotation)) # affine batch shape: (b, l)
|
| 619 |
|
| 620 |
+
residue_frame = backbone_affine(coords.float()) # affine batch shape: (b, l)
|
| 621 |
residue_frame = Affine3D.from_tensor(
|
| 622 |
residue_frame.tensor.where(
|
| 623 |
coord_mask[..., None],
|
| 624 |
missing_frame.tensor,
|
| 625 |
)
|
| 626 |
+
) # affine batch shape: (b, l)
|
| 627 |
+
return residue_frame, coord_mask # affine batch shape (b, l), mask (b, l)
|
fastplms/models/esmfold2/esmfold2_aligner.py
CHANGED
|
@@ -2,11 +2,11 @@
|
|
| 2 |
|
| 3 |
from __future__ import annotations
|
| 4 |
|
| 5 |
-
from dataclasses import Field, replace
|
| 6 |
-
from typing import Any, ClassVar, Protocol, TypeVar
|
| 7 |
-
|
| 8 |
import numpy as np
|
| 9 |
import torch
|
|
|
|
|
|
|
|
|
|
| 10 |
from torch import Tensor
|
| 11 |
|
| 12 |
from .esmfold2_protein_structure import compute_affine_and_rmsd
|
|
@@ -30,18 +30,19 @@ AlignableT = TypeVar("AlignableT", bound=Alignable)
|
|
| 30 |
|
| 31 |
|
| 32 |
def _coordinate_batch(structure: Alignable) -> Tensor:
|
| 33 |
-
|
|
|
|
| 34 |
|
| 35 |
|
| 36 |
def _shared_atom_mask(mobile: Alignable, target: Alignable, backbone_only: bool) -> Tensor:
|
| 37 |
shared = np.asarray(mobile.atom37_mask, dtype=bool) & np.asarray(
|
| 38 |
target.atom37_mask,
|
| 39 |
dtype=bool,
|
| 40 |
-
)
|
| 41 |
if backbone_only:
|
| 42 |
-
shared = shared.copy()
|
| 43 |
-
shared[:, 3:] = False
|
| 44 |
-
return torch.from_numpy(shared).unsqueeze(0)
|
| 45 |
|
| 46 |
|
| 47 |
class Aligner:
|
|
@@ -57,16 +58,16 @@ class Aligner:
|
|
| 57 |
if len(mobile) != len(target):
|
| 58 |
raise AssertionError("mobile and target must contain the same residue count")
|
| 59 |
|
| 60 |
-
mobile_coordinates = _coordinate_batch(mobile)
|
| 61 |
-
target_coordinates = _coordinate_batch(target)
|
| 62 |
if use_reflection:
|
| 63 |
-
target_coordinates = -target_coordinates
|
| 64 |
-
atom_mask = _shared_atom_mask(mobile, target, only_use_backbone)
|
| 65 |
self._affine3D, rmsd = compute_affine_and_rmsd(
|
| 66 |
mobile_coordinates,
|
| 67 |
target_coordinates,
|
| 68 |
atom_exists_mask=atom_mask,
|
| 69 |
-
)
|
| 70 |
self._rmsd = rmsd.item()
|
| 71 |
|
| 72 |
@property
|
|
@@ -76,12 +77,12 @@ class Aligner:
|
|
| 76 |
def apply(self, mobile: AlignableT) -> AlignableT:
|
| 77 |
"""Return a dataclass copy with all present atom coordinates aligned."""
|
| 78 |
|
| 79 |
-
present = np.asarray(mobile.atom37_mask, dtype=bool)
|
| 80 |
packed = torch.as_tensor(
|
| 81 |
mobile.atom37_positions[present],
|
| 82 |
dtype=torch.float32,
|
| 83 |
-
).unsqueeze(0)
|
| 84 |
-
aligned = self._affine3D.apply(packed).squeeze(0).cpu().numpy()
|
| 85 |
-
atom37_positions = np.full_like(mobile.atom37_positions, np.nan)
|
| 86 |
-
atom37_positions[present] = aligned
|
| 87 |
return replace(mobile, atom37_positions=atom37_positions)
|
|
|
|
| 2 |
|
| 3 |
from __future__ import annotations
|
| 4 |
|
|
|
|
|
|
|
|
|
|
| 5 |
import numpy as np
|
| 6 |
import torch
|
| 7 |
+
|
| 8 |
+
from dataclasses import Field, replace
|
| 9 |
+
from typing import Any, ClassVar, Protocol, TypeVar
|
| 10 |
from torch import Tensor
|
| 11 |
|
| 12 |
from .esmfold2_protein_structure import compute_affine_and_rmsd
|
|
|
|
| 30 |
|
| 31 |
|
| 32 |
def _coordinate_batch(structure: Alignable) -> Tensor:
|
| 33 |
+
# l is the residue count; the atom37 table has three Cartesian coordinates.
|
| 34 |
+
return torch.as_tensor(structure.atom37_positions, dtype=torch.double).unsqueeze(0) # (1, l, 37, 3)
|
| 35 |
|
| 36 |
|
| 37 |
def _shared_atom_mask(mobile: Alignable, target: Alignable, backbone_only: bool) -> Tensor:
|
| 38 |
shared = np.asarray(mobile.atom37_mask, dtype=bool) & np.asarray(
|
| 39 |
target.atom37_mask,
|
| 40 |
dtype=bool,
|
| 41 |
+
) # (l, 37)
|
| 42 |
if backbone_only:
|
| 43 |
+
shared = shared.copy() # (l, 37)
|
| 44 |
+
shared[:, 3:] = False # (l, 34); retain N, CA, C.
|
| 45 |
+
return torch.from_numpy(shared).unsqueeze(0) # (1, l, 37)
|
| 46 |
|
| 47 |
|
| 48 |
class Aligner:
|
|
|
|
| 58 |
if len(mobile) != len(target):
|
| 59 |
raise AssertionError("mobile and target must contain the same residue count")
|
| 60 |
|
| 61 |
+
mobile_coordinates = _coordinate_batch(mobile) # (1, l, 37, 3)
|
| 62 |
+
target_coordinates = _coordinate_batch(target) # (1, l, 37, 3)
|
| 63 |
if use_reflection:
|
| 64 |
+
target_coordinates = -target_coordinates # (1, l, 37, 3)
|
| 65 |
+
atom_mask = _shared_atom_mask(mobile, target, only_use_backbone) # (1, l, 37)
|
| 66 |
self._affine3D, rmsd = compute_affine_and_rmsd(
|
| 67 |
mobile_coordinates,
|
| 68 |
target_coordinates,
|
| 69 |
atom_exists_mask=atom_mask,
|
| 70 |
+
) # affine shape: (1, 1); rmsd: ()
|
| 71 |
self._rmsd = rmsd.item()
|
| 72 |
|
| 73 |
@property
|
|
|
|
| 77 |
def apply(self, mobile: AlignableT) -> AlignableT:
|
| 78 |
"""Return a dataclass copy with all present atom coordinates aligned."""
|
| 79 |
|
| 80 |
+
present = np.asarray(mobile.atom37_mask, dtype=bool) # (l, 37)
|
| 81 |
packed = torch.as_tensor(
|
| 82 |
mobile.atom37_positions[present],
|
| 83 |
dtype=torch.float32,
|
| 84 |
+
).unsqueeze(0) # (1, n, 3), n is the number of present atoms.
|
| 85 |
+
aligned = self._affine3D.apply(packed).squeeze(0).cpu().numpy() # (n, 3)
|
| 86 |
+
atom37_positions = np.full_like(mobile.atom37_positions, np.nan) # (l, 37, 3)
|
| 87 |
+
atom37_positions[present] = aligned # (n, 3)
|
| 88 |
return replace(mobile, atom37_positions=atom37_positions)
|
fastplms/models/esmfold2/esmfold2_atom_indexer.py
CHANGED
|
@@ -2,11 +2,11 @@
|
|
| 2 |
|
| 3 |
from __future__ import annotations
|
| 4 |
|
|
|
|
|
|
|
| 5 |
from operator import attrgetter
|
| 6 |
from typing import Any
|
| 7 |
|
| 8 |
-
import numpy as np
|
| 9 |
-
|
| 10 |
from .esmfold2_protein_structure import index_by_atom_name
|
| 11 |
|
| 12 |
|
|
@@ -19,7 +19,7 @@ class AtomIndexer:
|
|
| 19 |
|
| 20 |
__slots__ = ("_get_property", "dim", "property", "structure")
|
| 21 |
|
| 22 |
-
def __init__(self, structure: Any, property: str, dim: int):
|
| 23 |
self.structure = structure
|
| 24 |
self.property = property
|
| 25 |
self.dim = dim
|
|
|
|
| 2 |
|
| 3 |
from __future__ import annotations
|
| 4 |
|
| 5 |
+
import numpy as np
|
| 6 |
+
|
| 7 |
from operator import attrgetter
|
| 8 |
from typing import Any
|
| 9 |
|
|
|
|
|
|
|
| 10 |
from .esmfold2_protein_structure import index_by_atom_name
|
| 11 |
|
| 12 |
|
|
|
|
| 19 |
|
| 20 |
__slots__ = ("_get_property", "dim", "property", "structure")
|
| 21 |
|
| 22 |
+
def __init__(self, structure: Any, property: str, dim: int) -> None:
|
| 23 |
self.structure = structure
|
| 24 |
self.property = property
|
| 25 |
self.dim = dim
|
fastplms/models/esmfold2/esmfold2_conformers.py
CHANGED
|
@@ -11,21 +11,21 @@ import os
|
|
| 11 |
import pickle
|
| 12 |
import stat
|
| 13 |
import tempfile
|
|
|
|
|
|
|
| 14 |
from collections.abc import Iterator
|
| 15 |
from contextlib import contextmanager
|
| 16 |
from dataclasses import dataclass
|
| 17 |
from hashlib import file_digest
|
| 18 |
from pathlib import Path
|
| 19 |
from typing import Any, BinaryIO
|
| 20 |
-
|
| 21 |
-
import numpy as np
|
| 22 |
from huggingface_hub import hf_hub_download
|
| 23 |
from huggingface_hub.constants import HF_HUB_CACHE
|
| 24 |
|
| 25 |
from fastplms.registry import RuntimeAsset, get_model_registry
|
| 26 |
-
|
| 27 |
from .esmfold2_constants import RES_TYPE_TO_CCD
|
| 28 |
|
|
|
|
| 29 |
_CCD_ENVIRONMENT_VARIABLE = "ESMCFOLD_CCD_PATH"
|
| 30 |
_CCD_ASSET_ID = "esmfold2_ccd"
|
| 31 |
|
|
|
|
| 11 |
import pickle
|
| 12 |
import stat
|
| 13 |
import tempfile
|
| 14 |
+
import numpy as np
|
| 15 |
+
|
| 16 |
from collections.abc import Iterator
|
| 17 |
from contextlib import contextmanager
|
| 18 |
from dataclasses import dataclass
|
| 19 |
from hashlib import file_digest
|
| 20 |
from pathlib import Path
|
| 21 |
from typing import Any, BinaryIO
|
|
|
|
|
|
|
| 22 |
from huggingface_hub import hf_hub_download
|
| 23 |
from huggingface_hub.constants import HF_HUB_CACHE
|
| 24 |
|
| 25 |
from fastplms.registry import RuntimeAsset, get_model_registry
|
|
|
|
| 26 |
from .esmfold2_constants import RES_TYPE_TO_CCD
|
| 27 |
|
| 28 |
+
|
| 29 |
_CCD_ENVIRONMENT_VARIABLE = "ESMCFOLD_CCD_PATH"
|
| 30 |
_CCD_ASSET_ID = "esmfold2_ccd"
|
| 31 |
|
fastplms/models/esmfold2/esmfold2_input_builder.py
CHANGED
|
@@ -2,14 +2,15 @@
|
|
| 2 |
|
| 3 |
from __future__ import annotations
|
| 4 |
|
|
|
|
|
|
|
| 5 |
from collections.abc import Sequence
|
| 6 |
from dataclasses import dataclass
|
| 7 |
from typing import Any, TypeAlias
|
| 8 |
|
| 9 |
-
import numpy as np
|
| 10 |
-
|
| 11 |
from .esmfold2_msa import MSA
|
| 12 |
|
|
|
|
| 13 |
MSAInput: TypeAlias = MSA | None
|
| 14 |
|
| 15 |
|
|
|
|
| 2 |
|
| 3 |
from __future__ import annotations
|
| 4 |
|
| 5 |
+
import numpy as np
|
| 6 |
+
|
| 7 |
from collections.abc import Sequence
|
| 8 |
from dataclasses import dataclass
|
| 9 |
from typing import Any, TypeAlias
|
| 10 |
|
|
|
|
|
|
|
| 11 |
from .esmfold2_msa import MSA
|
| 12 |
|
| 13 |
+
|
| 14 |
MSAInput: TypeAlias = MSA | None
|
| 15 |
|
| 16 |
|
fastplms/models/esmfold2/esmfold2_metrics.py
CHANGED
|
@@ -5,6 +5,7 @@ from __future__ import annotations
|
|
| 5 |
import numpy as np
|
| 6 |
import torch
|
| 7 |
import torch.nn.functional as F
|
|
|
|
| 8 |
from torch import Tensor
|
| 9 |
from torch.amp import autocast # type: ignore
|
| 10 |
|
|
@@ -18,8 +19,9 @@ from .esmfold2_protein_structure import (
|
|
| 18 |
|
| 19 |
|
| 20 |
def _distance_matrix(positions: Tensor, eps: float) -> Tensor:
|
| 21 |
-
|
| 22 |
-
|
|
|
|
| 23 |
|
| 24 |
|
| 25 |
def compute_lddt_from_dmat(
|
|
@@ -31,20 +33,21 @@ def compute_lddt_from_dmat(
|
|
| 31 |
per_residue: bool = True,
|
| 32 |
) -> Tensor:
|
| 33 |
"""Score distance matrices ``D_pred`` and ``D_true`` with shape (..., l, l)."""
|
|
|
|
| 34 |
|
| 35 |
sequence_length = dmat_true.size(-1)
|
| 36 |
-
identity = torch.eye(sequence_length, device=dmat_true.device)
|
| 37 |
-
scored_pairs = (dmat_true < cutoff) * pairwise_mask * (1.0 - identity)
|
| 38 |
-
absolute_error = torch.abs(dmat_true - dmat_pred)
|
| 39 |
score = (
|
| 40 |
(absolute_error < 0.5).type(absolute_error.dtype)
|
| 41 |
+ (absolute_error < 1.0).type(absolute_error.dtype)
|
| 42 |
+ (absolute_error < 2.0).type(absolute_error.dtype)
|
| 43 |
+ (absolute_error < 4.0).type(absolute_error.dtype)
|
| 44 |
-
) * 0.25
|
| 45 |
dimensions = (-1,) if per_residue else (-2, -1)
|
| 46 |
-
normalization = 1.0 / (eps + scored_pairs.sum(dim=dimensions))
|
| 47 |
-
return normalization * (eps + (scored_pairs * score).sum(dim=dimensions))
|
| 48 |
|
| 49 |
|
| 50 |
def compute_lddt(
|
|
@@ -58,16 +61,18 @@ def compute_lddt(
|
|
| 58 |
sequence_id: Tensor | None = None,
|
| 59 |
) -> Tensor:
|
| 60 |
"""Compute lDDT from coordinate tensors and atom masks."""
|
|
|
|
|
|
|
| 61 |
|
| 62 |
-
expanded_mask = all_atom_mask[..., None]
|
| 63 |
-
true_distances = _distance_matrix(all_atom_positions, eps)
|
| 64 |
-
predicted_distances = _distance_matrix(all_atom_pred_pos, eps)
|
| 65 |
-
pair_mask = expanded_mask * expanded_mask.transpose(-2, -1)
|
| 66 |
if pairwise_all_atom_mask is not None:
|
| 67 |
-
pair_mask = pair_mask * pairwise_all_atom_mask
|
| 68 |
if sequence_id is not None:
|
| 69 |
-
same_sequence = sequence_id[..., None] == sequence_id[..., None, :]
|
| 70 |
-
pair_mask = pair_mask * same_sequence.type_as(pair_mask)
|
| 71 |
return compute_lddt_from_dmat(
|
| 72 |
predicted_distances,
|
| 73 |
true_distances,
|
|
@@ -75,7 +80,7 @@ def compute_lddt(
|
|
| 75 |
cutoff=cutoff,
|
| 76 |
eps=eps,
|
| 77 |
per_residue=per_residue,
|
| 78 |
-
)
|
| 79 |
|
| 80 |
|
| 81 |
def compute_lddt_ca(
|
|
@@ -88,11 +93,13 @@ def compute_lddt_ca(
|
|
| 88 |
sequence_id: Tensor | None = None,
|
| 89 |
) -> Tensor:
|
| 90 |
"""Compute lDDT using only C-alpha coordinates."""
|
|
|
|
|
|
|
| 91 |
|
| 92 |
ca_index = residue_constants.atom_order["CA"]
|
| 93 |
predicted_ca = (
|
| 94 |
all_atom_pred_pos if all_atom_pred_pos.dim() == 3 else all_atom_pred_pos[..., ca_index, :]
|
| 95 |
-
)
|
| 96 |
return compute_lddt(
|
| 97 |
predicted_ca,
|
| 98 |
all_atom_positions[..., ca_index, :],
|
|
@@ -101,7 +108,7 @@ def compute_lddt_ca(
|
|
| 101 |
eps=eps,
|
| 102 |
per_residue=per_residue,
|
| 103 |
sequence_id=sequence_id,
|
| 104 |
-
)
|
| 105 |
|
| 106 |
|
| 107 |
@torch.no_grad()
|
|
@@ -114,22 +121,24 @@ def compute_rmsd(
|
|
| 114 |
reduction: str = "batch",
|
| 115 |
) -> Tensor:
|
| 116 |
"""Align ``X`` to ``Y`` and compute RMSD."""
|
|
|
|
|
|
|
| 117 |
|
| 118 |
centered_mobile, _, centered_target, _, rotation, counts = compute_alignment_tensors(
|
| 119 |
mobile,
|
| 120 |
target,
|
| 121 |
atom_exists_mask,
|
| 122 |
sequence_id,
|
| 123 |
-
)
|
| 124 |
rmsd = compute_rmsd_no_alignment(
|
| 125 |
torch.matmul(centered_mobile, rotation),
|
| 126 |
centered_target,
|
| 127 |
counts,
|
| 128 |
reduction=reduction,
|
| 129 |
-
)
|
| 130 |
if reduction == "per_residue" and sequence_id is not None:
|
| 131 |
-
return binpack(rmsd, sequence_id, pad_value=0)
|
| 132 |
-
return rmsd
|
| 133 |
|
| 134 |
|
| 135 |
def compute_gdt_ts(
|
|
@@ -140,36 +149,39 @@ def compute_gdt_ts(
|
|
| 140 |
reduction: str = "per_sample",
|
| 141 |
) -> Tensor:
|
| 142 |
"""Align ``X`` to ``Y`` and compute GDT-TS."""
|
|
|
|
|
|
|
| 143 |
|
| 144 |
if atom_exists_mask is None:
|
| 145 |
-
atom_exists_mask = torch.isfinite(target).all(dim=-1)
|
| 146 |
centered_mobile, _, centered_target, _, rotation, _ = compute_alignment_tensors(
|
| 147 |
mobile,
|
| 148 |
target,
|
| 149 |
atom_exists_mask,
|
| 150 |
sequence_id,
|
| 151 |
-
)
|
| 152 |
if sequence_id is not None:
|
| 153 |
-
atom_exists_mask = unbinpack(atom_exists_mask, sequence_id, pad_value=False)
|
| 154 |
return compute_gdt_ts_no_alignment(
|
| 155 |
torch.matmul(centered_mobile, rotation),
|
| 156 |
centered_target,
|
| 157 |
atom_exists_mask,
|
| 158 |
reduction,
|
| 159 |
-
)
|
| 160 |
|
| 161 |
|
| 162 |
def _batched_contacts(predictions: Tensor, targets: Tensor) -> tuple[Tensor, Tensor]:
|
|
|
|
| 163 |
if predictions.dim() == 2:
|
| 164 |
-
predictions = predictions.unsqueeze(0)
|
| 165 |
if targets.dim() == 2:
|
| 166 |
-
targets = targets.unsqueeze(0)
|
| 167 |
if predictions.size() != targets.size():
|
| 168 |
raise ValueError(
|
| 169 |
f"Size mismatch. Received predictions of size {predictions.size()}, "
|
| 170 |
f"targets of size {targets.size()}"
|
| 171 |
)
|
| 172 |
-
return predictions, targets
|
| 173 |
|
| 174 |
|
| 175 |
def _valid_contact_mask(
|
|
@@ -178,14 +190,15 @@ def _valid_contact_mask(
|
|
| 178 |
minsep: int,
|
| 179 |
maxsep: int | None,
|
| 180 |
) -> Tensor:
|
|
|
|
| 181 |
sequence_length = targets.shape[-1]
|
| 182 |
-
positions = torch.arange(sequence_length, device=targets.device)
|
| 183 |
-
separation = (positions.unsqueeze(0) - positions.unsqueeze(1)).unsqueeze(0)
|
| 184 |
-
valid = (separation >= minsep) & (targets >= 0)
|
| 185 |
if maxsep is not None:
|
| 186 |
-
valid &= separation < maxsep
|
| 187 |
-
within_length = positions.unsqueeze(0) < src_lengths.unsqueeze(1)
|
| 188 |
-
return valid & within_length.unsqueeze(1) & within_length.unsqueeze(2)
|
| 189 |
|
| 190 |
|
| 191 |
def contact_precision(
|
|
@@ -197,8 +210,10 @@ def contact_precision(
|
|
| 197 |
override_length: int | None = None,
|
| 198 |
) -> dict[str, Tensor]:
|
| 199 |
"""Compute P@L, P@L/5, and binned area for contact probabilities."""
|
|
|
|
|
|
|
| 200 |
|
| 201 |
-
predictions, targets = _batched_contacts(predictions, targets)
|
| 202 |
batch_size, sequence_length, _ = predictions.shape
|
| 203 |
if src_lengths is None:
|
| 204 |
src_lengths = torch.full(
|
|
@@ -206,30 +221,30 @@ def contact_precision(
|
|
| 206 |
sequence_length,
|
| 207 |
dtype=torch.long,
|
| 208 |
device=predictions.device,
|
| 209 |
-
)
|
| 210 |
-
valid = _valid_contact_mask(targets, src_lengths, minsep, maxsep)
|
| 211 |
-
masked_predictions = predictions.masked_fill(~valid, float("-inf"))
|
| 212 |
-
row_index, column_index = np.triu_indices(sequence_length, minsep)
|
| 213 |
-
upper_predictions = masked_predictions[:, row_index, column_index]
|
| 214 |
-
upper_targets = targets[:, row_index, column_index]
|
| 215 |
|
| 216 |
topk = sequence_length if override_length is None else max(sequence_length, override_length)
|
| 217 |
-
ranked_indices = upper_predictions.argsort(dim=-1, descending=True)[:, :topk]
|
| 218 |
-
batch_indices = torch.arange(batch_size, device=ranked_indices.device).unsqueeze(1)
|
| 219 |
-
ranked_targets = upper_targets[batch_indices, ranked_indices]
|
| 220 |
if ranked_targets.size(1) < topk:
|
| 221 |
-
ranked_targets = F.pad(ranked_targets, [0, topk - ranked_targets.size(1)])
|
| 222 |
-
cumulative_contacts = ranked_targets.type_as(predictions).cumsum(dim=-1)
|
| 223 |
|
| 224 |
-
gather_lengths = src_lengths.unsqueeze(1)
|
| 225 |
if override_length is not None:
|
| 226 |
-
gather_lengths = override_length * torch.ones_like(gather_lengths)
|
| 227 |
-
fractions = torch.arange(0.1, 1.1, 0.1, device=predictions.device).unsqueeze(0)
|
| 228 |
-
gather_indices = (fractions * gather_lengths).type(torch.long).sub(1).clamp_min(0)
|
| 229 |
-
cumulative_bins = cumulative_contacts.gather(1, gather_indices)
|
| 230 |
-
precisions = cumulative_bins / (gather_indices + 1).type_as(cumulative_bins)
|
| 231 |
return {
|
| 232 |
"AUC": precisions.mean(dim=-1),
|
| 233 |
"P@L": precisions[:, 9],
|
| 234 |
"P@L5": precisions[:, 1],
|
| 235 |
-
}
|
|
|
|
| 5 |
import numpy as np
|
| 6 |
import torch
|
| 7 |
import torch.nn.functional as F
|
| 8 |
+
|
| 9 |
from torch import Tensor
|
| 10 |
from torch.amp import autocast # type: ignore
|
| 11 |
|
|
|
|
| 19 |
|
| 20 |
|
| 21 |
def _distance_matrix(positions: Tensor, eps: float) -> Tensor:
|
| 22 |
+
# positions: (..., n, 3); n is the number of points in each distance matrix.
|
| 23 |
+
displacement = positions[..., None, :] - positions[..., None, :, :] # (..., n, n, 3)
|
| 24 |
+
return torch.sqrt(eps + torch.sum(displacement**2, dim=-1)) # (..., n, n)
|
| 25 |
|
| 26 |
|
| 27 |
def compute_lddt_from_dmat(
|
|
|
|
| 33 |
per_residue: bool = True,
|
| 34 |
) -> Tensor:
|
| 35 |
"""Score distance matrices ``D_pred`` and ``D_true`` with shape (..., l, l)."""
|
| 36 |
+
# dmat_pred, dmat_true, pairwise_mask: (..., l, l); cutoff broadcasts with distances.
|
| 37 |
|
| 38 |
sequence_length = dmat_true.size(-1)
|
| 39 |
+
identity = torch.eye(sequence_length, device=dmat_true.device) # (l, l)
|
| 40 |
+
scored_pairs = (dmat_true < cutoff) * pairwise_mask * (1.0 - identity) # (..., l, l)
|
| 41 |
+
absolute_error = torch.abs(dmat_true - dmat_pred) # (..., l, l)
|
| 42 |
score = (
|
| 43 |
(absolute_error < 0.5).type(absolute_error.dtype)
|
| 44 |
+ (absolute_error < 1.0).type(absolute_error.dtype)
|
| 45 |
+ (absolute_error < 2.0).type(absolute_error.dtype)
|
| 46 |
+ (absolute_error < 4.0).type(absolute_error.dtype)
|
| 47 |
+
) * 0.25 # (..., l, l)
|
| 48 |
dimensions = (-1,) if per_residue else (-2, -1)
|
| 49 |
+
normalization = 1.0 / (eps + scored_pairs.sum(dim=dimensions)) # (..., l) if per_residue, otherwise (...)
|
| 50 |
+
return normalization * (eps + (scored_pairs * score).sum(dim=dimensions)) # (..., l) if per_residue, otherwise (...)
|
| 51 |
|
| 52 |
|
| 53 |
def compute_lddt(
|
|
|
|
| 61 |
sequence_id: Tensor | None = None,
|
| 62 |
) -> Tensor:
|
| 63 |
"""Compute lDDT from coordinate tensors and atom masks."""
|
| 64 |
+
# positions: (..., n, 3); all_atom_mask: (..., n); pairwise mask: (..., n, n).
|
| 65 |
+
# sequence_id, when supplied, follows the point axes (..., n).
|
| 66 |
|
| 67 |
+
expanded_mask = all_atom_mask[..., None] # (..., n, 1)
|
| 68 |
+
true_distances = _distance_matrix(all_atom_positions, eps) # (..., n, n)
|
| 69 |
+
predicted_distances = _distance_matrix(all_atom_pred_pos, eps) # (..., n, n)
|
| 70 |
+
pair_mask = expanded_mask * expanded_mask.transpose(-2, -1) # (..., n, n)
|
| 71 |
if pairwise_all_atom_mask is not None:
|
| 72 |
+
pair_mask = pair_mask * pairwise_all_atom_mask # (..., n, n)
|
| 73 |
if sequence_id is not None:
|
| 74 |
+
same_sequence = sequence_id[..., None] == sequence_id[..., None, :] # (..., n, n)
|
| 75 |
+
pair_mask = pair_mask * same_sequence.type_as(pair_mask) # (..., n, n)
|
| 76 |
return compute_lddt_from_dmat(
|
| 77 |
predicted_distances,
|
| 78 |
true_distances,
|
|
|
|
| 80 |
cutoff=cutoff,
|
| 81 |
eps=eps,
|
| 82 |
per_residue=per_residue,
|
| 83 |
+
) # (..., n) if per_residue, otherwise (...)
|
| 84 |
|
| 85 |
|
| 86 |
def compute_lddt_ca(
|
|
|
|
| 93 |
sequence_id: Tensor | None = None,
|
| 94 |
) -> Tensor:
|
| 95 |
"""Compute lDDT using only C-alpha coordinates."""
|
| 96 |
+
# True coordinates/mask: (..., l, n_atoms, 3) / (..., l, n_atoms).
|
| 97 |
+
# Predicted rank-three input is treated as CA-only; otherwise its CA atom axis is selected.
|
| 98 |
|
| 99 |
ca_index = residue_constants.atom_order["CA"]
|
| 100 |
predicted_ca = (
|
| 101 |
all_atom_pred_pos if all_atom_pred_pos.dim() == 3 else all_atom_pred_pos[..., ca_index, :]
|
| 102 |
+
) # (..., l, 3)
|
| 103 |
return compute_lddt(
|
| 104 |
predicted_ca,
|
| 105 |
all_atom_positions[..., ca_index, :],
|
|
|
|
| 108 |
eps=eps,
|
| 109 |
per_residue=per_residue,
|
| 110 |
sequence_id=sequence_id,
|
| 111 |
+
) # (..., l) if per_residue, otherwise (...)
|
| 112 |
|
| 113 |
|
| 114 |
@torch.no_grad()
|
|
|
|
| 121 |
reduction: str = "batch",
|
| 122 |
) -> Tensor:
|
| 123 |
"""Align ``X`` to ``Y`` and compute RMSD."""
|
| 124 |
+
# mobile/target: (b, n, 3) or (b, l, n_atoms, 3); masks omit xyz.
|
| 125 |
+
# b_eff counts unpacked sequences when sequence_id is provided; n counts flattened atoms.
|
| 126 |
|
| 127 |
centered_mobile, _, centered_target, _, rotation, counts = compute_alignment_tensors(
|
| 128 |
mobile,
|
| 129 |
target,
|
| 130 |
atom_exists_mask,
|
| 131 |
sequence_id,
|
| 132 |
+
) # coordinates (b_eff, n, 3), centroids (b_eff, 1, 3), rotation (b_eff, 3, 3), counts (b_eff, 1)
|
| 133 |
rmsd = compute_rmsd_no_alignment(
|
| 134 |
torch.matmul(centered_mobile, rotation),
|
| 135 |
centered_target,
|
| 136 |
counts,
|
| 137 |
reduction=reduction,
|
| 138 |
+
) # (b_eff, n / 3) per_residue; (b_eff,) per_sample; () batch
|
| 139 |
if reduction == "per_residue" and sequence_id is not None:
|
| 140 |
+
return binpack(rmsd, sequence_id, pad_value=0) # (b, packed_length)
|
| 141 |
+
return rmsd # shape selected by reduction above
|
| 142 |
|
| 143 |
|
| 144 |
def compute_gdt_ts(
|
|
|
|
| 149 |
reduction: str = "per_sample",
|
| 150 |
) -> Tensor:
|
| 151 |
"""Align ``X`` to ``Y`` and compute GDT-TS."""
|
| 152 |
+
# mobile/target: batched xyz coordinates; masks omit xyz.
|
| 153 |
+
# b_eff counts unpacked sequences when sequence_id is provided; n counts flattened atoms.
|
| 154 |
|
| 155 |
if atom_exists_mask is None:
|
| 156 |
+
atom_exists_mask = torch.isfinite(target).all(dim=-1) # target.shape[:-1]
|
| 157 |
centered_mobile, _, centered_target, _, rotation, _ = compute_alignment_tensors(
|
| 158 |
mobile,
|
| 159 |
target,
|
| 160 |
atom_exists_mask,
|
| 161 |
sequence_id,
|
| 162 |
+
) # coordinates (b_eff, n, 3), centroids (b_eff, 1, 3), rotation (b_eff, 3, 3), counts (b_eff, 1)
|
| 163 |
if sequence_id is not None:
|
| 164 |
+
atom_exists_mask = unbinpack(atom_exists_mask, sequence_id, pad_value=False) # (b_eff, max_sequence_length, *original_mask.shape[2:])
|
| 165 |
return compute_gdt_ts_no_alignment(
|
| 166 |
torch.matmul(centered_mobile, rotation),
|
| 167 |
centered_target,
|
| 168 |
atom_exists_mask,
|
| 169 |
reduction,
|
| 170 |
+
) # (b_eff,) per_sample; () batch
|
| 171 |
|
| 172 |
|
| 173 |
def _batched_contacts(predictions: Tensor, targets: Tensor) -> tuple[Tensor, Tensor]:
|
| 174 |
+
# predictions, targets: (l, l) or (b, l, l).
|
| 175 |
if predictions.dim() == 2:
|
| 176 |
+
predictions = predictions.unsqueeze(0) # (1, l, l)
|
| 177 |
if targets.dim() == 2:
|
| 178 |
+
targets = targets.unsqueeze(0) # (1, l, l)
|
| 179 |
if predictions.size() != targets.size():
|
| 180 |
raise ValueError(
|
| 181 |
f"Size mismatch. Received predictions of size {predictions.size()}, "
|
| 182 |
f"targets of size {targets.size()}"
|
| 183 |
)
|
| 184 |
+
return predictions, targets # each (b, l, l)
|
| 185 |
|
| 186 |
|
| 187 |
def _valid_contact_mask(
|
|
|
|
| 190 |
minsep: int,
|
| 191 |
maxsep: int | None,
|
| 192 |
) -> Tensor:
|
| 193 |
+
# targets: (b, l, l); src_lengths: (b,).
|
| 194 |
sequence_length = targets.shape[-1]
|
| 195 |
+
positions = torch.arange(sequence_length, device=targets.device) # (l,)
|
| 196 |
+
separation = (positions.unsqueeze(0) - positions.unsqueeze(1)).unsqueeze(0) # (1, l, l)
|
| 197 |
+
valid = (separation >= minsep) & (targets >= 0) # (b, l, l)
|
| 198 |
if maxsep is not None:
|
| 199 |
+
valid &= separation < maxsep # (b, l, l)
|
| 200 |
+
within_length = positions.unsqueeze(0) < src_lengths.unsqueeze(1) # (b, l)
|
| 201 |
+
return valid & within_length.unsqueeze(1) & within_length.unsqueeze(2) # (b, l, l)
|
| 202 |
|
| 203 |
|
| 204 |
def contact_precision(
|
|
|
|
| 210 |
override_length: int | None = None,
|
| 211 |
) -> dict[str, Tensor]:
|
| 212 |
"""Compute P@L, P@L/5, and binned area for contact probabilities."""
|
| 213 |
+
# predictions, targets: (l, l) or (b, l, l); src_lengths: (b,) or None.
|
| 214 |
+
# n_upper_pairs counts the entries returned by triu_indices(l, minsep).
|
| 215 |
|
| 216 |
+
predictions, targets = _batched_contacts(predictions, targets) # each (b, l, l)
|
| 217 |
batch_size, sequence_length, _ = predictions.shape
|
| 218 |
if src_lengths is None:
|
| 219 |
src_lengths = torch.full(
|
|
|
|
| 221 |
sequence_length,
|
| 222 |
dtype=torch.long,
|
| 223 |
device=predictions.device,
|
| 224 |
+
) # (b,)
|
| 225 |
+
valid = _valid_contact_mask(targets, src_lengths, minsep, maxsep) # (b, l, l)
|
| 226 |
+
masked_predictions = predictions.masked_fill(~valid, float("-inf")) # (b, l, l)
|
| 227 |
+
row_index, column_index = np.triu_indices(sequence_length, minsep) # each (n_upper_pairs,)
|
| 228 |
+
upper_predictions = masked_predictions[:, row_index, column_index] # (b, n_upper_pairs)
|
| 229 |
+
upper_targets = targets[:, row_index, column_index] # (b, n_upper_pairs)
|
| 230 |
|
| 231 |
topk = sequence_length if override_length is None else max(sequence_length, override_length)
|
| 232 |
+
ranked_indices = upper_predictions.argsort(dim=-1, descending=True)[:, :topk] # (b, min(topk, n_upper_pairs))
|
| 233 |
+
batch_indices = torch.arange(batch_size, device=ranked_indices.device).unsqueeze(1) # (b, 1)
|
| 234 |
+
ranked_targets = upper_targets[batch_indices, ranked_indices] # (b, min(topk, n_upper_pairs))
|
| 235 |
if ranked_targets.size(1) < topk:
|
| 236 |
+
ranked_targets = F.pad(ranked_targets, [0, topk - ranked_targets.size(1)]) # (b, topk)
|
| 237 |
+
cumulative_contacts = ranked_targets.type_as(predictions).cumsum(dim=-1) # (b, topk)
|
| 238 |
|
| 239 |
+
gather_lengths = src_lengths.unsqueeze(1) # (b, 1)
|
| 240 |
if override_length is not None:
|
| 241 |
+
gather_lengths = override_length * torch.ones_like(gather_lengths) # (b, 1)
|
| 242 |
+
fractions = torch.arange(0.1, 1.1, 0.1, device=predictions.device).unsqueeze(0) # (1, 10)
|
| 243 |
+
gather_indices = (fractions * gather_lengths).type(torch.long).sub(1).clamp_min(0) # (b, 10)
|
| 244 |
+
cumulative_bins = cumulative_contacts.gather(1, gather_indices) # (b, 10)
|
| 245 |
+
precisions = cumulative_bins / (gather_indices + 1).type_as(cumulative_bins) # (b, 10)
|
| 246 |
return {
|
| 247 |
"AUC": precisions.mean(dim=-1),
|
| 248 |
"P@L": precisions[:, 9],
|
| 249 |
"P@L5": precisions[:, 1],
|
| 250 |
+
} # each metric: (b,)
|
fastplms/models/esmfold2/esmfold2_misc.py
CHANGED
|
@@ -6,6 +6,10 @@ module therefore performs no device selection, compilation, or remote access.
|
|
| 6 |
|
| 7 |
from __future__ import annotations
|
| 8 |
|
|
|
|
|
|
|
|
|
|
|
|
|
| 9 |
from collections import defaultdict
|
| 10 |
from collections.abc import Generator, Iterable, Sequence
|
| 11 |
from contextlib import AbstractContextManager, nullcontext
|
|
@@ -14,13 +18,10 @@ from io import BytesIO
|
|
| 14 |
from typing import Any, Protocol, TypeVar, runtime_checkable
|
| 15 |
from warnings import warn
|
| 16 |
|
| 17 |
-
import numpy as np
|
| 18 |
-
import torch
|
| 19 |
-
import zstandard
|
| 20 |
-
|
| 21 |
from .esmfold2_constants_esm3 import CHAIN_BREAK_STR
|
| 22 |
from .esmfold2_utils_types import FunctionAnnotation
|
| 23 |
|
|
|
|
| 24 |
MAX_SUPPORTED_DISTANCE = 1e6
|
| 25 |
|
| 26 |
TSequence = TypeVar("TSequence", bound=Sequence)
|
|
@@ -50,44 +51,47 @@ def fp32_autocast_context(
|
|
| 50 |
|
| 51 |
def maybe_tensor(value, convert_none_to_nan: bool = False) -> torch.Tensor | None:
|
| 52 |
"""Convert an optional array-like value to a tensor."""
|
|
|
|
| 53 |
|
| 54 |
if value is None:
|
| 55 |
return None
|
| 56 |
if isinstance(value, torch.Tensor):
|
| 57 |
-
return value
|
| 58 |
if isinstance(value, list) and all(isinstance(element, torch.Tensor) for element in value):
|
| 59 |
-
return torch.stack(value)
|
| 60 |
if convert_none_to_nan:
|
| 61 |
-
value = np.asarray(value, dtype=np.float32)
|
| 62 |
-
value = np.where(value is None, np.nan, value)
|
| 63 |
-
return torch.tensor(value)
|
| 64 |
|
| 65 |
|
| 66 |
def maybe_list(value, convert_nan_to_none: bool = False) -> list | None:
|
| 67 |
"""Convert an optional tensor or NumPy array to nested Python lists."""
|
|
|
|
| 68 |
|
| 69 |
if value is None:
|
| 70 |
return None
|
| 71 |
if not convert_nan_to_none:
|
| 72 |
return value.tolist()
|
| 73 |
if isinstance(value, torch.Tensor):
|
| 74 |
-
nan_mask = torch.isnan(value).cpu().numpy()
|
| 75 |
-
array = value.cpu().numpy().astype(object)
|
| 76 |
elif isinstance(value, np.ndarray):
|
| 77 |
-
nan_mask = np.isnan(value)
|
| 78 |
-
array = value.astype(object)
|
| 79 |
else:
|
| 80 |
raise TypeError("maybe_list can only work with torch.tensor or np.ndarray.")
|
| 81 |
-
array[nan_mask] = None
|
| 82 |
return array.tolist()
|
| 83 |
|
| 84 |
|
| 85 |
def replace_inf(data):
|
| 86 |
"""Replace infinite array values by the ESM API sentinel value."""
|
|
|
|
| 87 |
|
| 88 |
if data is None:
|
| 89 |
return None
|
| 90 |
-
array = np.asarray(data, dtype=np.float32)
|
| 91 |
return np.where(np.isinf(array), 1000, array).tolist()
|
| 92 |
|
| 93 |
|
|
@@ -118,9 +122,10 @@ def slice_any_object(
|
|
| 118 |
idx: int | list[int] | slice | np.ndarray,
|
| 119 |
) -> TSequence:
|
| 120 |
"""Slice tensors, arrays, dataclasses, and ordinary Python sequences."""
|
|
|
|
| 121 |
|
| 122 |
if isinstance(obj, (np.ndarray, torch.Tensor)) or is_dataclass(obj):
|
| 123 |
-
return obj[idx] # type: ignore[index,return-value]
|
| 124 |
return slice_python_object_as_numpy(obj, idx)
|
| 125 |
|
| 126 |
|
|
@@ -172,20 +177,21 @@ def concat_objects(objs: Sequence[Any], separator: Any | None = None):
|
|
| 172 |
objs
|
| 173 |
if separator is None
|
| 174 |
else list(iterate_with_intermediate(objs, np.array([separator])))
|
| 175 |
-
)
|
| 176 |
-
return np.concatenate(pieces)
|
| 177 |
if isinstance(first, torch.Tensor):
|
| 178 |
pieces = (
|
| 179 |
objs
|
| 180 |
if separator is None
|
| 181 |
else list(iterate_with_intermediate(objs, torch.tensor([separator])))
|
| 182 |
-
)
|
| 183 |
-
return torch.cat(pieces) # type: ignore[arg-type]
|
| 184 |
raise TypeError(type(first))
|
| 185 |
|
| 186 |
|
| 187 |
-
def rbf(values, v_min, v_max, n_bins=16):
|
| 188 |
"""Encode values against evenly spaced radial basis centers."""
|
|
|
|
| 189 |
|
| 190 |
centers = torch.linspace(
|
| 191 |
v_min,
|
|
@@ -193,34 +199,39 @@ def rbf(values, v_min, v_max, n_bins=16):
|
|
| 193 |
n_bins,
|
| 194 |
dtype=values.dtype,
|
| 195 |
device=values.device,
|
| 196 |
-
)
|
| 197 |
-
centers = centers.reshape((1,) * values.ndim + (-1,))
|
| 198 |
-
standardized = (values.unsqueeze(-1) - centers) / ((v_max - v_min) / n_bins)
|
| 199 |
-
return torch.exp(-(standardized**2))
|
| 200 |
|
| 201 |
|
| 202 |
-
def batched_gather(
|
|
|
|
|
|
|
| 203 |
"""Gather along one data dimension while retaining leading batch axes."""
|
|
|
|
|
|
|
| 204 |
|
| 205 |
batch_indices = []
|
| 206 |
index_rank = len(inds.shape)
|
| 207 |
for axis, size in enumerate(data.shape[:no_batch_dims]):
|
| 208 |
shape = (1,) * axis + (-1,) + (1,) * (index_rank - axis - 1)
|
| 209 |
-
batch_indices.append(torch.arange(size).view(*shape))
|
| 210 |
tail = [slice(None)] * (len(data.shape) - no_batch_dims)
|
| 211 |
tail[dim - no_batch_dims if dim >= 0 else dim] = inds
|
| 212 |
-
return data[tuple(batch_indices + tail)]
|
| 213 |
|
| 214 |
|
| 215 |
def node_gather(s: torch.Tensor, edges: torch.Tensor) -> torch.Tensor:
|
| 216 |
"""Gather node features for each row of an edge-index tensor."""
|
|
|
|
| 217 |
|
| 218 |
return batched_gather(
|
| 219 |
s.unsqueeze(-3),
|
| 220 |
edges,
|
| 221 |
-2,
|
| 222 |
no_batch_dims=len(s.shape) - 1,
|
| 223 |
-
)
|
| 224 |
|
| 225 |
|
| 226 |
def knn_graph(
|
|
@@ -230,31 +241,32 @@ def knn_graph(
|
|
| 230 |
sequence_id: torch.Tensor,
|
| 231 |
*,
|
| 232 |
no_knn: int,
|
| 233 |
-
):
|
| 234 |
"""Build nearest-neighbor edges, using sequence distance for missing geometry."""
|
|
|
|
| 235 |
|
| 236 |
length = coords.shape[-2]
|
| 237 |
-
coords = coords.nan_to_num()
|
| 238 |
-
missing_pair = ~(coord_mask[..., None, :] & coord_mask[..., :, None])
|
| 239 |
-
excluded_pair = padding_mask[..., None, :] | padding_mask[..., :, None]
|
| 240 |
if sequence_id is not None:
|
| 241 |
-
excluded_pair |= sequence_id.unsqueeze(1) != sequence_id.unsqueeze(2)
|
| 242 |
|
| 243 |
-
distances = (coords.unsqueeze(-2) - coords.unsqueeze(-3)).norm(dim=-1)
|
| 244 |
-
residue_index = torch.arange(length, device=coords.device)
|
| 245 |
-
sequence_distance = (residue_index.unsqueeze(-1) - residue_index.unsqueeze(-2)).abs()
|
| 246 |
if not (distances[~missing_pair] < MAX_SUPPORTED_DISTANCE).all():
|
| 247 |
raise ValueError(
|
| 248 |
"Coordinate pairwise distances exceed max supported distance "
|
| 249 |
f"({MAX_SUPPORTED_DISTANCE}). "
|
| 250 |
)
|
| 251 |
|
| 252 |
-
rank_distance = sequence_distance.to(distances.dtype).mul(1e2).add(MAX_SUPPORTED_DISTANCE)
|
| 253 |
-
rank_distance = rank_distance.where(missing_pair, distances)
|
| 254 |
-
rank_distance = rank_distance.masked_fill(excluded_pair, torch.inf)
|
| 255 |
-
sorted_distance, sorted_edge = rank_distance.sort(dim=-1, descending=False)
|
| 256 |
width = min(no_knn, length)
|
| 257 |
-
return sorted_edge[..., :width], sorted_distance[..., :width].isfinite()
|
| 258 |
|
| 259 |
|
| 260 |
def stack_variable_length_tensors(
|
|
@@ -263,6 +275,7 @@ def stack_variable_length_tensors(
|
|
| 263 |
dtype: torch.dtype | None = None,
|
| 264 |
) -> torch.Tensor:
|
| 265 |
"""Pad arbitrary tensor dimensions to their maxima, then stack."""
|
|
|
|
| 266 |
|
| 267 |
output_shape = [
|
| 268 |
len(sequences),
|
|
@@ -273,56 +286,58 @@ def stack_variable_length_tensors(
|
|
| 273 |
constant_value,
|
| 274 |
dtype=sequences[0].dtype if dtype is None else dtype,
|
| 275 |
device=sequences[0].device,
|
| 276 |
-
)
|
| 277 |
for destination, source in zip(output, sequences, strict=True):
|
| 278 |
-
destination[tuple(slice(size) for size in source.shape)] = source
|
| 279 |
-
return output
|
| 280 |
|
| 281 |
|
| 282 |
def binpack(
|
| 283 |
tensor: torch.Tensor,
|
| 284 |
sequence_id: torch.Tensor | None,
|
| 285 |
pad_value: int | float,
|
| 286 |
-
):
|
| 287 |
"""Scatter a sequence-major tensor into the packed layout described by IDs."""
|
|
|
|
| 288 |
|
| 289 |
if sequence_id is None:
|
| 290 |
-
return tensor
|
| 291 |
-
sequence_counts = sequence_id.max(dim=-1).values + 1
|
| 292 |
output = torch.full(
|
| 293 |
sequence_id.shape + tensor.shape[2:],
|
| 294 |
fill_value=pad_value,
|
| 295 |
dtype=tensor.dtype,
|
| 296 |
device=tensor.device,
|
| 297 |
-
)
|
| 298 |
source_index = 0
|
| 299 |
for batch_index, (batch_ids, count) in enumerate(
|
| 300 |
zip(sequence_id, sequence_counts, strict=True)
|
| 301 |
):
|
| 302 |
for seqid in range(count):
|
| 303 |
-
selection = batch_ids == seqid
|
| 304 |
-
output[batch_index, selection] = tensor[source_index, : selection.sum()]
|
| 305 |
source_index += 1
|
| 306 |
-
return output
|
| 307 |
|
| 308 |
|
| 309 |
def unbinpack(
|
| 310 |
tensor: torch.Tensor,
|
| 311 |
sequence_id: torch.Tensor | None,
|
| 312 |
pad_value: int | float,
|
| 313 |
-
):
|
| 314 |
"""Restore sequence-major rows from a packed tensor and its sequence IDs."""
|
|
|
|
| 315 |
|
| 316 |
if sequence_id is None:
|
| 317 |
-
return tensor
|
| 318 |
rows = []
|
| 319 |
-
sequence_counts = sequence_id.max(dim=-1).values + 1
|
| 320 |
for batch_index, (batch_ids, count) in enumerate(
|
| 321 |
zip(sequence_id, sequence_counts, strict=True)
|
| 322 |
):
|
| 323 |
for seqid in range(count):
|
| 324 |
-
rows.append(tensor[batch_index, batch_ids == seqid])
|
| 325 |
-
return stack_variable_length_tensors(rows, pad_value)
|
| 326 |
|
| 327 |
|
| 328 |
def merge_ranges(
|
|
@@ -386,7 +401,7 @@ def get_chainbreak_boundaries_from_sequence(
|
|
| 386 |
boundaries.extend((index, index + 1))
|
| 387 |
boundaries.append(len(sequence))
|
| 388 |
assert len(boundaries) % 2 == 0
|
| 389 |
-
return np.asarray(boundaries).reshape(-1, 2)
|
| 390 |
|
| 391 |
|
| 392 |
def deserialize_tensors(data: bytes) -> Any:
|
|
|
|
| 6 |
|
| 7 |
from __future__ import annotations
|
| 8 |
|
| 9 |
+
import numpy as np
|
| 10 |
+
import torch
|
| 11 |
+
import zstandard
|
| 12 |
+
|
| 13 |
from collections import defaultdict
|
| 14 |
from collections.abc import Generator, Iterable, Sequence
|
| 15 |
from contextlib import AbstractContextManager, nullcontext
|
|
|
|
| 18 |
from typing import Any, Protocol, TypeVar, runtime_checkable
|
| 19 |
from warnings import warn
|
| 20 |
|
|
|
|
|
|
|
|
|
|
|
|
|
| 21 |
from .esmfold2_constants_esm3 import CHAIN_BREAK_STR
|
| 22 |
from .esmfold2_utils_types import FunctionAnnotation
|
| 23 |
|
| 24 |
+
|
| 25 |
MAX_SUPPORTED_DISTANCE = 1e6
|
| 26 |
|
| 27 |
TSequence = TypeVar("TSequence", bound=Sequence)
|
|
|
|
| 51 |
|
| 52 |
def maybe_tensor(value, convert_none_to_nan: bool = False) -> torch.Tensor | None:
|
| 53 |
"""Convert an optional array-like value to a tensor."""
|
| 54 |
+
# value: array-like arbitrary shape, or a list of identically shaped tensors.
|
| 55 |
|
| 56 |
if value is None:
|
| 57 |
return None
|
| 58 |
if isinstance(value, torch.Tensor):
|
| 59 |
+
return value # value.shape
|
| 60 |
if isinstance(value, list) and all(isinstance(element, torch.Tensor) for element in value):
|
| 61 |
+
return torch.stack(value) # (n_values, *element_shape)
|
| 62 |
if convert_none_to_nan:
|
| 63 |
+
value = np.asarray(value, dtype=np.float32) # shape inferred from the nested input
|
| 64 |
+
value = np.where(value is None, np.nan, value) # value.shape
|
| 65 |
+
return torch.tensor(value) # shape inferred from the array-like input
|
| 66 |
|
| 67 |
|
| 68 |
def maybe_list(value, convert_nan_to_none: bool = False) -> list | None:
|
| 69 |
"""Convert an optional tensor or NumPy array to nested Python lists."""
|
| 70 |
+
# value: arbitrary-shaped array/tensor; element order and nesting are retained.
|
| 71 |
|
| 72 |
if value is None:
|
| 73 |
return None
|
| 74 |
if not convert_nan_to_none:
|
| 75 |
return value.tolist()
|
| 76 |
if isinstance(value, torch.Tensor):
|
| 77 |
+
nan_mask = torch.isnan(value).cpu().numpy() # value.shape
|
| 78 |
+
array = value.cpu().numpy().astype(object) # value.shape
|
| 79 |
elif isinstance(value, np.ndarray):
|
| 80 |
+
nan_mask = np.isnan(value) # value.shape
|
| 81 |
+
array = value.astype(object) # value.shape
|
| 82 |
else:
|
| 83 |
raise TypeError("maybe_list can only work with torch.tensor or np.ndarray.")
|
| 84 |
+
array[nan_mask] = None # (n_nan,) selected elements
|
| 85 |
return array.tolist()
|
| 86 |
|
| 87 |
|
| 88 |
def replace_inf(data):
|
| 89 |
"""Replace infinite array values by the ESM API sentinel value."""
|
| 90 |
+
# data: array-like arbitrary shape; the returned list retains its nesting.
|
| 91 |
|
| 92 |
if data is None:
|
| 93 |
return None
|
| 94 |
+
array = np.asarray(data, dtype=np.float32) # shape inferred from data
|
| 95 |
return np.where(np.isinf(array), 1000, array).tolist()
|
| 96 |
|
| 97 |
|
|
|
|
| 122 |
idx: int | list[int] | slice | np.ndarray,
|
| 123 |
) -> TSequence:
|
| 124 |
"""Slice tensors, arrays, dataclasses, and ordinary Python sequences."""
|
| 125 |
+
# Array shape is caller-defined; idx follows that array type's indexing rules.
|
| 126 |
|
| 127 |
if isinstance(obj, (np.ndarray, torch.Tensor)) or is_dataclass(obj):
|
| 128 |
+
return obj[idx] # type: ignore[index,return-value]; shape determined by NumPy/Torch indexing when obj is an array
|
| 129 |
return slice_python_object_as_numpy(obj, idx)
|
| 130 |
|
| 131 |
|
|
|
|
| 177 |
objs
|
| 178 |
if separator is None
|
| 179 |
else list(iterate_with_intermediate(objs, np.array([separator])))
|
| 180 |
+
) # arrays with a common trailing shape, interleaved with a separator if supplied
|
| 181 |
+
return np.concatenate(pieces) # (sum of leading lengths, *common_trailing_shape)
|
| 182 |
if isinstance(first, torch.Tensor):
|
| 183 |
pieces = (
|
| 184 |
objs
|
| 185 |
if separator is None
|
| 186 |
else list(iterate_with_intermediate(objs, torch.tensor([separator])))
|
| 187 |
+
) # tensors with a common trailing shape, interleaved with a separator if supplied
|
| 188 |
+
return torch.cat(pieces) # type: ignore[arg-type]; (sum of leading lengths, *common_trailing_shape)
|
| 189 |
raise TypeError(type(first))
|
| 190 |
|
| 191 |
|
| 192 |
+
def rbf(values: torch.Tensor, v_min: float, v_max: float, n_bins: int = 16) -> torch.Tensor:
|
| 193 |
"""Encode values against evenly spaced radial basis centers."""
|
| 194 |
+
# values: arbitrary shape; the output appends a final n_bins axis.
|
| 195 |
|
| 196 |
centers = torch.linspace(
|
| 197 |
v_min,
|
|
|
|
| 199 |
n_bins,
|
| 200 |
dtype=values.dtype,
|
| 201 |
device=values.device,
|
| 202 |
+
) # (n_bins,)
|
| 203 |
+
centers = centers.reshape((1,) * values.ndim + (-1,)) # (1, ..., 1, n_bins); values.ndim leading singleton axes
|
| 204 |
+
standardized = (values.unsqueeze(-1) - centers) / ((v_max - v_min) / n_bins) # (*values.shape, n_bins)
|
| 205 |
+
return torch.exp(-(standardized**2)) # (*values.shape, n_bins)
|
| 206 |
|
| 207 |
|
| 208 |
+
def batched_gather(
|
| 209 |
+
data: torch.Tensor, inds: torch.Tensor, dim: int = 0, no_batch_dims: int = 0
|
| 210 |
+
) -> torch.Tensor:
|
| 211 |
"""Gather along one data dimension while retaining leading batch axes."""
|
| 212 |
+
# data/inds ranks are caller-defined. The first no_batch_dims axes index together;
|
| 213 |
+
# remaining advanced-index axes broadcast, while slice axes retain their data lengths.
|
| 214 |
|
| 215 |
batch_indices = []
|
| 216 |
index_rank = len(inds.shape)
|
| 217 |
for axis, size in enumerate(data.shape[:no_batch_dims]):
|
| 218 |
shape = (1,) * axis + (-1,) + (1,) * (index_rank - axis - 1)
|
| 219 |
+
batch_indices.append(torch.arange(size).view(*shape)) # index-rank tensor; only the batch axis has length size
|
| 220 |
tail = [slice(None)] * (len(data.shape) - no_batch_dims)
|
| 221 |
tail[dim - no_batch_dims if dim >= 0 else dim] = inds
|
| 222 |
+
return data[tuple(batch_indices + tail)] # broadcast advanced-index shape plus retained data slice axes
|
| 223 |
|
| 224 |
|
| 225 |
def node_gather(s: torch.Tensor, edges: torch.Tensor) -> torch.Tensor:
|
| 226 |
"""Gather node features for each row of an edge-index tensor."""
|
| 227 |
+
# s: (..., n_nodes, d); edges: (..., n_nodes, n_neighbors).
|
| 228 |
|
| 229 |
return batched_gather(
|
| 230 |
s.unsqueeze(-3),
|
| 231 |
edges,
|
| 232 |
-2,
|
| 233 |
no_batch_dims=len(s.shape) - 1,
|
| 234 |
+
) # (..., n_nodes, n_neighbors, d)
|
| 235 |
|
| 236 |
|
| 237 |
def knn_graph(
|
|
|
|
| 241 |
sequence_id: torch.Tensor,
|
| 242 |
*,
|
| 243 |
no_knn: int,
|
| 244 |
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
| 245 |
"""Build nearest-neighbor edges, using sequence distance for missing geometry."""
|
| 246 |
+
# coords: (..., l, 3); masks: (..., l). With sequence_id, inputs use batch shape (b, l).
|
| 247 |
|
| 248 |
length = coords.shape[-2]
|
| 249 |
+
coords = coords.nan_to_num() # (..., l, 3)
|
| 250 |
+
missing_pair = ~(coord_mask[..., None, :] & coord_mask[..., :, None]) # (..., l, l)
|
| 251 |
+
excluded_pair = padding_mask[..., None, :] | padding_mask[..., :, None] # (..., l, l)
|
| 252 |
if sequence_id is not None:
|
| 253 |
+
excluded_pair |= sequence_id.unsqueeze(1) != sequence_id.unsqueeze(2) # (b, l, l)
|
| 254 |
|
| 255 |
+
distances = (coords.unsqueeze(-2) - coords.unsqueeze(-3)).norm(dim=-1) # (..., l, l)
|
| 256 |
+
residue_index = torch.arange(length, device=coords.device) # (l,)
|
| 257 |
+
sequence_distance = (residue_index.unsqueeze(-1) - residue_index.unsqueeze(-2)).abs() # (l, l)
|
| 258 |
if not (distances[~missing_pair] < MAX_SUPPORTED_DISTANCE).all():
|
| 259 |
raise ValueError(
|
| 260 |
"Coordinate pairwise distances exceed max supported distance "
|
| 261 |
f"({MAX_SUPPORTED_DISTANCE}). "
|
| 262 |
)
|
| 263 |
|
| 264 |
+
rank_distance = sequence_distance.to(distances.dtype).mul(1e2).add(MAX_SUPPORTED_DISTANCE) # (l, l)
|
| 265 |
+
rank_distance = rank_distance.where(missing_pair, distances) # (..., l, l)
|
| 266 |
+
rank_distance = rank_distance.masked_fill(excluded_pair, torch.inf) # (..., l, l)
|
| 267 |
+
sorted_distance, sorted_edge = rank_distance.sort(dim=-1, descending=False) # each (..., l, l)
|
| 268 |
width = min(no_knn, length)
|
| 269 |
+
return sorted_edge[..., :width], sorted_distance[..., :width].isfinite() # each (..., l, min(no_knn, l))
|
| 270 |
|
| 271 |
|
| 272 |
def stack_variable_length_tensors(
|
|
|
|
| 275 |
dtype: torch.dtype | None = None,
|
| 276 |
) -> torch.Tensor:
|
| 277 |
"""Pad arbitrary tensor dimensions to their maxima, then stack."""
|
| 278 |
+
# Each sequence has the same rank; its axis lengths may differ. Padding uses each axis maximum.
|
| 279 |
|
| 280 |
output_shape = [
|
| 281 |
len(sequences),
|
|
|
|
| 286 |
constant_value,
|
| 287 |
dtype=sequences[0].dtype if dtype is None else dtype,
|
| 288 |
device=sequences[0].device,
|
| 289 |
+
) # (n_sequences, *axiswise_maximum_shape)
|
| 290 |
for destination, source in zip(output, sequences, strict=True):
|
| 291 |
+
destination[tuple(slice(size) for size in source.shape)] = source # source.shape slice of destination
|
| 292 |
+
return output # (n_sequences, *axiswise_maximum_shape)
|
| 293 |
|
| 294 |
|
| 295 |
def binpack(
|
| 296 |
tensor: torch.Tensor,
|
| 297 |
sequence_id: torch.Tensor | None,
|
| 298 |
pad_value: int | float,
|
| 299 |
+
) -> torch.Tensor:
|
| 300 |
"""Scatter a sequence-major tensor into the packed layout described by IDs."""
|
| 301 |
+
# tensor: (n_unpacked_sequences, max_length, ...); sequence_id: (b, packed_length) or None.
|
| 302 |
|
| 303 |
if sequence_id is None:
|
| 304 |
+
return tensor # tensor.shape
|
| 305 |
+
sequence_counts = sequence_id.max(dim=-1).values + 1 # (b,)
|
| 306 |
output = torch.full(
|
| 307 |
sequence_id.shape + tensor.shape[2:],
|
| 308 |
fill_value=pad_value,
|
| 309 |
dtype=tensor.dtype,
|
| 310 |
device=tensor.device,
|
| 311 |
+
) # (b, packed_length, *tensor.shape[2:])
|
| 312 |
source_index = 0
|
| 313 |
for batch_index, (batch_ids, count) in enumerate(
|
| 314 |
zip(sequence_id, sequence_counts, strict=True)
|
| 315 |
):
|
| 316 |
for seqid in range(count):
|
| 317 |
+
selection = batch_ids == seqid # (packed_length,)
|
| 318 |
+
output[batch_index, selection] = tensor[source_index, : selection.sum()] # (n_selected, *tensor.shape[2:]) selected rows
|
| 319 |
source_index += 1
|
| 320 |
+
return output # (b, packed_length, *tensor.shape[2:])
|
| 321 |
|
| 322 |
|
| 323 |
def unbinpack(
|
| 324 |
tensor: torch.Tensor,
|
| 325 |
sequence_id: torch.Tensor | None,
|
| 326 |
pad_value: int | float,
|
| 327 |
+
) -> torch.Tensor:
|
| 328 |
"""Restore sequence-major rows from a packed tensor and its sequence IDs."""
|
| 329 |
+
# tensor: (b, packed_length, ...); sequence_id: (b, packed_length) or None.
|
| 330 |
|
| 331 |
if sequence_id is None:
|
| 332 |
+
return tensor # tensor.shape
|
| 333 |
rows = []
|
| 334 |
+
sequence_counts = sequence_id.max(dim=-1).values + 1 # (b,)
|
| 335 |
for batch_index, (batch_ids, count) in enumerate(
|
| 336 |
zip(sequence_id, sequence_counts, strict=True)
|
| 337 |
):
|
| 338 |
for seqid in range(count):
|
| 339 |
+
rows.append(tensor[batch_index, batch_ids == seqid]) # (selected_sequence_length, *tensor.shape[2:])
|
| 340 |
+
return stack_variable_length_tensors(rows, pad_value) # (n_unpacked_sequences, max_sequence_length, *tensor.shape[2:])
|
| 341 |
|
| 342 |
|
| 343 |
def merge_ranges(
|
|
|
|
| 401 |
boundaries.extend((index, index + 1))
|
| 402 |
boundaries.append(len(sequence))
|
| 403 |
assert len(boundaries) % 2 == 0
|
| 404 |
+
return np.asarray(boundaries).reshape(-1, 2) # (n_chains, 2)
|
| 405 |
|
| 406 |
|
| 407 |
def deserialize_tensors(data: bytes) -> Any:
|
fastplms/models/esmfold2/esmfold2_mmcif_parsing.py
CHANGED
|
@@ -5,17 +5,18 @@ from __future__ import annotations
|
|
| 5 |
import functools
|
| 6 |
import io
|
| 7 |
import os
|
| 8 |
-
from contextlib import suppress
|
| 9 |
-
from dataclasses import dataclass
|
| 10 |
-
from datetime import datetime
|
| 11 |
-
|
| 12 |
import biotite.structure as bs
|
| 13 |
import biotite.structure.io.pdbx as pdbx
|
| 14 |
import numpy as np
|
|
|
|
|
|
|
|
|
|
|
|
|
| 15 |
from biotite.structure.io.pdbx import CIFColumn, CIFData, CIFFile
|
| 16 |
|
| 17 |
from . import esmfold2_residue_constants as residue_constants
|
| 18 |
|
|
|
|
| 19 |
PathOrBuffer = str | os.PathLike | io.StringIO
|
| 20 |
|
| 21 |
PLDDT_B_FACTOR_SCALE = 100.0
|
|
|
|
| 5 |
import functools
|
| 6 |
import io
|
| 7 |
import os
|
|
|
|
|
|
|
|
|
|
|
|
|
| 8 |
import biotite.structure as bs
|
| 9 |
import biotite.structure.io.pdbx as pdbx
|
| 10 |
import numpy as np
|
| 11 |
+
|
| 12 |
+
from contextlib import suppress
|
| 13 |
+
from dataclasses import dataclass
|
| 14 |
+
from datetime import datetime
|
| 15 |
from biotite.structure.io.pdbx import CIFColumn, CIFData, CIFFile
|
| 16 |
|
| 17 |
from . import esmfold2_residue_constants as residue_constants
|
| 18 |
|
| 19 |
+
|
| 20 |
PathOrBuffer = str | os.PathLike | io.StringIO
|
| 21 |
|
| 22 |
PLDDT_B_FACTOR_SCALE = 100.0
|
fastplms/models/esmfold2/esmfold2_molecular_complex.py
CHANGED
|
@@ -11,18 +11,18 @@ from __future__ import annotations
|
|
| 11 |
import io
|
| 12 |
import os
|
| 13 |
import re
|
| 14 |
-
from dataclasses import asdict, dataclass
|
| 15 |
-
from pathlib import Path
|
| 16 |
-
from subprocess import check_output
|
| 17 |
-
from tempfile import TemporaryDirectory
|
| 18 |
-
from typing import TYPE_CHECKING, Any
|
| 19 |
-
|
| 20 |
import biotite.structure as bs
|
| 21 |
import biotite.structure.io.pdbx as pdbx
|
| 22 |
import brotli
|
| 23 |
import msgpack
|
| 24 |
import numpy as np
|
| 25 |
import torch
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 26 |
from biotite.structure.io.pdbx import (
|
| 27 |
CIFCategory,
|
| 28 |
CIFColumn,
|
|
|
|
| 11 |
import io
|
| 12 |
import os
|
| 13 |
import re
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 14 |
import biotite.structure as bs
|
| 15 |
import biotite.structure.io.pdbx as pdbx
|
| 16 |
import brotli
|
| 17 |
import msgpack
|
| 18 |
import numpy as np
|
| 19 |
import torch
|
| 20 |
+
|
| 21 |
+
from dataclasses import asdict, dataclass
|
| 22 |
+
from pathlib import Path
|
| 23 |
+
from subprocess import check_output
|
| 24 |
+
from tempfile import TemporaryDirectory
|
| 25 |
+
from typing import TYPE_CHECKING, Any
|
| 26 |
from biotite.structure.io.pdbx import (
|
| 27 |
CIFCategory,
|
| 28 |
CIFColumn,
|
fastplms/models/esmfold2/esmfold2_msa.py
CHANGED
|
@@ -4,13 +4,13 @@ from __future__ import annotations
|
|
| 4 |
|
| 5 |
import dataclasses
|
| 6 |
import string
|
|
|
|
|
|
|
| 7 |
from collections.abc import Sequence
|
| 8 |
from dataclasses import dataclass
|
| 9 |
from functools import cached_property
|
| 10 |
from itertools import islice
|
| 11 |
from typing import Any
|
| 12 |
-
|
| 13 |
-
import numpy as np
|
| 14 |
from Bio import SeqIO
|
| 15 |
from scipy.spatial.distance import cdist
|
| 16 |
|
|
@@ -20,6 +20,7 @@ from .esmfold2_parsing import FastaEntry, read_sequences, write_sequences
|
|
| 20 |
from .esmfold2_sequential_dataclass import SequentialDataclass
|
| 21 |
from .esmfold2_system import PathOrBuffer
|
| 22 |
|
|
|
|
| 23 |
_A3M_INSERTION_DELETE_TABLE = str.maketrans(
|
| 24 |
dict.fromkeys(string.ascii_lowercase + ".")
|
| 25 |
)
|
|
|
|
| 4 |
|
| 5 |
import dataclasses
|
| 6 |
import string
|
| 7 |
+
import numpy as np
|
| 8 |
+
|
| 9 |
from collections.abc import Sequence
|
| 10 |
from dataclasses import dataclass
|
| 11 |
from functools import cached_property
|
| 12 |
from itertools import islice
|
| 13 |
from typing import Any
|
|
|
|
|
|
|
| 14 |
from Bio import SeqIO
|
| 15 |
from scipy.spatial.distance import cdist
|
| 16 |
|
|
|
|
| 20 |
from .esmfold2_sequential_dataclass import SequentialDataclass
|
| 21 |
from .esmfold2_system import PathOrBuffer
|
| 22 |
|
| 23 |
+
|
| 24 |
_A3M_INSERTION_DELETE_TABLE = str.maketrans(
|
| 25 |
dict.fromkeys(string.ascii_lowercase + ".")
|
| 26 |
)
|
fastplms/models/esmfold2/esmfold2_msa_filter_sequences.py
CHANGED
|
@@ -4,22 +4,24 @@ from __future__ import annotations
|
|
| 4 |
|
| 5 |
import os
|
| 6 |
import tempfile
|
| 7 |
-
from pathlib import Path
|
| 8 |
-
|
| 9 |
import numpy as np
|
| 10 |
|
|
|
|
|
|
|
| 11 |
from .esmfold2_system import run_subprocess_with_errorcheck
|
| 12 |
|
| 13 |
|
| 14 |
def _byte_matrix(array: np.ndarray) -> np.ndarray:
|
| 15 |
"""Return a two-dimensional byte view used for Hamming comparisons."""
|
| 16 |
|
| 17 |
-
|
| 18 |
-
|
|
|
|
| 19 |
|
| 20 |
|
| 21 |
def _hamming_to_all(query: np.ndarray, sequences: np.ndarray) -> np.ndarray:
|
| 22 |
-
|
|
|
|
| 23 |
|
| 24 |
|
| 25 |
def greedy_select_indices(array: np.ndarray, num_seqs: int, mode: str = "max") -> list[int]:
|
|
@@ -48,20 +50,20 @@ def greedy_select_indices(array: np.ndarray, num_seqs: int, mode: str = "max") -
|
|
| 48 |
if depth <= num_seqs:
|
| 49 |
return list(range(depth))
|
| 50 |
|
| 51 |
-
sequences = _byte_matrix(array)
|
| 52 |
selected = [0]
|
| 53 |
-
available = np.ones(depth, dtype=bool)
|
| 54 |
-
available[0] = False
|
| 55 |
-
distance_sum = _hamming_to_all(sequences[0], sequences)
|
| 56 |
choose = np.argmax if mode == "max" else np.argmin
|
| 57 |
|
| 58 |
while len(selected) < num_seqs:
|
| 59 |
-
candidates = np.flatnonzero(available)
|
| 60 |
-
candidate_scores = distance_sum[candidates] / len(selected)
|
| 61 |
next_index = int(candidates[int(choose(candidate_scores))])
|
| 62 |
selected.append(next_index)
|
| 63 |
-
available[next_index] = False
|
| 64 |
-
distance_sum += _hamming_to_all(sequences[next_index], sequences)
|
| 65 |
return sorted(selected)
|
| 66 |
|
| 67 |
|
|
|
|
| 4 |
|
| 5 |
import os
|
| 6 |
import tempfile
|
|
|
|
|
|
|
| 7 |
import numpy as np
|
| 8 |
|
| 9 |
+
from pathlib import Path
|
| 10 |
+
|
| 11 |
from .esmfold2_system import run_subprocess_with_errorcheck
|
| 12 |
|
| 13 |
|
| 14 |
def _byte_matrix(array: np.ndarray) -> np.ndarray:
|
| 15 |
"""Return a two-dimensional byte view used for Hamming comparisons."""
|
| 16 |
|
| 17 |
+
# array: (n, l); w is its row width in bytes after viewing its dtype.
|
| 18 |
+
matrix = np.asarray(array).view(np.uint8) # (n, w)
|
| 19 |
+
return matrix.reshape(matrix.shape[0], -1) # (n, w)
|
| 20 |
|
| 21 |
|
| 22 |
def _hamming_to_all(query: np.ndarray, sequences: np.ndarray) -> np.ndarray:
|
| 23 |
+
# query: (w,); sequences: (n, w).
|
| 24 |
+
return np.not_equal(sequences, query).mean(axis=1, dtype=np.float64) # (n,)
|
| 25 |
|
| 26 |
|
| 27 |
def greedy_select_indices(array: np.ndarray, num_seqs: int, mode: str = "max") -> list[int]:
|
|
|
|
| 50 |
if depth <= num_seqs:
|
| 51 |
return list(range(depth))
|
| 52 |
|
| 53 |
+
sequences = _byte_matrix(array) # (n, w)
|
| 54 |
selected = [0]
|
| 55 |
+
available = np.ones(depth, dtype=bool) # (n,)
|
| 56 |
+
available[0] = False # ()
|
| 57 |
+
distance_sum = _hamming_to_all(sequences[0], sequences) # (n,)
|
| 58 |
choose = np.argmax if mode == "max" else np.argmin
|
| 59 |
|
| 60 |
while len(selected) < num_seqs:
|
| 61 |
+
candidates = np.flatnonzero(available) # (n_remaining,)
|
| 62 |
+
candidate_scores = distance_sum[candidates] / len(selected) # (n_remaining,)
|
| 63 |
next_index = int(candidates[int(choose(candidate_scores))])
|
| 64 |
selected.append(next_index)
|
| 65 |
+
available[next_index] = False # ()
|
| 66 |
+
distance_sum += _hamming_to_all(sequences[next_index], sequences) # (n,)
|
| 67 |
return sorted(selected)
|
| 68 |
|
| 69 |
|
fastplms/models/esmfold2/esmfold2_normalize_coordinates.py
CHANGED
|
@@ -2,23 +2,25 @@
|
|
| 2 |
|
| 3 |
from __future__ import annotations
|
| 4 |
|
| 5 |
-
from typing import TypeVar
|
| 6 |
-
|
| 7 |
import numpy as np
|
| 8 |
import torch
|
|
|
|
|
|
|
| 9 |
from torch import Tensor
|
| 10 |
|
| 11 |
from . import esmfold2_residue_constants as residue_constants
|
| 12 |
from .esmfold2_affine3d import Affine3D
|
| 13 |
|
|
|
|
| 14 |
ArrayOrTensor = TypeVar("ArrayOrTensor", np.ndarray, Tensor)
|
| 15 |
|
| 16 |
|
| 17 |
def atom3_to_backbone_frames(bb_positions: Tensor) -> Affine3D:
|
| 18 |
"""Construct a frame from N, C-alpha, and C positions in ``X``."""
|
| 19 |
|
| 20 |
-
|
| 21 |
-
|
|
|
|
| 22 |
|
| 23 |
|
| 24 |
def index_by_atom_name(
|
|
@@ -26,42 +28,43 @@ def index_by_atom_name(
|
|
| 26 |
atom_names: str | list[str],
|
| 27 |
dim: int = -2,
|
| 28 |
) -> ArrayOrTensor:
|
| 29 |
-
"""Select
|
| 30 |
|
| 31 |
single_atom = isinstance(atom_names, str)
|
| 32 |
names = [atom_names] if single_atom else atom_names
|
| 33 |
indices = [residue_constants.atom_order[name] for name in names]
|
| 34 |
axis = dim % atom37.ndim
|
| 35 |
if isinstance(atom37, Tensor):
|
| 36 |
-
index = torch.tensor(indices, dtype=torch.long, device=atom37.device)
|
| 37 |
-
selected = torch.index_select(atom37, axis, index)
|
| 38 |
else:
|
| 39 |
-
selected = np.take(atom37, indices, axis=axis)
|
| 40 |
return selected.squeeze(axis) if single_atom else selected # type: ignore[return-value]
|
| 41 |
|
| 42 |
|
| 43 |
def get_protein_normalization_frame(coords: Tensor) -> Affine3D:
|
| 44 |
-
"""Build one frame
|
| 45 |
|
| 46 |
-
backbone = index_by_atom_name(coords, ["N", "CA", "C"], dim=-2)
|
| 47 |
-
residue_is_valid = torch.isfinite(backbone).all(dim=-1).all(dim=-1)
|
| 48 |
-
weights = residue_is_valid[..., None, None]
|
| 49 |
-
coordinate_sum = backbone.masked_fill(~weights, 0).sum(dim=-3)
|
| 50 |
-
count = residue_is_valid.sum(dim=-1)[..., None, None]
|
| 51 |
-
mean_backbone = coordinate_sum / (count + 1e-8)
|
| 52 |
-
return atom3_to_backbone_frames(mean_backbone.float())
|
| 53 |
|
| 54 |
|
| 55 |
def apply_frame_to_coords(coords: Tensor, frame: Affine3D) -> Tensor:
|
| 56 |
"""Express atom coordinates ``X`` in the inverse of ``frame``."""
|
| 57 |
|
| 58 |
-
|
| 59 |
-
|
| 60 |
-
|
| 61 |
-
|
|
|
|
| 62 |
|
| 63 |
|
| 64 |
def normalize_coordinates(coords: Tensor) -> Tensor:
|
| 65 |
"""Normalize ``X`` with shape (..., l, 37, 3) to its backbone frame."""
|
| 66 |
|
| 67 |
-
return apply_frame_to_coords(coords, get_protein_normalization_frame(coords))
|
|
|
|
| 2 |
|
| 3 |
from __future__ import annotations
|
| 4 |
|
|
|
|
|
|
|
| 5 |
import numpy as np
|
| 6 |
import torch
|
| 7 |
+
|
| 8 |
+
from typing import TypeVar
|
| 9 |
from torch import Tensor
|
| 10 |
|
| 11 |
from . import esmfold2_residue_constants as residue_constants
|
| 12 |
from .esmfold2_affine3d import Affine3D
|
| 13 |
|
| 14 |
+
|
| 15 |
ArrayOrTensor = TypeVar("ArrayOrTensor", np.ndarray, Tensor)
|
| 16 |
|
| 17 |
|
| 18 |
def atom3_to_backbone_frames(bb_positions: Tensor) -> Affine3D:
|
| 19 |
"""Construct a frame from N, C-alpha, and C positions in ``X``."""
|
| 20 |
|
| 21 |
+
# bb_positions: (..., 3, 3), ordered N, CA, C on the penultimate axis.
|
| 22 |
+
n_position, ca_position, c_position = bb_positions.unbind(dim=-2) # each (..., 3)
|
| 23 |
+
return Affine3D.from_graham_schmidt(c_position, ca_position, n_position) # affine shape: (...,)
|
| 24 |
|
| 25 |
|
| 26 |
def index_by_atom_name(
|
|
|
|
| 28 |
atom_names: str | list[str],
|
| 29 |
dim: int = -2,
|
| 30 |
) -> ArrayOrTensor:
|
| 31 |
+
"""Select named atoms, replacing axis ``dim`` by their count (or removing it)."""
|
| 32 |
|
| 33 |
single_atom = isinstance(atom_names, str)
|
| 34 |
names = [atom_names] if single_atom else atom_names
|
| 35 |
indices = [residue_constants.atom_order[name] for name in names]
|
| 36 |
axis = dim % atom37.ndim
|
| 37 |
if isinstance(atom37, Tensor):
|
| 38 |
+
index = torch.tensor(indices, dtype=torch.long, device=atom37.device) # (n_names,)
|
| 39 |
+
selected = torch.index_select(atom37, axis, index) # atom37.shape with axis = n_names
|
| 40 |
else:
|
| 41 |
+
selected = np.take(atom37, indices, axis=axis) # atom37.shape with axis = n_names
|
| 42 |
return selected.squeeze(axis) if single_atom else selected # type: ignore[return-value]
|
| 43 |
|
| 44 |
|
| 45 |
def get_protein_normalization_frame(coords: Tensor) -> Affine3D:
|
| 46 |
+
"""Build one frame per batch item from coordinates of shape (..., l, 37, 3)."""
|
| 47 |
|
| 48 |
+
backbone = index_by_atom_name(coords, ["N", "CA", "C"], dim=-2) # (..., l, 3, 3)
|
| 49 |
+
residue_is_valid = torch.isfinite(backbone).all(dim=-1).all(dim=-1) # (..., l)
|
| 50 |
+
weights = residue_is_valid[..., None, None] # (..., l, 1, 1)
|
| 51 |
+
coordinate_sum = backbone.masked_fill(~weights, 0).sum(dim=-3) # (..., 3, 3)
|
| 52 |
+
count = residue_is_valid.sum(dim=-1)[..., None, None] # (..., 1, 1)
|
| 53 |
+
mean_backbone = coordinate_sum / (count + 1e-8) # (..., 3, 3)
|
| 54 |
+
return atom3_to_backbone_frames(mean_backbone.float()) # affine shape: (...,)
|
| 55 |
|
| 56 |
|
| 57 |
def apply_frame_to_coords(coords: Tensor, frame: Affine3D) -> Tensor:
|
| 58 |
"""Express atom coordinates ``X`` in the inverse of ``frame``."""
|
| 59 |
|
| 60 |
+
# coords: (..., l, 37, 3); frame shape: (...,).
|
| 61 |
+
transformed = frame[..., None, None].invert().apply(coords) # (..., l, 37, 3)
|
| 62 |
+
frame_is_valid = frame.trans.norm(dim=-1) > 0 # (...,)
|
| 63 |
+
normalized = torch.where(frame_is_valid[..., None, None, None], transformed, coords) # (..., l, 37, 3)
|
| 64 |
+
return normalized.masked_fill(torch.isinf(coords), torch.inf) # (..., l, 37, 3)
|
| 65 |
|
| 66 |
|
| 67 |
def normalize_coordinates(coords: Tensor) -> Tensor:
|
| 68 |
"""Normalize ``X`` with shape (..., l, 37, 3) to its backbone frame."""
|
| 69 |
|
| 70 |
+
return apply_frame_to_coords(coords, get_protein_normalization_frame(coords)) # (..., l, 37, 3)
|
fastplms/models/esmfold2/esmfold2_output.py
CHANGED
|
@@ -2,14 +2,14 @@
|
|
| 2 |
|
| 3 |
from __future__ import annotations
|
| 4 |
|
|
|
|
|
|
|
|
|
|
| 5 |
from collections.abc import Iterable
|
| 6 |
from dataclasses import dataclass, field
|
| 7 |
from itertools import groupby
|
| 8 |
from typing import Any
|
| 9 |
|
| 10 |
-
import numpy as np
|
| 11 |
-
import torch
|
| 12 |
-
|
| 13 |
from .esmfold2_constants import ELEMENT_NUMBER_TO_SYMBOL, MOL_TYPE_NONPOLYMER
|
| 14 |
from .esmfold2_molecular_complex import MolecularComplex, MolecularComplexMetadata
|
| 15 |
|
|
@@ -92,17 +92,18 @@ def build_molecular_complex_from_features(
|
|
| 92 |
tokens are collapsed into one non-polymer residue per chain.
|
| 93 |
"""
|
| 94 |
|
| 95 |
-
|
| 96 |
-
|
| 97 |
-
|
| 98 |
-
|
| 99 |
-
|
|
|
|
| 100 |
records = _ComplexRecords()
|
| 101 |
|
| 102 |
-
def decode_atoms(tokens: Iterable[Any]):
|
| 103 |
for token in tokens:
|
| 104 |
for atom_index in range(token.atom_start, token.atom_start + token.atom_count):
|
| 105 |
-
if
|
| 106 |
yield (
|
| 107 |
X[atom_index].tolist(),
|
| 108 |
get_element_symbol(int(elements[atom_index])),
|
|
|
|
| 2 |
|
| 3 |
from __future__ import annotations
|
| 4 |
|
| 5 |
+
import numpy as np
|
| 6 |
+
import torch
|
| 7 |
+
|
| 8 |
from collections.abc import Iterable
|
| 9 |
from dataclasses import dataclass, field
|
| 10 |
from itertools import groupby
|
| 11 |
from typing import Any
|
| 12 |
|
|
|
|
|
|
|
|
|
|
| 13 |
from .esmfold2_constants import ELEMENT_NUMBER_TO_SYMBOL, MOL_TYPE_NONPOLYMER
|
| 14 |
from .esmfold2_molecular_complex import MolecularComplex, MolecularComplexMetadata
|
| 15 |
|
|
|
|
| 92 |
tokens are collapsed into one non-polymer residue per chain.
|
| 93 |
"""
|
| 94 |
|
| 95 |
+
# a counts padded atoms; l counts model tokens (ligand tokens may be atoms).
|
| 96 |
+
present_atoms = atom_mask.bool().cpu().numpy() # (a,)
|
| 97 |
+
X = coords.float().cpu().numpy() # (a, 3)
|
| 98 |
+
atom_names = ref_atom_name_chars.cpu().numpy() # (a, 4)
|
| 99 |
+
elements = ref_element.cpu().numpy() # (a,)
|
| 100 |
+
confidence = None if plddt is None else plddt.float().cpu().numpy() # (l,) or None
|
| 101 |
records = _ComplexRecords()
|
| 102 |
|
| 103 |
+
def decode_atoms(tokens: Iterable[Any]) -> Iterable[tuple[list[float], str, str]]:
|
| 104 |
for token in tokens:
|
| 105 |
for atom_index in range(token.atom_start, token.atom_start + token.atom_count):
|
| 106 |
+
if present_atoms[atom_index]:
|
| 107 |
yield (
|
| 108 |
X[atom_index].tolist(),
|
| 109 |
get_element_symbol(int(elements[atom_index])),
|
fastplms/models/esmfold2/esmfold2_paired_msa.py
CHANGED
|
@@ -3,10 +3,10 @@
|
|
| 3 |
from __future__ import annotations
|
| 4 |
|
| 5 |
import re
|
| 6 |
-
from dataclasses import dataclass
|
| 7 |
-
|
| 8 |
import numpy as np
|
| 9 |
|
|
|
|
|
|
|
| 10 |
from .esmfold2_constants import (
|
| 11 |
MSA_GAP_TOKEN_ID,
|
| 12 |
PROTEIN_3TO1,
|
|
@@ -15,6 +15,7 @@ from .esmfold2_constants import (
|
|
| 15 |
)
|
| 16 |
from .esmfold2_msa import MSA
|
| 17 |
|
|
|
|
| 18 |
_TAXONOMY_PATTERN = re.compile(r"key=(-?\d+)")
|
| 19 |
|
| 20 |
|
|
|
|
| 3 |
from __future__ import annotations
|
| 4 |
|
| 5 |
import re
|
|
|
|
|
|
|
| 6 |
import numpy as np
|
| 7 |
|
| 8 |
+
from dataclasses import dataclass
|
| 9 |
+
|
| 10 |
from .esmfold2_constants import (
|
| 11 |
MSA_GAP_TOKEN_ID,
|
| 12 |
PROTEIN_3TO1,
|
|
|
|
| 15 |
)
|
| 16 |
from .esmfold2_msa import MSA
|
| 17 |
|
| 18 |
+
|
| 19 |
_TAXONOMY_PATTERN = re.compile(r"key=(-?\d+)")
|
| 20 |
|
| 21 |
|
fastplms/models/esmfold2/esmfold2_parsing.py
CHANGED
|
@@ -4,8 +4,9 @@ from __future__ import annotations
|
|
| 4 |
|
| 5 |
import gzip
|
| 6 |
import io
|
|
|
|
| 7 |
from collections.abc import Generator, Iterable
|
| 8 |
-
from contextlib import nullcontext
|
| 9 |
from pathlib import Path
|
| 10 |
from typing import NamedTuple, TextIO
|
| 11 |
|
|
@@ -45,7 +46,7 @@ def parse_fasta(text: str) -> Generator[FastaEntry, None, None]:
|
|
| 45 |
raise ValueError("Found no sequences in input")
|
| 46 |
|
| 47 |
|
| 48 |
-
def _open_reader(source: PathOrBuffer):
|
| 49 |
if isinstance(source, io.TextIOBase):
|
| 50 |
return nullcontext(source)
|
| 51 |
path = Path(source)
|
|
@@ -93,7 +94,7 @@ def append_fasta_sequence(header: str, sequence: str, path: str | Path) -> None:
|
|
| 93 |
handle.write(f">{header}\n{sequence}\n")
|
| 94 |
|
| 95 |
|
| 96 |
-
def _open_writer(destination: PathOrBuffer):
|
| 97 |
if isinstance(destination, io.TextIOBase):
|
| 98 |
return nullcontext(destination)
|
| 99 |
path = Path(destination)
|
|
|
|
| 4 |
|
| 5 |
import gzip
|
| 6 |
import io
|
| 7 |
+
|
| 8 |
from collections.abc import Generator, Iterable
|
| 9 |
+
from contextlib import AbstractContextManager, nullcontext
|
| 10 |
from pathlib import Path
|
| 11 |
from typing import NamedTuple, TextIO
|
| 12 |
|
|
|
|
| 46 |
raise ValueError("Found no sequences in input")
|
| 47 |
|
| 48 |
|
| 49 |
+
def _open_reader(source: PathOrBuffer) -> AbstractContextManager[TextIO]:
|
| 50 |
if isinstance(source, io.TextIOBase):
|
| 51 |
return nullcontext(source)
|
| 52 |
path = Path(source)
|
|
|
|
| 94 |
handle.write(f">{header}\n{sequence}\n")
|
| 95 |
|
| 96 |
|
| 97 |
+
def _open_writer(destination: PathOrBuffer) -> AbstractContextManager[TextIO]:
|
| 98 |
if isinstance(destination, io.TextIOBase):
|
| 99 |
return nullcontext(destination)
|
| 100 |
path = Path(destination)
|
fastplms/models/esmfold2/esmfold2_predicted_aligned_error.py
CHANGED
|
@@ -4,16 +4,19 @@ from __future__ import annotations
|
|
| 4 |
|
| 5 |
import torch
|
| 6 |
import torch.nn.functional as F
|
|
|
|
| 7 |
from torch import Tensor
|
| 8 |
|
| 9 |
from .esmfold2_affine3d import Affine3D
|
| 10 |
|
|
|
|
| 11 |
_CPU_DEVICE = torch.device("cpu")
|
| 12 |
|
| 13 |
|
| 14 |
def _compute_pae_masks(mask: Tensor) -> Tensor:
|
| 15 |
-
|
| 16 |
-
|
|
|
|
| 17 |
|
| 18 |
|
| 19 |
def _pae_bins(
|
|
@@ -23,19 +26,20 @@ def _pae_bins(
|
|
| 23 |
) -> Tensor:
|
| 24 |
"""Return the representative distance for each PAE probability bin."""
|
| 25 |
|
| 26 |
-
boundaries = torch.linspace(0, max_bin, steps=num_bins - 1, device=device)
|
| 27 |
width = max_bin / (num_bins - 2)
|
| 28 |
-
centers = boundaries + width / 2
|
| 29 |
-
overflow_center = centers[-1:] + width
|
| 30 |
-
return torch.cat((centers, overflow_center))
|
| 31 |
|
| 32 |
|
| 33 |
def _masked_probabilities(logits: Tensor, pair_mask: Tensor) -> Tensor:
|
|
|
|
| 34 |
masked_logits = logits.masked_fill(
|
| 35 |
~pair_mask.unsqueeze(-1),
|
| 36 |
torch.finfo(logits.dtype).min,
|
| 37 |
-
)
|
| 38 |
-
return masked_logits.softmax(dim=-1)
|
| 39 |
|
| 40 |
|
| 41 |
def masked_mean(
|
|
@@ -46,10 +50,11 @@ def masked_mean(
|
|
| 46 |
) -> Tensor:
|
| 47 |
"""Average values over true entries of a broadcast-compatible mask."""
|
| 48 |
|
| 49 |
-
|
| 50 |
-
|
| 51 |
-
|
| 52 |
-
|
|
|
|
| 53 |
|
| 54 |
|
| 55 |
def compute_predicted_aligned_error(
|
|
@@ -61,30 +66,31 @@ def compute_predicted_aligned_error(
|
|
| 61 |
"""Convert PAE logits ``X`` with shape (..., l, l, n) to distances."""
|
| 62 |
|
| 63 |
del sequence_id
|
| 64 |
-
pair_mask = _compute_pae_masks(aa_mask)
|
| 65 |
-
probabilities = _masked_probabilities(logits, pair_mask)
|
| 66 |
-
centers = _pae_bins(max_bin, logits.shape[-1], logits.device)
|
| 67 |
-
return torch.sum(probabilities * centers, dim=-1)
|
| 68 |
|
| 69 |
|
| 70 |
@torch.no_grad()
|
| 71 |
def compute_tm(logits: Tensor, aa_mask: Tensor, max_bin: float = 31.0) -> Tensor:
|
| 72 |
-
"""Estimate TM score from
|
| 73 |
|
| 74 |
-
pair_mask = _compute_pae_masks(aa_mask)
|
| 75 |
-
sequence_lengths = aa_mask.sum(dim=-1, keepdim=True)
|
| 76 |
-
centers = _pae_bins(max_bin, logits.shape[-1], logits.device)
|
| 77 |
-
distance_scale = 1.24 * (sequence_lengths.clamp_min(19) - 15) ** (1 / 3) - 1.8
|
| 78 |
-
tm_weights = 1.0 / (1 + (centers / distance_scale.unsqueeze(-1)) ** 2)
|
| 79 |
-
probabilities = _masked_probabilities(logits, pair_mask)
|
| 80 |
-
score_per_pair = torch.sum(probabilities * tm_weights.unsqueeze(-2), dim=-1)
|
| 81 |
-
score_per_anchor = masked_mean(pair_mask, score_per_pair, dim=-1)
|
| 82 |
-
return score_per_anchor.max(dim=-1).values
|
| 83 |
|
| 84 |
|
| 85 |
def _local_coordinates(frames: Affine3D) -> Tensor:
|
| 86 |
-
|
| 87 |
-
|
|
|
|
| 88 |
|
| 89 |
|
| 90 |
def tm_loss(
|
|
@@ -96,32 +102,32 @@ def tm_loss(
|
|
| 96 |
sequence_id: Tensor | None = None,
|
| 97 |
max_bin: float = 31,
|
| 98 |
) -> Tensor:
|
| 99 |
-
"""Cross-entropy loss for
|
| 100 |
|
| 101 |
del sequence_id
|
| 102 |
-
predicted_frames = Affine3D.from_tensor(pred_affine)
|
| 103 |
-
target_frames = Affine3D.from_tensor(targ_affine)
|
| 104 |
with torch.no_grad():
|
| 105 |
squared_error = (
|
| 106 |
(_local_coordinates(predicted_frames) - _local_coordinates(target_frames))
|
| 107 |
.square()
|
| 108 |
.sum(dim=-1)
|
| 109 |
-
)
|
| 110 |
boundaries = torch.linspace(
|
| 111 |
0,
|
| 112 |
max_bin,
|
| 113 |
logits.shape[-1] - 1,
|
| 114 |
device=logits.device,
|
| 115 |
-
).square()
|
| 116 |
-
target_bins = (squared_error[..., None] > boundaries).sum(dim=-1).long()
|
| 117 |
|
| 118 |
cross_entropy = F.cross_entropy(
|
| 119 |
logits.movedim(3, 1),
|
| 120 |
target_bins,
|
| 121 |
reduction="none",
|
| 122 |
-
)
|
| 123 |
-
pair_mask = _compute_pae_masks(targ_mask)
|
| 124 |
-
loss_per_sample = masked_mean(pair_mask, cross_entropy, dim=(-1, -2))
|
| 125 |
if tm_mask is None:
|
| 126 |
-
return loss_per_sample.mean()
|
| 127 |
-
return masked_mean(tm_mask, loss_per_sample)
|
|
|
|
| 4 |
|
| 5 |
import torch
|
| 6 |
import torch.nn.functional as F
|
| 7 |
+
|
| 8 |
from torch import Tensor
|
| 9 |
|
| 10 |
from .esmfold2_affine3d import Affine3D
|
| 11 |
|
| 12 |
+
|
| 13 |
_CPU_DEVICE = torch.device("cpu")
|
| 14 |
|
| 15 |
|
| 16 |
def _compute_pae_masks(mask: Tensor) -> Tensor:
|
| 17 |
+
# mask: (..., l), where l is the residue count.
|
| 18 |
+
residue_mask = mask.bool() # (..., l)
|
| 19 |
+
return residue_mask.unsqueeze(-1) & residue_mask.unsqueeze(-2) # (..., l, l)
|
| 20 |
|
| 21 |
|
| 22 |
def _pae_bins(
|
|
|
|
| 26 |
) -> Tensor:
|
| 27 |
"""Return the representative distance for each PAE probability bin."""
|
| 28 |
|
| 29 |
+
boundaries = torch.linspace(0, max_bin, steps=num_bins - 1, device=device) # (n_bins - 1,)
|
| 30 |
width = max_bin / (num_bins - 2)
|
| 31 |
+
centers = boundaries + width / 2 # (n_bins - 1,)
|
| 32 |
+
overflow_center = centers[-1:] + width # (1,)
|
| 33 |
+
return torch.cat((centers, overflow_center)) # (n_bins,)
|
| 34 |
|
| 35 |
|
| 36 |
def _masked_probabilities(logits: Tensor, pair_mask: Tensor) -> Tensor:
|
| 37 |
+
# logits: (..., l, l, n_bins); pair_mask: (..., l, l).
|
| 38 |
masked_logits = logits.masked_fill(
|
| 39 |
~pair_mask.unsqueeze(-1),
|
| 40 |
torch.finfo(logits.dtype).min,
|
| 41 |
+
) # (..., l, l, n_bins)
|
| 42 |
+
return masked_logits.softmax(dim=-1) # (..., l, l, n_bins)
|
| 43 |
|
| 44 |
|
| 45 |
def masked_mean(
|
|
|
|
| 50 |
) -> Tensor:
|
| 51 |
"""Average values over true entries of a broadcast-compatible mask."""
|
| 52 |
|
| 53 |
+
# value has arbitrary shape; reduced_shape removes the axes named by dim.
|
| 54 |
+
weights = mask.expand_as(value) # value.shape
|
| 55 |
+
weighted_sum = torch.sum(weights * value, dim=dim) # reduced_shape
|
| 56 |
+
weight_sum = torch.sum(weights, dim=dim) # reduced_shape
|
| 57 |
+
return weighted_sum / (weight_sum + eps) # reduced_shape
|
| 58 |
|
| 59 |
|
| 60 |
def compute_predicted_aligned_error(
|
|
|
|
| 66 |
"""Convert PAE logits ``X`` with shape (..., l, l, n) to distances."""
|
| 67 |
|
| 68 |
del sequence_id
|
| 69 |
+
pair_mask = _compute_pae_masks(aa_mask) # (..., l, l)
|
| 70 |
+
probabilities = _masked_probabilities(logits, pair_mask) # (..., l, l, n_bins)
|
| 71 |
+
centers = _pae_bins(max_bin, logits.shape[-1], logits.device) # (n_bins,)
|
| 72 |
+
return torch.sum(probabilities * centers, dim=-1) # (..., l, l)
|
| 73 |
|
| 74 |
|
| 75 |
@torch.no_grad()
|
| 76 |
def compute_tm(logits: Tensor, aa_mask: Tensor, max_bin: float = 31.0) -> Tensor:
|
| 77 |
+
"""Estimate TM score from logits (..., l, l, n_bins) and residue mask (..., l)."""
|
| 78 |
|
| 79 |
+
pair_mask = _compute_pae_masks(aa_mask) # (..., l, l)
|
| 80 |
+
sequence_lengths = aa_mask.sum(dim=-1, keepdim=True) # (..., 1)
|
| 81 |
+
centers = _pae_bins(max_bin, logits.shape[-1], logits.device) # (n_bins,)
|
| 82 |
+
distance_scale = 1.24 * (sequence_lengths.clamp_min(19) - 15) ** (1 / 3) - 1.8 # (..., 1)
|
| 83 |
+
tm_weights = 1.0 / (1 + (centers / distance_scale.unsqueeze(-1)) ** 2) # (..., 1, n_bins)
|
| 84 |
+
probabilities = _masked_probabilities(logits, pair_mask) # (..., l, l, n_bins)
|
| 85 |
+
score_per_pair = torch.sum(probabilities * tm_weights.unsqueeze(-2), dim=-1) # (..., l, l)
|
| 86 |
+
score_per_anchor = masked_mean(pair_mask, score_per_pair, dim=-1) # (..., l)
|
| 87 |
+
return score_per_anchor.max(dim=-1).values # (...,)
|
| 88 |
|
| 89 |
|
| 90 |
def _local_coordinates(frames: Affine3D) -> Tensor:
|
| 91 |
+
# frames.shape: (..., l); trans: (..., l, 3).
|
| 92 |
+
origins = frames.trans[..., None, :, :] # (..., 1, l, 3)
|
| 93 |
+
return frames.invert()[..., None].apply(origins) # (..., l, l, 3)
|
| 94 |
|
| 95 |
|
| 96 |
def tm_loss(
|
|
|
|
| 102 |
sequence_id: Tensor | None = None,
|
| 103 |
max_bin: float = 31,
|
| 104 |
) -> Tensor:
|
| 105 |
+
"""Cross-entropy loss for logits (b, l, l, n_bins) and residue frames (b, l)."""
|
| 106 |
|
| 107 |
del sequence_id
|
| 108 |
+
predicted_frames = Affine3D.from_tensor(pred_affine) # frame shape: (b, l)
|
| 109 |
+
target_frames = Affine3D.from_tensor(targ_affine) # frame shape: (b, l)
|
| 110 |
with torch.no_grad():
|
| 111 |
squared_error = (
|
| 112 |
(_local_coordinates(predicted_frames) - _local_coordinates(target_frames))
|
| 113 |
.square()
|
| 114 |
.sum(dim=-1)
|
| 115 |
+
) # (b, l, l)
|
| 116 |
boundaries = torch.linspace(
|
| 117 |
0,
|
| 118 |
max_bin,
|
| 119 |
logits.shape[-1] - 1,
|
| 120 |
device=logits.device,
|
| 121 |
+
).square() # (n_bins - 1,)
|
| 122 |
+
target_bins = (squared_error[..., None] > boundaries).sum(dim=-1).long() # (b, l, l)
|
| 123 |
|
| 124 |
cross_entropy = F.cross_entropy(
|
| 125 |
logits.movedim(3, 1),
|
| 126 |
target_bins,
|
| 127 |
reduction="none",
|
| 128 |
+
) # (b, l, l)
|
| 129 |
+
pair_mask = _compute_pae_masks(targ_mask) # (b, l, l)
|
| 130 |
+
loss_per_sample = masked_mean(pair_mask, cross_entropy, dim=(-1, -2)) # (b,)
|
| 131 |
if tm_mask is None:
|
| 132 |
+
return loss_per_sample.mean() # ()
|
| 133 |
+
return masked_mean(tm_mask, loss_per_sample) # ()
|
fastplms/models/esmfold2/esmfold2_prepare_input.py
CHANGED
|
@@ -10,15 +10,15 @@ from __future__ import annotations
|
|
| 10 |
|
| 11 |
import math
|
| 12 |
import warnings
|
|
|
|
|
|
|
|
|
|
| 13 |
from collections import defaultdict
|
| 14 |
from contextlib import suppress
|
| 15 |
from dataclasses import dataclass, field
|
| 16 |
from itertools import combinations
|
| 17 |
from typing import Any
|
| 18 |
|
| 19 |
-
import numpy as np
|
| 20 |
-
import torch
|
| 21 |
-
|
| 22 |
from .esmfold2_conformers import (
|
| 23 |
get_ccd_leaving_atoms,
|
| 24 |
get_idealized_atom_pos,
|
|
@@ -62,6 +62,7 @@ from .esmfold2_types import (
|
|
| 62 |
StructurePredictionInput,
|
| 63 |
)
|
| 64 |
|
|
|
|
| 65 |
_ZERO_POS = np.zeros(3, dtype=np.float32)
|
| 66 |
_ENCODE_ATOM_NAME_CACHE: dict[str, list[int]] = {}
|
| 67 |
_ELEMENT_ATOMIC_NUM_CACHE: dict[str, int] = {}
|
|
|
|
| 10 |
|
| 11 |
import math
|
| 12 |
import warnings
|
| 13 |
+
import numpy as np
|
| 14 |
+
import torch
|
| 15 |
+
|
| 16 |
from collections import defaultdict
|
| 17 |
from contextlib import suppress
|
| 18 |
from dataclasses import dataclass, field
|
| 19 |
from itertools import combinations
|
| 20 |
from typing import Any
|
| 21 |
|
|
|
|
|
|
|
|
|
|
| 22 |
from .esmfold2_conformers import (
|
| 23 |
get_ccd_leaving_atoms,
|
| 24 |
get_idealized_atom_pos,
|
|
|
|
| 62 |
StructurePredictionInput,
|
| 63 |
)
|
| 64 |
|
| 65 |
+
|
| 66 |
_ZERO_POS = np.zeros(3, dtype=np.float32)
|
| 67 |
_ENCODE_ATOM_NAME_CACHE: dict[str, list[int]] = {}
|
| 68 |
_ELEMENT_ATOMIC_NUM_CACHE: dict[str, int] = {}
|
fastplms/models/esmfold2/esmfold2_processor.py
CHANGED
|
@@ -2,13 +2,13 @@
|
|
| 2 |
|
| 3 |
from __future__ import annotations
|
| 4 |
|
|
|
|
|
|
|
|
|
|
| 5 |
from collections.abc import Mapping
|
| 6 |
from dataclasses import dataclass
|
| 7 |
from pathlib import Path
|
| 8 |
from typing import Any
|
| 9 |
-
|
| 10 |
-
import numpy as np
|
| 11 |
-
import torch
|
| 12 |
from torch import Tensor
|
| 13 |
from tqdm.auto import tqdm
|
| 14 |
|
|
@@ -20,6 +20,7 @@ from .esmfold2_types import MSA, Modification, ProteinInput, StructurePrediction
|
|
| 20 |
from .modeling_esmfold2_common import MSA_CONDITIONING_INPUT_NAMES
|
| 21 |
from .reproducibility import seed_context
|
| 22 |
|
|
|
|
| 23 |
# Backward-compatible private alias for the pinned parity helpers. New callers
|
| 24 |
# should import ``seed_context`` from the public ``fastplms.models.esmfold2``
|
| 25 |
# package instead of reaching into implementation modules.
|
|
|
|
| 2 |
|
| 3 |
from __future__ import annotations
|
| 4 |
|
| 5 |
+
import numpy as np
|
| 6 |
+
import torch
|
| 7 |
+
|
| 8 |
from collections.abc import Mapping
|
| 9 |
from dataclasses import dataclass
|
| 10 |
from pathlib import Path
|
| 11 |
from typing import Any
|
|
|
|
|
|
|
|
|
|
| 12 |
from torch import Tensor
|
| 13 |
from tqdm.auto import tqdm
|
| 14 |
|
|
|
|
| 20 |
from .modeling_esmfold2_common import MSA_CONDITIONING_INPUT_NAMES
|
| 21 |
from .reproducibility import seed_context
|
| 22 |
|
| 23 |
+
|
| 24 |
# Backward-compatible private alias for the pinned parity helpers. New callers
|
| 25 |
# should import ``seed_context`` from the public ``fastplms.models.esmfold2``
|
| 26 |
# package instead of reaching into implementation modules.
|
fastplms/models/esmfold2/esmfold2_protein_chain.py
CHANGED
|
@@ -4,18 +4,18 @@ from __future__ import annotations
|
|
| 4 |
|
| 5 |
import io
|
| 6 |
import warnings
|
| 7 |
-
from collections.abc import Mapping, Sequence
|
| 8 |
-
from dataclasses import asdict, dataclass, replace
|
| 9 |
-
from functools import cached_property
|
| 10 |
-
from pathlib import Path
|
| 11 |
-
from typing import Any
|
| 12 |
-
|
| 13 |
import biotite.structure as bs
|
| 14 |
import brotli
|
| 15 |
import msgpack
|
| 16 |
import msgpack_numpy
|
| 17 |
import numpy as np
|
| 18 |
import torch
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 19 |
from biotite.database import rcsb
|
| 20 |
from biotite.structure.io.pdb import PDBFile
|
| 21 |
from biotite.structure.io.pdbx import CIFCategory, CIFColumn, CIFData, CIFFile
|
|
@@ -42,6 +42,7 @@ from .esmfold2_normalize_coordinates import (
|
|
| 42 |
from .esmfold2_protein_structure import index_by_atom_name
|
| 43 |
from .esmfold2_utils_types import PathOrBuffer
|
| 44 |
|
|
|
|
| 45 |
CHAIN_ID_CONST = "A"
|
| 46 |
|
| 47 |
|
|
@@ -70,27 +71,30 @@ def infer_cb(
|
|
| 70 |
dihedral: float = -2.143,
|
| 71 |
):
|
| 72 |
"""Infer C-beta coordinates from C, N, and C-alpha coordinates."""
|
|
|
|
| 73 |
|
| 74 |
def normalize(X: np.ndarray) -> np.ndarray:
|
| 75 |
-
|
|
|
|
| 76 |
|
| 77 |
with np.errstate(invalid="ignore"):
|
| 78 |
-
n_to_ca = N - Ca
|
| 79 |
-
n_to_c = N - C
|
| 80 |
-
axis = normalize(n_to_ca)
|
| 81 |
-
normal = normalize(np.cross(n_to_c, axis))
|
| 82 |
-
basis = (axis, np.cross(normal, axis), normal)
|
| 83 |
offsets = (
|
| 84 |
bond_length * np.cos(bond_angle),
|
| 85 |
bond_length * np.sin(bond_angle) * np.cos(dihedral),
|
| 86 |
-bond_length * np.sin(bond_angle) * np.sin(dihedral),
|
| 87 |
-
)
|
| 88 |
-
return Ca + sum(vector * offset for vector, offset in zip(basis, offsets, strict=False))
|
| 89 |
|
| 90 |
|
| 91 |
def chain_to_ndarray(
|
| 92 |
atom_array: bs.AtomArray, mmcif: MmcifWrapper, chain_id: str, is_predicted=False
|
| 93 |
):
|
|
|
|
| 94 |
if not isinstance(atom_array, bs.AtomArray):
|
| 95 |
raise TypeError("atom_array must be a biotite AtomArray.")
|
| 96 |
if not isinstance(mmcif, MmcifWrapper):
|
|
@@ -106,14 +110,14 @@ def chain_to_ndarray(
|
|
| 106 |
num_res = len(mmcif.chain_to_seqres[chain_id])
|
| 107 |
sequence = mmcif.chain_to_seqres[chain_id]
|
| 108 |
|
| 109 |
-
atom_positions = np.full([num_res, residue_constants.atom_type_num, 3], np.nan)
|
| 110 |
-
atom_mask = np.full([num_res, residue_constants.atom_type_num], False, dtype=bool)
|
| 111 |
-
residue_index = np.full([num_res], -1, dtype=np.int64)
|
| 112 |
-
insertion_code = np.full([num_res], "", dtype="<U4")
|
| 113 |
|
| 114 |
-
confidence = np.ones([num_res], dtype=np.float32)
|
| 115 |
|
| 116 |
-
chain = atom_array[atom_array.chain_id == chain_id]
|
| 117 |
if not isinstance(chain, bs.AtomArray):
|
| 118 |
raise RuntimeError("Biotite selection did not return an AtomArray.")
|
| 119 |
for res_index in range(num_res):
|
|
@@ -122,13 +126,13 @@ def chain_to_ndarray(
|
|
| 122 |
if res_at_position.residue_number is None:
|
| 123 |
continue
|
| 124 |
|
| 125 |
-
residue_index[res_index] = res_at_position.residue_number
|
| 126 |
-
insertion_code[res_index] = res_at_position.insertion_code
|
| 127 |
res = chain[
|
| 128 |
(chain.res_id == res_at_position.residue_number)
|
| 129 |
& (chain.ins_code == res_at_position.insertion_code)
|
| 130 |
& (chain.hetero == res_at_position.hetflag)
|
| 131 |
-
]
|
| 132 |
if not isinstance(res, bs.AtomArray):
|
| 133 |
raise RuntimeError("Biotite residue selection did not return an AtomArray.")
|
| 134 |
|
|
@@ -140,10 +144,10 @@ def chain_to_ndarray(
|
|
| 140 |
atom_name = "SD"
|
| 141 |
|
| 142 |
if atom_name in residue_constants.atom_order:
|
| 143 |
-
atom_positions[res_index, residue_constants.atom_order[atom_name]] = atom.coord
|
| 144 |
-
atom_mask[res_index, residue_constants.atom_order[atom_name]] = True
|
| 145 |
if is_predicted and atom_name == "CA":
|
| 146 |
-
confidence[res_index] = atom.b_factor / PLDDT_B_FACTOR_SCALE
|
| 147 |
|
| 148 |
if not sequence or not all(sequence):
|
| 149 |
raise ValueError("Some residue name was not specified correctly.")
|
|
@@ -155,22 +159,23 @@ def chain_to_ndarray(
|
|
| 155 |
insertion_code,
|
| 156 |
confidence,
|
| 157 |
entity_id,
|
| 158 |
-
)
|
| 159 |
|
| 160 |
|
| 161 |
@dataclass(frozen=True)
|
| 162 |
class ProteinChain:
|
| 163 |
"""Dataclass with atom37 representation of a single protein chain."""
|
| 164 |
|
|
|
|
| 165 |
id: str
|
| 166 |
sequence: str
|
| 167 |
chain_id: str # author chain id - mutable
|
| 168 |
entity_id: int | None
|
| 169 |
-
residue_index: np.ndarray
|
| 170 |
-
insertion_code: np.ndarray
|
| 171 |
-
atom37_positions: np.ndarray
|
| 172 |
-
atom37_mask: np.ndarray
|
| 173 |
-
confidence: np.ndarray
|
| 174 |
mmcif: MmcifWrapper | None = None
|
| 175 |
atom37_confidence: np.ndarray | None = None # P has shape (l, 37).
|
| 176 |
|
|
@@ -186,7 +191,7 @@ class ProteinChain:
|
|
| 186 |
"""Yield every protein chain represented in an mmCIF structure."""
|
| 187 |
mmcif = path if isinstance(path, MmcifWrapper) else MmcifWrapper.read(path, id)
|
| 188 |
for chain in bs.chain_iter(mmcif.structure):
|
| 189 |
-
chain = chain[bs.filter_amino_acids(chain) & ~chain.hetero]
|
| 190 |
if len(chain) == 0:
|
| 191 |
continue
|
| 192 |
chain_id = chain.chain_id[0]
|
|
@@ -206,7 +211,7 @@ class ProteinChain:
|
|
| 206 |
insertion_code,
|
| 207 |
confidence,
|
| 208 |
_,
|
| 209 |
-
) = chain_to_ndarray(chain, mmcif, chain_id, is_predicted)
|
| 210 |
if not all(sequence):
|
| 211 |
raise ValueError("Some residue name was not specified correctly.")
|
| 212 |
|
|
@@ -280,7 +285,7 @@ class ProteinChain:
|
|
| 280 |
stacklevel=2,
|
| 281 |
)
|
| 282 |
|
| 283 |
-
atom_array = mmcif.structure
|
| 284 |
(
|
| 285 |
sequence,
|
| 286 |
atom_positions,
|
|
@@ -289,7 +294,7 @@ class ProteinChain:
|
|
| 289 |
insertion_code,
|
| 290 |
confidence,
|
| 291 |
_,
|
| 292 |
-
) = chain_to_ndarray(atom_array, mmcif, chain_id, is_predicted)
|
| 293 |
if not all(sequence):
|
| 294 |
raise ValueError("Some residue name was not specified correctly.")
|
| 295 |
|
|
@@ -319,15 +324,17 @@ class ProteinChain:
|
|
| 319 |
insertion_code: np.ndarray | None = None,
|
| 320 |
confidence: np.ndarray | torch.Tensor | None = None,
|
| 321 |
):
|
|
|
|
|
|
|
| 322 |
if isinstance(atom37_positions, torch.Tensor):
|
| 323 |
-
atom37_positions = atom37_positions.cpu().numpy()
|
| 324 |
if atom37_positions.ndim == 4:
|
| 325 |
if atom37_positions.shape[0] != 1:
|
| 326 |
raise ValueError(
|
| 327 |
"Cannot handle batched inputs, atom37_positions has shape "
|
| 328 |
f"{atom37_positions.shape}"
|
| 329 |
)
|
| 330 |
-
atom37_positions = atom37_positions[0]
|
| 331 |
|
| 332 |
if not isinstance(atom37_positions, np.ndarray):
|
| 333 |
raise TypeError("atom37_positions must be a NumPy array or Torch tensor.")
|
|
@@ -338,7 +345,7 @@ class ProteinChain:
|
|
| 338 |
)
|
| 339 |
seqlen = atom37_positions.shape[0]
|
| 340 |
|
| 341 |
-
atom_mask = np.isfinite(atom37_positions).all(-1)
|
| 342 |
|
| 343 |
if id is None:
|
| 344 |
id = ""
|
|
@@ -350,32 +357,32 @@ class ProteinChain:
|
|
| 350 |
chain_id = "A"
|
| 351 |
|
| 352 |
if residue_index is None:
|
| 353 |
-
residue_index = np.arange(1, seqlen + 1)
|
| 354 |
elif isinstance(residue_index, torch.Tensor):
|
| 355 |
-
residue_index = residue_index.cpu().numpy()
|
| 356 |
if residue_index.ndim == 2:
|
| 357 |
if residue_index.shape[0] != 1:
|
| 358 |
raise ValueError(
|
| 359 |
"Cannot handle batched inputs, residue_index has shape "
|
| 360 |
f"{residue_index.shape}"
|
| 361 |
)
|
| 362 |
-
residue_index = residue_index[0]
|
| 363 |
if not isinstance(residue_index, np.ndarray):
|
| 364 |
raise TypeError("residue_index must be a NumPy array or Torch tensor.")
|
| 365 |
|
| 366 |
if insertion_code is None:
|
| 367 |
-
insertion_code = np.array(["" for _ in range(seqlen)])
|
| 368 |
|
| 369 |
if confidence is None:
|
| 370 |
-
confidence = np.ones(seqlen, dtype=np.float32)
|
| 371 |
elif isinstance(confidence, torch.Tensor):
|
| 372 |
-
confidence = confidence.cpu().numpy()
|
| 373 |
if confidence.ndim == 2:
|
| 374 |
if confidence.shape[0] != 1:
|
| 375 |
raise ValueError(
|
| 376 |
f"Cannot handle batched inputs, confidence has shape {confidence.shape}"
|
| 377 |
)
|
| 378 |
-
confidence = confidence[0]
|
| 379 |
if not isinstance(confidence, np.ndarray):
|
| 380 |
raise TypeError("confidence must be a NumPy array or Torch tensor.")
|
| 381 |
|
|
@@ -404,15 +411,16 @@ class ProteinChain:
|
|
| 404 |
|
| 405 |
This function passes all kwargs to from_atom37.
|
| 406 |
"""
|
|
|
|
| 407 |
if isinstance(backbone_atom_coordinates, torch.Tensor):
|
| 408 |
-
backbone_atom_coordinates = backbone_atom_coordinates.cpu().numpy()
|
| 409 |
if backbone_atom_coordinates.ndim == 4:
|
| 410 |
if backbone_atom_coordinates.shape[0] != 1:
|
| 411 |
raise ValueError(
|
| 412 |
f"Cannot handle batched inputs, backbone_atom_coordinates has "
|
| 413 |
f"shape {backbone_atom_coordinates.shape}"
|
| 414 |
)
|
| 415 |
-
backbone_atom_coordinates = backbone_atom_coordinates[0]
|
| 416 |
|
| 417 |
if not isinstance(backbone_atom_coordinates, np.ndarray):
|
| 418 |
raise TypeError(
|
|
@@ -431,8 +439,8 @@ class ProteinChain:
|
|
| 431 |
(backbone_atom_coordinates.shape[0], 37, 3),
|
| 432 |
np.inf,
|
| 433 |
dtype=backbone_atom_coordinates.dtype,
|
| 434 |
-
)
|
| 435 |
-
atom37_positions[:, :3, :] = backbone_atom_coordinates
|
| 436 |
|
| 437 |
return cls.from_atom37(atom37_positions=atom37_positions, **kwargs)
|
| 438 |
|
|
@@ -464,7 +472,7 @@ class ProteinChain:
|
|
| 464 |
case _:
|
| 465 |
file_id = "null"
|
| 466 |
|
| 467 |
-
atom_array = PDBFile.read(path).get_structure(model=1, extra_fields=["b_factor"])
|
| 468 |
if len(atom_array) == 0:
|
| 469 |
raise ValueError("PDB contains no atoms.")
|
| 470 |
if chain_id == "detect":
|
|
@@ -473,7 +481,7 @@ class ProteinChain:
|
|
| 473 |
bs.filter_amino_acids(atom_array)
|
| 474 |
& ~atom_array.hetero
|
| 475 |
& (atom_array.chain_id == chain_id)
|
| 476 |
-
]
|
| 477 |
if len(atom_array) == 0:
|
| 478 |
raise ValueError(f"PDB contains no amino-acid atoms for chain {chain_id!r}.")
|
| 479 |
|
|
@@ -487,17 +495,17 @@ class ProteinChain:
|
|
| 487 |
|
| 488 |
atom_positions = np.full(
|
| 489 |
[num_res, residue_constants.atom_type_num, 3], np.nan, dtype=np.float32
|
| 490 |
-
)
|
| 491 |
-
atom_mask = np.full([num_res, residue_constants.atom_type_num], False, dtype=bool)
|
| 492 |
-
residue_index = np.full([num_res], -1, dtype=np.int64)
|
| 493 |
-
insertion_code = np.full([num_res], "", dtype="<U4")
|
| 494 |
|
| 495 |
-
confidence = np.ones([num_res], dtype=np.float32)
|
| 496 |
|
| 497 |
for i, res in enumerate(bs.residue_iter(atom_array)):
|
| 498 |
res_index = res[0].res_id
|
| 499 |
-
residue_index[i] = res_index
|
| 500 |
-
insertion_code[i] = res[0].ins_code
|
| 501 |
|
| 502 |
# Atom level features
|
| 503 |
for atom in res:
|
|
@@ -507,10 +515,10 @@ class ProteinChain:
|
|
| 507 |
atom_name = "SD"
|
| 508 |
|
| 509 |
if atom_name in residue_constants.atom_order:
|
| 510 |
-
atom_positions[i, residue_constants.atom_order[atom_name]] = atom.coord
|
| 511 |
-
atom_mask[i, residue_constants.atom_order[atom_name]] = True
|
| 512 |
if is_predicted and atom_name == "CA":
|
| 513 |
-
confidence[i] = atom.b_factor / PLDDT_B_FACTOR_SCALE
|
| 514 |
|
| 515 |
if not sequence or not all(sequence):
|
| 516 |
raise ValueError("Some residue name was not specified correctly.")
|
|
@@ -567,7 +575,7 @@ class ProteinChain:
|
|
| 567 |
) -> ProteinChain:
|
| 568 |
"""A simple converter from bs.AtomArray -> ProteinChain.
|
| 569 |
Uses PDB file format as intermediate."""
|
| 570 |
-
atom_array = atom_array.copy()
|
| 571 |
atom_array.box = None # remove surrounding box, from_pdb won't handle this
|
| 572 |
pdb_file = PDBFile() # pyright: ignore
|
| 573 |
pdb_file.set_structure(atom_array)
|
|
@@ -596,7 +604,7 @@ class ProteinChain:
|
|
| 596 |
"residue_index": self.residue_index,
|
| 597 |
"insertion_code": self.insertion_code,
|
| 598 |
"confidence": self.confidence,
|
| 599 |
-
}
|
| 600 |
for name, values in aligned.items():
|
| 601 |
if not isinstance(values, np.ndarray):
|
| 602 |
raise TypeError(f"{name} must be a NumPy array, got {type(values).__name__}.")
|
|
@@ -636,7 +644,7 @@ class ProteinChain:
|
|
| 636 |
)
|
| 637 |
if not np.issubdtype(self.confidence.dtype, np.number):
|
| 638 |
raise TypeError("confidence must use a numeric dtype.")
|
| 639 |
-
atom37_confidence = self.atom37_confidence
|
| 640 |
if atom37_confidence is not None and not isinstance(atom37_confidence, np.ndarray):
|
| 641 |
raise TypeError("atom37_confidence must be a NumPy array when provided.")
|
| 642 |
if (
|
|
@@ -686,7 +694,7 @@ class ProteinChain:
|
|
| 686 |
self.atom37_confidence[res_idx_i, i]
|
| 687 |
if self.atom37_confidence is not None
|
| 688 |
else conf
|
| 689 |
-
)
|
| 690 |
atom = bs.Atom(
|
| 691 |
coord=pos,
|
| 692 |
chain_id="A" if self.chain_id is None else self.chain_id,
|
|
@@ -700,7 +708,7 @@ class ProteinChain:
|
|
| 700 |
occupancy=1.0,
|
| 701 |
)
|
| 702 |
atoms.append(atom)
|
| 703 |
-
return bs.array(atoms)
|
| 704 |
|
| 705 |
# Coordinate transformations and dataset adapters
|
| 706 |
def get_normalization_frame(self) -> Affine3D:
|
|
@@ -711,8 +719,8 @@ class ProteinChain:
|
|
| 711 |
Returns:
|
| 712 |
Affine3D: [] tensor of Affine3D frame
|
| 713 |
"""
|
| 714 |
-
coords = torch.from_numpy(self.atom37_positions)
|
| 715 |
-
frame = get_protein_normalization_frame(coords)
|
| 716 |
|
| 717 |
return frame
|
| 718 |
|
|
@@ -725,9 +733,10 @@ class ProteinChain:
|
|
| 725 |
Returns:
|
| 726 |
ProteinChain: Transformed protein chain
|
| 727 |
"""
|
| 728 |
-
|
| 729 |
-
coords =
|
| 730 |
-
|
|
|
|
| 731 |
return replace(self, atom37_positions=atom37_positions)
|
| 732 |
|
| 733 |
def normalize_coordinates(self) -> ProteinChain:
|
|
@@ -736,36 +745,36 @@ class ProteinChain:
|
|
| 736 |
|
| 737 |
def infer_oxygen(self) -> ProteinChain:
|
| 738 |
"""Oxygen position is fixed given N, CA, C atoms. Infer it if not provided."""
|
| 739 |
-
O_missing_indices = np.argwhere(~np.isfinite(self.atoms["O"]).all(axis=1)).squeeze()
|
| 740 |
|
| 741 |
-
O_vector = torch.tensor([0.6240, -1.0613, 0.0103], dtype=torch.float32)
|
| 742 |
-
N, CA, C = torch.from_numpy(self.atoms[["N", "CA", "C"]]).float().unbind(dim=1)
|
| 743 |
-
N = torch.roll(N, -3)
|
| 744 |
-
N[..., -1, :] = torch.nan
|
| 745 |
|
| 746 |
# Get the frame defined by the CA-C-N atom
|
| 747 |
-
frames = Affine3D.from_graham_schmidt(CA, C, N)
|
| 748 |
-
oxygen_coordinates = frames.apply(O_vector)
|
| 749 |
-
atom37_positions = self.atom37_positions.copy()
|
| 750 |
-
atom37_mask = self.atom37_mask.copy()
|
| 751 |
|
| 752 |
atom37_positions[O_missing_indices, residue_constants.atom_order["O"]] = oxygen_coordinates[
|
| 753 |
O_missing_indices
|
| 754 |
-
].numpy()
|
| 755 |
atom37_mask[O_missing_indices, residue_constants.atom_order["O"]] = ~np.isnan(
|
| 756 |
atom37_positions[O_missing_indices, residue_constants.atom_order["O"]]
|
| 757 |
-
).any(-1)
|
| 758 |
new_chain = replace(self, atom37_positions=atom37_positions, atom37_mask=atom37_mask)
|
| 759 |
return new_chain
|
| 760 |
|
| 761 |
@cached_property
|
| 762 |
def inferred_cbeta(self) -> np.ndarray:
|
| 763 |
"""Infer cbeta positions based on N, C, CA."""
|
| 764 |
-
N, CA, C = np.moveaxis(self.atoms[["N", "CA", "C"]], 1, 0)
|
| 765 |
# See usage in trDesign codebase.
|
| 766 |
# https://github.com/gjoni/trDesign/blob/f2d5930b472e77bfacc2f437b3966e7a708a8d37/02-GD/utils.py#L140
|
| 767 |
-
CB = infer_cb(C, N, CA, 1.522, 1.927, -2.143)
|
| 768 |
-
return CB
|
| 769 |
|
| 770 |
def infer_cbeta(self, infer_cbeta_for_glycine: bool = False) -> ProteinChain:
|
| 771 |
"""Return a new chain with inferred CB atoms at all residues except GLY.
|
|
@@ -780,30 +789,30 @@ class ProteinChain:
|
|
| 780 |
calculation between two designs for a given structural template, w/
|
| 781 |
CB atoms.
|
| 782 |
"""
|
| 783 |
-
atom37_positions = self.atom37_positions.copy()
|
| 784 |
-
atom37_mask = self.atom37_mask.copy()
|
| 785 |
|
| 786 |
-
inferred_cbeta_positions = self.inferred_cbeta
|
| 787 |
if not infer_cbeta_for_glycine:
|
| 788 |
-
inferred_cbeta_positions[np.array(list(self.sequence)) == "G", :] = np.nan
|
| 789 |
|
| 790 |
-
atom37_positions[:, residue_constants.atom_order["CB"]] = inferred_cbeta_positions
|
| 791 |
atom37_mask[:, residue_constants.atom_order["CB"]] = ~np.isnan(
|
| 792 |
atom37_positions[:, residue_constants.atom_order["CB"]]
|
| 793 |
-
).any(-1)
|
| 794 |
new_chain = replace(self, atom37_positions=atom37_positions, atom37_mask=atom37_mask)
|
| 795 |
return new_chain
|
| 796 |
|
| 797 |
@cached_property
|
| 798 |
def pdist_CA(self) -> np.ndarray:
|
| 799 |
-
CA = self.atoms["CA"]
|
| 800 |
-
pdist_CA = squareform(pdist(CA))
|
| 801 |
-
return pdist_CA
|
| 802 |
|
| 803 |
@cached_property
|
| 804 |
def pdist_CB(self) -> np.ndarray:
|
| 805 |
-
pdist_CB = squareform(pdist(self.inferred_cbeta))
|
| 806 |
-
return pdist_CB
|
| 807 |
|
| 808 |
@classmethod
|
| 809 |
def as_complex(cls, chains: Sequence[ProteinChain]):
|
|
@@ -824,23 +833,24 @@ class ProteinChain:
|
|
| 824 |
"atom37_positions": np.full([1, 37, 3], np.inf),
|
| 825 |
"atom37_mask": np.zeros([1, 37], dtype=bool),
|
| 826 |
"confidence": np.array([0]),
|
| 827 |
-
}
|
| 828 |
|
| 829 |
def join_arrays(arrays: Sequence[np.ndarray], sep: np.ndarray):
|
|
|
|
| 830 |
if use_chainbreak:
|
| 831 |
full_array = []
|
| 832 |
for array in arrays:
|
| 833 |
full_array.append(array)
|
| 834 |
full_array.append(sep)
|
| 835 |
full_array = full_array[:-1]
|
| 836 |
-
return np.concatenate(full_array, 0)
|
| 837 |
else:
|
| 838 |
-
return np.concatenate(arrays, 0)
|
| 839 |
|
| 840 |
array_args: dict[str, np.ndarray] = {
|
| 841 |
name: join_arrays([getattr(chain, name) for chain in chains], sep)
|
| 842 |
for name, sep in sep_tokens.items()
|
| 843 |
-
}
|
| 844 |
|
| 845 |
chain_break = residue_constants.CHAIN_BREAK_TOKEN if use_chainbreak else ""
|
| 846 |
return cls(
|
|
@@ -868,15 +878,15 @@ class ProteinChain:
|
|
| 868 |
raise ValueError(
|
| 869 |
f"Non-polymer {nonpolymer.comp_id!r} has no coordinate table."
|
| 870 |
)
|
| 871 |
-
chain_coords = self.atom37_positions[self.atom37_mask]
|
| 872 |
-
distance = cdist(nonpolymer_array.coord, chain_coords)
|
| 873 |
|
| 874 |
-
is_contact = distance < 5
|
| 875 |
if not is_contact.any():
|
| 876 |
continue
|
| 877 |
-
contacting_atoms = np.where(is_contact.any(0))[0]
|
| 878 |
-
chain_index = np.where(self.atom37_mask)[0]
|
| 879 |
-
contacting_residues = np.unique(chain_index[contacting_atoms])
|
| 880 |
|
| 881 |
result = {
|
| 882 |
"ligand": nonpolymer.name,
|
|
@@ -890,7 +900,7 @@ class ProteinChain:
|
|
| 890 |
self, indices: list[int | str], ignore_x_mismatch: bool = False
|
| 891 |
) -> ProteinChain:
|
| 892 |
numeric_indices = [idx if isinstance(idx, int) else int(idx[1:]) for idx in indices]
|
| 893 |
-
mask = np.isin(self.residue_index, numeric_indices)
|
| 894 |
new = self[mask]
|
| 895 |
mismatches = []
|
| 896 |
for aa, idx in zip(new.sequence, indices, strict=False):
|
|
@@ -922,20 +932,22 @@ class ProteinChain:
|
|
| 922 |
# Convert to tensors and add batch dimension
|
| 923 |
coordinates = (
|
| 924 |
torch.from_numpy(self.atom37_positions).float().unsqueeze(0)
|
| 925 |
-
) # X has shape (1, l, 37, 3).
|
| 926 |
-
plddt = torch.from_numpy(self.confidence).float().unsqueeze(0) # P: (1, l)
|
| 927 |
residue_index = (
|
| 928 |
torch.from_numpy(self.residue_index).long().unsqueeze(0)
|
| 929 |
-
) # R has shape (1, l).
|
| 930 |
|
| 931 |
-
return coordinates, plddt, residue_index
|
| 932 |
|
| 933 |
# Sequence access, interchange, and compact storage
|
| 934 |
def __getitem__(self, idx: int | list[int] | slice | np.ndarray | torch.Tensor):
|
|
|
|
|
|
|
| 935 |
if isinstance(idx, int):
|
| 936 |
idx = [idx]
|
| 937 |
if isinstance(idx, torch.Tensor):
|
| 938 |
-
idx = idx.cpu().numpy()
|
| 939 |
|
| 940 |
sequence = slice_python_object_as_numpy(self.sequence, idx)
|
| 941 |
return replace(
|
|
@@ -955,11 +967,11 @@ class ProteinChain:
|
|
| 955 |
return len(self.sequence)
|
| 956 |
|
| 957 |
def cbeta_contacts(self, distance_threshold: float = 8.0) -> np.ndarray:
|
| 958 |
-
distance = self.pdist_CB
|
| 959 |
-
contacts = (distance < distance_threshold).astype(np.int64)
|
| 960 |
-
contacts[np.isnan(distance)] = -1
|
| 961 |
np.fill_diagonal(contacts, -1)
|
| 962 |
-
return contacts
|
| 963 |
|
| 964 |
def to_pdb(self, path: PathOrBuffer, include_insertions: bool = True):
|
| 965 |
"""Dssp works better w/o insertions."""
|
|
@@ -989,7 +1001,7 @@ class ProteinChain:
|
|
| 989 |
"mode": CIFColumn(data=CIFData(array=np.array(["global", "local"]), dtype=np.str_)),
|
| 990 |
"name": CIFColumn(data=CIFData(array=np.array(["pLDDT", "pLDDT"]), dtype=np.str_)),
|
| 991 |
},
|
| 992 |
-
)
|
| 993 |
|
| 994 |
# table is a duplicate of data already in the atom array, but
|
| 995 |
# needed by molstar to render pLDDT / confidence
|
|
@@ -1039,10 +1051,10 @@ class ProteinChain:
|
|
| 1039 |
need more than 2**32 residues..."""
|
| 1040 |
dct = {k: v for k, v in asdict(self).items() if k not in ["mmcif"]}
|
| 1041 |
if backbone_only:
|
| 1042 |
-
dct["atom37_mask"][:, 3:] = False
|
| 1043 |
-
dct["atom37_positions"] = dct["atom37_positions"][dct["atom37_mask"]]
|
| 1044 |
if dct.get("atom37_confidence") is not None:
|
| 1045 |
-
dct["atom37_confidence"] = dct["atom37_confidence"][dct["atom37_mask"]]
|
| 1046 |
else:
|
| 1047 |
dct.pop("atom37_confidence", None)
|
| 1048 |
|
|
@@ -1050,9 +1062,9 @@ class ProteinChain:
|
|
| 1050 |
if isinstance(v, np.ndarray):
|
| 1051 |
match v.dtype:
|
| 1052 |
case np.int64:
|
| 1053 |
-
dct[k] = v.astype(np.int32)
|
| 1054 |
case np.float64 | np.float32:
|
| 1055 |
-
dct[k] = v.astype(np.float16)
|
| 1056 |
case _:
|
| 1057 |
pass
|
| 1058 |
if json_serializable:
|
|
@@ -1074,15 +1086,15 @@ class ProteinChain:
|
|
| 1074 |
|
| 1075 |
for k, v in dct.items():
|
| 1076 |
if isinstance(v, list):
|
| 1077 |
-
dct[k] = np.array(v)
|
| 1078 |
|
| 1079 |
-
atom37 = np.full((*dct["atom37_mask"].shape, 3), np.nan)
|
| 1080 |
-
atom37[dct["atom37_mask"]] = dct["atom37_positions"]
|
| 1081 |
-
dct["atom37_positions"] = atom37
|
| 1082 |
if "atom37_confidence" in dct:
|
| 1083 |
-
atom37_conf = np.full(dct["atom37_mask"].shape, np.nan, dtype=np.float32)
|
| 1084 |
-
atom37_conf[dct["atom37_mask"]] = dct["atom37_confidence"]
|
| 1085 |
-
dct["atom37_confidence"] = atom37_conf
|
| 1086 |
dct = {
|
| 1087 |
k: (
|
| 1088 |
v.astype(np.float32)
|
|
@@ -1091,7 +1103,7 @@ class ProteinChain:
|
|
| 1091 |
)
|
| 1092 |
for k, v in dct.items()
|
| 1093 |
if not (k == "atom37_confidence" and v is None)
|
| 1094 |
-
}
|
| 1095 |
return cls(**dct, mmcif=None)
|
| 1096 |
|
| 1097 |
@classmethod
|
|
@@ -1111,10 +1123,10 @@ class ProteinChain:
|
|
| 1111 |
|
| 1112 |
# Surface and structural comparison metrics
|
| 1113 |
def sasa(self, by_residue: bool = True):
|
| 1114 |
-
arr = self.atom_array_no_insertions
|
| 1115 |
if len(arr) == 0:
|
| 1116 |
raise ValueError("SASA requires at least one resolved atom.")
|
| 1117 |
-
sasa_per_atom = bs.sasa(arr) # type: ignore
|
| 1118 |
if by_residue:
|
| 1119 |
# Sum per-atom SASA into residue "bins", with np.bincount.
|
| 1120 |
if arr.res_id is None:
|
|
@@ -1127,12 +1139,12 @@ class ProteinChain:
|
|
| 1127 |
np.bincount(arr.res_id, weights=sasa_per_atom)[1:],
|
| 1128 |
np.zeros(num_trailing_residues),
|
| 1129 |
]
|
| 1130 |
-
)
|
| 1131 |
-
sasa_per_residue[~self.atom37_mask.any(-1)] = np.nan
|
| 1132 |
if len(sasa_per_residue) != len(self):
|
| 1133 |
raise RuntimeError("Residue SASA output does not align with the protein chain.")
|
| 1134 |
-
return sasa_per_residue
|
| 1135 |
-
return sasa_per_atom
|
| 1136 |
|
| 1137 |
def sap_score(self, aggregation: str = "atom") -> np.ndarray:
|
| 1138 |
"""Compute per-atom spatial aggregation propensity (SAP).
|
|
@@ -1141,7 +1153,7 @@ class ProteinChain:
|
|
| 1141 |
Protein aggregation sums positive atom scores, following Lauer et al. 2011.
|
| 1142 |
"""
|
| 1143 |
sap_radius = 5.0
|
| 1144 |
-
arr = self.atom_array_no_insertions
|
| 1145 |
if len(arr) == 0:
|
| 1146 |
raise ValueError("SAP requires at least one resolved atom.")
|
| 1147 |
|
|
@@ -1150,42 +1162,42 @@ class ProteinChain:
|
|
| 1150 |
raise RuntimeError(f"Biotite AtomArray is missing required {name!r} data.")
|
| 1151 |
|
| 1152 |
# compute SASA and residue-specific properties
|
| 1153 |
-
sasa_per_atom = self.sasa(by_residue=False)
|
| 1154 |
resid_to_resname = dict(zip(arr.res_id, arr.res_name, strict=False))
|
| 1155 |
|
| 1156 |
-
max_side_chain_asa = np.full(len(self), np.nan)
|
| 1157 |
-
res_hydrophobicity = np.full(len(self), np.nan)
|
| 1158 |
-
resolved_res_mask = self.atom37_mask.any(-1)
|
| 1159 |
num_trailing_residues = len(self) - arr.res_id.max()
|
| 1160 |
|
| 1161 |
max_side_chain_asa[resolved_res_mask] = np.array(
|
| 1162 |
[residue_constants.side_chain_asa[resid_to_resname[i]] for i in np.unique(arr.res_id)]
|
| 1163 |
-
)
|
| 1164 |
res_hydrophobicity[resolved_res_mask] = np.array(
|
| 1165 |
[residue_constants.hydrophobicity[resid_to_resname[i]] for i in np.unique(arr.res_id)]
|
| 1166 |
-
)
|
| 1167 |
|
| 1168 |
# compute SAP score
|
| 1169 |
-
is_side_chain = ~bs.filter_peptide_backbone(arr)
|
| 1170 |
-
sasa_per_atom[is_side_chain] = 0
|
| 1171 |
kdtree = KDTree(arr.coord)
|
| 1172 |
neighbors = kdtree.query_ball_tree(kdtree, sap_radius, p=2.0)
|
| 1173 |
-
sap_by_atom = np.zeros_like(sasa_per_atom)
|
| 1174 |
for i, nn_list in enumerate(neighbors):
|
| 1175 |
-
saa_nn = np.zeros_like(sasa_per_atom)
|
| 1176 |
-
saa_nn[nn_list] = sasa_per_atom[nn_list]
|
| 1177 |
sasa_within_r = np.concatenate(
|
| 1178 |
[
|
| 1179 |
np.bincount(arr.res_id, weights=saa_nn)[1:],
|
| 1180 |
np.zeros(num_trailing_residues),
|
| 1181 |
]
|
| 1182 |
-
)
|
| 1183 |
-
sap = np.nansum((sasa_within_r / max_side_chain_asa) * res_hydrophobicity)
|
| 1184 |
-
sap_by_atom[i] = sap
|
| 1185 |
|
| 1186 |
match aggregation:
|
| 1187 |
case "atom":
|
| 1188 |
-
return sap_by_atom
|
| 1189 |
case "residue":
|
| 1190 |
sap_by_residue = np.concatenate(
|
| 1191 |
[
|
|
@@ -1195,11 +1207,11 @@ class ProteinChain:
|
|
| 1195 |
) / (
|
| 1196 |
np.concatenate([np.bincount(arr.res_id)[1:], np.zeros(num_trailing_residues)])
|
| 1197 |
+ 1e-8
|
| 1198 |
-
)
|
| 1199 |
-
sap_by_residue[~resolved_res_mask] = np.nan
|
| 1200 |
if len(sap_by_residue) != len(self):
|
| 1201 |
raise RuntimeError("Residue SAP output does not align with the protein chain.")
|
| 1202 |
-
return sap_by_residue
|
| 1203 |
case "protein":
|
| 1204 |
return sum(sap_by_atom[sap_by_atom > 0]) # pyright: ignore[reportReturnType]
|
| 1205 |
case _:
|
|
@@ -1216,19 +1228,19 @@ class ProteinChain:
|
|
| 1216 |
|
| 1217 |
# https://www.mdpi.com/2073-4352/11/12/1539
|
| 1218 |
# The non-overlapping-atom approximation can produce globularity above one.
|
| 1219 |
-
mask = self.atom37_mask.any(-1)
|
| 1220 |
-
points = self.atom37_positions[self.atom37_mask]
|
| 1221 |
sequence = [aa for aa, m in zip(self.sequence, mask, strict=False) if m] # type: ignore
|
| 1222 |
-
A, _ = self._mvee(points, tol=1e-3)
|
| 1223 |
-
mvee_volume = (4 * np.pi) / (3 * np.sqrt(np.linalg.det(A)))
|
| 1224 |
volume = sum(residue_constants.amino_acid_volumes[x] for x in sequence)
|
| 1225 |
ratio = volume / mvee_volume
|
| 1226 |
|
| 1227 |
# The paper compares the ellipsoidal profile with scalar t, a measurement
|
| 1228 |
# of elongation. We want a single number, so we multiply by 1/(2t), so
|
| 1229 |
# that value is normalized between 0-1
|
| 1230 |
-
eigenvalues = np.linalg.eigvals(A)
|
| 1231 |
-
R = 1 / np.sqrt(eigenvalues)
|
| 1232 |
# ellipsoid radii length triangle inequality coefficient
|
| 1233 |
t = max(R[0] / (R[1] + R[2]), R[1] / (R[0] + R[2]), R[2] / (R[0] + R[1]))
|
| 1234 |
elongation_metric = 1 / max(t, 1)
|
|
@@ -1239,50 +1251,51 @@ class ProteinChain:
|
|
| 1239 |
# Finds minimum volume enclosing ellipsoid of a set of points.
|
| 1240 |
# Returns A, c where the ellipse is defined as:
|
| 1241 |
# (x-c).T @ A @ (x-c) = 1
|
|
|
|
| 1242 |
hull = ConvexHull(P)
|
| 1243 |
-
P = P[hull.vertices]
|
| 1244 |
-
P = P.T
|
| 1245 |
|
| 1246 |
# Data points
|
| 1247 |
d, n = P.shape
|
| 1248 |
-
Q = np.zeros((d + 1, n))
|
| 1249 |
-
Q[:d, :] = P[:d, :n]
|
| 1250 |
-
Q[d, :] = np.ones((1, n))
|
| 1251 |
|
| 1252 |
# Initializations
|
| 1253 |
count = 1
|
| 1254 |
err = 1.0
|
| 1255 |
-
u = np.full((n, 1), 1 / n) # First iteration.
|
| 1256 |
|
| 1257 |
# Khachiyan Algorithm
|
| 1258 |
for _ in range(max_iter):
|
| 1259 |
-
X = Q.dot(np.diag(u.squeeze())) @ Q.T
|
| 1260 |
-
M = np.diag(Q.T @ np.linalg.inv(X) @ Q)
|
| 1261 |
-
maximum, j = np.max(M), np.argmax(M)
|
| 1262 |
step_size = (maximum - d - 1) / ((d + 1) * (maximum - 1))
|
| 1263 |
-
new_u = (1 - step_size) * u
|
| 1264 |
-
new_u[j] += step_size
|
| 1265 |
count += 1
|
| 1266 |
-
err = np.linalg.norm(new_u - u)
|
| 1267 |
-
u = new_u
|
| 1268 |
if err < tol:
|
| 1269 |
break
|
| 1270 |
else:
|
| 1271 |
raise ValueError("MVEE did not converge")
|
| 1272 |
|
| 1273 |
d = P.shape[0] # Fixed: use P.shape[0] instead of P.shape
|
| 1274 |
-
U = np.diag(u.squeeze())
|
| 1275 |
|
| 1276 |
# The A matrix for the ellipse
|
| 1277 |
-
A = (1 / d) * np.linalg.inv(P @ U @ P.T - (P @ u) @ (P @ u).T)
|
| 1278 |
|
| 1279 |
# Center of the ellipse
|
| 1280 |
-
c = P @ u
|
| 1281 |
|
| 1282 |
-
return A, c
|
| 1283 |
|
| 1284 |
def radius_of_gyration(self):
|
| 1285 |
-
arr = self.atom_array_no_insertions
|
| 1286 |
return bs.gyration_radius(arr)
|
| 1287 |
|
| 1288 |
def align(
|
|
@@ -1375,8 +1388,8 @@ class ProteinChain:
|
|
| 1375 |
torch.tensor(native.atom37_positions[target_inds]).unsqueeze(0),
|
| 1376 |
torch.tensor(native.atom37_mask[mobile_inds]).unsqueeze(0),
|
| 1377 |
**kwargs,
|
| 1378 |
-
)
|
| 1379 |
-
return float(lddt) if lddt.numel() == 1 else lddt.numpy().flatten()
|
| 1380 |
|
| 1381 |
def gdt_ts(
|
| 1382 |
self,
|
|
@@ -1409,12 +1422,12 @@ class ProteinChain:
|
|
| 1409 |
& index_by_atom_name(target.atom37_mask[target_inds], "CA", dim=-1)
|
| 1410 |
).unsqueeze(0),
|
| 1411 |
**kwargs,
|
| 1412 |
-
)
|
| 1413 |
-
return float(gdt_ts) if gdt_ts.numel() == 1 else gdt_ts.numpy().flatten()
|
| 1414 |
|
| 1415 |
@cached_property
|
| 1416 |
def residue_index_no_insertions(self) -> np.ndarray:
|
| 1417 |
-
return self.residue_index + np.cumsum(self.insertion_code != "")
|
| 1418 |
|
| 1419 |
@cached_property
|
| 1420 |
def atom_array_no_insertions(self) -> bs.AtomArray:
|
|
@@ -1433,7 +1446,7 @@ class ProteinChain:
|
|
| 1433 |
self.atom37_confidence[res_idx, i]
|
| 1434 |
if self.atom37_confidence is not None
|
| 1435 |
else conf
|
| 1436 |
-
)
|
| 1437 |
atom = bs.Atom(
|
| 1438 |
coord=pos,
|
| 1439 |
# hard coded to as we currently only support single chain structures
|
|
@@ -1447,4 +1460,4 @@ class ProteinChain:
|
|
| 1447 |
occupancy=1.0,
|
| 1448 |
)
|
| 1449 |
atoms.append(atom)
|
| 1450 |
-
return bs.array(atoms)
|
|
|
|
| 4 |
|
| 5 |
import io
|
| 6 |
import warnings
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 7 |
import biotite.structure as bs
|
| 8 |
import brotli
|
| 9 |
import msgpack
|
| 10 |
import msgpack_numpy
|
| 11 |
import numpy as np
|
| 12 |
import torch
|
| 13 |
+
|
| 14 |
+
from collections.abc import Mapping, Sequence
|
| 15 |
+
from dataclasses import asdict, dataclass, replace
|
| 16 |
+
from functools import cached_property
|
| 17 |
+
from pathlib import Path
|
| 18 |
+
from typing import Any
|
| 19 |
from biotite.database import rcsb
|
| 20 |
from biotite.structure.io.pdb import PDBFile
|
| 21 |
from biotite.structure.io.pdbx import CIFCategory, CIFColumn, CIFData, CIFFile
|
|
|
|
| 42 |
from .esmfold2_protein_structure import index_by_atom_name
|
| 43 |
from .esmfold2_utils_types import PathOrBuffer
|
| 44 |
|
| 45 |
+
|
| 46 |
CHAIN_ID_CONST = "A"
|
| 47 |
|
| 48 |
|
|
|
|
| 71 |
dihedral: float = -2.143,
|
| 72 |
):
|
| 73 |
"""Infer C-beta coordinates from C, N, and C-alpha coordinates."""
|
| 74 |
+
# C, N, Ca: (..., 3); xyz coordinates with broadcast-compatible leading axes.
|
| 75 |
|
| 76 |
def normalize(X: np.ndarray) -> np.ndarray:
|
| 77 |
+
# X: (..., 3); normalize each xyz vector independently.
|
| 78 |
+
return X / np.sqrt(np.square(X).sum(-1, keepdims=True) + 1e-8) # X.shape; normalized along xyz
|
| 79 |
|
| 80 |
with np.errstate(invalid="ignore"):
|
| 81 |
+
n_to_ca = N - Ca # (..., 3)
|
| 82 |
+
n_to_c = N - C # (..., 3)
|
| 83 |
+
axis = normalize(n_to_ca) # (..., 3)
|
| 84 |
+
normal = normalize(np.cross(n_to_c, axis)) # (..., 3)
|
| 85 |
+
basis = (axis, np.cross(normal, axis), normal) # three arrays (..., 3)
|
| 86 |
offsets = (
|
| 87 |
bond_length * np.cos(bond_angle),
|
| 88 |
bond_length * np.sin(bond_angle) * np.cos(dihedral),
|
| 89 |
-bond_length * np.sin(bond_angle) * np.sin(dihedral),
|
| 90 |
+
) # three NumPy scalars, each ()
|
| 91 |
+
return Ca + sum(vector * offset for vector, offset in zip(basis, offsets, strict=False)) # (..., 3)
|
| 92 |
|
| 93 |
|
| 94 |
def chain_to_ndarray(
|
| 95 |
atom_array: bs.AtomArray, mmcif: MmcifWrapper, chain_id: str, is_predicted=False
|
| 96 |
):
|
| 97 |
+
# atom_array: n input atoms; l is the selected chain's sequence length.
|
| 98 |
if not isinstance(atom_array, bs.AtomArray):
|
| 99 |
raise TypeError("atom_array must be a biotite AtomArray.")
|
| 100 |
if not isinstance(mmcif, MmcifWrapper):
|
|
|
|
| 110 |
num_res = len(mmcif.chain_to_seqres[chain_id])
|
| 111 |
sequence = mmcif.chain_to_seqres[chain_id]
|
| 112 |
|
| 113 |
+
atom_positions = np.full([num_res, residue_constants.atom_type_num, 3], np.nan) # (l, 37, 3)
|
| 114 |
+
atom_mask = np.full([num_res, residue_constants.atom_type_num], False, dtype=bool) # (l, 37)
|
| 115 |
+
residue_index = np.full([num_res], -1, dtype=np.int64) # (l,)
|
| 116 |
+
insertion_code = np.full([num_res], "", dtype="<U4") # (l,)
|
| 117 |
|
| 118 |
+
confidence = np.ones([num_res], dtype=np.float32) # (l,)
|
| 119 |
|
| 120 |
+
chain = atom_array[atom_array.chain_id == chain_id] # AtomArray with n_chain_atoms entries
|
| 121 |
if not isinstance(chain, bs.AtomArray):
|
| 122 |
raise RuntimeError("Biotite selection did not return an AtomArray.")
|
| 123 |
for res_index in range(num_res):
|
|
|
|
| 126 |
if res_at_position.residue_number is None:
|
| 127 |
continue
|
| 128 |
|
| 129 |
+
residue_index[res_index] = res_at_position.residue_number # scalar element in (l,)
|
| 130 |
+
insertion_code[res_index] = res_at_position.insertion_code # scalar element in (l,)
|
| 131 |
res = chain[
|
| 132 |
(chain.res_id == res_at_position.residue_number)
|
| 133 |
& (chain.ins_code == res_at_position.insertion_code)
|
| 134 |
& (chain.hetero == res_at_position.hetflag)
|
| 135 |
+
] # AtomArray with n_residue_atoms entries
|
| 136 |
if not isinstance(res, bs.AtomArray):
|
| 137 |
raise RuntimeError("Biotite residue selection did not return an AtomArray.")
|
| 138 |
|
|
|
|
| 144 |
atom_name = "SD"
|
| 145 |
|
| 146 |
if atom_name in residue_constants.atom_order:
|
| 147 |
+
atom_positions[res_index, residue_constants.atom_order[atom_name]] = atom.coord # (3,) xyz vector in (l, 37, 3)
|
| 148 |
+
atom_mask[res_index, residue_constants.atom_order[atom_name]] = True # scalar element in (l, 37)
|
| 149 |
if is_predicted and atom_name == "CA":
|
| 150 |
+
confidence[res_index] = atom.b_factor / PLDDT_B_FACTOR_SCALE # scalar element in (l,)
|
| 151 |
|
| 152 |
if not sequence or not all(sequence):
|
| 153 |
raise ValueError("Some residue name was not specified correctly.")
|
|
|
|
| 159 |
insertion_code,
|
| 160 |
confidence,
|
| 161 |
entity_id,
|
| 162 |
+
) # array fields: (l, 37, 3), (l, 37), (l,), (l,), (l,)
|
| 163 |
|
| 164 |
|
| 165 |
@dataclass(frozen=True)
|
| 166 |
class ProteinChain:
|
| 167 |
"""Dataclass with atom37 representation of a single protein chain."""
|
| 168 |
|
| 169 |
+
# l is len(sequence), including separator rows when present.
|
| 170 |
id: str
|
| 171 |
sequence: str
|
| 172 |
chain_id: str # author chain id - mutable
|
| 173 |
entity_id: int | None
|
| 174 |
+
residue_index: np.ndarray # (l,)
|
| 175 |
+
insertion_code: np.ndarray # (l,)
|
| 176 |
+
atom37_positions: np.ndarray # (l, 37, 3)
|
| 177 |
+
atom37_mask: np.ndarray # (l, 37)
|
| 178 |
+
confidence: np.ndarray # (l,)
|
| 179 |
mmcif: MmcifWrapper | None = None
|
| 180 |
atom37_confidence: np.ndarray | None = None # P has shape (l, 37).
|
| 181 |
|
|
|
|
| 191 |
"""Yield every protein chain represented in an mmCIF structure."""
|
| 192 |
mmcif = path if isinstance(path, MmcifWrapper) else MmcifWrapper.read(path, id)
|
| 193 |
for chain in bs.chain_iter(mmcif.structure):
|
| 194 |
+
chain = chain[bs.filter_amino_acids(chain) & ~chain.hetero] # AtomArray with n_protein_atoms entries
|
| 195 |
if len(chain) == 0:
|
| 196 |
continue
|
| 197 |
chain_id = chain.chain_id[0]
|
|
|
|
| 211 |
insertion_code,
|
| 212 |
confidence,
|
| 213 |
_,
|
| 214 |
+
) = chain_to_ndarray(chain, mmcif, chain_id, is_predicted) # array fields: (l, 37, 3), (l, 37), (l,), (l,), (l,)
|
| 215 |
if not all(sequence):
|
| 216 |
raise ValueError("Some residue name was not specified correctly.")
|
| 217 |
|
|
|
|
| 285 |
stacklevel=2,
|
| 286 |
)
|
| 287 |
|
| 288 |
+
atom_array = mmcif.structure # AtomArray with n_structure_atoms entries
|
| 289 |
(
|
| 290 |
sequence,
|
| 291 |
atom_positions,
|
|
|
|
| 294 |
insertion_code,
|
| 295 |
confidence,
|
| 296 |
_,
|
| 297 |
+
) = chain_to_ndarray(atom_array, mmcif, chain_id, is_predicted) # array fields: (l, 37, 3), (l, 37), (l,), (l,), (l,)
|
| 298 |
if not all(sequence):
|
| 299 |
raise ValueError("Some residue name was not specified correctly.")
|
| 300 |
|
|
|
|
| 324 |
insertion_code: np.ndarray | None = None,
|
| 325 |
confidence: np.ndarray | torch.Tensor | None = None,
|
| 326 |
):
|
| 327 |
+
# atom37_positions: (l, 37, 3) or (1, l, 37, 3).
|
| 328 |
+
# Optional residue indices and confidence are (l,) or (1, l).
|
| 329 |
if isinstance(atom37_positions, torch.Tensor):
|
| 330 |
+
atom37_positions = atom37_positions.cpu().numpy() # same incoming shape, (l, 37, 3) or (1, l, 37, 3)
|
| 331 |
if atom37_positions.ndim == 4:
|
| 332 |
if atom37_positions.shape[0] != 1:
|
| 333 |
raise ValueError(
|
| 334 |
"Cannot handle batched inputs, atom37_positions has shape "
|
| 335 |
f"{atom37_positions.shape}"
|
| 336 |
)
|
| 337 |
+
atom37_positions = atom37_positions[0] # (l, 37, 3)
|
| 338 |
|
| 339 |
if not isinstance(atom37_positions, np.ndarray):
|
| 340 |
raise TypeError("atom37_positions must be a NumPy array or Torch tensor.")
|
|
|
|
| 345 |
)
|
| 346 |
seqlen = atom37_positions.shape[0]
|
| 347 |
|
| 348 |
+
atom_mask = np.isfinite(atom37_positions).all(-1) # (l, 37)
|
| 349 |
|
| 350 |
if id is None:
|
| 351 |
id = ""
|
|
|
|
| 357 |
chain_id = "A"
|
| 358 |
|
| 359 |
if residue_index is None:
|
| 360 |
+
residue_index = np.arange(1, seqlen + 1) # (l,)
|
| 361 |
elif isinstance(residue_index, torch.Tensor):
|
| 362 |
+
residue_index = residue_index.cpu().numpy() # same incoming index shape
|
| 363 |
if residue_index.ndim == 2:
|
| 364 |
if residue_index.shape[0] != 1:
|
| 365 |
raise ValueError(
|
| 366 |
"Cannot handle batched inputs, residue_index has shape "
|
| 367 |
f"{residue_index.shape}"
|
| 368 |
)
|
| 369 |
+
residue_index = residue_index[0] # (l,)
|
| 370 |
if not isinstance(residue_index, np.ndarray):
|
| 371 |
raise TypeError("residue_index must be a NumPy array or Torch tensor.")
|
| 372 |
|
| 373 |
if insertion_code is None:
|
| 374 |
+
insertion_code = np.array(["" for _ in range(seqlen)]) # (l,)
|
| 375 |
|
| 376 |
if confidence is None:
|
| 377 |
+
confidence = np.ones(seqlen, dtype=np.float32) # (l,)
|
| 378 |
elif isinstance(confidence, torch.Tensor):
|
| 379 |
+
confidence = confidence.cpu().numpy() # same incoming confidence shape
|
| 380 |
if confidence.ndim == 2:
|
| 381 |
if confidence.shape[0] != 1:
|
| 382 |
raise ValueError(
|
| 383 |
f"Cannot handle batched inputs, confidence has shape {confidence.shape}"
|
| 384 |
)
|
| 385 |
+
confidence = confidence[0] # (l,)
|
| 386 |
if not isinstance(confidence, np.ndarray):
|
| 387 |
raise TypeError("confidence must be a NumPy array or Torch tensor.")
|
| 388 |
|
|
|
|
| 411 |
|
| 412 |
This function passes all kwargs to from_atom37.
|
| 413 |
"""
|
| 414 |
+
# backbone_atom_coordinates: (l, 3, 3) or (1, l, 3, 3), atom then xyz axes.
|
| 415 |
if isinstance(backbone_atom_coordinates, torch.Tensor):
|
| 416 |
+
backbone_atom_coordinates = backbone_atom_coordinates.cpu().numpy() # same incoming shape, (l, 3, 3) or (1, l, 3, 3)
|
| 417 |
if backbone_atom_coordinates.ndim == 4:
|
| 418 |
if backbone_atom_coordinates.shape[0] != 1:
|
| 419 |
raise ValueError(
|
| 420 |
f"Cannot handle batched inputs, backbone_atom_coordinates has "
|
| 421 |
f"shape {backbone_atom_coordinates.shape}"
|
| 422 |
)
|
| 423 |
+
backbone_atom_coordinates = backbone_atom_coordinates[0] # (l, 3, 3)
|
| 424 |
|
| 425 |
if not isinstance(backbone_atom_coordinates, np.ndarray):
|
| 426 |
raise TypeError(
|
|
|
|
| 439 |
(backbone_atom_coordinates.shape[0], 37, 3),
|
| 440 |
np.inf,
|
| 441 |
dtype=backbone_atom_coordinates.dtype,
|
| 442 |
+
) # (l, 37, 3)
|
| 443 |
+
atom37_positions[:, :3, :] = backbone_atom_coordinates # (l, 3, 3) backbone slice
|
| 444 |
|
| 445 |
return cls.from_atom37(atom37_positions=atom37_positions, **kwargs)
|
| 446 |
|
|
|
|
| 472 |
case _:
|
| 473 |
file_id = "null"
|
| 474 |
|
| 475 |
+
atom_array = PDBFile.read(path).get_structure(model=1, extra_fields=["b_factor"]) # AtomArray with n_file_atoms entries
|
| 476 |
if len(atom_array) == 0:
|
| 477 |
raise ValueError("PDB contains no atoms.")
|
| 478 |
if chain_id == "detect":
|
|
|
|
| 481 |
bs.filter_amino_acids(atom_array)
|
| 482 |
& ~atom_array.hetero
|
| 483 |
& (atom_array.chain_id == chain_id)
|
| 484 |
+
] # AtomArray with n_selected_protein_atoms entries
|
| 485 |
if len(atom_array) == 0:
|
| 486 |
raise ValueError(f"PDB contains no amino-acid atoms for chain {chain_id!r}.")
|
| 487 |
|
|
|
|
| 495 |
|
| 496 |
atom_positions = np.full(
|
| 497 |
[num_res, residue_constants.atom_type_num, 3], np.nan, dtype=np.float32
|
| 498 |
+
) # (l, 37, 3)
|
| 499 |
+
atom_mask = np.full([num_res, residue_constants.atom_type_num], False, dtype=bool) # (l, 37)
|
| 500 |
+
residue_index = np.full([num_res], -1, dtype=np.int64) # (l,)
|
| 501 |
+
insertion_code = np.full([num_res], "", dtype="<U4") # (l,)
|
| 502 |
|
| 503 |
+
confidence = np.ones([num_res], dtype=np.float32) # (l,)
|
| 504 |
|
| 505 |
for i, res in enumerate(bs.residue_iter(atom_array)):
|
| 506 |
res_index = res[0].res_id
|
| 507 |
+
residue_index[i] = res_index # scalar element in (l,)
|
| 508 |
+
insertion_code[i] = res[0].ins_code # scalar element in (l,)
|
| 509 |
|
| 510 |
# Atom level features
|
| 511 |
for atom in res:
|
|
|
|
| 515 |
atom_name = "SD"
|
| 516 |
|
| 517 |
if atom_name in residue_constants.atom_order:
|
| 518 |
+
atom_positions[i, residue_constants.atom_order[atom_name]] = atom.coord # (3,) xyz vector in (l, 37, 3)
|
| 519 |
+
atom_mask[i, residue_constants.atom_order[atom_name]] = True # scalar element in (l, 37)
|
| 520 |
if is_predicted and atom_name == "CA":
|
| 521 |
+
confidence[i] = atom.b_factor / PLDDT_B_FACTOR_SCALE # scalar element in (l,)
|
| 522 |
|
| 523 |
if not sequence or not all(sequence):
|
| 524 |
raise ValueError("Some residue name was not specified correctly.")
|
|
|
|
| 575 |
) -> ProteinChain:
|
| 576 |
"""A simple converter from bs.AtomArray -> ProteinChain.
|
| 577 |
Uses PDB file format as intermediate."""
|
| 578 |
+
atom_array = atom_array.copy() # AtomArray copy with unchanged atom count
|
| 579 |
atom_array.box = None # remove surrounding box, from_pdb won't handle this
|
| 580 |
pdb_file = PDBFile() # pyright: ignore
|
| 581 |
pdb_file.set_structure(atom_array)
|
|
|
|
| 604 |
"residue_index": self.residue_index,
|
| 605 |
"insertion_code": self.insertion_code,
|
| 606 |
"confidence": self.confidence,
|
| 607 |
+
} # arrays share l: positions (l, 37, 3), masks (l, 37), residue fields (l,)
|
| 608 |
for name, values in aligned.items():
|
| 609 |
if not isinstance(values, np.ndarray):
|
| 610 |
raise TypeError(f"{name} must be a NumPy array, got {type(values).__name__}.")
|
|
|
|
| 644 |
)
|
| 645 |
if not np.issubdtype(self.confidence.dtype, np.number):
|
| 646 |
raise TypeError("confidence must use a numeric dtype.")
|
| 647 |
+
atom37_confidence = self.atom37_confidence # (l, 37) or None
|
| 648 |
if atom37_confidence is not None and not isinstance(atom37_confidence, np.ndarray):
|
| 649 |
raise TypeError("atom37_confidence must be a NumPy array when provided.")
|
| 650 |
if (
|
|
|
|
| 694 |
self.atom37_confidence[res_idx_i, i]
|
| 695 |
if self.atom37_confidence is not None
|
| 696 |
else conf
|
| 697 |
+
) # scalar atom confidence
|
| 698 |
atom = bs.Atom(
|
| 699 |
coord=pos,
|
| 700 |
chain_id="A" if self.chain_id is None else self.chain_id,
|
|
|
|
| 708 |
occupancy=1.0,
|
| 709 |
)
|
| 710 |
atoms.append(atom)
|
| 711 |
+
return bs.array(atoms) # AtomArray with n_present_atoms entries
|
| 712 |
|
| 713 |
# Coordinate transformations and dataset adapters
|
| 714 |
def get_normalization_frame(self) -> Affine3D:
|
|
|
|
| 719 |
Returns:
|
| 720 |
Affine3D: [] tensor of Affine3D frame
|
| 721 |
"""
|
| 722 |
+
coords = torch.from_numpy(self.atom37_positions) # (l, 37, 3)
|
| 723 |
+
frame = get_protein_normalization_frame(coords) # one normalization frame
|
| 724 |
|
| 725 |
return frame
|
| 726 |
|
|
|
|
| 733 |
Returns:
|
| 734 |
ProteinChain: Transformed protein chain
|
| 735 |
"""
|
| 736 |
+
# frame is a rigid transform broadcast-compatible with (l, 37, 3) coordinates.
|
| 737 |
+
coords = torch.from_numpy(self.atom37_positions).to(frame.trans.dtype) # (l, 37, 3)
|
| 738 |
+
coords = apply_frame_to_coords(coords, frame) # (l, 37, 3)
|
| 739 |
+
atom37_positions = coords.numpy() # (l, 37, 3)
|
| 740 |
return replace(self, atom37_positions=atom37_positions)
|
| 741 |
|
| 742 |
def normalize_coordinates(self) -> ProteinChain:
|
|
|
|
| 745 |
|
| 746 |
def infer_oxygen(self) -> ProteinChain:
|
| 747 |
"""Oxygen position is fixed given N, CA, C atoms. Infer it if not provided."""
|
| 748 |
+
O_missing_indices = np.argwhere(~np.isfinite(self.atoms["O"]).all(axis=1)).squeeze() # (n_missing,) or () when exactly one oxygen is missing
|
| 749 |
|
| 750 |
+
O_vector = torch.tensor([0.6240, -1.0613, 0.0103], dtype=torch.float32) # (3,)
|
| 751 |
+
N, CA, C = torch.from_numpy(self.atoms[["N", "CA", "C"]]).float().unbind(dim=1) # each (l, 3)
|
| 752 |
+
N = torch.roll(N, -3) # (l, 3); torch.roll keeps the original shape
|
| 753 |
+
N[..., -1, :] = torch.nan # (3,) xyz row
|
| 754 |
|
| 755 |
# Get the frame defined by the CA-C-N atom
|
| 756 |
+
frames = Affine3D.from_graham_schmidt(CA, C, N) # affine batch shape: (l,)
|
| 757 |
+
oxygen_coordinates = frames.apply(O_vector) # (l, 3)
|
| 758 |
+
atom37_positions = self.atom37_positions.copy() # (l, 37, 3)
|
| 759 |
+
atom37_mask = self.atom37_mask.copy() # (l, 37)
|
| 760 |
|
| 761 |
atom37_positions[O_missing_indices, residue_constants.atom_order["O"]] = oxygen_coordinates[
|
| 762 |
O_missing_indices
|
| 763 |
+
].numpy() # (n_missing, 3) or (3,) selected oxygen coordinates
|
| 764 |
atom37_mask[O_missing_indices, residue_constants.atom_order["O"]] = ~np.isnan(
|
| 765 |
atom37_positions[O_missing_indices, residue_constants.atom_order["O"]]
|
| 766 |
+
).any(-1) # (n_missing,) or () selected oxygen mask
|
| 767 |
new_chain = replace(self, atom37_positions=atom37_positions, atom37_mask=atom37_mask)
|
| 768 |
return new_chain
|
| 769 |
|
| 770 |
@cached_property
|
| 771 |
def inferred_cbeta(self) -> np.ndarray:
|
| 772 |
"""Infer cbeta positions based on N, C, CA."""
|
| 773 |
+
N, CA, C = np.moveaxis(self.atoms[["N", "CA", "C"]], 1, 0) # each (l, 3)
|
| 774 |
# See usage in trDesign codebase.
|
| 775 |
# https://github.com/gjoni/trDesign/blob/f2d5930b472e77bfacc2f437b3966e7a708a8d37/02-GD/utils.py#L140
|
| 776 |
+
CB = infer_cb(C, N, CA, 1.522, 1.927, -2.143) # (l, 3)
|
| 777 |
+
return CB # (l, 3)
|
| 778 |
|
| 779 |
def infer_cbeta(self, infer_cbeta_for_glycine: bool = False) -> ProteinChain:
|
| 780 |
"""Return a new chain with inferred CB atoms at all residues except GLY.
|
|
|
|
| 789 |
calculation between two designs for a given structural template, w/
|
| 790 |
CB atoms.
|
| 791 |
"""
|
| 792 |
+
atom37_positions = self.atom37_positions.copy() # (l, 37, 3)
|
| 793 |
+
atom37_mask = self.atom37_mask.copy() # (l, 37)
|
| 794 |
|
| 795 |
+
inferred_cbeta_positions = self.inferred_cbeta # (l, 3)
|
| 796 |
if not infer_cbeta_for_glycine:
|
| 797 |
+
inferred_cbeta_positions[np.array(list(self.sequence)) == "G", :] = np.nan # (n_glycine, 3) selected rows
|
| 798 |
|
| 799 |
+
atom37_positions[:, residue_constants.atom_order["CB"]] = inferred_cbeta_positions # (l, 3) C-beta slice
|
| 800 |
atom37_mask[:, residue_constants.atom_order["CB"]] = ~np.isnan(
|
| 801 |
atom37_positions[:, residue_constants.atom_order["CB"]]
|
| 802 |
+
).any(-1) # (l,) C-beta mask
|
| 803 |
new_chain = replace(self, atom37_positions=atom37_positions, atom37_mask=atom37_mask)
|
| 804 |
return new_chain
|
| 805 |
|
| 806 |
@cached_property
|
| 807 |
def pdist_CA(self) -> np.ndarray:
|
| 808 |
+
CA = self.atoms["CA"] # (l, 3)
|
| 809 |
+
pdist_CA = squareform(pdist(CA)) # (l, l)
|
| 810 |
+
return pdist_CA # (l, l)
|
| 811 |
|
| 812 |
@cached_property
|
| 813 |
def pdist_CB(self) -> np.ndarray:
|
| 814 |
+
pdist_CB = squareform(pdist(self.inferred_cbeta)) # (l, l)
|
| 815 |
+
return pdist_CB # (l, l)
|
| 816 |
|
| 817 |
@classmethod
|
| 818 |
def as_complex(cls, chains: Sequence[ProteinChain]):
|
|
|
|
| 833 |
"atom37_positions": np.full([1, 37, 3], np.inf),
|
| 834 |
"atom37_mask": np.zeros([1, 37], dtype=bool),
|
| 835 |
"confidence": np.array([0]),
|
| 836 |
+
} # one-residue separator arrays: (1,), (1, 37, 3), or (1, 37)
|
| 837 |
|
| 838 |
def join_arrays(arrays: Sequence[np.ndarray], sep: np.ndarray):
|
| 839 |
+
# arrays: (l_i, *trailing_shape); separator: (1, *trailing_shape).
|
| 840 |
if use_chainbreak:
|
| 841 |
full_array = []
|
| 842 |
for array in arrays:
|
| 843 |
full_array.append(array)
|
| 844 |
full_array.append(sep)
|
| 845 |
full_array = full_array[:-1]
|
| 846 |
+
return np.concatenate(full_array, 0) # (sum(chain_lengths) + n_chains - 1, *trailing_shape)
|
| 847 |
else:
|
| 848 |
+
return np.concatenate(arrays, 0) # (sum(chain_lengths), *trailing_shape)
|
| 849 |
|
| 850 |
array_args: dict[str, np.ndarray] = {
|
| 851 |
name: join_arrays([getattr(chain, name) for chain in chains], sep)
|
| 852 |
for name, sep in sep_tokens.items()
|
| 853 |
+
} # each array retains its trailing atom/xyz axes
|
| 854 |
|
| 855 |
chain_break = residue_constants.CHAIN_BREAK_TOKEN if use_chainbreak else ""
|
| 856 |
return cls(
|
|
|
|
| 878 |
raise ValueError(
|
| 879 |
f"Non-polymer {nonpolymer.comp_id!r} has no coordinate table."
|
| 880 |
)
|
| 881 |
+
chain_coords = self.atom37_positions[self.atom37_mask] # (n_present_atoms, 3)
|
| 882 |
+
distance = cdist(nonpolymer_array.coord, chain_coords) # (n_ligand_atoms, n_present_atoms)
|
| 883 |
|
| 884 |
+
is_contact = distance < 5 # (n_ligand_atoms, n_present_atoms)
|
| 885 |
if not is_contact.any():
|
| 886 |
continue
|
| 887 |
+
contacting_atoms = np.where(is_contact.any(0))[0] # (n_contacting_atoms,)
|
| 888 |
+
chain_index = np.where(self.atom37_mask)[0] # (n_present_atoms,)
|
| 889 |
+
contacting_residues = np.unique(chain_index[contacting_atoms]) # (n_contacting_residues,)
|
| 890 |
|
| 891 |
result = {
|
| 892 |
"ligand": nonpolymer.name,
|
|
|
|
| 900 |
self, indices: list[int | str], ignore_x_mismatch: bool = False
|
| 901 |
) -> ProteinChain:
|
| 902 |
numeric_indices = [idx if isinstance(idx, int) else int(idx[1:]) for idx in indices]
|
| 903 |
+
mask = np.isin(self.residue_index, numeric_indices) # (l,)
|
| 904 |
new = self[mask]
|
| 905 |
mismatches = []
|
| 906 |
for aa, idx in zip(new.sequence, indices, strict=False):
|
|
|
|
| 932 |
# Convert to tensors and add batch dimension
|
| 933 |
coordinates = (
|
| 934 |
torch.from_numpy(self.atom37_positions).float().unsqueeze(0)
|
| 935 |
+
) # X has shape (1, l, 37, 3).; (1, l, 37, 3)
|
| 936 |
+
plddt = torch.from_numpy(self.confidence).float().unsqueeze(0) # P: (1, l); (1, l)
|
| 937 |
residue_index = (
|
| 938 |
torch.from_numpy(self.residue_index).long().unsqueeze(0)
|
| 939 |
+
) # R has shape (1, l).; (1, l)
|
| 940 |
|
| 941 |
+
return coordinates, plddt, residue_index # (1, l, 37, 3), (1, l), (1, l)
|
| 942 |
|
| 943 |
# Sequence access, interchange, and compact storage
|
| 944 |
def __getitem__(self, idx: int | list[int] | slice | np.ndarray | torch.Tensor):
|
| 945 |
+
# idx selects residues; an integer is promoted to a length-one index.
|
| 946 |
+
# Returned fields retain atom/xyz trailing axes with the selected residue count.
|
| 947 |
if isinstance(idx, int):
|
| 948 |
idx = [idx]
|
| 949 |
if isinstance(idx, torch.Tensor):
|
| 950 |
+
idx = idx.cpu().numpy() # same index shape
|
| 951 |
|
| 952 |
sequence = slice_python_object_as_numpy(self.sequence, idx)
|
| 953 |
return replace(
|
|
|
|
| 967 |
return len(self.sequence)
|
| 968 |
|
| 969 |
def cbeta_contacts(self, distance_threshold: float = 8.0) -> np.ndarray:
|
| 970 |
+
distance = self.pdist_CB # (l, l)
|
| 971 |
+
contacts = (distance < distance_threshold).astype(np.int64) # (l, l)
|
| 972 |
+
contacts[np.isnan(distance)] = -1 # (n_missing_pairs,) selected entries
|
| 973 |
np.fill_diagonal(contacts, -1)
|
| 974 |
+
return contacts # (l, l)
|
| 975 |
|
| 976 |
def to_pdb(self, path: PathOrBuffer, include_insertions: bool = True):
|
| 977 |
"""Dssp works better w/o insertions."""
|
|
|
|
| 1001 |
"mode": CIFColumn(data=CIFData(array=np.array(["global", "local"]), dtype=np.str_)),
|
| 1002 |
"name": CIFColumn(data=CIFData(array=np.array(["pLDDT", "pLDDT"]), dtype=np.str_)),
|
| 1003 |
},
|
| 1004 |
+
) # each CIF metric column has shape (2,)
|
| 1005 |
|
| 1006 |
# table is a duplicate of data already in the atom array, but
|
| 1007 |
# needed by molstar to render pLDDT / confidence
|
|
|
|
| 1051 |
need more than 2**32 residues..."""
|
| 1052 |
dct = {k: v for k, v in asdict(self).items() if k not in ["mmcif"]}
|
| 1053 |
if backbone_only:
|
| 1054 |
+
dct["atom37_mask"][:, 3:] = False # (l, 34) mask slice for atoms beyond N/CA/C
|
| 1055 |
+
dct["atom37_positions"] = dct["atom37_positions"][dct["atom37_mask"]] # (n_present_atoms, 3)
|
| 1056 |
if dct.get("atom37_confidence") is not None:
|
| 1057 |
+
dct["atom37_confidence"] = dct["atom37_confidence"][dct["atom37_mask"]] # (n_present_atoms,)
|
| 1058 |
else:
|
| 1059 |
dct.pop("atom37_confidence", None)
|
| 1060 |
|
|
|
|
| 1062 |
if isinstance(v, np.ndarray):
|
| 1063 |
match v.dtype:
|
| 1064 |
case np.int64:
|
| 1065 |
+
dct[k] = v.astype(np.int32) # v.shape
|
| 1066 |
case np.float64 | np.float32:
|
| 1067 |
+
dct[k] = v.astype(np.float16) # v.shape
|
| 1068 |
case _:
|
| 1069 |
pass
|
| 1070 |
if json_serializable:
|
|
|
|
| 1086 |
|
| 1087 |
for k, v in dct.items():
|
| 1088 |
if isinstance(v, list):
|
| 1089 |
+
dct[k] = np.array(v) # shape inferred from serialized nested list
|
| 1090 |
|
| 1091 |
+
atom37 = np.full((*dct["atom37_mask"].shape, 3), np.nan) # (l, 37, 3)
|
| 1092 |
+
atom37[dct["atom37_mask"]] = dct["atom37_positions"] # (n_present_atoms, 3) selected coordinates
|
| 1093 |
+
dct["atom37_positions"] = atom37 # (l, 37, 3)
|
| 1094 |
if "atom37_confidence" in dct:
|
| 1095 |
+
atom37_conf = np.full(dct["atom37_mask"].shape, np.nan, dtype=np.float32) # (l, 37)
|
| 1096 |
+
atom37_conf[dct["atom37_mask"]] = dct["atom37_confidence"] # (n_present_atoms,) selected confidence values
|
| 1097 |
+
dct["atom37_confidence"] = atom37_conf # (l, 37)
|
| 1098 |
dct = {
|
| 1099 |
k: (
|
| 1100 |
v.astype(np.float32)
|
|
|
|
| 1103 |
)
|
| 1104 |
for k, v in dct.items()
|
| 1105 |
if not (k == "atom37_confidence" and v is None)
|
| 1106 |
+
} # each converted array retains its serialized field shape
|
| 1107 |
return cls(**dct, mmcif=None)
|
| 1108 |
|
| 1109 |
@classmethod
|
|
|
|
| 1123 |
|
| 1124 |
# Surface and structural comparison metrics
|
| 1125 |
def sasa(self, by_residue: bool = True):
|
| 1126 |
+
arr = self.atom_array_no_insertions # AtomArray with n_present_atoms entries
|
| 1127 |
if len(arr) == 0:
|
| 1128 |
raise ValueError("SASA requires at least one resolved atom.")
|
| 1129 |
+
sasa_per_atom = bs.sasa(arr) # type: ignore; (n_present_atoms,)
|
| 1130 |
if by_residue:
|
| 1131 |
# Sum per-atom SASA into residue "bins", with np.bincount.
|
| 1132 |
if arr.res_id is None:
|
|
|
|
| 1139 |
np.bincount(arr.res_id, weights=sasa_per_atom)[1:],
|
| 1140 |
np.zeros(num_trailing_residues),
|
| 1141 |
]
|
| 1142 |
+
) # (l,)
|
| 1143 |
+
sasa_per_residue[~self.atom37_mask.any(-1)] = np.nan # (n_missing_residues,) selected entries
|
| 1144 |
if len(sasa_per_residue) != len(self):
|
| 1145 |
raise RuntimeError("Residue SASA output does not align with the protein chain.")
|
| 1146 |
+
return sasa_per_residue # (l,)
|
| 1147 |
+
return sasa_per_atom # (n_present_atoms,)
|
| 1148 |
|
| 1149 |
def sap_score(self, aggregation: str = "atom") -> np.ndarray:
|
| 1150 |
"""Compute per-atom spatial aggregation propensity (SAP).
|
|
|
|
| 1153 |
Protein aggregation sums positive atom scores, following Lauer et al. 2011.
|
| 1154 |
"""
|
| 1155 |
sap_radius = 5.0
|
| 1156 |
+
arr = self.atom_array_no_insertions # AtomArray with n_present_atoms entries
|
| 1157 |
if len(arr) == 0:
|
| 1158 |
raise ValueError("SAP requires at least one resolved atom.")
|
| 1159 |
|
|
|
|
| 1162 |
raise RuntimeError(f"Biotite AtomArray is missing required {name!r} data.")
|
| 1163 |
|
| 1164 |
# compute SASA and residue-specific properties
|
| 1165 |
+
sasa_per_atom = self.sasa(by_residue=False) # (n_present_atoms,)
|
| 1166 |
resid_to_resname = dict(zip(arr.res_id, arr.res_name, strict=False))
|
| 1167 |
|
| 1168 |
+
max_side_chain_asa = np.full(len(self), np.nan) # (l,)
|
| 1169 |
+
res_hydrophobicity = np.full(len(self), np.nan) # (l,)
|
| 1170 |
+
resolved_res_mask = self.atom37_mask.any(-1) # (l,)
|
| 1171 |
num_trailing_residues = len(self) - arr.res_id.max()
|
| 1172 |
|
| 1173 |
max_side_chain_asa[resolved_res_mask] = np.array(
|
| 1174 |
[residue_constants.side_chain_asa[resid_to_resname[i]] for i in np.unique(arr.res_id)]
|
| 1175 |
+
) # (n_resolved_residues,) selected entries
|
| 1176 |
res_hydrophobicity[resolved_res_mask] = np.array(
|
| 1177 |
[residue_constants.hydrophobicity[resid_to_resname[i]] for i in np.unique(arr.res_id)]
|
| 1178 |
+
) # (n_resolved_residues,) selected entries
|
| 1179 |
|
| 1180 |
# compute SAP score
|
| 1181 |
+
is_side_chain = ~bs.filter_peptide_backbone(arr) # (n_present_atoms,)
|
| 1182 |
+
sasa_per_atom[is_side_chain] = 0 # (n_selected_atoms,) selected entries
|
| 1183 |
kdtree = KDTree(arr.coord)
|
| 1184 |
neighbors = kdtree.query_ball_tree(kdtree, sap_radius, p=2.0)
|
| 1185 |
+
sap_by_atom = np.zeros_like(sasa_per_atom) # (n_present_atoms,)
|
| 1186 |
for i, nn_list in enumerate(neighbors):
|
| 1187 |
+
saa_nn = np.zeros_like(sasa_per_atom) # (n_present_atoms,)
|
| 1188 |
+
saa_nn[nn_list] = sasa_per_atom[nn_list] # (n_neighbors,) selected entries
|
| 1189 |
sasa_within_r = np.concatenate(
|
| 1190 |
[
|
| 1191 |
np.bincount(arr.res_id, weights=saa_nn)[1:],
|
| 1192 |
np.zeros(num_trailing_residues),
|
| 1193 |
]
|
| 1194 |
+
) # (l,)
|
| 1195 |
+
sap = np.nansum((sasa_within_r / max_side_chain_asa) * res_hydrophobicity) # scalar NumPy reduction
|
| 1196 |
+
sap_by_atom[i] = sap # scalar entry in (n_present_atoms,)
|
| 1197 |
|
| 1198 |
match aggregation:
|
| 1199 |
case "atom":
|
| 1200 |
+
return sap_by_atom # (n_present_atoms,)
|
| 1201 |
case "residue":
|
| 1202 |
sap_by_residue = np.concatenate(
|
| 1203 |
[
|
|
|
|
| 1207 |
) / (
|
| 1208 |
np.concatenate([np.bincount(arr.res_id)[1:], np.zeros(num_trailing_residues)])
|
| 1209 |
+ 1e-8
|
| 1210 |
+
) # (l,)
|
| 1211 |
+
sap_by_residue[~resolved_res_mask] = np.nan # (n_missing_residues,) selected entries
|
| 1212 |
if len(sap_by_residue) != len(self):
|
| 1213 |
raise RuntimeError("Residue SAP output does not align with the protein chain.")
|
| 1214 |
+
return sap_by_residue # (l,)
|
| 1215 |
case "protein":
|
| 1216 |
return sum(sap_by_atom[sap_by_atom > 0]) # pyright: ignore[reportReturnType]
|
| 1217 |
case _:
|
|
|
|
| 1228 |
|
| 1229 |
# https://www.mdpi.com/2073-4352/11/12/1539
|
| 1230 |
# The non-overlapping-atom approximation can produce globularity above one.
|
| 1231 |
+
mask = self.atom37_mask.any(-1) # (l,)
|
| 1232 |
+
points = self.atom37_positions[self.atom37_mask] # (n_present_atoms, 3)
|
| 1233 |
sequence = [aa for aa, m in zip(self.sequence, mask, strict=False) if m] # type: ignore
|
| 1234 |
+
A, _ = self._mvee(points, tol=1e-3) # A: (3, 3), center: (3, 1)
|
| 1235 |
+
mvee_volume = (4 * np.pi) / (3 * np.sqrt(np.linalg.det(A))) # NumPy scalar ()
|
| 1236 |
volume = sum(residue_constants.amino_acid_volumes[x] for x in sequence)
|
| 1237 |
ratio = volume / mvee_volume
|
| 1238 |
|
| 1239 |
# The paper compares the ellipsoidal profile with scalar t, a measurement
|
| 1240 |
# of elongation. We want a single number, so we multiply by 1/(2t), so
|
| 1241 |
# that value is normalized between 0-1
|
| 1242 |
+
eigenvalues = np.linalg.eigvals(A) # (3,)
|
| 1243 |
+
R = 1 / np.sqrt(eigenvalues) # (3,)
|
| 1244 |
# ellipsoid radii length triangle inequality coefficient
|
| 1245 |
t = max(R[0] / (R[1] + R[2]), R[1] / (R[0] + R[2]), R[2] / (R[0] + R[1]))
|
| 1246 |
elongation_metric = 1 / max(t, 1)
|
|
|
|
| 1251 |
# Finds minimum volume enclosing ellipsoid of a set of points.
|
| 1252 |
# Returns A, c where the ellipse is defined as:
|
| 1253 |
# (x-c).T @ A @ (x-c) = 1
|
| 1254 |
+
# P: (n_input_points, d); d is coordinate dimension. Hull selection changes only point count.
|
| 1255 |
hull = ConvexHull(P)
|
| 1256 |
+
P = P[hull.vertices] # (n_hull_vertices, d)
|
| 1257 |
+
P = P.T # (d, n); n = n_hull_vertices
|
| 1258 |
|
| 1259 |
# Data points
|
| 1260 |
d, n = P.shape
|
| 1261 |
+
Q = np.zeros((d + 1, n)) # (d + 1, n)
|
| 1262 |
+
Q[:d, :] = P[:d, :n] # (d, n) slice
|
| 1263 |
+
Q[d, :] = np.ones((1, n)) # (n,) row, assigned from (1, n)
|
| 1264 |
|
| 1265 |
# Initializations
|
| 1266 |
count = 1
|
| 1267 |
err = 1.0
|
| 1268 |
+
u = np.full((n, 1), 1 / n) # First iteration.; (n, 1)
|
| 1269 |
|
| 1270 |
# Khachiyan Algorithm
|
| 1271 |
for _ in range(max_iter):
|
| 1272 |
+
X = Q.dot(np.diag(u.squeeze())) @ Q.T # (d + 1, d + 1)
|
| 1273 |
+
M = np.diag(Q.T @ np.linalg.inv(X) @ Q) # (n,)
|
| 1274 |
+
maximum, j = np.max(M), np.argmax(M) # each NumPy scalar ()
|
| 1275 |
step_size = (maximum - d - 1) / ((d + 1) * (maximum - 1))
|
| 1276 |
+
new_u = (1 - step_size) * u # (n, 1)
|
| 1277 |
+
new_u[j] += step_size # (1,) selected row
|
| 1278 |
count += 1
|
| 1279 |
+
err = np.linalg.norm(new_u - u) # NumPy scalar ()
|
| 1280 |
+
u = new_u # (n, 1)
|
| 1281 |
if err < tol:
|
| 1282 |
break
|
| 1283 |
else:
|
| 1284 |
raise ValueError("MVEE did not converge")
|
| 1285 |
|
| 1286 |
d = P.shape[0] # Fixed: use P.shape[0] instead of P.shape
|
| 1287 |
+
U = np.diag(u.squeeze()) # (n, n)
|
| 1288 |
|
| 1289 |
# The A matrix for the ellipse
|
| 1290 |
+
A = (1 / d) * np.linalg.inv(P @ U @ P.T - (P @ u) @ (P @ u).T) # (d, d)
|
| 1291 |
|
| 1292 |
# Center of the ellipse
|
| 1293 |
+
c = P @ u # (d, 1)
|
| 1294 |
|
| 1295 |
+
return A, c # (d, d), (d, 1)
|
| 1296 |
|
| 1297 |
def radius_of_gyration(self):
|
| 1298 |
+
arr = self.atom_array_no_insertions # AtomArray with n_present_atoms entries
|
| 1299 |
return bs.gyration_radius(arr)
|
| 1300 |
|
| 1301 |
def align(
|
|
|
|
| 1388 |
torch.tensor(native.atom37_positions[target_inds]).unsqueeze(0),
|
| 1389 |
torch.tensor(native.atom37_mask[mobile_inds]).unsqueeze(0),
|
| 1390 |
**kwargs,
|
| 1391 |
+
) # score shape follows selected coordinate axes and per_residue
|
| 1392 |
+
return float(lddt) if lddt.numel() == 1 else lddt.numpy().flatten() # scalar or (lddt.numel(),)
|
| 1393 |
|
| 1394 |
def gdt_ts(
|
| 1395 |
self,
|
|
|
|
| 1422 |
& index_by_atom_name(target.atom37_mask[target_inds], "CA", dim=-1)
|
| 1423 |
).unsqueeze(0),
|
| 1424 |
**kwargs,
|
| 1425 |
+
) # () or (n_samples,), selected by reduction
|
| 1426 |
+
return float(gdt_ts) if gdt_ts.numel() == 1 else gdt_ts.numpy().flatten() # scalar or (gdt_ts.numel(),)
|
| 1427 |
|
| 1428 |
@cached_property
|
| 1429 |
def residue_index_no_insertions(self) -> np.ndarray:
|
| 1430 |
+
return self.residue_index + np.cumsum(self.insertion_code != "") # (l,)
|
| 1431 |
|
| 1432 |
@cached_property
|
| 1433 |
def atom_array_no_insertions(self) -> bs.AtomArray:
|
|
|
|
| 1446 |
self.atom37_confidence[res_idx, i]
|
| 1447 |
if self.atom37_confidence is not None
|
| 1448 |
else conf
|
| 1449 |
+
) # scalar atom confidence
|
| 1450 |
atom = bs.Atom(
|
| 1451 |
coord=pos,
|
| 1452 |
# hard coded to as we currently only support single chain structures
|
|
|
|
| 1460 |
occupancy=1.0,
|
| 1461 |
)
|
| 1462 |
atoms.append(atom)
|
| 1463 |
+
return bs.array(atoms) # AtomArray with n_present_atoms entries
|
fastplms/models/esmfold2/esmfold2_protein_complex.py
CHANGED
|
@@ -7,6 +7,13 @@ import itertools
|
|
| 7 |
import random
|
| 8 |
import re
|
| 9 |
import warnings
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 10 |
from collections.abc import Iterable, Sequence
|
| 11 |
from dataclasses import asdict, dataclass, replace
|
| 12 |
from functools import cached_property
|
|
@@ -14,13 +21,6 @@ from pathlib import Path
|
|
| 14 |
from subprocess import check_output
|
| 15 |
from tempfile import TemporaryDirectory
|
| 16 |
from typing import Any
|
| 17 |
-
|
| 18 |
-
import biotite.structure as bs
|
| 19 |
-
import brotli
|
| 20 |
-
import msgpack
|
| 21 |
-
import msgpack_numpy
|
| 22 |
-
import numpy as np
|
| 23 |
-
import torch
|
| 24 |
from biotite.database import rcsb
|
| 25 |
from biotite.file import InvalidFileError
|
| 26 |
from biotite.structure.io.pdb import PDBFile
|
|
@@ -50,6 +50,7 @@ from .esmfold2_protein_chain import (
|
|
| 50 |
)
|
| 51 |
from .esmfold2_utils_types import PathOrBuffer
|
| 52 |
|
|
|
|
| 53 |
SINGLE_LETTER_CHAIN_IDS = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789"
|
| 54 |
|
| 55 |
|
|
@@ -73,14 +74,15 @@ def _parse_operation_expression(expression: str) -> list[tuple[str, ...]]:
|
|
| 73 |
|
| 74 |
def _apply_transformations_fast(chains, transformation_dict, operations):
|
| 75 |
"""Return transformed copies of each affected protein chain."""
|
|
|
|
| 76 |
transformed_chains = []
|
| 77 |
for chain in chains:
|
| 78 |
for operation in operations:
|
| 79 |
-
coordinates = chain.atom37_positions.copy()
|
| 80 |
for op_step in operation:
|
| 81 |
transform = transformation_dict[op_step]
|
| 82 |
-
coordinates = matrix_rotate(coordinates, transform.rotation)
|
| 83 |
-
coordinates += transform.target_translation
|
| 84 |
transformed_chains.append(replace(chain, atom37_positions=coordinates))
|
| 85 |
return transformed_chains
|
| 86 |
|
|
@@ -124,16 +126,17 @@ class DockQResult:
|
|
| 124 |
class ProteinComplex:
|
| 125 |
"""Dataclass with atom37 representation of an entire protein complex."""
|
| 126 |
|
|
|
|
| 127 |
id: str
|
| 128 |
sequence: str
|
| 129 |
entity_id: np.ndarray # entities map to unique sequences
|
| 130 |
chain_id: np.ndarray # multiple chains might share an entity id
|
| 131 |
sym_id: np.ndarray # complexes might be copies of the same chain
|
| 132 |
-
residue_index: np.ndarray
|
| 133 |
-
insertion_code: np.ndarray
|
| 134 |
-
atom37_positions: np.ndarray
|
| 135 |
-
atom37_mask: np.ndarray
|
| 136 |
-
confidence: np.ndarray
|
| 137 |
# This metadata is parsed from the MMCIF file. For synthetic data, we do a best effort.
|
| 138 |
metadata: ProteinComplexMetadata
|
| 139 |
atom37_confidence: np.ndarray | None = None # P has shape (l, 37).
|
|
@@ -141,25 +144,25 @@ class ProteinComplex:
|
|
| 141 |
# Coordinate completion, concatenation, and comparison
|
| 142 |
def infer_oxygen(self) -> ProteinComplex:
|
| 143 |
"""Oxygen position is fixed given N, CA, C atoms. Infer it if not provided."""
|
| 144 |
-
O_missing_indices = np.argwhere(~np.isfinite(self.atoms["O"]).all(axis=1)).squeeze()
|
| 145 |
|
| 146 |
-
O_vector = torch.tensor([0.6240, -1.0613, 0.0103], dtype=torch.float32)
|
| 147 |
-
N, CA, C = torch.from_numpy(self.atoms[["N", "CA", "C"]]).float().unbind(dim=1)
|
| 148 |
-
N = torch.roll(N, -3)
|
| 149 |
-
N[..., -1, :] = torch.nan
|
| 150 |
|
| 151 |
# Get the frame defined by the CA-C-N atom
|
| 152 |
-
frames = Affine3D.from_graham_schmidt(CA, C, N)
|
| 153 |
-
oxygen_coordinates = frames.apply(O_vector)
|
| 154 |
-
atom37_positions = self.atom37_positions.copy()
|
| 155 |
-
atom37_mask = self.atom37_mask.copy()
|
| 156 |
|
| 157 |
atom37_positions[O_missing_indices, residue_constants.atom_order["O"]] = oxygen_coordinates[
|
| 158 |
O_missing_indices
|
| 159 |
-
].numpy()
|
| 160 |
atom37_mask[O_missing_indices, residue_constants.atom_order["O"]] = ~np.isnan(
|
| 161 |
atom37_positions[O_missing_indices, residue_constants.atom_order["O"]]
|
| 162 |
-
).any(-1)
|
| 163 |
new_chain = replace(self, atom37_positions=atom37_positions, atom37_mask=atom37_mask)
|
| 164 |
return new_chain
|
| 165 |
|
|
@@ -176,20 +179,20 @@ class ProteinComplex:
|
|
| 176 |
calculation between two designs for a given structural template, w/
|
| 177 |
CB atoms.
|
| 178 |
"""
|
| 179 |
-
atom37_positions = self.atom37_positions.copy()
|
| 180 |
-
atom37_mask = self.atom37_mask.copy()
|
| 181 |
|
| 182 |
-
N, CA, C = np.moveaxis(self.atoms[["N", "CA", "C"]], 1, 0)
|
| 183 |
# See usage in trDesign codebase.
|
| 184 |
# https://github.com/gjoni/trDesign/blob/f2d5930b472e77bfacc2f437b3966e7a708a8d37/02-GD/utils.py#L140
|
| 185 |
-
inferred_cbeta_positions = infer_cb(C, N, CA, 1.522, 1.927, -2.143)
|
| 186 |
if not infer_cbeta_for_glycine:
|
| 187 |
-
inferred_cbeta_positions[np.array(list(self.sequence)) == "G", :] = np.nan
|
| 188 |
|
| 189 |
-
atom37_positions[:, residue_constants.atom_order["CB"]] = inferred_cbeta_positions
|
| 190 |
atom37_mask[:, residue_constants.atom_order["CB"]] = ~np.isnan(
|
| 191 |
atom37_positions[:, residue_constants.atom_order["CB"]]
|
| 192 |
-
).any(-1)
|
| 193 |
new_chain = replace(self, atom37_positions=atom37_positions, atom37_mask=atom37_mask)
|
| 194 |
return new_chain
|
| 195 |
|
|
@@ -281,8 +284,8 @@ class ProteinComplex:
|
|
| 281 |
torch.tensor(target.atom37_positions[target_inds]).unsqueeze(0),
|
| 282 |
torch.tensor(aligned.atom37_mask[mobile_inds]).unsqueeze(0),
|
| 283 |
**kwargs,
|
| 284 |
-
)
|
| 285 |
-
return float(lddt) if lddt.numel() == 1 else lddt.numpy().flatten()
|
| 286 |
|
| 287 |
def gdt_ts(
|
| 288 |
self,
|
|
@@ -317,8 +320,8 @@ class ProteinComplex:
|
|
| 317 |
& index_by_atom_name(target.atom37_mask[target_inds], "CA", dim=-1)
|
| 318 |
).unsqueeze(0),
|
| 319 |
**kwargs,
|
| 320 |
-
)
|
| 321 |
-
return float(gdt_ts) if gdt_ts.numel() == 1 else gdt_ts.numpy().flatten()
|
| 322 |
|
| 323 |
def dockq(self, native: ProteinComplex):
|
| 324 |
# This function uses dockqv2 to compute the DockQ score. Because it does a mapping
|
|
@@ -452,7 +455,7 @@ class ProteinComplex:
|
|
| 452 |
"entity_id": self.entity_id,
|
| 453 |
"chain_id": self.chain_id,
|
| 454 |
"sym_id": self.sym_id,
|
| 455 |
-
}
|
| 456 |
for name, values in aligned.items():
|
| 457 |
if not isinstance(values, np.ndarray):
|
| 458 |
raise TypeError(f"{name} must be a NumPy array, got {type(values).__name__}.")
|
|
@@ -489,7 +492,7 @@ class ProteinComplex:
|
|
| 489 |
)
|
| 490 |
if not np.issubdtype(self.confidence.dtype, np.number):
|
| 491 |
raise TypeError("confidence must use a numeric dtype.")
|
| 492 |
-
atom37_confidence = self.atom37_confidence
|
| 493 |
if atom37_confidence is not None and not isinstance(atom37_confidence, np.ndarray):
|
| 494 |
raise TypeError("atom37_confidence must be a NumPy array when provided.")
|
| 495 |
if (
|
|
@@ -506,6 +509,7 @@ class ProteinComplex:
|
|
| 506 |
NOTE: When slicing with a boolean mask, it's possible that the output array won't
|
| 507 |
be the expected length. This is because we do our best to preserve chainbreak tokens.
|
| 508 |
"""
|
|
|
|
| 509 |
|
| 510 |
if isinstance(idx, int):
|
| 511 |
idx = [idx]
|
|
@@ -513,8 +517,8 @@ class ProteinComplex:
|
|
| 513 |
raise ValueError("ProteinComplex doesn't supports indexing with lists of indices")
|
| 514 |
|
| 515 |
if isinstance(idx, np.ndarray):
|
| 516 |
-
is_chainbreak = np.asarray([s == "|" for s in self.sequence])
|
| 517 |
-
idx = idx.astype(bool) | is_chainbreak
|
| 518 |
|
| 519 |
complex = self._unsafe_slice(idx)
|
| 520 |
if len(complex) == 0:
|
|
@@ -524,17 +528,18 @@ class ProteinComplex:
|
|
| 524 |
chainbreak_runs = np.asarray(
|
| 525 |
[complex.sequence[i : i + 2] == "||" for i in range(len(complex.sequence) - 1)]
|
| 526 |
+ [complex.sequence[-1] == "|"]
|
| 527 |
-
)
|
| 528 |
# We should remove as many chainbreaks as possible from the start of the sequence
|
| 529 |
for i in range(len(chainbreak_runs)):
|
| 530 |
if complex.sequence[i] == "|":
|
| 531 |
-
chainbreak_runs[i] = True
|
| 532 |
else:
|
| 533 |
break
|
| 534 |
complex = complex._unsafe_slice(~chainbreak_runs)
|
| 535 |
return complex
|
| 536 |
|
| 537 |
def _unsafe_slice(self, idx: int | list[int] | slice | np.ndarray):
|
|
|
|
| 538 |
sequence = slice_python_object_as_numpy(self.sequence, idx)
|
| 539 |
return replace(
|
| 540 |
self,
|
|
@@ -569,7 +574,7 @@ class ProteinComplex:
|
|
| 569 |
|
| 570 |
@cached_property
|
| 571 |
def chain_lengths(self) -> np.ndarray:
|
| 572 |
-
return np.diff(self.chain_boundaries, axis=1).flatten()
|
| 573 |
|
| 574 |
@cached_property
|
| 575 |
def chain_boundaries(self) -> list[tuple[int, int]]:
|
|
@@ -664,11 +669,11 @@ class ProteinComplex:
|
|
| 664 |
# Iterate over chains, build KDTree for each chain
|
| 665 |
kdtrees = []
|
| 666 |
|
| 667 |
-
CA = self.atoms["CA"]
|
| 668 |
|
| 669 |
for start, end in self.chain_boundaries:
|
| 670 |
-
chain_CA = CA[start:end]
|
| 671 |
-
chain_CA = chain_CA[np.isfinite(chain_CA).all(axis=-1)]
|
| 672 |
kdtrees.append(KDTree(chain_CA))
|
| 673 |
|
| 674 |
return kdtrees
|
|
@@ -676,25 +681,25 @@ class ProteinComplex:
|
|
| 676 |
def chain_adjacency(self, cutoff: float = 8.0) -> np.ndarray:
|
| 677 |
# Compute adjacency matrix for protein complex
|
| 678 |
num_chains = self.num_chains
|
| 679 |
-
adjacency = np.zeros((num_chains, num_chains), dtype=bool)
|
| 680 |
for (i, kdtree), (j, kdtree2) in itertools.combinations(
|
| 681 |
enumerate(self.per_chain_kd_trees), 2
|
| 682 |
):
|
| 683 |
adj = kdtree.query_ball_tree(kdtree2, cutoff)
|
| 684 |
any_is_adjacent = any(len(a) > 0 for a in adj)
|
| 685 |
-
adjacency[i, j] = any_is_adjacent
|
| 686 |
-
adjacency[j, i] = any_is_adjacent
|
| 687 |
-
return adjacency
|
| 688 |
|
| 689 |
def chain_adjacency_by_index(self, index: int, cutoff: float = 8.0) -> np.ndarray:
|
| 690 |
num_chains = len(self.chain_boundaries)
|
| 691 |
-
adjacency = np.zeros(num_chains, dtype=bool)
|
| 692 |
for i, kdtree in enumerate(self.per_chain_kd_trees):
|
| 693 |
if i == index:
|
| 694 |
continue
|
| 695 |
adj = kdtree.query_ball_tree(self.per_chain_kd_trees[index], cutoff)
|
| 696 |
-
adjacency[i] = any(len(a) > 0 for a in adj)
|
| 697 |
-
return adjacency
|
| 698 |
|
| 699 |
def add_prefix_to_chain_ids(self, prefix: str) -> ProteinComplex:
|
| 700 |
"""Rename all chains in the complex with a given prefix.
|
|
@@ -715,7 +720,7 @@ class ProteinComplex:
|
|
| 715 |
|
| 716 |
def sasa(self, by_residue: bool = True):
|
| 717 |
chain = self.as_chain(force_conversion=True)
|
| 718 |
-
return chain.sasa(by_residue=by_residue)
|
| 719 |
|
| 720 |
def to_mmcif_string(self) -> str:
|
| 721 |
"""Convert the ProteinComplex to mmCIF format.
|
|
@@ -727,7 +732,7 @@ class ProteinComplex:
|
|
| 727 |
# Collect all atoms from all chains
|
| 728 |
all_atoms = []
|
| 729 |
for chain in self.chain_iter():
|
| 730 |
-
chain_atom_array = chain.atom_array
|
| 731 |
# Convert AtomArray to list of atoms and add to collection
|
| 732 |
all_atoms.extend(chain_atom_array)
|
| 733 |
|
|
@@ -735,7 +740,7 @@ class ProteinComplex:
|
|
| 735 |
if not all_atoms:
|
| 736 |
raise ValueError("No atoms found in protein complex")
|
| 737 |
|
| 738 |
-
atom_array = bs.array(all_atoms)
|
| 739 |
|
| 740 |
# Create CIF file
|
| 741 |
f = CIFFile()
|
|
@@ -786,7 +791,7 @@ class ProteinComplex:
|
|
| 786 |
data=CIFData(array=np.array(entity_descriptions), dtype=np.str_)
|
| 787 |
),
|
| 788 |
},
|
| 789 |
-
)
|
| 790 |
|
| 791 |
# Create _entity_poly section
|
| 792 |
poly_entity_ids = []
|
|
@@ -814,7 +819,7 @@ class ProteinComplex:
|
|
| 814 |
data=CIFData(array=np.array(poly_sequences), dtype=np.str_)
|
| 815 |
),
|
| 816 |
},
|
| 817 |
-
)
|
| 818 |
|
| 819 |
# Create _struct_asym section
|
| 820 |
asym_ids = []
|
|
@@ -835,18 +840,18 @@ class ProteinComplex:
|
|
| 835 |
),
|
| 836 |
"details": CIFColumn(data=CIFData(array=np.array(asym_details), dtype=np.str_)),
|
| 837 |
},
|
| 838 |
-
)
|
| 839 |
|
| 840 |
# Construction, PDB interchange, and compact storage
|
| 841 |
@classmethod
|
| 842 |
def from_pdb(
|
| 843 |
cls, path: PathOrBuffer, id: str | None = None, is_predicted: bool = False
|
| 844 |
) -> ProteinComplex:
|
| 845 |
-
atom_array = PDBFile.read(path).get_structure(model=1, extra_fields=["b_factor"])
|
| 846 |
|
| 847 |
chains = []
|
| 848 |
for chain in bs.chain_iter(atom_array):
|
| 849 |
-
chain = chain[~chain.hetero]
|
| 850 |
if len(chain) == 0:
|
| 851 |
continue
|
| 852 |
chains.append(ProteinChain.from_atomarray(chain, id, is_predicted))
|
|
@@ -855,8 +860,8 @@ class ProteinComplex:
|
|
| 855 |
def to_pdb(self, path: PathOrBuffer, include_insertions: bool = True):
|
| 856 |
atom_array = None
|
| 857 |
for chain in self.chain_iter():
|
| 858 |
-
carr = chain.atom_array if include_insertions else chain.atom_array_no_insertions
|
| 859 |
-
atom_array = carr if atom_array is None else atom_array + carr
|
| 860 |
f = PDBFile()
|
| 861 |
f.set_structure(atom_array)
|
| 862 |
f.write(path)
|
|
@@ -909,21 +914,21 @@ class ProteinComplex:
|
|
| 909 |
# Frozen dataclasses do not make their NumPy members immutable. Work on a
|
| 910 |
# private mask so requesting a compact backbone payload cannot clear the
|
| 911 |
# caller's side-chain atoms in-place.
|
| 912 |
-
atom37_mask = dct["atom37_mask"].copy()
|
| 913 |
-
atom37_mask[:, 3:] = False
|
| 914 |
-
dct["atom37_mask"] = atom37_mask
|
| 915 |
-
dct["atom37_positions"] = dct["atom37_positions"][dct["atom37_mask"]]
|
| 916 |
if dct.get("atom37_confidence") is not None:
|
| 917 |
-
dct["atom37_confidence"] = dct["atom37_confidence"][dct["atom37_mask"]]
|
| 918 |
else:
|
| 919 |
dct.pop("atom37_confidence", None)
|
| 920 |
for k, v in dct.items():
|
| 921 |
if isinstance(v, np.ndarray):
|
| 922 |
match v.dtype:
|
| 923 |
case np.int64:
|
| 924 |
-
dct[k] = v.astype(np.int32)
|
| 925 |
case np.float64 | np.float32:
|
| 926 |
-
dct[k] = v.astype(np.float16)
|
| 927 |
case _:
|
| 928 |
pass
|
| 929 |
if json_serializable:
|
|
@@ -948,15 +953,15 @@ class ProteinComplex:
|
|
| 948 |
|
| 949 |
for k, v in dct.items():
|
| 950 |
if isinstance(v, list):
|
| 951 |
-
dct[k] = np.array(v)
|
| 952 |
|
| 953 |
-
atom37 = np.full((*dct["atom37_mask"].shape, 3), np.nan)
|
| 954 |
-
atom37[dct["atom37_mask"]] = dct["atom37_positions"]
|
| 955 |
-
dct["atom37_positions"] = atom37
|
| 956 |
if "atom37_confidence" in dct:
|
| 957 |
-
atom37_conf = np.full(dct["atom37_mask"].shape, np.nan, dtype=np.float32)
|
| 958 |
-
atom37_conf[dct["atom37_mask"]] = dct["atom37_confidence"]
|
| 959 |
-
dct["atom37_confidence"] = atom37_conf
|
| 960 |
dct = {
|
| 961 |
k: (
|
| 962 |
v.astype(np.float32)
|
|
@@ -964,7 +969,7 @@ class ProteinComplex:
|
|
| 964 |
else v
|
| 965 |
)
|
| 966 |
for k, v in dct.items()
|
| 967 |
-
}
|
| 968 |
if "chain_boundaries" in dct:
|
| 969 |
del dct["chain_boundaries"]
|
| 970 |
if "chain_boundaries" in dct["metadata"]:
|
|
@@ -1030,12 +1035,13 @@ class ProteinComplex:
|
|
| 1030 |
|
| 1031 |
# TODO(roshan): Make a proper protein complex class
|
| 1032 |
def join_arrays(arrays: Sequence[np.ndarray], sep: np.ndarray):
|
|
|
|
| 1033 |
full_array = []
|
| 1034 |
for array in arrays:
|
| 1035 |
full_array.append(array)
|
| 1036 |
full_array.append(sep)
|
| 1037 |
full_array = full_array[:-1]
|
| 1038 |
-
return np.concatenate(full_array, 0)
|
| 1039 |
|
| 1040 |
sep_tokens = {
|
| 1041 |
"residue_index": np.array([-1]),
|
|
@@ -1043,22 +1049,22 @@ class ProteinComplex:
|
|
| 1043 |
"atom37_positions": np.full([1, 37, 3], np.nan),
|
| 1044 |
"atom37_mask": np.zeros([1, 37], dtype=bool),
|
| 1045 |
"confidence": np.array([0]),
|
| 1046 |
-
}
|
| 1047 |
|
| 1048 |
any_has_atom37_conf = any(c.atom37_confidence is not None for c in chains)
|
| 1049 |
if any_has_atom37_conf:
|
| 1050 |
-
sep_tokens["atom37_confidence"] = np.full([1, 37], np.nan, dtype=np.float32)
|
| 1051 |
|
| 1052 |
def _get_chain_attr(chain: ProteinChain, name: str) -> np.ndarray:
|
| 1053 |
-
val = getattr(chain, name)
|
| 1054 |
if val is None and name == "atom37_confidence":
|
| 1055 |
-
return np.full([len(chain), 37], np.nan, dtype=np.float32)
|
| 1056 |
-
return val
|
| 1057 |
|
| 1058 |
array_args: dict[str, np.ndarray] = {
|
| 1059 |
name: join_arrays([_get_chain_attr(chain, name) for chain in chains], sep)
|
| 1060 |
for name, sep in sep_tokens.items()
|
| 1061 |
-
}
|
| 1062 |
|
| 1063 |
multimer_arrays = []
|
| 1064 |
chain2num_max = -1
|
|
@@ -1070,7 +1076,7 @@ class ProteinComplex:
|
|
| 1070 |
num_res = c.residue_index.shape[0]
|
| 1071 |
if c.chain_id not in chain2num:
|
| 1072 |
chain2num[c.chain_id] = (chain2num_max := chain2num_max + 1)
|
| 1073 |
-
chain_id_array = np.full([num_res], chain2num[c.chain_id], dtype=np.int64)
|
| 1074 |
|
| 1075 |
if c.entity_id is None:
|
| 1076 |
entity_num = (ent2num_max := ent2num_max + 1)
|
|
@@ -1078,9 +1084,9 @@ class ProteinComplex:
|
|
| 1078 |
if c.entity_id not in ent2num:
|
| 1079 |
ent2num[c.entity_id] = (ent2num_max := ent2num_max + 1)
|
| 1080 |
entity_num = ent2num[c.entity_id]
|
| 1081 |
-
entity_id_array = np.full([num_res], entity_num, dtype=np.int64)
|
| 1082 |
|
| 1083 |
-
sym_id_array = np.full([num_res], i, dtype=np.int64)
|
| 1084 |
|
| 1085 |
multimer_arrays.append(
|
| 1086 |
{
|
|
@@ -1092,11 +1098,11 @@ class ProteinComplex:
|
|
| 1092 |
|
| 1093 |
total_index += num_res + 1
|
| 1094 |
|
| 1095 |
-
sep = np.array([-1])
|
| 1096 |
update = {
|
| 1097 |
name: join_arrays([dct[name] for dct in multimer_arrays], sep=sep)
|
| 1098 |
for name in ["chain_id", "entity_id", "sym_id"]
|
| 1099 |
-
}
|
| 1100 |
array_args.update(update)
|
| 1101 |
|
| 1102 |
metadata = ProteinComplexMetadata(
|
|
@@ -1153,7 +1159,7 @@ def get_assembly_fast(
|
|
| 1153 |
]
|
| 1154 |
if len(structure) == 0:
|
| 1155 |
raise NoProteinError
|
| 1156 |
-
unique_asym_ids = np.unique(structure.label_asym_id) # type: ignore
|
| 1157 |
asym2chain = {}
|
| 1158 |
asym2auth = {}
|
| 1159 |
for asym_id in unique_asym_ids:
|
|
@@ -1167,7 +1173,7 @@ def get_assembly_fast(
|
|
| 1167 |
insertion_code,
|
| 1168 |
confidence,
|
| 1169 |
entity_id,
|
| 1170 |
-
) = chain_to_ndarray(sub_structure, mmcif, chain_id, False)
|
| 1171 |
|
| 1172 |
asym2chain[asym_id] = ProteinChain(
|
| 1173 |
id=mmcif.id or "unknown",
|
|
@@ -1217,12 +1223,13 @@ def get_assembly_fast(
|
|
| 1217 |
|
| 1218 |
|
| 1219 |
def protein_chain_to_protein_complex(chain: ProteinChain) -> ProteinComplex:
|
|
|
|
| 1220 |
if "|" not in chain.sequence:
|
| 1221 |
return ProteinComplex.from_chains([chain])
|
| 1222 |
-
chain_breaks = np.array(list(chain.sequence)) == "|"
|
| 1223 |
-
chain_break_inds = np.where(chain_breaks)[0]
|
| 1224 |
-
chain_break_inds = np.concatenate([[0], chain_break_inds, [len(chain)]])
|
| 1225 |
-
chain_break_inds = np.array(list(itertools.pairwise(chain_break_inds)))
|
| 1226 |
complex_chains = []
|
| 1227 |
for start, end in chain_break_inds:
|
| 1228 |
if start != 0:
|
|
|
|
| 7 |
import random
|
| 8 |
import re
|
| 9 |
import warnings
|
| 10 |
+
import biotite.structure as bs
|
| 11 |
+
import brotli
|
| 12 |
+
import msgpack
|
| 13 |
+
import msgpack_numpy
|
| 14 |
+
import numpy as np
|
| 15 |
+
import torch
|
| 16 |
+
|
| 17 |
from collections.abc import Iterable, Sequence
|
| 18 |
from dataclasses import asdict, dataclass, replace
|
| 19 |
from functools import cached_property
|
|
|
|
| 21 |
from subprocess import check_output
|
| 22 |
from tempfile import TemporaryDirectory
|
| 23 |
from typing import Any
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 24 |
from biotite.database import rcsb
|
| 25 |
from biotite.file import InvalidFileError
|
| 26 |
from biotite.structure.io.pdb import PDBFile
|
|
|
|
| 50 |
)
|
| 51 |
from .esmfold2_utils_types import PathOrBuffer
|
| 52 |
|
| 53 |
+
|
| 54 |
SINGLE_LETTER_CHAIN_IDS = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789"
|
| 55 |
|
| 56 |
|
|
|
|
| 74 |
|
| 75 |
def _apply_transformations_fast(chains, transformation_dict, operations):
|
| 76 |
"""Return transformed copies of each affected protein chain."""
|
| 77 |
+
# Each chain supplies coordinates (l, 37, 3); rotations are (3, 3), translations are (3,).
|
| 78 |
transformed_chains = []
|
| 79 |
for chain in chains:
|
| 80 |
for operation in operations:
|
| 81 |
+
coordinates = chain.atom37_positions.copy() # (l, 37, 3)
|
| 82 |
for op_step in operation:
|
| 83 |
transform = transformation_dict[op_step]
|
| 84 |
+
coordinates = matrix_rotate(coordinates, transform.rotation) # (l, 37, 3)
|
| 85 |
+
coordinates += transform.target_translation # (l, 37, 3)
|
| 86 |
transformed_chains.append(replace(chain, atom37_positions=coordinates))
|
| 87 |
return transformed_chains
|
| 88 |
|
|
|
|
| 126 |
class ProteinComplex:
|
| 127 |
"""Dataclass with atom37 representation of an entire protein complex."""
|
| 128 |
|
| 129 |
+
# l is len(sequence), including separator rows when present.
|
| 130 |
id: str
|
| 131 |
sequence: str
|
| 132 |
entity_id: np.ndarray # entities map to unique sequences
|
| 133 |
chain_id: np.ndarray # multiple chains might share an entity id
|
| 134 |
sym_id: np.ndarray # complexes might be copies of the same chain
|
| 135 |
+
residue_index: np.ndarray # (l,)
|
| 136 |
+
insertion_code: np.ndarray # (l,)
|
| 137 |
+
atom37_positions: np.ndarray # (l, 37, 3)
|
| 138 |
+
atom37_mask: np.ndarray # (l, 37)
|
| 139 |
+
confidence: np.ndarray # (l,)
|
| 140 |
# This metadata is parsed from the MMCIF file. For synthetic data, we do a best effort.
|
| 141 |
metadata: ProteinComplexMetadata
|
| 142 |
atom37_confidence: np.ndarray | None = None # P has shape (l, 37).
|
|
|
|
| 144 |
# Coordinate completion, concatenation, and comparison
|
| 145 |
def infer_oxygen(self) -> ProteinComplex:
|
| 146 |
"""Oxygen position is fixed given N, CA, C atoms. Infer it if not provided."""
|
| 147 |
+
O_missing_indices = np.argwhere(~np.isfinite(self.atoms["O"]).all(axis=1)).squeeze() # (n_missing,) or () when exactly one oxygen is missing
|
| 148 |
|
| 149 |
+
O_vector = torch.tensor([0.6240, -1.0613, 0.0103], dtype=torch.float32) # (3,)
|
| 150 |
+
N, CA, C = torch.from_numpy(self.atoms[["N", "CA", "C"]]).float().unbind(dim=1) # each (l, 3)
|
| 151 |
+
N = torch.roll(N, -3) # (l, 3); torch.roll keeps the original shape
|
| 152 |
+
N[..., -1, :] = torch.nan # (3,) xyz row
|
| 153 |
|
| 154 |
# Get the frame defined by the CA-C-N atom
|
| 155 |
+
frames = Affine3D.from_graham_schmidt(CA, C, N) # affine batch shape: (l,)
|
| 156 |
+
oxygen_coordinates = frames.apply(O_vector) # (l, 3)
|
| 157 |
+
atom37_positions = self.atom37_positions.copy() # (l, 37, 3)
|
| 158 |
+
atom37_mask = self.atom37_mask.copy() # (l, 37)
|
| 159 |
|
| 160 |
atom37_positions[O_missing_indices, residue_constants.atom_order["O"]] = oxygen_coordinates[
|
| 161 |
O_missing_indices
|
| 162 |
+
].numpy() # (n_missing, 3) or (3,) selected oxygen coordinates
|
| 163 |
atom37_mask[O_missing_indices, residue_constants.atom_order["O"]] = ~np.isnan(
|
| 164 |
atom37_positions[O_missing_indices, residue_constants.atom_order["O"]]
|
| 165 |
+
).any(-1) # (n_missing,) or () selected oxygen mask
|
| 166 |
new_chain = replace(self, atom37_positions=atom37_positions, atom37_mask=atom37_mask)
|
| 167 |
return new_chain
|
| 168 |
|
|
|
|
| 179 |
calculation between two designs for a given structural template, w/
|
| 180 |
CB atoms.
|
| 181 |
"""
|
| 182 |
+
atom37_positions = self.atom37_positions.copy() # (l, 37, 3)
|
| 183 |
+
atom37_mask = self.atom37_mask.copy() # (l, 37)
|
| 184 |
|
| 185 |
+
N, CA, C = np.moveaxis(self.atoms[["N", "CA", "C"]], 1, 0) # each (l, 3)
|
| 186 |
# See usage in trDesign codebase.
|
| 187 |
# https://github.com/gjoni/trDesign/blob/f2d5930b472e77bfacc2f437b3966e7a708a8d37/02-GD/utils.py#L140
|
| 188 |
+
inferred_cbeta_positions = infer_cb(C, N, CA, 1.522, 1.927, -2.143) # (l, 3)
|
| 189 |
if not infer_cbeta_for_glycine:
|
| 190 |
+
inferred_cbeta_positions[np.array(list(self.sequence)) == "G", :] = np.nan # (n_glycine, 3) selected rows
|
| 191 |
|
| 192 |
+
atom37_positions[:, residue_constants.atom_order["CB"]] = inferred_cbeta_positions # (l, 3) C-beta slice
|
| 193 |
atom37_mask[:, residue_constants.atom_order["CB"]] = ~np.isnan(
|
| 194 |
atom37_positions[:, residue_constants.atom_order["CB"]]
|
| 195 |
+
).any(-1) # (l,) C-beta mask
|
| 196 |
new_chain = replace(self, atom37_positions=atom37_positions, atom37_mask=atom37_mask)
|
| 197 |
return new_chain
|
| 198 |
|
|
|
|
| 284 |
torch.tensor(target.atom37_positions[target_inds]).unsqueeze(0),
|
| 285 |
torch.tensor(aligned.atom37_mask[mobile_inds]).unsqueeze(0),
|
| 286 |
**kwargs,
|
| 287 |
+
) # score shape follows selected coordinate axes and per_residue
|
| 288 |
+
return float(lddt) if lddt.numel() == 1 else lddt.numpy().flatten() # scalar or (lddt.numel(),)
|
| 289 |
|
| 290 |
def gdt_ts(
|
| 291 |
self,
|
|
|
|
| 320 |
& index_by_atom_name(target.atom37_mask[target_inds], "CA", dim=-1)
|
| 321 |
).unsqueeze(0),
|
| 322 |
**kwargs,
|
| 323 |
+
) # () or (n_samples,), selected by reduction
|
| 324 |
+
return float(gdt_ts) if gdt_ts.numel() == 1 else gdt_ts.numpy().flatten() # scalar or (gdt_ts.numel(),)
|
| 325 |
|
| 326 |
def dockq(self, native: ProteinComplex):
|
| 327 |
# This function uses dockqv2 to compute the DockQ score. Because it does a mapping
|
|
|
|
| 455 |
"entity_id": self.entity_id,
|
| 456 |
"chain_id": self.chain_id,
|
| 457 |
"sym_id": self.sym_id,
|
| 458 |
+
} # arrays share l: positions (l, 37, 3), masks (l, 37), residue fields (l,)
|
| 459 |
for name, values in aligned.items():
|
| 460 |
if not isinstance(values, np.ndarray):
|
| 461 |
raise TypeError(f"{name} must be a NumPy array, got {type(values).__name__}.")
|
|
|
|
| 492 |
)
|
| 493 |
if not np.issubdtype(self.confidence.dtype, np.number):
|
| 494 |
raise TypeError("confidence must use a numeric dtype.")
|
| 495 |
+
atom37_confidence = self.atom37_confidence # (l, 37) or None
|
| 496 |
if atom37_confidence is not None and not isinstance(atom37_confidence, np.ndarray):
|
| 497 |
raise TypeError("atom37_confidence must be a NumPy array when provided.")
|
| 498 |
if (
|
|
|
|
| 509 |
NOTE: When slicing with a boolean mask, it's possible that the output array won't
|
| 510 |
be the expected length. This is because we do our best to preserve chainbreak tokens.
|
| 511 |
"""
|
| 512 |
+
# idx selects residue rows; masks retain separator rows before repeated separators are removed.
|
| 513 |
|
| 514 |
if isinstance(idx, int):
|
| 515 |
idx = [idx]
|
|
|
|
| 517 |
raise ValueError("ProteinComplex doesn't supports indexing with lists of indices")
|
| 518 |
|
| 519 |
if isinstance(idx, np.ndarray):
|
| 520 |
+
is_chainbreak = np.asarray([s == "|" for s in self.sequence]) # (l,)
|
| 521 |
+
idx = idx.astype(bool) | is_chainbreak # (l,)
|
| 522 |
|
| 523 |
complex = self._unsafe_slice(idx)
|
| 524 |
if len(complex) == 0:
|
|
|
|
| 528 |
chainbreak_runs = np.asarray(
|
| 529 |
[complex.sequence[i : i + 2] == "||" for i in range(len(complex.sequence) - 1)]
|
| 530 |
+ [complex.sequence[-1] == "|"]
|
| 531 |
+
) # (selected_length,)
|
| 532 |
# We should remove as many chainbreaks as possible from the start of the sequence
|
| 533 |
for i in range(len(chainbreak_runs)):
|
| 534 |
if complex.sequence[i] == "|":
|
| 535 |
+
chainbreak_runs[i] = True # scalar mask element
|
| 536 |
else:
|
| 537 |
break
|
| 538 |
complex = complex._unsafe_slice(~chainbreak_runs)
|
| 539 |
return complex
|
| 540 |
|
| 541 |
def _unsafe_slice(self, idx: int | list[int] | slice | np.ndarray):
|
| 542 |
+
# idx selects residue rows; atom/xyz trailing axes and aligned field lengths are retained.
|
| 543 |
sequence = slice_python_object_as_numpy(self.sequence, idx)
|
| 544 |
return replace(
|
| 545 |
self,
|
|
|
|
| 574 |
|
| 575 |
@cached_property
|
| 576 |
def chain_lengths(self) -> np.ndarray:
|
| 577 |
+
return np.diff(self.chain_boundaries, axis=1).flatten() # (n_chains,)
|
| 578 |
|
| 579 |
@cached_property
|
| 580 |
def chain_boundaries(self) -> list[tuple[int, int]]:
|
|
|
|
| 669 |
# Iterate over chains, build KDTree for each chain
|
| 670 |
kdtrees = []
|
| 671 |
|
| 672 |
+
CA = self.atoms["CA"] # (l, 3)
|
| 673 |
|
| 674 |
for start, end in self.chain_boundaries:
|
| 675 |
+
chain_CA = CA[start:end] # (chain_length, 3)
|
| 676 |
+
chain_CA = chain_CA[np.isfinite(chain_CA).all(axis=-1)] # (n_finite_ca, 3)
|
| 677 |
kdtrees.append(KDTree(chain_CA))
|
| 678 |
|
| 679 |
return kdtrees
|
|
|
|
| 681 |
def chain_adjacency(self, cutoff: float = 8.0) -> np.ndarray:
|
| 682 |
# Compute adjacency matrix for protein complex
|
| 683 |
num_chains = self.num_chains
|
| 684 |
+
adjacency = np.zeros((num_chains, num_chains), dtype=bool) # (n_chains, n_chains)
|
| 685 |
for (i, kdtree), (j, kdtree2) in itertools.combinations(
|
| 686 |
enumerate(self.per_chain_kd_trees), 2
|
| 687 |
):
|
| 688 |
adj = kdtree.query_ball_tree(kdtree2, cutoff)
|
| 689 |
any_is_adjacent = any(len(a) > 0 for a in adj)
|
| 690 |
+
adjacency[i, j] = any_is_adjacent # scalar matrix entry
|
| 691 |
+
adjacency[j, i] = any_is_adjacent # scalar matrix entry
|
| 692 |
+
return adjacency # (n_chains, n_chains)
|
| 693 |
|
| 694 |
def chain_adjacency_by_index(self, index: int, cutoff: float = 8.0) -> np.ndarray:
|
| 695 |
num_chains = len(self.chain_boundaries)
|
| 696 |
+
adjacency = np.zeros(num_chains, dtype=bool) # (n_chains,)
|
| 697 |
for i, kdtree in enumerate(self.per_chain_kd_trees):
|
| 698 |
if i == index:
|
| 699 |
continue
|
| 700 |
adj = kdtree.query_ball_tree(self.per_chain_kd_trees[index], cutoff)
|
| 701 |
+
adjacency[i] = any(len(a) > 0 for a in adj) # scalar vector entry
|
| 702 |
+
return adjacency # (n_chains,)
|
| 703 |
|
| 704 |
def add_prefix_to_chain_ids(self, prefix: str) -> ProteinComplex:
|
| 705 |
"""Rename all chains in the complex with a given prefix.
|
|
|
|
| 720 |
|
| 721 |
def sasa(self, by_residue: bool = True):
|
| 722 |
chain = self.as_chain(force_conversion=True)
|
| 723 |
+
return chain.sasa(by_residue=by_residue) # (l,) if by_residue, otherwise (n_present_atoms,)
|
| 724 |
|
| 725 |
def to_mmcif_string(self) -> str:
|
| 726 |
"""Convert the ProteinComplex to mmCIF format.
|
|
|
|
| 732 |
# Collect all atoms from all chains
|
| 733 |
all_atoms = []
|
| 734 |
for chain in self.chain_iter():
|
| 735 |
+
chain_atom_array = chain.atom_array # AtomArray with n_chain_atoms entries
|
| 736 |
# Convert AtomArray to list of atoms and add to collection
|
| 737 |
all_atoms.extend(chain_atom_array)
|
| 738 |
|
|
|
|
| 740 |
if not all_atoms:
|
| 741 |
raise ValueError("No atoms found in protein complex")
|
| 742 |
|
| 743 |
+
atom_array = bs.array(all_atoms) # AtomArray with total n_present_atoms entries
|
| 744 |
|
| 745 |
# Create CIF file
|
| 746 |
f = CIFFile()
|
|
|
|
| 791 |
data=CIFData(array=np.array(entity_descriptions), dtype=np.str_)
|
| 792 |
),
|
| 793 |
},
|
| 794 |
+
) # each entity column has shape (n_entities,)
|
| 795 |
|
| 796 |
# Create _entity_poly section
|
| 797 |
poly_entity_ids = []
|
|
|
|
| 819 |
data=CIFData(array=np.array(poly_sequences), dtype=np.str_)
|
| 820 |
),
|
| 821 |
},
|
| 822 |
+
) # each polymer column has shape (n_polymer_entities,)
|
| 823 |
|
| 824 |
# Create _struct_asym section
|
| 825 |
asym_ids = []
|
|
|
|
| 840 |
),
|
| 841 |
"details": CIFColumn(data=CIFData(array=np.array(asym_details), dtype=np.str_)),
|
| 842 |
},
|
| 843 |
+
) # each asym column has shape (n_chains,)
|
| 844 |
|
| 845 |
# Construction, PDB interchange, and compact storage
|
| 846 |
@classmethod
|
| 847 |
def from_pdb(
|
| 848 |
cls, path: PathOrBuffer, id: str | None = None, is_predicted: bool = False
|
| 849 |
) -> ProteinComplex:
|
| 850 |
+
atom_array = PDBFile.read(path).get_structure(model=1, extra_fields=["b_factor"]) # AtomArray with n_file_atoms entries
|
| 851 |
|
| 852 |
chains = []
|
| 853 |
for chain in bs.chain_iter(atom_array):
|
| 854 |
+
chain = chain[~chain.hetero] # AtomArray with n_nonhetero_atoms entries
|
| 855 |
if len(chain) == 0:
|
| 856 |
continue
|
| 857 |
chains.append(ProteinChain.from_atomarray(chain, id, is_predicted))
|
|
|
|
| 860 |
def to_pdb(self, path: PathOrBuffer, include_insertions: bool = True):
|
| 861 |
atom_array = None
|
| 862 |
for chain in self.chain_iter():
|
| 863 |
+
carr = chain.atom_array if include_insertions else chain.atom_array_no_insertions # AtomArray with n_chain_atoms entries
|
| 864 |
+
atom_array = carr if atom_array is None else atom_array + carr # AtomArray containing accumulated chain atoms
|
| 865 |
f = PDBFile()
|
| 866 |
f.set_structure(atom_array)
|
| 867 |
f.write(path)
|
|
|
|
| 914 |
# Frozen dataclasses do not make their NumPy members immutable. Work on a
|
| 915 |
# private mask so requesting a compact backbone payload cannot clear the
|
| 916 |
# caller's side-chain atoms in-place.
|
| 917 |
+
atom37_mask = dct["atom37_mask"].copy() # (l, 37)
|
| 918 |
+
atom37_mask[:, 3:] = False # (l, 34) mask slice for atoms beyond N/CA/C
|
| 919 |
+
dct["atom37_mask"] = atom37_mask # (l, 37)
|
| 920 |
+
dct["atom37_positions"] = dct["atom37_positions"][dct["atom37_mask"]] # (n_present_atoms, 3)
|
| 921 |
if dct.get("atom37_confidence") is not None:
|
| 922 |
+
dct["atom37_confidence"] = dct["atom37_confidence"][dct["atom37_mask"]] # (n_present_atoms,)
|
| 923 |
else:
|
| 924 |
dct.pop("atom37_confidence", None)
|
| 925 |
for k, v in dct.items():
|
| 926 |
if isinstance(v, np.ndarray):
|
| 927 |
match v.dtype:
|
| 928 |
case np.int64:
|
| 929 |
+
dct[k] = v.astype(np.int32) # v.shape
|
| 930 |
case np.float64 | np.float32:
|
| 931 |
+
dct[k] = v.astype(np.float16) # v.shape
|
| 932 |
case _:
|
| 933 |
pass
|
| 934 |
if json_serializable:
|
|
|
|
| 953 |
|
| 954 |
for k, v in dct.items():
|
| 955 |
if isinstance(v, list):
|
| 956 |
+
dct[k] = np.array(v) # shape inferred from serialized nested list
|
| 957 |
|
| 958 |
+
atom37 = np.full((*dct["atom37_mask"].shape, 3), np.nan) # (l, 37, 3)
|
| 959 |
+
atom37[dct["atom37_mask"]] = dct["atom37_positions"] # (n_present_atoms, 3) selected coordinates
|
| 960 |
+
dct["atom37_positions"] = atom37 # (l, 37, 3)
|
| 961 |
if "atom37_confidence" in dct:
|
| 962 |
+
atom37_conf = np.full(dct["atom37_mask"].shape, np.nan, dtype=np.float32) # (l, 37)
|
| 963 |
+
atom37_conf[dct["atom37_mask"]] = dct["atom37_confidence"] # (n_present_atoms,) selected confidence values
|
| 964 |
+
dct["atom37_confidence"] = atom37_conf # (l, 37)
|
| 965 |
dct = {
|
| 966 |
k: (
|
| 967 |
v.astype(np.float32)
|
|
|
|
| 969 |
else v
|
| 970 |
)
|
| 971 |
for k, v in dct.items()
|
| 972 |
+
} # each converted array retains its serialized field shape
|
| 973 |
if "chain_boundaries" in dct:
|
| 974 |
del dct["chain_boundaries"]
|
| 975 |
if "chain_boundaries" in dct["metadata"]:
|
|
|
|
| 1035 |
|
| 1036 |
# TODO(roshan): Make a proper protein complex class
|
| 1037 |
def join_arrays(arrays: Sequence[np.ndarray], sep: np.ndarray):
|
| 1038 |
+
# arrays: (l_i, *trailing_shape); sep: (1, *trailing_shape).
|
| 1039 |
full_array = []
|
| 1040 |
for array in arrays:
|
| 1041 |
full_array.append(array)
|
| 1042 |
full_array.append(sep)
|
| 1043 |
full_array = full_array[:-1]
|
| 1044 |
+
return np.concatenate(full_array, 0) # (sum(chain_lengths) + n_chains - 1, *trailing_shape)
|
| 1045 |
|
| 1046 |
sep_tokens = {
|
| 1047 |
"residue_index": np.array([-1]),
|
|
|
|
| 1049 |
"atom37_positions": np.full([1, 37, 3], np.nan),
|
| 1050 |
"atom37_mask": np.zeros([1, 37], dtype=bool),
|
| 1051 |
"confidence": np.array([0]),
|
| 1052 |
+
} # one-residue separator arrays: (1,), (1, 37, 3), or (1, 37)
|
| 1053 |
|
| 1054 |
any_has_atom37_conf = any(c.atom37_confidence is not None for c in chains)
|
| 1055 |
if any_has_atom37_conf:
|
| 1056 |
+
sep_tokens["atom37_confidence"] = np.full([1, 37], np.nan, dtype=np.float32) # (1, 37)
|
| 1057 |
|
| 1058 |
def _get_chain_attr(chain: ProteinChain, name: str) -> np.ndarray:
|
| 1059 |
+
val = getattr(chain, name) # (chain_length, *field_trailing_shape) or None
|
| 1060 |
if val is None and name == "atom37_confidence":
|
| 1061 |
+
return np.full([len(chain), 37], np.nan, dtype=np.float32) # (chain_length, 37)
|
| 1062 |
+
return val # (chain_length, *field_trailing_shape)
|
| 1063 |
|
| 1064 |
array_args: dict[str, np.ndarray] = {
|
| 1065 |
name: join_arrays([_get_chain_attr(chain, name) for chain in chains], sep)
|
| 1066 |
for name, sep in sep_tokens.items()
|
| 1067 |
+
} # fields retain atom/xyz axes; first axis includes chain separators
|
| 1068 |
|
| 1069 |
multimer_arrays = []
|
| 1070 |
chain2num_max = -1
|
|
|
|
| 1076 |
num_res = c.residue_index.shape[0]
|
| 1077 |
if c.chain_id not in chain2num:
|
| 1078 |
chain2num[c.chain_id] = (chain2num_max := chain2num_max + 1)
|
| 1079 |
+
chain_id_array = np.full([num_res], chain2num[c.chain_id], dtype=np.int64) # (chain_length,)
|
| 1080 |
|
| 1081 |
if c.entity_id is None:
|
| 1082 |
entity_num = (ent2num_max := ent2num_max + 1)
|
|
|
|
| 1084 |
if c.entity_id not in ent2num:
|
| 1085 |
ent2num[c.entity_id] = (ent2num_max := ent2num_max + 1)
|
| 1086 |
entity_num = ent2num[c.entity_id]
|
| 1087 |
+
entity_id_array = np.full([num_res], entity_num, dtype=np.int64) # (chain_length,)
|
| 1088 |
|
| 1089 |
+
sym_id_array = np.full([num_res], i, dtype=np.int64) # (chain_length,)
|
| 1090 |
|
| 1091 |
multimer_arrays.append(
|
| 1092 |
{
|
|
|
|
| 1098 |
|
| 1099 |
total_index += num_res + 1
|
| 1100 |
|
| 1101 |
+
sep = np.array([-1]) # (1,)
|
| 1102 |
update = {
|
| 1103 |
name: join_arrays([dct[name] for dct in multimer_arrays], sep=sep)
|
| 1104 |
for name in ["chain_id", "entity_id", "sym_id"]
|
| 1105 |
+
} # each field: (sum(chain_lengths) + n_chains - 1,)
|
| 1106 |
array_args.update(update)
|
| 1107 |
|
| 1108 |
metadata = ProteinComplexMetadata(
|
|
|
|
| 1159 |
]
|
| 1160 |
if len(structure) == 0:
|
| 1161 |
raise NoProteinError
|
| 1162 |
+
unique_asym_ids = np.unique(structure.label_asym_id) # type: ignore; (n_unique_asym_ids,)
|
| 1163 |
asym2chain = {}
|
| 1164 |
asym2auth = {}
|
| 1165 |
for asym_id in unique_asym_ids:
|
|
|
|
| 1173 |
insertion_code,
|
| 1174 |
confidence,
|
| 1175 |
entity_id,
|
| 1176 |
+
) = chain_to_ndarray(sub_structure, mmcif, chain_id, False) # array fields: (l, 37, 3), (l, 37), (l,), (l,), (l,)
|
| 1177 |
|
| 1178 |
asym2chain[asym_id] = ProteinChain(
|
| 1179 |
id=mmcif.id or "unknown",
|
|
|
|
| 1223 |
|
| 1224 |
|
| 1225 |
def protein_chain_to_protein_complex(chain: ProteinChain) -> ProteinComplex:
|
| 1226 |
+
# chain fields share residue axis l; splitting removes chain-break separator rows.
|
| 1227 |
if "|" not in chain.sequence:
|
| 1228 |
return ProteinComplex.from_chains([chain])
|
| 1229 |
+
chain_breaks = np.array(list(chain.sequence)) == "|" # (l,)
|
| 1230 |
+
chain_break_inds = np.where(chain_breaks)[0] # (n_chainbreaks,)
|
| 1231 |
+
chain_break_inds = np.concatenate([[0], chain_break_inds, [len(chain)]]) # (n_chainbreaks + 2,)
|
| 1232 |
+
chain_break_inds = np.array(list(itertools.pairwise(chain_break_inds))) # (n_chainbreaks + 1, 2)
|
| 1233 |
complex_chains = []
|
| 1234 |
for start, end in chain_break_inds:
|
| 1235 |
if start != 0:
|
fastplms/models/esmfold2/esmfold2_protein_structure.py
CHANGED
|
@@ -2,12 +2,12 @@
|
|
| 2 |
|
| 3 |
from __future__ import annotations
|
| 4 |
|
| 5 |
-
from collections.abc import Callable
|
| 6 |
-
from typing import TypeVar
|
| 7 |
-
|
| 8 |
import numpy as np
|
| 9 |
import torch
|
| 10 |
import torch.nn.functional as F
|
|
|
|
|
|
|
|
|
|
| 11 |
from torch import Tensor
|
| 12 |
from torch.amp import autocast # type: ignore
|
| 13 |
|
|
@@ -15,6 +15,7 @@ from .esmfold2_affine3d import Affine3D
|
|
| 15 |
from .esmfold2_misc import unbinpack
|
| 16 |
from .esmfold2_normalize_coordinates import index_by_atom_name
|
| 17 |
|
|
|
|
| 18 |
ArrayOrTensor = TypeVar("ArrayOrTensor", np.ndarray, Tensor)
|
| 19 |
|
| 20 |
|
|
@@ -24,7 +25,7 @@ def _coordinate_operations(
|
|
| 24 |
if isinstance(coordinates, np.ndarray):
|
| 25 |
|
| 26 |
def normalize(X: ArrayOrTensor) -> ArrayOrTensor:
|
| 27 |
-
return X / np.linalg.norm(X, axis=-1, keepdims=True)
|
| 28 |
|
| 29 |
return normalize, np.cross
|
| 30 |
return F.normalize, torch.cross # type: ignore[return-value]
|
|
@@ -42,16 +43,17 @@ def infer_cbeta_from_atom37(
|
|
| 42 |
dihedral in radians used by the checkpoint's training geometry.
|
| 43 |
"""
|
| 44 |
|
| 45 |
-
|
| 46 |
-
|
| 47 |
-
|
|
|
|
| 48 |
normalize, cross = _coordinate_operations(atom37)
|
| 49 |
with np.errstate(invalid="ignore"):
|
| 50 |
-
n_to_ca = n_position - ca_position
|
| 51 |
-
n_to_c = n_position - c_position
|
| 52 |
-
unit_n_to_ca = normalize(n_to_ca)
|
| 53 |
-
normal = normalize(cross(n_to_c, unit_n_to_ca))
|
| 54 |
-
basis = [unit_n_to_ca, cross(normal, unit_n_to_ca), normal]
|
| 55 |
coefficients = [
|
| 56 |
bond_length * np.cos(bond_angle),
|
| 57 |
bond_length * np.sin(bond_angle) * np.cos(dihedral),
|
|
@@ -59,8 +61,8 @@ def infer_cbeta_from_atom37(
|
|
| 59 |
]
|
| 60 |
offset = sum(
|
| 61 |
vector * coefficient for vector, coefficient in zip(basis, coefficients, strict=True)
|
| 62 |
-
)
|
| 63 |
-
return ca_position + offset
|
| 64 |
|
| 65 |
|
| 66 |
def _unpack_alignment_inputs(
|
|
@@ -71,12 +73,13 @@ def _unpack_alignment_inputs(
|
|
| 71 |
) -> tuple[Tensor, Tensor, Tensor | None]:
|
| 72 |
if sequence_id is None:
|
| 73 |
return mobile, target, atom_mask
|
| 74 |
-
|
| 75 |
-
|
|
|
|
| 76 |
if atom_mask is None:
|
| 77 |
-
unpacked_mask = torch.isfinite(unpacked_target).all(dim=-1)
|
| 78 |
else:
|
| 79 |
-
unpacked_mask = unbinpack(atom_mask, sequence_id, pad_value=0)
|
| 80 |
return unpacked_mobile, unpacked_target, unpacked_mask
|
| 81 |
|
| 82 |
|
|
@@ -85,13 +88,14 @@ def _flatten_atom_axes(
|
|
| 85 |
target: Tensor,
|
| 86 |
atom_mask: Tensor | None,
|
| 87 |
) -> tuple[Tensor, Tensor, Tensor | None]:
|
|
|
|
| 88 |
b = mobile.shape[0]
|
| 89 |
-
flat_mobile = mobile.view(b, -1, 3) if mobile.dim() == 4 else mobile
|
| 90 |
-
flat_target = target.view(b, -1, 3) if target.dim() == 4 else target
|
| 91 |
-
flat_mask = atom_mask
|
| 92 |
if flat_mask is not None and flat_mask.dim() == 3:
|
| 93 |
-
flat_mask = flat_mask.view(b, -1)
|
| 94 |
-
return flat_mobile, flat_target, flat_mask
|
| 95 |
|
| 96 |
|
| 97 |
def _masked_coordinates(
|
|
@@ -104,12 +108,12 @@ def _masked_coordinates(
|
|
| 104 |
mobile.shape[:2],
|
| 105 |
dtype=torch.bool,
|
| 106 |
device=mobile.device,
|
| 107 |
-
)
|
| 108 |
return mobile, target, atom_mask
|
| 109 |
-
expanded_mask = atom_mask.unsqueeze(-1)
|
| 110 |
return (
|
| 111 |
-
mobile.masked_fill(~expanded_mask, 0),
|
| 112 |
-
target.masked_fill(~expanded_mask, 0),
|
| 113 |
atom_mask,
|
| 114 |
)
|
| 115 |
|
|
@@ -147,18 +151,19 @@ def compute_alignment_tensors(
|
|
| 147 |
atom_exists_mask,
|
| 148 |
)
|
| 149 |
|
| 150 |
-
|
| 151 |
-
|
| 152 |
-
|
| 153 |
-
|
| 154 |
-
|
| 155 |
-
|
| 156 |
-
|
| 157 |
-
|
| 158 |
-
|
| 159 |
-
|
| 160 |
-
|
| 161 |
-
|
|
|
|
| 162 |
return (
|
| 163 |
centered_mobile,
|
| 164 |
centroid_mobile,
|
|
@@ -185,16 +190,17 @@ def compute_rmsd_no_alignment(
|
|
| 185 |
"""Measure RMSD after alignment using a declared reduction."""
|
| 186 |
|
| 187 |
_validate_reduction(reduction, ("per_residue", "per_sample", "batch"))
|
| 188 |
-
|
|
|
|
| 189 |
if reduction == "per_residue":
|
| 190 |
-
mean_squared_error = difference.square().view(difference.size(0), -1, 9).mean(-1)
|
| 191 |
else:
|
| 192 |
-
mean_squared_error = difference.square().sum(dim=(1, 2)) / num_valid_atoms.squeeze(-1)
|
| 193 |
-
rmsd = torch.sqrt(mean_squared_error)
|
| 194 |
if reduction in {"per_residue", "per_sample"}:
|
| 195 |
return rmsd
|
| 196 |
-
valid_samples = num_valid_atoms.squeeze(-1) > 0
|
| 197 |
-
return rmsd.masked_fill(~valid_samples, 0).sum() / (valid_samples.sum() + 1e-8)
|
| 198 |
|
| 199 |
|
| 200 |
@torch.no_grad()
|
|
@@ -215,19 +221,19 @@ def compute_affine_and_rmsd(
|
|
| 215 |
rotation,
|
| 216 |
num_valid_atoms,
|
| 217 |
) = compute_alignment_tensors(mobile, target, atom_exists_mask, sequence_id)
|
| 218 |
-
translation = torch.matmul(-centroid_mobile, rotation) + centroid_target
|
| 219 |
affine = Affine3D.from_tensor_pair(
|
| 220 |
translation,
|
| 221 |
-
rotation.unsqueeze(dim=-3).transpose(-2, -1),
|
| 222 |
)
|
| 223 |
-
rotated_mobile = torch.matmul(centered_mobile, rotation)
|
| 224 |
rmsd = compute_rmsd_no_alignment(
|
| 225 |
rotated_mobile,
|
| 226 |
centered_target,
|
| 227 |
num_valid_atoms,
|
| 228 |
reduction="batch",
|
| 229 |
)
|
| 230 |
-
return affine, rmsd
|
| 231 |
|
| 232 |
|
| 233 |
def compute_gdt_ts_no_alignment(
|
|
@@ -240,12 +246,13 @@ def compute_gdt_ts_no_alignment(
|
|
| 240 |
|
| 241 |
_validate_reduction(reduction, ("per_sample", "batch"))
|
| 242 |
if atom_exists_mask is None:
|
| 243 |
-
atom_exists_mask = torch.isfinite(target).all(dim=-1)
|
| 244 |
-
|
| 245 |
-
|
| 246 |
-
|
| 247 |
-
|
| 248 |
-
|
| 249 |
-
|
| 250 |
-
|
| 251 |
-
|
|
|
|
|
|
| 2 |
|
| 3 |
from __future__ import annotations
|
| 4 |
|
|
|
|
|
|
|
|
|
|
| 5 |
import numpy as np
|
| 6 |
import torch
|
| 7 |
import torch.nn.functional as F
|
| 8 |
+
|
| 9 |
+
from collections.abc import Callable
|
| 10 |
+
from typing import TypeVar
|
| 11 |
from torch import Tensor
|
| 12 |
from torch.amp import autocast # type: ignore
|
| 13 |
|
|
|
|
| 15 |
from .esmfold2_misc import unbinpack
|
| 16 |
from .esmfold2_normalize_coordinates import index_by_atom_name
|
| 17 |
|
| 18 |
+
|
| 19 |
ArrayOrTensor = TypeVar("ArrayOrTensor", np.ndarray, Tensor)
|
| 20 |
|
| 21 |
|
|
|
|
| 25 |
if isinstance(coordinates, np.ndarray):
|
| 26 |
|
| 27 |
def normalize(X: ArrayOrTensor) -> ArrayOrTensor:
|
| 28 |
+
return X / np.linalg.norm(X, axis=-1, keepdims=True) # X.shape; normalize last axis.
|
| 29 |
|
| 30 |
return normalize, np.cross
|
| 31 |
return F.normalize, torch.cross # type: ignore[return-value]
|
|
|
|
| 43 |
dihedral in radians used by the checkpoint's training geometry.
|
| 44 |
"""
|
| 45 |
|
| 46 |
+
# atom37: (..., 37, 3); ... contains optional batch and residue axes.
|
| 47 |
+
n_position = index_by_atom_name(atom37, "N", dim=-2) # (..., 3)
|
| 48 |
+
ca_position = index_by_atom_name(atom37, "CA", dim=-2) # (..., 3)
|
| 49 |
+
c_position = index_by_atom_name(atom37, "C", dim=-2) # (..., 3)
|
| 50 |
normalize, cross = _coordinate_operations(atom37)
|
| 51 |
with np.errstate(invalid="ignore"):
|
| 52 |
+
n_to_ca = n_position - ca_position # (..., 3)
|
| 53 |
+
n_to_c = n_position - c_position # (..., 3)
|
| 54 |
+
unit_n_to_ca = normalize(n_to_ca) # (..., 3)
|
| 55 |
+
normal = normalize(cross(n_to_c, unit_n_to_ca)) # (..., 3)
|
| 56 |
+
basis = [unit_n_to_ca, cross(normal, unit_n_to_ca), normal] # three arrays (..., 3)
|
| 57 |
coefficients = [
|
| 58 |
bond_length * np.cos(bond_angle),
|
| 59 |
bond_length * np.sin(bond_angle) * np.cos(dihedral),
|
|
|
|
| 61 |
]
|
| 62 |
offset = sum(
|
| 63 |
vector * coefficient for vector, coefficient in zip(basis, coefficients, strict=True)
|
| 64 |
+
) # (..., 3)
|
| 65 |
+
return ca_position + offset # (..., 3)
|
| 66 |
|
| 67 |
|
| 68 |
def _unpack_alignment_inputs(
|
|
|
|
| 73 |
) -> tuple[Tensor, Tensor, Tensor | None]:
|
| 74 |
if sequence_id is None:
|
| 75 |
return mobile, target, atom_mask
|
| 76 |
+
# Packed (b, l, ..., 3) coordinates become (n_sequences, max_length, ..., 3).
|
| 77 |
+
unpacked_mobile = unbinpack(mobile, sequence_id, pad_value=torch.nan) # unpacked coordinate shape
|
| 78 |
+
unpacked_target = unbinpack(target, sequence_id, pad_value=torch.nan) # unpacked coordinate shape
|
| 79 |
if atom_mask is None:
|
| 80 |
+
unpacked_mask = torch.isfinite(unpacked_target).all(dim=-1) # unpacked shape without xyz
|
| 81 |
else:
|
| 82 |
+
unpacked_mask = unbinpack(atom_mask, sequence_id, pad_value=0) # unpacked shape without xyz
|
| 83 |
return unpacked_mobile, unpacked_target, unpacked_mask
|
| 84 |
|
| 85 |
|
|
|
|
| 88 |
target: Tensor,
|
| 89 |
atom_mask: Tensor | None,
|
| 90 |
) -> tuple[Tensor, Tensor, Tensor | None]:
|
| 91 |
+
# Inputs: (b, l, n_atoms, 3) or (b, n, 3); n = l * n_atoms after flattening.
|
| 92 |
b = mobile.shape[0]
|
| 93 |
+
flat_mobile = mobile.view(b, -1, 3) if mobile.dim() == 4 else mobile # (b, n, 3)
|
| 94 |
+
flat_target = target.view(b, -1, 3) if target.dim() == 4 else target # (b, n, 3)
|
| 95 |
+
flat_mask = atom_mask # (b, l, n_atoms), (b, n), or None
|
| 96 |
if flat_mask is not None and flat_mask.dim() == 3:
|
| 97 |
+
flat_mask = flat_mask.view(b, -1) # (b, n)
|
| 98 |
+
return flat_mobile, flat_target, flat_mask # (b, n, 3), (b, n, 3), (b, n) or None
|
| 99 |
|
| 100 |
|
| 101 |
def _masked_coordinates(
|
|
|
|
| 108 |
mobile.shape[:2],
|
| 109 |
dtype=torch.bool,
|
| 110 |
device=mobile.device,
|
| 111 |
+
) # (b, n)
|
| 112 |
return mobile, target, atom_mask
|
| 113 |
+
expanded_mask = atom_mask.unsqueeze(-1) # (b, n, 1)
|
| 114 |
return (
|
| 115 |
+
mobile.masked_fill(~expanded_mask, 0), # (b, n, 3)
|
| 116 |
+
target.masked_fill(~expanded_mask, 0), # (b, n, 3)
|
| 117 |
atom_mask,
|
| 118 |
)
|
| 119 |
|
|
|
|
| 151 |
atom_exists_mask,
|
| 152 |
)
|
| 153 |
|
| 154 |
+
# b now counts unpacked sequences if sequence_id was supplied; n counts atoms.
|
| 155 |
+
num_valid_atoms = atom_exists_mask.sum(dim=-1, keepdim=True) # (b, 1)
|
| 156 |
+
centroid_mobile = mobile.sum(dim=-2, keepdim=True) / num_valid_atoms.unsqueeze(-1) # (b, 1, 3)
|
| 157 |
+
centroid_target = target.sum(dim=-2, keepdim=True) / num_valid_atoms.unsqueeze(-1) # (b, 1, 3)
|
| 158 |
+
centroid_mobile[num_valid_atoms == 0] = 0 # (n_empty, 3)
|
| 159 |
+
centroid_target[num_valid_atoms == 0] = 0 # (n_empty, 3)
|
| 160 |
+
|
| 161 |
+
expanded_mask = atom_exists_mask.unsqueeze(-1) # (b, n, 1)
|
| 162 |
+
centered_mobile = (mobile - centroid_mobile).masked_fill(~expanded_mask, 0) # (b, n, 3)
|
| 163 |
+
centered_target = (target - centroid_target).masked_fill(~expanded_mask, 0) # (b, n, 3)
|
| 164 |
+
covariance = torch.matmul(centered_mobile.transpose(1, 2), centered_target) # (b, 3, 3)
|
| 165 |
+
left_vectors, _, right_vectors = torch.svd(covariance) # (b, 3, 3), (b, 3), (b, 3, 3)
|
| 166 |
+
rotation = torch.matmul(left_vectors, right_vectors.transpose(1, 2)) # (b, 3, 3)
|
| 167 |
return (
|
| 168 |
centered_mobile,
|
| 169 |
centroid_mobile,
|
|
|
|
| 190 |
"""Measure RMSD after alignment using a declared reduction."""
|
| 191 |
|
| 192 |
_validate_reduction(reduction, ("per_residue", "per_sample", "batch"))
|
| 193 |
+
# aligned/target: (b, n, 3); num_valid_atoms: (b, 1).
|
| 194 |
+
difference = aligned - target # (b, n, 3)
|
| 195 |
if reduction == "per_residue":
|
| 196 |
+
mean_squared_error = difference.square().view(difference.size(0), -1, 9).mean(-1) # (b, n / 3)
|
| 197 |
else:
|
| 198 |
+
mean_squared_error = difference.square().sum(dim=(1, 2)) / num_valid_atoms.squeeze(-1) # (b,)
|
| 199 |
+
rmsd = torch.sqrt(mean_squared_error) # (b, n / 3) for per_residue, otherwise (b,)
|
| 200 |
if reduction in {"per_residue", "per_sample"}:
|
| 201 |
return rmsd
|
| 202 |
+
valid_samples = num_valid_atoms.squeeze(-1) > 0 # (b,)
|
| 203 |
+
return rmsd.masked_fill(~valid_samples, 0).sum() / (valid_samples.sum() + 1e-8) # ()
|
| 204 |
|
| 205 |
|
| 206 |
@torch.no_grad()
|
|
|
|
| 221 |
rotation,
|
| 222 |
num_valid_atoms,
|
| 223 |
) = compute_alignment_tensors(mobile, target, atom_exists_mask, sequence_id)
|
| 224 |
+
translation = torch.matmul(-centroid_mobile, rotation) + centroid_target # (b, 1, 3)
|
| 225 |
affine = Affine3D.from_tensor_pair(
|
| 226 |
translation,
|
| 227 |
+
rotation.unsqueeze(dim=-3).transpose(-2, -1), # (b, 1, 3, 3)
|
| 228 |
)
|
| 229 |
+
rotated_mobile = torch.matmul(centered_mobile, rotation) # (b, n, 3)
|
| 230 |
rmsd = compute_rmsd_no_alignment(
|
| 231 |
rotated_mobile,
|
| 232 |
centered_target,
|
| 233 |
num_valid_atoms,
|
| 234 |
reduction="batch",
|
| 235 |
)
|
| 236 |
+
return affine, rmsd # affine shape: (b, 1); rmsd: ()
|
| 237 |
|
| 238 |
|
| 239 |
def compute_gdt_ts_no_alignment(
|
|
|
|
| 246 |
|
| 247 |
_validate_reduction(reduction, ("per_sample", "batch"))
|
| 248 |
if atom_exists_mask is None:
|
| 249 |
+
atom_exists_mask = torch.isfinite(target).all(dim=-1) # (b, n)
|
| 250 |
+
# aligned/target: (b, n, 3); atom_exists_mask: (b, n).
|
| 251 |
+
deviation = torch.linalg.vector_norm(aligned - target, dim=-1) # (b, n)
|
| 252 |
+
counts = atom_exists_mask.sum(dim=-1) # (b,)
|
| 253 |
+
score_1 = ((deviation < 1) * atom_exists_mask).sum(dim=-1) / counts # (b,)
|
| 254 |
+
score_2 = ((deviation < 2) * atom_exists_mask).sum(dim=-1) / counts # (b,)
|
| 255 |
+
score_4 = ((deviation < 4) * atom_exists_mask).sum(dim=-1) / counts # (b,)
|
| 256 |
+
score_8 = ((deviation < 8) * atom_exists_mask).sum(dim=-1) / counts # (b,)
|
| 257 |
+
score = (score_1 + score_2 + score_4 + score_8) * 0.25 # (b,)
|
| 258 |
+
return score.mean() if reduction == "batch" else score # () for batch, otherwise (b,)
|
fastplms/models/esmfold2/esmfold2_residue_constants.py
CHANGED
|
@@ -24,12 +24,13 @@ exactly against the pinned Biohub implementation.
|
|
| 24 |
from __future__ import annotations
|
| 25 |
|
| 26 |
import functools
|
|
|
|
|
|
|
| 27 |
from collections import defaultdict, namedtuple
|
| 28 |
from collections.abc import Mapping
|
| 29 |
from pathlib import Path
|
| 30 |
from typing import Any, cast
|
| 31 |
|
| 32 |
-
import numpy as np
|
| 33 |
|
| 34 |
ca_ca = 3.80209737096
|
| 35 |
chi_angles_atoms = {
|
|
|
|
| 24 |
from __future__ import annotations
|
| 25 |
|
| 26 |
import functools
|
| 27 |
+
import numpy as np
|
| 28 |
+
|
| 29 |
from collections import defaultdict, namedtuple
|
| 30 |
from collections.abc import Mapping
|
| 31 |
from pathlib import Path
|
| 32 |
from typing import Any, cast
|
| 33 |
|
|
|
|
| 34 |
|
| 35 |
ca_ca = 3.80209737096
|
| 36 |
chi_angles_atoms = {
|
fastplms/models/esmfold2/esmfold2_sequential_dataclass.py
CHANGED
|
@@ -2,15 +2,16 @@
|
|
| 2 |
|
| 3 |
from __future__ import annotations
|
| 4 |
|
|
|
|
|
|
|
| 5 |
from abc import ABC, abstractmethod
|
| 6 |
from collections.abc import Iterable
|
| 7 |
from dataclasses import Field, dataclass, fields, replace
|
| 8 |
from typing import Any, Self
|
| 9 |
|
| 10 |
-
import numpy as np
|
| 11 |
-
|
| 12 |
from .esmfold2_misc import concat_objects, slice_any_object
|
| 13 |
|
|
|
|
| 14 |
Index = int | list[int] | slice | np.ndarray
|
| 15 |
|
| 16 |
|
|
|
|
| 2 |
|
| 3 |
from __future__ import annotations
|
| 4 |
|
| 5 |
+
import numpy as np
|
| 6 |
+
|
| 7 |
from abc import ABC, abstractmethod
|
| 8 |
from collections.abc import Iterable
|
| 9 |
from dataclasses import Field, dataclass, fields, replace
|
| 10 |
from typing import Any, Self
|
| 11 |
|
|
|
|
|
|
|
| 12 |
from .esmfold2_misc import concat_objects, slice_any_object
|
| 13 |
|
| 14 |
+
|
| 15 |
Index = int | list[int] | slice | np.ndarray
|
| 16 |
|
| 17 |
|
fastplms/models/esmfold2/esmfold2_system.py
CHANGED
|
@@ -4,9 +4,11 @@ from __future__ import annotations
|
|
| 4 |
|
| 5 |
import io
|
| 6 |
import subprocess
|
|
|
|
| 7 |
from pathlib import Path
|
| 8 |
from typing import Any, TypeAlias
|
| 9 |
|
|
|
|
| 10 |
PathLike: TypeAlias = str | Path
|
| 11 |
PathOrBuffer: TypeAlias = PathLike | io.StringIO
|
| 12 |
|
|
|
|
| 4 |
|
| 5 |
import io
|
| 6 |
import subprocess
|
| 7 |
+
|
| 8 |
from pathlib import Path
|
| 9 |
from typing import Any, TypeAlias
|
| 10 |
|
| 11 |
+
|
| 12 |
PathLike: TypeAlias = str | Path
|
| 13 |
PathOrBuffer: TypeAlias = PathLike | io.StringIO
|
| 14 |
|
fastplms/models/esmfold2/esmfold2_types.py
CHANGED
|
@@ -6,6 +6,7 @@ from . import esmfold2_input_builder as _input_schema
|
|
| 6 |
from .esmfold2_msa import MSA
|
| 7 |
from .esmfold2_parsing import FastaEntry
|
| 8 |
|
|
|
|
| 9 |
Modification = _input_schema.Modification
|
| 10 |
ProteinInput = _input_schema.ProteinInput
|
| 11 |
RNAInput = _input_schema.RNAInput
|
|
|
|
| 6 |
from .esmfold2_msa import MSA
|
| 7 |
from .esmfold2_parsing import FastaEntry
|
| 8 |
|
| 9 |
+
|
| 10 |
Modification = _input_schema.Modification
|
| 11 |
ProteinInput = _input_schema.ProteinInput
|
| 12 |
RNAInput = _input_schema.RNAInput
|
fastplms/models/esmfold2/esmfold2_utils_types.py
CHANGED
|
@@ -9,9 +9,11 @@ from __future__ import annotations
|
|
| 9 |
|
| 10 |
import io
|
| 11 |
import os
|
|
|
|
| 12 |
from dataclasses import dataclass
|
| 13 |
from typing import TypeAlias
|
| 14 |
|
|
|
|
| 15 |
PathLike: TypeAlias = str | os.PathLike[str]
|
| 16 |
PathOrBuffer: TypeAlias = PathLike | io.TextIOBase
|
| 17 |
|
|
|
|
| 9 |
|
| 10 |
import io
|
| 11 |
import os
|
| 12 |
+
|
| 13 |
from dataclasses import dataclass
|
| 14 |
from typing import TypeAlias
|
| 15 |
|
| 16 |
+
|
| 17 |
PathLike: TypeAlias = str | os.PathLike[str]
|
| 18 |
PathOrBuffer: TypeAlias = PathLike | io.TextIOBase
|
| 19 |
|
fastplms/models/esmfold2/modeling_esmfold2.py
CHANGED
|
@@ -2,10 +2,12 @@
|
|
| 2 |
|
| 3 |
Quickstart::
|
| 4 |
|
| 5 |
-
from
|
|
|
|
| 6 |
|
| 7 |
-
model =
|
| 8 |
-
|
|
|
|
| 9 |
|
| 10 |
For multi-chain, ligand, and MSA inputs, use ``model.input_types`` together
|
| 11 |
with ``model.fold(...)`` or ``model.prepare_structure_input(...)``.
|
|
@@ -17,15 +19,15 @@ import gc
|
|
| 17 |
import importlib
|
| 18 |
import importlib.metadata
|
| 19 |
import math
|
|
|
|
|
|
|
|
|
|
|
|
|
| 20 |
from collections.abc import Mapping
|
| 21 |
from contextlib import contextmanager
|
| 22 |
from dataclasses import asdict, dataclass
|
| 23 |
from pathlib import Path
|
| 24 |
from typing import Any, ClassVar, Literal, cast
|
| 25 |
-
|
| 26 |
-
import torch
|
| 27 |
-
import torch.nn as nn
|
| 28 |
-
import torch.nn.functional as F
|
| 29 |
from torch import Tensor
|
| 30 |
from tqdm.auto import tqdm
|
| 31 |
from transformers.modeling_outputs import ModelOutput
|
|
@@ -33,6 +35,7 @@ from transformers.modeling_utils import PreTrainedModel
|
|
| 33 |
|
| 34 |
from ...attention import get_attn_implementation, set_config_attn_implementation
|
| 35 |
|
|
|
|
| 36 |
try:
|
| 37 |
from fastplms.models.ttt import FastPLMTestTimeTrainingMixin, TTTConfig
|
| 38 |
except ModuleNotFoundError as error:
|
|
@@ -195,6 +198,7 @@ class _ESMFold2ESMplusplusAdapter(nn.Module):
|
|
| 195 |
compute_sae: bool = True,
|
| 196 |
normalize_sae: bool = False,
|
| 197 |
):
|
|
|
|
| 198 |
del return_dict, compute_sae, normalize_sae
|
| 199 |
output = self.model(
|
| 200 |
input_ids=input_ids,
|
|
@@ -206,14 +210,14 @@ class _ESMFold2ESMplusplusAdapter(nn.Module):
|
|
| 206 |
esmfold2_hidden_states=True,
|
| 207 |
)
|
| 208 |
if output_hidden_states:
|
| 209 |
-
hidden_states = output.hidden_states
|
| 210 |
if hidden_states is None:
|
| 211 |
raise RuntimeError("ESM++ did not return requested hidden states.")
|
| 212 |
if isinstance(hidden_states, torch.Tensor):
|
| 213 |
-
output.hidden_states = hidden_states
|
| 214 |
else:
|
| 215 |
-
output.hidden_states = torch.stack(tuple(hidden_states), dim=0)
|
| 216 |
-
return output
|
| 217 |
|
| 218 |
|
| 219 |
def _load_fastplms_esmplusplus_for_esmfold2(
|
|
@@ -547,14 +551,15 @@ class PairTransition(nn.Module):
|
|
| 547 |
self._chunk_size = chunk_size
|
| 548 |
|
| 549 |
def forward(self, x: Tensor) -> Tensor:
|
|
|
|
| 550 |
if self._chunk_size is None or x.shape[1] <= self._chunk_size:
|
| 551 |
-
return self.ffn(self.norm(x))
|
| 552 |
out: list[Tensor] = []
|
| 553 |
for s in range(0, x.shape[1], self._chunk_size):
|
| 554 |
e = min(s + self._chunk_size, x.shape[1])
|
| 555 |
-
sl = x[:, s:e]
|
| 556 |
out.append(self.ffn(self.norm(sl)))
|
| 557 |
-
return torch.cat(out, dim=1)
|
| 558 |
|
| 559 |
|
| 560 |
class ConfidenceHead(nn.Module):
|
|
@@ -569,8 +574,8 @@ class ConfidenceHead(nn.Module):
|
|
| 569 |
d_pair = config.d_pair
|
| 570 |
d_inputs = config.inputs.d_inputs
|
| 571 |
|
| 572 |
-
boundaries = torch.linspace(ch.min_dist, ch.max_dist, ch.distogram_bins - 1)
|
| 573 |
-
self.register_buffer("boundaries", boundaries)
|
| 574 |
self.dist_bin_pairwise_embed = nn.Embedding(ch.distogram_bins, d_pair)
|
| 575 |
|
| 576 |
self.s_norm = nn.LayerNorm(d_single)
|
|
@@ -594,7 +599,7 @@ class ConfidenceHead(nn.Module):
|
|
| 594 |
max_atoms_per_token = 23
|
| 595 |
self.plddt_weight = nn.Parameter(
|
| 596 |
torch.zeros(max_atoms_per_token, d_single, ch.num_plddt_bins)
|
| 597 |
-
)
|
| 598 |
|
| 599 |
self.pae_ln = nn.LayerNorm(d_pair)
|
| 600 |
self.pae_head = nn.Linear(d_pair, ch.num_pae_bins, bias=False)
|
|
@@ -604,7 +609,7 @@ class ConfidenceHead(nn.Module):
|
|
| 604 |
|
| 605 |
self.resolved_ln = nn.LayerNorm(d_single)
|
| 606 |
# 2 = resolved logits ([unresolved, resolved]).
|
| 607 |
-
self.resolved_weight = nn.Parameter(torch.zeros(max_atoms_per_token, d_single, 2))
|
| 608 |
|
| 609 |
def set_kernel_backend(self, backend: str | None) -> None:
|
| 610 |
self.folding_trunk.set_kernel_backend(backend)
|
|
@@ -614,14 +619,16 @@ class ConfidenceHead(nn.Module):
|
|
| 614 |
|
| 615 |
@staticmethod
|
| 616 |
def _repeat_batch(x: Tensor, num_diffusion_samples: int) -> Tensor:
|
| 617 |
-
|
|
|
|
| 618 |
|
| 619 |
@staticmethod
|
| 620 |
def _flatten_sample_axis(x: Tensor) -> Tensor:
|
|
|
|
| 621 |
if x.ndim == 4:
|
| 622 |
b, mult, n, c = x.shape
|
| 623 |
-
return x.reshape(b * mult, n, c)
|
| 624 |
-
return x
|
| 625 |
|
| 626 |
def forward(
|
| 627 |
self,
|
|
@@ -638,55 +645,56 @@ class ConfidenceHead(nn.Module):
|
|
| 638 |
relative_position_encoding: Tensor | None = None,
|
| 639 |
token_bonds_encoding: Tensor | None = None,
|
| 640 |
) -> dict[str, Tensor]:
|
| 641 |
-
|
|
|
|
| 642 |
|
| 643 |
-
z_base = self.z_norm(z)
|
| 644 |
if relative_position_encoding is not None:
|
| 645 |
-
z_base = z_base + relative_position_encoding
|
| 646 |
if token_bonds_encoding is not None:
|
| 647 |
-
z_base = z_base + token_bonds_encoding
|
| 648 |
-
z_base = z_base + self.s_to_z(s_inputs_normed).unsqueeze(2)
|
| 649 |
-
z_base = z_base + self.s_to_z_transpose(s_inputs_normed).unsqueeze(1)
|
| 650 |
z_base = z_base + self.s_to_z_prod_out(
|
| 651 |
self.s_to_z_prod_in1(s_inputs_normed)[:, :, None, :]
|
| 652 |
* self.s_to_z_prod_in2(s_inputs_normed)[:, None, :, :]
|
| 653 |
-
)
|
| 654 |
-
|
| 655 |
-
pair = self._repeat_batch(z_base, num_diffusion_samples)
|
| 656 |
-
x_pred_flat = self._flatten_sample_axis(x_pred)
|
| 657 |
-
atom_to_token_m = self._repeat_batch(atom_to_token, num_diffusion_samples)
|
| 658 |
-
atom_mask_m = self._repeat_batch(atom_attention_mask, num_diffusion_samples)
|
| 659 |
-
rep_idx_m = self._repeat_batch(distogram_atom_idx, num_diffusion_samples).long()
|
| 660 |
-
mask = self._repeat_batch(token_attention_mask, num_diffusion_samples)
|
| 661 |
expanded_batch_size = pair.shape[0]
|
| 662 |
|
| 663 |
-
rep_coords = gather_rep_atom_coords(x_pred_flat, rep_idx_m)
|
| 664 |
rep_distances = torch.cdist(
|
| 665 |
rep_coords, rep_coords, compute_mode="donot_use_mm_for_euclid_dist"
|
| 666 |
-
)
|
| 667 |
-
distogram_bins = (rep_distances.unsqueeze(-1) > self.boundaries).sum(dim=-1).long()
|
| 668 |
-
pair = pair + self.dist_bin_pairwise_embed(distogram_bins)
|
| 669 |
|
| 670 |
-
pair_mask = mask[:, :, None].float() * mask[:, None, :].float()
|
| 671 |
|
| 672 |
# FoldingTrunk handles the bf16 cast internally during inference so
|
| 673 |
# each block's fused trimul engages. In-place residual avoids an
|
| 674 |
# extra fp32 pair allocation.
|
| 675 |
with torch.amp.autocast("cuda", enabled=pair.is_cuda, dtype=torch.bfloat16):
|
| 676 |
-
pair_delta = self.folding_trunk(pair, pair_attention_mask=pair_mask)
|
| 677 |
-
pair.add_(pair_delta.float())
|
| 678 |
del pair_delta
|
| 679 |
-
single = self.row_attention_pooling(pair, mask)
|
| 680 |
|
| 681 |
-
atom_mask_f = atom_mask_m.float()
|
| 682 |
-
s_at_atoms = gather_token_to_atom(single, atom_to_token_m)
|
| 683 |
-
s_at_atoms_ln = self.plddt_ln(s_at_atoms)
|
| 684 |
|
| 685 |
-
intra_idx = _compute_intra_token_idx(atom_to_token_m)
|
| 686 |
-
intra_idx = intra_idx.clamp(max=self.plddt_weight.shape[0] - 1)
|
| 687 |
-
w_plddt = self.plddt_weight[intra_idx]
|
| 688 |
-
plddt_logits = torch.einsum("...c,...cb->...b", s_at_atoms_ln, w_plddt)
|
| 689 |
-
plddt_per_atom = _categorical_mean(plddt_logits, start=0.0, end=1.0)
|
| 690 |
|
| 691 |
sequence_length = single.shape[1]
|
| 692 |
plddt_sum = torch.zeros(
|
|
@@ -694,80 +702,80 @@ class ConfidenceHead(nn.Module):
|
|
| 694 |
sequence_length,
|
| 695 |
device=single.device,
|
| 696 |
dtype=plddt_per_atom.dtype,
|
| 697 |
-
)
|
| 698 |
atom_count = torch.zeros(
|
| 699 |
expanded_batch_size,
|
| 700 |
sequence_length,
|
| 701 |
device=single.device,
|
| 702 |
dtype=plddt_per_atom.dtype,
|
| 703 |
-
)
|
| 704 |
-
atom_mask_t = atom_mask_f.to(plddt_per_atom.dtype)
|
| 705 |
-
plddt_sum.scatter_add_(1, atom_to_token_m, plddt_per_atom * atom_mask_t)
|
| 706 |
-
atom_count.scatter_add_(1, atom_to_token_m, atom_mask_t)
|
| 707 |
-
plddt = plddt_sum / atom_count.clamp(min=1e-6)
|
| 708 |
|
| 709 |
complex_plddt = (plddt_per_atom * atom_mask_f).sum(dim=-1) / (
|
| 710 |
atom_mask_f.sum(dim=-1) + _EPS
|
| 711 |
-
)
|
| 712 |
|
| 713 |
-
expanded_type = self._repeat_batch(mol_type, num_diffusion_samples)
|
| 714 |
-
expanded_asym = self._repeat_batch(asym_id, num_diffusion_samples)
|
| 715 |
-
is_ligand = (expanded_type == _NONPOLYMER_ID).float()
|
| 716 |
-
inter_chain = (expanded_asym.unsqueeze(-1) != expanded_asym.unsqueeze(-2)).float()
|
| 717 |
-
near_contact = (rep_distances < 8).float()
|
| 718 |
interface_per_token = (near_contact * inter_chain * (1.0 - is_ligand).unsqueeze(-1)).amax(
|
| 719 |
dim=-1
|
| 720 |
-
)
|
| 721 |
iplddt_weight = torch.where(
|
| 722 |
is_ligand.bool(),
|
| 723 |
torch.full_like(interface_per_token, 2.0),
|
| 724 |
interface_per_token,
|
| 725 |
-
)
|
| 726 |
iplddt_weight_atoms = gather_token_to_atom(
|
| 727 |
iplddt_weight.unsqueeze(-1), atom_to_token_m
|
| 728 |
-
).squeeze(-1)
|
| 729 |
-
atom_iplddt_w = atom_mask_f * iplddt_weight_atoms
|
| 730 |
complex_iplddt = (plddt_per_atom * atom_iplddt_w).sum(dim=-1) / (
|
| 731 |
atom_iplddt_w.sum(dim=-1) + _EPS
|
| 732 |
-
)
|
| 733 |
|
| 734 |
-
plddt_ca = plddt_per_atom.gather(1, rep_idx_m)
|
| 735 |
|
| 736 |
# PAE
|
| 737 |
-
pae_logits = self.pae_head(self.pae_ln(pair))
|
| 738 |
-
pae = _categorical_mean(pae_logits, start=0.0, end=32.0).detach()
|
| 739 |
|
| 740 |
# PDE
|
| 741 |
-
pde_logits = self.pde_head(self.pde_ln(pair))
|
| 742 |
-
pde = _categorical_mean(pde_logits, start=0.0, end=32.0).detach()
|
| 743 |
|
| 744 |
# Resolved (per-atom binary).
|
| 745 |
-
s_at_atoms_res = self.resolved_ln(s_at_atoms)
|
| 746 |
-
w_res = self.resolved_weight[intra_idx]
|
| 747 |
-
resolved_logits = torch.einsum("...c,...cb->...b", s_at_atoms_res, w_res)
|
| 748 |
|
| 749 |
# pTM / ipTM from pae_logits.
|
| 750 |
n_bins = pae_logits.shape[-1]
|
| 751 |
bin_width = 32.0 / n_bins
|
| 752 |
-
bin_centers = torch.arange(0.5 * bin_width, 32.0, bin_width, device=pae_logits.device)
|
| 753 |
-
mask_f = mask.float()
|
| 754 |
-
n_residues = mask_f.sum(dim=-1, keepdim=True)
|
| 755 |
-
d0 = 1.24 * (n_residues.clamp(min=19) - 15) ** (1 / 3) - 1.8
|
| 756 |
-
tm_per_bin = 1 / (1 + (bin_centers / d0) ** 2)
|
| 757 |
-
pae_probs = F.softmax(pae_logits, dim=-1)
|
| 758 |
-
tm_expected = (pae_probs * tm_per_bin[:, None, None, :]).sum(dim=-1)
|
| 759 |
-
|
| 760 |
-
pair_mask_2d = mask_f.unsqueeze(-1) * mask_f.unsqueeze(-2)
|
| 761 |
-
ptm_per_row = (tm_expected * pair_mask_2d).sum(dim=-1) / (pair_mask_2d.sum(dim=-1) + _EPS)
|
| 762 |
-
ptm = ptm_per_row.max(dim=-1).values
|
| 763 |
|
| 764 |
inter_chain_mask = (
|
| 765 |
expanded_asym.unsqueeze(-1) != expanded_asym.unsqueeze(-2)
|
| 766 |
-
).float() * pair_mask_2d
|
| 767 |
iptm_per_row = (tm_expected * inter_chain_mask).sum(dim=-1) / (
|
| 768 |
inter_chain_mask.sum(dim=-1) + _EPS
|
| 769 |
-
)
|
| 770 |
-
iptm = iptm_per_row.max(dim=-1).values
|
| 771 |
|
| 772 |
max_chain_id = int(expanded_asym.max().item()) if expanded_batch_size > 0 else 0
|
| 773 |
n_chains = max_chain_id + 1
|
|
@@ -777,16 +785,16 @@ class ConfidenceHead(nn.Module):
|
|
| 777 |
n_chains,
|
| 778 |
device=tm_expected.device,
|
| 779 |
dtype=tm_expected.dtype,
|
| 780 |
-
)
|
| 781 |
for c1 in range(n_chains):
|
| 782 |
-
chain_c1 = (expanded_asym == c1).float() * mask_f
|
| 783 |
if chain_c1.sum() == 0:
|
| 784 |
continue
|
| 785 |
for c2 in range(n_chains):
|
| 786 |
-
chain_c2 = (expanded_asym == c2).float() * mask_f
|
| 787 |
-
pair_m = chain_c1.unsqueeze(-1) * chain_c2.unsqueeze(-2)
|
| 788 |
-
denom = pair_m.sum(dim=(-1, -2)) + _EPS
|
| 789 |
-
pair_chains_iptm[:, c1, c2] = (tm_expected * pair_m).sum(dim=(-1, -2)) / denom
|
| 790 |
|
| 791 |
return {
|
| 792 |
"plddt_logits": plddt_logits,
|
|
@@ -803,7 +811,7 @@ class ConfidenceHead(nn.Module):
|
|
| 803 |
"ptm": ptm.detach(),
|
| 804 |
"iptm": iptm.detach(),
|
| 805 |
"pair_chains_iptm": pair_chains_iptm.detach(),
|
| 806 |
-
}
|
| 807 |
|
| 808 |
|
| 809 |
def _inverse_softplus(value: float) -> float:
|
|
@@ -834,9 +842,9 @@ def _convert_esmc_attention_outputs_to_te(module: nn.Module) -> tuple[str, ...]:
|
|
| 834 |
device=child.weight.device,
|
| 835 |
)
|
| 836 |
with torch.no_grad():
|
| 837 |
-
replacement.weight.copy_(child.weight)
|
| 838 |
if child.bias is not None:
|
| 839 |
-
replacement.bias.copy_(child.bias)
|
| 840 |
replacement.eval().requires_grad_(False)
|
| 841 |
setattr(owner, name, replacement)
|
| 842 |
converted.append(path)
|
|
@@ -953,15 +961,15 @@ class ESMFold2Model(
|
|
| 953 |
self.lm_encoder = None
|
| 954 |
|
| 955 |
self.parcae_input_norm = nn.LayerNorm(d_pair)
|
| 956 |
-
self.parcae_log_a = nn.Parameter(torch.zeros(d_pair))
|
| 957 |
parcae_decay_init = math.sqrt(1.0 / 5.0)
|
| 958 |
parcae_delta_init = -math.log(parcae_decay_init)
|
| 959 |
self.parcae_log_delta = nn.Parameter(
|
| 960 |
torch.full((d_pair,), _inverse_softplus(parcae_delta_init), dtype=torch.float32)
|
| 961 |
-
)
|
| 962 |
-
self.parcae_b_cont = nn.Parameter(torch.eye(d_pair))
|
| 963 |
self.parcae_readout = nn.Linear(d_pair, d_pair, bias=False)
|
| 964 |
-
nn.init.eye_(self.parcae_readout.weight)
|
| 965 |
self.parcae_coda = FoldingTrunk(
|
| 966 |
n_layers=config.parcae.coda_n_layers, d_pair=d_pair, expansion_ratio=4
|
| 967 |
)
|
|
@@ -1101,9 +1109,10 @@ class ESMFold2Model(
|
|
| 1101 |
input_ids: torch.Tensor | None = None,
|
| 1102 |
**kwargs,
|
| 1103 |
) -> torch.Tensor:
|
|
|
|
| 1104 |
del kwargs
|
| 1105 |
if input_ids is not None:
|
| 1106 |
-
return input_ids
|
| 1107 |
if seq is None:
|
| 1108 |
raise ValueError("Pass either seq or input_ids for ESMFold2 TTT.")
|
| 1109 |
sequences = [seq] if isinstance(seq, str) else seq
|
|
@@ -1122,13 +1131,13 @@ class ESMFold2Model(
|
|
| 1122 |
(len(encoded), max_len),
|
| 1123 |
SEQUENCE_PAD_TOKEN,
|
| 1124 |
dtype=torch.long,
|
| 1125 |
-
)
|
| 1126 |
for row, token_ids in enumerate(encoded):
|
| 1127 |
input_tensor[row, : len(token_ids)] = torch.tensor(
|
| 1128 |
token_ids,
|
| 1129 |
dtype=torch.long,
|
| 1130 |
-
)
|
| 1131 |
-
return input_tensor
|
| 1132 |
|
| 1133 |
def _ttt_mask_token(self) -> int:
|
| 1134 |
return SEQUENCE_MASK_TOKEN
|
|
@@ -1142,18 +1151,20 @@ class ESMFold2Model(
|
|
| 1142 |
SEQUENCE_STANDARD_AA_MAX_TOKEN,
|
| 1143 |
device=input_ids.device,
|
| 1144 |
dtype=input_ids.dtype,
|
| 1145 |
-
)
|
| 1146 |
|
| 1147 |
def _ttt_non_special_mask(self, input_ids: torch.Tensor) -> torch.Tensor:
|
|
|
|
| 1148 |
return (input_ids >= SEQUENCE_STANDARD_AA_MIN_TOKEN) & (
|
| 1149 |
input_ids < SEQUENCE_STANDARD_AA_MAX_TOKEN
|
| 1150 |
-
)
|
| 1151 |
|
| 1152 |
def _ttt_predict_logits(
|
| 1153 |
self,
|
| 1154 |
batch: torch.Tensor | dict[str, torch.Tensor],
|
| 1155 |
**kwargs,
|
| 1156 |
) -> torch.Tensor:
|
|
|
|
| 1157 |
del kwargs
|
| 1158 |
if not isinstance(batch, torch.Tensor):
|
| 1159 |
raise TypeError("ESMFold2 TTT expects input_ids tensors.")
|
|
@@ -1163,14 +1174,14 @@ class ESMFold2Model(
|
|
| 1163 |
self._ensure_ttt_lm_head()
|
| 1164 |
if self._ttt_lm_head is None:
|
| 1165 |
raise RuntimeError("ESMFold2 TTT MLM head initialization failed.")
|
| 1166 |
-
attention_mask = batch.ne(SEQUENCE_PAD_TOKEN)
|
| 1167 |
output = self._esmc(
|
| 1168 |
input_ids=batch,
|
| 1169 |
attention_mask=attention_mask,
|
| 1170 |
return_dict=True,
|
| 1171 |
compute_sae=False,
|
| 1172 |
)
|
| 1173 |
-
return self._ttt_lm_head(output.last_hidden_state)
|
| 1174 |
|
| 1175 |
@classmethod
|
| 1176 |
def from_pretrained(
|
|
@@ -1287,6 +1298,7 @@ class ESMFold2Model(
|
|
| 1287 |
lm_mask_pct: float = 0.0,
|
| 1288 |
verbose: bool = False,
|
| 1289 |
) -> Tensor:
|
|
|
|
| 1290 |
if self._esmc_fp8 and torch.is_grad_enabled():
|
| 1291 |
_reload_esmc_bf16_for_gradients(
|
| 1292 |
self,
|
|
@@ -1312,9 +1324,9 @@ class ESMFold2Model(
|
|
| 1312 |
pad_to_multiple=pad_to,
|
| 1313 |
lm_mask_pct=lm_mask_pct,
|
| 1314 |
mask_token_id=SEQUENCE_MASK_TOKEN,
|
| 1315 |
-
)
|
| 1316 |
progress.update()
|
| 1317 |
-
return result
|
| 1318 |
return compute_lm_hidden_states(
|
| 1319 |
self._esmc,
|
| 1320 |
input_ids,
|
|
@@ -1325,19 +1337,20 @@ class ESMFold2Model(
|
|
| 1325 |
pad_to_multiple=pad_to,
|
| 1326 |
lm_mask_pct=lm_mask_pct,
|
| 1327 |
mask_token_id=SEQUENCE_MASK_TOKEN,
|
| 1328 |
-
)
|
| 1329 |
|
| 1330 |
def _discretized_dynamics(self) -> tuple[Tensor, Tensor]:
|
| 1331 |
-
delta = F.softplus(self.parcae_log_delta)
|
| 1332 |
-
a = torch.exp(-delta * torch.exp(self.parcae_log_a))
|
| 1333 |
-
b = delta[:, None] * self.parcae_b_cont
|
| 1334 |
-
return a, b
|
| 1335 |
|
| 1336 |
def _init_pair_state(self, ref: Tensor) -> Tensor:
|
|
|
|
| 1337 |
std = math.sqrt(2.0 / (5.0 * ref.shape[-1]))
|
| 1338 |
-
state = torch.empty_like(ref, dtype=torch.float32)
|
| 1339 |
-
nn.init.trunc_normal_(state, mean=0.0, std=std, a=-3 * std, b=3 * std)
|
| 1340 |
-
return state.to(dtype=ref.dtype)
|
| 1341 |
|
| 1342 |
def _run_one_loop(
|
| 1343 |
self,
|
|
@@ -1356,6 +1369,7 @@ class ESMFold2Model(
|
|
| 1356 |
# otherwise leaks about 2 GB of l^2 * c_z data into distogram/sample scope.
|
| 1357 |
# training=True forces dropout under eval(), matching the per-loop
|
| 1358 |
# dropout strategy used at train time.
|
|
|
|
| 1359 |
lm_cfg = self.config.lm_encoder
|
| 1360 |
_per_loop_lm_dropout = (
|
| 1361 |
lm_z is not None
|
|
@@ -1377,19 +1391,19 @@ class ESMFold2Model(
|
|
| 1377 |
if _per_loop_lm_dropout:
|
| 1378 |
if lm_z is None:
|
| 1379 |
raise RuntimeError("Per-loop LM dropout requires LM pair features.")
|
| 1380 |
-
lm_z_i: Tensor | None = F.dropout(lm_z, p=_lm_dropout_p, training=True)
|
| 1381 |
else:
|
| 1382 |
-
lm_z_i = lm_z
|
| 1383 |
|
| 1384 |
-
refined_lm_z: Tensor | None = None
|
| 1385 |
if lm_z_i is not None and self.lm_encoder is not None:
|
| 1386 |
refined_lm_z = self.lm_encoder(
|
| 1387 |
lm_z_i.to(z_init.dtype), pair_attention_mask=pair_mask
|
| 1388 |
-
)
|
| 1389 |
|
| 1390 |
-
z_inject_pair = z_init
|
| 1391 |
if lm_z_i is not None and self.lm_encoder is None:
|
| 1392 |
-
z_inject_pair = z_inject_pair + lm_z_i.to(z_inject_pair.dtype)
|
| 1393 |
|
| 1394 |
if self.msa_encoder is not None and _msa_inputs is not None:
|
| 1395 |
msa_i, mask_i, hd_i, dv_i = maybe_subsample_msa(
|
|
@@ -1399,26 +1413,26 @@ class ESMFold2Model(
|
|
| 1399 |
_msa_inputs["deletion_value"],
|
| 1400 |
max_depth=_msa_inputs["max_depth"],
|
| 1401 |
enabled=_msa_inputs["subsample_enabled"],
|
| 1402 |
-
)
|
| 1403 |
b_msa, m, l_msa = msa_i.shape
|
| 1404 |
-
msa_oh = F.one_hot(msa_i.permute(0, 2, 1).long(), num_classes=NUM_RES_TYPES).float()
|
| 1405 |
msa_attn = (
|
| 1406 |
mask_i.permute(0, 2, 1).float()
|
| 1407 |
if mask_i is not None
|
| 1408 |
else tok_mask[:, :, None].expand(-1, -1, m).float()
|
| 1409 |
-
)
|
| 1410 |
# Bias-free MSAEncoder.embed requires zeroed padding.
|
| 1411 |
-
msa_oh = msa_oh * msa_attn.unsqueeze(-1)
|
| 1412 |
hd = (
|
| 1413 |
hd_i.permute(0, 2, 1).float()
|
| 1414 |
if hd_i is not None
|
| 1415 |
else torch.zeros(b_msa, l_msa, m, device=msa_i.device)
|
| 1416 |
-
)
|
| 1417 |
dv = (
|
| 1418 |
dv_i.permute(0, 2, 1).float()
|
| 1419 |
if dv_i is not None
|
| 1420 |
else torch.zeros(b_msa, l_msa, m, device=msa_i.device)
|
| 1421 |
-
)
|
| 1422 |
msa_pair = self.msa_encoder(
|
| 1423 |
x_pair=z_inject_pair,
|
| 1424 |
x_inputs=_msa_inputs["x_inputs"],
|
|
@@ -1426,19 +1440,19 @@ class ESMFold2Model(
|
|
| 1426 |
has_deletion=hd,
|
| 1427 |
deletion_value=dv,
|
| 1428 |
msa_attention_mask=msa_attn,
|
| 1429 |
-
).to(z_inject_pair.dtype)
|
| 1430 |
z_inject_pair = (
|
| 1431 |
msa_pair if self.config.msa_encoder_overwrite else (z_inject_pair + msa_pair)
|
| 1432 |
-
)
|
| 1433 |
|
| 1434 |
if refined_lm_z is not None:
|
| 1435 |
-
z_inject_pair = z_inject_pair + refined_lm_z.to(z_inject_pair.dtype)
|
| 1436 |
|
| 1437 |
-
injected_pair = self.parcae_input_norm(z_inject_pair)
|
| 1438 |
-
z = a * z + F.linear(injected_pair.to(z.dtype), b_mat)
|
| 1439 |
-
z = self.folding_trunk(z, pair_attention_mask=pair_mask)
|
| 1440 |
|
| 1441 |
-
return z
|
| 1442 |
|
| 1443 |
def forward(
|
| 1444 |
self,
|
|
@@ -1488,6 +1502,7 @@ class ESMFold2Model(
|
|
| 1488 |
disto_cond_mask: Tensor | None = None,
|
| 1489 |
verbose: bool = False,
|
| 1490 |
) -> ESMFold2Output | tuple[Any, ...]:
|
|
|
|
| 1491 |
output_hidden_states, return_dict = _resolve_structure_output_controls(
|
| 1492 |
self.config,
|
| 1493 |
output_attentions=output_attentions,
|
|
@@ -1508,9 +1523,9 @@ class ESMFold2Model(
|
|
| 1508 |
disto_cond_mask=disto_cond_mask,
|
| 1509 |
)
|
| 1510 |
del gt_coords, is_resolved, frames_idx
|
| 1511 |
-
tok_mask = token_attention_mask
|
| 1512 |
-
atm_mask = atom_attention_mask
|
| 1513 |
-
disto_idx = distogram_atom_idx
|
| 1514 |
|
| 1515 |
n_loops: int = num_loops if num_loops is not None else self.config.num_loops
|
| 1516 |
n_samples: int = (
|
|
@@ -1521,37 +1536,37 @@ class ESMFold2Model(
|
|
| 1521 |
total_steps = max(1, n_loops + 1)
|
| 1522 |
|
| 1523 |
if res_type.dim() == 2:
|
| 1524 |
-
res_type_oh = F.one_hot(res_type.long(), num_classes=NUM_RES_TYPES).float()
|
| 1525 |
-
res_type_oh = res_type_oh * tok_mask.unsqueeze(-1).float()
|
| 1526 |
else:
|
| 1527 |
-
res_type_oh = res_type.float()
|
| 1528 |
|
| 1529 |
if msa is not None:
|
| 1530 |
-
msa_oh_profile = F.one_hot(msa.long(), num_classes=NUM_RES_TYPES).float()
|
| 1531 |
if msa_attention_mask is not None:
|
| 1532 |
-
mask_f = msa_attention_mask.float().unsqueeze(-1)
|
| 1533 |
-
msa_oh_profile = msa_oh_profile * mask_f
|
| 1534 |
-
valid_seq_count = msa_attention_mask.float().sum(dim=1).clamp(min=1)
|
| 1535 |
-
profile = msa_oh_profile.sum(dim=1) / valid_seq_count.unsqueeze(-1)
|
| 1536 |
else:
|
| 1537 |
-
profile = msa_oh_profile.mean(dim=1)
|
| 1538 |
else:
|
| 1539 |
-
profile = res_type_oh
|
| 1540 |
|
| 1541 |
if deletion_mean is None:
|
| 1542 |
deletion_mean = torch.zeros(
|
| 1543 |
res_type.shape[0], res_type.shape[1], device=res_type.device
|
| 1544 |
-
)
|
| 1545 |
|
| 1546 |
-
ref_element_oh = F.one_hot(ref_element.long(), num_classes=MAX_ATOMIC_NUMBER).float()
|
| 1547 |
ref_atom_name_chars_oh = F.one_hot(
|
| 1548 |
ref_atom_name_chars.long(), num_classes=CHAR_VOCAB_SIZE
|
| 1549 |
-
).float()
|
| 1550 |
# Bias-free downstream Linears require zeroed padding.
|
| 1551 |
-
atm_mask_f = atm_mask.float()
|
| 1552 |
-
ref_element_oh = ref_element_oh * atm_mask_f.unsqueeze(-1)
|
| 1553 |
-
ref_atom_name_chars_oh = ref_atom_name_chars_oh * atm_mask_f.unsqueeze(-1).unsqueeze(-1)
|
| 1554 |
-
atom_to_token = atom_to_token * atm_mask.long()
|
| 1555 |
|
| 1556 |
use_amp = ref_pos.device.type == "cuda"
|
| 1557 |
with torch.amp.autocast("cuda", enabled=use_amp, dtype=torch.bfloat16):
|
|
@@ -1566,9 +1581,9 @@ class ESMFold2Model(
|
|
| 1566 |
ref_element=ref_element_oh,
|
| 1567 |
ref_atom_name_chars=ref_atom_name_chars_oh,
|
| 1568 |
atom_to_token=atom_to_token,
|
| 1569 |
-
)
|
| 1570 |
|
| 1571 |
-
z_init = self.z_init_1(x_inputs).unsqueeze(2) + self.z_init_2(x_inputs).unsqueeze(1)
|
| 1572 |
|
| 1573 |
relative_position_encoding = self.rel_pos(
|
| 1574 |
residue_index=residue_index,
|
|
@@ -1576,9 +1591,9 @@ class ESMFold2Model(
|
|
| 1576 |
sym_id=sym_id,
|
| 1577 |
entity_id=entity_id,
|
| 1578 |
token_index=token_index,
|
| 1579 |
-
)
|
| 1580 |
-
token_bonds_encoding = self.token_bonds(token_bonds.float())
|
| 1581 |
-
z_init = z_init + relative_position_encoding + token_bonds_encoding
|
| 1582 |
|
| 1583 |
if lm_hidden_states is None and input_ids is not None and self._esmc is not None:
|
| 1584 |
lm_hidden_states = self._compute_lm_hidden_states(
|
|
@@ -1589,26 +1604,26 @@ class ESMFold2Model(
|
|
| 1589 |
tok_mask,
|
| 1590 |
lm_mask_pct=(self.config.lm_mask_pct if lm_mask_pct is None else lm_mask_pct),
|
| 1591 |
verbose=verbose,
|
| 1592 |
-
)
|
| 1593 |
-
lm_z: Tensor | None = None
|
| 1594 |
if lm_hidden_states is not None:
|
| 1595 |
-
lm_z = self.language_model(lm_hidden_states.detach())
|
| 1596 |
del lm_hidden_states
|
| 1597 |
|
| 1598 |
-
pair_mask = tok_mask[:, :, None].float() * tok_mask[:, None, :].float()
|
| 1599 |
|
| 1600 |
-
z = self._init_pair_state(z_init)
|
| 1601 |
|
| 1602 |
-
a, b = self._discretized_dynamics()
|
| 1603 |
-
a = a.view(1, 1, 1, -1).to(device=z.device, dtype=z.dtype)
|
| 1604 |
-
b_mat = b.to(device=z.device, dtype=z.dtype)
|
| 1605 |
|
| 1606 |
_msa_inputs: dict | None = None
|
| 1607 |
if self.msa_encoder is not None and msa is not None:
|
| 1608 |
msa_attention_mask = maybe_apply_msa_column_masking(
|
| 1609 |
msa_attention_mask,
|
| 1610 |
msa_column_mask_rate,
|
| 1611 |
-
)
|
| 1612 |
_msa_inputs = dict(
|
| 1613 |
x_inputs=x_inputs,
|
| 1614 |
msa=msa,
|
|
@@ -1631,14 +1646,14 @@ class ESMFold2Model(
|
|
| 1631 |
tok_mask=tok_mask,
|
| 1632 |
total_steps=total_steps,
|
| 1633 |
verbose=verbose,
|
| 1634 |
-
)
|
| 1635 |
del z_init, lm_z, _msa_inputs, a, b_mat
|
| 1636 |
|
| 1637 |
-
z = self.parcae_readout(z)
|
| 1638 |
-
z = self.parcae_coda(z, pair_attention_mask=pair_mask)
|
| 1639 |
|
| 1640 |
-
z = z.float()
|
| 1641 |
-
distogram_logits = self.distogram_head(z + z.transpose(-2, -3))
|
| 1642 |
|
| 1643 |
structure_output = self.structure_head.sample(
|
| 1644 |
z_trunk=z,
|
|
@@ -1666,13 +1681,13 @@ class ESMFold2Model(
|
|
| 1666 |
return_atom_repr=False,
|
| 1667 |
denoising_early_exit_rmsd=(0.10 if early_exit else None),
|
| 1668 |
verbose=verbose,
|
| 1669 |
-
)
|
| 1670 |
|
| 1671 |
-
sample_coords = structure_output["sample_atom_coords"]
|
| 1672 |
if sample_coords is None:
|
| 1673 |
raise RuntimeError("ESMFold2 structure sampling did not return coordinates.")
|
| 1674 |
output: dict[str, Tensor] = {"distogram_logits": distogram_logits}
|
| 1675 |
-
output["sample_atom_coords"] = sample_coords
|
| 1676 |
|
| 1677 |
confidence_config = self.config.confidence_head
|
| 1678 |
confidence_enabled = (
|
|
@@ -1697,7 +1712,7 @@ class ESMFold2Model(
|
|
| 1697 |
num_diffusion_samples=n_samples,
|
| 1698 |
relative_position_encoding=relative_position_encoding.detach(),
|
| 1699 |
token_bonds_encoding=token_bonds_encoding.detach(),
|
| 1700 |
-
)
|
| 1701 |
progress.update()
|
| 1702 |
else:
|
| 1703 |
confidence_output = self.confidence_head(
|
|
@@ -1713,18 +1728,18 @@ class ESMFold2Model(
|
|
| 1713 |
num_diffusion_samples=n_samples,
|
| 1714 |
relative_position_encoding=relative_position_encoding.detach(),
|
| 1715 |
token_bonds_encoding=token_bonds_encoding.detach(),
|
| 1716 |
-
)
|
| 1717 |
output.update(confidence_output)
|
| 1718 |
-
output["atom_pad_mask"] = atm_mask.unsqueeze(0) if atm_mask.dim() == 1 else atm_mask
|
| 1719 |
-
output["residue_index"] = residue_index
|
| 1720 |
-
output["entity_id"] = entity_id
|
| 1721 |
return _finalize_structure_output(
|
| 1722 |
output,
|
| 1723 |
token_input_state=x_inputs,
|
| 1724 |
pair_state=z,
|
| 1725 |
output_hidden_states=output_hidden_states,
|
| 1726 |
return_dict=return_dict,
|
| 1727 |
-
)
|
| 1728 |
|
| 1729 |
@torch.no_grad()
|
| 1730 |
def infer_protein(self, seq: str, **forward_kwargs) -> ESMFold2Output:
|
|
@@ -1743,7 +1758,7 @@ class ESMFold2Model(
|
|
| 1743 |
if not self.config.msa_conditioning:
|
| 1744 |
for name in MSA_CONDITIONING_INPUT_NAMES:
|
| 1745 |
features.pop(name, None)
|
| 1746 |
-
features = {k: v.to(self.device) for k, v in features.items()}
|
| 1747 |
return self(**features, **forward_kwargs, return_dict=True)
|
| 1748 |
|
| 1749 |
@property
|
|
@@ -2056,14 +2071,15 @@ class MSAEncoderBlock(nn.Module):
|
|
| 2056 |
msa_attention_mask: Tensor,
|
| 2057 |
pair_attention_mask: Tensor,
|
| 2058 |
) -> tuple[Tensor, Tensor]:
|
| 2059 |
-
|
|
|
|
| 2060 |
if not self.is_final_block:
|
| 2061 |
-
m = m + self.msa_pair_weighted_averaging(m, pair, pair_attention_mask)
|
| 2062 |
-
m = m + self.msa_transition(m)
|
| 2063 |
-
pair = pair + self.tri_mul_out(pair, mask=pair_attention_mask)
|
| 2064 |
-
pair = pair + self.tri_mul_in(pair, mask=pair_attention_mask)
|
| 2065 |
-
pair = pair + self.pair_transition(pair)
|
| 2066 |
-
return m, pair
|
| 2067 |
|
| 2068 |
|
| 2069 |
class MSAEncoder(nn.Module):
|
|
@@ -2110,12 +2126,13 @@ class MSAEncoder(nn.Module):
|
|
| 2110 |
msa_attention_mask: Tensor,
|
| 2111 |
) -> Tensor:
|
| 2112 |
# Every input tensor is pre-transposed to shape (b, l, m, ...) before this call.
|
|
|
|
| 2113 |
m_feat = torch.cat(
|
| 2114 |
[msa_oh, has_deletion.unsqueeze(-1), deletion_value.unsqueeze(-1)], dim=-1
|
| 2115 |
-
)
|
| 2116 |
-
m = self.embed(m_feat) + self.project_inputs(x_inputs).unsqueeze(2)
|
| 2117 |
-
tok_mask = msa_attention_mask[:, :, 0].bool()
|
| 2118 |
-
pair_attention_mask = tok_mask.unsqueeze(2) & tok_mask.unsqueeze(1)
|
| 2119 |
for block in self.blocks:
|
| 2120 |
-
m, x_pair = block(m, x_pair, msa_attention_mask, pair_attention_mask)
|
| 2121 |
-
return x_pair
|
|
|
|
| 2 |
|
| 3 |
Quickstart::
|
| 4 |
|
| 5 |
+
from pathlib import Path
|
| 6 |
+
from transformers import AutoModel
|
| 7 |
|
| 8 |
+
model = AutoModel.from_pretrained("Synthyra/ESMFold2", trust_remote_code=True).cuda().eval()
|
| 9 |
+
structure = model.infer_protein_as_pdb("MQIFVKTLTGKT")
|
| 10 |
+
Path("structure.pdb").write_text(structure, encoding="utf-8")
|
| 11 |
|
| 12 |
For multi-chain, ligand, and MSA inputs, use ``model.input_types`` together
|
| 13 |
with ``model.fold(...)`` or ``model.prepare_structure_input(...)``.
|
|
|
|
| 19 |
import importlib
|
| 20 |
import importlib.metadata
|
| 21 |
import math
|
| 22 |
+
import torch
|
| 23 |
+
import torch.nn as nn
|
| 24 |
+
import torch.nn.functional as F
|
| 25 |
+
|
| 26 |
from collections.abc import Mapping
|
| 27 |
from contextlib import contextmanager
|
| 28 |
from dataclasses import asdict, dataclass
|
| 29 |
from pathlib import Path
|
| 30 |
from typing import Any, ClassVar, Literal, cast
|
|
|
|
|
|
|
|
|
|
|
|
|
| 31 |
from torch import Tensor
|
| 32 |
from tqdm.auto import tqdm
|
| 33 |
from transformers.modeling_outputs import ModelOutput
|
|
|
|
| 35 |
|
| 36 |
from ...attention import get_attn_implementation, set_config_attn_implementation
|
| 37 |
|
| 38 |
+
|
| 39 |
try:
|
| 40 |
from fastplms.models.ttt import FastPLMTestTimeTrainingMixin, TTTConfig
|
| 41 |
except ModuleNotFoundError as error:
|
|
|
|
| 198 |
compute_sae: bool = True,
|
| 199 |
normalize_sae: bool = False,
|
| 200 |
):
|
| 201 |
+
# input_ids and optional masks/sequence IDs: (b, t).
|
| 202 |
del return_dict, compute_sae, normalize_sae
|
| 203 |
output = self.model(
|
| 204 |
input_ids=input_ids,
|
|
|
|
| 210 |
esmfold2_hidden_states=True,
|
| 211 |
)
|
| 212 |
if output_hidden_states:
|
| 213 |
+
hidden_states = output.hidden_states # Tensor (n_states, b, t, d_lm), or sequence of (b, t, d_lm) tensors
|
| 214 |
if hidden_states is None:
|
| 215 |
raise RuntimeError("ESM++ did not return requested hidden states.")
|
| 216 |
if isinstance(hidden_states, torch.Tensor):
|
| 217 |
+
output.hidden_states = hidden_states # (n_states, b, t, d_lm)
|
| 218 |
else:
|
| 219 |
+
output.hidden_states = torch.stack(tuple(hidden_states), dim=0) # (n_states, b, t, d_lm)
|
| 220 |
+
return output # model output; hidden_states stacked on the leading state axis when requested
|
| 221 |
|
| 222 |
|
| 223 |
def _load_fastplms_esmplusplus_for_esmfold2(
|
|
|
|
| 551 |
self._chunk_size = chunk_size
|
| 552 |
|
| 553 |
def forward(self, x: Tensor) -> Tensor:
|
| 554 |
+
# x: (b, l, ..., d_model); l_c is the current chunk width.
|
| 555 |
if self._chunk_size is None or x.shape[1] <= self._chunk_size:
|
| 556 |
+
return self.ffn(self.norm(x)) # x.shape
|
| 557 |
out: list[Tensor] = []
|
| 558 |
for s in range(0, x.shape[1], self._chunk_size):
|
| 559 |
e = min(s + self._chunk_size, x.shape[1])
|
| 560 |
+
sl = x[:, s:e] # (b, l_c, ..., d_model)
|
| 561 |
out.append(self.ffn(self.norm(sl)))
|
| 562 |
+
return torch.cat(out, dim=1) # x.shape
|
| 563 |
|
| 564 |
|
| 565 |
class ConfidenceHead(nn.Module):
|
|
|
|
| 574 |
d_pair = config.d_pair
|
| 575 |
d_inputs = config.inputs.d_inputs
|
| 576 |
|
| 577 |
+
boundaries = torch.linspace(ch.min_dist, ch.max_dist, ch.distogram_bins - 1) # (distogram_bins - 1,)
|
| 578 |
+
self.register_buffer("boundaries", boundaries) # (distogram_bins - 1,)
|
| 579 |
self.dist_bin_pairwise_embed = nn.Embedding(ch.distogram_bins, d_pair)
|
| 580 |
|
| 581 |
self.s_norm = nn.LayerNorm(d_single)
|
|
|
|
| 599 |
max_atoms_per_token = 23
|
| 600 |
self.plddt_weight = nn.Parameter(
|
| 601 |
torch.zeros(max_atoms_per_token, d_single, ch.num_plddt_bins)
|
| 602 |
+
) # (23, d_single, n_plddt_bins)
|
| 603 |
|
| 604 |
self.pae_ln = nn.LayerNorm(d_pair)
|
| 605 |
self.pae_head = nn.Linear(d_pair, ch.num_pae_bins, bias=False)
|
|
|
|
| 609 |
|
| 610 |
self.resolved_ln = nn.LayerNorm(d_single)
|
| 611 |
# 2 = resolved logits ([unresolved, resolved]).
|
| 612 |
+
self.resolved_weight = nn.Parameter(torch.zeros(max_atoms_per_token, d_single, 2)) # (23, d_single, 2)
|
| 613 |
|
| 614 |
def set_kernel_backend(self, backend: str | None) -> None:
|
| 615 |
self.folding_trunk.set_kernel_backend(backend)
|
|
|
|
| 619 |
|
| 620 |
@staticmethod
|
| 621 |
def _repeat_batch(x: Tensor, num_diffusion_samples: int) -> Tensor:
|
| 622 |
+
# x: (b, ...); output repeats the batch axis by samples.
|
| 623 |
+
return x if num_diffusion_samples == 1 else x.repeat_interleave(num_diffusion_samples, 0) # (b * samples, ...), including samples = 1
|
| 624 |
|
| 625 |
@staticmethod
|
| 626 |
def _flatten_sample_axis(x: Tensor) -> Tensor:
|
| 627 |
+
# x: (b, samples, n, c) or an already flattened tensor.
|
| 628 |
if x.ndim == 4:
|
| 629 |
b, mult, n, c = x.shape
|
| 630 |
+
return x.reshape(b * mult, n, c) # (b * samples, n, c) for 4D input; otherwise x.shape
|
| 631 |
+
return x # (b * samples, n, c) for 4D input; otherwise x.shape
|
| 632 |
|
| 633 |
def forward(
|
| 634 |
self,
|
|
|
|
| 645 |
relative_position_encoding: Tensor | None = None,
|
| 646 |
token_bonds_encoding: Tensor | None = None,
|
| 647 |
) -> dict[str, Tensor]:
|
| 648 |
+
# s_inputs: (b, l, d_inputs); z: (b, l, l, d_pair); x_pred: (bs, a, 3) or (b, samples, a, 3). bs = b * samples.
|
| 649 |
+
s_inputs_normed = self.s_inputs_norm(s_inputs) # (b, l, d_inputs)
|
| 650 |
|
| 651 |
+
z_base = self.z_norm(z) # (b, l, l, d_pair)
|
| 652 |
if relative_position_encoding is not None:
|
| 653 |
+
z_base = z_base + relative_position_encoding # (b, l, l, d_pair)
|
| 654 |
if token_bonds_encoding is not None:
|
| 655 |
+
z_base = z_base + token_bonds_encoding # (b, l, l, d_pair)
|
| 656 |
+
z_base = z_base + self.s_to_z(s_inputs_normed).unsqueeze(2) # (b, l, l, d_pair)
|
| 657 |
+
z_base = z_base + self.s_to_z_transpose(s_inputs_normed).unsqueeze(1) # (b, l, l, d_pair)
|
| 658 |
z_base = z_base + self.s_to_z_prod_out(
|
| 659 |
self.s_to_z_prod_in1(s_inputs_normed)[:, :, None, :]
|
| 660 |
* self.s_to_z_prod_in2(s_inputs_normed)[:, None, :, :]
|
| 661 |
+
) # (b, l, l, d_pair)
|
| 662 |
+
|
| 663 |
+
pair = self._repeat_batch(z_base, num_diffusion_samples) # (bs, l, l, d_pair)
|
| 664 |
+
x_pred_flat = self._flatten_sample_axis(x_pred) # (bs, a, 3)
|
| 665 |
+
atom_to_token_m = self._repeat_batch(atom_to_token, num_diffusion_samples) # (bs, a)
|
| 666 |
+
atom_mask_m = self._repeat_batch(atom_attention_mask, num_diffusion_samples) # (bs, a)
|
| 667 |
+
rep_idx_m = self._repeat_batch(distogram_atom_idx, num_diffusion_samples).long() # (bs, l)
|
| 668 |
+
mask = self._repeat_batch(token_attention_mask, num_diffusion_samples) # (bs, l)
|
| 669 |
expanded_batch_size = pair.shape[0]
|
| 670 |
|
| 671 |
+
rep_coords = gather_rep_atom_coords(x_pred_flat, rep_idx_m) # (bs, l, 3)
|
| 672 |
rep_distances = torch.cdist(
|
| 673 |
rep_coords, rep_coords, compute_mode="donot_use_mm_for_euclid_dist"
|
| 674 |
+
) # (bs, l, l)
|
| 675 |
+
distogram_bins = (rep_distances.unsqueeze(-1) > self.boundaries).sum(dim=-1).long() # (bs, l, l)
|
| 676 |
+
pair = pair + self.dist_bin_pairwise_embed(distogram_bins) # (bs, l, l, d_pair)
|
| 677 |
|
| 678 |
+
pair_mask = mask[:, :, None].float() * mask[:, None, :].float() # (bs, l, l)
|
| 679 |
|
| 680 |
# FoldingTrunk handles the bf16 cast internally during inference so
|
| 681 |
# each block's fused trimul engages. In-place residual avoids an
|
| 682 |
# extra fp32 pair allocation.
|
| 683 |
with torch.amp.autocast("cuda", enabled=pair.is_cuda, dtype=torch.bfloat16):
|
| 684 |
+
pair_delta = self.folding_trunk(pair, pair_attention_mask=pair_mask) # (bs, l, l, d_pair)
|
| 685 |
+
pair.add_(pair_delta.float()) # (bs, l, l, d_pair)
|
| 686 |
del pair_delta
|
| 687 |
+
single = self.row_attention_pooling(pair, mask) # (bs, l, d_single)
|
| 688 |
|
| 689 |
+
atom_mask_f = atom_mask_m.float() # (bs, a)
|
| 690 |
+
s_at_atoms = gather_token_to_atom(single, atom_to_token_m) # (bs, a, d_single)
|
| 691 |
+
s_at_atoms_ln = self.plddt_ln(s_at_atoms) # (bs, a, d_single)
|
| 692 |
|
| 693 |
+
intra_idx = _compute_intra_token_idx(atom_to_token_m) # (bs, a)
|
| 694 |
+
intra_idx = intra_idx.clamp(max=self.plddt_weight.shape[0] - 1) # (bs, a)
|
| 695 |
+
w_plddt = self.plddt_weight[intra_idx] # (bs, a, d_single, n_plddt_bins)
|
| 696 |
+
plddt_logits = torch.einsum("...c,...cb->...b", s_at_atoms_ln, w_plddt) # (bs, a, n_plddt_bins)
|
| 697 |
+
plddt_per_atom = _categorical_mean(plddt_logits, start=0.0, end=1.0) # (bs, a)
|
| 698 |
|
| 699 |
sequence_length = single.shape[1]
|
| 700 |
plddt_sum = torch.zeros(
|
|
|
|
| 702 |
sequence_length,
|
| 703 |
device=single.device,
|
| 704 |
dtype=plddt_per_atom.dtype,
|
| 705 |
+
) # (bs, l)
|
| 706 |
atom_count = torch.zeros(
|
| 707 |
expanded_batch_size,
|
| 708 |
sequence_length,
|
| 709 |
device=single.device,
|
| 710 |
dtype=plddt_per_atom.dtype,
|
| 711 |
+
) # (bs, l)
|
| 712 |
+
atom_mask_t = atom_mask_f.to(plddt_per_atom.dtype) # (bs, a)
|
| 713 |
+
plddt_sum.scatter_add_(1, atom_to_token_m, plddt_per_atom * atom_mask_t) # (bs, l)
|
| 714 |
+
atom_count.scatter_add_(1, atom_to_token_m, atom_mask_t) # (bs, l)
|
| 715 |
+
plddt = plddt_sum / atom_count.clamp(min=1e-6) # (bs, l)
|
| 716 |
|
| 717 |
complex_plddt = (plddt_per_atom * atom_mask_f).sum(dim=-1) / (
|
| 718 |
atom_mask_f.sum(dim=-1) + _EPS
|
| 719 |
+
) # (bs,)
|
| 720 |
|
| 721 |
+
expanded_type = self._repeat_batch(mol_type, num_diffusion_samples) # (bs, l)
|
| 722 |
+
expanded_asym = self._repeat_batch(asym_id, num_diffusion_samples) # (bs, l)
|
| 723 |
+
is_ligand = (expanded_type == _NONPOLYMER_ID).float() # (bs, l)
|
| 724 |
+
inter_chain = (expanded_asym.unsqueeze(-1) != expanded_asym.unsqueeze(-2)).float() # (bs, l, l)
|
| 725 |
+
near_contact = (rep_distances < 8).float() # (bs, l, l)
|
| 726 |
interface_per_token = (near_contact * inter_chain * (1.0 - is_ligand).unsqueeze(-1)).amax(
|
| 727 |
dim=-1
|
| 728 |
+
) # (bs, l)
|
| 729 |
iplddt_weight = torch.where(
|
| 730 |
is_ligand.bool(),
|
| 731 |
torch.full_like(interface_per_token, 2.0),
|
| 732 |
interface_per_token,
|
| 733 |
+
) # (bs, l)
|
| 734 |
iplddt_weight_atoms = gather_token_to_atom(
|
| 735 |
iplddt_weight.unsqueeze(-1), atom_to_token_m
|
| 736 |
+
).squeeze(-1) # (bs, a)
|
| 737 |
+
atom_iplddt_w = atom_mask_f * iplddt_weight_atoms # (bs, a)
|
| 738 |
complex_iplddt = (plddt_per_atom * atom_iplddt_w).sum(dim=-1) / (
|
| 739 |
atom_iplddt_w.sum(dim=-1) + _EPS
|
| 740 |
+
) # (bs,)
|
| 741 |
|
| 742 |
+
plddt_ca = plddt_per_atom.gather(1, rep_idx_m) # (bs, l)
|
| 743 |
|
| 744 |
# PAE
|
| 745 |
+
pae_logits = self.pae_head(self.pae_ln(pair)) # (bs, l, l, n_pae_bins)
|
| 746 |
+
pae = _categorical_mean(pae_logits, start=0.0, end=32.0).detach() # (bs, l, l)
|
| 747 |
|
| 748 |
# PDE
|
| 749 |
+
pde_logits = self.pde_head(self.pde_ln(pair)) # (bs, l, l, n_pde_bins)
|
| 750 |
+
pde = _categorical_mean(pde_logits, start=0.0, end=32.0).detach() # (bs, l, l)
|
| 751 |
|
| 752 |
# Resolved (per-atom binary).
|
| 753 |
+
s_at_atoms_res = self.resolved_ln(s_at_atoms) # (bs, a, d_single)
|
| 754 |
+
w_res = self.resolved_weight[intra_idx] # (bs, a, d_single, 2)
|
| 755 |
+
resolved_logits = torch.einsum("...c,...cb->...b", s_at_atoms_res, w_res) # (bs, a, 2)
|
| 756 |
|
| 757 |
# pTM / ipTM from pae_logits.
|
| 758 |
n_bins = pae_logits.shape[-1]
|
| 759 |
bin_width = 32.0 / n_bins
|
| 760 |
+
bin_centers = torch.arange(0.5 * bin_width, 32.0, bin_width, device=pae_logits.device) # (n_pae_bins,)
|
| 761 |
+
mask_f = mask.float() # (bs, l)
|
| 762 |
+
n_residues = mask_f.sum(dim=-1, keepdim=True) # (bs, 1)
|
| 763 |
+
d0 = 1.24 * (n_residues.clamp(min=19) - 15) ** (1 / 3) - 1.8 # (bs, 1)
|
| 764 |
+
tm_per_bin = 1 / (1 + (bin_centers / d0) ** 2) # (bs, n_pae_bins)
|
| 765 |
+
pae_probs = F.softmax(pae_logits, dim=-1) # (bs, l, l, n_pae_bins)
|
| 766 |
+
tm_expected = (pae_probs * tm_per_bin[:, None, None, :]).sum(dim=-1) # (bs, l, l)
|
| 767 |
+
|
| 768 |
+
pair_mask_2d = mask_f.unsqueeze(-1) * mask_f.unsqueeze(-2) # (bs, l, l)
|
| 769 |
+
ptm_per_row = (tm_expected * pair_mask_2d).sum(dim=-1) / (pair_mask_2d.sum(dim=-1) + _EPS) # (bs, l)
|
| 770 |
+
ptm = ptm_per_row.max(dim=-1).values # (bs,)
|
| 771 |
|
| 772 |
inter_chain_mask = (
|
| 773 |
expanded_asym.unsqueeze(-1) != expanded_asym.unsqueeze(-2)
|
| 774 |
+
).float() * pair_mask_2d # (bs, l, l)
|
| 775 |
iptm_per_row = (tm_expected * inter_chain_mask).sum(dim=-1) / (
|
| 776 |
inter_chain_mask.sum(dim=-1) + _EPS
|
| 777 |
+
) # (bs, l)
|
| 778 |
+
iptm = iptm_per_row.max(dim=-1).values # (bs,)
|
| 779 |
|
| 780 |
max_chain_id = int(expanded_asym.max().item()) if expanded_batch_size > 0 else 0
|
| 781 |
n_chains = max_chain_id + 1
|
|
|
|
| 785 |
n_chains,
|
| 786 |
device=tm_expected.device,
|
| 787 |
dtype=tm_expected.dtype,
|
| 788 |
+
) # (bs, n_chains, n_chains)
|
| 789 |
for c1 in range(n_chains):
|
| 790 |
+
chain_c1 = (expanded_asym == c1).float() * mask_f # (bs, l)
|
| 791 |
if chain_c1.sum() == 0:
|
| 792 |
continue
|
| 793 |
for c2 in range(n_chains):
|
| 794 |
+
chain_c2 = (expanded_asym == c2).float() * mask_f # (bs, l)
|
| 795 |
+
pair_m = chain_c1.unsqueeze(-1) * chain_c2.unsqueeze(-2) # (bs, l, l)
|
| 796 |
+
denom = pair_m.sum(dim=(-1, -2)) + _EPS # (bs,)
|
| 797 |
+
pair_chains_iptm[:, c1, c2] = (tm_expected * pair_m).sum(dim=(-1, -2)) / denom # (bs,)
|
| 798 |
|
| 799 |
return {
|
| 800 |
"plddt_logits": plddt_logits,
|
|
|
|
| 811 |
"ptm": ptm.detach(),
|
| 812 |
"iptm": iptm.detach(),
|
| 813 |
"pair_chains_iptm": pair_chains_iptm.detach(),
|
| 814 |
+
} # mapping of confidence tensors with shapes traced above
|
| 815 |
|
| 816 |
|
| 817 |
def _inverse_softplus(value: float) -> float:
|
|
|
|
| 842 |
device=child.weight.device,
|
| 843 |
)
|
| 844 |
with torch.no_grad():
|
| 845 |
+
replacement.weight.copy_(child.weight) # child.weight.shape
|
| 846 |
if child.bias is not None:
|
| 847 |
+
replacement.bias.copy_(child.bias) # child.bias.shape
|
| 848 |
replacement.eval().requires_grad_(False)
|
| 849 |
setattr(owner, name, replacement)
|
| 850 |
converted.append(path)
|
|
|
|
| 961 |
self.lm_encoder = None
|
| 962 |
|
| 963 |
self.parcae_input_norm = nn.LayerNorm(d_pair)
|
| 964 |
+
self.parcae_log_a = nn.Parameter(torch.zeros(d_pair)) # (d_pair,)
|
| 965 |
parcae_decay_init = math.sqrt(1.0 / 5.0)
|
| 966 |
parcae_delta_init = -math.log(parcae_decay_init)
|
| 967 |
self.parcae_log_delta = nn.Parameter(
|
| 968 |
torch.full((d_pair,), _inverse_softplus(parcae_delta_init), dtype=torch.float32)
|
| 969 |
+
) # (d_pair,)
|
| 970 |
+
self.parcae_b_cont = nn.Parameter(torch.eye(d_pair)) # (d_pair, d_pair)
|
| 971 |
self.parcae_readout = nn.Linear(d_pair, d_pair, bias=False)
|
| 972 |
+
nn.init.eye_(self.parcae_readout.weight) # (d_pair, d_pair)
|
| 973 |
self.parcae_coda = FoldingTrunk(
|
| 974 |
n_layers=config.parcae.coda_n_layers, d_pair=d_pair, expansion_ratio=4
|
| 975 |
)
|
|
|
|
| 1109 |
input_ids: torch.Tensor | None = None,
|
| 1110 |
**kwargs,
|
| 1111 |
) -> torch.Tensor:
|
| 1112 |
+
# Encoded batch: b sequences, padded token width t including BOS/EOS.
|
| 1113 |
del kwargs
|
| 1114 |
if input_ids is not None:
|
| 1115 |
+
return input_ids # (b, t), or caller input_ids.shape
|
| 1116 |
if seq is None:
|
| 1117 |
raise ValueError("Pass either seq or input_ids for ESMFold2 TTT.")
|
| 1118 |
sequences = [seq] if isinstance(seq, str) else seq
|
|
|
|
| 1131 |
(len(encoded), max_len),
|
| 1132 |
SEQUENCE_PAD_TOKEN,
|
| 1133 |
dtype=torch.long,
|
| 1134 |
+
) # (b, t)
|
| 1135 |
for row, token_ids in enumerate(encoded):
|
| 1136 |
input_tensor[row, : len(token_ids)] = torch.tensor(
|
| 1137 |
token_ids,
|
| 1138 |
dtype=torch.long,
|
| 1139 |
+
) # (t_i,)
|
| 1140 |
+
return input_tensor # (b, t), or caller input_ids.shape
|
| 1141 |
|
| 1142 |
def _ttt_mask_token(self) -> int:
|
| 1143 |
return SEQUENCE_MASK_TOKEN
|
|
|
|
| 1151 |
SEQUENCE_STANDARD_AA_MAX_TOKEN,
|
| 1152 |
device=input_ids.device,
|
| 1153 |
dtype=input_ids.dtype,
|
| 1154 |
+
) # (SEQUENCE_STANDARD_AA_MAX_TOKEN - SEQUENCE_STANDARD_AA_MIN_TOKEN,)
|
| 1155 |
|
| 1156 |
def _ttt_non_special_mask(self, input_ids: torch.Tensor) -> torch.Tensor:
|
| 1157 |
+
# input_ids: arbitrary token-ID shape.
|
| 1158 |
return (input_ids >= SEQUENCE_STANDARD_AA_MIN_TOKEN) & (
|
| 1159 |
input_ids < SEQUENCE_STANDARD_AA_MAX_TOKEN
|
| 1160 |
+
) # input_ids.shape
|
| 1161 |
|
| 1162 |
def _ttt_predict_logits(
|
| 1163 |
self,
|
| 1164 |
batch: torch.Tensor | dict[str, torch.Tensor],
|
| 1165 |
**kwargs,
|
| 1166 |
) -> torch.Tensor:
|
| 1167 |
+
# batch: (b, t) token IDs; backbone output last_hidden_state: (b, t, d_lm).
|
| 1168 |
del kwargs
|
| 1169 |
if not isinstance(batch, torch.Tensor):
|
| 1170 |
raise TypeError("ESMFold2 TTT expects input_ids tensors.")
|
|
|
|
| 1174 |
self._ensure_ttt_lm_head()
|
| 1175 |
if self._ttt_lm_head is None:
|
| 1176 |
raise RuntimeError("ESMFold2 TTT MLM head initialization failed.")
|
| 1177 |
+
attention_mask = batch.ne(SEQUENCE_PAD_TOKEN) # (b, t)
|
| 1178 |
output = self._esmc(
|
| 1179 |
input_ids=batch,
|
| 1180 |
attention_mask=attention_mask,
|
| 1181 |
return_dict=True,
|
| 1182 |
compute_sae=False,
|
| 1183 |
)
|
| 1184 |
+
return self._ttt_lm_head(output.last_hidden_state) # (b, t, vocab_size)
|
| 1185 |
|
| 1186 |
@classmethod
|
| 1187 |
def from_pretrained(
|
|
|
|
| 1298 |
lm_mask_pct: float = 0.0,
|
| 1299 |
verbose: bool = False,
|
| 1300 |
) -> Tensor:
|
| 1301 |
+
# Input tensors: (b, l); n_states and d_lm come from the loaded backbone.
|
| 1302 |
if self._esmc_fp8 and torch.is_grad_enabled():
|
| 1303 |
_reload_esmc_bf16_for_gradients(
|
| 1304 |
self,
|
|
|
|
| 1324 |
pad_to_multiple=pad_to,
|
| 1325 |
lm_mask_pct=lm_mask_pct,
|
| 1326 |
mask_token_id=SEQUENCE_MASK_TOKEN,
|
| 1327 |
+
) # (b, l, n_states, d_lm)
|
| 1328 |
progress.update()
|
| 1329 |
+
return result # (b, l, n_states, d_lm)
|
| 1330 |
return compute_lm_hidden_states(
|
| 1331 |
self._esmc,
|
| 1332 |
input_ids,
|
|
|
|
| 1337 |
pad_to_multiple=pad_to,
|
| 1338 |
lm_mask_pct=lm_mask_pct,
|
| 1339 |
mask_token_id=SEQUENCE_MASK_TOKEN,
|
| 1340 |
+
) # (b, l, n_states, d_lm)
|
| 1341 |
|
| 1342 |
def _discretized_dynamics(self) -> tuple[Tensor, Tensor]:
|
| 1343 |
+
delta = F.softplus(self.parcae_log_delta) # (d_pair,)
|
| 1344 |
+
a = torch.exp(-delta * torch.exp(self.parcae_log_a)) # (d_pair,)
|
| 1345 |
+
b = delta[:, None] * self.parcae_b_cont # (d_pair, d_pair)
|
| 1346 |
+
return a, b # (d_pair,), (d_pair, d_pair)
|
| 1347 |
|
| 1348 |
def _init_pair_state(self, ref: Tensor) -> Tensor:
|
| 1349 |
+
# ref: (b, l, l, d_pair).
|
| 1350 |
std = math.sqrt(2.0 / (5.0 * ref.shape[-1]))
|
| 1351 |
+
state = torch.empty_like(ref, dtype=torch.float32) # ref.shape
|
| 1352 |
+
nn.init.trunc_normal_(state, mean=0.0, std=std, a=-3 * std, b=3 * std) # ref.shape
|
| 1353 |
+
return state.to(dtype=ref.dtype) # ref.shape
|
| 1354 |
|
| 1355 |
def _run_one_loop(
|
| 1356 |
self,
|
|
|
|
| 1369 |
# otherwise leaks about 2 GB of l^2 * c_z data into distogram/sample scope.
|
| 1370 |
# training=True forces dropout under eval(), matching the per-loop
|
| 1371 |
# dropout strategy used at train time.
|
| 1372 |
+
# Pair states: (b, l, l, d_pair); pair_mask: (b, l, l); tok_mask: (b, l). MSA depth m may be subsampled.
|
| 1373 |
lm_cfg = self.config.lm_encoder
|
| 1374 |
_per_loop_lm_dropout = (
|
| 1375 |
lm_z is not None
|
|
|
|
| 1391 |
if _per_loop_lm_dropout:
|
| 1392 |
if lm_z is None:
|
| 1393 |
raise RuntimeError("Per-loop LM dropout requires LM pair features.")
|
| 1394 |
+
lm_z_i: Tensor | None = F.dropout(lm_z, p=_lm_dropout_p, training=True) # (b, l, l, d_pair) or None
|
| 1395 |
else:
|
| 1396 |
+
lm_z_i = lm_z # (b, l, l, d_pair) or None
|
| 1397 |
|
| 1398 |
+
refined_lm_z: Tensor | None = None # (b, l, l, d_pair) or None
|
| 1399 |
if lm_z_i is not None and self.lm_encoder is not None:
|
| 1400 |
refined_lm_z = self.lm_encoder(
|
| 1401 |
lm_z_i.to(z_init.dtype), pair_attention_mask=pair_mask
|
| 1402 |
+
) # (b, l, l, d_pair) or None
|
| 1403 |
|
| 1404 |
+
z_inject_pair = z_init # (b, l, l, d_pair)
|
| 1405 |
if lm_z_i is not None and self.lm_encoder is None:
|
| 1406 |
+
z_inject_pair = z_inject_pair + lm_z_i.to(z_inject_pair.dtype) # (b, l, l, d_pair)
|
| 1407 |
|
| 1408 |
if self.msa_encoder is not None and _msa_inputs is not None:
|
| 1409 |
msa_i, mask_i, hd_i, dv_i = maybe_subsample_msa(
|
|
|
|
| 1413 |
_msa_inputs["deletion_value"],
|
| 1414 |
max_depth=_msa_inputs["max_depth"],
|
| 1415 |
enabled=_msa_inputs["subsample_enabled"],
|
| 1416 |
+
) # each (b, m, l); masks/deletion tensors may be None
|
| 1417 |
b_msa, m, l_msa = msa_i.shape
|
| 1418 |
+
msa_oh = F.one_hot(msa_i.permute(0, 2, 1).long(), num_classes=NUM_RES_TYPES).float() # (b, l, m, 33)
|
| 1419 |
msa_attn = (
|
| 1420 |
mask_i.permute(0, 2, 1).float()
|
| 1421 |
if mask_i is not None
|
| 1422 |
else tok_mask[:, :, None].expand(-1, -1, m).float()
|
| 1423 |
+
) # (b, l, m)
|
| 1424 |
# Bias-free MSAEncoder.embed requires zeroed padding.
|
| 1425 |
+
msa_oh = msa_oh * msa_attn.unsqueeze(-1) # (b, l, m, 33)
|
| 1426 |
hd = (
|
| 1427 |
hd_i.permute(0, 2, 1).float()
|
| 1428 |
if hd_i is not None
|
| 1429 |
else torch.zeros(b_msa, l_msa, m, device=msa_i.device)
|
| 1430 |
+
) # (b, l, m)
|
| 1431 |
dv = (
|
| 1432 |
dv_i.permute(0, 2, 1).float()
|
| 1433 |
if dv_i is not None
|
| 1434 |
else torch.zeros(b_msa, l_msa, m, device=msa_i.device)
|
| 1435 |
+
) # (b, l, m)
|
| 1436 |
msa_pair = self.msa_encoder(
|
| 1437 |
x_pair=z_inject_pair,
|
| 1438 |
x_inputs=_msa_inputs["x_inputs"],
|
|
|
|
| 1440 |
has_deletion=hd,
|
| 1441 |
deletion_value=dv,
|
| 1442 |
msa_attention_mask=msa_attn,
|
| 1443 |
+
).to(z_inject_pair.dtype) # (b, l, l, d_pair)
|
| 1444 |
z_inject_pair = (
|
| 1445 |
msa_pair if self.config.msa_encoder_overwrite else (z_inject_pair + msa_pair)
|
| 1446 |
+
) # (b, l, l, d_pair)
|
| 1447 |
|
| 1448 |
if refined_lm_z is not None:
|
| 1449 |
+
z_inject_pair = z_inject_pair + refined_lm_z.to(z_inject_pair.dtype) # (b, l, l, d_pair)
|
| 1450 |
|
| 1451 |
+
injected_pair = self.parcae_input_norm(z_inject_pair) # (b, l, l, d_pair)
|
| 1452 |
+
z = a * z + F.linear(injected_pair.to(z.dtype), b_mat) # (b, l, l, d_pair)
|
| 1453 |
+
z = self.folding_trunk(z, pair_attention_mask=pair_mask) # (b, l, l, d_pair)
|
| 1454 |
|
| 1455 |
+
return z # (b, l, l, d_pair)
|
| 1456 |
|
| 1457 |
def forward(
|
| 1458 |
self,
|
|
|
|
| 1502 |
disto_cond_mask: Tensor | None = None,
|
| 1503 |
verbose: bool = False,
|
| 1504 |
) -> ESMFold2Output | tuple[Any, ...]:
|
| 1505 |
+
# Token IDs/masks: (b, l); atom IDs/masks: (b, a); ref_pos: (b, a, 3); chars: (b, a, 4); MSA: (b, m, l); bs = b * samples.
|
| 1506 |
output_hidden_states, return_dict = _resolve_structure_output_controls(
|
| 1507 |
self.config,
|
| 1508 |
output_attentions=output_attentions,
|
|
|
|
| 1523 |
disto_cond_mask=disto_cond_mask,
|
| 1524 |
)
|
| 1525 |
del gt_coords, is_resolved, frames_idx
|
| 1526 |
+
tok_mask = token_attention_mask # (b, l)
|
| 1527 |
+
atm_mask = atom_attention_mask # (b, a)
|
| 1528 |
+
disto_idx = distogram_atom_idx # (b, l)
|
| 1529 |
|
| 1530 |
n_loops: int = num_loops if num_loops is not None else self.config.num_loops
|
| 1531 |
n_samples: int = (
|
|
|
|
| 1536 |
total_steps = max(1, n_loops + 1)
|
| 1537 |
|
| 1538 |
if res_type.dim() == 2:
|
| 1539 |
+
res_type_oh = F.one_hot(res_type.long(), num_classes=NUM_RES_TYPES).float() # (b, l, 33)
|
| 1540 |
+
res_type_oh = res_type_oh * tok_mask.unsqueeze(-1).float() # (b, l, 33)
|
| 1541 |
else:
|
| 1542 |
+
res_type_oh = res_type.float() # (b, l, 33)
|
| 1543 |
|
| 1544 |
if msa is not None:
|
| 1545 |
+
msa_oh_profile = F.one_hot(msa.long(), num_classes=NUM_RES_TYPES).float() # (b, m, l, 33)
|
| 1546 |
if msa_attention_mask is not None:
|
| 1547 |
+
mask_f = msa_attention_mask.float().unsqueeze(-1) # (b, m, l, 1)
|
| 1548 |
+
msa_oh_profile = msa_oh_profile * mask_f # (b, m, l, 33)
|
| 1549 |
+
valid_seq_count = msa_attention_mask.float().sum(dim=1).clamp(min=1) # (b, l)
|
| 1550 |
+
profile = msa_oh_profile.sum(dim=1) / valid_seq_count.unsqueeze(-1) # (b, l, 33)
|
| 1551 |
else:
|
| 1552 |
+
profile = msa_oh_profile.mean(dim=1) # (b, l, 33)
|
| 1553 |
else:
|
| 1554 |
+
profile = res_type_oh # (b, l, 33)
|
| 1555 |
|
| 1556 |
if deletion_mean is None:
|
| 1557 |
deletion_mean = torch.zeros(
|
| 1558 |
res_type.shape[0], res_type.shape[1], device=res_type.device
|
| 1559 |
+
) # (b, l)
|
| 1560 |
|
| 1561 |
+
ref_element_oh = F.one_hot(ref_element.long(), num_classes=MAX_ATOMIC_NUMBER).float() # (b, a, 128)
|
| 1562 |
ref_atom_name_chars_oh = F.one_hot(
|
| 1563 |
ref_atom_name_chars.long(), num_classes=CHAR_VOCAB_SIZE
|
| 1564 |
+
).float() # (b, a, 4, 64)
|
| 1565 |
# Bias-free downstream Linears require zeroed padding.
|
| 1566 |
+
atm_mask_f = atm_mask.float() # (b, a)
|
| 1567 |
+
ref_element_oh = ref_element_oh * atm_mask_f.unsqueeze(-1) # (b, a, 128)
|
| 1568 |
+
ref_atom_name_chars_oh = ref_atom_name_chars_oh * atm_mask_f.unsqueeze(-1).unsqueeze(-1) # (b, a, 4, 64)
|
| 1569 |
+
atom_to_token = atom_to_token * atm_mask.long() # (b, a)
|
| 1570 |
|
| 1571 |
use_amp = ref_pos.device.type == "cuda"
|
| 1572 |
with torch.amp.autocast("cuda", enabled=use_amp, dtype=torch.bfloat16):
|
|
|
|
| 1581 |
ref_element=ref_element_oh,
|
| 1582 |
ref_atom_name_chars=ref_atom_name_chars_oh,
|
| 1583 |
atom_to_token=atom_to_token,
|
| 1584 |
+
) # (b, l, d_inputs)
|
| 1585 |
|
| 1586 |
+
z_init = self.z_init_1(x_inputs).unsqueeze(2) + self.z_init_2(x_inputs).unsqueeze(1) # (b, l, l, d_pair)
|
| 1587 |
|
| 1588 |
relative_position_encoding = self.rel_pos(
|
| 1589 |
residue_index=residue_index,
|
|
|
|
| 1591 |
sym_id=sym_id,
|
| 1592 |
entity_id=entity_id,
|
| 1593 |
token_index=token_index,
|
| 1594 |
+
) # (b, l, l, d_pair)
|
| 1595 |
+
token_bonds_encoding = self.token_bonds(token_bonds.float()) # (b, l, l, d_pair)
|
| 1596 |
+
z_init = z_init + relative_position_encoding + token_bonds_encoding # (b, l, l, d_pair)
|
| 1597 |
|
| 1598 |
if lm_hidden_states is None and input_ids is not None and self._esmc is not None:
|
| 1599 |
lm_hidden_states = self._compute_lm_hidden_states(
|
|
|
|
| 1604 |
tok_mask,
|
| 1605 |
lm_mask_pct=(self.config.lm_mask_pct if lm_mask_pct is None else lm_mask_pct),
|
| 1606 |
verbose=verbose,
|
| 1607 |
+
) # (b, l, n_states, d_lm)
|
| 1608 |
+
lm_z: Tensor | None = None # (b, l, l, d_pair) or None
|
| 1609 |
if lm_hidden_states is not None:
|
| 1610 |
+
lm_z = self.language_model(lm_hidden_states.detach()) # (b, l, l, d_pair) or None
|
| 1611 |
del lm_hidden_states
|
| 1612 |
|
| 1613 |
+
pair_mask = tok_mask[:, :, None].float() * tok_mask[:, None, :].float() # (b, l, l)
|
| 1614 |
|
| 1615 |
+
z = self._init_pair_state(z_init) # (b, l, l, d_pair)
|
| 1616 |
|
| 1617 |
+
a, b = self._discretized_dynamics() # (d_pair,), (d_pair, d_pair)
|
| 1618 |
+
a = a.view(1, 1, 1, -1).to(device=z.device, dtype=z.dtype) # (1, 1, 1, d_pair)
|
| 1619 |
+
b_mat = b.to(device=z.device, dtype=z.dtype) # (d_pair, d_pair)
|
| 1620 |
|
| 1621 |
_msa_inputs: dict | None = None
|
| 1622 |
if self.msa_encoder is not None and msa is not None:
|
| 1623 |
msa_attention_mask = maybe_apply_msa_column_masking(
|
| 1624 |
msa_attention_mask,
|
| 1625 |
msa_column_mask_rate,
|
| 1626 |
+
) # (b, m, l)
|
| 1627 |
_msa_inputs = dict(
|
| 1628 |
x_inputs=x_inputs,
|
| 1629 |
msa=msa,
|
|
|
|
| 1646 |
tok_mask=tok_mask,
|
| 1647 |
total_steps=total_steps,
|
| 1648 |
verbose=verbose,
|
| 1649 |
+
) # (b, l, l, d_pair)
|
| 1650 |
del z_init, lm_z, _msa_inputs, a, b_mat
|
| 1651 |
|
| 1652 |
+
z = self.parcae_readout(z) # (b, l, l, d_pair)
|
| 1653 |
+
z = self.parcae_coda(z, pair_attention_mask=pair_mask) # (b, l, l, d_pair)
|
| 1654 |
|
| 1655 |
+
z = z.float() # (b, l, l, d_pair)
|
| 1656 |
+
distogram_logits = self.distogram_head(z + z.transpose(-2, -3)) # (b, l, l, n_distogram_bins)
|
| 1657 |
|
| 1658 |
structure_output = self.structure_head.sample(
|
| 1659 |
z_trunk=z,
|
|
|
|
| 1681 |
return_atom_repr=False,
|
| 1682 |
denoising_early_exit_rmsd=(0.10 if early_exit else None),
|
| 1683 |
verbose=verbose,
|
| 1684 |
+
) # tensor mapping follows the called head's shape contract
|
| 1685 |
|
| 1686 |
+
sample_coords = structure_output["sample_atom_coords"] # (bs, a, 3), or explicit (b, samples, a, 3)
|
| 1687 |
if sample_coords is None:
|
| 1688 |
raise RuntimeError("ESMFold2 structure sampling did not return coordinates.")
|
| 1689 |
output: dict[str, Tensor] = {"distogram_logits": distogram_logits}
|
| 1690 |
+
output["sample_atom_coords"] = sample_coords # sample_coords.shape
|
| 1691 |
|
| 1692 |
confidence_config = self.config.confidence_head
|
| 1693 |
confidence_enabled = (
|
|
|
|
| 1712 |
num_diffusion_samples=n_samples,
|
| 1713 |
relative_position_encoding=relative_position_encoding.detach(),
|
| 1714 |
token_bonds_encoding=token_bonds_encoding.detach(),
|
| 1715 |
+
) # tensor mapping follows the called head's shape contract
|
| 1716 |
progress.update()
|
| 1717 |
else:
|
| 1718 |
confidence_output = self.confidence_head(
|
|
|
|
| 1728 |
num_diffusion_samples=n_samples,
|
| 1729 |
relative_position_encoding=relative_position_encoding.detach(),
|
| 1730 |
token_bonds_encoding=token_bonds_encoding.detach(),
|
| 1731 |
+
) # tensor mapping follows the called head's shape contract
|
| 1732 |
output.update(confidence_output)
|
| 1733 |
+
output["atom_pad_mask"] = atm_mask.unsqueeze(0) if atm_mask.dim() == 1 else atm_mask # (b, a)
|
| 1734 |
+
output["residue_index"] = residue_index # (b, l)
|
| 1735 |
+
output["entity_id"] = entity_id # (b, l)
|
| 1736 |
return _finalize_structure_output(
|
| 1737 |
output,
|
| 1738 |
token_input_state=x_inputs,
|
| 1739 |
pair_state=z,
|
| 1740 |
output_hidden_states=output_hidden_states,
|
| 1741 |
return_dict=return_dict,
|
| 1742 |
+
) # ESMFold2Output/tuple retaining the traced tensor shapes
|
| 1743 |
|
| 1744 |
@torch.no_grad()
|
| 1745 |
def infer_protein(self, seq: str, **forward_kwargs) -> ESMFold2Output:
|
|
|
|
| 1758 |
if not self.config.msa_conditioning:
|
| 1759 |
for name in MSA_CONDITIONING_INPUT_NAMES:
|
| 1760 |
features.pop(name, None)
|
| 1761 |
+
features = {k: v.to(self.device) for k, v in features.items()} # every feature retains its shape
|
| 1762 |
return self(**features, **forward_kwargs, return_dict=True)
|
| 1763 |
|
| 1764 |
@property
|
|
|
|
| 2071 |
msa_attention_mask: Tensor,
|
| 2072 |
pair_attention_mask: Tensor,
|
| 2073 |
) -> tuple[Tensor, Tensor]:
|
| 2074 |
+
# m: (b, l, m_depth, d_msa); pair: (b, l, l, d_pair); corresponding masks omit the feature axis.
|
| 2075 |
+
pair = pair + self.outer_product_mean(m, msa_attention_mask) # (b, l, l, d_pair)
|
| 2076 |
if not self.is_final_block:
|
| 2077 |
+
m = m + self.msa_pair_weighted_averaging(m, pair, pair_attention_mask) # (b, l, m_depth, d_msa)
|
| 2078 |
+
m = m + self.msa_transition(m) # (b, l, m_depth, d_msa)
|
| 2079 |
+
pair = pair + self.tri_mul_out(pair, mask=pair_attention_mask) # (b, l, l, d_pair)
|
| 2080 |
+
pair = pair + self.tri_mul_in(pair, mask=pair_attention_mask) # (b, l, l, d_pair)
|
| 2081 |
+
pair = pair + self.pair_transition(pair) # (b, l, l, d_pair)
|
| 2082 |
+
return m, pair # (b, l, m_depth, d_msa), (b, l, l, d_pair)
|
| 2083 |
|
| 2084 |
|
| 2085 |
class MSAEncoder(nn.Module):
|
|
|
|
| 2126 |
msa_attention_mask: Tensor,
|
| 2127 |
) -> Tensor:
|
| 2128 |
# Every input tensor is pre-transposed to shape (b, l, m, ...) before this call.
|
| 2129 |
+
# x_pair: (b, l, l, d_pair); x_inputs: (b, l, d_inputs); MSA features: (b, l, m, 33), deletion/mask: (b, l, m).
|
| 2130 |
m_feat = torch.cat(
|
| 2131 |
[msa_oh, has_deletion.unsqueeze(-1), deletion_value.unsqueeze(-1)], dim=-1
|
| 2132 |
+
) # (b, l, m, 35)
|
| 2133 |
+
m = self.embed(m_feat) + self.project_inputs(x_inputs).unsqueeze(2) # (b, l, m, d_msa)
|
| 2134 |
+
tok_mask = msa_attention_mask[:, :, 0].bool() # (b, l)
|
| 2135 |
+
pair_attention_mask = tok_mask.unsqueeze(2) & tok_mask.unsqueeze(1) # (b, l, l)
|
| 2136 |
for block in self.blocks:
|
| 2137 |
+
m, x_pair = block(m, x_pair, msa_attention_mask, pair_attention_mask) # (b, l, m, d_msa), (b, l, l, d_pair)
|
| 2138 |
+
return x_pair # (b, l, l, d_pair)
|
fastplms/models/esmfold2/modeling_esmfold2_classification.py
CHANGED
|
@@ -2,10 +2,10 @@
|
|
| 2 |
|
| 3 |
from __future__ import annotations
|
| 4 |
|
| 5 |
-
from typing import Any, Literal
|
| 6 |
-
|
| 7 |
import torch
|
| 8 |
import torch.nn as nn
|
|
|
|
|
|
|
| 9 |
from torch import Tensor
|
| 10 |
|
| 11 |
from ..classification_probe import SequenceClassificationProbe, TokenClassificationProbe
|
|
@@ -94,46 +94,47 @@ class _ESMFold2ClassificationMixin:
|
|
| 94 |
if not sequences:
|
| 95 |
raise ValueError("prepare_classifier_inputs requires at least one sequence.")
|
| 96 |
encoded = [_encode_single_chain(sequence) for sequence in sequences]
|
| 97 |
-
sequence_length = max(map(len, encoded))
|
| 98 |
input_ids = torch.full(
|
| 99 |
(len(encoded), sequence_length),
|
| 100 |
SEQUENCE_PAD_TOKEN,
|
| 101 |
dtype=torch.long,
|
| 102 |
device=self.device,
|
| 103 |
-
)
|
| 104 |
-
attention_mask = torch.zeros_like(input_ids, dtype=torch.bool)
|
| 105 |
for batch_index, token_ids in enumerate(encoded):
|
| 106 |
-
length = len(token_ids)
|
| 107 |
input_ids[batch_index, :length] = torch.tensor(
|
| 108 |
token_ids, dtype=torch.long, device=self.device
|
| 109 |
-
)
|
| 110 |
-
attention_mask[batch_index, :length] = True
|
| 111 |
-
return {"input_ids": input_ids, "attention_mask": attention_mask}
|
| 112 |
|
| 113 |
def _classifier_embeddings(
|
| 114 |
self, input_ids: Tensor, attention_mask: Tensor | None
|
| 115 |
) -> tuple[Tensor, Tensor]:
|
|
|
|
| 116 |
if input_ids.ndim != 2:
|
| 117 |
raise ValueError(
|
| 118 |
"ESMFold2 classifier input_ids must have shape (batch, residue), "
|
| 119 |
f"got {tuple(input_ids.shape)}."
|
| 120 |
)
|
| 121 |
if attention_mask is None:
|
| 122 |
-
attention_mask = input_ids.ne(SEQUENCE_PAD_TOKEN)
|
| 123 |
elif attention_mask.shape != input_ids.shape:
|
| 124 |
raise ValueError(
|
| 125 |
"ESMFold2 classifier attention_mask must match input_ids, got "
|
| 126 |
f"{tuple(attention_mask.shape)} and {tuple(input_ids.shape)}."
|
| 127 |
)
|
| 128 |
-
residue_mask = attention_mask.to(device=input_ids.device, dtype=torch.bool)
|
| 129 |
if not residue_mask.any(dim=1).all():
|
| 130 |
raise ValueError("Every ESMFold2 classifier input must contain a protein residue.")
|
| 131 |
if input_ids.masked_select(residue_mask).eq(SEQUENCE_PAD_TOKEN).any():
|
| 132 |
raise ValueError("ESMFold2 classifier padding tokens cannot be attended residues.")
|
| 133 |
-
residue_ids = input_ids.masked_select(residue_mask)
|
| 134 |
valid_residue_ids = torch.tensor(
|
| 135 |
sorted(_VALID_RESIDUE_IDS), dtype=input_ids.dtype, device=input_ids.device
|
| 136 |
-
)
|
| 137 |
if not torch.isin(residue_ids, valid_residue_ids).all():
|
| 138 |
raise ValueError(
|
| 139 |
"ESMFold2 classifiers accept residue-only single-chain protein inputs."
|
|
@@ -142,9 +143,9 @@ class _ESMFold2ClassificationMixin:
|
|
| 142 |
batch_size, sequence_length = input_ids.shape
|
| 143 |
residue_index = torch.arange(sequence_length, device=input_ids.device).expand(
|
| 144 |
batch_size, -1
|
| 145 |
-
)
|
| 146 |
-
asym_id = torch.zeros_like(input_ids)
|
| 147 |
-
mol_type = torch.zeros_like(input_ids)
|
| 148 |
with torch.no_grad():
|
| 149 |
hidden_states = self._compute_lm_hidden_states(
|
| 150 |
input_ids,
|
|
@@ -152,9 +153,9 @@ class _ESMFold2ClassificationMixin:
|
|
| 152 |
residue_index,
|
| 153 |
mol_type,
|
| 154 |
residue_mask,
|
| 155 |
-
)
|
| 156 |
-
embeddings = self.project_esmc_hidden_states(hidden_states, residue_mask)
|
| 157 |
-
return embeddings, residue_mask
|
| 158 |
|
| 159 |
def _classifier_forward(
|
| 160 |
self,
|
|
@@ -165,7 +166,7 @@ class _ESMFold2ClassificationMixin:
|
|
| 165 |
output_hidden_states: bool | None = None,
|
| 166 |
return_dict: bool | None = None,
|
| 167 |
):
|
| 168 |
-
embeddings, residue_mask = self._classifier_embeddings(input_ids, attention_mask)
|
| 169 |
return self.classifier(
|
| 170 |
embeddings,
|
| 171 |
attention_mask=residue_mask,
|
|
|
|
| 2 |
|
| 3 |
from __future__ import annotations
|
| 4 |
|
|
|
|
|
|
|
| 5 |
import torch
|
| 6 |
import torch.nn as nn
|
| 7 |
+
|
| 8 |
+
from typing import Any, Literal
|
| 9 |
from torch import Tensor
|
| 10 |
|
| 11 |
from ..classification_probe import SequenceClassificationProbe, TokenClassificationProbe
|
|
|
|
| 94 |
if not sequences:
|
| 95 |
raise ValueError("prepare_classifier_inputs requires at least one sequence.")
|
| 96 |
encoded = [_encode_single_chain(sequence) for sequence in sequences]
|
| 97 |
+
sequence_length = max(map(len, encoded)) # l; b = len(encoded).
|
| 98 |
input_ids = torch.full(
|
| 99 |
(len(encoded), sequence_length),
|
| 100 |
SEQUENCE_PAD_TOKEN,
|
| 101 |
dtype=torch.long,
|
| 102 |
device=self.device,
|
| 103 |
+
) # (b, l)
|
| 104 |
+
attention_mask = torch.zeros_like(input_ids, dtype=torch.bool) # (b, l)
|
| 105 |
for batch_index, token_ids in enumerate(encoded):
|
| 106 |
+
length = len(token_ids) # l_i
|
| 107 |
input_ids[batch_index, :length] = torch.tensor(
|
| 108 |
token_ids, dtype=torch.long, device=self.device
|
| 109 |
+
) # (l_i,)
|
| 110 |
+
attention_mask[batch_index, :length] = True # (l_i,)
|
| 111 |
+
return {"input_ids": input_ids, "attention_mask": attention_mask} # both (b, l)
|
| 112 |
|
| 113 |
def _classifier_embeddings(
|
| 114 |
self, input_ids: Tensor, attention_mask: Tensor | None
|
| 115 |
) -> tuple[Tensor, Tensor]:
|
| 116 |
+
# input_ids: (b, l); attention_mask: (b, l) or None.
|
| 117 |
if input_ids.ndim != 2:
|
| 118 |
raise ValueError(
|
| 119 |
"ESMFold2 classifier input_ids must have shape (batch, residue), "
|
| 120 |
f"got {tuple(input_ids.shape)}."
|
| 121 |
)
|
| 122 |
if attention_mask is None:
|
| 123 |
+
attention_mask = input_ids.ne(SEQUENCE_PAD_TOKEN) # (b, l)
|
| 124 |
elif attention_mask.shape != input_ids.shape:
|
| 125 |
raise ValueError(
|
| 126 |
"ESMFold2 classifier attention_mask must match input_ids, got "
|
| 127 |
f"{tuple(attention_mask.shape)} and {tuple(input_ids.shape)}."
|
| 128 |
)
|
| 129 |
+
residue_mask = attention_mask.to(device=input_ids.device, dtype=torch.bool) # (b, l)
|
| 130 |
if not residue_mask.any(dim=1).all():
|
| 131 |
raise ValueError("Every ESMFold2 classifier input must contain a protein residue.")
|
| 132 |
if input_ids.masked_select(residue_mask).eq(SEQUENCE_PAD_TOKEN).any():
|
| 133 |
raise ValueError("ESMFold2 classifier padding tokens cannot be attended residues.")
|
| 134 |
+
residue_ids = input_ids.masked_select(residue_mask) # (n_present,)
|
| 135 |
valid_residue_ids = torch.tensor(
|
| 136 |
sorted(_VALID_RESIDUE_IDS), dtype=input_ids.dtype, device=input_ids.device
|
| 137 |
+
) # (n_valid_ids,)
|
| 138 |
if not torch.isin(residue_ids, valid_residue_ids).all():
|
| 139 |
raise ValueError(
|
| 140 |
"ESMFold2 classifiers accept residue-only single-chain protein inputs."
|
|
|
|
| 143 |
batch_size, sequence_length = input_ids.shape
|
| 144 |
residue_index = torch.arange(sequence_length, device=input_ids.device).expand(
|
| 145 |
batch_size, -1
|
| 146 |
+
) # (b, l)
|
| 147 |
+
asym_id = torch.zeros_like(input_ids) # (b, l)
|
| 148 |
+
mol_type = torch.zeros_like(input_ids) # (b, l)
|
| 149 |
with torch.no_grad():
|
| 150 |
hidden_states = self._compute_lm_hidden_states(
|
| 151 |
input_ids,
|
|
|
|
| 153 |
residue_index,
|
| 154 |
mol_type,
|
| 155 |
residue_mask,
|
| 156 |
+
) # (b, l, n_states, d_lm)
|
| 157 |
+
embeddings = self.project_esmc_hidden_states(hidden_states, residue_mask) # (b, l, d_pair)
|
| 158 |
+
return embeddings, residue_mask # (b, l, d_pair), (b, l)
|
| 159 |
|
| 160 |
def _classifier_forward(
|
| 161 |
self,
|
|
|
|
| 166 |
output_hidden_states: bool | None = None,
|
| 167 |
return_dict: bool | None = None,
|
| 168 |
):
|
| 169 |
+
embeddings, residue_mask = self._classifier_embeddings(input_ids, attention_mask) # (b, l, d_pair), (b, l)
|
| 170 |
return self.classifier(
|
| 171 |
embeddings,
|
| 172 |
attention_mask=residue_mask,
|
fastplms/models/esmfold2/modeling_esmfold2_common.py
CHANGED
|
@@ -10,13 +10,13 @@
|
|
| 10 |
from __future__ import annotations
|
| 11 |
|
| 12 |
import importlib
|
| 13 |
-
from functools import partial
|
| 14 |
-
from importlib.util import find_spec
|
| 15 |
-
from typing import Any, ClassVar, cast
|
| 16 |
-
|
| 17 |
import torch
|
| 18 |
import torch.nn as nn
|
| 19 |
import torch.nn.functional as F
|
|
|
|
|
|
|
|
|
|
|
|
|
| 20 |
from torch import Tensor
|
| 21 |
from torch.utils.checkpoint import checkpoint
|
| 22 |
from tqdm.auto import tqdm
|
|
@@ -24,6 +24,7 @@ from tqdm.auto import tqdm
|
|
| 24 |
from .configuration_esmfold2 import ESMFold2Config
|
| 25 |
from .reproducibility import seed_context
|
| 26 |
|
|
|
|
| 27 |
_seed_context = seed_context
|
| 28 |
|
| 29 |
try:
|
|
@@ -197,15 +198,16 @@ class DropoutResidual(nn.Module):
|
|
| 197 |
self._impl = nn.Dropout(r)
|
| 198 |
|
| 199 |
def forward(self, residual: Tensor, delta: Tensor) -> Tensor:
|
|
|
|
| 200 |
if self._use_fused_kernels:
|
| 201 |
-
return self._impl(residual, delta)
|
| 202 |
-
# The unfused
|
| 203 |
if self._r == 0.0 or not self.training:
|
| 204 |
-
return residual + delta
|
| 205 |
shape = list(delta.shape)
|
| 206 |
shape[self._batch_dim] = 1
|
| 207 |
-
mask = self._impl(delta.new_ones(shape))
|
| 208 |
-
return residual + delta * mask
|
| 209 |
|
| 210 |
|
| 211 |
# ---------------------------------------------------------------------------
|
|
@@ -217,7 +219,7 @@ XYZ_DIMS: int = 3
|
|
| 217 |
MAX_ATOMIC_NUMBER: int = 128
|
| 218 |
|
| 219 |
# Input feature dim = 3 + 1 + 1 + 128 + 64*4 = 389
|
| 220 |
-
ATOM_FEATURE_DIM: int = XYZ_DIMS + 1 + 1 + MAX_ATOMIC_NUMBER + CHAR_VOCAB_SIZE * MAX_CHARS
|
| 221 |
|
| 222 |
|
| 223 |
NUM_RES_TYPES: int = 33
|
|
@@ -246,39 +248,41 @@ def maybe_subsample_msa(
|
|
| 246 |
max_depth: int | None,
|
| 247 |
enabled: bool,
|
| 248 |
) -> tuple[Tensor, Tensor | None, Tensor | None, Tensor | None]:
|
|
|
|
| 249 |
if not enabled or max_depth is None:
|
| 250 |
-
return msa, msa_attention_mask, has_deletion, deletion_value
|
| 251 |
|
| 252 |
depth = msa.size(1)
|
| 253 |
if depth <= 1 or depth <= max_depth:
|
| 254 |
-
return msa, msa_attention_mask, has_deletion, deletion_value
|
| 255 |
|
| 256 |
-
indices = torch.zeros(max_depth, dtype=torch.long, device=msa.device)
|
| 257 |
-
indices[1:] = torch.randperm(depth - 1, device=msa.device)[: max_depth - 1] + 1
|
| 258 |
-
indices = indices.sort().values
|
| 259 |
|
| 260 |
-
msa = msa[:, indices]
|
| 261 |
if msa_attention_mask is not None:
|
| 262 |
-
msa_attention_mask = msa_attention_mask[:, indices]
|
| 263 |
if has_deletion is not None:
|
| 264 |
-
has_deletion = has_deletion[:, indices]
|
| 265 |
if deletion_value is not None:
|
| 266 |
-
deletion_value = deletion_value[:, indices]
|
| 267 |
-
return msa, msa_attention_mask, has_deletion, deletion_value
|
| 268 |
|
| 269 |
|
| 270 |
def maybe_apply_msa_column_masking(
|
| 271 |
msa_attention_mask: Tensor | None,
|
| 272 |
rate: float,
|
| 273 |
) -> Tensor | None:
|
|
|
|
| 274 |
if msa_attention_mask is None or rate <= 0.0 or msa_attention_mask.size(1) <= 1:
|
| 275 |
-
return msa_attention_mask
|
| 276 |
|
| 277 |
batch_size, _, length = msa_attention_mask.shape
|
| 278 |
-
col_keep = torch.rand(batch_size, length, device=msa_attention_mask.device) >= rate
|
| 279 |
-
col_keep = col_keep.unsqueeze(1).expand_as(msa_attention_mask).clone()
|
| 280 |
-
col_keep[:, 0, :] = True
|
| 281 |
-
return msa_attention_mask.bool() & col_keep
|
| 282 |
|
| 283 |
|
| 284 |
# ===========================================================================
|
|
@@ -296,8 +300,8 @@ def gather_token_to_atom(token_features: Tensor, atom_to_token_idx: Tensor) -> T
|
|
| 296 |
Returns:
|
| 297 |
X with shape (b, a, d).
|
| 298 |
"""
|
| 299 |
-
idx = atom_to_token_idx.unsqueeze(-1).expand(-1, -1, token_features.size(-1))
|
| 300 |
-
return torch.gather(token_features, 1, idx)
|
| 301 |
|
| 302 |
|
| 303 |
def scatter_atom_to_token(
|
|
@@ -319,20 +323,20 @@ def scatter_atom_to_token(
|
|
| 319 |
"""
|
| 320 |
batch_size, n_atoms, d_model = atom_features.shape
|
| 321 |
n_out = n_tokens
|
| 322 |
-
idx = atom_to_token_idx
|
| 323 |
if atom_mask is not None:
|
| 324 |
-
idx = torch.where(atom_mask, atom_to_token_idx, n_tokens)
|
| 325 |
n_out = n_tokens + 1
|
| 326 |
-
idx_expanded = idx.unsqueeze(-1).expand(batch_size, n_atoms, d_model)
|
| 327 |
out = torch.zeros(
|
| 328 |
batch_size,
|
| 329 |
n_out,
|
| 330 |
d_model,
|
| 331 |
device=atom_features.device,
|
| 332 |
dtype=atom_features.dtype,
|
| 333 |
-
)
|
| 334 |
-
out.scatter_reduce_(1, idx_expanded, atom_features, reduce="mean", include_self=False)
|
| 335 |
-
return out[:, :n_tokens, :]
|
| 336 |
|
| 337 |
|
| 338 |
def gather_rep_atom_coords(coords: Tensor, rep_atom_idx: Tensor) -> Tensor:
|
|
@@ -345,8 +349,8 @@ def gather_rep_atom_coords(coords: Tensor, rep_atom_idx: Tensor) -> Tensor:
|
|
| 345 |
Returns:
|
| 346 |
X with shape (b, l, 3).
|
| 347 |
"""
|
| 348 |
-
idx = rep_atom_idx.unsqueeze(-1).expand(-1, -1, coords.size(-1))
|
| 349 |
-
return torch.gather(coords, 1, idx)
|
| 350 |
|
| 351 |
|
| 352 |
def _compute_intra_token_idx(atom_to_token: Tensor) -> Tensor:
|
|
@@ -362,12 +366,12 @@ def _compute_intra_token_idx(atom_to_token: Tensor) -> Tensor:
|
|
| 362 |
Index tensor I with shape (b, a) and values from zero through
|
| 363 |
``max_atoms_per_token - 1``.
|
| 364 |
"""
|
| 365 |
-
same_as_prev = F.pad(atom_to_token[:, 1:] == atom_to_token[:, :-1], (1, 0), value=False)
|
| 366 |
-
ones = torch.ones_like(atom_to_token)
|
| 367 |
-
cumsum = torch.cumsum(ones, dim=-1)
|
| 368 |
-
group_start = cumsum.masked_fill(same_as_prev, 0)
|
| 369 |
-
group_start = torch.cummax(group_start, dim=-1).values
|
| 370 |
-
return cumsum - group_start
|
| 371 |
|
| 372 |
|
| 373 |
def _categorical_mean(logits: Tensor, start: float, end: float) -> Tensor:
|
|
@@ -384,9 +388,9 @@ def _categorical_mean(logits: Tensor, start: float, end: float) -> Tensor:
|
|
| 384 |
Expected value tensor Y with shape (...).
|
| 385 |
"""
|
| 386 |
n_bins = logits.shape[-1]
|
| 387 |
-
edges = torch.linspace(start, end, n_bins + 1, device=logits.device, dtype=torch.float32)
|
| 388 |
v_bins = (edges[:-1] + edges[1:]) / 2 # V_bin has shape (n_bins,).
|
| 389 |
-
return (logits.float().softmax(-1) @ v_bins.unsqueeze(1)).squeeze(-1)
|
| 390 |
|
| 391 |
|
| 392 |
# ===========================================================================
|
|
@@ -403,16 +407,17 @@ class RowAttentionPooling(nn.Module):
|
|
| 403 |
self.out_proj = nn.Linear(d_pair, d_single, bias=False)
|
| 404 |
|
| 405 |
def forward(self, z: Tensor, mask: Tensor) -> Tensor:
|
| 406 |
-
|
|
|
|
| 407 |
mask_bias = torch.where(
|
| 408 |
mask[:, None, :].bool(),
|
| 409 |
torch.zeros_like(scores),
|
| 410 |
torch.full_like(scores, -1e9),
|
| 411 |
-
)
|
| 412 |
-
scores = scores + mask_bias
|
| 413 |
-
weights = F.softmax(scores, dim=-1)
|
| 414 |
-
pooled = torch.einsum("bnm,bnmd->bnd", weights, z)
|
| 415 |
-
return self.out_proj(pooled)
|
| 416 |
|
| 417 |
|
| 418 |
# ===========================================================================
|
|
@@ -460,6 +465,7 @@ class InputsEmbedder(nn.Module):
|
|
| 460 |
X with shape (b, l, d_inputs), concatenating atom encoding,
|
| 461 |
aatype, profile, and deletion mean.
|
| 462 |
"""
|
|
|
|
| 463 |
a, _q, _c, _attn_params, _intermediates = self.atom_attention_encoder(
|
| 464 |
ref_pos=ref_pos,
|
| 465 |
atom_attention_mask=atom_attention_mask,
|
|
@@ -468,8 +474,8 @@ class InputsEmbedder(nn.Module):
|
|
| 468 |
ref_element=ref_element,
|
| 469 |
ref_atom_name_chars=ref_atom_name_chars,
|
| 470 |
atom_to_token=atom_to_token,
|
| 471 |
-
)
|
| 472 |
-
return torch.cat([a, aatype, profile, deletion_mean.unsqueeze(-1)], dim=-1)
|
| 473 |
|
| 474 |
|
| 475 |
# ===========================================================================
|
|
@@ -511,38 +517,39 @@ class ResIdxAsymIdSymIdEntityIdEncoding(nn.Module):
|
|
| 511 |
entity_id: Tensor,
|
| 512 |
token_index: Tensor,
|
| 513 |
) -> Tensor:
|
| 514 |
-
|
| 515 |
-
|
| 516 |
-
|
|
|
|
| 517 |
|
| 518 |
-
dij_residue = residue_index.unsqueeze(2) - residue_index.unsqueeze(1)
|
| 519 |
dij_residue = torch.clip(
|
| 520 |
dij_residue + self.n_relative_residx_bins,
|
| 521 |
0,
|
| 522 |
2 * self.n_relative_residx_bins,
|
| 523 |
-
)
|
| 524 |
-
dij_residue = torch.where(bij_same_chain, dij_residue, 2 * self.n_relative_residx_bins + 1)
|
| 525 |
-
aij_rel_pos = F.one_hot(dij_residue, 2 * self.n_relative_residx_bins + 2)
|
| 526 |
|
| 527 |
dij_token = torch.clip(
|
| 528 |
token_index.unsqueeze(2) - token_index.unsqueeze(1) + self.n_relative_residx_bins,
|
| 529 |
0,
|
| 530 |
2 * self.n_relative_residx_bins,
|
| 531 |
-
)
|
| 532 |
dij_token = torch.where(
|
| 533 |
bij_same_chain & bij_same_residue,
|
| 534 |
dij_token,
|
| 535 |
2 * self.n_relative_residx_bins + 1,
|
| 536 |
-
)
|
| 537 |
-
aij_rel_token = F.one_hot(dij_token, 2 * self.n_relative_residx_bins + 2)
|
| 538 |
|
| 539 |
dij_chain = torch.clip(
|
| 540 |
sym_id.unsqueeze(2) - sym_id.unsqueeze(1) + self.n_relative_chain_bins,
|
| 541 |
0,
|
| 542 |
2 * self.n_relative_chain_bins,
|
| 543 |
-
)
|
| 544 |
-
dij_chain = torch.where(bij_same_chain, 2 * self.n_relative_chain_bins + 1, dij_chain)
|
| 545 |
-
aij_rel_chain = F.one_hot(dij_chain, 2 * self.n_relative_chain_bins + 2)
|
| 546 |
|
| 547 |
feats = torch.cat(
|
| 548 |
[
|
|
@@ -552,9 +559,9 @@ class ResIdxAsymIdSymIdEntityIdEncoding(nn.Module):
|
|
| 552 |
aij_rel_chain.float(),
|
| 553 |
],
|
| 554 |
dim=-1,
|
| 555 |
-
)
|
| 556 |
|
| 557 |
-
return self.embed(feats)
|
| 558 |
|
| 559 |
|
| 560 |
# ===========================================================================
|
|
@@ -575,12 +582,13 @@ class SingleToPair(nn.Module):
|
|
| 575 |
)
|
| 576 |
|
| 577 |
def forward(self, x: Tensor) -> Tensor:
|
| 578 |
-
x
|
|
|
|
| 579 |
x = torch.cat(
|
| 580 |
[(x.unsqueeze(2) * x.unsqueeze(1)), (x.unsqueeze(2) - x.unsqueeze(1))],
|
| 581 |
dim=3,
|
| 582 |
-
)
|
| 583 |
-
return self.output_mlp(x)
|
| 584 |
|
| 585 |
|
| 586 |
# ===========================================================================
|
|
@@ -604,7 +612,7 @@ class LanguageModelShim(nn.Module):
|
|
| 604 |
self.base_z_linear = nn.Sequential(
|
| 605 |
nn.LayerNorm(d_model), nn.Linear(d_model, d_z, bias=False)
|
| 606 |
)
|
| 607 |
-
self.base_z_combine = nn.Parameter(torch.zeros(num_layers + 1))
|
| 608 |
|
| 609 |
def project_sequence(
|
| 610 |
self,
|
|
@@ -641,12 +649,12 @@ class LanguageModelShim(nn.Module):
|
|
| 641 |
# Match the learned projection parameters at this explicit boundary;
|
| 642 |
# this preserves the official BF16 path and leaves FP32 models exact.
|
| 643 |
projection_dtype = cast(nn.LayerNorm, self.base_z_linear[0]).weight.dtype
|
| 644 |
-
hidden_states = hidden_states.to(dtype=projection_dtype)
|
| 645 |
-
projected_states = self.base_z_linear(hidden_states)
|
| 646 |
-
layer_weights = self.base_z_combine.softmax(dim=0)
|
| 647 |
# Preserve Biohub's matmul path exactly so checkpoint inference does
|
| 648 |
# not change through a different reduction order.
|
| 649 |
-
projected = layer_weights @ projected_states
|
| 650 |
if residue_mask is not None:
|
| 651 |
if residue_mask.shape != hidden_states.shape[:2]:
|
| 652 |
raise ValueError(
|
|
@@ -655,8 +663,8 @@ class LanguageModelShim(nn.Module):
|
|
| 655 |
)
|
| 656 |
projected = projected * residue_mask.to(
|
| 657 |
device=projected.device, dtype=projected.dtype
|
| 658 |
-
).unsqueeze(-1)
|
| 659 |
-
return projected
|
| 660 |
|
| 661 |
def forward(self, hidden_states: Tensor, *, lm_dropout: float = 0.0) -> Tensor:
|
| 662 |
"""Project pre-computed ESMC hidden states to pair representation.
|
|
@@ -669,11 +677,11 @@ class LanguageModelShim(nn.Module):
|
|
| 669 |
Returns:
|
| 670 |
Z_pair with shape ``(b, l, l, d_pair)``.
|
| 671 |
"""
|
| 672 |
-
lm_z = self.project_sequence(hidden_states)
|
| 673 |
-
lm_z = self.base_z_mlp(lm_z)
|
| 674 |
if lm_dropout > 0:
|
| 675 |
-
lm_z = F.dropout(lm_z, p=lm_dropout, training=True)
|
| 676 |
-
return lm_z
|
| 677 |
|
| 678 |
|
| 679 |
# ===========================================================================
|
|
@@ -700,62 +708,63 @@ def compute_lm_hidden_states(
|
|
| 700 |
was trained on per-residue inputs, not per-atom), then scatter the
|
| 701 |
hidden states back to the per-token layout.
|
| 702 |
"""
|
|
|
|
| 703 |
b_size, l_size = input_ids.shape
|
| 704 |
device = input_ids.device
|
| 705 |
-
protein_mask = (mol_type == 0) & token_mask
|
| 706 |
|
| 707 |
lm_input_list = []
|
| 708 |
lm_lengths = []
|
| 709 |
# Per-batch maps from (original protein-token index) to (LM input position).
|
| 710 |
expand_maps: list[Tensor] = []
|
| 711 |
for batch_index in range(b_size):
|
| 712 |
-
mask_b = protein_mask[batch_index]
|
| 713 |
-
ids_b = input_ids[batch_index][mask_b]
|
| 714 |
-
asym_b = asym_id[batch_index][mask_b]
|
| 715 |
-
res_b = residue_index[batch_index][mask_b]
|
| 716 |
|
| 717 |
# Collapse: keep first token per (asym_id, residue_index) key, in
|
| 718 |
# input order. ``inverse`` maps each original protein-token to its
|
| 719 |
# collapsed residue index.
|
| 720 |
-
keys = torch.stack((asym_b, res_b), dim=1)
|
| 721 |
-
unique_keys, inverse = torch.unique(keys, dim=0, return_inverse=True)
|
| 722 |
n_unique = unique_keys.size(0)
|
| 723 |
-
token_positions = torch.arange(keys.size(0), device=device, dtype=torch.long)
|
| 724 |
-
first_pos = torch.full((n_unique,), keys.size(0), device=device, dtype=torch.long)
|
| 725 |
-
first_pos.scatter_reduce_(0, inverse, token_positions, reduce="amin", include_self=True)
|
| 726 |
-
ordered = torch.argsort(first_pos)
|
| 727 |
-
first_pos_ordered = first_pos[ordered]
|
| 728 |
-
ids_collapsed = ids_b[first_pos_ordered]
|
| 729 |
-
asym_collapsed = asym_b[first_pos_ordered]
|
| 730 |
-
remap = torch.empty_like(ordered)
|
| 731 |
-
remap[ordered] = torch.arange(n_unique, device=device, dtype=torch.long)
|
| 732 |
-
inverse_ordered = remap[inverse]
|
| 733 |
-
|
| 734 |
-
chain_ids = asym_collapsed.unique(sorted=True)
|
| 735 |
# [BOS] chain1 [EOS BOS] chain2 ... [EOS]
|
| 736 |
-
parts: list[Tensor] = [torch.tensor([0], device=device, dtype=ids_b.dtype)]
|
| 737 |
# Per-chain LM positions accumulate; track them for the expand map.
|
| 738 |
-
per_token_lm_pos = torch.empty(n_unique, device=device, dtype=torch.long)
|
| 739 |
cursor = 1 # position 0 is the leading BOS
|
| 740 |
for i, cid in enumerate(chain_ids):
|
| 741 |
-
in_chain = (asym_collapsed == cid).nonzero(as_tuple=True)[0]
|
| 742 |
parts.append(ids_collapsed[in_chain])
|
| 743 |
per_token_lm_pos[in_chain] = torch.arange(
|
| 744 |
cursor, cursor + in_chain.shape[0], device=device, dtype=torch.long
|
| 745 |
-
)
|
| 746 |
cursor += in_chain.shape[0]
|
| 747 |
if i < len(chain_ids) - 1:
|
| 748 |
parts.append(torch.tensor([2, 0], device=device, dtype=ids_b.dtype))
|
| 749 |
cursor += 2 # EOS + BOS
|
| 750 |
parts.append(torch.tensor([2], device=device, dtype=ids_b.dtype))
|
| 751 |
-
lm_seq = torch.cat(parts)
|
| 752 |
lm_input_list.append(lm_seq)
|
| 753 |
lm_lengths.append(lm_seq.shape[0])
|
| 754 |
|
| 755 |
# Map each original protein-token position to its LM input position.
|
| 756 |
-
prot_pos_b = mask_b.nonzero(as_tuple=True)[0]
|
| 757 |
-
expand_map = torch.full((l_size,), -1, device=device, dtype=torch.long)
|
| 758 |
-
expand_map[prot_pos_b] = per_token_lm_pos[inverse_ordered]
|
| 759 |
expand_maps.append(expand_map)
|
| 760 |
|
| 761 |
# Pad the language-model input to its longest sequence. FP8 callers round
|
|
@@ -768,32 +777,32 @@ def compute_lm_hidden_states(
|
|
| 768 |
1,
|
| 769 |
device=device,
|
| 770 |
dtype=input_ids.dtype, # PAD=1
|
| 771 |
-
)
|
| 772 |
for batch_index in range(b_size):
|
| 773 |
-
lm_input_ids[batch_index, : lm_lengths[batch_index]] = lm_input_list[batch_index]
|
| 774 |
|
| 775 |
# sequence_id for chain-aware attention; PAD tokens get -1 (no attention).
|
| 776 |
-
sequence_id = (lm_input_ids == 0).cumsum(dim=1) - 1 # BOS=0
|
| 777 |
-
sequence_id = sequence_id.masked_fill(lm_input_ids == 1, -1) # PAD=1
|
| 778 |
|
| 779 |
if lm_mask_pct > 0.0:
|
| 780 |
-
special = (lm_input_ids == 0) | (lm_input_ids == 1) | (lm_input_ids == 2)
|
| 781 |
-
do_mask = (torch.rand(lm_input_ids.shape, device=device) < lm_mask_pct) & ~special
|
| 782 |
-
lm_input_ids = lm_input_ids.masked_fill(do_mask, mask_token_id)
|
| 783 |
|
| 784 |
with torch.inference_mode():
|
| 785 |
esmc_out = esmc(input_ids=lm_input_ids, sequence_id=sequence_id, output_hidden_states=True)
|
| 786 |
|
| 787 |
-
hidden_stack = esmc_out.hidden_states
|
| 788 |
n_states, _, _, d_model = hidden_stack.shape
|
| 789 |
-
result = torch.zeros(b_size, l_size, n_states, d_model, device=device, dtype=hidden_stack.dtype)
|
| 790 |
for batch_index in range(b_size):
|
| 791 |
-
M_i = protein_mask[batch_index]
|
| 792 |
-
positions = expand_maps[batch_index][M_i]
|
| 793 |
-
gathered = hidden_stack[:, batch_index, positions, :].permute(1, 0, 2)
|
| 794 |
-
result[batch_index, M_i.nonzero(as_tuple=True)[0]] = gathered
|
| 795 |
|
| 796 |
-
return result.detach()
|
| 797 |
|
| 798 |
|
| 799 |
# ===========================================================================
|
|
@@ -844,7 +853,8 @@ class TriangleMultiplicativeBlock(nn.Module):
|
|
| 844 |
return self.flow
|
| 845 |
|
| 846 |
def _triangular_contract(self, left_stream: Tensor, right_stream: Tensor) -> Tensor:
|
| 847 |
-
|
|
|
|
| 848 |
|
| 849 |
def _triangular_contract_chunked(
|
| 850 |
self, left_stream: Tensor, right_stream: Tensor, chunk_size: int
|
|
@@ -869,17 +879,18 @@ class TriangleMultiplicativeBlock(nn.Module):
|
|
| 869 |
chunks = []
|
| 870 |
for start in range(0, length, chunk_size):
|
| 871 |
rows = left_rows[:, :, start : start + chunk_size] # (b, d, i_c, k)
|
| 872 |
-
product = torch.bmm(rows.reshape(batch_size * channels, -1, inner), right_columns)
|
| 873 |
product = product.view(batch_size, channels, rows.shape[2], -1) # (b, d, i_c, j)
|
| 874 |
chunks.append(product.permute(0, 2, 3, 1)) # (b, i_c, j, d)
|
| 875 |
-
return torch.cat(chunks, dim=1)
|
| 876 |
|
| 877 |
def forward(self, pair_grid: Tensor, visibility: Tensor | None = None) -> Tensor:
|
|
|
|
| 878 |
if visibility is None:
|
| 879 |
-
visibility = pair_grid.new_ones(pair_grid.shape[:-1])
|
| 880 |
|
| 881 |
if self._use_kernels:
|
| 882 |
-
p_in_weight, g_in_weight = self.split_kernel_weights()
|
| 883 |
return _cue_tri_mul( # type: ignore[misc]
|
| 884 |
pair_grid,
|
| 885 |
direction=self._kernel_flow_direction(),
|
|
@@ -893,38 +904,38 @@ class TriangleMultiplicativeBlock(nn.Module):
|
|
| 893 |
p_out_weight=self.proj_emit.weight,
|
| 894 |
g_out_weight=self.proj_gate.weight,
|
| 895 |
eps=_EPS,
|
| 896 |
-
)
|
| 897 |
|
| 898 |
# Every tensor below is as large as the pair representation or larger, and this
|
| 899 |
# block sets the peak memory of a fold. Each name is dropped once it is dead, so
|
| 900 |
# the allocator can reuse its buffer. No value changes.
|
| 901 |
-
normalized_grid = self.norm_start(pair_grid)
|
| 902 |
bundled = self.proj_bundle(normalized_grid) # (b, l, l, 4 * d)
|
| 903 |
-
signal, gate_logits = bundled.split(2 * self.latent_channels, dim=-1)
|
| 904 |
routed = signal * torch.sigmoid(gate_logits) # (b, l, l, 2 * d)
|
| 905 |
# The two views would keep the whole projection alive.
|
| 906 |
del bundled, signal, gate_logits
|
| 907 |
-
routed = routed * visibility.unsqueeze(-1)
|
| 908 |
|
| 909 |
left_stream, right_stream = routed.float().chunk(2, dim=-1) # each (b, l, l, d)
|
| 910 |
if torch.is_autocast_enabled(left_stream.device.type):
|
| 911 |
# The contraction is an autocast operation. Casting its inputs here, as it
|
| 912 |
# would, lets the full-precision product go before the contraction runs.
|
| 913 |
autocast_dtype = torch.get_autocast_dtype(left_stream.device.type)
|
| 914 |
-
left_stream = left_stream.to(autocast_dtype)
|
| 915 |
-
right_stream = right_stream.to(autocast_dtype)
|
| 916 |
del routed
|
| 917 |
if self._chunk_size is not None:
|
| 918 |
contracted = self._triangular_contract_chunked(
|
| 919 |
left_stream, right_stream, self._chunk_size
|
| 920 |
-
)
|
| 921 |
else:
|
| 922 |
-
contracted = self._triangular_contract(left_stream, right_stream)
|
| 923 |
del left_stream, right_stream
|
| 924 |
-
mixed = self.proj_emit(self.norm_mix(contracted))
|
| 925 |
del contracted
|
| 926 |
-
output_gate = torch.sigmoid(self.proj_gate(normalized_grid))
|
| 927 |
-
return mixed * output_gate
|
| 928 |
|
| 929 |
|
| 930 |
class TriangleMultiplicativeUpdate(nn.Module):
|
|
@@ -947,7 +958,8 @@ class TriangleMultiplicativeUpdate(nn.Module):
|
|
| 947 |
self._engine.set_chunk_size(chunk_size)
|
| 948 |
|
| 949 |
def forward(self, z: Tensor, mask: Tensor | None = None) -> Tensor:
|
| 950 |
-
|
|
|
|
| 951 |
|
| 952 |
|
| 953 |
# ===========================================================================
|
|
@@ -989,11 +1001,11 @@ class Transition(nn.Module):
|
|
| 989 |
dtype=dtype,
|
| 990 |
)
|
| 991 |
with torch.no_grad():
|
| 992 |
-
fused.LN_W.copy_(self.norm.weight)
|
| 993 |
if has_ln_bias:
|
| 994 |
fused.LN_B.copy_(self.norm.bias) # type: ignore[union-attr]
|
| 995 |
# FusedLNLinearSwiGLU.W12 is (d_model, 2*d_inner); transpose nn.Linear once.
|
| 996 |
-
fused.W12.copy_(self.ffn.w12.weight.t().contiguous())
|
| 997 |
self._fused_swiglu = fused.eval().requires_grad_(False)
|
| 998 |
else:
|
| 999 |
self._fused_swiglu = None
|
|
@@ -1005,50 +1017,53 @@ class Transition(nn.Module):
|
|
| 1005 |
|
| 1006 |
def _swiglu_pre_w3(self, x_normed: Tensor) -> Tensor:
|
| 1007 |
"""SwiGLU through silu(x1)*x2, before the final w3."""
|
|
|
|
| 1008 |
ffn = self.ffn
|
| 1009 |
-
x12 = ffn.w12(x_normed)
|
| 1010 |
-
x1, x2 = x12.split(ffn.hidden_features, dim=-1)
|
| 1011 |
-
return F.silu(x1) * x2
|
| 1012 |
|
| 1013 |
def _addmm_residual(self, x: Tensor, hidden: Tensor) -> Tensor:
|
| 1014 |
"""x + w3(hidden) via single cuBLAS addmm: avoids transition-output allocation."""
|
|
|
|
| 1015 |
ffn = self.ffn
|
| 1016 |
x_shape = x.shape
|
| 1017 |
out = torch.addmm(
|
| 1018 |
x.contiguous().view(-1, x_shape[-1]),
|
| 1019 |
hidden.view(-1, hidden.shape[-1]),
|
| 1020 |
ffn.w3.weight.t(),
|
| 1021 |
-
)
|
| 1022 |
-
return out.view(x_shape)
|
| 1023 |
|
| 1024 |
def forward(self, x: Tensor) -> Tensor:
|
| 1025 |
# Inference-only fast path (addmm-fused residual + pre-alloc out)
|
| 1026 |
#: diverges bit-exactly from ``x + ffn(norm(x))`` so we only use
|
| 1027 |
# it when grad is disabled (binder-design / bit-exact tests run
|
| 1028 |
# with grad on and need the reference path).
|
|
|
|
| 1029 |
if not torch.is_grad_enabled() and self._can_use_fused_path(x):
|
| 1030 |
fused = self._fused_swiglu
|
| 1031 |
assert fused is not None
|
| 1032 |
pre_w3 = fused
|
| 1033 |
if self._chunk_size is None or x.shape[1] <= self._chunk_size:
|
| 1034 |
-
hidden = pre_w3(x)
|
| 1035 |
-
return self._addmm_residual(x, hidden)
|
| 1036 |
-
out = torch.empty_like(x)
|
| 1037 |
for s in range(0, x.shape[1], self._chunk_size):
|
| 1038 |
e = min(s + self._chunk_size, x.shape[1])
|
| 1039 |
-
sl = x[:, s:e]
|
| 1040 |
-
hidden = pre_w3(sl)
|
| 1041 |
-
out[:, s:e] = self._addmm_residual(sl, hidden)
|
| 1042 |
-
return out
|
| 1043 |
# Reference path: bit-exact with main: x + ffn(norm(x)).
|
| 1044 |
if self._chunk_size is None or x.shape[1] <= self._chunk_size:
|
| 1045 |
-
return x + self.ffn(self.norm(x))
|
| 1046 |
out_list: list[Tensor] = []
|
| 1047 |
for s in range(0, x.shape[1], self._chunk_size):
|
| 1048 |
e = min(s + self._chunk_size, x.shape[1])
|
| 1049 |
-
sl = x[:, s:e]
|
| 1050 |
out_list.append(sl + self.ffn(self.norm(sl)))
|
| 1051 |
-
return torch.cat(out_list, dim=1)
|
| 1052 |
|
| 1053 |
|
| 1054 |
class PairUpdateBlock(nn.Module):
|
|
@@ -1087,12 +1102,13 @@ class PairUpdateBlock(nn.Module):
|
|
| 1087 |
self, pair: Tensor, direction: str, pair_attention_mask: Tensor | None
|
| 1088 |
) -> Tensor:
|
| 1089 |
"""Fused TriMul+residual call; weights from the corresponding engine."""
|
|
|
|
| 1090 |
tri = self.tri_mul_out if direction == "outgoing" else self.tri_mul_in
|
| 1091 |
engine: TriangleMultiplicativeBlock = tri._engine # type: ignore[assignment]
|
| 1092 |
-
p_in_weight, g_in_weight = engine.split_kernel_weights()
|
| 1093 |
|
| 1094 |
def _bf16(t: Tensor) -> Tensor:
|
| 1095 |
-
return t if t.dtype == torch.bfloat16 else t.to(torch.bfloat16)
|
| 1096 |
|
| 1097 |
return _fused_trimul_with_residual( # type: ignore[misc]
|
| 1098 |
pair,
|
|
@@ -1109,17 +1125,18 @@ class PairUpdateBlock(nn.Module):
|
|
| 1109 |
g_out_weight=_bf16(engine.proj_gate.weight),
|
| 1110 |
mask=pair_attention_mask,
|
| 1111 |
eps=_EPS,
|
| 1112 |
-
)
|
| 1113 |
|
| 1114 |
def forward(self, pair: Tensor, pair_attention_mask: Tensor | None = None) -> Tensor:
|
|
|
|
| 1115 |
if self._can_use_fused_trimul_with_residual(pair):
|
| 1116 |
-
pair = self._fused_trimul_with_residual(pair, "outgoing", pair_attention_mask)
|
| 1117 |
-
pair = self._fused_trimul_with_residual(pair, "incoming", pair_attention_mask)
|
| 1118 |
else:
|
| 1119 |
-
pair = self.row_drop(pair, self.tri_mul_out(pair, mask=pair_attention_mask))
|
| 1120 |
-
pair = self.row_drop(pair, self.tri_mul_in(pair, mask=pair_attention_mask))
|
| 1121 |
-
pair = self.pair_transition(pair)
|
| 1122 |
-
return pair
|
| 1123 |
|
| 1124 |
|
| 1125 |
class FoldingTrunk(nn.Module):
|
|
@@ -1145,22 +1162,23 @@ class FoldingTrunk(nn.Module):
|
|
| 1145 |
def forward(self, pair: Tensor, pair_attention_mask: Tensor | None = None) -> Tensor:
|
| 1146 |
# Cast the pair tensor to BF16 when the fused triangle backend is enabled
|
| 1147 |
# (its bwd kernel requires bf16). Other backends keep the input dtype.
|
|
|
|
| 1148 |
orig_dtype = pair.dtype
|
| 1149 |
fused_on = (
|
| 1150 |
len(self.blocks) > 0
|
| 1151 |
and getattr(self.blocks[0], "_kernel_backend", None) == BACKEND_FUSED
|
| 1152 |
)
|
| 1153 |
if pair.is_cuda and fused_on and orig_dtype != torch.bfloat16:
|
| 1154 |
-
pair = pair.to(torch.bfloat16)
|
| 1155 |
for block in self.blocks:
|
| 1156 |
fn = partial(block, pair_attention_mask=pair_attention_mask)
|
| 1157 |
if torch.is_grad_enabled():
|
| 1158 |
-
pair = checkpoint(fn, pair, use_reentrant=False) # pyright: ignore
|
| 1159 |
else:
|
| 1160 |
-
pair = fn(pair)
|
| 1161 |
if pair.dtype != orig_dtype:
|
| 1162 |
-
pair = pair.to(orig_dtype)
|
| 1163 |
-
return pair
|
| 1164 |
|
| 1165 |
|
| 1166 |
# ===========================================================================
|
|
@@ -1201,28 +1219,29 @@ class OuterProductMean(nn.Module):
|
|
| 1201 |
self._chunk_size = chunk_size
|
| 1202 |
|
| 1203 |
def forward(self, m: Tensor, msa_attention_mask: Tensor) -> Tensor:
|
| 1204 |
-
|
| 1205 |
-
|
| 1206 |
-
|
| 1207 |
-
|
| 1208 |
-
|
|
|
|
| 1209 |
if self._chunk_size is None:
|
| 1210 |
-
outer = torch.einsum("bimc,bjmd->bijcd", a, b).flatten(-2)
|
| 1211 |
if self.divide_outer_before_proj:
|
| 1212 |
-
return self.Wout(outer / n_valid)
|
| 1213 |
-
return self.Wout(outer) / n_valid
|
| 1214 |
# Chunk along the left (i) axis so the peak einsum intermediate is
|
| 1215 |
# X uses shape (b, chunk, l, c, d) instead of (b, l, l, c, d).
|
| 1216 |
length = a.shape[1]
|
| 1217 |
out_chunks: list[Tensor] = []
|
| 1218 |
for start in range(0, length, self._chunk_size):
|
| 1219 |
end = min(start + self._chunk_size, length)
|
| 1220 |
-
outer_chunk = torch.einsum("bimc,bjmd->bijcd", a[:, start:end], b).flatten(-2)
|
| 1221 |
if self.divide_outer_before_proj:
|
| 1222 |
out_chunks.append(self.Wout(outer_chunk / n_valid[:, start:end]))
|
| 1223 |
else:
|
| 1224 |
out_chunks.append(self.Wout(outer_chunk) / n_valid[:, start:end])
|
| 1225 |
-
return torch.cat(out_chunks, dim=1)
|
| 1226 |
|
| 1227 |
|
| 1228 |
class MSAPairWeightedAveraging(nn.Module):
|
|
@@ -1252,18 +1271,18 @@ class MSAPairWeightedAveraging(nn.Module):
|
|
| 1252 |
batch_size, length, depth, _ = msa_repr.shape
|
| 1253 |
n_heads, head_width = self.n_heads, self.head_width
|
| 1254 |
|
| 1255 |
-
msa_normed = self.norm_single(msa_repr)
|
| 1256 |
bias = self.compute_bias(pair_repr) # A has shape (b, l, l, n_heads).
|
| 1257 |
-
bias.masked_fill_(~pair_attention_mask.unsqueeze(-1).bool(), -1e5)
|
| 1258 |
-
attn = torch.softmax(bias, dim=-2) # softmax over j
|
| 1259 |
|
| 1260 |
-
v = self.Wv(msa_normed).reshape(batch_size, length, depth, n_heads, head_width)
|
| 1261 |
gate = torch.sigmoid(self.Wgate(msa_normed)).reshape(
|
| 1262 |
batch_size, length, depth, n_heads, head_width
|
| 1263 |
-
)
|
| 1264 |
|
| 1265 |
-
output = torch.einsum("bijh,bjmhd,bimhd->bimhd", attn, v, gate)
|
| 1266 |
-
return self.Wout(output.reshape(batch_size, length, depth, n_heads * head_width))
|
| 1267 |
|
| 1268 |
|
| 1269 |
# ===========================================================================
|
|
@@ -1283,10 +1302,11 @@ class TransitionLayer(nn.Module):
|
|
| 1283 |
self.out_proj = nn.Linear(hidden, d_model, bias=False)
|
| 1284 |
|
| 1285 |
def forward(self, x: Tensor) -> Tensor:
|
| 1286 |
-
x =
|
| 1287 |
-
|
| 1288 |
-
|
| 1289 |
-
|
|
|
|
| 1290 |
|
| 1291 |
|
| 1292 |
# ===========================================================================
|
|
@@ -1302,14 +1322,15 @@ class AdaptiveLayerNorm(nn.Module):
|
|
| 1302 |
self.d_model = d_model
|
| 1303 |
self.d_cond = d_cond
|
| 1304 |
self.eps = eps
|
| 1305 |
-
self.s_scale = nn.Parameter(torch.ones(d_cond))
|
| 1306 |
self.s_gate = nn.Linear(d_cond, d_model, bias=True)
|
| 1307 |
self.s_shift = nn.Linear(d_cond, d_model, bias=False)
|
| 1308 |
|
| 1309 |
def forward(self, a: Tensor, s: Tensor) -> Tensor:
|
| 1310 |
-
|
| 1311 |
-
|
| 1312 |
-
|
|
|
|
| 1313 |
|
| 1314 |
|
| 1315 |
# ===========================================================================
|
|
@@ -1326,12 +1347,13 @@ class FourierEmbedding(nn.Module):
|
|
| 1326 |
def __init__(self, c: int) -> None:
|
| 1327 |
super().__init__()
|
| 1328 |
self.c = c
|
| 1329 |
-
self.register_buffer("w", torch.randn(c))
|
| 1330 |
-
self.register_buffer("b", torch.randn(c))
|
| 1331 |
|
| 1332 |
def forward(self, t_hat: Tensor) -> Tensor:
|
| 1333 |
-
|
| 1334 |
-
|
|
|
|
| 1335 |
|
| 1336 |
|
| 1337 |
# ===========================================================================
|
|
@@ -1360,14 +1382,15 @@ class SwiGLU(nn.Module):
|
|
| 1360 |
self.hidden_features = hidden_features
|
| 1361 |
|
| 1362 |
def forward(self, x: Tensor) -> Tensor:
|
| 1363 |
-
|
| 1364 |
-
|
| 1365 |
-
|
|
|
|
| 1366 |
# Without autograd the product can reuse the activation's buffer. On a pair tensor
|
| 1367 |
# that buffer is twice the pair representation. The values are the same either way.
|
| 1368 |
-
hidden = hidden * x2 if torch.is_grad_enabled() else hidden.mul_(x2)
|
| 1369 |
del x12, x1, x2
|
| 1370 |
-
return self.w3(hidden)
|
| 1371 |
|
| 1372 |
|
| 1373 |
class SwiGLUMLP(SwiGLU):
|
|
@@ -1386,8 +1409,9 @@ class SwiGLUMLP(SwiGLU):
|
|
| 1386 |
|
| 1387 |
|
| 1388 |
def _rotate_half(x: Tensor) -> Tensor:
|
| 1389 |
-
|
| 1390 |
-
|
|
|
|
| 1391 |
|
| 1392 |
|
| 1393 |
def apply_rotary_emb_3d(x: Tensor, cos: Tensor, sin: Tensor) -> Tensor:
|
|
@@ -1399,12 +1423,12 @@ def apply_rotary_emb_3d(x: Tensor, cos: Tensor, sin: Tensor) -> Tensor:
|
|
| 1399 |
sin: S with shape (b, l, d / 2).
|
| 1400 |
"""
|
| 1401 |
ro_dim = cos.shape[-1] * 2
|
| 1402 |
-
cos = cos.unsqueeze(2).repeat(1, 1, 1, 2)
|
| 1403 |
-
sin = sin.unsqueeze(2).repeat(1, 1, 1, 2)
|
| 1404 |
return torch.cat(
|
| 1405 |
[x[..., :ro_dim] * cos + _rotate_half(x[..., :ro_dim]) * sin, x[..., ro_dim:]],
|
| 1406 |
dim=-1,
|
| 1407 |
-
)
|
| 1408 |
|
| 1409 |
|
| 1410 |
@torch.compiler.disable
|
|
@@ -1418,6 +1442,7 @@ def build_3d_rope(
|
|
| 1418 |
uid_base_freq: float = 10.0,
|
| 1419 |
) -> tuple[Tensor, Tensor]:
|
| 1420 |
"""Build cos/sin for 3D RoPE + UID RoPE."""
|
|
|
|
| 1421 |
device = ref_pos.device
|
| 1422 |
batch_size, n_atoms = ref_pos.shape[:2]
|
| 1423 |
half_dim = head_dim // 2
|
|
@@ -1429,21 +1454,21 @@ def build_3d_rope(
|
|
| 1429 |
torch.arange(0, n_spatial_per_axis, dtype=torch.float32, device=device)
|
| 1430 |
/ n_spatial_per_axis
|
| 1431 |
)
|
| 1432 |
-
)
|
| 1433 |
uid_inv_freq = 1.0 / (
|
| 1434 |
uid_base_freq
|
| 1435 |
** (torch.arange(0, n_uid_pairs, dtype=torch.float32, device=device) / n_uid_pairs)
|
| 1436 |
-
)
|
| 1437 |
|
| 1438 |
-
pos_f32 = ref_pos.float()
|
| 1439 |
-
spatial_freqs = torch.einsum("bna,k->bnak", pos_f32, spatial_inv_freq)
|
| 1440 |
-
spatial_freqs = spatial_freqs.reshape(batch_size, n_atoms, n_spatial_total)
|
| 1441 |
|
| 1442 |
-
uid_f32 = ref_space_uid.float()
|
| 1443 |
-
uid_freqs = torch.einsum("bn,k->bnk", uid_f32, uid_inv_freq)
|
| 1444 |
|
| 1445 |
n_active = n_spatial_total + n_uid_pairs
|
| 1446 |
-
freqs = torch.cat([spatial_freqs, uid_freqs], dim=-1)
|
| 1447 |
|
| 1448 |
if n_active < half_dim:
|
| 1449 |
padding = torch.zeros(
|
|
@@ -1452,16 +1477,17 @@ def build_3d_rope(
|
|
| 1452 |
half_dim - n_active,
|
| 1453 |
device=device,
|
| 1454 |
dtype=torch.float32,
|
| 1455 |
-
)
|
| 1456 |
-
freqs = torch.cat([freqs, padding], dim=-1)
|
| 1457 |
|
| 1458 |
-
cos = freqs.cos().to(torch.bfloat16)
|
| 1459 |
-
sin = freqs.sin().to(torch.bfloat16)
|
| 1460 |
-
return cos, sin
|
| 1461 |
|
| 1462 |
|
| 1463 |
def qk_norm(x: Tensor) -> Tensor:
|
| 1464 |
-
|
|
|
|
| 1465 |
|
| 1466 |
|
| 1467 |
# ===========================================================================
|
|
@@ -1479,9 +1505,10 @@ class SwiGLUFFN(nn.Module):
|
|
| 1479 |
self.w_down = nn.Linear(hidden_size, d_model, bias=False)
|
| 1480 |
|
| 1481 |
def forward(self, x: Tensor) -> Tensor:
|
| 1482 |
-
|
| 1483 |
-
|
| 1484 |
-
|
|
|
|
| 1485 |
|
| 1486 |
|
| 1487 |
# ===========================================================================
|
|
@@ -1524,7 +1551,7 @@ class SWA3DRoPEAttention(nn.Module):
|
|
| 1524 |
# indices: (t,) flat positions of real atoms; cu_seqlens: (b + 1,) int32 row offsets.
|
| 1525 |
indices, cu_seqlens, max_seqlen = attention_params[2:5]
|
| 1526 |
flat_shape = (batch_size * n_atoms, self.n_heads, self.head_dim)
|
| 1527 |
-
q, k, v = q.reshape(flat_shape), k.reshape(flat_shape), v.reshape(flat_shape)
|
| 1528 |
has_padding = indices.shape[0] != batch_size * n_atoms
|
| 1529 |
if has_padding:
|
| 1530 |
q, k, v = q[indices], k[indices], v[indices] # each (t, h, d_h)
|
|
@@ -1545,46 +1572,47 @@ class SWA3DRoPEAttention(nn.Module):
|
|
| 1545 |
) # (t, h, d_h)
|
| 1546 |
if has_padding:
|
| 1547 |
out = attended.new_zeros(flat_shape) # (b * n_atoms, h, d_h)
|
| 1548 |
-
out[indices] = attended
|
| 1549 |
else:
|
| 1550 |
-
out = attended
|
| 1551 |
-
return out.view(batch_size, n_atoms, self.n_heads, self.head_dim)
|
| 1552 |
|
| 1553 |
def forward(self, x: Tensor, attention_params: tuple) -> Tensor:
|
|
|
|
| 1554 |
batch_size, n_atoms = x.shape[:2]
|
| 1555 |
-
cos, sin = attention_params[0], attention_params[1]
|
| 1556 |
|
| 1557 |
-
x_input = x
|
| 1558 |
-
qkv = self.Wqkv(x)
|
| 1559 |
-
qkv = qkv.view(batch_size, n_atoms, 3, self.n_heads, self.head_dim).permute(2, 0, 1, 3, 4)
|
| 1560 |
-
q, k, v = qkv.unbind(0)
|
| 1561 |
-
q, k = qk_norm(q), qk_norm(k)
|
| 1562 |
|
| 1563 |
-
q = apply_rotary_emb_3d(q, cos, sin)
|
| 1564 |
-
k = apply_rotary_emb_3d(k, cos, sin)
|
| 1565 |
|
| 1566 |
input_dtype = q.dtype
|
| 1567 |
if q.dtype not in (torch.float16, torch.bfloat16):
|
| 1568 |
-
q, k, v = q.bfloat16(), k.bfloat16(), v.bfloat16()
|
| 1569 |
|
| 1570 |
# ESMFold2 does not advertise FlashAttention. Keep this atom path on
|
| 1571 |
# PyTorch. Models that advertise FlashAttention dispatch through the
|
| 1572 |
# precompiled Hugging Face kernels interface in fastplms.attention.
|
| 1573 |
if self._atom_attention == ATOM_ATTENTION_WINDOWED:
|
| 1574 |
-
out = self._windowed_attention(q, k, v, attention_params)
|
| 1575 |
else:
|
| 1576 |
-
q_t = q.transpose(1, 2)
|
| 1577 |
-
k_t = k.transpose(1, 2)
|
| 1578 |
-
v_t = v.transpose(1, 2)
|
| 1579 |
-
attn = torch.matmul(q_t, k_t.transpose(-2, -1)) * self.scale
|
| 1580 |
-
attn = F.softmax(attn, dim=-1)
|
| 1581 |
-
out = torch.matmul(attn, v_t).transpose(1, 2)
|
| 1582 |
|
| 1583 |
out = out.to(input_dtype).reshape( # type: ignore[union-attr]
|
| 1584 |
batch_size, n_atoms, -1
|
| 1585 |
-
)
|
| 1586 |
-
out = out * torch.sigmoid(self.gate_proj(x_input))
|
| 1587 |
-
return self.out_proj(out)
|
| 1588 |
|
| 1589 |
|
| 1590 |
# ===========================================================================
|
|
@@ -1593,11 +1621,13 @@ class SWA3DRoPEAttention(nn.Module):
|
|
| 1593 |
|
| 1594 |
|
| 1595 |
def _rms_adaln_raw(x: Tensor, scale: Tensor, shift: Tensor) -> Tensor:
|
| 1596 |
-
|
|
|
|
| 1597 |
|
| 1598 |
|
| 1599 |
def _gated_residual_raw(x: Tensor, gate: Tensor, y: Tensor) -> Tensor:
|
| 1600 |
-
|
|
|
|
| 1601 |
|
| 1602 |
|
| 1603 |
class SWAAtomBlock(nn.Module):
|
|
@@ -1619,7 +1649,7 @@ class SWAAtomBlock(nn.Module):
|
|
| 1619 |
self.ffn_norm = nn.RMSNorm(d_atom, elementwise_affine=False)
|
| 1620 |
|
| 1621 |
adaln_linear = nn.Linear(d_atom, 6 * d_atom, bias=False)
|
| 1622 |
-
nn.init.zeros_(adaln_linear.weight)
|
| 1623 |
self.adaln_modulation = nn.Sequential(nn.SiLU(), adaln_linear)
|
| 1624 |
|
| 1625 |
self.attn = SWA3DRoPEAttention(d_atom, n_heads, half_window=half_window)
|
|
@@ -1631,19 +1661,20 @@ class SWAAtomBlock(nn.Module):
|
|
| 1631 |
)
|
| 1632 |
|
| 1633 |
def forward(self, x: Tensor, c_l: Tensor, attention_params: tuple) -> Tensor:
|
| 1634 |
-
|
|
|
|
| 1635 |
if mod.dim() == 2:
|
| 1636 |
-
mod = mod.unsqueeze(1)
|
| 1637 |
-
shift_a, scale_a, gate_a, shift_f, scale_f, gate_f = mod.chunk(6, dim=-1)
|
| 1638 |
|
| 1639 |
-
attn_input = self._rms_adaln(x, scale_a, shift_a)
|
| 1640 |
-
attn_out = self.attn(attn_input, attention_params)
|
| 1641 |
-
x = self._gated_residual(x, gate_a, attn_out)
|
| 1642 |
|
| 1643 |
-
ffn_input = self._rms_adaln(x, scale_f, shift_f)
|
| 1644 |
-
ffn_out = self.ffn(ffn_input)
|
| 1645 |
-
x = self._gated_residual(x, gate_f, ffn_out)
|
| 1646 |
-
return x
|
| 1647 |
|
| 1648 |
|
| 1649 |
class SWAAtomTransformer(nn.Module):
|
|
@@ -1699,14 +1730,15 @@ class SWAAtomTransformer(nn.Module):
|
|
| 1699 |
attention_params: tuple,
|
| 1700 |
return_intermediates: bool = False,
|
| 1701 |
) -> Tensor | tuple[Tensor, list[Tensor]]:
|
|
|
|
| 1702 |
intermediates: list[Tensor] = []
|
| 1703 |
for block in self.blocks:
|
| 1704 |
-
q_l = block(q_l, c_l, attention_params)
|
| 1705 |
if return_intermediates:
|
| 1706 |
intermediates.append(q_l)
|
| 1707 |
if return_intermediates:
|
| 1708 |
-
return q_l, intermediates
|
| 1709 |
-
return q_l
|
| 1710 |
|
| 1711 |
|
| 1712 |
# ===========================================================================
|
|
@@ -1721,13 +1753,14 @@ def _prepare_atom_encoder_metadata(
|
|
| 1721 |
num_diffusion_samples: int,
|
| 1722 |
) -> tuple[Tensor, Tensor, Tensor, int, int]:
|
| 1723 |
"""Prepare mask-derived atom metadata outside compiled diffusion graphs."""
|
| 1724 |
-
|
| 1725 |
-
|
| 1726 |
-
|
|
|
|
| 1727 |
max_seqlen = int(seqlens.max().item())
|
| 1728 |
-
cu_seqlens = F.pad(torch.cumsum(seqlens, dim=0, dtype=torch.int32), (1, 0))
|
| 1729 |
n_tokens = int(atom_to_token.max().item()) + 1
|
| 1730 |
-
return mask_exp, indices, cu_seqlens, max_seqlen, n_tokens
|
| 1731 |
|
| 1732 |
|
| 1733 |
class ESMFold2AtomEncoder(nn.Module):
|
|
@@ -1805,6 +1838,8 @@ class ESMFold2AtomEncoder(nn.Module):
|
|
| 1805 |
``inference_cache`` caches step-invariant tensors (c_base, 3D RoPE,
|
| 1806 |
attention indices, n_tokens) across diffusion steps.
|
| 1807 |
"""
|
|
|
|
|
|
|
| 1808 |
batch_size, n_atoms = ref_pos.shape[:2]
|
| 1809 |
|
| 1810 |
layer_cache = None
|
|
@@ -1821,46 +1856,46 @@ class ESMFold2AtomEncoder(nn.Module):
|
|
| 1821 |
ref_atom_name_chars.reshape(batch_size, n_atoms, MAX_CHARS * CHAR_VOCAB_SIZE),
|
| 1822 |
],
|
| 1823 |
dim=-1,
|
| 1824 |
-
)
|
| 1825 |
-
c_base = self.atom_norm(self.atom_linear(atom_feats))
|
| 1826 |
-
cos, sin = self.atom_transformer._build_3d_rope(ref_pos, ref_space_uid)
|
| 1827 |
-
cos = cos.repeat_interleave(num_diffusion_samples, 0)
|
| 1828 |
-
sin = sin.repeat_interleave(num_diffusion_samples, 0)
|
| 1829 |
mask_exp, indices, cu_seqlens, max_seqlen, n_tokens = (
|
| 1830 |
_prepare_atom_encoder_metadata(
|
| 1831 |
atom_attention_mask,
|
| 1832 |
atom_to_token,
|
| 1833 |
num_diffusion_samples,
|
| 1834 |
)
|
| 1835 |
-
)
|
| 1836 |
attention_params = (cos, sin, indices, cu_seqlens, max_seqlen)
|
| 1837 |
if layer_cache is not None:
|
| 1838 |
-
layer_cache["c_base"] = c_base
|
| 1839 |
layer_cache["attention_params"] = attention_params
|
| 1840 |
-
layer_cache["mask_exp"] = mask_exp
|
| 1841 |
layer_cache["n_tokens"] = n_tokens
|
| 1842 |
layer_cache["atom_to_token_exp"] = atom_to_token.repeat_interleave(
|
| 1843 |
num_diffusion_samples, 0
|
| 1844 |
-
)
|
| 1845 |
else:
|
| 1846 |
-
c_base = layer_cache["c_base"]
|
| 1847 |
attention_params = layer_cache["attention_params"]
|
| 1848 |
-
mask_exp = layer_cache["mask_exp"]
|
| 1849 |
n_tokens = layer_cache["n_tokens"]
|
| 1850 |
|
| 1851 |
-
c = c_base
|
| 1852 |
|
| 1853 |
-
q = c
|
| 1854 |
|
| 1855 |
if self.structure_prediction and r_l is not None:
|
| 1856 |
-
q = q.repeat_interleave(num_diffusion_samples, 0)
|
| 1857 |
if pred_r1 is None:
|
| 1858 |
-
pred_r1 = torch.zeros_like(r_l)
|
| 1859 |
-
r_input = torch.cat([r_l, pred_r1], dim=-1)
|
| 1860 |
-
r_to_q = self.coords_linear(r_input)
|
| 1861 |
-
q = q + r_to_q
|
| 1862 |
|
| 1863 |
-
c = c.repeat_interleave(num_diffusion_samples, 0)
|
| 1864 |
|
| 1865 |
result = self.atom_transformer(
|
| 1866 |
q_l=q,
|
|
@@ -1869,19 +1904,19 @@ class ESMFold2AtomEncoder(nn.Module):
|
|
| 1869 |
return_intermediates=return_intermediates,
|
| 1870 |
)
|
| 1871 |
if return_intermediates:
|
| 1872 |
-
q, intermediates = result
|
| 1873 |
else:
|
| 1874 |
-
q = result
|
| 1875 |
intermediates = []
|
| 1876 |
|
| 1877 |
-
q_to_a = F.relu(self.atom_to_token_linear(q))
|
| 1878 |
if layer_cache is not None and "atom_to_token_exp" in layer_cache:
|
| 1879 |
-
atom_to_token_exp = layer_cache["atom_to_token_exp"]
|
| 1880 |
else:
|
| 1881 |
-
atom_to_token_exp = atom_to_token.repeat_interleave(num_diffusion_samples, 0)
|
| 1882 |
-
a = scatter_atom_to_token(q_to_a, atom_to_token_exp, n_tokens, atom_mask=mask_exp.bool())
|
| 1883 |
|
| 1884 |
-
return a, q, c, attention_params, intermediates
|
| 1885 |
|
| 1886 |
|
| 1887 |
# ===========================================================================
|
|
@@ -1935,10 +1970,11 @@ class ESMFold2AtomDecoder(nn.Module):
|
|
| 1935 |
return_intermediates: bool = False,
|
| 1936 |
) -> tuple[Tensor, list[Tensor]]:
|
| 1937 |
"""Returns (r_update, intermediates)."""
|
| 1938 |
-
|
| 1939 |
-
|
| 1940 |
-
a_to_q =
|
| 1941 |
-
|
|
|
|
| 1942 |
|
| 1943 |
result = self.atom_transformer(
|
| 1944 |
q_l=q_l,
|
|
@@ -1947,13 +1983,13 @@ class ESMFold2AtomDecoder(nn.Module):
|
|
| 1947 |
return_intermediates=return_intermediates,
|
| 1948 |
)
|
| 1949 |
if return_intermediates:
|
| 1950 |
-
q_l, intermediates = result
|
| 1951 |
else:
|
| 1952 |
-
q_l = result
|
| 1953 |
intermediates = []
|
| 1954 |
|
| 1955 |
-
r_l = self.output_linear(self.norm(q_l))
|
| 1956 |
-
return r_l, intermediates
|
| 1957 |
|
| 1958 |
|
| 1959 |
# ===========================================================================
|
|
@@ -1983,8 +2019,8 @@ class AttentionPairBias(nn.Module):
|
|
| 1983 |
self.adaln = AdaptiveLayerNorm(d_model, d_cond, eps=1e-5)
|
| 1984 |
self.out_gate = nn.Linear(d_cond, d_model, bias=True)
|
| 1985 |
# adaln init: weight=0, bias=-2
|
| 1986 |
-
nn.init.zeros_(self.out_gate.weight)
|
| 1987 |
-
nn.init.constant_(self.out_gate.bias, -2.0)
|
| 1988 |
else:
|
| 1989 |
self.pre_norm = nn.LayerNorm(d_model, eps=1e-5)
|
| 1990 |
|
|
@@ -2043,22 +2079,23 @@ class AttentionPairBias(nn.Module):
|
|
| 2043 |
conditions every denoising step on the same ``z``, so the PyTorch path
|
| 2044 |
projects the bias on the first step and reuses that tensor afterwards.
|
| 2045 |
"""
|
|
|
|
| 2046 |
bsz, n_queries, d_model = a.shape
|
| 2047 |
|
| 2048 |
-
x = self.adaln(a, s) if s is not None else self.pre_norm(a)
|
| 2049 |
|
| 2050 |
n_keys = x.shape[1]
|
| 2051 |
-
q = self.q_proj(x).view(bsz, n_queries, self.num_heads, self.head_dim)
|
| 2052 |
-
kv = self.kv_proj(x)
|
| 2053 |
-
k, v = kv.chunk(2, dim=-1)
|
| 2054 |
-
k = k.view(bsz, n_keys, self.num_heads, self.head_dim)
|
| 2055 |
-
v = v.view(bsz, n_keys, self.num_heads, self.head_dim)
|
| 2056 |
|
| 2057 |
use_fused_kernel = self._can_use_fused_pair_bias(z, n_queries, beta)
|
| 2058 |
use_cueq_kernel = not use_fused_kernel and self._can_use_cueq_pair_bias(z, n_queries, beta)
|
| 2059 |
-
cached_pair_bias = None
|
| 2060 |
if step_cache is not None and not use_fused_kernel and not use_cueq_kernel:
|
| 2061 |
-
cached_pair_bias = step_cache.get("pair_bias")
|
| 2062 |
|
| 2063 |
# Expand z for num_diffusion_samples, unless its projection is already cached.
|
| 2064 |
if (
|
|
@@ -2073,21 +2110,21 @@ class AttentionPairBias(nn.Module):
|
|
| 2073 |
and attention_mask.shape[0] != bsz
|
| 2074 |
and num_diffusion_samples > 1
|
| 2075 |
):
|
| 2076 |
-
attention_mask = attention_mask.repeat_interleave(num_diffusion_samples, dim=0)
|
| 2077 |
|
| 2078 |
if use_fused_kernel:
|
| 2079 |
kernel_mask = (
|
| 2080 |
attention_mask
|
| 2081 |
if attention_mask is not None
|
| 2082 |
else torch.ones(bsz, n_queries, device=a.device, dtype=torch.bool)
|
| 2083 |
-
)
|
| 2084 |
-
pair_norm_w = self.pair_norm.weight
|
| 2085 |
pair_norm_b = (
|
| 2086 |
self.pair_norm.bias
|
| 2087 |
if self.pair_norm.bias is not None
|
| 2088 |
else torch.zeros_like(pair_norm_w)
|
| 2089 |
-
)
|
| 2090 |
-
z_bf = z if z.dtype == torch.bfloat16 else z.to(torch.bfloat16)
|
| 2091 |
bias = _fused_pair_bias( # type: ignore[misc]
|
| 2092 |
z_bf,
|
| 2093 |
kernel_mask,
|
|
@@ -2095,26 +2132,26 @@ class AttentionPairBias(nn.Module):
|
|
| 2095 |
num_heads=self.num_heads,
|
| 2096 |
pair_norm_w=pair_norm_w,
|
| 2097 |
pair_norm_b=pair_norm_b,
|
| 2098 |
-
) #
|
| 2099 |
-
q_bhqd = q.transpose(1, 2)
|
| 2100 |
-
k_bhqd = k.transpose(1, 2)
|
| 2101 |
-
v_bhqd = v.transpose(1, 2)
|
| 2102 |
attn_out = F.scaled_dot_product_attention(
|
| 2103 |
q_bhqd, k_bhqd, v_bhqd, attn_mask=bias.to(q_bhqd.dtype)
|
| 2104 |
-
)
|
| 2105 |
-
g = torch.sigmoid(self.g_proj(x)).view(bsz, n_queries, self.num_heads, self.head_dim)
|
| 2106 |
-
ctx = g * attn_out.transpose(1, 2)
|
| 2107 |
-
out = self.out_proj(ctx.reshape(bsz, n_queries, d_model))
|
| 2108 |
if s is not None:
|
| 2109 |
-
out = torch.sigmoid(self.out_gate(s)) * out
|
| 2110 |
-
return out
|
| 2111 |
|
| 2112 |
if use_cueq_kernel:
|
| 2113 |
kernel_mask = (
|
| 2114 |
attention_mask
|
| 2115 |
if attention_mask is not None
|
| 2116 |
else torch.ones(bsz, n_queries, device=a.device, dtype=torch.bool)
|
| 2117 |
-
)
|
| 2118 |
out, _ = _cue_attn_pair_bias( # type: ignore[misc]
|
| 2119 |
s=x,
|
| 2120 |
q=q.transpose(1, 2),
|
|
@@ -2130,36 +2167,36 @@ class AttentionPairBias(nn.Module):
|
|
| 2130 |
b_ln_z=self.pair_norm.bias,
|
| 2131 |
return_z_proj=False,
|
| 2132 |
is_cached_z_proj=False,
|
| 2133 |
-
)
|
| 2134 |
else:
|
| 2135 |
# Standard attention with pair bias
|
| 2136 |
-
g = torch.sigmoid(self.g_proj(x)).view(bsz, n_queries, self.num_heads, self.head_dim)
|
| 2137 |
|
| 2138 |
-
logits = torch.einsum("... i h d, ... j h d -> ... i j h", q, k) * self.scale
|
| 2139 |
|
| 2140 |
if cached_pair_bias is not None:
|
| 2141 |
pair_bias = cached_pair_bias # (b * samples, n, n, h)
|
| 2142 |
elif z.dim() == 4:
|
| 2143 |
pair_bias = self.pair_bias_proj(self.pair_norm(z)) # (b * samples, n, n, h)
|
| 2144 |
if step_cache is not None:
|
| 2145 |
-
step_cache["pair_bias"] = pair_bias
|
| 2146 |
else:
|
| 2147 |
pair_bias = z.unsqueeze(-1) # (b * samples, n, n, 1), a precomputed bias
|
| 2148 |
-
logits = logits + pair_bias.to(dtype=logits.dtype)
|
| 2149 |
|
| 2150 |
if attention_mask is not None:
|
| 2151 |
min_val = torch.finfo(logits.dtype).min
|
| 2152 |
-
mask_bias = torch.where(attention_mask.bool()[:, None, :, None], 0.0, min_val)
|
| 2153 |
-
logits = logits + mask_bias.to(dtype=logits.dtype)
|
| 2154 |
|
| 2155 |
-
attn = torch.softmax(logits, dim=-2).to(dtype=v.dtype)
|
| 2156 |
-
ctx = torch.einsum("... i j h, ... j h d -> ... i h d", attn, v)
|
| 2157 |
-
ctx = g * ctx
|
| 2158 |
-
out = self.out_proj(ctx.reshape(bsz, n_queries, d_model))
|
| 2159 |
|
| 2160 |
if s is not None:
|
| 2161 |
-
out = torch.sigmoid(self.out_gate(s)) * out
|
| 2162 |
-
return out
|
| 2163 |
|
| 2164 |
|
| 2165 |
# ===========================================================================
|
|
@@ -2184,8 +2221,8 @@ class ConditionedTransitionBlock(nn.Module):
|
|
| 2184 |
if use_conditioning:
|
| 2185 |
self.adaln = AdaptiveLayerNorm(d_model, d_cond, eps=1e-5)
|
| 2186 |
self.output_gate = nn.Linear(d_cond, d_model, bias=True)
|
| 2187 |
-
nn.init.zeros_(self.output_gate.weight)
|
| 2188 |
-
nn.init.constant_(self.output_gate.bias, -2.0)
|
| 2189 |
else:
|
| 2190 |
self.pre_norm = nn.LayerNorm(d_model, eps=1e-5)
|
| 2191 |
|
|
@@ -2193,15 +2230,16 @@ class ConditionedTransitionBlock(nn.Module):
|
|
| 2193 |
self.lin_out = nn.Linear(hidden, d_model, bias=False)
|
| 2194 |
|
| 2195 |
def forward(self, a: Tensor, s: Tensor | None) -> Tensor:
|
| 2196 |
-
|
|
|
|
| 2197 |
|
| 2198 |
-
swish_a, swish_b = self.lin_swish(x).chunk(2, dim=-1)
|
| 2199 |
-
b = F.silu(swish_a) * swish_b
|
| 2200 |
-
out = self.lin_out(b)
|
| 2201 |
|
| 2202 |
if s is not None:
|
| 2203 |
-
out = torch.sigmoid(self.output_gate(s)) * out
|
| 2204 |
-
return out
|
| 2205 |
|
| 2206 |
|
| 2207 |
# ===========================================================================
|
|
@@ -2269,11 +2307,12 @@ class DiffusionTransformer(nn.Module):
|
|
| 2269 |
``inference_cache`` must span only calls that share ``z``, as one
|
| 2270 |
``sample`` call does; each block then keeps its pair bias across steps.
|
| 2271 |
"""
|
|
|
|
| 2272 |
intermediates: list[Tensor] = []
|
| 2273 |
block_caches: dict[int, dict[str, Tensor]] | None = None
|
| 2274 |
if inference_cache is not None:
|
| 2275 |
block_caches = inference_cache.setdefault("token_pair_bias", {})
|
| 2276 |
-
x = a
|
| 2277 |
for block_index, (attn, transition) in enumerate(
|
| 2278 |
zip(self.attn_blocks, self.transition_blocks, strict=True)
|
| 2279 |
):
|
|
@@ -2286,11 +2325,11 @@ class DiffusionTransformer(nn.Module):
|
|
| 2286 |
attention_mask=attention_mask,
|
| 2287 |
num_diffusion_samples=num_diffusion_samples,
|
| 2288 |
step_cache=step_cache,
|
| 2289 |
-
)
|
| 2290 |
-
x = x + transition(x, s)
|
| 2291 |
if return_intermediates:
|
| 2292 |
intermediates.append(x)
|
| 2293 |
-
return x, intermediates
|
| 2294 |
|
| 2295 |
|
| 2296 |
# ===========================================================================
|
|
@@ -2343,45 +2382,46 @@ class DiffusionConditioning(nn.Module):
|
|
| 2343 |
num_diffusion_samples: int = 1,
|
| 2344 |
inference_cache: dict[str, Tensor] | None = None,
|
| 2345 |
) -> tuple[Tensor, Tensor]:
|
|
|
|
| 2346 |
sigma = self.sigma_data if sigma_data is None else float(sigma_data)
|
| 2347 |
base_batch = z_trunk.shape[0]
|
| 2348 |
target_batch = base_batch * num_diffusion_samples
|
| 2349 |
|
| 2350 |
# z conditioning (cached across diffusion steps: independent of t_hat)
|
| 2351 |
if inference_cache is not None and "z" in inference_cache:
|
| 2352 |
-
z = inference_cache["z"]
|
| 2353 |
else:
|
| 2354 |
-
z_rel = relative_position_encoding.to(dtype=torch.float32)
|
| 2355 |
-
z = torch.cat([z_trunk.to(dtype=torch.float32), z_rel], dim=-1)
|
| 2356 |
-
z = self.z_proj(self.z_input_norm(z))
|
| 2357 |
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
|
| 2358 |
for block in self.z_transitions:
|
| 2359 |
-
z = z + block(z)
|
| 2360 |
if inference_cache is not None:
|
| 2361 |
-
inference_cache["z"] = z
|
| 2362 |
|
| 2363 |
# s conditioning
|
| 2364 |
-
s_inputs_eff = s_inputs
|
| 2365 |
if s_inputs_eff.shape[0] != target_batch:
|
| 2366 |
-
s_inputs_eff = s_inputs_eff.repeat_interleave(num_diffusion_samples, 0)
|
| 2367 |
|
| 2368 |
-
s = self.s_proj(self.s_input_norm(s_inputs_eff.to(dtype=torch.float32)))
|
| 2369 |
|
| 2370 |
# Noise embedding
|
| 2371 |
-
t = torch.as_tensor(t_hat, dtype=torch.float32, device=s.device).reshape(-1)
|
| 2372 |
if t.numel() == 1:
|
| 2373 |
-
t = t.expand(target_batch)
|
| 2374 |
elif t.shape[0] != target_batch:
|
| 2375 |
-
t = t.repeat_interleave(num_diffusion_samples, 0)
|
| 2376 |
-
t_noise = 0.25 * torch.log((t / sigma).clamp(min=1e-20))
|
| 2377 |
-
n = self.fourier(t_noise)
|
| 2378 |
-
n = self.noise_proj(self.noise_norm(n))
|
| 2379 |
-
s = s + n.unsqueeze(1)
|
| 2380 |
|
| 2381 |
for block in self.s_transitions:
|
| 2382 |
-
s = s + block(s)
|
| 2383 |
|
| 2384 |
-
return s, z
|
| 2385 |
|
| 2386 |
|
| 2387 |
# ===========================================================================
|
|
@@ -2453,7 +2493,7 @@ class DiffusionModule(nn.Module):
|
|
| 2453 |
)
|
| 2454 |
|
| 2455 |
self.s_to_token = nn.Linear(c_token, c_token, bias=False)
|
| 2456 |
-
nn.init.zeros_(self.s_to_token.weight)
|
| 2457 |
|
| 2458 |
# Token transformer (DiffusionTransformer with pair bias)
|
| 2459 |
self.token_transformer = DiffusionTransformer(
|
|
@@ -2499,11 +2539,12 @@ class DiffusionModule(nn.Module):
|
|
| 2499 |
return_atom_repr: bool = False,
|
| 2500 |
inference_cache: dict[str, Tensor] | None = None,
|
| 2501 |
) -> dict[str, Tensor | None]:
|
|
|
|
| 2502 |
bsz = x_noisy.shape[0]
|
| 2503 |
sigma = self.sigma_data if sigma_data is None else float(sigma_data)
|
| 2504 |
-
t = torch.as_tensor(t_hat, dtype=torch.float32, device=x_noisy.device).reshape(-1)
|
| 2505 |
if t.numel() == 1:
|
| 2506 |
-
t = t.expand(bsz)
|
| 2507 |
|
| 2508 |
# Step 1: conditioning (pair z is cached across diffusion steps)
|
| 2509 |
s, z = self.conditioning(
|
|
@@ -2515,11 +2556,11 @@ class DiffusionModule(nn.Module):
|
|
| 2515 |
sigma_data=sigma,
|
| 2516 |
num_diffusion_samples=num_diffusion_samples,
|
| 2517 |
inference_cache=inference_cache,
|
| 2518 |
-
)
|
| 2519 |
|
| 2520 |
# Step 2: normalize noisy coords
|
| 2521 |
-
denom = torch.sqrt(t * t + sigma * sigma)
|
| 2522 |
-
r_noisy = x_noisy / denom[:, None, None]
|
| 2523 |
|
| 2524 |
# Step 3: atom encoder
|
| 2525 |
a, q_skip, c_skip, p_skip, enc_intermediates = self.atom_encoder(
|
|
@@ -2535,10 +2576,10 @@ class DiffusionModule(nn.Module):
|
|
| 2535 |
num_diffusion_samples=num_diffusion_samples,
|
| 2536 |
return_intermediates=return_atom_repr,
|
| 2537 |
inference_cache=inference_cache,
|
| 2538 |
-
)
|
| 2539 |
|
| 2540 |
# Step 4: add conditioned s
|
| 2541 |
-
a = a + self.s_to_token(self.s_step_norm(s))
|
| 2542 |
|
| 2543 |
# Step 5: token transformer
|
| 2544 |
a, _ = self.token_transformer(
|
|
@@ -2549,10 +2590,10 @@ class DiffusionModule(nn.Module):
|
|
| 2549 |
attention_mask=token_attention_mask,
|
| 2550 |
num_diffusion_samples=num_diffusion_samples,
|
| 2551 |
inference_cache=inference_cache,
|
| 2552 |
-
)
|
| 2553 |
|
| 2554 |
# Step 6: token norm
|
| 2555 |
-
a = self.token_norm(a)
|
| 2556 |
|
| 2557 |
# Step 7: atom decoder
|
| 2558 |
r_update, dec_intermediates = self.atom_decoder(
|
|
@@ -2564,26 +2605,26 @@ class DiffusionModule(nn.Module):
|
|
| 2564 |
atom_attention_mask=ref_mask,
|
| 2565 |
num_diffusion_samples=num_diffusion_samples,
|
| 2566 |
return_intermediates=return_atom_repr,
|
| 2567 |
-
)
|
| 2568 |
|
| 2569 |
# Step 8: compute denoised output
|
| 2570 |
sigma2 = sigma * sigma
|
| 2571 |
-
t2 = t * t
|
| 2572 |
-
out = (sigma2 / (sigma2 + t2))[:, None, None] * x_noisy
|
| 2573 |
-
out = out + ((sigma * t) / torch.sqrt(sigma2 + t2))[:, None, None] * r_update
|
| 2574 |
|
| 2575 |
# Collect atom intermediates from encoder + decoder
|
| 2576 |
-
atom_intermediates: Tensor | None = None
|
| 2577 |
if return_atom_repr:
|
| 2578 |
all_ints = enc_intermediates + dec_intermediates
|
| 2579 |
if all_ints:
|
| 2580 |
-
atom_intermediates = torch.stack(all_ints, dim=2)
|
| 2581 |
|
| 2582 |
return {
|
| 2583 |
"x_denoised": out,
|
| 2584 |
"token_repr": a if return_token_repr else None,
|
| 2585 |
"atom_intermediates": atom_intermediates,
|
| 2586 |
-
}
|
| 2587 |
|
| 2588 |
|
| 2589 |
# ===========================================================================
|
|
@@ -2647,24 +2688,24 @@ class DiffusionStructureHead(nn.Module):
|
|
| 2647 |
[self.inference_s_max * self.sigma_data, 0.0],
|
| 2648 |
device=device,
|
| 2649 |
dtype=torch.float32,
|
| 2650 |
-
)
|
| 2651 |
p = float(self.inference_p)
|
| 2652 |
inv_p = 1.0 / p
|
| 2653 |
-
k = torch.arange(steps, device=device, dtype=torch.float32)
|
| 2654 |
base = self.inference_s_max**inv_p + (k / (steps - 1)) * (
|
| 2655 |
self.inference_s_min**inv_p - self.inference_s_max**inv_p
|
| 2656 |
-
)
|
| 2657 |
-
schedule = self.sigma_data * base.pow(p)
|
| 2658 |
-
return F.pad(schedule, (0, 1), value=0.0)
|
| 2659 |
|
| 2660 |
@staticmethod
|
| 2661 |
def _random_rotations(n: int, dtype: torch.dtype, device: torch.device) -> Tensor:
|
| 2662 |
-
q = torch.randn((n, 4), dtype=dtype, device=device)
|
| 2663 |
-
scale = torch.sqrt((q * q).sum(dim=1))
|
| 2664 |
-
signs = torch.where(q[:, 0] < 0, -scale, scale)
|
| 2665 |
-
q = q / signs[:, None]
|
| 2666 |
-
r, i, j, k = torch.unbind(q, dim=-1)
|
| 2667 |
-
two_s = 2.0 / (q * q).sum(dim=-1)
|
| 2668 |
return torch.stack(
|
| 2669 |
(
|
| 2670 |
1 - two_s * (j * j + k * k),
|
|
@@ -2678,51 +2719,53 @@ class DiffusionStructureHead(nn.Module):
|
|
| 2678 |
1 - two_s * (i * i + j * j),
|
| 2679 |
),
|
| 2680 |
dim=-1,
|
| 2681 |
-
).reshape(n, 3, 3)
|
| 2682 |
|
| 2683 |
def _center_random_augmentation(
|
| 2684 |
self, x: Tensor, atom_mask: Tensor, second_coords: Tensor | None = None
|
| 2685 |
) -> tuple[Tensor, Tensor | None]:
|
| 2686 |
"""Algorithm 19: center + random rotation + translation."""
|
|
|
|
| 2687 |
bsz = x.shape[0]
|
| 2688 |
mask = atom_mask.unsqueeze(-1) # M has shape (b, a, 1).
|
| 2689 |
-
denom = mask.sum(dim=1, keepdim=True).clamp(min=1)
|
| 2690 |
-
mean = (x * mask).sum(dim=1, keepdim=True) / denom
|
| 2691 |
-
x = x - mean
|
| 2692 |
if second_coords is not None:
|
| 2693 |
-
second_coords = second_coords - mean
|
| 2694 |
|
| 2695 |
-
r = self._random_rotations(bsz, x.dtype, x.device)
|
| 2696 |
-
x = torch.einsum("bmd,bds->bms", x, r)
|
| 2697 |
if second_coords is not None:
|
| 2698 |
-
second_coords = torch.einsum("bmd,bds->bms", second_coords, r)
|
| 2699 |
|
| 2700 |
-
t = torch.randn_like(x[:, 0:1, :])
|
| 2701 |
-
x = x + t
|
| 2702 |
if second_coords is not None:
|
| 2703 |
-
second_coords = second_coords + t
|
| 2704 |
-
return x, second_coords
|
| 2705 |
|
| 2706 |
@staticmethod
|
| 2707 |
def _weighted_rigid_align(x: Tensor, x_gt: Tensor, w: Tensor, mask: Tensor) -> Tensor:
|
| 2708 |
"""Kabsch alignment: align x to x_gt with weights w."""
|
|
|
|
| 2709 |
w = (mask * w).unsqueeze(-1) # W has shape (b, n, 1).
|
| 2710 |
-
denom = w.sum(dim=-2, keepdim=True).clamp(min=1e-8)
|
| 2711 |
-
mu = (x * w).sum(dim=-2, keepdim=True) / denom
|
| 2712 |
-
mu_gt = (x_gt * w).sum(dim=-2, keepdim=True) / denom
|
| 2713 |
-
x_c = x - mu
|
| 2714 |
-
xgt_c = x_gt - mu_gt
|
| 2715 |
-
covariance = torch.einsum("bni,bnj->bij", w * xgt_c, x_c)
|
| 2716 |
-
covariance_f32 = covariance.float()
|
| 2717 |
u, _, vh = torch.linalg.svd(
|
| 2718 |
covariance_f32, driver="gesvd" if covariance_f32.is_cuda else None
|
| 2719 |
-
)
|
| 2720 |
-
det = torch.linalg.det(u @ vh)
|
| 2721 |
-
ones = torch.ones_like(det)
|
| 2722 |
rotation = (u @ torch.diag_embed(torch.stack([ones, ones, det], dim=-1)) @ vh).to(
|
| 2723 |
covariance.dtype
|
| 2724 |
-
)
|
| 2725 |
-
return x_c @ rotation.transpose(-1, -2) + mu_gt
|
| 2726 |
|
| 2727 |
# ------------------------------------------------------------------
|
| 2728 |
# Sampling
|
|
@@ -2766,6 +2809,7 @@ class DiffusionStructureHead(nn.Module):
|
|
| 2766 |
so we inflate the underlying schedule length here to land back at the
|
| 2767 |
requested step count post-truncation.
|
| 2768 |
"""
|
|
|
|
| 2769 |
n_atoms = tok_idx.shape[1]
|
| 2770 |
device = s_inputs.device
|
| 2771 |
target_batch = s_inputs.shape[0] * num_diffusion_samples
|
|
@@ -2774,26 +2818,26 @@ class DiffusionStructureHead(nn.Module):
|
|
| 2774 |
|
| 2775 |
steps = self.inference_num_steps if num_sampling_steps is None else int(num_sampling_steps)
|
| 2776 |
|
| 2777 |
-
schedule = self.inference_noise_schedule(steps, device)
|
| 2778 |
if max_inference_sigma is not None:
|
| 2779 |
-
schedule = schedule[schedule <= float(max_inference_sigma)]
|
| 2780 |
-
schedule = F.pad(schedule, (1, 0), value=float(max_inference_sigma))
|
| 2781 |
|
| 2782 |
lam = self.noise_scale if noise_scale is None else float(noise_scale)
|
| 2783 |
eta = self.step_scale if step_scale is None else float(step_scale)
|
| 2784 |
|
| 2785 |
-
x = schedule[0] * torch.randn(target_batch, n_atoms, 3, device=device, dtype=torch.float32)
|
| 2786 |
-
atom_mask = ref_mask.repeat_interleave(num_diffusion_samples, 0).float()
|
| 2787 |
|
| 2788 |
gammas = torch.where(
|
| 2789 |
schedule > self.gamma_min,
|
| 2790 |
torch.full_like(schedule, self.gamma_0),
|
| 2791 |
torch.zeros_like(schedule),
|
| 2792 |
-
)
|
| 2793 |
|
| 2794 |
-
x_denoised_prev: Tensor | None = None
|
| 2795 |
-
token_repr: Tensor | None = None
|
| 2796 |
-
diff_atom_intermediates: Tensor | None = None
|
| 2797 |
|
| 2798 |
step_pairs = list(zip(schedule[:-1], schedule[1:], gammas[1:], strict=True))
|
| 2799 |
num_steps = len(step_pairs)
|
|
@@ -2810,12 +2854,12 @@ class DiffusionStructureHead(nn.Module):
|
|
| 2810 |
for step_idx, (sigma_tm, sigma_t, gamma) in enumerate(step_iterator):
|
| 2811 |
x, x_denoised_prev = self._center_random_augmentation(
|
| 2812 |
x, atom_mask, second_coords=x_denoised_prev
|
| 2813 |
-
)
|
| 2814 |
|
| 2815 |
sigma_tm_val = float(sigma_tm.item())
|
| 2816 |
t_hat_val = sigma_tm_val * (1.0 + float(gamma.item()))
|
| 2817 |
eps_std = lam * max(t_hat_val**2 - sigma_tm_val**2, 0.0) ** 0.5
|
| 2818 |
-
x_noisy = x + eps_std * torch.randn_like(x)
|
| 2819 |
|
| 2820 |
is_last_step = step_idx == num_steps - 1
|
| 2821 |
request_atom_repr = return_atom_repr and (
|
|
@@ -2846,24 +2890,24 @@ class DiffusionStructureHead(nn.Module):
|
|
| 2846 |
return_token_repr=True,
|
| 2847 |
return_atom_repr=request_atom_repr,
|
| 2848 |
inference_cache=inference_cache,
|
| 2849 |
-
)
|
| 2850 |
|
| 2851 |
-
x_denoised = dm_out["x_denoised"]
|
| 2852 |
-
token_repr = dm_out["token_repr"]
|
| 2853 |
if request_atom_repr:
|
| 2854 |
-
diff_atom_intermediates = dm_out.get("atom_intermediates")
|
| 2855 |
|
| 2856 |
# Reverse diffusion alignment (Kabsch)
|
| 2857 |
with torch.autocast(device_type="cuda", enabled=False):
|
| 2858 |
x_noisy = self._weighted_rigid_align(
|
| 2859 |
x_noisy.float(), x_denoised.float(), atom_mask, atom_mask
|
| 2860 |
-
)
|
| 2861 |
-
x_noisy = x_noisy.to(dtype=x_denoised.dtype)
|
| 2862 |
|
| 2863 |
# ODE/SDE step
|
| 2864 |
sigma_t_val = float(sigma_t.item())
|
| 2865 |
-
denoised_over_sigma = (x_noisy - x_denoised) / t_hat_val
|
| 2866 |
-
x = x_noisy + eta * (sigma_t_val - t_hat_val) * denoised_over_sigma
|
| 2867 |
|
| 2868 |
# Denoising early-exit: stop when consecutive predictions converge
|
| 2869 |
if (
|
|
@@ -2877,17 +2921,17 @@ class DiffusionStructureHead(nn.Module):
|
|
| 2877 |
x_denoised.float(),
|
| 2878 |
atom_mask,
|
| 2879 |
atom_mask,
|
| 2880 |
-
)
|
| 2881 |
-
diff = (x_denoised.float() - aligned) * atom_mask.unsqueeze(-1)
|
| 2882 |
per_sample_rmsd = (
|
| 2883 |
diff.pow(2).sum(dim=(-1, -2)) / atom_mask.sum(dim=-1).clamp(min=1)
|
| 2884 |
-
).sqrt()
|
| 2885 |
if per_sample_rmsd.max().item() < denoising_early_exit_rmsd:
|
| 2886 |
-
x = x_denoised
|
| 2887 |
-
x_denoised_prev = x_denoised
|
| 2888 |
break
|
| 2889 |
|
| 2890 |
-
x_denoised_prev = x_denoised
|
| 2891 |
|
| 2892 |
result: dict[str, Tensor | None] = {
|
| 2893 |
"sample_atom_coords": x,
|
|
@@ -2895,4 +2939,4 @@ class DiffusionStructureHead(nn.Module):
|
|
| 2895 |
}
|
| 2896 |
if return_atom_repr:
|
| 2897 |
result["diff_atom_intermediates"] = diff_atom_intermediates
|
| 2898 |
-
return result
|
|
|
|
| 10 |
from __future__ import annotations
|
| 11 |
|
| 12 |
import importlib
|
|
|
|
|
|
|
|
|
|
|
|
|
| 13 |
import torch
|
| 14 |
import torch.nn as nn
|
| 15 |
import torch.nn.functional as F
|
| 16 |
+
|
| 17 |
+
from functools import partial
|
| 18 |
+
from importlib.util import find_spec
|
| 19 |
+
from typing import Any, ClassVar, cast
|
| 20 |
from torch import Tensor
|
| 21 |
from torch.utils.checkpoint import checkpoint
|
| 22 |
from tqdm.auto import tqdm
|
|
|
|
| 24 |
from .configuration_esmfold2 import ESMFold2Config
|
| 25 |
from .reproducibility import seed_context
|
| 26 |
|
| 27 |
+
|
| 28 |
_seed_context = seed_context
|
| 29 |
|
| 30 |
try:
|
|
|
|
| 198 |
self._impl = nn.Dropout(r)
|
| 199 |
|
| 200 |
def forward(self, residual: Tensor, delta: Tensor) -> Tensor:
|
| 201 |
+
# residual, delta: same shape; dropout shares delta axis self._batch_dim.
|
| 202 |
if self._use_fused_kernels:
|
| 203 |
+
return self._impl(residual, delta) # delta.shape
|
| 204 |
+
# The unfused mask shares the selected row/column axis and retains the other dimensions.
|
| 205 |
if self._r == 0.0 or not self.training:
|
| 206 |
+
return residual + delta # delta.shape
|
| 207 |
shape = list(delta.shape)
|
| 208 |
shape[self._batch_dim] = 1
|
| 209 |
+
mask = self._impl(delta.new_ones(shape)) # delta.shape with shared axis set to 1
|
| 210 |
+
return residual + delta * mask # delta.shape
|
| 211 |
|
| 212 |
|
| 213 |
# ---------------------------------------------------------------------------
|
|
|
|
| 219 |
MAX_ATOMIC_NUMBER: int = 128
|
| 220 |
|
| 221 |
# Input feature dim = 3 + 1 + 1 + 128 + 64*4 = 389
|
| 222 |
+
ATOM_FEATURE_DIM: int = XYZ_DIMS + 1 + 1 + MAX_ATOMIC_NUMBER + CHAR_VOCAB_SIZE * MAX_CHARS # d_atom_features
|
| 223 |
|
| 224 |
|
| 225 |
NUM_RES_TYPES: int = 33
|
|
|
|
| 248 |
max_depth: int | None,
|
| 249 |
enabled: bool,
|
| 250 |
) -> tuple[Tensor, Tensor | None, Tensor | None, Tensor | None]:
|
| 251 |
+
# MSA tensors: (b, m, l); k = max_depth for subsampled rows.
|
| 252 |
if not enabled or max_depth is None:
|
| 253 |
+
return msa, msa_attention_mask, has_deletion, deletion_value # MSA tensors retain (b, selected_depth, l); optional tensors remain None
|
| 254 |
|
| 255 |
depth = msa.size(1)
|
| 256 |
if depth <= 1 or depth <= max_depth:
|
| 257 |
+
return msa, msa_attention_mask, has_deletion, deletion_value # MSA tensors retain (b, selected_depth, l); optional tensors remain None
|
| 258 |
|
| 259 |
+
indices = torch.zeros(max_depth, dtype=torch.long, device=msa.device) # (k,)
|
| 260 |
+
indices[1:] = torch.randperm(depth - 1, device=msa.device)[: max_depth - 1] + 1 # (k - 1,)
|
| 261 |
+
indices = indices.sort().values # (k,)
|
| 262 |
|
| 263 |
+
msa = msa[:, indices] # (b, k, l)
|
| 264 |
if msa_attention_mask is not None:
|
| 265 |
+
msa_attention_mask = msa_attention_mask[:, indices] # (b, k, l)
|
| 266 |
if has_deletion is not None:
|
| 267 |
+
has_deletion = has_deletion[:, indices] # (b, k, l)
|
| 268 |
if deletion_value is not None:
|
| 269 |
+
deletion_value = deletion_value[:, indices] # (b, k, l)
|
| 270 |
+
return msa, msa_attention_mask, has_deletion, deletion_value # MSA tensors retain (b, selected_depth, l); optional tensors remain None
|
| 271 |
|
| 272 |
|
| 273 |
def maybe_apply_msa_column_masking(
|
| 274 |
msa_attention_mask: Tensor | None,
|
| 275 |
rate: float,
|
| 276 |
) -> Tensor | None:
|
| 277 |
+
# msa_attention_mask: (b, m, l) or None.
|
| 278 |
if msa_attention_mask is None or rate <= 0.0 or msa_attention_mask.size(1) <= 1:
|
| 279 |
+
return msa_attention_mask # (b, m, l) or None
|
| 280 |
|
| 281 |
batch_size, _, length = msa_attention_mask.shape
|
| 282 |
+
col_keep = torch.rand(batch_size, length, device=msa_attention_mask.device) >= rate # (b, l)
|
| 283 |
+
col_keep = col_keep.unsqueeze(1).expand_as(msa_attention_mask).clone() # (b, m, l)
|
| 284 |
+
col_keep[:, 0, :] = True # (b, l)
|
| 285 |
+
return msa_attention_mask.bool() & col_keep # (b, m, l) or None
|
| 286 |
|
| 287 |
|
| 288 |
# ===========================================================================
|
|
|
|
| 300 |
Returns:
|
| 301 |
X with shape (b, a, d).
|
| 302 |
"""
|
| 303 |
+
idx = atom_to_token_idx.unsqueeze(-1).expand(-1, -1, token_features.size(-1)) # (b, a, d)
|
| 304 |
+
return torch.gather(token_features, 1, idx) # (b, a, d)
|
| 305 |
|
| 306 |
|
| 307 |
def scatter_atom_to_token(
|
|
|
|
| 323 |
"""
|
| 324 |
batch_size, n_atoms, d_model = atom_features.shape
|
| 325 |
n_out = n_tokens
|
| 326 |
+
idx = atom_to_token_idx # (b, a)
|
| 327 |
if atom_mask is not None:
|
| 328 |
+
idx = torch.where(atom_mask, atom_to_token_idx, n_tokens) # (b, a)
|
| 329 |
n_out = n_tokens + 1
|
| 330 |
+
idx_expanded = idx.unsqueeze(-1).expand(batch_size, n_atoms, d_model) # (b, a, d)
|
| 331 |
out = torch.zeros(
|
| 332 |
batch_size,
|
| 333 |
n_out,
|
| 334 |
d_model,
|
| 335 |
device=atom_features.device,
|
| 336 |
dtype=atom_features.dtype,
|
| 337 |
+
) # (b, n_out, d)
|
| 338 |
+
out.scatter_reduce_(1, idx_expanded, atom_features, reduce="mean", include_self=False) # (b, n_out, d)
|
| 339 |
+
return out[:, :n_tokens, :] # (b, l, d)
|
| 340 |
|
| 341 |
|
| 342 |
def gather_rep_atom_coords(coords: Tensor, rep_atom_idx: Tensor) -> Tensor:
|
|
|
|
| 349 |
Returns:
|
| 350 |
X with shape (b, l, 3).
|
| 351 |
"""
|
| 352 |
+
idx = rep_atom_idx.unsqueeze(-1).expand(-1, -1, coords.size(-1)) # (b, l, 3)
|
| 353 |
+
return torch.gather(coords, 1, idx) # (b, l, 3)
|
| 354 |
|
| 355 |
|
| 356 |
def _compute_intra_token_idx(atom_to_token: Tensor) -> Tensor:
|
|
|
|
| 366 |
Index tensor I with shape (b, a) and values from zero through
|
| 367 |
``max_atoms_per_token - 1``.
|
| 368 |
"""
|
| 369 |
+
same_as_prev = F.pad(atom_to_token[:, 1:] == atom_to_token[:, :-1], (1, 0), value=False) # (b, a)
|
| 370 |
+
ones = torch.ones_like(atom_to_token) # (b, a)
|
| 371 |
+
cumsum = torch.cumsum(ones, dim=-1) # (b, a)
|
| 372 |
+
group_start = cumsum.masked_fill(same_as_prev, 0) # (b, a)
|
| 373 |
+
group_start = torch.cummax(group_start, dim=-1).values # (b, a)
|
| 374 |
+
return cumsum - group_start # (b, a)
|
| 375 |
|
| 376 |
|
| 377 |
def _categorical_mean(logits: Tensor, start: float, end: float) -> Tensor:
|
|
|
|
| 388 |
Expected value tensor Y with shape (...).
|
| 389 |
"""
|
| 390 |
n_bins = logits.shape[-1]
|
| 391 |
+
edges = torch.linspace(start, end, n_bins + 1, device=logits.device, dtype=torch.float32) # (n_bins + 1,)
|
| 392 |
v_bins = (edges[:-1] + edges[1:]) / 2 # V_bin has shape (n_bins,).
|
| 393 |
+
return (logits.float().softmax(-1) @ v_bins.unsqueeze(1)).squeeze(-1) # logits.shape[:-1]
|
| 394 |
|
| 395 |
|
| 396 |
# ===========================================================================
|
|
|
|
| 407 |
self.out_proj = nn.Linear(d_pair, d_single, bias=False)
|
| 408 |
|
| 409 |
def forward(self, z: Tensor, mask: Tensor) -> Tensor:
|
| 410 |
+
# z: (b, l, l, d_pair); mask: (b, l).
|
| 411 |
+
scores = self.attn_proj(z).squeeze(-1) # (b, l, l)
|
| 412 |
mask_bias = torch.where(
|
| 413 |
mask[:, None, :].bool(),
|
| 414 |
torch.zeros_like(scores),
|
| 415 |
torch.full_like(scores, -1e9),
|
| 416 |
+
) # (b, l, l)
|
| 417 |
+
scores = scores + mask_bias # (b, l, l)
|
| 418 |
+
weights = F.softmax(scores, dim=-1) # (b, l, l)
|
| 419 |
+
pooled = torch.einsum("bnm,bnmd->bnd", weights, z) # (b, l, d_pair)
|
| 420 |
+
return self.out_proj(pooled) # (b, l, d_single)
|
| 421 |
|
| 422 |
|
| 423 |
# ===========================================================================
|
|
|
|
| 465 |
X with shape (b, l, d_inputs), concatenating atom encoding,
|
| 466 |
aatype, profile, and deletion mean.
|
| 467 |
"""
|
| 468 |
+
# aatype/profile: (b, l, 33); deletion_mean: (b, l); atom features use a atoms.
|
| 469 |
a, _q, _c, _attn_params, _intermediates = self.atom_attention_encoder(
|
| 470 |
ref_pos=ref_pos,
|
| 471 |
atom_attention_mask=atom_attention_mask,
|
|
|
|
| 474 |
ref_element=ref_element,
|
| 475 |
ref_atom_name_chars=ref_atom_name_chars,
|
| 476 |
atom_to_token=atom_to_token,
|
| 477 |
+
) # a: (b, l, d_token / 2); _q/_c: (b, a, d_atom)
|
| 478 |
+
return torch.cat([a, aatype, profile, deletion_mean.unsqueeze(-1)], dim=-1) # (b, l, d_token / 2 + 67)
|
| 479 |
|
| 480 |
|
| 481 |
# ===========================================================================
|
|
|
|
| 517 |
entity_id: Tensor,
|
| 518 |
token_index: Tensor,
|
| 519 |
) -> Tensor:
|
| 520 |
+
# All input IDs: (b, l); r/c are relative residue/chain bin counts.
|
| 521 |
+
bij_same_chain = asym_id.unsqueeze(2) == asym_id.unsqueeze(1) # (b, l, l)
|
| 522 |
+
bij_same_residue = residue_index.unsqueeze(2) == residue_index.unsqueeze(1) # (b, l, l)
|
| 523 |
+
bij_same_entity = entity_id.unsqueeze(2) == entity_id.unsqueeze(1) # (b, l, l)
|
| 524 |
|
| 525 |
+
dij_residue = residue_index.unsqueeze(2) - residue_index.unsqueeze(1) # (b, l, l)
|
| 526 |
dij_residue = torch.clip(
|
| 527 |
dij_residue + self.n_relative_residx_bins,
|
| 528 |
0,
|
| 529 |
2 * self.n_relative_residx_bins,
|
| 530 |
+
) # (b, l, l)
|
| 531 |
+
dij_residue = torch.where(bij_same_chain, dij_residue, 2 * self.n_relative_residx_bins + 1) # (b, l, l)
|
| 532 |
+
aij_rel_pos = F.one_hot(dij_residue, 2 * self.n_relative_residx_bins + 2) # (b, l, l, 2 * r + 2)
|
| 533 |
|
| 534 |
dij_token = torch.clip(
|
| 535 |
token_index.unsqueeze(2) - token_index.unsqueeze(1) + self.n_relative_residx_bins,
|
| 536 |
0,
|
| 537 |
2 * self.n_relative_residx_bins,
|
| 538 |
+
) # (b, l, l)
|
| 539 |
dij_token = torch.where(
|
| 540 |
bij_same_chain & bij_same_residue,
|
| 541 |
dij_token,
|
| 542 |
2 * self.n_relative_residx_bins + 1,
|
| 543 |
+
) # (b, l, l)
|
| 544 |
+
aij_rel_token = F.one_hot(dij_token, 2 * self.n_relative_residx_bins + 2) # (b, l, l, 2 * r + 2)
|
| 545 |
|
| 546 |
dij_chain = torch.clip(
|
| 547 |
sym_id.unsqueeze(2) - sym_id.unsqueeze(1) + self.n_relative_chain_bins,
|
| 548 |
0,
|
| 549 |
2 * self.n_relative_chain_bins,
|
| 550 |
+
) # (b, l, l)
|
| 551 |
+
dij_chain = torch.where(bij_same_chain, 2 * self.n_relative_chain_bins + 1, dij_chain) # (b, l, l)
|
| 552 |
+
aij_rel_chain = F.one_hot(dij_chain, 2 * self.n_relative_chain_bins + 2) # (b, l, l, 2 * c + 2)
|
| 553 |
|
| 554 |
feats = torch.cat(
|
| 555 |
[
|
|
|
|
| 559 |
aij_rel_chain.float(),
|
| 560 |
],
|
| 561 |
dim=-1,
|
| 562 |
+
) # (b, l, l, 2 * (2 * r + 2) + 1 + 2 * c + 2)
|
| 563 |
|
| 564 |
+
return self.embed(feats) # (b, l, l, d_pair)
|
| 565 |
|
| 566 |
|
| 567 |
# ===========================================================================
|
|
|
|
| 582 |
)
|
| 583 |
|
| 584 |
def forward(self, x: Tensor) -> Tensor:
|
| 585 |
+
# x: (b, l, input_dim); d_down is downproject.out_features.
|
| 586 |
+
x = self.downproject(x) # (b, l, d_down)
|
| 587 |
x = torch.cat(
|
| 588 |
[(x.unsqueeze(2) * x.unsqueeze(1)), (x.unsqueeze(2) - x.unsqueeze(1))],
|
| 589 |
dim=3,
|
| 590 |
+
) # (b, l, l, 2 * d_down)
|
| 591 |
+
return self.output_mlp(x) # (b, l, l, output_dim)
|
| 592 |
|
| 593 |
|
| 594 |
# ===========================================================================
|
|
|
|
| 612 |
self.base_z_linear = nn.Sequential(
|
| 613 |
nn.LayerNorm(d_model), nn.Linear(d_model, d_z, bias=False)
|
| 614 |
)
|
| 615 |
+
self.base_z_combine = nn.Parameter(torch.zeros(num_layers + 1)) # (num_layers + 1,)
|
| 616 |
|
| 617 |
def project_sequence(
|
| 618 |
self,
|
|
|
|
| 649 |
# Match the learned projection parameters at this explicit boundary;
|
| 650 |
# this preserves the official BF16 path and leaves FP32 models exact.
|
| 651 |
projection_dtype = cast(nn.LayerNorm, self.base_z_linear[0]).weight.dtype
|
| 652 |
+
hidden_states = hidden_states.to(dtype=projection_dtype) # (b, l, n_layers + 1, d_model)
|
| 653 |
+
projected_states = self.base_z_linear(hidden_states) # (b, l, n_layers + 1, d_z)
|
| 654 |
+
layer_weights = self.base_z_combine.softmax(dim=0) # (n_layers + 1,)
|
| 655 |
# Preserve Biohub's matmul path exactly so checkpoint inference does
|
| 656 |
# not change through a different reduction order.
|
| 657 |
+
projected = layer_weights @ projected_states # (b, l, d_z)
|
| 658 |
if residue_mask is not None:
|
| 659 |
if residue_mask.shape != hidden_states.shape[:2]:
|
| 660 |
raise ValueError(
|
|
|
|
| 663 |
)
|
| 664 |
projected = projected * residue_mask.to(
|
| 665 |
device=projected.device, dtype=projected.dtype
|
| 666 |
+
).unsqueeze(-1) # (b, l, d_z)
|
| 667 |
+
return projected # (b, l, d_z)
|
| 668 |
|
| 669 |
def forward(self, hidden_states: Tensor, *, lm_dropout: float = 0.0) -> Tensor:
|
| 670 |
"""Project pre-computed ESMC hidden states to pair representation.
|
|
|
|
| 677 |
Returns:
|
| 678 |
Z_pair with shape ``(b, l, l, d_pair)``.
|
| 679 |
"""
|
| 680 |
+
lm_z = self.project_sequence(hidden_states) # (b, l, d_z)
|
| 681 |
+
lm_z = self.base_z_mlp(lm_z) # (b, l, l, d_z)
|
| 682 |
if lm_dropout > 0:
|
| 683 |
+
lm_z = F.dropout(lm_z, p=lm_dropout, training=True) # (b, l, l, d_z)
|
| 684 |
+
return lm_z # (b, l, l, d_z)
|
| 685 |
|
| 686 |
|
| 687 |
# ===========================================================================
|
|
|
|
| 708 |
was trained on per-residue inputs, not per-atom), then scatter the
|
| 709 |
hidden states back to the per-token layout.
|
| 710 |
"""
|
| 711 |
+
# Input IDs/masks: (b, l). Per sample p protein tokens collapse to u residues; t is padded LM length.
|
| 712 |
b_size, l_size = input_ids.shape
|
| 713 |
device = input_ids.device
|
| 714 |
+
protein_mask = (mol_type == 0) & token_mask # (b, l)
|
| 715 |
|
| 716 |
lm_input_list = []
|
| 717 |
lm_lengths = []
|
| 718 |
# Per-batch maps from (original protein-token index) to (LM input position).
|
| 719 |
expand_maps: list[Tensor] = []
|
| 720 |
for batch_index in range(b_size):
|
| 721 |
+
mask_b = protein_mask[batch_index] # (l,)
|
| 722 |
+
ids_b = input_ids[batch_index][mask_b] # (p,)
|
| 723 |
+
asym_b = asym_id[batch_index][mask_b] # (p,)
|
| 724 |
+
res_b = residue_index[batch_index][mask_b] # (p,)
|
| 725 |
|
| 726 |
# Collapse: keep first token per (asym_id, residue_index) key, in
|
| 727 |
# input order. ``inverse`` maps each original protein-token to its
|
| 728 |
# collapsed residue index.
|
| 729 |
+
keys = torch.stack((asym_b, res_b), dim=1) # (p, 2)
|
| 730 |
+
unique_keys, inverse = torch.unique(keys, dim=0, return_inverse=True) # (u, 2), (p,)
|
| 731 |
n_unique = unique_keys.size(0)
|
| 732 |
+
token_positions = torch.arange(keys.size(0), device=device, dtype=torch.long) # (p,)
|
| 733 |
+
first_pos = torch.full((n_unique,), keys.size(0), device=device, dtype=torch.long) # (u,)
|
| 734 |
+
first_pos.scatter_reduce_(0, inverse, token_positions, reduce="amin", include_self=True) # (u,)
|
| 735 |
+
ordered = torch.argsort(first_pos) # (u,)
|
| 736 |
+
first_pos_ordered = first_pos[ordered] # (u,)
|
| 737 |
+
ids_collapsed = ids_b[first_pos_ordered] # (u,)
|
| 738 |
+
asym_collapsed = asym_b[first_pos_ordered] # (u,)
|
| 739 |
+
remap = torch.empty_like(ordered) # (u,)
|
| 740 |
+
remap[ordered] = torch.arange(n_unique, device=device, dtype=torch.long) # (u,)
|
| 741 |
+
inverse_ordered = remap[inverse] # (p,)
|
| 742 |
+
|
| 743 |
+
chain_ids = asym_collapsed.unique(sorted=True) # (n_chains,)
|
| 744 |
# [BOS] chain1 [EOS BOS] chain2 ... [EOS]
|
| 745 |
+
parts: list[Tensor] = [torch.tensor([0], device=device, dtype=ids_b.dtype)] # list of 1D token tensors
|
| 746 |
# Per-chain LM positions accumulate; track them for the expand map.
|
| 747 |
+
per_token_lm_pos = torch.empty(n_unique, device=device, dtype=torch.long) # (u,)
|
| 748 |
cursor = 1 # position 0 is the leading BOS
|
| 749 |
for i, cid in enumerate(chain_ids):
|
| 750 |
+
in_chain = (asym_collapsed == cid).nonzero(as_tuple=True)[0] # (u_chain,)
|
| 751 |
parts.append(ids_collapsed[in_chain])
|
| 752 |
per_token_lm_pos[in_chain] = torch.arange(
|
| 753 |
cursor, cursor + in_chain.shape[0], device=device, dtype=torch.long
|
| 754 |
+
) # (u_chain,)
|
| 755 |
cursor += in_chain.shape[0]
|
| 756 |
if i < len(chain_ids) - 1:
|
| 757 |
parts.append(torch.tensor([2, 0], device=device, dtype=ids_b.dtype))
|
| 758 |
cursor += 2 # EOS + BOS
|
| 759 |
parts.append(torch.tensor([2], device=device, dtype=ids_b.dtype))
|
| 760 |
+
lm_seq = torch.cat(parts) # (t_i,)
|
| 761 |
lm_input_list.append(lm_seq)
|
| 762 |
lm_lengths.append(lm_seq.shape[0])
|
| 763 |
|
| 764 |
# Map each original protein-token position to its LM input position.
|
| 765 |
+
prot_pos_b = mask_b.nonzero(as_tuple=True)[0] # (p,)
|
| 766 |
+
expand_map = torch.full((l_size,), -1, device=device, dtype=torch.long) # (l,)
|
| 767 |
+
expand_map[prot_pos_b] = per_token_lm_pos[inverse_ordered] # (p,)
|
| 768 |
expand_maps.append(expand_map)
|
| 769 |
|
| 770 |
# Pad the language-model input to its longest sequence. FP8 callers round
|
|
|
|
| 777 |
1,
|
| 778 |
device=device,
|
| 779 |
dtype=input_ids.dtype, # PAD=1
|
| 780 |
+
) # (b, t)
|
| 781 |
for batch_index in range(b_size):
|
| 782 |
+
lm_input_ids[batch_index, : lm_lengths[batch_index]] = lm_input_list[batch_index] # (t_i,)
|
| 783 |
|
| 784 |
# sequence_id for chain-aware attention; PAD tokens get -1 (no attention).
|
| 785 |
+
sequence_id = (lm_input_ids == 0).cumsum(dim=1) - 1 # BOS=0; (b, t)
|
| 786 |
+
sequence_id = sequence_id.masked_fill(lm_input_ids == 1, -1) # PAD=1; (b, t)
|
| 787 |
|
| 788 |
if lm_mask_pct > 0.0:
|
| 789 |
+
special = (lm_input_ids == 0) | (lm_input_ids == 1) | (lm_input_ids == 2) # (b, t)
|
| 790 |
+
do_mask = (torch.rand(lm_input_ids.shape, device=device) < lm_mask_pct) & ~special # (b, t)
|
| 791 |
+
lm_input_ids = lm_input_ids.masked_fill(do_mask, mask_token_id) # (b, t)
|
| 792 |
|
| 793 |
with torch.inference_mode():
|
| 794 |
esmc_out = esmc(input_ids=lm_input_ids, sequence_id=sequence_id, output_hidden_states=True)
|
| 795 |
|
| 796 |
+
hidden_stack = esmc_out.hidden_states # (n_states, b, t, d_model)
|
| 797 |
n_states, _, _, d_model = hidden_stack.shape
|
| 798 |
+
result = torch.zeros(b_size, l_size, n_states, d_model, device=device, dtype=hidden_stack.dtype) # (b, l, n_states, d_model)
|
| 799 |
for batch_index in range(b_size):
|
| 800 |
+
M_i = protein_mask[batch_index] # (l,)
|
| 801 |
+
positions = expand_maps[batch_index][M_i] # (p,)
|
| 802 |
+
gathered = hidden_stack[:, batch_index, positions, :].permute(1, 0, 2) # (p, n_states, d_model)
|
| 803 |
+
result[batch_index, M_i.nonzero(as_tuple=True)[0]] = gathered # (p, n_states, d_model)
|
| 804 |
|
| 805 |
+
return result.detach() # (b, l, n_states, d_model)
|
| 806 |
|
| 807 |
|
| 808 |
# ===========================================================================
|
|
|
|
| 853 |
return self.flow
|
| 854 |
|
| 855 |
def _triangular_contract(self, left_stream: Tensor, right_stream: Tensor) -> Tensor:
|
| 856 |
+
# Streams: (b, l, l, d_latent); equation chooses incoming/outgoing contraction.
|
| 857 |
+
return torch.einsum(self._einsum_equation, left_stream, right_stream) # (b, l, l, d_latent)
|
| 858 |
|
| 859 |
def _triangular_contract_chunked(
|
| 860 |
self, left_stream: Tensor, right_stream: Tensor, chunk_size: int
|
|
|
|
| 879 |
chunks = []
|
| 880 |
for start in range(0, length, chunk_size):
|
| 881 |
rows = left_rows[:, :, start : start + chunk_size] # (b, d, i_c, k)
|
| 882 |
+
product = torch.bmm(rows.reshape(batch_size * channels, -1, inner), right_columns) # (b * d, i_c, j)
|
| 883 |
product = product.view(batch_size, channels, rows.shape[2], -1) # (b, d, i_c, j)
|
| 884 |
chunks.append(product.permute(0, 2, 3, 1)) # (b, i_c, j, d)
|
| 885 |
+
return torch.cat(chunks, dim=1) # (b, l, l, d)
|
| 886 |
|
| 887 |
def forward(self, pair_grid: Tensor, visibility: Tensor | None = None) -> Tensor:
|
| 888 |
+
# pair_grid: (b, l, l, d_input); visibility: (b, l, l); d = latent_channels.
|
| 889 |
if visibility is None:
|
| 890 |
+
visibility = pair_grid.new_ones(pair_grid.shape[:-1]) # (b, l, l)
|
| 891 |
|
| 892 |
if self._use_kernels:
|
| 893 |
+
p_in_weight, g_in_weight = self.split_kernel_weights() # each (2 * d, d_input)
|
| 894 |
return _cue_tri_mul( # type: ignore[misc]
|
| 895 |
pair_grid,
|
| 896 |
direction=self._kernel_flow_direction(),
|
|
|
|
| 904 |
p_out_weight=self.proj_emit.weight,
|
| 905 |
g_out_weight=self.proj_gate.weight,
|
| 906 |
eps=_EPS,
|
| 907 |
+
) # (b, l, l, d_input)
|
| 908 |
|
| 909 |
# Every tensor below is as large as the pair representation or larger, and this
|
| 910 |
# block sets the peak memory of a fold. Each name is dropped once it is dead, so
|
| 911 |
# the allocator can reuse its buffer. No value changes.
|
| 912 |
+
normalized_grid = self.norm_start(pair_grid) # (b, l, l, d_input)
|
| 913 |
bundled = self.proj_bundle(normalized_grid) # (b, l, l, 4 * d)
|
| 914 |
+
signal, gate_logits = bundled.split(2 * self.latent_channels, dim=-1) # each (b, l, l, 2 * d)
|
| 915 |
routed = signal * torch.sigmoid(gate_logits) # (b, l, l, 2 * d)
|
| 916 |
# The two views would keep the whole projection alive.
|
| 917 |
del bundled, signal, gate_logits
|
| 918 |
+
routed = routed * visibility.unsqueeze(-1) # (b, l, l, 2 * d)
|
| 919 |
|
| 920 |
left_stream, right_stream = routed.float().chunk(2, dim=-1) # each (b, l, l, d)
|
| 921 |
if torch.is_autocast_enabled(left_stream.device.type):
|
| 922 |
# The contraction is an autocast operation. Casting its inputs here, as it
|
| 923 |
# would, lets the full-precision product go before the contraction runs.
|
| 924 |
autocast_dtype = torch.get_autocast_dtype(left_stream.device.type)
|
| 925 |
+
left_stream = left_stream.to(autocast_dtype) # (b, l, l, d)
|
| 926 |
+
right_stream = right_stream.to(autocast_dtype) # (b, l, l, d)
|
| 927 |
del routed
|
| 928 |
if self._chunk_size is not None:
|
| 929 |
contracted = self._triangular_contract_chunked(
|
| 930 |
left_stream, right_stream, self._chunk_size
|
| 931 |
+
) # (b, l, l, d)
|
| 932 |
else:
|
| 933 |
+
contracted = self._triangular_contract(left_stream, right_stream) # (b, l, l, d)
|
| 934 |
del left_stream, right_stream
|
| 935 |
+
mixed = self.proj_emit(self.norm_mix(contracted)) # (b, l, l, d_input)
|
| 936 |
del contracted
|
| 937 |
+
output_gate = torch.sigmoid(self.proj_gate(normalized_grid)) # (b, l, l, d_input)
|
| 938 |
+
return mixed * output_gate # (b, l, l, d_input)
|
| 939 |
|
| 940 |
|
| 941 |
class TriangleMultiplicativeUpdate(nn.Module):
|
|
|
|
| 958 |
self._engine.set_chunk_size(chunk_size)
|
| 959 |
|
| 960 |
def forward(self, z: Tensor, mask: Tensor | None = None) -> Tensor:
|
| 961 |
+
# z: (b, l, l, d_pair); mask: (b, l, l) or None.
|
| 962 |
+
return self._engine(z, visibility=mask) # (b, l, l, d_pair)
|
| 963 |
|
| 964 |
|
| 965 |
# ===========================================================================
|
|
|
|
| 1001 |
dtype=dtype,
|
| 1002 |
)
|
| 1003 |
with torch.no_grad():
|
| 1004 |
+
fused.LN_W.copy_(self.norm.weight) # (d_model,)
|
| 1005 |
if has_ln_bias:
|
| 1006 |
fused.LN_B.copy_(self.norm.bias) # type: ignore[union-attr]
|
| 1007 |
# FusedLNLinearSwiGLU.W12 is (d_model, 2*d_inner); transpose nn.Linear once.
|
| 1008 |
+
fused.W12.copy_(self.ffn.w12.weight.t().contiguous()) # (d_model, 2 * d_inner)
|
| 1009 |
self._fused_swiglu = fused.eval().requires_grad_(False)
|
| 1010 |
else:
|
| 1011 |
self._fused_swiglu = None
|
|
|
|
| 1017 |
|
| 1018 |
def _swiglu_pre_w3(self, x_normed: Tensor) -> Tensor:
|
| 1019 |
"""SwiGLU through silu(x1)*x2, before the final w3."""
|
| 1020 |
+
# x_normed: (..., d_model); d_inner = ffn.hidden_features.
|
| 1021 |
ffn = self.ffn
|
| 1022 |
+
x12 = ffn.w12(x_normed) # (..., 2 * d_inner)
|
| 1023 |
+
x1, x2 = x12.split(ffn.hidden_features, dim=-1) # each (..., d_inner)
|
| 1024 |
+
return F.silu(x1) * x2 # (..., d_inner)
|
| 1025 |
|
| 1026 |
def _addmm_residual(self, x: Tensor, hidden: Tensor) -> Tensor:
|
| 1027 |
"""x + w3(hidden) via single cuBLAS addmm: avoids transition-output allocation."""
|
| 1028 |
+
# x: (..., d_model); hidden: (..., d_inner).
|
| 1029 |
ffn = self.ffn
|
| 1030 |
x_shape = x.shape
|
| 1031 |
out = torch.addmm(
|
| 1032 |
x.contiguous().view(-1, x_shape[-1]),
|
| 1033 |
hidden.view(-1, hidden.shape[-1]),
|
| 1034 |
ffn.w3.weight.t(),
|
| 1035 |
+
) # (product(x.shape[:-1]), d_model)
|
| 1036 |
+
return out.view(x_shape) # x.shape
|
| 1037 |
|
| 1038 |
def forward(self, x: Tensor) -> Tensor:
|
| 1039 |
# Inference-only fast path (addmm-fused residual + pre-alloc out)
|
| 1040 |
#: diverges bit-exactly from ``x + ffn(norm(x))`` so we only use
|
| 1041 |
# it when grad is disabled (binder-design / bit-exact tests run
|
| 1042 |
# with grad on and need the reference path).
|
| 1043 |
+
# x: (b, l, ..., d_model); chunk width l_c <= _chunk_size.
|
| 1044 |
if not torch.is_grad_enabled() and self._can_use_fused_path(x):
|
| 1045 |
fused = self._fused_swiglu
|
| 1046 |
assert fused is not None
|
| 1047 |
pre_w3 = fused
|
| 1048 |
if self._chunk_size is None or x.shape[1] <= self._chunk_size:
|
| 1049 |
+
hidden = pre_w3(x) # (b, l, ..., d_inner)
|
| 1050 |
+
return self._addmm_residual(x, hidden) # x.shape
|
| 1051 |
+
out = torch.empty_like(x) # x.shape
|
| 1052 |
for s in range(0, x.shape[1], self._chunk_size):
|
| 1053 |
e = min(s + self._chunk_size, x.shape[1])
|
| 1054 |
+
sl = x[:, s:e] # (b, l_c, ..., d_model)
|
| 1055 |
+
hidden = pre_w3(sl) # (b, l_c, ..., d_inner)
|
| 1056 |
+
out[:, s:e] = self._addmm_residual(sl, hidden) # (b, l_c, ..., d_model)
|
| 1057 |
+
return out # x.shape
|
| 1058 |
# Reference path: bit-exact with main: x + ffn(norm(x)).
|
| 1059 |
if self._chunk_size is None or x.shape[1] <= self._chunk_size:
|
| 1060 |
+
return x + self.ffn(self.norm(x)) # x.shape
|
| 1061 |
out_list: list[Tensor] = []
|
| 1062 |
for s in range(0, x.shape[1], self._chunk_size):
|
| 1063 |
e = min(s + self._chunk_size, x.shape[1])
|
| 1064 |
+
sl = x[:, s:e] # (b, l_c, ..., d_model)
|
| 1065 |
out_list.append(sl + self.ffn(self.norm(sl)))
|
| 1066 |
+
return torch.cat(out_list, dim=1) # x.shape
|
| 1067 |
|
| 1068 |
|
| 1069 |
class PairUpdateBlock(nn.Module):
|
|
|
|
| 1102 |
self, pair: Tensor, direction: str, pair_attention_mask: Tensor | None
|
| 1103 |
) -> Tensor:
|
| 1104 |
"""Fused TriMul+residual call; weights from the corresponding engine."""
|
| 1105 |
+
# pair: (b, l, l, d_pair); pair_attention_mask: (b, l, l) or None.
|
| 1106 |
tri = self.tri_mul_out if direction == "outgoing" else self.tri_mul_in
|
| 1107 |
engine: TriangleMultiplicativeBlock = tri._engine # type: ignore[assignment]
|
| 1108 |
+
p_in_weight, g_in_weight = engine.split_kernel_weights() # each (2 * d_pair, d_pair)
|
| 1109 |
|
| 1110 |
def _bf16(t: Tensor) -> Tensor:
|
| 1111 |
+
return t if t.dtype == torch.bfloat16 else t.to(torch.bfloat16) # t.shape
|
| 1112 |
|
| 1113 |
return _fused_trimul_with_residual( # type: ignore[misc]
|
| 1114 |
pair,
|
|
|
|
| 1125 |
g_out_weight=_bf16(engine.proj_gate.weight),
|
| 1126 |
mask=pair_attention_mask,
|
| 1127 |
eps=_EPS,
|
| 1128 |
+
) # pair.shape
|
| 1129 |
|
| 1130 |
def forward(self, pair: Tensor, pair_attention_mask: Tensor | None = None) -> Tensor:
|
| 1131 |
+
# pair: (b, l, l, d_pair); pair_attention_mask: (b, l, l) or None.
|
| 1132 |
if self._can_use_fused_trimul_with_residual(pair):
|
| 1133 |
+
pair = self._fused_trimul_with_residual(pair, "outgoing", pair_attention_mask) # (b, l, l, d_pair)
|
| 1134 |
+
pair = self._fused_trimul_with_residual(pair, "incoming", pair_attention_mask) # (b, l, l, d_pair)
|
| 1135 |
else:
|
| 1136 |
+
pair = self.row_drop(pair, self.tri_mul_out(pair, mask=pair_attention_mask)) # (b, l, l, d_pair)
|
| 1137 |
+
pair = self.row_drop(pair, self.tri_mul_in(pair, mask=pair_attention_mask)) # (b, l, l, d_pair)
|
| 1138 |
+
pair = self.pair_transition(pair) # (b, l, l, d_pair)
|
| 1139 |
+
return pair # (b, l, l, d_pair)
|
| 1140 |
|
| 1141 |
|
| 1142 |
class FoldingTrunk(nn.Module):
|
|
|
|
| 1162 |
def forward(self, pair: Tensor, pair_attention_mask: Tensor | None = None) -> Tensor:
|
| 1163 |
# Cast the pair tensor to BF16 when the fused triangle backend is enabled
|
| 1164 |
# (its bwd kernel requires bf16). Other backends keep the input dtype.
|
| 1165 |
+
# pair: (b, l, l, d_pair); pair_attention_mask: (b, l, l) or None.
|
| 1166 |
orig_dtype = pair.dtype
|
| 1167 |
fused_on = (
|
| 1168 |
len(self.blocks) > 0
|
| 1169 |
and getattr(self.blocks[0], "_kernel_backend", None) == BACKEND_FUSED
|
| 1170 |
)
|
| 1171 |
if pair.is_cuda and fused_on and orig_dtype != torch.bfloat16:
|
| 1172 |
+
pair = pair.to(torch.bfloat16) # (b, l, l, d_pair)
|
| 1173 |
for block in self.blocks:
|
| 1174 |
fn = partial(block, pair_attention_mask=pair_attention_mask)
|
| 1175 |
if torch.is_grad_enabled():
|
| 1176 |
+
pair = checkpoint(fn, pair, use_reentrant=False) # pyright: ignore; (b, l, l, d_pair)
|
| 1177 |
else:
|
| 1178 |
+
pair = fn(pair) # (b, l, l, d_pair)
|
| 1179 |
if pair.dtype != orig_dtype:
|
| 1180 |
+
pair = pair.to(orig_dtype) # (b, l, l, d_pair)
|
| 1181 |
+
return pair # (b, l, l, d_pair)
|
| 1182 |
|
| 1183 |
|
| 1184 |
# ===========================================================================
|
|
|
|
| 1219 |
self._chunk_size = chunk_size
|
| 1220 |
|
| 1221 |
def forward(self, m: Tensor, msa_attention_mask: Tensor) -> Tensor:
|
| 1222 |
+
# m: (b, l, m_depth, d_msa); msa_attention_mask: (b, l, m_depth); d_h = d_hidden.
|
| 1223 |
+
m_norm = self.norm(m) # (b, l, m_depth, d_msa)
|
| 1224 |
+
x = self.W(m_norm) * msa_attention_mask.unsqueeze(-1).to(m_norm.dtype) # (b, l, m_depth, 2 * d_h)
|
| 1225 |
+
a, b = x.chunk(2, dim=-1) # each (batch, l, m_depth, d_h)
|
| 1226 |
+
mask_f = msa_attention_mask.to(a.dtype) # (batch, l, m_depth)
|
| 1227 |
+
n_valid = (mask_f @ mask_f.transpose(-1, -2)).unsqueeze(-1).clamp(min=1.0) # (batch, l, l, 1)
|
| 1228 |
if self._chunk_size is None:
|
| 1229 |
+
outer = torch.einsum("bimc,bjmd->bijcd", a, b).flatten(-2) # (batch, l, l, d_h * d_h)
|
| 1230 |
if self.divide_outer_before_proj:
|
| 1231 |
+
return self.Wout(outer / n_valid) # (batch, l, l, d_pair)
|
| 1232 |
+
return self.Wout(outer) / n_valid # (batch, l, l, d_pair)
|
| 1233 |
# Chunk along the left (i) axis so the peak einsum intermediate is
|
| 1234 |
# X uses shape (b, chunk, l, c, d) instead of (b, l, l, c, d).
|
| 1235 |
length = a.shape[1]
|
| 1236 |
out_chunks: list[Tensor] = []
|
| 1237 |
for start in range(0, length, self._chunk_size):
|
| 1238 |
end = min(start + self._chunk_size, length)
|
| 1239 |
+
outer_chunk = torch.einsum("bimc,bjmd->bijcd", a[:, start:end], b).flatten(-2) # (batch, l_c, l, d_h * d_h)
|
| 1240 |
if self.divide_outer_before_proj:
|
| 1241 |
out_chunks.append(self.Wout(outer_chunk / n_valid[:, start:end]))
|
| 1242 |
else:
|
| 1243 |
out_chunks.append(self.Wout(outer_chunk) / n_valid[:, start:end])
|
| 1244 |
+
return torch.cat(out_chunks, dim=1) # (batch, l, l, d_pair)
|
| 1245 |
|
| 1246 |
|
| 1247 |
class MSAPairWeightedAveraging(nn.Module):
|
|
|
|
| 1271 |
batch_size, length, depth, _ = msa_repr.shape
|
| 1272 |
n_heads, head_width = self.n_heads, self.head_width
|
| 1273 |
|
| 1274 |
+
msa_normed = self.norm_single(msa_repr) # (b, l, m, d_msa)
|
| 1275 |
bias = self.compute_bias(pair_repr) # A has shape (b, l, l, n_heads).
|
| 1276 |
+
bias.masked_fill_(~pair_attention_mask.unsqueeze(-1).bool(), -1e5) # (b, l, l, n_heads)
|
| 1277 |
+
attn = torch.softmax(bias, dim=-2) # softmax over j; (b, l, l, n_heads)
|
| 1278 |
|
| 1279 |
+
v = self.Wv(msa_normed).reshape(batch_size, length, depth, n_heads, head_width) # (b, l, m, n_heads, head_width)
|
| 1280 |
gate = torch.sigmoid(self.Wgate(msa_normed)).reshape(
|
| 1281 |
batch_size, length, depth, n_heads, head_width
|
| 1282 |
+
) # (b, l, m, n_heads, head_width)
|
| 1283 |
|
| 1284 |
+
output = torch.einsum("bijh,bjmhd,bimhd->bimhd", attn, v, gate) # (b, l, m, n_heads, head_width)
|
| 1285 |
+
return self.Wout(output.reshape(batch_size, length, depth, n_heads * head_width)) # (b, l, m, d_msa)
|
| 1286 |
|
| 1287 |
|
| 1288 |
# ===========================================================================
|
|
|
|
| 1302 |
self.out_proj = nn.Linear(hidden, d_model, bias=False)
|
| 1303 |
|
| 1304 |
def forward(self, x: Tensor) -> Tensor:
|
| 1305 |
+
# x: (..., d_model); d_hidden = n * d_model.
|
| 1306 |
+
x = self.norm(x) # (..., d_model)
|
| 1307 |
+
a = self.a_proj(x) # (..., d_hidden)
|
| 1308 |
+
b = self.b_proj(x) # (..., d_hidden)
|
| 1309 |
+
return self.out_proj(F.silu(a) * b) # (..., d_model)
|
| 1310 |
|
| 1311 |
|
| 1312 |
# ===========================================================================
|
|
|
|
| 1322 |
self.d_model = d_model
|
| 1323 |
self.d_cond = d_cond
|
| 1324 |
self.eps = eps
|
| 1325 |
+
self.s_scale = nn.Parameter(torch.ones(d_cond)) # (d_cond,)
|
| 1326 |
self.s_gate = nn.Linear(d_cond, d_model, bias=True)
|
| 1327 |
self.s_shift = nn.Linear(d_cond, d_model, bias=False)
|
| 1328 |
|
| 1329 |
def forward(self, a: Tensor, s: Tensor) -> Tensor:
|
| 1330 |
+
# a: (..., d_model); s: (..., d_cond) with broadcast-compatible leading axes.
|
| 1331 |
+
a_norm = F.layer_norm(a, (self.d_model,), None, None, self.eps) # a.shape
|
| 1332 |
+
s_norm = F.layer_norm(s, (self.d_cond,), self.s_scale, None, self.eps) # s.shape
|
| 1333 |
+
return torch.sigmoid(self.s_gate(s_norm)) * a_norm + self.s_shift(s_norm) # broadcast leading shape + (d_model,)
|
| 1334 |
|
| 1335 |
|
| 1336 |
# ===========================================================================
|
|
|
|
| 1347 |
def __init__(self, c: int) -> None:
|
| 1348 |
super().__init__()
|
| 1349 |
self.c = c
|
| 1350 |
+
self.register_buffer("w", torch.randn(c)) # (c,)
|
| 1351 |
+
self.register_buffer("b", torch.randn(c)) # (c,)
|
| 1352 |
|
| 1353 |
def forward(self, t_hat: Tensor) -> Tensor:
|
| 1354 |
+
# t_hat: scalar or arbitrary noise-time tensor; n = t_hat.numel().
|
| 1355 |
+
t = torch.as_tensor(t_hat, device=self.w.device, dtype=self.w.dtype).reshape(-1) # (n,)
|
| 1356 |
+
return torch.cos(2.0 * torch.pi * (t[:, None] * self.w[None, :] + self.b[None, :])) # (n, c)
|
| 1357 |
|
| 1358 |
|
| 1359 |
# ===========================================================================
|
|
|
|
| 1382 |
self.hidden_features = hidden_features
|
| 1383 |
|
| 1384 |
def forward(self, x: Tensor) -> Tensor:
|
| 1385 |
+
# x: (..., in_features); d_hidden = hidden_features.
|
| 1386 |
+
x12 = self.w12(x) # (..., 2 * d_hidden)
|
| 1387 |
+
x1, x2 = x12.split(self.hidden_features, dim=-1) # each (..., d_hidden)
|
| 1388 |
+
hidden = F.silu(x1) # (..., d_hidden)
|
| 1389 |
# Without autograd the product can reuse the activation's buffer. On a pair tensor
|
| 1390 |
# that buffer is twice the pair representation. The values are the same either way.
|
| 1391 |
+
hidden = hidden * x2 if torch.is_grad_enabled() else hidden.mul_(x2) # (..., d_hidden)
|
| 1392 |
del x12, x1, x2
|
| 1393 |
+
return self.w3(hidden) # (..., out_features)
|
| 1394 |
|
| 1395 |
|
| 1396 |
class SwiGLUMLP(SwiGLU):
|
|
|
|
| 1409 |
|
| 1410 |
|
| 1411 |
def _rotate_half(x: Tensor) -> Tensor:
|
| 1412 |
+
# x: (..., d_rot), with an even final width.
|
| 1413 |
+
x1, x2 = x.chunk(2, dim=-1) # each (..., d_rot / 2)
|
| 1414 |
+
return torch.cat((-x2, x1), dim=-1) # x.shape
|
| 1415 |
|
| 1416 |
|
| 1417 |
def apply_rotary_emb_3d(x: Tensor, cos: Tensor, sin: Tensor) -> Tensor:
|
|
|
|
| 1423 |
sin: S with shape (b, l, d / 2).
|
| 1424 |
"""
|
| 1425 |
ro_dim = cos.shape[-1] * 2
|
| 1426 |
+
cos = cos.unsqueeze(2).repeat(1, 1, 1, 2) # (b, l, 1, ro_dim)
|
| 1427 |
+
sin = sin.unsqueeze(2).repeat(1, 1, 1, 2) # (b, l, 1, ro_dim)
|
| 1428 |
return torch.cat(
|
| 1429 |
[x[..., :ro_dim] * cos + _rotate_half(x[..., :ro_dim]) * sin, x[..., ro_dim:]],
|
| 1430 |
dim=-1,
|
| 1431 |
+
) # (b, l, h, d)
|
| 1432 |
|
| 1433 |
|
| 1434 |
@torch.compiler.disable
|
|
|
|
| 1442 |
uid_base_freq: float = 10.0,
|
| 1443 |
) -> tuple[Tensor, Tensor]:
|
| 1444 |
"""Build cos/sin for 3D RoPE + UID RoPE."""
|
| 1445 |
+
# ref_pos: (b, a, 3); ref_space_uid: (b, a); s/u = spatial/UID pair counts.
|
| 1446 |
device = ref_pos.device
|
| 1447 |
batch_size, n_atoms = ref_pos.shape[:2]
|
| 1448 |
half_dim = head_dim // 2
|
|
|
|
| 1454 |
torch.arange(0, n_spatial_per_axis, dtype=torch.float32, device=device)
|
| 1455 |
/ n_spatial_per_axis
|
| 1456 |
)
|
| 1457 |
+
) # (s,)
|
| 1458 |
uid_inv_freq = 1.0 / (
|
| 1459 |
uid_base_freq
|
| 1460 |
** (torch.arange(0, n_uid_pairs, dtype=torch.float32, device=device) / n_uid_pairs)
|
| 1461 |
+
) # (u,)
|
| 1462 |
|
| 1463 |
+
pos_f32 = ref_pos.float() # (b, a, 3)
|
| 1464 |
+
spatial_freqs = torch.einsum("bna,k->bnak", pos_f32, spatial_inv_freq) # (b, a, 3, s)
|
| 1465 |
+
spatial_freqs = spatial_freqs.reshape(batch_size, n_atoms, n_spatial_total) # (b, a, 3 * s)
|
| 1466 |
|
| 1467 |
+
uid_f32 = ref_space_uid.float() # (b, a)
|
| 1468 |
+
uid_freqs = torch.einsum("bn,k->bnk", uid_f32, uid_inv_freq) # (b, a, u)
|
| 1469 |
|
| 1470 |
n_active = n_spatial_total + n_uid_pairs
|
| 1471 |
+
freqs = torch.cat([spatial_freqs, uid_freqs], dim=-1) # (b, a, 3 * s + u)
|
| 1472 |
|
| 1473 |
if n_active < half_dim:
|
| 1474 |
padding = torch.zeros(
|
|
|
|
| 1477 |
half_dim - n_active,
|
| 1478 |
device=device,
|
| 1479 |
dtype=torch.float32,
|
| 1480 |
+
) # (b, a, half_dim - n_active)
|
| 1481 |
+
freqs = torch.cat([freqs, padding], dim=-1) # (b, a, half_dim)
|
| 1482 |
|
| 1483 |
+
cos = freqs.cos().to(torch.bfloat16) # freqs.shape
|
| 1484 |
+
sin = freqs.sin().to(torch.bfloat16) # freqs.shape
|
| 1485 |
+
return cos, sin # each (b, a, max(3 * s + u, half_dim))
|
| 1486 |
|
| 1487 |
|
| 1488 |
def qk_norm(x: Tensor) -> Tensor:
|
| 1489 |
+
# x: arbitrary leading dimensions and a final head-width axis.
|
| 1490 |
+
return F.rms_norm(x, (x.size(-1),)).to(x.dtype) # x.shape
|
| 1491 |
|
| 1492 |
|
| 1493 |
# ===========================================================================
|
|
|
|
| 1505 |
self.w_down = nn.Linear(hidden_size, d_model, bias=False)
|
| 1506 |
|
| 1507 |
def forward(self, x: Tensor) -> Tensor:
|
| 1508 |
+
# x: (..., d_model); d_hidden = w_down.in_features.
|
| 1509 |
+
x = x.to(self.w_up.weight.dtype) # (..., d_model)
|
| 1510 |
+
x1, x2 = self.w_up(x).chunk(2, dim=-1) # each (..., d_hidden)
|
| 1511 |
+
return self.w_down(F.silu(x1) * x2) # (..., d_model)
|
| 1512 |
|
| 1513 |
|
| 1514 |
# ===========================================================================
|
|
|
|
| 1551 |
# indices: (t,) flat positions of real atoms; cu_seqlens: (b + 1,) int32 row offsets.
|
| 1552 |
indices, cu_seqlens, max_seqlen = attention_params[2:5]
|
| 1553 |
flat_shape = (batch_size * n_atoms, self.n_heads, self.head_dim)
|
| 1554 |
+
q, k, v = q.reshape(flat_shape), k.reshape(flat_shape), v.reshape(flat_shape) # each (b * n_atoms, h, d_h)
|
| 1555 |
has_padding = indices.shape[0] != batch_size * n_atoms
|
| 1556 |
if has_padding:
|
| 1557 |
q, k, v = q[indices], k[indices], v[indices] # each (t, h, d_h)
|
|
|
|
| 1572 |
) # (t, h, d_h)
|
| 1573 |
if has_padding:
|
| 1574 |
out = attended.new_zeros(flat_shape) # (b * n_atoms, h, d_h)
|
| 1575 |
+
out[indices] = attended # (t, h, d_h)
|
| 1576 |
else:
|
| 1577 |
+
out = attended # (b * n_atoms, h, d_h)
|
| 1578 |
+
return out.view(batch_size, n_atoms, self.n_heads, self.head_dim) # (b, n_atoms, h, d_h)
|
| 1579 |
|
| 1580 |
def forward(self, x: Tensor, attention_params: tuple) -> Tensor:
|
| 1581 |
+
# x: (b, a, d_model); h = n_heads, d_h = head_dim; r is rotary-pair count.
|
| 1582 |
batch_size, n_atoms = x.shape[:2]
|
| 1583 |
+
cos, sin = attention_params[0], attention_params[1] # each (b, a, r)
|
| 1584 |
|
| 1585 |
+
x_input = x # (b, a, d_model)
|
| 1586 |
+
qkv = self.Wqkv(x) # (b, a, 3 * d_model)
|
| 1587 |
+
qkv = qkv.view(batch_size, n_atoms, 3, self.n_heads, self.head_dim).permute(2, 0, 1, 3, 4) # (3, b, a, h, d_h)
|
| 1588 |
+
q, k, v = qkv.unbind(0) # each (b, a, h, d_h)
|
| 1589 |
+
q, k = qk_norm(q), qk_norm(k) # each (b, a, h, d_h)
|
| 1590 |
|
| 1591 |
+
q = apply_rotary_emb_3d(q, cos, sin) # (b, a, h, d_h)
|
| 1592 |
+
k = apply_rotary_emb_3d(k, cos, sin) # (b, a, h, d_h)
|
| 1593 |
|
| 1594 |
input_dtype = q.dtype
|
| 1595 |
if q.dtype not in (torch.float16, torch.bfloat16):
|
| 1596 |
+
q, k, v = q.bfloat16(), k.bfloat16(), v.bfloat16() # each (b, a, h, d_h)
|
| 1597 |
|
| 1598 |
# ESMFold2 does not advertise FlashAttention. Keep this atom path on
|
| 1599 |
# PyTorch. Models that advertise FlashAttention dispatch through the
|
| 1600 |
# precompiled Hugging Face kernels interface in fastplms.attention.
|
| 1601 |
if self._atom_attention == ATOM_ATTENTION_WINDOWED:
|
| 1602 |
+
out = self._windowed_attention(q, k, v, attention_params) # (b, a, h, d_h)
|
| 1603 |
else:
|
| 1604 |
+
q_t = q.transpose(1, 2) # (b, h, a, d_h)
|
| 1605 |
+
k_t = k.transpose(1, 2) # (b, h, a, d_h)
|
| 1606 |
+
v_t = v.transpose(1, 2) # (b, h, a, d_h)
|
| 1607 |
+
attn = torch.matmul(q_t, k_t.transpose(-2, -1)) * self.scale # (b, h, a, a)
|
| 1608 |
+
attn = F.softmax(attn, dim=-1) # (b, h, a, a)
|
| 1609 |
+
out = torch.matmul(attn, v_t).transpose(1, 2) # (b, a, h, d_h)
|
| 1610 |
|
| 1611 |
out = out.to(input_dtype).reshape( # type: ignore[union-attr]
|
| 1612 |
batch_size, n_atoms, -1
|
| 1613 |
+
) # (b, a, d_model)
|
| 1614 |
+
out = out * torch.sigmoid(self.gate_proj(x_input)) # (b, a, d_model)
|
| 1615 |
+
return self.out_proj(out) # (b, a, d_model)
|
| 1616 |
|
| 1617 |
|
| 1618 |
# ===========================================================================
|
|
|
|
| 1621 |
|
| 1622 |
|
| 1623 |
def _rms_adaln_raw(x: Tensor, scale: Tensor, shift: Tensor) -> Tensor:
|
| 1624 |
+
# x, scale, shift: broadcast-compatible arrays; normalize x final axis.
|
| 1625 |
+
return F.rms_norm(x, (x.shape[-1],)) * (1 + scale) + shift # broadcast(x.shape, scale.shape, shift.shape)
|
| 1626 |
|
| 1627 |
|
| 1628 |
def _gated_residual_raw(x: Tensor, gate: Tensor, y: Tensor) -> Tensor:
|
| 1629 |
+
# x, gate, y: broadcast-compatible arrays.
|
| 1630 |
+
return x + gate * y # broadcast(x.shape, gate.shape, y.shape)
|
| 1631 |
|
| 1632 |
|
| 1633 |
class SWAAtomBlock(nn.Module):
|
|
|
|
| 1649 |
self.ffn_norm = nn.RMSNorm(d_atom, elementwise_affine=False)
|
| 1650 |
|
| 1651 |
adaln_linear = nn.Linear(d_atom, 6 * d_atom, bias=False)
|
| 1652 |
+
nn.init.zeros_(adaln_linear.weight) # (6 * d_atom, d_atom)
|
| 1653 |
self.adaln_modulation = nn.Sequential(nn.SiLU(), adaln_linear)
|
| 1654 |
|
| 1655 |
self.attn = SWA3DRoPEAttention(d_atom, n_heads, half_window=half_window)
|
|
|
|
| 1661 |
)
|
| 1662 |
|
| 1663 |
def forward(self, x: Tensor, c_l: Tensor, attention_params: tuple) -> Tensor:
|
| 1664 |
+
# x: (b, a, d_atom); c_l: (b, d_atom) or (b, a, d_atom).
|
| 1665 |
+
mod = self.adaln_modulation(c_l) # c_l.shape[:-1] + (6 * d_atom,)
|
| 1666 |
if mod.dim() == 2:
|
| 1667 |
+
mod = mod.unsqueeze(1) # (b, 1, 6 * d_atom)
|
| 1668 |
+
shift_a, scale_a, gate_a, shift_f, scale_f, gate_f = mod.chunk(6, dim=-1) # each (b, 1 or a, d_atom)
|
| 1669 |
|
| 1670 |
+
attn_input = self._rms_adaln(x, scale_a, shift_a) # (b, a, d_atom)
|
| 1671 |
+
attn_out = self.attn(attn_input, attention_params) # (b, a, d_atom)
|
| 1672 |
+
x = self._gated_residual(x, gate_a, attn_out) # (b, a, d_atom)
|
| 1673 |
|
| 1674 |
+
ffn_input = self._rms_adaln(x, scale_f, shift_f) # (b, a, d_atom)
|
| 1675 |
+
ffn_out = self.ffn(ffn_input) # (b, a, d_atom)
|
| 1676 |
+
x = self._gated_residual(x, gate_f, ffn_out) # (b, a, d_atom)
|
| 1677 |
+
return x # (b, a, d_atom)
|
| 1678 |
|
| 1679 |
|
| 1680 |
class SWAAtomTransformer(nn.Module):
|
|
|
|
| 1730 |
attention_params: tuple,
|
| 1731 |
return_intermediates: bool = False,
|
| 1732 |
) -> Tensor | tuple[Tensor, list[Tensor]]:
|
| 1733 |
+
# q_l/c_l: (b, a, d_atom); each saved intermediate has the same shape.
|
| 1734 |
intermediates: list[Tensor] = []
|
| 1735 |
for block in self.blocks:
|
| 1736 |
+
q_l = block(q_l, c_l, attention_params) # (b, a, d_atom)
|
| 1737 |
if return_intermediates:
|
| 1738 |
intermediates.append(q_l)
|
| 1739 |
if return_intermediates:
|
| 1740 |
+
return q_l, intermediates # q_l: (b, a, d_atom); optional list contains tensors of that shape
|
| 1741 |
+
return q_l # q_l: (b, a, d_atom); optional list contains tensors of that shape
|
| 1742 |
|
| 1743 |
|
| 1744 |
# ===========================================================================
|
|
|
|
| 1753 |
num_diffusion_samples: int,
|
| 1754 |
) -> tuple[Tensor, Tensor, Tensor, int, int]:
|
| 1755 |
"""Prepare mask-derived atom metadata outside compiled diffusion graphs."""
|
| 1756 |
+
# atom_attention_mask/atom_to_token: (b, a); bs = b * num_diffusion_samples.
|
| 1757 |
+
mask_exp = atom_attention_mask.repeat_interleave(num_diffusion_samples, 0) # (bs, a)
|
| 1758 |
+
seqlens = mask_exp.sum(dim=-1, dtype=torch.int32) # (bs,)
|
| 1759 |
+
indices = torch.nonzero(mask_exp.flatten(), as_tuple=False).flatten() # (n_present,)
|
| 1760 |
max_seqlen = int(seqlens.max().item())
|
| 1761 |
+
cu_seqlens = F.pad(torch.cumsum(seqlens, dim=0, dtype=torch.int32), (1, 0)) # (bs + 1,)
|
| 1762 |
n_tokens = int(atom_to_token.max().item()) + 1
|
| 1763 |
+
return mask_exp, indices, cu_seqlens, max_seqlen, n_tokens # (bs, a), (n_present,), (bs + 1,), scalar integers
|
| 1764 |
|
| 1765 |
|
| 1766 |
class ESMFold2AtomEncoder(nn.Module):
|
|
|
|
| 1838 |
``inference_cache`` caches step-invariant tensors (c_base, 3D RoPE,
|
| 1839 |
attention indices, n_tokens) across diffusion steps.
|
| 1840 |
"""
|
| 1841 |
+
# ref_pos: (b, a, 3); ref_element: (b, a, 128); chars: (b, a, 4, 64); other atom inputs: (b, a).
|
| 1842 |
+
# bs = b * samples; d_out = d_token for structure prediction, otherwise d_token / 2.
|
| 1843 |
batch_size, n_atoms = ref_pos.shape[:2]
|
| 1844 |
|
| 1845 |
layer_cache = None
|
|
|
|
| 1856 |
ref_atom_name_chars.reshape(batch_size, n_atoms, MAX_CHARS * CHAR_VOCAB_SIZE),
|
| 1857 |
],
|
| 1858 |
dim=-1,
|
| 1859 |
+
) # (b, a, d_atom_features)
|
| 1860 |
+
c_base = self.atom_norm(self.atom_linear(atom_feats)) # (b, a, d_atom)
|
| 1861 |
+
cos, sin = self.atom_transformer._build_3d_rope(ref_pos, ref_space_uid) # each (b, a, rotary_pairs)
|
| 1862 |
+
cos = cos.repeat_interleave(num_diffusion_samples, 0) # (bs, a, rotary_pairs)
|
| 1863 |
+
sin = sin.repeat_interleave(num_diffusion_samples, 0) # (bs, a, rotary_pairs)
|
| 1864 |
mask_exp, indices, cu_seqlens, max_seqlen, n_tokens = (
|
| 1865 |
_prepare_atom_encoder_metadata(
|
| 1866 |
atom_attention_mask,
|
| 1867 |
atom_to_token,
|
| 1868 |
num_diffusion_samples,
|
| 1869 |
)
|
| 1870 |
+
) # (bs, a), (n_present,), (bs + 1,), integers
|
| 1871 |
attention_params = (cos, sin, indices, cu_seqlens, max_seqlen)
|
| 1872 |
if layer_cache is not None:
|
| 1873 |
+
layer_cache["c_base"] = c_base # (b, a, d_atom)
|
| 1874 |
layer_cache["attention_params"] = attention_params
|
| 1875 |
+
layer_cache["mask_exp"] = mask_exp # (bs, a)
|
| 1876 |
layer_cache["n_tokens"] = n_tokens
|
| 1877 |
layer_cache["atom_to_token_exp"] = atom_to_token.repeat_interleave(
|
| 1878 |
num_diffusion_samples, 0
|
| 1879 |
+
) # (bs, a)
|
| 1880 |
else:
|
| 1881 |
+
c_base = layer_cache["c_base"] # (b, a, d_atom)
|
| 1882 |
attention_params = layer_cache["attention_params"]
|
| 1883 |
+
mask_exp = layer_cache["mask_exp"] # (bs, a)
|
| 1884 |
n_tokens = layer_cache["n_tokens"]
|
| 1885 |
|
| 1886 |
+
c = c_base # (b, a, d_atom)
|
| 1887 |
|
| 1888 |
+
q = c # (b, a, d_atom)
|
| 1889 |
|
| 1890 |
if self.structure_prediction and r_l is not None:
|
| 1891 |
+
q = q.repeat_interleave(num_diffusion_samples, 0) # (bs, a, d_atom)
|
| 1892 |
if pred_r1 is None:
|
| 1893 |
+
pred_r1 = torch.zeros_like(r_l) # r_l.shape = (bs, a, 3)
|
| 1894 |
+
r_input = torch.cat([r_l, pred_r1], dim=-1) # (bs, a, 6)
|
| 1895 |
+
r_to_q = self.coords_linear(r_input) # (bs, a, d_atom)
|
| 1896 |
+
q = q + r_to_q # (bs, a, d_atom)
|
| 1897 |
|
| 1898 |
+
c = c.repeat_interleave(num_diffusion_samples, 0) # (bs, a, d_atom)
|
| 1899 |
|
| 1900 |
result = self.atom_transformer(
|
| 1901 |
q_l=q,
|
|
|
|
| 1904 |
return_intermediates=return_intermediates,
|
| 1905 |
)
|
| 1906 |
if return_intermediates:
|
| 1907 |
+
q, intermediates = result # q: (bs, a, d_atom); list of same-shaped tensors
|
| 1908 |
else:
|
| 1909 |
+
q = result # (bs, a, d_atom)
|
| 1910 |
intermediates = []
|
| 1911 |
|
| 1912 |
+
q_to_a = F.relu(self.atom_to_token_linear(q)) # (bs, a, d_out)
|
| 1913 |
if layer_cache is not None and "atom_to_token_exp" in layer_cache:
|
| 1914 |
+
atom_to_token_exp = layer_cache["atom_to_token_exp"] # (bs, a)
|
| 1915 |
else:
|
| 1916 |
+
atom_to_token_exp = atom_to_token.repeat_interleave(num_diffusion_samples, 0) # (bs, a)
|
| 1917 |
+
a = scatter_atom_to_token(q_to_a, atom_to_token_exp, n_tokens, atom_mask=mask_exp.bool()) # (bs, n_tokens, d_out)
|
| 1918 |
|
| 1919 |
+
return a, q, c, attention_params, intermediates # a: (bs, l, d_out); q/c: (bs, a, d_atom); metadata tuple and intermediate list
|
| 1920 |
|
| 1921 |
|
| 1922 |
# ===========================================================================
|
|
|
|
| 1970 |
return_intermediates: bool = False,
|
| 1971 |
) -> tuple[Tensor, list[Tensor]]:
|
| 1972 |
"""Returns (r_update, intermediates)."""
|
| 1973 |
+
# a_i: (bs, l, d_token); q_l/c_l: (bs, a, d_atom); atom_to_token: (b, a); bs = b * samples.
|
| 1974 |
+
atom_to_token_exp = atom_to_token.repeat_interleave(num_diffusion_samples, 0) # (bs, a)
|
| 1975 |
+
a_to_q = self.token_to_atom_linear(a_i) # (bs, l, d_atom)
|
| 1976 |
+
a_to_q = gather_token_to_atom(a_to_q, atom_to_token_exp) # (bs, a, d_atom)
|
| 1977 |
+
q_l = q_l + a_to_q # (bs, a, d_atom)
|
| 1978 |
|
| 1979 |
result = self.atom_transformer(
|
| 1980 |
q_l=q_l,
|
|
|
|
| 1983 |
return_intermediates=return_intermediates,
|
| 1984 |
)
|
| 1985 |
if return_intermediates:
|
| 1986 |
+
q_l, intermediates = result # q_l: (bs, a, d_atom); list of same-shaped tensors
|
| 1987 |
else:
|
| 1988 |
+
q_l = result # (bs, a, d_atom)
|
| 1989 |
intermediates = []
|
| 1990 |
|
| 1991 |
+
r_l = self.output_linear(self.norm(q_l)) # (bs, a, 3)
|
| 1992 |
+
return r_l, intermediates # r_l: (bs, a, 3); intermediate tensors: (bs, a, d_atom)
|
| 1993 |
|
| 1994 |
|
| 1995 |
# ===========================================================================
|
|
|
|
| 2019 |
self.adaln = AdaptiveLayerNorm(d_model, d_cond, eps=1e-5)
|
| 2020 |
self.out_gate = nn.Linear(d_cond, d_model, bias=True)
|
| 2021 |
# adaln init: weight=0, bias=-2
|
| 2022 |
+
nn.init.zeros_(self.out_gate.weight) # (d_model, d_cond)
|
| 2023 |
+
nn.init.constant_(self.out_gate.bias, -2.0) # (d_model,)
|
| 2024 |
else:
|
| 2025 |
self.pre_norm = nn.LayerNorm(d_model, eps=1e-5)
|
| 2026 |
|
|
|
|
| 2079 |
conditions every denoising step on the same ``z``, so the PyTorch path
|
| 2080 |
projects the bias on the first step and reuses that tensor afterwards.
|
| 2081 |
"""
|
| 2082 |
+
# a: (bs, l, d_model); s: (bs, l, d_cond) or None; z: (b, l, l, d_pair) or (bs, l, l); bs = b * samples.
|
| 2083 |
bsz, n_queries, d_model = a.shape
|
| 2084 |
|
| 2085 |
+
x = self.adaln(a, s) if s is not None else self.pre_norm(a) # (bs, l, d_model)
|
| 2086 |
|
| 2087 |
n_keys = x.shape[1]
|
| 2088 |
+
q = self.q_proj(x).view(bsz, n_queries, self.num_heads, self.head_dim) # (bs, l, h, d_h)
|
| 2089 |
+
kv = self.kv_proj(x) # (bs, l, 2 * d_model)
|
| 2090 |
+
k, v = kv.chunk(2, dim=-1) # each (bs, l, d_model)
|
| 2091 |
+
k = k.view(bsz, n_keys, self.num_heads, self.head_dim) # (bs, l, h, d_h)
|
| 2092 |
+
v = v.view(bsz, n_keys, self.num_heads, self.head_dim) # (bs, l, h, d_h)
|
| 2093 |
|
| 2094 |
use_fused_kernel = self._can_use_fused_pair_bias(z, n_queries, beta)
|
| 2095 |
use_cueq_kernel = not use_fused_kernel and self._can_use_cueq_pair_bias(z, n_queries, beta)
|
| 2096 |
+
cached_pair_bias = None # (bs, l, l, h) or None
|
| 2097 |
if step_cache is not None and not use_fused_kernel and not use_cueq_kernel:
|
| 2098 |
+
cached_pair_bias = step_cache.get("pair_bias") # (bs, l, l, h) or None
|
| 2099 |
|
| 2100 |
# Expand z for num_diffusion_samples, unless its projection is already cached.
|
| 2101 |
if (
|
|
|
|
| 2110 |
and attention_mask.shape[0] != bsz
|
| 2111 |
and num_diffusion_samples > 1
|
| 2112 |
):
|
| 2113 |
+
attention_mask = attention_mask.repeat_interleave(num_diffusion_samples, dim=0) # (bs, l)
|
| 2114 |
|
| 2115 |
if use_fused_kernel:
|
| 2116 |
kernel_mask = (
|
| 2117 |
attention_mask
|
| 2118 |
if attention_mask is not None
|
| 2119 |
else torch.ones(bsz, n_queries, device=a.device, dtype=torch.bool)
|
| 2120 |
+
) # (bs, l)
|
| 2121 |
+
pair_norm_w = self.pair_norm.weight # (d_pair,)
|
| 2122 |
pair_norm_b = (
|
| 2123 |
self.pair_norm.bias
|
| 2124 |
if self.pair_norm.bias is not None
|
| 2125 |
else torch.zeros_like(pair_norm_w)
|
| 2126 |
+
) # (d_pair,)
|
| 2127 |
+
z_bf = z if z.dtype == torch.bfloat16 else z.to(torch.bfloat16) # (bs, l, l, d_pair)
|
| 2128 |
bias = _fused_pair_bias( # type: ignore[misc]
|
| 2129 |
z_bf,
|
| 2130 |
kernel_mask,
|
|
|
|
| 2132 |
num_heads=self.num_heads,
|
| 2133 |
pair_norm_w=pair_norm_w,
|
| 2134 |
pair_norm_b=pair_norm_b,
|
| 2135 |
+
) # bias: (bs, h, l, l)
|
| 2136 |
+
q_bhqd = q.transpose(1, 2) # (bs, h, l, d_h)
|
| 2137 |
+
k_bhqd = k.transpose(1, 2) # (bs, h, l, d_h)
|
| 2138 |
+
v_bhqd = v.transpose(1, 2) # (bs, h, l, d_h)
|
| 2139 |
attn_out = F.scaled_dot_product_attention(
|
| 2140 |
q_bhqd, k_bhqd, v_bhqd, attn_mask=bias.to(q_bhqd.dtype)
|
| 2141 |
+
) # (bs, h, l, d_h)
|
| 2142 |
+
g = torch.sigmoid(self.g_proj(x)).view(bsz, n_queries, self.num_heads, self.head_dim) # (bs, l, h, d_h)
|
| 2143 |
+
ctx = g * attn_out.transpose(1, 2) # (bs, l, h, d_h)
|
| 2144 |
+
out = self.out_proj(ctx.reshape(bsz, n_queries, d_model)) # (bs, l, d_model)
|
| 2145 |
if s is not None:
|
| 2146 |
+
out = torch.sigmoid(self.out_gate(s)) * out # (bs, l, d_model)
|
| 2147 |
+
return out # (bs, l, d_model)
|
| 2148 |
|
| 2149 |
if use_cueq_kernel:
|
| 2150 |
kernel_mask = (
|
| 2151 |
attention_mask
|
| 2152 |
if attention_mask is not None
|
| 2153 |
else torch.ones(bsz, n_queries, device=a.device, dtype=torch.bool)
|
| 2154 |
+
) # (bs, l)
|
| 2155 |
out, _ = _cue_attn_pair_bias( # type: ignore[misc]
|
| 2156 |
s=x,
|
| 2157 |
q=q.transpose(1, 2),
|
|
|
|
| 2167 |
b_ln_z=self.pair_norm.bias,
|
| 2168 |
return_z_proj=False,
|
| 2169 |
is_cached_z_proj=False,
|
| 2170 |
+
) # out: (bs, l, d_model); unused kernel auxiliary
|
| 2171 |
else:
|
| 2172 |
# Standard attention with pair bias
|
| 2173 |
+
g = torch.sigmoid(self.g_proj(x)).view(bsz, n_queries, self.num_heads, self.head_dim) # (bs, l, h, d_h)
|
| 2174 |
|
| 2175 |
+
logits = torch.einsum("... i h d, ... j h d -> ... i j h", q, k) * self.scale # (bs, l, l, h)
|
| 2176 |
|
| 2177 |
if cached_pair_bias is not None:
|
| 2178 |
pair_bias = cached_pair_bias # (b * samples, n, n, h)
|
| 2179 |
elif z.dim() == 4:
|
| 2180 |
pair_bias = self.pair_bias_proj(self.pair_norm(z)) # (b * samples, n, n, h)
|
| 2181 |
if step_cache is not None:
|
| 2182 |
+
step_cache["pair_bias"] = pair_bias # (bs, l, l, h)
|
| 2183 |
else:
|
| 2184 |
pair_bias = z.unsqueeze(-1) # (b * samples, n, n, 1), a precomputed bias
|
| 2185 |
+
logits = logits + pair_bias.to(dtype=logits.dtype) # (bs, l, l, h)
|
| 2186 |
|
| 2187 |
if attention_mask is not None:
|
| 2188 |
min_val = torch.finfo(logits.dtype).min
|
| 2189 |
+
mask_bias = torch.where(attention_mask.bool()[:, None, :, None], 0.0, min_val) # (bs, 1, l, 1)
|
| 2190 |
+
logits = logits + mask_bias.to(dtype=logits.dtype) # (bs, l, l, h)
|
| 2191 |
|
| 2192 |
+
attn = torch.softmax(logits, dim=-2).to(dtype=v.dtype) # (bs, l, l, h)
|
| 2193 |
+
ctx = torch.einsum("... i j h, ... j h d -> ... i h d", attn, v) # (bs, l, h, d_h)
|
| 2194 |
+
ctx = g * ctx # (bs, l, h, d_h)
|
| 2195 |
+
out = self.out_proj(ctx.reshape(bsz, n_queries, d_model)) # (bs, l, d_model)
|
| 2196 |
|
| 2197 |
if s is not None:
|
| 2198 |
+
out = torch.sigmoid(self.out_gate(s)) * out # (bs, l, d_model)
|
| 2199 |
+
return out # (bs, l, d_model)
|
| 2200 |
|
| 2201 |
|
| 2202 |
# ===========================================================================
|
|
|
|
| 2221 |
if use_conditioning:
|
| 2222 |
self.adaln = AdaptiveLayerNorm(d_model, d_cond, eps=1e-5)
|
| 2223 |
self.output_gate = nn.Linear(d_cond, d_model, bias=True)
|
| 2224 |
+
nn.init.zeros_(self.output_gate.weight) # (d_model, d_cond)
|
| 2225 |
+
nn.init.constant_(self.output_gate.bias, -2.0) # (d_model,)
|
| 2226 |
else:
|
| 2227 |
self.pre_norm = nn.LayerNorm(d_model, eps=1e-5)
|
| 2228 |
|
|
|
|
| 2230 |
self.lin_out = nn.Linear(hidden, d_model, bias=False)
|
| 2231 |
|
| 2232 |
def forward(self, a: Tensor, s: Tensor | None) -> Tensor:
|
| 2233 |
+
# a: (..., d_model); s: (..., d_cond) or None; d_hidden = lin_out.in_features.
|
| 2234 |
+
x = self.adaln(a, s) if s is not None else self.pre_norm(a) # (..., d_model)
|
| 2235 |
|
| 2236 |
+
swish_a, swish_b = self.lin_swish(x).chunk(2, dim=-1) # each (..., d_hidden)
|
| 2237 |
+
b = F.silu(swish_a) * swish_b # (..., d_hidden)
|
| 2238 |
+
out = self.lin_out(b) # (..., d_model)
|
| 2239 |
|
| 2240 |
if s is not None:
|
| 2241 |
+
out = torch.sigmoid(self.output_gate(s)) * out # (..., d_model)
|
| 2242 |
+
return out # (..., d_model)
|
| 2243 |
|
| 2244 |
|
| 2245 |
# ===========================================================================
|
|
|
|
| 2307 |
``inference_cache`` must span only calls that share ``z``, as one
|
| 2308 |
``sample`` call does; each block then keeps its pair bias across steps.
|
| 2309 |
"""
|
| 2310 |
+
# a: (bs, l, d_model); s: (bs, l, d_cond) or None; z follows AttentionPairBias contract.
|
| 2311 |
intermediates: list[Tensor] = []
|
| 2312 |
block_caches: dict[int, dict[str, Tensor]] | None = None
|
| 2313 |
if inference_cache is not None:
|
| 2314 |
block_caches = inference_cache.setdefault("token_pair_bias", {})
|
| 2315 |
+
x = a # (bs, l, d_model)
|
| 2316 |
for block_index, (attn, transition) in enumerate(
|
| 2317 |
zip(self.attn_blocks, self.transition_blocks, strict=True)
|
| 2318 |
):
|
|
|
|
| 2325 |
attention_mask=attention_mask,
|
| 2326 |
num_diffusion_samples=num_diffusion_samples,
|
| 2327 |
step_cache=step_cache,
|
| 2328 |
+
) # (bs, l, d_model)
|
| 2329 |
+
x = x + transition(x, s) # (bs, l, d_model)
|
| 2330 |
if return_intermediates:
|
| 2331 |
intermediates.append(x)
|
| 2332 |
+
return x, intermediates # x: (bs, l, d_model); list of same-shaped intermediate tensors
|
| 2333 |
|
| 2334 |
|
| 2335 |
# ===========================================================================
|
|
|
|
| 2382 |
num_diffusion_samples: int = 1,
|
| 2383 |
inference_cache: dict[str, Tensor] | None = None,
|
| 2384 |
) -> tuple[Tensor, Tensor]:
|
| 2385 |
+
# z_trunk/relative_position_encoding: (b, l, l, c_z); s_inputs: (b or bs, l, c_s_inputs); bs = b * samples.
|
| 2386 |
sigma = self.sigma_data if sigma_data is None else float(sigma_data)
|
| 2387 |
base_batch = z_trunk.shape[0]
|
| 2388 |
target_batch = base_batch * num_diffusion_samples
|
| 2389 |
|
| 2390 |
# z conditioning (cached across diffusion steps: independent of t_hat)
|
| 2391 |
if inference_cache is not None and "z" in inference_cache:
|
| 2392 |
+
z = inference_cache["z"] # (b, l, l, c_z)
|
| 2393 |
else:
|
| 2394 |
+
z_rel = relative_position_encoding.to(dtype=torch.float32) # (b, l, l, c_z)
|
| 2395 |
+
z = torch.cat([z_trunk.to(dtype=torch.float32), z_rel], dim=-1) # (b, l, l, 2 * c_z)
|
| 2396 |
+
z = self.z_proj(self.z_input_norm(z)) # (b, l, l, c_z)
|
| 2397 |
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
|
| 2398 |
for block in self.z_transitions:
|
| 2399 |
+
z = z + block(z) # (b, l, l, c_z)
|
| 2400 |
if inference_cache is not None:
|
| 2401 |
+
inference_cache["z"] = z # (b, l, l, c_z)
|
| 2402 |
|
| 2403 |
# s conditioning
|
| 2404 |
+
s_inputs_eff = s_inputs # (b or bs, l, c_s_inputs)
|
| 2405 |
if s_inputs_eff.shape[0] != target_batch:
|
| 2406 |
+
s_inputs_eff = s_inputs_eff.repeat_interleave(num_diffusion_samples, 0) # (bs, l, c_s_inputs)
|
| 2407 |
|
| 2408 |
+
s = self.s_proj(self.s_input_norm(s_inputs_eff.to(dtype=torch.float32))) # (bs, l, c_s)
|
| 2409 |
|
| 2410 |
# Noise embedding
|
| 2411 |
+
t = torch.as_tensor(t_hat, dtype=torch.float32, device=s.device).reshape(-1) # (t_hat.numel(),)
|
| 2412 |
if t.numel() == 1:
|
| 2413 |
+
t = t.expand(target_batch) # (bs,)
|
| 2414 |
elif t.shape[0] != target_batch:
|
| 2415 |
+
t = t.repeat_interleave(num_diffusion_samples, 0) # (bs,)
|
| 2416 |
+
t_noise = 0.25 * torch.log((t / sigma).clamp(min=1e-20)) # (bs,)
|
| 2417 |
+
n = self.fourier(t_noise) # (bs, fourier_dim)
|
| 2418 |
+
n = self.noise_proj(self.noise_norm(n)) # (bs, c_s)
|
| 2419 |
+
s = s + n.unsqueeze(1) # (bs, l, c_s)
|
| 2420 |
|
| 2421 |
for block in self.s_transitions:
|
| 2422 |
+
s = s + block(s) # (bs, l, c_s)
|
| 2423 |
|
| 2424 |
+
return s, z # s: (bs, l, c_s); z: (b, l, l, c_z)
|
| 2425 |
|
| 2426 |
|
| 2427 |
# ===========================================================================
|
|
|
|
| 2493 |
)
|
| 2494 |
|
| 2495 |
self.s_to_token = nn.Linear(c_token, c_token, bias=False)
|
| 2496 |
+
nn.init.zeros_(self.s_to_token.weight) # (c_token, c_token)
|
| 2497 |
|
| 2498 |
# Token transformer (DiffusionTransformer with pair bias)
|
| 2499 |
self.token_transformer = DiffusionTransformer(
|
|
|
|
| 2539 |
return_atom_repr: bool = False,
|
| 2540 |
inference_cache: dict[str, Tensor] | None = None,
|
| 2541 |
) -> dict[str, Tensor | None]:
|
| 2542 |
+
# x_noisy: (bs, a, 3); ref/ID tensors retain base batch b; bs = b * samples; l tokens.
|
| 2543 |
bsz = x_noisy.shape[0]
|
| 2544 |
sigma = self.sigma_data if sigma_data is None else float(sigma_data)
|
| 2545 |
+
t = torch.as_tensor(t_hat, dtype=torch.float32, device=x_noisy.device).reshape(-1) # (t_hat.numel(),)
|
| 2546 |
if t.numel() == 1:
|
| 2547 |
+
t = t.expand(bsz) # (bs,)
|
| 2548 |
|
| 2549 |
# Step 1: conditioning (pair z is cached across diffusion steps)
|
| 2550 |
s, z = self.conditioning(
|
|
|
|
| 2556 |
sigma_data=sigma,
|
| 2557 |
num_diffusion_samples=num_diffusion_samples,
|
| 2558 |
inference_cache=inference_cache,
|
| 2559 |
+
) # (bs, l, c_token), (b, l, l, c_z)
|
| 2560 |
|
| 2561 |
# Step 2: normalize noisy coords
|
| 2562 |
+
denom = torch.sqrt(t * t + sigma * sigma) # (bs,)
|
| 2563 |
+
r_noisy = x_noisy / denom[:, None, None] # (bs, a, 3)
|
| 2564 |
|
| 2565 |
# Step 3: atom encoder
|
| 2566 |
a, q_skip, c_skip, p_skip, enc_intermediates = self.atom_encoder(
|
|
|
|
| 2576 |
num_diffusion_samples=num_diffusion_samples,
|
| 2577 |
return_intermediates=return_atom_repr,
|
| 2578 |
inference_cache=inference_cache,
|
| 2579 |
+
) # a: (bs, l, c_token); q/c: (bs, a, c_atom); metadata and atom intermediate list
|
| 2580 |
|
| 2581 |
# Step 4: add conditioned s
|
| 2582 |
+
a = a + self.s_to_token(self.s_step_norm(s)) # (bs, l, c_token)
|
| 2583 |
|
| 2584 |
# Step 5: token transformer
|
| 2585 |
a, _ = self.token_transformer(
|
|
|
|
| 2590 |
attention_mask=token_attention_mask,
|
| 2591 |
num_diffusion_samples=num_diffusion_samples,
|
| 2592 |
inference_cache=inference_cache,
|
| 2593 |
+
) # a: (bs, l, c_token); unused intermediate list
|
| 2594 |
|
| 2595 |
# Step 6: token norm
|
| 2596 |
+
a = self.token_norm(a) # (bs, l, c_token)
|
| 2597 |
|
| 2598 |
# Step 7: atom decoder
|
| 2599 |
r_update, dec_intermediates = self.atom_decoder(
|
|
|
|
| 2605 |
atom_attention_mask=ref_mask,
|
| 2606 |
num_diffusion_samples=num_diffusion_samples,
|
| 2607 |
return_intermediates=return_atom_repr,
|
| 2608 |
+
) # r_update: (bs, a, 3); atom intermediate list
|
| 2609 |
|
| 2610 |
# Step 8: compute denoised output
|
| 2611 |
sigma2 = sigma * sigma
|
| 2612 |
+
t2 = t * t # (bs,)
|
| 2613 |
+
out = (sigma2 / (sigma2 + t2))[:, None, None] * x_noisy # (bs, a, 3)
|
| 2614 |
+
out = out + ((sigma * t) / torch.sqrt(sigma2 + t2))[:, None, None] * r_update # (bs, a, 3)
|
| 2615 |
|
| 2616 |
# Collect atom intermediates from encoder + decoder
|
| 2617 |
+
atom_intermediates: Tensor | None = None # (bs, a, n_atom_blocks, c_atom) or None
|
| 2618 |
if return_atom_repr:
|
| 2619 |
all_ints = enc_intermediates + dec_intermediates
|
| 2620 |
if all_ints:
|
| 2621 |
+
atom_intermediates = torch.stack(all_ints, dim=2) # (bs, a, n_atom_blocks, c_atom) or None
|
| 2622 |
|
| 2623 |
return {
|
| 2624 |
"x_denoised": out,
|
| 2625 |
"token_repr": a if return_token_repr else None,
|
| 2626 |
"atom_intermediates": atom_intermediates,
|
| 2627 |
+
} # mapping: x_denoised (bs, a, 3); token_repr (bs, l, c_token) or None; atom intermediates as above
|
| 2628 |
|
| 2629 |
|
| 2630 |
# ===========================================================================
|
|
|
|
| 2688 |
[self.inference_s_max * self.sigma_data, 0.0],
|
| 2689 |
device=device,
|
| 2690 |
dtype=torch.float32,
|
| 2691 |
+
) # (steps + 1,)
|
| 2692 |
p = float(self.inference_p)
|
| 2693 |
inv_p = 1.0 / p
|
| 2694 |
+
k = torch.arange(steps, device=device, dtype=torch.float32) # (steps,)
|
| 2695 |
base = self.inference_s_max**inv_p + (k / (steps - 1)) * (
|
| 2696 |
self.inference_s_min**inv_p - self.inference_s_max**inv_p
|
| 2697 |
+
) # (steps,)
|
| 2698 |
+
schedule = self.sigma_data * base.pow(p) # (steps,)
|
| 2699 |
+
return F.pad(schedule, (0, 1), value=0.0) # (steps + 1,)
|
| 2700 |
|
| 2701 |
@staticmethod
|
| 2702 |
def _random_rotations(n: int, dtype: torch.dtype, device: torch.device) -> Tensor:
|
| 2703 |
+
q = torch.randn((n, 4), dtype=dtype, device=device) # (n, 4)
|
| 2704 |
+
scale = torch.sqrt((q * q).sum(dim=1)) # (n,)
|
| 2705 |
+
signs = torch.where(q[:, 0] < 0, -scale, scale) # (n,)
|
| 2706 |
+
q = q / signs[:, None] # (n, 4)
|
| 2707 |
+
r, i, j, k = torch.unbind(q, dim=-1) # each (n,)
|
| 2708 |
+
two_s = 2.0 / (q * q).sum(dim=-1) # (n,)
|
| 2709 |
return torch.stack(
|
| 2710 |
(
|
| 2711 |
1 - two_s * (j * j + k * k),
|
|
|
|
| 2719 |
1 - two_s * (i * i + j * j),
|
| 2720 |
),
|
| 2721 |
dim=-1,
|
| 2722 |
+
).reshape(n, 3, 3) # (n, 3, 3)
|
| 2723 |
|
| 2724 |
def _center_random_augmentation(
|
| 2725 |
self, x: Tensor, atom_mask: Tensor, second_coords: Tensor | None = None
|
| 2726 |
) -> tuple[Tensor, Tensor | None]:
|
| 2727 |
"""Algorithm 19: center + random rotation + translation."""
|
| 2728 |
+
# x/second_coords: (b, a, 3); atom_mask: (b, a).
|
| 2729 |
bsz = x.shape[0]
|
| 2730 |
mask = atom_mask.unsqueeze(-1) # M has shape (b, a, 1).
|
| 2731 |
+
denom = mask.sum(dim=1, keepdim=True).clamp(min=1) # (b, 1, 1)
|
| 2732 |
+
mean = (x * mask).sum(dim=1, keepdim=True) / denom # (b, 1, 3)
|
| 2733 |
+
x = x - mean # (b, a, 3)
|
| 2734 |
if second_coords is not None:
|
| 2735 |
+
second_coords = second_coords - mean # (b, a, 3)
|
| 2736 |
|
| 2737 |
+
r = self._random_rotations(bsz, x.dtype, x.device) # (b, 3, 3)
|
| 2738 |
+
x = torch.einsum("bmd,bds->bms", x, r) # (b, a, 3)
|
| 2739 |
if second_coords is not None:
|
| 2740 |
+
second_coords = torch.einsum("bmd,bds->bms", second_coords, r) # (b, a, 3)
|
| 2741 |
|
| 2742 |
+
t = torch.randn_like(x[:, 0:1, :]) # (b, 1, 3)
|
| 2743 |
+
x = x + t # (b, a, 3)
|
| 2744 |
if second_coords is not None:
|
| 2745 |
+
second_coords = second_coords + t # (b, a, 3)
|
| 2746 |
+
return x, second_coords # each (b, a, 3), or second_coords None
|
| 2747 |
|
| 2748 |
@staticmethod
|
| 2749 |
def _weighted_rigid_align(x: Tensor, x_gt: Tensor, w: Tensor, mask: Tensor) -> Tensor:
|
| 2750 |
"""Kabsch alignment: align x to x_gt with weights w."""
|
| 2751 |
+
# x/x_gt: (b, n, 3); w/mask: (b, n).
|
| 2752 |
w = (mask * w).unsqueeze(-1) # W has shape (b, n, 1).
|
| 2753 |
+
denom = w.sum(dim=-2, keepdim=True).clamp(min=1e-8) # (b, 1, 1)
|
| 2754 |
+
mu = (x * w).sum(dim=-2, keepdim=True) / denom # (b, 1, 3)
|
| 2755 |
+
mu_gt = (x_gt * w).sum(dim=-2, keepdim=True) / denom # (b, 1, 3)
|
| 2756 |
+
x_c = x - mu # (b, n, 3)
|
| 2757 |
+
xgt_c = x_gt - mu_gt # (b, n, 3)
|
| 2758 |
+
covariance = torch.einsum("bni,bnj->bij", w * xgt_c, x_c) # (b, 3, 3)
|
| 2759 |
+
covariance_f32 = covariance.float() # (b, 3, 3)
|
| 2760 |
u, _, vh = torch.linalg.svd(
|
| 2761 |
covariance_f32, driver="gesvd" if covariance_f32.is_cuda else None
|
| 2762 |
+
) # (b, 3, 3), (b, 3), (b, 3, 3)
|
| 2763 |
+
det = torch.linalg.det(u @ vh) # (b,)
|
| 2764 |
+
ones = torch.ones_like(det) # (b,)
|
| 2765 |
rotation = (u @ torch.diag_embed(torch.stack([ones, ones, det], dim=-1)) @ vh).to(
|
| 2766 |
covariance.dtype
|
| 2767 |
+
) # (b, 3, 3)
|
| 2768 |
+
return x_c @ rotation.transpose(-1, -2) + mu_gt # (b, n, 3)
|
| 2769 |
|
| 2770 |
# ------------------------------------------------------------------
|
| 2771 |
# Sampling
|
|
|
|
| 2809 |
so we inflate the underlying schedule length here to land back at the
|
| 2810 |
requested step count post-truncation.
|
| 2811 |
"""
|
| 2812 |
+
# z_trunk: (b, l, l, c_z); s_inputs: (b, l, c_s_inputs); atom features: b by a; bs = b * samples.
|
| 2813 |
n_atoms = tok_idx.shape[1]
|
| 2814 |
device = s_inputs.device
|
| 2815 |
target_batch = s_inputs.shape[0] * num_diffusion_samples
|
|
|
|
| 2818 |
|
| 2819 |
steps = self.inference_num_steps if num_sampling_steps is None else int(num_sampling_steps)
|
| 2820 |
|
| 2821 |
+
schedule = self.inference_noise_schedule(steps, device) # (steps + 1,)
|
| 2822 |
if max_inference_sigma is not None:
|
| 2823 |
+
schedule = schedule[schedule <= float(max_inference_sigma)] # (n_below_cap,)
|
| 2824 |
+
schedule = F.pad(schedule, (1, 0), value=float(max_inference_sigma)) # (n_below_cap + 1,)
|
| 2825 |
|
| 2826 |
lam = self.noise_scale if noise_scale is None else float(noise_scale)
|
| 2827 |
eta = self.step_scale if step_scale is None else float(step_scale)
|
| 2828 |
|
| 2829 |
+
x = schedule[0] * torch.randn(target_batch, n_atoms, 3, device=device, dtype=torch.float32) # (bs, a, 3)
|
| 2830 |
+
atom_mask = ref_mask.repeat_interleave(num_diffusion_samples, 0).float() # (bs, a)
|
| 2831 |
|
| 2832 |
gammas = torch.where(
|
| 2833 |
schedule > self.gamma_min,
|
| 2834 |
torch.full_like(schedule, self.gamma_0),
|
| 2835 |
torch.zeros_like(schedule),
|
| 2836 |
+
) # schedule.shape
|
| 2837 |
|
| 2838 |
+
x_denoised_prev: Tensor | None = None # (bs, a, 3) or None
|
| 2839 |
+
token_repr: Tensor | None = None # (bs, l, c_token) or None
|
| 2840 |
+
diff_atom_intermediates: Tensor | None = None # (bs, a, n_blocks, c_atom) or None
|
| 2841 |
|
| 2842 |
step_pairs = list(zip(schedule[:-1], schedule[1:], gammas[1:], strict=True))
|
| 2843 |
num_steps = len(step_pairs)
|
|
|
|
| 2854 |
for step_idx, (sigma_tm, sigma_t, gamma) in enumerate(step_iterator):
|
| 2855 |
x, x_denoised_prev = self._center_random_augmentation(
|
| 2856 |
x, atom_mask, second_coords=x_denoised_prev
|
| 2857 |
+
) # each (bs, a, 3), second may be None
|
| 2858 |
|
| 2859 |
sigma_tm_val = float(sigma_tm.item())
|
| 2860 |
t_hat_val = sigma_tm_val * (1.0 + float(gamma.item()))
|
| 2861 |
eps_std = lam * max(t_hat_val**2 - sigma_tm_val**2, 0.0) ** 0.5
|
| 2862 |
+
x_noisy = x + eps_std * torch.randn_like(x) # (bs, a, 3)
|
| 2863 |
|
| 2864 |
is_last_step = step_idx == num_steps - 1
|
| 2865 |
request_atom_repr = return_atom_repr and (
|
|
|
|
| 2890 |
return_token_repr=True,
|
| 2891 |
return_atom_repr=request_atom_repr,
|
| 2892 |
inference_cache=inference_cache,
|
| 2893 |
+
) # tensor mapping follows the called head's shape contract
|
| 2894 |
|
| 2895 |
+
x_denoised = dm_out["x_denoised"] # (bs, a, 3)
|
| 2896 |
+
token_repr = dm_out["token_repr"] # (bs, l, c_token) or None
|
| 2897 |
if request_atom_repr:
|
| 2898 |
+
diff_atom_intermediates = dm_out.get("atom_intermediates") # (bs, a, n_blocks, c_atom) or None
|
| 2899 |
|
| 2900 |
# Reverse diffusion alignment (Kabsch)
|
| 2901 |
with torch.autocast(device_type="cuda", enabled=False):
|
| 2902 |
x_noisy = self._weighted_rigid_align(
|
| 2903 |
x_noisy.float(), x_denoised.float(), atom_mask, atom_mask
|
| 2904 |
+
) # (bs, a, 3)
|
| 2905 |
+
x_noisy = x_noisy.to(dtype=x_denoised.dtype) # (bs, a, 3)
|
| 2906 |
|
| 2907 |
# ODE/SDE step
|
| 2908 |
sigma_t_val = float(sigma_t.item())
|
| 2909 |
+
denoised_over_sigma = (x_noisy - x_denoised) / t_hat_val # (bs, a, 3)
|
| 2910 |
+
x = x_noisy + eta * (sigma_t_val - t_hat_val) * denoised_over_sigma # (bs, a, 3)
|
| 2911 |
|
| 2912 |
# Denoising early-exit: stop when consecutive predictions converge
|
| 2913 |
if (
|
|
|
|
| 2921 |
x_denoised.float(),
|
| 2922 |
atom_mask,
|
| 2923 |
atom_mask,
|
| 2924 |
+
) # (bs, a, 3)
|
| 2925 |
+
diff = (x_denoised.float() - aligned) * atom_mask.unsqueeze(-1) # (bs, a, 3)
|
| 2926 |
per_sample_rmsd = (
|
| 2927 |
diff.pow(2).sum(dim=(-1, -2)) / atom_mask.sum(dim=-1).clamp(min=1)
|
| 2928 |
+
).sqrt() # (bs,)
|
| 2929 |
if per_sample_rmsd.max().item() < denoising_early_exit_rmsd:
|
| 2930 |
+
x = x_denoised # (bs, a, 3)
|
| 2931 |
+
x_denoised_prev = x_denoised # (bs, a, 3) or None
|
| 2932 |
break
|
| 2933 |
|
| 2934 |
+
x_denoised_prev = x_denoised # (bs, a, 3) or None
|
| 2935 |
|
| 2936 |
result: dict[str, Tensor | None] = {
|
| 2937 |
"sample_atom_coords": x,
|
|
|
|
| 2939 |
}
|
| 2940 |
if return_atom_repr:
|
| 2941 |
result["diff_atom_intermediates"] = diff_atom_intermediates
|
| 2942 |
+
return result # coordinate/token/optional atom-intermediate mapping with the shapes above
|
fastplms/models/esmfold2/modeling_esmfold2_experimental.py
CHANGED
|
@@ -9,13 +9,13 @@ re-injection and a different confidence/MSA stack.
|
|
| 9 |
from __future__ import annotations
|
| 10 |
|
| 11 |
import gc
|
| 12 |
-
from collections.abc import Mapping
|
| 13 |
-
from pathlib import Path
|
| 14 |
-
from typing import Any, ClassVar, cast
|
| 15 |
-
|
| 16 |
import torch
|
| 17 |
import torch.nn as nn
|
| 18 |
import torch.nn.functional as F
|
|
|
|
|
|
|
|
|
|
|
|
|
| 19 |
from torch import Tensor
|
| 20 |
from tqdm.auto import tqdm
|
| 21 |
from transformers.modeling_utils import PreTrainedModel
|
|
@@ -66,6 +66,7 @@ from .modeling_esmfold2_common import (
|
|
| 66 |
validate_prepared_auxiliary_inputs,
|
| 67 |
)
|
| 68 |
|
|
|
|
| 69 |
_EPS = 1e-5
|
| 70 |
_NONPOLYMER_ID = 3
|
| 71 |
|
|
@@ -82,8 +83,8 @@ class ConfidenceHead(nn.Module):
|
|
| 82 |
d_pair = config.d_pair
|
| 83 |
d_inputs = config.inputs.d_inputs
|
| 84 |
|
| 85 |
-
boundaries = torch.linspace(ch.min_dist, ch.max_dist, ch.distogram_bins - 1)
|
| 86 |
-
self.register_buffer("boundaries", boundaries)
|
| 87 |
self.dist_bin_pairwise_embed = nn.Embedding(ch.distogram_bins, d_pair)
|
| 88 |
|
| 89 |
self.s_norm = nn.LayerNorm(d_single)
|
|
@@ -105,7 +106,7 @@ class ConfidenceHead(nn.Module):
|
|
| 105 |
max_atoms_per_token = 23
|
| 106 |
self.plddt_weight = nn.Parameter(
|
| 107 |
torch.zeros(max_atoms_per_token, d_single, ch.num_plddt_bins)
|
| 108 |
-
)
|
| 109 |
self.pae_head = nn.Linear(d_pair, ch.num_pae_bins, bias=False)
|
| 110 |
|
| 111 |
def set_kernel_backend(self, backend: str | None) -> None:
|
|
@@ -117,16 +118,18 @@ class ConfidenceHead(nn.Module):
|
|
| 117 |
|
| 118 |
@staticmethod
|
| 119 |
def _repeat_batch(x: Tensor, num_diffusion_samples: int) -> Tensor:
|
|
|
|
| 120 |
if num_diffusion_samples == 1:
|
| 121 |
-
return x
|
| 122 |
-
return x.repeat_interleave(num_diffusion_samples, 0)
|
| 123 |
|
| 124 |
@staticmethod
|
| 125 |
def _flatten_sample_axis(x: Tensor) -> Tensor:
|
|
|
|
| 126 |
if x.ndim == 4:
|
| 127 |
b, mult, n, c = x.shape
|
| 128 |
-
return x.reshape(b * mult, n, c)
|
| 129 |
-
return x
|
| 130 |
|
| 131 |
def forward(
|
| 132 |
self,
|
|
@@ -143,109 +146,110 @@ class ConfidenceHead(nn.Module):
|
|
| 143 |
relative_position_encoding: Tensor | None = None,
|
| 144 |
token_bonds_encoding: Tensor | None = None,
|
| 145 |
) -> dict[str, Tensor]:
|
| 146 |
-
|
| 147 |
-
|
|
|
|
| 148 |
if relative_position_encoding is not None:
|
| 149 |
-
z_base = z_base + relative_position_encoding
|
| 150 |
if token_bonds_encoding is not None:
|
| 151 |
-
z_base = z_base + token_bonds_encoding
|
| 152 |
-
z_base = z_base + self.s_to_z(s_inputs_normed).unsqueeze(2)
|
| 153 |
-
z_base = z_base + self.s_to_z_transpose(s_inputs_normed).unsqueeze(1)
|
| 154 |
z_base = z_base + self.s_to_z_prod_out(
|
| 155 |
self.s_to_z_prod_in1(s_inputs_normed)[:, :, None, :]
|
| 156 |
* self.s_to_z_prod_in2(s_inputs_normed)[:, None, :, :]
|
| 157 |
-
)
|
| 158 |
-
|
| 159 |
-
pair = self._repeat_batch(z_base, num_diffusion_samples)
|
| 160 |
-
x_pred_flat = self._flatten_sample_axis(x_pred)
|
| 161 |
-
atom_to_token_m = self._repeat_batch(atom_to_token, num_diffusion_samples)
|
| 162 |
-
atom_mask_m = self._repeat_batch(atom_attention_mask, num_diffusion_samples)
|
| 163 |
-
rep_idx_m = self._repeat_batch(distogram_atom_idx, num_diffusion_samples).long()
|
| 164 |
-
mask = self._repeat_batch(token_attention_mask, num_diffusion_samples)
|
| 165 |
batch_mult = pair.shape[0]
|
| 166 |
|
| 167 |
-
rep_coords = gather_rep_atom_coords(x_pred_flat, rep_idx_m)
|
| 168 |
rep_distances = torch.cdist(
|
| 169 |
rep_coords, rep_coords, compute_mode="donot_use_mm_for_euclid_dist"
|
| 170 |
-
)
|
| 171 |
-
distogram_bins = (rep_distances.unsqueeze(-1) > self.boundaries).sum(dim=-1).long()
|
| 172 |
-
pair = pair + self.dist_bin_pairwise_embed(distogram_bins)
|
| 173 |
-
|
| 174 |
-
pair_mask = mask[:, :, None].float() * mask[:, None, :].float()
|
| 175 |
-
pair = pair + self.folding_trunk(pair, pair_attention_mask=pair_mask)
|
| 176 |
-
single = self.row_attention_pooling(pair, mask)
|
| 177 |
-
|
| 178 |
-
atom_mask_f = atom_mask_m.float()
|
| 179 |
-
s_at_atoms = gather_token_to_atom(single, atom_to_token_m)
|
| 180 |
-
s_at_atoms = self.plddt_ln(s_at_atoms)
|
| 181 |
-
intra_idx = _compute_intra_token_idx(atom_to_token_m)
|
| 182 |
-
intra_idx = intra_idx.clamp(max=self.plddt_weight.shape[0] - 1)
|
| 183 |
-
plddt_weight = self.plddt_weight[intra_idx]
|
| 184 |
-
plddt_logits = torch.einsum("...c,...cb->...b", s_at_atoms, plddt_weight)
|
| 185 |
-
plddt_per_atom = _categorical_mean(plddt_logits, start=0.0, end=1.0)
|
| 186 |
|
| 187 |
length = single.shape[1]
|
| 188 |
plddt_sum = torch.zeros(
|
| 189 |
batch_mult, length, device=single.device, dtype=plddt_per_atom.dtype
|
| 190 |
-
)
|
| 191 |
atom_count = torch.zeros(
|
| 192 |
batch_mult, length, device=single.device, dtype=plddt_per_atom.dtype
|
| 193 |
-
)
|
| 194 |
-
atom_mask_t = atom_mask_f.to(plddt_per_atom.dtype)
|
| 195 |
-
plddt_sum.scatter_add_(1, atom_to_token_m, plddt_per_atom * atom_mask_t)
|
| 196 |
-
atom_count.scatter_add_(1, atom_to_token_m, atom_mask_t)
|
| 197 |
-
plddt = plddt_sum / atom_count.clamp(min=1e-6)
|
| 198 |
|
| 199 |
complex_plddt = (plddt_per_atom * atom_mask_f).sum(dim=-1) / (
|
| 200 |
atom_mask_f.sum(dim=-1) + _EPS
|
| 201 |
-
)
|
| 202 |
|
| 203 |
-
expanded_type = self._repeat_batch(mol_type, num_diffusion_samples)
|
| 204 |
-
expanded_asym = self._repeat_batch(asym_id, num_diffusion_samples)
|
| 205 |
-
is_ligand = (expanded_type == _NONPOLYMER_ID).float()
|
| 206 |
-
inter_chain = (expanded_asym.unsqueeze(-1) != expanded_asym.unsqueeze(-2)).float()
|
| 207 |
-
near_contact = (rep_distances < 8).float()
|
| 208 |
interface_per_token = (near_contact * inter_chain * (1.0 - is_ligand).unsqueeze(-1)).amax(
|
| 209 |
dim=-1
|
| 210 |
-
)
|
| 211 |
iplddt_weight = torch.where(
|
| 212 |
is_ligand.bool(),
|
| 213 |
torch.full_like(interface_per_token, 2.0),
|
| 214 |
interface_per_token,
|
| 215 |
-
)
|
| 216 |
iplddt_weight_atoms = gather_token_to_atom(
|
| 217 |
iplddt_weight.unsqueeze(-1), atom_to_token_m
|
| 218 |
-
).squeeze(-1)
|
| 219 |
-
atom_iplddt_w = atom_mask_f * iplddt_weight_atoms
|
| 220 |
complex_iplddt = (plddt_per_atom * atom_iplddt_w).sum(dim=-1) / (
|
| 221 |
atom_iplddt_w.sum(dim=-1) + _EPS
|
| 222 |
-
)
|
| 223 |
-
plddt_ca = plddt_per_atom.gather(1, rep_idx_m)
|
| 224 |
|
| 225 |
-
pae_logits = self.pae_head(pair)
|
| 226 |
-
pae = _categorical_mean(pae_logits, start=0.0, end=32.0).detach()
|
| 227 |
|
| 228 |
n_bins = pae_logits.shape[-1]
|
| 229 |
bin_width = 32.0 / n_bins
|
| 230 |
-
bin_centers = torch.arange(0.5 * bin_width, 32.0, bin_width, device=pae_logits.device)
|
| 231 |
-
mask_f = mask.float()
|
| 232 |
-
n_res = mask_f.sum(dim=-1, keepdim=True)
|
| 233 |
-
d0 = 1.24 * (n_res.clamp(min=19) - 15) ** (1 / 3) - 1.8
|
| 234 |
-
tm_per_bin = 1 / (1 + (bin_centers / d0) ** 2)
|
| 235 |
-
pae_probs = F.softmax(pae_logits, dim=-1)
|
| 236 |
-
tm_expected = (pae_probs * tm_per_bin[:, None, None, :]).sum(dim=-1)
|
| 237 |
-
|
| 238 |
-
pair_mask_2d = mask_f.unsqueeze(-1) * mask_f.unsqueeze(-2)
|
| 239 |
-
ptm_per_row = (tm_expected * pair_mask_2d).sum(dim=-1) / (pair_mask_2d.sum(dim=-1) + _EPS)
|
| 240 |
-
ptm = ptm_per_row.max(dim=-1).values
|
| 241 |
|
| 242 |
inter_chain_mask = (
|
| 243 |
expanded_asym.unsqueeze(-1) != expanded_asym.unsqueeze(-2)
|
| 244 |
-
).float() * pair_mask_2d
|
| 245 |
iptm_per_row = (tm_expected * inter_chain_mask).sum(dim=-1) / (
|
| 246 |
inter_chain_mask.sum(dim=-1) + _EPS
|
| 247 |
-
)
|
| 248 |
-
iptm = iptm_per_row.max(dim=-1).values
|
| 249 |
|
| 250 |
max_chain_id = int(expanded_asym.max().item()) if batch_mult > 0 else 0
|
| 251 |
n_chains = max_chain_id + 1
|
|
@@ -255,16 +259,16 @@ class ConfidenceHead(nn.Module):
|
|
| 255 |
n_chains,
|
| 256 |
device=tm_expected.device,
|
| 257 |
dtype=tm_expected.dtype,
|
| 258 |
-
)
|
| 259 |
for c1 in range(n_chains):
|
| 260 |
-
chain_c1 = (expanded_asym == c1).float() * mask_f
|
| 261 |
if chain_c1.sum() == 0:
|
| 262 |
continue
|
| 263 |
for c2 in range(n_chains):
|
| 264 |
-
chain_c2 = (expanded_asym == c2).float() * mask_f
|
| 265 |
-
pair_m = chain_c1.unsqueeze(-1) * chain_c2.unsqueeze(-2)
|
| 266 |
-
denom = pair_m.sum(dim=(-1, -2)) + _EPS
|
| 267 |
-
pair_chains_iptm[:, c1, c2] = (tm_expected * pair_m).sum(dim=(-1, -2)) / denom
|
| 268 |
|
| 269 |
return {
|
| 270 |
"plddt_logits": plddt_logits,
|
|
@@ -278,7 +282,7 @@ class ConfidenceHead(nn.Module):
|
|
| 278 |
"ptm": ptm.detach(),
|
| 279 |
"iptm": iptm.detach(),
|
| 280 |
"pair_chains_iptm": pair_chains_iptm.detach(),
|
| 281 |
-
}
|
| 282 |
|
| 283 |
|
| 284 |
class _TransitionFFN(nn.Module):
|
|
@@ -288,7 +292,8 @@ class _TransitionFFN(nn.Module):
|
|
| 288 |
self.ffn = SwiGLUMLP(d_model, expansion_ratio=expansion_ratio, bias=False)
|
| 289 |
|
| 290 |
def forward(self, x: Tensor) -> Tensor:
|
| 291 |
-
|
|
|
|
| 292 |
|
| 293 |
|
| 294 |
class MSAEncoderBlock(nn.Module):
|
|
@@ -327,44 +332,45 @@ class MSAEncoderBlock(nn.Module):
|
|
| 327 |
pair_attention_mask: Tensor,
|
| 328 |
msa_track_mask: Tensor | None = None,
|
| 329 |
) -> tuple[Tensor, Tensor]:
|
|
|
|
| 330 |
mask4d = (
|
| 331 |
msa_track_mask[:, None, None, None].to(dtype=msa_repr.dtype)
|
| 332 |
if msa_track_mask is not None
|
| 333 |
else None
|
| 334 |
-
)
|
| 335 |
|
| 336 |
-
pair_mask4d = mask4d[:, :, :1] if mask4d is not None else None
|
| 337 |
|
| 338 |
-
msa_update = self.msa_pair_weighted_averaging(msa_repr, pair_repr, pair_attention_mask)
|
| 339 |
if mask4d is not None:
|
| 340 |
-
msa_update = msa_update * mask4d
|
| 341 |
-
msa_repr = msa_repr + msa_update
|
| 342 |
|
| 343 |
-
msa_transition = self.msa_transition(msa_repr)
|
| 344 |
if mask4d is not None:
|
| 345 |
-
msa_transition = msa_transition * mask4d
|
| 346 |
-
msa_repr = msa_repr + msa_transition
|
| 347 |
|
| 348 |
-
pair_opm = self.outer_product_mean(msa_repr, msa_attention_mask)
|
| 349 |
if pair_mask4d is not None:
|
| 350 |
-
pair_opm = pair_opm * pair_mask4d
|
| 351 |
-
pair_repr = pair_repr + pair_opm
|
| 352 |
|
| 353 |
-
pair_out = self.tri_mul_out(pair_repr, mask=pair_attention_mask)
|
| 354 |
if pair_mask4d is not None:
|
| 355 |
-
pair_out = pair_out * pair_mask4d
|
| 356 |
-
pair_repr = pair_repr + pair_out
|
| 357 |
|
| 358 |
-
pair_in = self.tri_mul_in(pair_repr, mask=pair_attention_mask)
|
| 359 |
if pair_mask4d is not None:
|
| 360 |
-
pair_in = pair_in * pair_mask4d
|
| 361 |
-
pair_repr = pair_repr + pair_in
|
| 362 |
|
| 363 |
-
pair_transition = self.pair_transition(pair_repr)
|
| 364 |
if pair_mask4d is not None:
|
| 365 |
-
pair_transition = pair_transition * pair_mask4d
|
| 366 |
-
pair_repr = pair_repr + pair_transition
|
| 367 |
-
return msa_repr, pair_repr
|
| 368 |
|
| 369 |
|
| 370 |
class MSAEncoder(nn.Module):
|
|
@@ -407,18 +413,19 @@ class MSAEncoder(nn.Module):
|
|
| 407 |
deletion_value: Tensor,
|
| 408 |
msa_attention_mask: Tensor,
|
| 409 |
) -> Tensor:
|
|
|
|
| 410 |
batch_size, _, depth = msa_attention_mask.shape
|
| 411 |
m_feat = torch.cat(
|
| 412 |
[msa_oh, has_deletion.unsqueeze(-1), deletion_value.unsqueeze(-1)],
|
| 413 |
dim=-1,
|
| 414 |
-
)
|
| 415 |
-
m = self.embed(m_feat) + self.project_inputs(x_inputs).unsqueeze(2)
|
| 416 |
if depth > 1:
|
| 417 |
-
msa_track_mask = msa_attention_mask[:, :, 1:].any(dim=(1, 2))
|
| 418 |
else:
|
| 419 |
-
msa_track_mask = torch.zeros(batch_size, dtype=torch.bool, device=x_pair.device)
|
| 420 |
-
tok_mask = msa_attention_mask[:, :, 0]
|
| 421 |
-
pair_attention_mask = tok_mask.unsqueeze(2) * tok_mask.unsqueeze(1)
|
| 422 |
for block in self.blocks:
|
| 423 |
m, x_pair = cast(MSAEncoderBlock, block)(
|
| 424 |
m,
|
|
@@ -426,8 +433,8 @@ class MSAEncoder(nn.Module):
|
|
| 426 |
msa_attention_mask,
|
| 427 |
pair_attention_mask,
|
| 428 |
msa_track_mask,
|
| 429 |
-
)
|
| 430 |
-
return x_pair * msa_track_mask[:, None, None, None].to(dtype=x_pair.dtype)
|
| 431 |
|
| 432 |
|
| 433 |
class ESMFold2ExperimentalModel(ESMFold2EmbeddingMixin, ESMFold2AttentionMixin, PreTrainedModel):
|
|
@@ -477,7 +484,7 @@ class ESMFold2ExperimentalModel(ESMFold2EmbeddingMixin, ESMFold2AttentionMixin,
|
|
| 477 |
self.pair_loop_proj = nn.Sequential(
|
| 478 |
nn.LayerNorm(d_pair), nn.Linear(d_pair, d_pair, bias=False)
|
| 479 |
)
|
| 480 |
-
nn.init.zeros_(cast(nn.Linear, self.pair_loop_proj[1]).weight)
|
| 481 |
|
| 482 |
self.structure_head = DiffusionStructureHead(config)
|
| 483 |
self.distogram_head = nn.Linear(d_pair, config.structure_head.distogram_bins, bias=True)
|
|
@@ -648,6 +655,7 @@ class ESMFold2ExperimentalModel(ESMFold2EmbeddingMixin, ESMFold2AttentionMixin,
|
|
| 648 |
tok_mask: Tensor,
|
| 649 |
verbose: bool = False,
|
| 650 |
) -> Tensor:
|
|
|
|
| 651 |
if self._esmc_fp8 and torch.is_grad_enabled():
|
| 652 |
_reload_esmc_bf16_for_gradients(
|
| 653 |
self,
|
|
@@ -670,9 +678,9 @@ class ESMFold2ExperimentalModel(ESMFold2EmbeddingMixin, ESMFold2AttentionMixin,
|
|
| 670 |
mol_type,
|
| 671 |
tok_mask,
|
| 672 |
pad_to_multiple=pad_to,
|
| 673 |
-
)
|
| 674 |
progress.update()
|
| 675 |
-
return result
|
| 676 |
return compute_lm_hidden_states(
|
| 677 |
self._esmc,
|
| 678 |
input_ids,
|
|
@@ -681,7 +689,7 @@ class ESMFold2ExperimentalModel(ESMFold2EmbeddingMixin, ESMFold2AttentionMixin,
|
|
| 681 |
mol_type,
|
| 682 |
tok_mask,
|
| 683 |
pad_to_multiple=pad_to,
|
| 684 |
-
)
|
| 685 |
|
| 686 |
def forward(
|
| 687 |
self,
|
|
@@ -731,6 +739,7 @@ class ESMFold2ExperimentalModel(ESMFold2EmbeddingMixin, ESMFold2AttentionMixin,
|
|
| 731 |
disto_cond_mask: Tensor | None = None,
|
| 732 |
verbose: bool = False,
|
| 733 |
) -> ESMFold2Output | tuple[Any, ...]:
|
|
|
|
| 734 |
output_hidden_states, return_dict = _resolve_structure_output_controls(
|
| 735 |
self.config,
|
| 736 |
output_attentions=output_attentions,
|
|
@@ -751,8 +760,8 @@ class ESMFold2ExperimentalModel(ESMFold2EmbeddingMixin, ESMFold2AttentionMixin,
|
|
| 751 |
disto_cond_mask=disto_cond_mask,
|
| 752 |
)
|
| 753 |
del gt_coords, is_resolved, frames_idx
|
| 754 |
-
tok_mask = token_attention_mask
|
| 755 |
-
atm_mask = atom_attention_mask
|
| 756 |
n_loops = num_loops if num_loops is not None else self.config.num_loops
|
| 757 |
n_samples = (
|
| 758 |
num_diffusion_samples
|
|
@@ -761,46 +770,46 @@ class ESMFold2ExperimentalModel(ESMFold2EmbeddingMixin, ESMFold2AttentionMixin,
|
|
| 761 |
)
|
| 762 |
|
| 763 |
if res_type.dim() == 2:
|
| 764 |
-
res_type_oh = F.one_hot(res_type.long(), num_classes=NUM_RES_TYPES).float()
|
| 765 |
-
res_type_oh = res_type_oh * tok_mask.unsqueeze(-1).float()
|
| 766 |
else:
|
| 767 |
-
res_type_oh = res_type.float()
|
| 768 |
|
| 769 |
if msa is not None:
|
| 770 |
-
msa_oh_profile = F.one_hot(msa.long(), num_classes=NUM_RES_TYPES).float()
|
| 771 |
if msa_attention_mask is not None:
|
| 772 |
-
mask_f = msa_attention_mask.float().unsqueeze(-1)
|
| 773 |
-
msa_oh_profile = msa_oh_profile * mask_f
|
| 774 |
-
valid_seq_count = msa_attention_mask.float().sum(dim=1).clamp(min=1)
|
| 775 |
-
profile = msa_oh_profile.sum(dim=1) / valid_seq_count.unsqueeze(-1)
|
| 776 |
else:
|
| 777 |
-
profile = msa_oh_profile.mean(dim=1)
|
| 778 |
else:
|
| 779 |
-
profile = res_type_oh
|
| 780 |
|
| 781 |
if res_type_soft is not None:
|
| 782 |
-
res_type_oh = res_type_soft.float()
|
| 783 |
if not self.config.disable_msa_features and provide_soft_sequence_to_msa_and_profile:
|
| 784 |
-
profile = res_type_oh
|
| 785 |
-
msa = res_type_oh.unsqueeze(1)
|
| 786 |
-
msa_attention_mask = tok_mask.unsqueeze(1)
|
| 787 |
|
| 788 |
if deletion_mean is None:
|
| 789 |
deletion_mean = torch.zeros(
|
| 790 |
res_type.shape[0], res_type.shape[1], device=res_type.device
|
| 791 |
-
)
|
| 792 |
if self.config.disable_msa_features:
|
| 793 |
-
profile = torch.zeros_like(profile)
|
| 794 |
-
deletion_mean = torch.zeros_like(deletion_mean)
|
| 795 |
|
| 796 |
-
ref_element_oh = F.one_hot(ref_element.long(), num_classes=MAX_ATOMIC_NUMBER).float()
|
| 797 |
ref_atom_name_chars_oh = F.one_hot(
|
| 798 |
ref_atom_name_chars.long(), num_classes=CHAR_VOCAB_SIZE
|
| 799 |
-
).float()
|
| 800 |
-
atm_mask_f = atm_mask.float()
|
| 801 |
-
ref_element_oh = ref_element_oh * atm_mask_f.unsqueeze(-1)
|
| 802 |
-
ref_atom_name_chars_oh = ref_atom_name_chars_oh * atm_mask_f.unsqueeze(-1).unsqueeze(-1)
|
| 803 |
-
atom_to_token = atom_to_token * atm_mask.long()
|
| 804 |
|
| 805 |
use_amp = ref_pos.device.type == "cuda"
|
| 806 |
with torch.amp.autocast("cuda", enabled=use_amp, dtype=torch.bfloat16):
|
|
@@ -815,58 +824,58 @@ class ESMFold2ExperimentalModel(ESMFold2EmbeddingMixin, ESMFold2AttentionMixin,
|
|
| 815 |
ref_element=ref_element_oh,
|
| 816 |
ref_atom_name_chars=ref_atom_name_chars_oh,
|
| 817 |
atom_to_token=atom_to_token,
|
| 818 |
-
)
|
| 819 |
|
| 820 |
-
z_init = self.z_init_1(x_inputs).unsqueeze(2) + self.z_init_2(x_inputs).unsqueeze(1)
|
| 821 |
relative_position_encoding = self.rel_pos(
|
| 822 |
residue_index=residue_index,
|
| 823 |
asym_id=asym_id,
|
| 824 |
sym_id=sym_id,
|
| 825 |
entity_id=entity_id,
|
| 826 |
token_index=token_index,
|
| 827 |
-
)
|
| 828 |
-
token_bonds_encoding = self.token_bonds(token_bonds.float())
|
| 829 |
-
z_init = z_init + relative_position_encoding + token_bonds_encoding
|
| 830 |
|
| 831 |
if lm_hidden_states is None and input_ids is not None and self._esmc is not None:
|
| 832 |
lm_hidden_states = self._compute_lm_hidden_states(
|
| 833 |
input_ids, asym_id, residue_index, mol_type, tok_mask, verbose=verbose
|
| 834 |
-
)
|
| 835 |
if lm_hidden_states is not None:
|
| 836 |
lm_dropout = (
|
| 837 |
self.config.lm_dropout
|
| 838 |
if self.config.force_lm_dropout_during_inference or self.training
|
| 839 |
else 0.0
|
| 840 |
)
|
| 841 |
-
lm_z = self.language_model(lm_hidden_states.detach(), lm_dropout=lm_dropout)
|
| 842 |
-
z_init = z_init + lm_z.to(z_init.dtype)
|
| 843 |
|
| 844 |
msa_kwargs: dict[str, Tensor] | None = None
|
| 845 |
if self.msa_encoder is not None and msa is not None:
|
| 846 |
if msa.dim() == 4:
|
| 847 |
batch_msa, depth, length_msa, _ = msa.shape
|
| 848 |
-
msa_oh = msa.permute(0, 2, 1, 3).float()
|
| 849 |
else:
|
| 850 |
batch_msa, depth, length_msa = msa.shape
|
| 851 |
msa_oh = F.one_hot(
|
| 852 |
msa.permute(0, 2, 1).long(), num_classes=NUM_RES_TYPES
|
| 853 |
-
).float()
|
| 854 |
msa_attn = (
|
| 855 |
msa_attention_mask.permute(0, 2, 1).float()
|
| 856 |
if msa_attention_mask is not None
|
| 857 |
else tok_mask[:, :, None].expand(-1, -1, depth).float()
|
| 858 |
-
)
|
| 859 |
-
msa_oh = msa_oh * msa_attn.unsqueeze(-1)
|
| 860 |
hd = (
|
| 861 |
has_deletion.permute(0, 2, 1).float()
|
| 862 |
if has_deletion is not None
|
| 863 |
else torch.zeros(batch_msa, length_msa, depth, device=msa.device)
|
| 864 |
-
)
|
| 865 |
dv = (
|
| 866 |
deletion_value.permute(0, 2, 1).float()
|
| 867 |
if deletion_value is not None
|
| 868 |
else torch.zeros(batch_msa, length_msa, depth, device=msa.device)
|
| 869 |
-
)
|
| 870 |
msa_kwargs = {
|
| 871 |
"x_inputs": x_inputs,
|
| 872 |
"msa_oh": msa_oh,
|
|
@@ -875,10 +884,10 @@ class ESMFold2ExperimentalModel(ESMFold2EmbeddingMixin, ESMFold2AttentionMixin,
|
|
| 875 |
"msa_attention_mask": msa_attn,
|
| 876 |
}
|
| 877 |
|
| 878 |
-
pair_mask = tok_mask[:, :, None].float() * tok_mask[:, None, :].float()
|
| 879 |
-
z = torch.zeros_like(z_init)
|
| 880 |
-
prev_pair: Tensor | None = None
|
| 881 |
-
prev_disto_probs: Tensor | None = None
|
| 882 |
loop_iterator = range(n_loops + 1)
|
| 883 |
if verbose:
|
| 884 |
loop_iterator = tqdm(
|
|
@@ -889,32 +898,32 @@ class ESMFold2ExperimentalModel(ESMFold2EmbeddingMixin, ESMFold2AttentionMixin,
|
|
| 889 |
)
|
| 890 |
|
| 891 |
for loop_num in loop_iterator:
|
| 892 |
-
z = z_init + self.pair_loop_proj(z)
|
| 893 |
if msa_kwargs is not None and self.msa_encoder is not None:
|
| 894 |
-
z = z + self.msa_encoder(x_pair=z, **msa_kwargs).to(z.dtype)
|
| 895 |
-
z = self.folding_trunk(z, pair_attention_mask=pair_mask)
|
| 896 |
|
| 897 |
if early_exit and loop_num < n_loops:
|
| 898 |
l2_converged = False
|
| 899 |
if prev_pair is not None and loop_num > 0:
|
| 900 |
rel_l2 = (
|
| 901 |
z.float() - prev_pair.float()
|
| 902 |
-
).norm() / prev_pair.float().norm().clamp(min=1e-8)
|
| 903 |
l2_converged = rel_l2.item() < 0.25
|
| 904 |
-
prev_pair = z.detach().clone()
|
| 905 |
-
sym_z = z.float() + z.float().transpose(-2, -3)
|
| 906 |
-
cur_probs = F.softmax(self.distogram_head(sym_z).float(), dim=-1)
|
| 907 |
if prev_disto_probs is not None and loop_num > 0:
|
| 908 |
kl_per_pair = (
|
| 909 |
cur_probs
|
| 910 |
* (cur_probs.clamp(min=1e-8) / prev_disto_probs.clamp(min=1e-8)).log()
|
| 911 |
-
).sum(-1)
|
| 912 |
-
kl = (kl_per_pair + kl_per_pair.transpose(-1, -2)).mean() / 2
|
| 913 |
if l2_converged or kl.item() < 0.05:
|
| 914 |
break
|
| 915 |
-
prev_disto_probs = cur_probs.detach()
|
| 916 |
|
| 917 |
-
distogram_logits = self.distogram_head(z + z.transpose(-2, -3))
|
| 918 |
|
| 919 |
with torch.no_grad(), _seed_context(seed):
|
| 920 |
structure_output = self.structure_head.sample(
|
|
@@ -943,8 +952,8 @@ class ESMFold2ExperimentalModel(ESMFold2EmbeddingMixin, ESMFold2AttentionMixin,
|
|
| 943 |
return_atom_repr=False,
|
| 944 |
denoising_early_exit_rmsd=(0.10 if early_exit else None),
|
| 945 |
verbose=verbose,
|
| 946 |
-
)
|
| 947 |
-
sample_coords = structure_output["sample_atom_coords"]
|
| 948 |
if sample_coords is None:
|
| 949 |
raise RuntimeError("ESMFold2 structure sampling did not return coordinates.")
|
| 950 |
if sample_coords.ndim == 4:
|
|
@@ -953,15 +962,15 @@ class ESMFold2ExperimentalModel(ESMFold2EmbeddingMixin, ESMFold2AttentionMixin,
|
|
| 953 |
batch * sample_count,
|
| 954 |
atom_count,
|
| 955 |
coord_dim,
|
| 956 |
-
)
|
| 957 |
-
rep_idx = distogram_atom_idx.repeat_interleave(sample_count, 0).long()
|
| 958 |
else:
|
| 959 |
-
sample_coords_for_gather = sample_coords
|
| 960 |
-
rep_idx = distogram_atom_idx.long()
|
| 961 |
representative_atom_coords = gather_rep_atom_coords(
|
| 962 |
sample_coords_for_gather,
|
| 963 |
rep_idx,
|
| 964 |
-
)
|
| 965 |
|
| 966 |
output: dict[str, Tensor] = {
|
| 967 |
"distogram_logits": distogram_logits,
|
|
@@ -990,7 +999,7 @@ class ESMFold2ExperimentalModel(ESMFold2EmbeddingMixin, ESMFold2AttentionMixin,
|
|
| 990 |
num_diffusion_samples=n_samples,
|
| 991 |
relative_position_encoding=relative_position_encoding.detach(),
|
| 992 |
token_bonds_encoding=token_bonds_encoding.detach(),
|
| 993 |
-
)
|
| 994 |
progress.update()
|
| 995 |
else:
|
| 996 |
confidence_output = self.confidence_head(
|
|
@@ -1006,18 +1015,18 @@ class ESMFold2ExperimentalModel(ESMFold2EmbeddingMixin, ESMFold2AttentionMixin,
|
|
| 1006 |
num_diffusion_samples=n_samples,
|
| 1007 |
relative_position_encoding=relative_position_encoding.detach(),
|
| 1008 |
token_bonds_encoding=token_bonds_encoding.detach(),
|
| 1009 |
-
)
|
| 1010 |
output.update(confidence_output)
|
| 1011 |
-
output["atom_pad_mask"] = atm_mask.unsqueeze(0) if atm_mask.dim() == 1 else atm_mask
|
| 1012 |
-
output["residue_index"] = residue_index
|
| 1013 |
-
output["entity_id"] = entity_id
|
| 1014 |
return _finalize_structure_output(
|
| 1015 |
output,
|
| 1016 |
token_input_state=x_inputs,
|
| 1017 |
pair_state=z,
|
| 1018 |
output_hidden_states=output_hidden_states,
|
| 1019 |
return_dict=return_dict,
|
| 1020 |
-
)
|
| 1021 |
|
| 1022 |
@property
|
| 1023 |
def input_builder(self):
|
|
@@ -1058,7 +1067,7 @@ class ESMFold2ExperimentalModel(ESMFold2EmbeddingMixin, ESMFold2AttentionMixin,
|
|
| 1058 |
if not self.config.msa_conditioning:
|
| 1059 |
for name in MSA_CONDITIONING_INPUT_NAMES:
|
| 1060 |
features.pop(name, None)
|
| 1061 |
-
features = {name: tensor.to(self.device) for name, tensor in features.items()}
|
| 1062 |
output = self(**features, **forward_kwargs, return_dict=True)
|
| 1063 |
for name in (
|
| 1064 |
"res_type",
|
|
|
|
| 9 |
from __future__ import annotations
|
| 10 |
|
| 11 |
import gc
|
|
|
|
|
|
|
|
|
|
|
|
|
| 12 |
import torch
|
| 13 |
import torch.nn as nn
|
| 14 |
import torch.nn.functional as F
|
| 15 |
+
|
| 16 |
+
from collections.abc import Mapping
|
| 17 |
+
from pathlib import Path
|
| 18 |
+
from typing import Any, ClassVar, cast
|
| 19 |
from torch import Tensor
|
| 20 |
from tqdm.auto import tqdm
|
| 21 |
from transformers.modeling_utils import PreTrainedModel
|
|
|
|
| 66 |
validate_prepared_auxiliary_inputs,
|
| 67 |
)
|
| 68 |
|
| 69 |
+
|
| 70 |
_EPS = 1e-5
|
| 71 |
_NONPOLYMER_ID = 3
|
| 72 |
|
|
|
|
| 83 |
d_pair = config.d_pair
|
| 84 |
d_inputs = config.inputs.d_inputs
|
| 85 |
|
| 86 |
+
boundaries = torch.linspace(ch.min_dist, ch.max_dist, ch.distogram_bins - 1) # (distogram_bins - 1,)
|
| 87 |
+
self.register_buffer("boundaries", boundaries) # (distogram_bins - 1,)
|
| 88 |
self.dist_bin_pairwise_embed = nn.Embedding(ch.distogram_bins, d_pair)
|
| 89 |
|
| 90 |
self.s_norm = nn.LayerNorm(d_single)
|
|
|
|
| 106 |
max_atoms_per_token = 23
|
| 107 |
self.plddt_weight = nn.Parameter(
|
| 108 |
torch.zeros(max_atoms_per_token, d_single, ch.num_plddt_bins)
|
| 109 |
+
) # (23, d_single, n_plddt_bins)
|
| 110 |
self.pae_head = nn.Linear(d_pair, ch.num_pae_bins, bias=False)
|
| 111 |
|
| 112 |
def set_kernel_backend(self, backend: str | None) -> None:
|
|
|
|
| 118 |
|
| 119 |
@staticmethod
|
| 120 |
def _repeat_batch(x: Tensor, num_diffusion_samples: int) -> Tensor:
|
| 121 |
+
# x: (b, ...); output repeats the batch axis by samples.
|
| 122 |
if num_diffusion_samples == 1:
|
| 123 |
+
return x # (b * samples, ...), including samples = 1
|
| 124 |
+
return x.repeat_interleave(num_diffusion_samples, 0) # (b * samples, ...), including samples = 1
|
| 125 |
|
| 126 |
@staticmethod
|
| 127 |
def _flatten_sample_axis(x: Tensor) -> Tensor:
|
| 128 |
+
# x: (b, samples, n, c) or an already flattened tensor.
|
| 129 |
if x.ndim == 4:
|
| 130 |
b, mult, n, c = x.shape
|
| 131 |
+
return x.reshape(b * mult, n, c) # (b * samples, n, c) for 4D input; otherwise x.shape
|
| 132 |
+
return x # (b * samples, n, c) for 4D input; otherwise x.shape
|
| 133 |
|
| 134 |
def forward(
|
| 135 |
self,
|
|
|
|
| 146 |
relative_position_encoding: Tensor | None = None,
|
| 147 |
token_bonds_encoding: Tensor | None = None,
|
| 148 |
) -> dict[str, Tensor]:
|
| 149 |
+
# s_inputs: (b, l, d_inputs); z: (b, l, l, d_pair); x_pred: (bs, a, 3) or (b, samples, a, 3). bs = b * samples.
|
| 150 |
+
s_inputs_normed = self.s_inputs_norm(s_inputs) # (b, l, d_inputs)
|
| 151 |
+
z_base = self.z_norm(z) # (b, l, l, d_pair)
|
| 152 |
if relative_position_encoding is not None:
|
| 153 |
+
z_base = z_base + relative_position_encoding # (b, l, l, d_pair)
|
| 154 |
if token_bonds_encoding is not None:
|
| 155 |
+
z_base = z_base + token_bonds_encoding # (b, l, l, d_pair)
|
| 156 |
+
z_base = z_base + self.s_to_z(s_inputs_normed).unsqueeze(2) # (b, l, l, d_pair)
|
| 157 |
+
z_base = z_base + self.s_to_z_transpose(s_inputs_normed).unsqueeze(1) # (b, l, l, d_pair)
|
| 158 |
z_base = z_base + self.s_to_z_prod_out(
|
| 159 |
self.s_to_z_prod_in1(s_inputs_normed)[:, :, None, :]
|
| 160 |
* self.s_to_z_prod_in2(s_inputs_normed)[:, None, :, :]
|
| 161 |
+
) # (b, l, l, d_pair)
|
| 162 |
+
|
| 163 |
+
pair = self._repeat_batch(z_base, num_diffusion_samples) # (bs, l, l, d_pair)
|
| 164 |
+
x_pred_flat = self._flatten_sample_axis(x_pred) # (bs, a, 3)
|
| 165 |
+
atom_to_token_m = self._repeat_batch(atom_to_token, num_diffusion_samples) # (bs, a)
|
| 166 |
+
atom_mask_m = self._repeat_batch(atom_attention_mask, num_diffusion_samples) # (bs, a)
|
| 167 |
+
rep_idx_m = self._repeat_batch(distogram_atom_idx, num_diffusion_samples).long() # (bs, l)
|
| 168 |
+
mask = self._repeat_batch(token_attention_mask, num_diffusion_samples) # (bs, l)
|
| 169 |
batch_mult = pair.shape[0]
|
| 170 |
|
| 171 |
+
rep_coords = gather_rep_atom_coords(x_pred_flat, rep_idx_m) # (bs, l, 3)
|
| 172 |
rep_distances = torch.cdist(
|
| 173 |
rep_coords, rep_coords, compute_mode="donot_use_mm_for_euclid_dist"
|
| 174 |
+
) # (bs, l, l)
|
| 175 |
+
distogram_bins = (rep_distances.unsqueeze(-1) > self.boundaries).sum(dim=-1).long() # (bs, l, l)
|
| 176 |
+
pair = pair + self.dist_bin_pairwise_embed(distogram_bins) # (bs, l, l, d_pair)
|
| 177 |
+
|
| 178 |
+
pair_mask = mask[:, :, None].float() * mask[:, None, :].float() # (bs, l, l)
|
| 179 |
+
pair = pair + self.folding_trunk(pair, pair_attention_mask=pair_mask) # (bs, l, l, d_pair)
|
| 180 |
+
single = self.row_attention_pooling(pair, mask) # (bs, l, d_single)
|
| 181 |
+
|
| 182 |
+
atom_mask_f = atom_mask_m.float() # (bs, a)
|
| 183 |
+
s_at_atoms = gather_token_to_atom(single, atom_to_token_m) # (bs, a, d_single)
|
| 184 |
+
s_at_atoms = self.plddt_ln(s_at_atoms) # (bs, a, d_single)
|
| 185 |
+
intra_idx = _compute_intra_token_idx(atom_to_token_m) # (bs, a)
|
| 186 |
+
intra_idx = intra_idx.clamp(max=self.plddt_weight.shape[0] - 1) # (bs, a)
|
| 187 |
+
plddt_weight = self.plddt_weight[intra_idx] # (bs, a, d_single, n_plddt_bins)
|
| 188 |
+
plddt_logits = torch.einsum("...c,...cb->...b", s_at_atoms, plddt_weight) # (bs, a, n_plddt_bins)
|
| 189 |
+
plddt_per_atom = _categorical_mean(plddt_logits, start=0.0, end=1.0) # (bs, a)
|
| 190 |
|
| 191 |
length = single.shape[1]
|
| 192 |
plddt_sum = torch.zeros(
|
| 193 |
batch_mult, length, device=single.device, dtype=plddt_per_atom.dtype
|
| 194 |
+
) # (bs, l)
|
| 195 |
atom_count = torch.zeros(
|
| 196 |
batch_mult, length, device=single.device, dtype=plddt_per_atom.dtype
|
| 197 |
+
) # (bs, l)
|
| 198 |
+
atom_mask_t = atom_mask_f.to(plddt_per_atom.dtype) # (bs, a)
|
| 199 |
+
plddt_sum.scatter_add_(1, atom_to_token_m, plddt_per_atom * atom_mask_t) # (bs, l)
|
| 200 |
+
atom_count.scatter_add_(1, atom_to_token_m, atom_mask_t) # (bs, l)
|
| 201 |
+
plddt = plddt_sum / atom_count.clamp(min=1e-6) # (bs, l)
|
| 202 |
|
| 203 |
complex_plddt = (plddt_per_atom * atom_mask_f).sum(dim=-1) / (
|
| 204 |
atom_mask_f.sum(dim=-1) + _EPS
|
| 205 |
+
) # (bs,)
|
| 206 |
|
| 207 |
+
expanded_type = self._repeat_batch(mol_type, num_diffusion_samples) # (bs, l)
|
| 208 |
+
expanded_asym = self._repeat_batch(asym_id, num_diffusion_samples) # (bs, l)
|
| 209 |
+
is_ligand = (expanded_type == _NONPOLYMER_ID).float() # (bs, l)
|
| 210 |
+
inter_chain = (expanded_asym.unsqueeze(-1) != expanded_asym.unsqueeze(-2)).float() # (bs, l, l)
|
| 211 |
+
near_contact = (rep_distances < 8).float() # (bs, l, l)
|
| 212 |
interface_per_token = (near_contact * inter_chain * (1.0 - is_ligand).unsqueeze(-1)).amax(
|
| 213 |
dim=-1
|
| 214 |
+
) # (bs, l)
|
| 215 |
iplddt_weight = torch.where(
|
| 216 |
is_ligand.bool(),
|
| 217 |
torch.full_like(interface_per_token, 2.0),
|
| 218 |
interface_per_token,
|
| 219 |
+
) # (bs, l)
|
| 220 |
iplddt_weight_atoms = gather_token_to_atom(
|
| 221 |
iplddt_weight.unsqueeze(-1), atom_to_token_m
|
| 222 |
+
).squeeze(-1) # (bs, a)
|
| 223 |
+
atom_iplddt_w = atom_mask_f * iplddt_weight_atoms # (bs, a)
|
| 224 |
complex_iplddt = (plddt_per_atom * atom_iplddt_w).sum(dim=-1) / (
|
| 225 |
atom_iplddt_w.sum(dim=-1) + _EPS
|
| 226 |
+
) # (bs,)
|
| 227 |
+
plddt_ca = plddt_per_atom.gather(1, rep_idx_m) # (bs, l)
|
| 228 |
|
| 229 |
+
pae_logits = self.pae_head(pair) # (bs, l, l, n_pae_bins)
|
| 230 |
+
pae = _categorical_mean(pae_logits, start=0.0, end=32.0).detach() # (bs, l, l)
|
| 231 |
|
| 232 |
n_bins = pae_logits.shape[-1]
|
| 233 |
bin_width = 32.0 / n_bins
|
| 234 |
+
bin_centers = torch.arange(0.5 * bin_width, 32.0, bin_width, device=pae_logits.device) # (n_pae_bins,)
|
| 235 |
+
mask_f = mask.float() # (bs, l)
|
| 236 |
+
n_res = mask_f.sum(dim=-1, keepdim=True) # (bs, 1)
|
| 237 |
+
d0 = 1.24 * (n_res.clamp(min=19) - 15) ** (1 / 3) - 1.8 # (bs, 1)
|
| 238 |
+
tm_per_bin = 1 / (1 + (bin_centers / d0) ** 2) # (bs, n_pae_bins)
|
| 239 |
+
pae_probs = F.softmax(pae_logits, dim=-1) # (bs, l, l, n_pae_bins)
|
| 240 |
+
tm_expected = (pae_probs * tm_per_bin[:, None, None, :]).sum(dim=-1) # (bs, l, l)
|
| 241 |
+
|
| 242 |
+
pair_mask_2d = mask_f.unsqueeze(-1) * mask_f.unsqueeze(-2) # (bs, l, l)
|
| 243 |
+
ptm_per_row = (tm_expected * pair_mask_2d).sum(dim=-1) / (pair_mask_2d.sum(dim=-1) + _EPS) # (bs, l)
|
| 244 |
+
ptm = ptm_per_row.max(dim=-1).values # (bs,)
|
| 245 |
|
| 246 |
inter_chain_mask = (
|
| 247 |
expanded_asym.unsqueeze(-1) != expanded_asym.unsqueeze(-2)
|
| 248 |
+
).float() * pair_mask_2d # (bs, l, l)
|
| 249 |
iptm_per_row = (tm_expected * inter_chain_mask).sum(dim=-1) / (
|
| 250 |
inter_chain_mask.sum(dim=-1) + _EPS
|
| 251 |
+
) # (bs, l)
|
| 252 |
+
iptm = iptm_per_row.max(dim=-1).values # (bs,)
|
| 253 |
|
| 254 |
max_chain_id = int(expanded_asym.max().item()) if batch_mult > 0 else 0
|
| 255 |
n_chains = max_chain_id + 1
|
|
|
|
| 259 |
n_chains,
|
| 260 |
device=tm_expected.device,
|
| 261 |
dtype=tm_expected.dtype,
|
| 262 |
+
) # (bs, n_chains, n_chains)
|
| 263 |
for c1 in range(n_chains):
|
| 264 |
+
chain_c1 = (expanded_asym == c1).float() * mask_f # (bs, l)
|
| 265 |
if chain_c1.sum() == 0:
|
| 266 |
continue
|
| 267 |
for c2 in range(n_chains):
|
| 268 |
+
chain_c2 = (expanded_asym == c2).float() * mask_f # (bs, l)
|
| 269 |
+
pair_m = chain_c1.unsqueeze(-1) * chain_c2.unsqueeze(-2) # (bs, l, l)
|
| 270 |
+
denom = pair_m.sum(dim=(-1, -2)) + _EPS # (bs,)
|
| 271 |
+
pair_chains_iptm[:, c1, c2] = (tm_expected * pair_m).sum(dim=(-1, -2)) / denom # (bs,)
|
| 272 |
|
| 273 |
return {
|
| 274 |
"plddt_logits": plddt_logits,
|
|
|
|
| 282 |
"ptm": ptm.detach(),
|
| 283 |
"iptm": iptm.detach(),
|
| 284 |
"pair_chains_iptm": pair_chains_iptm.detach(),
|
| 285 |
+
} # mapping of confidence tensors with shapes traced above
|
| 286 |
|
| 287 |
|
| 288 |
class _TransitionFFN(nn.Module):
|
|
|
|
| 292 |
self.ffn = SwiGLUMLP(d_model, expansion_ratio=expansion_ratio, bias=False)
|
| 293 |
|
| 294 |
def forward(self, x: Tensor) -> Tensor:
|
| 295 |
+
# x: (..., d_model).
|
| 296 |
+
return self.ffn(self.norm(x)) # x.shape
|
| 297 |
|
| 298 |
|
| 299 |
class MSAEncoderBlock(nn.Module):
|
|
|
|
| 332 |
pair_attention_mask: Tensor,
|
| 333 |
msa_track_mask: Tensor | None = None,
|
| 334 |
) -> tuple[Tensor, Tensor]:
|
| 335 |
+
# msa_repr: (b, l, m, d_msa); pair_repr: (b, l, l, d_pair); msa_track_mask: (b,) or None.
|
| 336 |
mask4d = (
|
| 337 |
msa_track_mask[:, None, None, None].to(dtype=msa_repr.dtype)
|
| 338 |
if msa_track_mask is not None
|
| 339 |
else None
|
| 340 |
+
) # (b, 1, 1, 1) or None
|
| 341 |
|
| 342 |
+
pair_mask4d = mask4d[:, :, :1] if mask4d is not None else None # (b, 1, 1, 1) or None
|
| 343 |
|
| 344 |
+
msa_update = self.msa_pair_weighted_averaging(msa_repr, pair_repr, pair_attention_mask) # (b, l, m, d_msa)
|
| 345 |
if mask4d is not None:
|
| 346 |
+
msa_update = msa_update * mask4d # (b, l, m, d_msa)
|
| 347 |
+
msa_repr = msa_repr + msa_update # (b, l, m, d_msa)
|
| 348 |
|
| 349 |
+
msa_transition = self.msa_transition(msa_repr) # (b, l, m, d_msa)
|
| 350 |
if mask4d is not None:
|
| 351 |
+
msa_transition = msa_transition * mask4d # (b, l, m, d_msa)
|
| 352 |
+
msa_repr = msa_repr + msa_transition # (b, l, m, d_msa)
|
| 353 |
|
| 354 |
+
pair_opm = self.outer_product_mean(msa_repr, msa_attention_mask) # (b, l, l, d_pair)
|
| 355 |
if pair_mask4d is not None:
|
| 356 |
+
pair_opm = pair_opm * pair_mask4d # (b, l, l, d_pair)
|
| 357 |
+
pair_repr = pair_repr + pair_opm # (b, l, l, d_pair)
|
| 358 |
|
| 359 |
+
pair_out = self.tri_mul_out(pair_repr, mask=pair_attention_mask) # (b, l, l, d_pair)
|
| 360 |
if pair_mask4d is not None:
|
| 361 |
+
pair_out = pair_out * pair_mask4d # (b, l, l, d_pair)
|
| 362 |
+
pair_repr = pair_repr + pair_out # (b, l, l, d_pair)
|
| 363 |
|
| 364 |
+
pair_in = self.tri_mul_in(pair_repr, mask=pair_attention_mask) # (b, l, l, d_pair)
|
| 365 |
if pair_mask4d is not None:
|
| 366 |
+
pair_in = pair_in * pair_mask4d # (b, l, l, d_pair)
|
| 367 |
+
pair_repr = pair_repr + pair_in # (b, l, l, d_pair)
|
| 368 |
|
| 369 |
+
pair_transition = self.pair_transition(pair_repr) # (b, l, l, d_pair)
|
| 370 |
if pair_mask4d is not None:
|
| 371 |
+
pair_transition = pair_transition * pair_mask4d # (b, l, l, d_pair)
|
| 372 |
+
pair_repr = pair_repr + pair_transition # (b, l, l, d_pair)
|
| 373 |
+
return msa_repr, pair_repr # (b, l, m, d_msa), (b, l, l, d_pair)
|
| 374 |
|
| 375 |
|
| 376 |
class MSAEncoder(nn.Module):
|
|
|
|
| 413 |
deletion_value: Tensor,
|
| 414 |
msa_attention_mask: Tensor,
|
| 415 |
) -> Tensor:
|
| 416 |
+
# x_pair: (b, l, l, d_pair); x_inputs: (b, l, d_inputs); MSA features: (b, l, m, 33), deletion/mask: (b, l, m).
|
| 417 |
batch_size, _, depth = msa_attention_mask.shape
|
| 418 |
m_feat = torch.cat(
|
| 419 |
[msa_oh, has_deletion.unsqueeze(-1), deletion_value.unsqueeze(-1)],
|
| 420 |
dim=-1,
|
| 421 |
+
) # (b, l, m, 35)
|
| 422 |
+
m = self.embed(m_feat) + self.project_inputs(x_inputs).unsqueeze(2) # (b, l, m, d_msa)
|
| 423 |
if depth > 1:
|
| 424 |
+
msa_track_mask = msa_attention_mask[:, :, 1:].any(dim=(1, 2)) # (b,)
|
| 425 |
else:
|
| 426 |
+
msa_track_mask = torch.zeros(batch_size, dtype=torch.bool, device=x_pair.device) # (b,)
|
| 427 |
+
tok_mask = msa_attention_mask[:, :, 0] # (b, l)
|
| 428 |
+
pair_attention_mask = tok_mask.unsqueeze(2) * tok_mask.unsqueeze(1) # (b, l, l)
|
| 429 |
for block in self.blocks:
|
| 430 |
m, x_pair = cast(MSAEncoderBlock, block)(
|
| 431 |
m,
|
|
|
|
| 433 |
msa_attention_mask,
|
| 434 |
pair_attention_mask,
|
| 435 |
msa_track_mask,
|
| 436 |
+
) # (b, l, m, d_msa), (b, l, l, d_pair)
|
| 437 |
+
return x_pair * msa_track_mask[:, None, None, None].to(dtype=x_pair.dtype) # (b, l, l, d_pair)
|
| 438 |
|
| 439 |
|
| 440 |
class ESMFold2ExperimentalModel(ESMFold2EmbeddingMixin, ESMFold2AttentionMixin, PreTrainedModel):
|
|
|
|
| 484 |
self.pair_loop_proj = nn.Sequential(
|
| 485 |
nn.LayerNorm(d_pair), nn.Linear(d_pair, d_pair, bias=False)
|
| 486 |
)
|
| 487 |
+
nn.init.zeros_(cast(nn.Linear, self.pair_loop_proj[1]).weight) # (d_pair, d_pair)
|
| 488 |
|
| 489 |
self.structure_head = DiffusionStructureHead(config)
|
| 490 |
self.distogram_head = nn.Linear(d_pair, config.structure_head.distogram_bins, bias=True)
|
|
|
|
| 655 |
tok_mask: Tensor,
|
| 656 |
verbose: bool = False,
|
| 657 |
) -> Tensor:
|
| 658 |
+
# Input tensors: (b, l); n_states and d_lm come from the loaded backbone.
|
| 659 |
if self._esmc_fp8 and torch.is_grad_enabled():
|
| 660 |
_reload_esmc_bf16_for_gradients(
|
| 661 |
self,
|
|
|
|
| 678 |
mol_type,
|
| 679 |
tok_mask,
|
| 680 |
pad_to_multiple=pad_to,
|
| 681 |
+
) # (b, l, n_states, d_lm)
|
| 682 |
progress.update()
|
| 683 |
+
return result # (b, l, n_states, d_lm)
|
| 684 |
return compute_lm_hidden_states(
|
| 685 |
self._esmc,
|
| 686 |
input_ids,
|
|
|
|
| 689 |
mol_type,
|
| 690 |
tok_mask,
|
| 691 |
pad_to_multiple=pad_to,
|
| 692 |
+
) # (b, l, n_states, d_lm)
|
| 693 |
|
| 694 |
def forward(
|
| 695 |
self,
|
|
|
|
| 739 |
disto_cond_mask: Tensor | None = None,
|
| 740 |
verbose: bool = False,
|
| 741 |
) -> ESMFold2Output | tuple[Any, ...]:
|
| 742 |
+
# Token IDs/masks: (b, l); atom IDs/masks: (b, a); ref_pos: (b, a, 3); chars: (b, a, 4); MSA: (b, m, l); bs = b * samples.
|
| 743 |
output_hidden_states, return_dict = _resolve_structure_output_controls(
|
| 744 |
self.config,
|
| 745 |
output_attentions=output_attentions,
|
|
|
|
| 760 |
disto_cond_mask=disto_cond_mask,
|
| 761 |
)
|
| 762 |
del gt_coords, is_resolved, frames_idx
|
| 763 |
+
tok_mask = token_attention_mask # (b, l)
|
| 764 |
+
atm_mask = atom_attention_mask # (b, a)
|
| 765 |
n_loops = num_loops if num_loops is not None else self.config.num_loops
|
| 766 |
n_samples = (
|
| 767 |
num_diffusion_samples
|
|
|
|
| 770 |
)
|
| 771 |
|
| 772 |
if res_type.dim() == 2:
|
| 773 |
+
res_type_oh = F.one_hot(res_type.long(), num_classes=NUM_RES_TYPES).float() # (b, l, 33)
|
| 774 |
+
res_type_oh = res_type_oh * tok_mask.unsqueeze(-1).float() # (b, l, 33)
|
| 775 |
else:
|
| 776 |
+
res_type_oh = res_type.float() # (b, l, 33)
|
| 777 |
|
| 778 |
if msa is not None:
|
| 779 |
+
msa_oh_profile = F.one_hot(msa.long(), num_classes=NUM_RES_TYPES).float() # (b, m, l, 33)
|
| 780 |
if msa_attention_mask is not None:
|
| 781 |
+
mask_f = msa_attention_mask.float().unsqueeze(-1) # (b, m, l, 1)
|
| 782 |
+
msa_oh_profile = msa_oh_profile * mask_f # (b, m, l, 33)
|
| 783 |
+
valid_seq_count = msa_attention_mask.float().sum(dim=1).clamp(min=1) # (b, l)
|
| 784 |
+
profile = msa_oh_profile.sum(dim=1) / valid_seq_count.unsqueeze(-1) # (b, l, 33)
|
| 785 |
else:
|
| 786 |
+
profile = msa_oh_profile.mean(dim=1) # (b, l, 33)
|
| 787 |
else:
|
| 788 |
+
profile = res_type_oh # (b, l, 33)
|
| 789 |
|
| 790 |
if res_type_soft is not None:
|
| 791 |
+
res_type_oh = res_type_soft.float() # (b, l, 33)
|
| 792 |
if not self.config.disable_msa_features and provide_soft_sequence_to_msa_and_profile:
|
| 793 |
+
profile = res_type_oh # (b, l, 33)
|
| 794 |
+
msa = res_type_oh.unsqueeze(1) # (b, 1, l, 33)
|
| 795 |
+
msa_attention_mask = tok_mask.unsqueeze(1) # (b, 1, l)
|
| 796 |
|
| 797 |
if deletion_mean is None:
|
| 798 |
deletion_mean = torch.zeros(
|
| 799 |
res_type.shape[0], res_type.shape[1], device=res_type.device
|
| 800 |
+
) # (b, l)
|
| 801 |
if self.config.disable_msa_features:
|
| 802 |
+
profile = torch.zeros_like(profile) # (b, l, 33)
|
| 803 |
+
deletion_mean = torch.zeros_like(deletion_mean) # (b, l)
|
| 804 |
|
| 805 |
+
ref_element_oh = F.one_hot(ref_element.long(), num_classes=MAX_ATOMIC_NUMBER).float() # (b, a, 128)
|
| 806 |
ref_atom_name_chars_oh = F.one_hot(
|
| 807 |
ref_atom_name_chars.long(), num_classes=CHAR_VOCAB_SIZE
|
| 808 |
+
).float() # (b, a, 4, 64)
|
| 809 |
+
atm_mask_f = atm_mask.float() # (b, a)
|
| 810 |
+
ref_element_oh = ref_element_oh * atm_mask_f.unsqueeze(-1) # (b, a, 128)
|
| 811 |
+
ref_atom_name_chars_oh = ref_atom_name_chars_oh * atm_mask_f.unsqueeze(-1).unsqueeze(-1) # (b, a, 4, 64)
|
| 812 |
+
atom_to_token = atom_to_token * atm_mask.long() # (b, a)
|
| 813 |
|
| 814 |
use_amp = ref_pos.device.type == "cuda"
|
| 815 |
with torch.amp.autocast("cuda", enabled=use_amp, dtype=torch.bfloat16):
|
|
|
|
| 824 |
ref_element=ref_element_oh,
|
| 825 |
ref_atom_name_chars=ref_atom_name_chars_oh,
|
| 826 |
atom_to_token=atom_to_token,
|
| 827 |
+
) # (b, l, d_inputs)
|
| 828 |
|
| 829 |
+
z_init = self.z_init_1(x_inputs).unsqueeze(2) + self.z_init_2(x_inputs).unsqueeze(1) # (b, l, l, d_pair)
|
| 830 |
relative_position_encoding = self.rel_pos(
|
| 831 |
residue_index=residue_index,
|
| 832 |
asym_id=asym_id,
|
| 833 |
sym_id=sym_id,
|
| 834 |
entity_id=entity_id,
|
| 835 |
token_index=token_index,
|
| 836 |
+
) # (b, l, l, d_pair)
|
| 837 |
+
token_bonds_encoding = self.token_bonds(token_bonds.float()) # (b, l, l, d_pair)
|
| 838 |
+
z_init = z_init + relative_position_encoding + token_bonds_encoding # (b, l, l, d_pair)
|
| 839 |
|
| 840 |
if lm_hidden_states is None and input_ids is not None and self._esmc is not None:
|
| 841 |
lm_hidden_states = self._compute_lm_hidden_states(
|
| 842 |
input_ids, asym_id, residue_index, mol_type, tok_mask, verbose=verbose
|
| 843 |
+
) # (b, l, n_states, d_lm)
|
| 844 |
if lm_hidden_states is not None:
|
| 845 |
lm_dropout = (
|
| 846 |
self.config.lm_dropout
|
| 847 |
if self.config.force_lm_dropout_during_inference or self.training
|
| 848 |
else 0.0
|
| 849 |
)
|
| 850 |
+
lm_z = self.language_model(lm_hidden_states.detach(), lm_dropout=lm_dropout) # (b, l, l, d_pair) or None
|
| 851 |
+
z_init = z_init + lm_z.to(z_init.dtype) # (b, l, l, d_pair)
|
| 852 |
|
| 853 |
msa_kwargs: dict[str, Tensor] | None = None
|
| 854 |
if self.msa_encoder is not None and msa is not None:
|
| 855 |
if msa.dim() == 4:
|
| 856 |
batch_msa, depth, length_msa, _ = msa.shape
|
| 857 |
+
msa_oh = msa.permute(0, 2, 1, 3).float() # (b, l, m, 33)
|
| 858 |
else:
|
| 859 |
batch_msa, depth, length_msa = msa.shape
|
| 860 |
msa_oh = F.one_hot(
|
| 861 |
msa.permute(0, 2, 1).long(), num_classes=NUM_RES_TYPES
|
| 862 |
+
).float() # (b, l, m, 33)
|
| 863 |
msa_attn = (
|
| 864 |
msa_attention_mask.permute(0, 2, 1).float()
|
| 865 |
if msa_attention_mask is not None
|
| 866 |
else tok_mask[:, :, None].expand(-1, -1, depth).float()
|
| 867 |
+
) # (b, l, m)
|
| 868 |
+
msa_oh = msa_oh * msa_attn.unsqueeze(-1) # (b, l, m, 33)
|
| 869 |
hd = (
|
| 870 |
has_deletion.permute(0, 2, 1).float()
|
| 871 |
if has_deletion is not None
|
| 872 |
else torch.zeros(batch_msa, length_msa, depth, device=msa.device)
|
| 873 |
+
) # (b, l, m)
|
| 874 |
dv = (
|
| 875 |
deletion_value.permute(0, 2, 1).float()
|
| 876 |
if deletion_value is not None
|
| 877 |
else torch.zeros(batch_msa, length_msa, depth, device=msa.device)
|
| 878 |
+
) # (b, l, m)
|
| 879 |
msa_kwargs = {
|
| 880 |
"x_inputs": x_inputs,
|
| 881 |
"msa_oh": msa_oh,
|
|
|
|
| 884 |
"msa_attention_mask": msa_attn,
|
| 885 |
}
|
| 886 |
|
| 887 |
+
pair_mask = tok_mask[:, :, None].float() * tok_mask[:, None, :].float() # (b, l, l)
|
| 888 |
+
z = torch.zeros_like(z_init) # (b, l, l, d_pair)
|
| 889 |
+
prev_pair: Tensor | None = None # (b, l, l, d_pair) or None
|
| 890 |
+
prev_disto_probs: Tensor | None = None # (b, l, l, n_distogram_bins) or None
|
| 891 |
loop_iterator = range(n_loops + 1)
|
| 892 |
if verbose:
|
| 893 |
loop_iterator = tqdm(
|
|
|
|
| 898 |
)
|
| 899 |
|
| 900 |
for loop_num in loop_iterator:
|
| 901 |
+
z = z_init + self.pair_loop_proj(z) # (b, l, l, d_pair)
|
| 902 |
if msa_kwargs is not None and self.msa_encoder is not None:
|
| 903 |
+
z = z + self.msa_encoder(x_pair=z, **msa_kwargs).to(z.dtype) # (b, l, l, d_pair)
|
| 904 |
+
z = self.folding_trunk(z, pair_attention_mask=pair_mask) # (b, l, l, d_pair)
|
| 905 |
|
| 906 |
if early_exit and loop_num < n_loops:
|
| 907 |
l2_converged = False
|
| 908 |
if prev_pair is not None and loop_num > 0:
|
| 909 |
rel_l2 = (
|
| 910 |
z.float() - prev_pair.float()
|
| 911 |
+
).norm() / prev_pair.float().norm().clamp(min=1e-8) # ()
|
| 912 |
l2_converged = rel_l2.item() < 0.25
|
| 913 |
+
prev_pair = z.detach().clone() # (b, l, l, d_pair) or None
|
| 914 |
+
sym_z = z.float() + z.float().transpose(-2, -3) # (b, l, l, d_pair)
|
| 915 |
+
cur_probs = F.softmax(self.distogram_head(sym_z).float(), dim=-1) # (b, l, l, n_distogram_bins)
|
| 916 |
if prev_disto_probs is not None and loop_num > 0:
|
| 917 |
kl_per_pair = (
|
| 918 |
cur_probs
|
| 919 |
* (cur_probs.clamp(min=1e-8) / prev_disto_probs.clamp(min=1e-8)).log()
|
| 920 |
+
).sum(-1) # (b, l, l)
|
| 921 |
+
kl = (kl_per_pair + kl_per_pair.transpose(-1, -2)).mean() / 2 # ()
|
| 922 |
if l2_converged or kl.item() < 0.05:
|
| 923 |
break
|
| 924 |
+
prev_disto_probs = cur_probs.detach() # (b, l, l, n_distogram_bins) or None
|
| 925 |
|
| 926 |
+
distogram_logits = self.distogram_head(z + z.transpose(-2, -3)) # (b, l, l, n_distogram_bins)
|
| 927 |
|
| 928 |
with torch.no_grad(), _seed_context(seed):
|
| 929 |
structure_output = self.structure_head.sample(
|
|
|
|
| 952 |
return_atom_repr=False,
|
| 953 |
denoising_early_exit_rmsd=(0.10 if early_exit else None),
|
| 954 |
verbose=verbose,
|
| 955 |
+
) # tensor mapping follows the called head's shape contract
|
| 956 |
+
sample_coords = structure_output["sample_atom_coords"] # (bs, a, 3), or explicit (b, samples, a, 3)
|
| 957 |
if sample_coords is None:
|
| 958 |
raise RuntimeError("ESMFold2 structure sampling did not return coordinates.")
|
| 959 |
if sample_coords.ndim == 4:
|
|
|
|
| 962 |
batch * sample_count,
|
| 963 |
atom_count,
|
| 964 |
coord_dim,
|
| 965 |
+
) # (bs, a, 3)
|
| 966 |
+
rep_idx = distogram_atom_idx.repeat_interleave(sample_count, 0).long() # (bs, l) for explicit sample axis, otherwise distogram_atom_idx.shape
|
| 967 |
else:
|
| 968 |
+
sample_coords_for_gather = sample_coords # (bs, a, 3)
|
| 969 |
+
rep_idx = distogram_atom_idx.long() # (bs, l) for explicit sample axis, otherwise distogram_atom_idx.shape
|
| 970 |
representative_atom_coords = gather_rep_atom_coords(
|
| 971 |
sample_coords_for_gather,
|
| 972 |
rep_idx,
|
| 973 |
+
) # (*rep_idx.shape, 3); gather retains the index batch size
|
| 974 |
|
| 975 |
output: dict[str, Tensor] = {
|
| 976 |
"distogram_logits": distogram_logits,
|
|
|
|
| 999 |
num_diffusion_samples=n_samples,
|
| 1000 |
relative_position_encoding=relative_position_encoding.detach(),
|
| 1001 |
token_bonds_encoding=token_bonds_encoding.detach(),
|
| 1002 |
+
) # tensor mapping follows the called head's shape contract
|
| 1003 |
progress.update()
|
| 1004 |
else:
|
| 1005 |
confidence_output = self.confidence_head(
|
|
|
|
| 1015 |
num_diffusion_samples=n_samples,
|
| 1016 |
relative_position_encoding=relative_position_encoding.detach(),
|
| 1017 |
token_bonds_encoding=token_bonds_encoding.detach(),
|
| 1018 |
+
) # tensor mapping follows the called head's shape contract
|
| 1019 |
output.update(confidence_output)
|
| 1020 |
+
output["atom_pad_mask"] = atm_mask.unsqueeze(0) if atm_mask.dim() == 1 else atm_mask # (b, a)
|
| 1021 |
+
output["residue_index"] = residue_index # (b, l)
|
| 1022 |
+
output["entity_id"] = entity_id # (b, l)
|
| 1023 |
return _finalize_structure_output(
|
| 1024 |
output,
|
| 1025 |
token_input_state=x_inputs,
|
| 1026 |
pair_state=z,
|
| 1027 |
output_hidden_states=output_hidden_states,
|
| 1028 |
return_dict=return_dict,
|
| 1029 |
+
) # ESMFold2Output/tuple retaining the traced tensor shapes
|
| 1030 |
|
| 1031 |
@property
|
| 1032 |
def input_builder(self):
|
|
|
|
| 1067 |
if not self.config.msa_conditioning:
|
| 1068 |
for name in MSA_CONDITIONING_INPUT_NAMES:
|
| 1069 |
features.pop(name, None)
|
| 1070 |
+
features = {name: tensor.to(self.device) for name, tensor in features.items()} # every feature retains its shape
|
| 1071 |
output = self(**features, **forward_kwargs, return_dict=True)
|
| 1072 |
for name in (
|
| 1073 |
"res_type",
|
fastplms/models/esmfold2/protein_utils.py
CHANGED
|
@@ -9,11 +9,11 @@ lazily from a provenance-bearing declarative package asset.
|
|
| 9 |
from __future__ import annotations
|
| 10 |
|
| 11 |
import json
|
|
|
|
|
|
|
| 12 |
from functools import cache
|
| 13 |
from importlib.resources import files
|
| 14 |
from typing import Any
|
| 15 |
-
|
| 16 |
-
import torch
|
| 17 |
from torch import Tensor
|
| 18 |
|
| 19 |
from .esmfold2_constants import (
|
|
@@ -27,6 +27,7 @@ from .esmfold2_constants import (
|
|
| 27 |
PROTEIN_UNK_RES_TYPE,
|
| 28 |
)
|
| 29 |
|
|
|
|
| 30 |
_GEOMETRY_ASSET = "protein_reference_geometry.json"
|
| 31 |
_GEOMETRY_SCHEMA = "fastplms.esmfold2.reference_geometry.v1"
|
| 32 |
|
|
@@ -122,57 +123,57 @@ def prepare_protein_features(sequence: str) -> dict[str, Tensor]:
|
|
| 122 |
raise ValueError("sequence must be non-empty")
|
| 123 |
|
| 124 |
atoms, residue_types, input_ids, representative_atoms = _residue_records(sequence)
|
| 125 |
-
sequence_length = len(sequence)
|
| 126 |
n_atoms = _padded_atom_count(len(atoms))
|
| 127 |
|
| 128 |
-
ref_pos = torch.zeros((n_atoms, 3), dtype=torch.float32)
|
| 129 |
-
ref_element = torch.zeros(n_atoms, dtype=torch.int64)
|
| 130 |
-
ref_charge = torch.zeros(n_atoms, dtype=torch.int8)
|
| 131 |
-
ref_atom_name_chars = torch.zeros((n_atoms, 4), dtype=torch.int64)
|
| 132 |
-
ref_space_uid = torch.zeros(n_atoms, dtype=torch.int64)
|
| 133 |
-
atom_attention_mask = torch.zeros(n_atoms, dtype=torch.bool)
|
| 134 |
-
atom_to_token = torch.zeros(n_atoms, dtype=torch.int64)
|
| 135 |
|
| 136 |
for atom_index, atom in enumerate(atoms):
|
| 137 |
token_index = atom["token_index"]
|
| 138 |
-
ref_pos[atom_index] = torch.tensor(atom["position"], dtype=torch.float32)
|
| 139 |
-
ref_element[atom_index] = ELEMENT_TO_ATOMIC_NUM[atom["element"]]
|
| 140 |
-
ref_charge[atom_index] = atom["charge"]
|
| 141 |
ref_atom_name_chars[atom_index] = torch.tensor(
|
| 142 |
_encode_atom_name(atom["name"]), dtype=torch.int64
|
| 143 |
-
)
|
| 144 |
-
ref_space_uid[atom_index] = token_index
|
| 145 |
-
atom_attention_mask[atom_index] = True
|
| 146 |
-
atom_to_token[atom_index] = token_index
|
| 147 |
|
| 148 |
-
residue_type_tensor = torch.tensor(residue_types, dtype=torch.int64)
|
| 149 |
-
msa = residue_type_tensor.unsqueeze(0)
|
| 150 |
features = {
|
| 151 |
-
"token_index": torch.arange(sequence_length, dtype=torch.int64),
|
| 152 |
-
"residue_index": torch.arange(sequence_length, dtype=torch.int64),
|
| 153 |
-
"asym_id": torch.zeros(sequence_length, dtype=torch.int64),
|
| 154 |
-
"sym_id": torch.zeros(sequence_length, dtype=torch.int64),
|
| 155 |
-
"entity_id": torch.ones(sequence_length, dtype=torch.int64),
|
| 156 |
-
"mol_type": torch.full((sequence_length,), MOL_TYPE_PROTEIN, dtype=torch.int64),
|
| 157 |
-
"res_type": residue_type_tensor,
|
| 158 |
-
"input_ids": torch.tensor(input_ids, dtype=torch.int64),
|
| 159 |
-
"token_bonds": torch.zeros((sequence_length, sequence_length, 1), dtype=torch.float32),
|
| 160 |
-
"token_attention_mask": torch.ones(sequence_length, dtype=torch.bool),
|
| 161 |
-
"ref_pos": ref_pos,
|
| 162 |
-
"ref_element": ref_element,
|
| 163 |
-
"ref_charge": ref_charge,
|
| 164 |
-
"ref_atom_name_chars": ref_atom_name_chars,
|
| 165 |
-
"ref_space_uid": ref_space_uid,
|
| 166 |
-
"atom_attention_mask": atom_attention_mask,
|
| 167 |
-
"atom_to_token": atom_to_token,
|
| 168 |
-
"distogram_atom_idx": torch.tensor(representative_atoms, dtype=torch.int64),
|
| 169 |
-
"msa": msa,
|
| 170 |
-
"msa_attention_mask": torch.ones_like(msa, dtype=torch.bool),
|
| 171 |
-
"has_deletion": torch.zeros_like(msa, dtype=torch.bool),
|
| 172 |
-
"deletion_value": torch.zeros_like(msa, dtype=torch.float32),
|
| 173 |
-
"deletion_mean": torch.zeros(sequence_length, dtype=torch.float32),
|
| 174 |
}
|
| 175 |
-
return {name: tensor.unsqueeze(0) for name, tensor in features.items()}
|
| 176 |
|
| 177 |
|
| 178 |
__all__ = ["prepare_protein_features"]
|
|
|
|
| 9 |
from __future__ import annotations
|
| 10 |
|
| 11 |
import json
|
| 12 |
+
import torch
|
| 13 |
+
|
| 14 |
from functools import cache
|
| 15 |
from importlib.resources import files
|
| 16 |
from typing import Any
|
|
|
|
|
|
|
| 17 |
from torch import Tensor
|
| 18 |
|
| 19 |
from .esmfold2_constants import (
|
|
|
|
| 27 |
PROTEIN_UNK_RES_TYPE,
|
| 28 |
)
|
| 29 |
|
| 30 |
+
|
| 31 |
_GEOMETRY_ASSET = "protein_reference_geometry.json"
|
| 32 |
_GEOMETRY_SCHEMA = "fastplms.esmfold2.reference_geometry.v1"
|
| 33 |
|
|
|
|
| 123 |
raise ValueError("sequence must be non-empty")
|
| 124 |
|
| 125 |
atoms, residue_types, input_ids, representative_atoms = _residue_records(sequence)
|
| 126 |
+
sequence_length = len(sequence) # l
|
| 127 |
n_atoms = _padded_atom_count(len(atoms))
|
| 128 |
|
| 129 |
+
ref_pos = torch.zeros((n_atoms, 3), dtype=torch.float32) # (n_atoms, 3)
|
| 130 |
+
ref_element = torch.zeros(n_atoms, dtype=torch.int64) # (n_atoms,)
|
| 131 |
+
ref_charge = torch.zeros(n_atoms, dtype=torch.int8) # (n_atoms,)
|
| 132 |
+
ref_atom_name_chars = torch.zeros((n_atoms, 4), dtype=torch.int64) # (n_atoms, 4)
|
| 133 |
+
ref_space_uid = torch.zeros(n_atoms, dtype=torch.int64) # (n_atoms,)
|
| 134 |
+
atom_attention_mask = torch.zeros(n_atoms, dtype=torch.bool) # (n_atoms,)
|
| 135 |
+
atom_to_token = torch.zeros(n_atoms, dtype=torch.int64) # (n_atoms,)
|
| 136 |
|
| 137 |
for atom_index, atom in enumerate(atoms):
|
| 138 |
token_index = atom["token_index"]
|
| 139 |
+
ref_pos[atom_index] = torch.tensor(atom["position"], dtype=torch.float32) # (3,)
|
| 140 |
+
ref_element[atom_index] = ELEMENT_TO_ATOMIC_NUM[atom["element"]] # ()
|
| 141 |
+
ref_charge[atom_index] = atom["charge"] # ()
|
| 142 |
ref_atom_name_chars[atom_index] = torch.tensor(
|
| 143 |
_encode_atom_name(atom["name"]), dtype=torch.int64
|
| 144 |
+
) # (4,)
|
| 145 |
+
ref_space_uid[atom_index] = token_index # ()
|
| 146 |
+
atom_attention_mask[atom_index] = True # ()
|
| 147 |
+
atom_to_token[atom_index] = token_index # ()
|
| 148 |
|
| 149 |
+
residue_type_tensor = torch.tensor(residue_types, dtype=torch.int64) # (l,)
|
| 150 |
+
msa = residue_type_tensor.unsqueeze(0) # (1, l)
|
| 151 |
features = {
|
| 152 |
+
"token_index": torch.arange(sequence_length, dtype=torch.int64), # (l,)
|
| 153 |
+
"residue_index": torch.arange(sequence_length, dtype=torch.int64), # (l,)
|
| 154 |
+
"asym_id": torch.zeros(sequence_length, dtype=torch.int64), # (l,)
|
| 155 |
+
"sym_id": torch.zeros(sequence_length, dtype=torch.int64), # (l,)
|
| 156 |
+
"entity_id": torch.ones(sequence_length, dtype=torch.int64), # (l,)
|
| 157 |
+
"mol_type": torch.full((sequence_length,), MOL_TYPE_PROTEIN, dtype=torch.int64), # (l,)
|
| 158 |
+
"res_type": residue_type_tensor, # (l,)
|
| 159 |
+
"input_ids": torch.tensor(input_ids, dtype=torch.int64), # (l,)
|
| 160 |
+
"token_bonds": torch.zeros((sequence_length, sequence_length, 1), dtype=torch.float32), # (l, l, 1)
|
| 161 |
+
"token_attention_mask": torch.ones(sequence_length, dtype=torch.bool), # (l,)
|
| 162 |
+
"ref_pos": ref_pos, # (n_atoms, 3)
|
| 163 |
+
"ref_element": ref_element, # (n_atoms,)
|
| 164 |
+
"ref_charge": ref_charge, # (n_atoms,)
|
| 165 |
+
"ref_atom_name_chars": ref_atom_name_chars, # (n_atoms, 4)
|
| 166 |
+
"ref_space_uid": ref_space_uid, # (n_atoms,)
|
| 167 |
+
"atom_attention_mask": atom_attention_mask, # (n_atoms,)
|
| 168 |
+
"atom_to_token": atom_to_token, # (n_atoms,)
|
| 169 |
+
"distogram_atom_idx": torch.tensor(representative_atoms, dtype=torch.int64), # (l,)
|
| 170 |
+
"msa": msa, # (1, l)
|
| 171 |
+
"msa_attention_mask": torch.ones_like(msa, dtype=torch.bool), # (1, l)
|
| 172 |
+
"has_deletion": torch.zeros_like(msa, dtype=torch.bool), # (1, l)
|
| 173 |
+
"deletion_value": torch.zeros_like(msa, dtype=torch.float32), # (1, l)
|
| 174 |
+
"deletion_mean": torch.zeros(sequence_length, dtype=torch.float32), # (l,)
|
| 175 |
}
|
| 176 |
+
return {name: tensor.unsqueeze(0) for name, tensor in features.items()} # each shape -> (1, *shape)
|
| 177 |
|
| 178 |
|
| 179 |
__all__ = ["prepare_protein_features"]
|
fastplms/models/esmfold2/reproducibility.py
CHANGED
|
@@ -3,13 +3,13 @@
|
|
| 3 |
from __future__ import annotations
|
| 4 |
|
| 5 |
import random
|
|
|
|
|
|
|
|
|
|
| 6 |
from collections.abc import Iterator
|
| 7 |
from contextlib import contextmanager
|
| 8 |
from dataclasses import dataclass
|
| 9 |
from typing import Any
|
| 10 |
-
|
| 11 |
-
import numpy as np
|
| 12 |
-
import torch
|
| 13 |
from torch import Tensor
|
| 14 |
|
| 15 |
|
|
|
|
| 3 |
from __future__ import annotations
|
| 4 |
|
| 5 |
import random
|
| 6 |
+
import numpy as np
|
| 7 |
+
import torch
|
| 8 |
+
|
| 9 |
from collections.abc import Iterator
|
| 10 |
from contextlib import contextmanager
|
| 11 |
from dataclasses import dataclass
|
| 12 |
from typing import Any
|
|
|
|
|
|
|
|
|
|
| 13 |
from torch import Tensor
|
| 14 |
|
| 15 |
|
fastplms/models/ttt.py
CHANGED
|
@@ -6,6 +6,7 @@ import numbers
|
|
| 6 |
import torch
|
| 7 |
import torch.nn as nn
|
| 8 |
import torch.nn.functional as F
|
|
|
|
| 9 |
from collections.abc import Iterator, Mapping
|
| 10 |
from dataclasses import asdict, dataclass, fields
|
| 11 |
from typing import Any
|
|
|
|
| 6 |
import torch
|
| 7 |
import torch.nn as nn
|
| 8 |
import torch.nn.functional as F
|
| 9 |
+
|
| 10 |
from collections.abc import Iterator, Mapping
|
| 11 |
from dataclasses import asdict, dataclass, fields
|
| 12 |
from typing import Any
|