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

Apply coding standards from 1cb5747 (files only)

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. README.md +10 -6
  2. fastplms/attention/_core.py +1 -0
  3. fastplms/attention/_kernel_lock.py +1 -0
  4. fastplms/attention/interfaces.py +1 -0
  5. fastplms/embeddings/batches.py +379 -378
  6. fastplms/embeddings/identity.py +508 -507
  7. fastplms/embeddings/inputs.py +264 -263
  8. fastplms/embeddings/output.py +215 -215
  9. fastplms/embeddings/pooling.py +1 -0
  10. fastplms/embeddings/runner.py +420 -419
  11. fastplms/embeddings/storage.py +1 -0
  12. fastplms/models/_esm_rotary.py +1 -0
  13. fastplms/models/classification_probe.py +21 -20
  14. fastplms/models/esm_plusplus/modeling_esm_plusplus.py +25 -25
  15. fastplms/models/esmfold2/__init__.py +1 -0
  16. fastplms/models/esmfold2/configuration_esmfold2.py +0 -1
  17. fastplms/models/esmfold2/embedding.py +1 -0
  18. fastplms/models/esmfold2/esmfold2_affine3d.py +134 -112
  19. fastplms/models/esmfold2/esmfold2_aligner.py +19 -18
  20. fastplms/models/esmfold2/esmfold2_atom_indexer.py +3 -3
  21. fastplms/models/esmfold2/esmfold2_conformers.py +3 -3
  22. fastplms/models/esmfold2/esmfold2_input_builder.py +3 -2
  23. fastplms/models/esmfold2/esmfold2_metrics.py +69 -54
  24. fastplms/models/esmfold2/esmfold2_misc.py +73 -58
  25. fastplms/models/esmfold2/esmfold2_mmcif_parsing.py +5 -4
  26. fastplms/models/esmfold2/esmfold2_molecular_complex.py +6 -6
  27. fastplms/models/esmfold2/esmfold2_msa.py +3 -2
  28. fastplms/models/esmfold2/esmfold2_msa_filter_sequences.py +15 -13
  29. fastplms/models/esmfold2/esmfold2_normalize_coordinates.py +24 -21
  30. fastplms/models/esmfold2/esmfold2_output.py +11 -10
  31. fastplms/models/esmfold2/esmfold2_paired_msa.py +3 -2
  32. fastplms/models/esmfold2/esmfold2_parsing.py +4 -3
  33. fastplms/models/esmfold2/esmfold2_predicted_aligned_error.py +45 -39
  34. fastplms/models/esmfold2/esmfold2_prepare_input.py +4 -3
  35. fastplms/models/esmfold2/esmfold2_processor.py +4 -3
  36. fastplms/models/esmfold2/esmfold2_protein_chain.py +200 -187
  37. fastplms/models/esmfold2/esmfold2_protein_complex.py +104 -97
  38. fastplms/models/esmfold2/esmfold2_protein_structure.py +65 -58
  39. fastplms/models/esmfold2/esmfold2_residue_constants.py +2 -1
  40. fastplms/models/esmfold2/esmfold2_sequential_dataclass.py +3 -2
  41. fastplms/models/esmfold2/esmfold2_system.py +2 -0
  42. fastplms/models/esmfold2/esmfold2_types.py +1 -0
  43. fastplms/models/esmfold2/esmfold2_utils_types.py +2 -0
  44. fastplms/models/esmfold2/modeling_esmfold2.py +219 -202
  45. fastplms/models/esmfold2/modeling_esmfold2_classification.py +21 -20
  46. fastplms/models/esmfold2/modeling_esmfold2_common.py +494 -450
  47. fastplms/models/esmfold2/modeling_esmfold2_experimental.py +194 -185
  48. fastplms/models/esmfold2/protein_utils.py +44 -43
  49. fastplms/models/esmfold2/reproducibility.py +3 -3
  50. 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
- from collections.abc import Callable, Iterator, Sequence
7
- from contextlib import contextmanager
8
- from dataclasses import dataclass, field
9
- from typing import Any
10
- from torch import Tensor
11
-
12
- from .identity import _model_device
13
- from .inputs import _planned_batches
14
- from .pooling import Pooler
15
- from .types import EmbeddingBatch, EmbeddingInput, EmbeddingRecord
16
-
17
-
18
- _MAX_PARTI_RESIDUES = 2_048
19
-
20
-
21
- def _validate_parti_length(M: Tensor) -> None:
22
- """Reject an oversized attention graph before model inference."""
23
-
24
- # M: (b, l)
25
- n_residues = int(M.to(dtype=torch.int64).sum(dim=1).max().item())
26
- if n_residues > _MAX_PARTI_RESIDUES:
27
- raise ValueError(f"parti supports at most {_MAX_PARTI_RESIDUES:,} biological residues.")
28
-
29
-
30
- def select_hidden_state_embeddings(
31
- last_hidden_state: Tensor,
32
- hidden_states: tuple[Tensor, ...] | None,
33
- *,
34
- hidden_state_index: int = -1,
35
- store_all_hidden_states: bool = False,
36
- ) -> Tensor:
37
- """Select one hidden state or stack every state without changing values."""
38
- # last_hidden_state and each hidden_states entry: (b, l, d)
39
- if store_all_hidden_states:
40
- if not hidden_states:
41
- raise ValueError("store_all_hidden_states requires model hidden states.")
42
- # H has shape (b, n, l, d), where n follows the model's output order.
43
- return torch.stack(hidden_states, dim=1) # (b, n, l, d)
44
- if hidden_state_index == -1:
45
- return last_hidden_state # (b, l, d)
46
- if not hidden_states:
47
- raise ValueError("hidden_state_index requires model hidden states.")
48
- return hidden_states[hidden_state_index] # (b, l, d)
49
-
50
-
51
- def _residue_embeddings(X: Tensor, M: Tensor) -> list[Tensor]:
52
- """Copy every sample's biological residues to the host in one transfer.
53
-
54
- Boolean indexing packs the selected rows in batch order, so splitting the
55
- packed rows by residue count gives the values that indexing each sample
56
- would. Each returned tensor owns its storage, as a per-sample copy does.
57
- """
58
- # X: (b, l, d); M: (b, l)
59
- residue_counts = M.sum(dim=1).tolist() # b counts r_i
60
- packed = X[M].detach().cpu() # (sum of r_i, d)
61
- return [sample.clone() for sample in torch.split(packed, residue_counts)] # each: (r_i, d)
62
-
63
-
64
- @contextmanager
65
- def _temporary_eval(model: Any) -> Iterator[None]:
66
- was_training = getattr(model, "training", None)
67
- eval_method = getattr(model, "eval", None)
68
- train_method = getattr(model, "train", None)
69
- if (
70
- not isinstance(was_training, bool)
71
- or not callable(eval_method)
72
- or not callable(train_method)
73
- ):
74
- yield
75
- return
76
- eval_method()
77
- try:
78
- yield
79
- finally:
80
- train_method(was_training)
81
-
82
-
83
- def _biological_residue_mask(
84
- input_ids: Tensor,
85
- attention_mask: Tensor,
86
- tokenizer: Any,
87
- ) -> Tensor:
88
- """Remove padding and tokenizer-declared special tokens from M."""
89
-
90
- # input_ids, attention_mask: (b, l)
91
- M = attention_mask.to(dtype=torch.bool) # (b, l)
92
- special_ids = tuple(int(token_id) for token_id in getattr(tokenizer, "all_special_ids", ()))
93
- if special_ids:
94
- specials = torch.tensor( # (n_special,)
95
- special_ids,
96
- device=input_ids.device,
97
- dtype=input_ids.dtype,
98
- )
99
- M = M & ~torch.isin(input_ids, specials) # (b, l)
100
- return M # (b, l)
101
-
102
-
103
- def _generic_embedding_batch(
104
- model: Any,
105
- sequences: list[str],
106
- *,
107
- tokenizer: Any | None,
108
- max_length: int | None,
109
- truncate: bool,
110
- need_attentions: bool,
111
- model_kwargs: dict[str, Any],
112
- ) -> EmbeddingBatch:
113
- config = getattr(model, "config", None)
114
- model_type = str(getattr(config, "model_type", "")).lower()
115
- if tokenizer is None:
116
- tokenizer = getattr(model, "tokenizer", None)
117
-
118
- if tokenizer is None and model_type == "e1":
119
- output = model._embed(sequences, return_attention_mask=True, **model_kwargs)
120
- if not isinstance(output, tuple) or len(output) != 2:
121
- raise TypeError("E1 _embed must return (X, residue_mask).")
122
- X, M = output # (b, l, d), (b, l)
123
- preparer = getattr(model, "prep_tokens", None)
124
- if preparer is not None and hasattr(preparer, "get_batch_kwargs"):
125
- prepared = preparer.get_batch_kwargs(sequences, device=X.device)
126
- input_ids = prepared["input_ids"] # (b, l)
127
- boundary_ids = preparer.boundary_token_ids.to( # (n_boundary,)
128
- device=input_ids.device, dtype=input_ids.dtype
129
- )
130
- # E1 wraps each raw sequence in BOS, context-label, terminal-label,
131
- # and EOS tokens. Only amino-acid rows are biological residues.
132
- M = M.to(dtype=torch.bool) & ~torch.isin(input_ids, boundary_ids) # (b, l)
133
- if need_attentions:
134
- raise ValueError("parti is not available for tokenizer-free E1 embedding.")
135
- return EmbeddingBatch( # X: (b, l, d); residue_mask: (b, l)
136
- X=X,
137
- residue_mask=M.to(dtype=torch.bool),
138
- )
139
- if tokenizer is None:
140
- raise ValueError("A tokenizer is required for this model's embedding path.")
141
-
142
- tokenize_kwargs: dict[str, Any] = {
143
- "return_tensors": "pt",
144
- "padding": True,
145
- "truncation": truncate,
146
- }
147
- if max_length is not None and truncate:
148
- # ``max_length`` is a biological-residue limit. Tokenizer limits include
149
- # boundary tokens, so reserve their declared width instead of dropping
150
- # residues at the exact boundary.
151
- special_token_count = 0
152
- num_special_tokens_to_add = getattr(tokenizer, "num_special_tokens_to_add", None)
153
- if callable(num_special_tokens_to_add):
154
- special_token_count = int(num_special_tokens_to_add(pair=False))
155
- tokenize_kwargs["max_length"] = max_length + special_token_count
156
- sequence_tokenizer = getattr(model, "_tokenize_sequence_batch", None)
157
- if callable(sequence_tokenizer):
158
- encoded = sequence_tokenizer(sequences, tokenizer=tokenizer, **tokenize_kwargs)
159
- else:
160
- encoded = tokenizer(sequences, **tokenize_kwargs)
161
- device = _model_device(model)
162
- input_ids = encoded["input_ids"].to(device) # (b, l)
163
- attention_mask = encoded.get( # (b, l)
164
- "attention_mask",
165
- input_ids.new_ones(input_ids.shape),
166
- ).to(device)
167
- M = _biological_residue_mask(input_ids, attention_mask, tokenizer) # (b, l)
168
- if need_attentions:
169
- # Validate l before either the backbone or its quadratic attention graph
170
- # is materialized. M has shape (b, l).
171
- _validate_parti_length(M)
172
- X = model._embed(input_ids, attention_mask, **model_kwargs) # (b, l, d)
173
- attentions = None
174
- if need_attentions:
175
- output = model(
176
- input_ids=input_ids,
177
- attention_mask=attention_mask,
178
- output_attentions=True,
179
- return_dict=True,
180
- )
181
- attentions = getattr(output, "attentions", None) # each: (b, h, l, l)
182
- if attentions is None:
183
- raise ValueError("The model did not return attentions required by parti.")
184
- return EmbeddingBatch( # X: (b, l, d); M: (b, l)
185
- X=X,
186
- residue_mask=M,
187
- attentions=attentions,
188
- )
189
-
190
-
191
- @dataclass(eq=False)
192
- class BatchExecutor:
193
- """Model and batch policy for one bounded embedding window at a time."""
194
-
195
- model: Any
196
- batch_size: int
197
- max_tokens_per_batch: int | None
198
- max_length: int | None
199
- truncate: bool
200
- model_kwargs: dict[str, Any]
201
- hidden_state_source: str
202
- normalized_decoder_inputs: tuple[str, ...] | None
203
- decoder_input_ids: Tensor | None
204
- decoder_attention_mask: Tensor | None
205
- _embedding_batch_fn: Callable[..., EmbeddingBatch] | None
206
- tokenizer: Any | None
207
- store_all_hidden_states: bool
208
- full_embeddings: bool
209
- dtype: torch.dtype | None
210
- pooler: Pooler | None
211
- attention_backend: str | None
212
- need_attentions: bool
213
- model_type: str = field(init=False)
214
- resolved_tokenizer: Any = field(init=False)
215
-
216
- def __post_init__(self) -> None:
217
- config = getattr(self.model, "config", None)
218
- self.model_type = str(getattr(config, "model_type", "")).lower()
219
- self.resolved_tokenizer = (
220
- self.tokenizer if self.tokenizer is not None else getattr(self.model, "tokenizer", None)
221
- )
222
-
223
- def run_window(
224
- self,
225
- window_records: Sequence[EmbeddingInput],
226
- *,
227
- window_start: int,
228
- ) -> tuple[list[EmbeddingRecord], dict[str, tuple[int, int]]]:
229
- """Restore source order after length-bucketed inference and pooling."""
230
-
231
- pool_slices: dict[str, tuple[int, int]] = {}
232
- window_results: dict[int, EmbeddingRecord] = {}
233
- for local_positions in _planned_batches(
234
- window_records,
235
- range(len(window_records)),
236
- batch_size=self.batch_size,
237
- max_tokens_per_batch=self.max_tokens_per_batch,
238
- max_length=self.max_length,
239
- truncate=self.truncate,
240
- ):
241
- batch_positions = [window_start + position for position in local_positions]
242
- batch_records = [window_records[position] for position in local_positions]
243
- sequences = [
244
- record.sequence[: self.max_length]
245
- if self.truncate and self.max_length is not None
246
- else record.sequence
247
- for record in batch_records
248
- ]
249
- batch_model_kwargs = dict(self.model_kwargs)
250
- if self.model_type == "fast_ankh" or self.hidden_state_source == "decoder":
251
- batch_model_kwargs["hidden_state_source"] = self.hidden_state_source
252
- if self.normalized_decoder_inputs is not None:
253
- batch_model_kwargs["decoder_inputs"] = [
254
- self.normalized_decoder_inputs[position] for position in batch_positions
255
- ]
256
- if self.decoder_input_ids is not None:
257
- # decoder_input_ids: (n_records, l_decoder)
258
- indices = torch.tensor( # (b,)
259
- batch_positions,
260
- device=self.decoder_input_ids.device,
261
- dtype=torch.long,
262
- )
263
- batch_model_kwargs["decoder_input_ids"] = ( # (b, l_decoder)
264
- self.decoder_input_ids.index_select(0, indices)
265
- )
266
- if self.decoder_attention_mask is not None:
267
- # decoder_attention_mask: (n_records, l_decoder)
268
- indices = torch.tensor( # (b,)
269
- batch_positions,
270
- device=self.decoder_attention_mask.device,
271
- dtype=torch.long,
272
- )
273
- batch_model_kwargs["decoder_attention_mask"] = (
274
- self.decoder_attention_mask.index_select(0, indices) # (b, l_decoder)
275
- )
276
- custom_batch = self._embedding_batch_fn or getattr(self.model, "_embedding_batch", None)
277
- if custom_batch is not None:
278
- if self.model_type == "fast_ankh":
279
- batch = custom_batch(
280
- sequences,
281
- tokenizer=self.resolved_tokenizer,
282
- max_length=self.max_length,
283
- truncate=self.truncate,
284
- need_attentions=self.need_attentions,
285
- **batch_model_kwargs,
286
- )
287
- else:
288
- batch = custom_batch(sequences, **batch_model_kwargs)
289
- if not isinstance(batch, EmbeddingBatch):
290
- raise TypeError("_embedding_batch must return EmbeddingBatch.")
291
- else:
292
- batch = _generic_embedding_batch(
293
- self.model,
294
- sequences,
295
- tokenizer=self.tokenizer,
296
- max_length=self.max_length,
297
- truncate=self.truncate,
298
- need_attentions=self.need_attentions,
299
- model_kwargs=batch_model_kwargs,
300
- )
301
- X = batch.X # (b, l, d) or (b, n_states, l, d)
302
- raw_mask = batch.residue_mask # (b, l)
303
- if not isinstance(X, Tensor) or not isinstance(raw_mask, Tensor):
304
- raise TypeError("Embedding batches must provide Tensor X and residue_mask.")
305
- if X.is_meta or raw_mask.is_meta:
306
- raise ValueError("Embedding batches cannot contain meta tensors.")
307
- if not X.is_floating_point():
308
- raise TypeError("Embedding batches must use a floating-point X dtype.")
309
- if raw_mask.is_complex() or not bool(torch.isfinite(raw_mask).all()):
310
- raise ValueError("Embedding residue_mask must contain finite binary values.")
311
- if not bool(((raw_mask == 0) | (raw_mask == 1)).all()):
312
- raise ValueError("Embedding residue_mask must contain finite binary values.")
313
- M = raw_mask.to(device=X.device, dtype=torch.bool) # (b, l)
314
- valid_X_shape = (
315
- X.ndim == 3
316
- and X.shape[0] == len(batch_records)
317
- and X.shape[-1] > 0
318
- and M.shape == X.shape[:2]
319
- )
320
- valid_all_states_shape = (
321
- X.ndim == 4
322
- and self.store_all_hidden_states
323
- and self.full_embeddings
324
- and X.shape[0] == len(batch_records)
325
- and X.shape[1] > 0
326
- and X.shape[-1] > 0
327
- and M.shape == (X.shape[0], X.shape[2])
328
- )
329
- if not (valid_X_shape or valid_all_states_shape):
330
- raise ValueError(
331
- "Embedding batches must provide X with shape (b, l, d), or "
332
- "(b, states, l, d) when storing all hidden states, and "
333
- "residue_mask with shape (b, l)."
334
- )
335
- if not bool(M.any(dim=1).all()):
336
- raise ValueError("Every embedding sample must contain a biological residue.")
337
- finite_selected = ( # X.shape
338
- torch.isfinite(X) | ~M.unsqueeze(-1)
339
- if X.ndim == 3
340
- else torch.isfinite(X) | ~M[:, None, :, None]
341
- )
342
- if not bool(finite_selected.all()):
343
- raise ValueError("Biological residue embeddings produced non-finite output.")
344
- if self.need_attentions:
345
- # Validate the biological graph only after mask integrity is established.
346
- _validate_parti_length(M)
347
- if self.dtype is not None:
348
- X = X.to(dtype=self.dtype) # unchanged shape
349
-
350
- if self.full_embeddings:
351
- if X.ndim == 4:
352
- values = [
353
- X_i[:, M_i, :].detach().cpu() # (n_states, r_i, d)
354
- for X_i, M_i in zip(X, M, strict=True)
355
- ]
356
- else:
357
- values = _residue_embeddings(X, M) # each: (r_i, d)
358
- else:
359
- if self.pooler is None:
360
- raise RuntimeError(
361
- "Pooled embedding output was requested without an initialized pooler."
362
- )
363
- Y = self.pooler( # (b, n_poolers * d)
364
- X,
365
- M,
366
- attentions=batch.attentions,
367
- attention_backend=self.attention_backend,
368
- )
369
- pool_slices = self.pooler.output_slices(X.shape[-1])
370
- values = list(Y.detach().cpu().unbind(0)) # each: (n_poolers * d,)
371
- for position, record, value in zip(batch_positions, batch_records, values, strict=True):
372
- window_results[position] = EmbeddingRecord(record.id, record.sequence, value)
373
-
374
- new_records = [
375
- window_results[position]
376
- for position in range(window_start, window_start + len(window_records))
377
- ]
378
- return new_records, pool_slices
 
 
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
- from collections.abc import Iterable, Mapping, Sequence
10
- from pathlib import Path
11
- from typing import Any
12
- from torch import Tensor
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
- from collections.abc import Iterable, Iterator, Mapping, Sequence
9
- from pathlib import Path
10
- from typing import overload
11
-
12
- from .types import EmbeddingInput
13
-
14
-
15
- def iter_fasta(path: str | Path) -> Iterator[EmbeddingInput]:
16
- """Yield FASTA records in source order without reading the file into memory."""
17
-
18
- identifier: str | None = None
19
- sequence_parts: list[str] = []
20
- found_record = False
21
- with Path(path).open("r", encoding="utf-8") as handle:
22
- for line_number, raw_line in enumerate(handle, start=1):
23
- line = raw_line.strip()
24
- if not line:
25
- continue
26
- if line.startswith(">"):
27
- if identifier is not None:
28
- found_record = True
29
- yield EmbeddingInput(identifier, "".join(sequence_parts))
30
- identifier = line[1:].strip().split(maxsplit=1)[0]
31
- if not identifier:
32
- raise ValueError(f"Missing FASTA identifier on line {line_number}.")
33
- sequence_parts = []
34
- else:
35
- if identifier is None:
36
- raise ValueError(
37
- f"Sequence data precedes the first FASTA header on line {line_number}."
38
- )
39
- sequence_parts.append("".join(line.split()))
40
- if identifier is not None:
41
- found_record = True
42
- yield EmbeddingInput(identifier, "".join(sequence_parts))
43
- if not found_record:
44
- raise ValueError(f"No FASTA records found in {path}.")
45
-
46
-
47
- def parse_fasta(path: str | Path) -> list[EmbeddingInput]:
48
- """Parse FASTA records while preserving identifiers, order, and duplicates."""
49
-
50
- return list(iter_fasta(path))
51
-
52
-
53
- def _normalize_input_item(
54
- position: int,
55
- item: str | EmbeddingInput | tuple[str, str],
56
- ) -> EmbeddingInput:
57
- if isinstance(item, EmbeddingInput):
58
- return item
59
- if isinstance(item, str):
60
- return EmbeddingInput(str(position), item)
61
- if isinstance(item, tuple) and len(item) == 2:
62
- return EmbeddingInput(str(item[0]), str(item[1]))
63
- raise TypeError(
64
- "inputs must contain sequences, EmbeddingInput values, or (id, sequence) tuples."
65
- )
66
-
67
-
68
- class _InputSpool(Sequence[EmbeddingInput]):
69
- """Immutable disk-backed normalized inputs with an incremental digest."""
70
-
71
- def __init__(
72
- self,
73
- values: Iterable[str | EmbeddingInput | tuple[str, str]],
74
- ) -> None:
75
- self._temporary: tempfile.TemporaryDirectory[str] | None = tempfile.TemporaryDirectory(
76
- prefix="fastplms-inputs-"
77
- )
78
- self.path = Path(self._temporary.name) / "inputs.sqlite"
79
- self._connection: sqlite3.Connection | None = sqlite3.connect(self.path)
80
- self._connection.execute(
81
- "CREATE TABLE inputs ("
82
- "position INTEGER PRIMARY KEY, input_id TEXT NOT NULL, sequence TEXT NOT NULL)"
83
- )
84
- digest = hashlib.sha256()
85
- count = 0
86
- pending: list[tuple[int, str, str]] = []
87
- try:
88
- for position, item in enumerate(values):
89
- record = _normalize_input_item(position, item)
90
- for value in (record.id, record.sequence):
91
- encoded = value.encode("utf-8")
92
- digest.update(len(encoded).to_bytes(8, "big"))
93
- digest.update(encoded)
94
- pending.append((position, record.id, record.sequence))
95
- count += 1
96
- if len(pending) == 1_024:
97
- self._connection.executemany("INSERT INTO inputs VALUES (?, ?, ?)", pending)
98
- pending.clear()
99
- if pending:
100
- self._connection.executemany("INSERT INTO inputs VALUES (?, ?, ?)", pending)
101
- if count == 0:
102
- raise ValueError("inputs must contain at least one sequence.")
103
- self._connection.commit()
104
- self._connection.close()
105
- self._connection = sqlite3.connect(
106
- f"{self.path.resolve().as_uri()}?mode=ro",
107
- uri=True,
108
- )
109
- except BaseException:
110
- self.close()
111
- raise
112
- digest.update(count.to_bytes(8, "big"))
113
- self.input_fingerprint = digest.hexdigest()
114
- self._count = count
115
-
116
- def _require_connection(self) -> sqlite3.Connection:
117
- if self._connection is None:
118
- raise RuntimeError("Input spool is closed.")
119
- return self._connection
120
-
121
- def __len__(self) -> int:
122
- return self._count
123
-
124
- def __iter__(self) -> Iterator[EmbeddingInput]:
125
- cursor = self._require_connection().execute(
126
- "SELECT input_id, sequence FROM inputs ORDER BY position"
127
- )
128
- while rows := cursor.fetchmany(1_024):
129
- for input_id, sequence in rows:
130
- yield EmbeddingInput(input_id, sequence)
131
-
132
- @overload
133
- def __getitem__(self, index: int, /) -> EmbeddingInput: ...
134
-
135
- @overload
136
- def __getitem__(self, index: slice, /) -> list[EmbeddingInput]: ...
137
-
138
- def __getitem__(self, index: int | slice) -> EmbeddingInput | list[EmbeddingInput]:
139
- connection = self._require_connection()
140
-
141
- if isinstance(index, slice):
142
- start, stop, step = index.indices(self._count)
143
- if step != 1:
144
- return [self[position] for position in range(start, stop, step)]
145
- rows = connection.execute(
146
- "SELECT input_id, sequence FROM inputs "
147
- "WHERE position >= ? AND position < ? ORDER BY position",
148
- (start, stop),
149
- ).fetchall()
150
- return [EmbeddingInput(input_id, sequence) for input_id, sequence in rows]
151
- position = index + self._count if index < 0 else index
152
- if position < 0 or position >= self._count:
153
- raise IndexError(index)
154
- row = connection.execute(
155
- "SELECT input_id, sequence FROM inputs WHERE position = ?", (position,)
156
- ).fetchone()
157
- if row is None:
158
- raise IndexError(index)
159
- return EmbeddingInput(row[0], row[1])
160
-
161
- def close(self) -> None:
162
- connection = getattr(self, "_connection", None)
163
- if connection is not None:
164
- connection.close()
165
- self._connection = None
166
- temporary = getattr(self, "_temporary", None)
167
- if temporary is not None:
168
- temporary.cleanup()
169
- self._temporary = None
170
-
171
- def __del__(self) -> None:
172
- self.close()
173
-
174
-
175
- def _normalize_inputs(
176
- inputs: (Iterable[str | EmbeddingInput | tuple[str, str]] | Mapping[str, str] | str | Path),
177
- *,
178
- disk_backed: bool,
179
- ) -> Sequence[EmbeddingInput]:
180
- is_fasta_path = isinstance(inputs, Path)
181
- if isinstance(inputs, str):
182
- try:
183
- is_fasta_path = Path(inputs).is_file()
184
- except OSError:
185
- is_fasta_path = False
186
- should_spool = disk_backed or is_fasta_path or not isinstance(inputs, (str, Sequence, Mapping))
187
- values: Iterable[str | EmbeddingInput | tuple[str, str]]
188
- if isinstance(inputs, Path):
189
- values = iter_fasta(inputs)
190
- elif isinstance(inputs, str):
191
- values = iter_fasta(inputs) if is_fasta_path else [inputs]
192
- elif isinstance(inputs, Mapping):
193
- values = inputs.items()
194
- else:
195
- values = inputs
196
- if should_spool:
197
- return _InputSpool(values)
198
- records: list[EmbeddingInput] = []
199
- for position, item in enumerate(values):
200
- records.append(_normalize_input_item(position, item))
201
- if not records:
202
- raise ValueError("inputs must contain at least one sequence.")
203
- return records
204
-
205
-
206
- def _validate_untruncated_lengths(
207
- records: Sequence[EmbeddingInput],
208
- *,
209
- max_length: int | None,
210
- truncate: bool,
211
- ) -> None:
212
- """Fail before inference when a biological-residue limit would be exceeded."""
213
-
214
- if max_length is None or truncate:
215
- return
216
- for position, record in enumerate(records):
217
- residue_count = len(record.sequence)
218
- if residue_count > max_length:
219
- raise ValueError(
220
- f"Input at position {position} with id {record.id!r} has "
221
- f"{residue_count} biological residues, exceeding max_length={max_length} "
222
- "while truncate=False."
223
- )
224
-
225
-
226
- def _planned_batches(
227
- records: Sequence[EmbeddingInput],
228
- positions: range,
229
- *,
230
- batch_size: int,
231
- max_tokens_per_batch: int | None,
232
- max_length: int | None,
233
- truncate: bool,
234
- ) -> Iterator[list[int]]:
235
- """Length-bucket one bounded window while retaining stable output positions."""
236
-
237
- def effective_length(position: int) -> int:
238
- length = len(records[position].sequence)
239
- return min(length, max_length) if truncate and max_length is not None else length
240
-
241
- ordered = sorted(positions, key=lambda position: (-effective_length(position), position))
242
- batch: list[int] = []
243
- longest = 0
244
- for position in ordered:
245
- length = effective_length(position)
246
- if max_tokens_per_batch is not None and length > max_tokens_per_batch:
247
- raise ValueError(
248
- f"Input at position {position} has {length} residues, exceeding "
249
- f"max_tokens_per_batch={max_tokens_per_batch}."
250
- )
251
- candidate_longest = max(longest, length)
252
- exceeds_tokens = (
253
- max_tokens_per_batch is not None
254
- and candidate_longest * (len(batch) + 1) > max_tokens_per_batch
255
- )
256
- if batch and (len(batch) >= batch_size or exceeds_tokens):
257
- yield batch
258
- batch = []
259
- longest = 0
260
- batch.append(position)
261
- longest = max(longest, length)
262
- if batch:
263
- yield batch
 
 
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
- from collections.abc import Callable, Iterable, Mapping, Sequence
7
- from pathlib import Path
8
- from typing import Any
 
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
- conjugate_sign = torch.tensor([1, -1, -1, -1], device=quaternion.device)
31
- return quaternion * conjugate_sign
 
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
- transformed = [func(component) for component in self.tensor.unbind(dim=-1)]
142
- return self._from_tensor(torch.stack(transformed, dim=-1))
 
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
- return _quat_rotation(self.normalized()._quats, points)
 
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
- return RotationMatrix(_graham_schmidt(x_axis, xy_plane, eps))
 
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
- x_axis = origin - neg_x_axis
461
- plane_direction = xy_plane - origin
 
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
- components = [func(value) for value in self.tensor.unbind(dim=-1)]
495
- return Affine3D.from_tensor(torch.stack(components, dim=-1))
 
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
- return self.rot.apply(points) + self.trans
 
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
- n, ca, c = positions.unbind(dim=-2)
574
- return Affine3D.from_graham_schmidt(c, ca, n)
 
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
- return torch.as_tensor(structure.atom37_positions, dtype=torch.double).unsqueeze(0)
 
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
- displacement = positions[..., None, :] - positions[..., None, :, :]
22
- return torch.sqrt(eps + torch.sum(displacement**2, dim=-1))
 
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(data, inds, dim=0, no_batch_dims=0):
 
 
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
- matrix = np.asarray(array).view(np.uint8)
18
- return matrix.reshape(matrix.shape[0], -1)
 
19
 
20
 
21
  def _hamming_to_all(query: np.ndarray, sequences: np.ndarray) -> np.ndarray:
22
- return np.not_equal(sequences, query).mean(axis=1, dtype=np.float64)
 
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
- n_position, ca_position, c_position = bb_positions.unbind(dim=-2)
21
- return Affine3D.from_graham_schmidt(c_position, ca_position, n_position)
 
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 one or more named atoms along an atom37 axis."""
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 from backbone coordinates ``X`` with shape (l, 37, 3)."""
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
- transformed = frame[..., None, None].invert().apply(coords)
59
- frame_is_valid = frame.trans.norm(dim=-1) > 0
60
- normalized = torch.where(frame_is_valid[..., None, None, None], transformed, coords)
61
- return normalized.masked_fill(torch.isinf(coords), torch.inf)
 
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
- M = atom_mask.bool().cpu().numpy()
96
- X = coords.float().cpu().numpy()
97
- atom_names = ref_atom_name_chars.cpu().numpy()
98
- elements = ref_element.cpu().numpy()
99
- confidence = None if plddt is None else plddt.float().cpu().numpy()
 
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 M[atom_index]:
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
- residue_mask = mask.bool()
16
- return residue_mask.unsqueeze(-1) & residue_mask.unsqueeze(-2)
 
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
- weights = mask.expand_as(value)
50
- weighted_sum = torch.sum(weights * value, dim=dim)
51
- weight_sum = torch.sum(weights, dim=dim)
52
- return weighted_sum / (weight_sum + eps)
 
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 pairwise PAE logits."""
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
- origins = frames.trans[..., None, :, :]
87
- return frames.invert()[..., None].apply(origins)
 
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 discretized aligned-position errors."""
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
- return X / np.sqrt(np.square(X).sum(-1, keepdims=True) + 1e-8)
 
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
- coords = torch.from_numpy(self.atom37_positions).to(frame.trans.dtype)
729
- coords = apply_frame_to_coords(coords, frame)
730
- atom37_positions = coords.numpy()
 
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
- n_position = index_by_atom_name(atom37, "N", dim=-2)
46
- ca_position = index_by_atom_name(atom37, "CA", dim=-2)
47
- c_position = index_by_atom_name(atom37, "C", dim=-2)
 
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
- unpacked_mobile = unbinpack(mobile, sequence_id, pad_value=torch.nan)
75
- unpacked_target = unbinpack(target, sequence_id, pad_value=torch.nan)
 
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
- num_valid_atoms = atom_exists_mask.sum(dim=-1, keepdim=True)
151
- centroid_mobile = mobile.sum(dim=-2, keepdim=True) / num_valid_atoms.unsqueeze(-1)
152
- centroid_target = target.sum(dim=-2, keepdim=True) / num_valid_atoms.unsqueeze(-1)
153
- centroid_mobile[num_valid_atoms == 0] = 0
154
- centroid_target[num_valid_atoms == 0] = 0
155
-
156
- expanded_mask = atom_exists_mask.unsqueeze(-1)
157
- centered_mobile = (mobile - centroid_mobile).masked_fill(~expanded_mask, 0)
158
- centered_target = (target - centroid_target).masked_fill(~expanded_mask, 0)
159
- covariance = torch.matmul(centered_mobile.transpose(1, 2), centered_target)
160
- left_vectors, _, right_vectors = torch.svd(covariance)
161
- rotation = torch.matmul(left_vectors, right_vectors.transpose(1, 2))
 
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
- difference = aligned - target
 
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
- deviation = torch.linalg.vector_norm(aligned - target, dim=-1)
245
- counts = atom_exists_mask.sum(dim=-1)
246
- score_1 = ((deviation < 1) * atom_exists_mask).sum(dim=-1) / counts
247
- score_2 = ((deviation < 2) * atom_exists_mask).sum(dim=-1) / counts
248
- score_4 = ((deviation < 4) * atom_exists_mask).sum(dim=-1) / counts
249
- score_8 = ((deviation < 8) * atom_exists_mask).sum(dim=-1) / counts
250
- score = (score_1 + score_2 + score_4 + score_8) * 0.25
251
- return score.mean() if reduction == "batch" else score
 
 
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 transformers import ESMFold2Model
 
6
 
7
- model = ESMFold2Model.from_pretrained("biohub/ESMFold2").cuda().eval()
8
- open("ubq.pdb", "w").write(model.infer_protein_as_pdb("MQIFVKTLTGKT..."))
 
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
- return x if num_diffusion_samples == 1 else x.repeat_interleave(num_diffusion_samples, 0)
 
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
- s_inputs_normed = self.s_inputs_norm(s_inputs)
 
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
- pair = pair + self.outer_product_mean(m, msa_attention_mask)
 
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 path broadcasts a row/column-shared mask M with shape (1, ...).
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
- scores = self.attn_proj(z).squeeze(-1)
 
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
- bij_same_chain = asym_id.unsqueeze(2) == asym_id.unsqueeze(1)
515
- bij_same_residue = residue_index.unsqueeze(2) == residue_index.unsqueeze(1)
516
- bij_same_entity = entity_id.unsqueeze(2) == entity_id.unsqueeze(1)
 
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 = self.downproject(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
- return torch.einsum(self._einsum_equation, left_stream, right_stream)
 
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
- return self._engine(z, visibility=mask)
 
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
- m_norm = self.norm(m)
1205
- x = self.W(m_norm) * msa_attention_mask.unsqueeze(-1).to(m_norm.dtype)
1206
- a, b = x.chunk(2, dim=-1)
1207
- mask_f = msa_attention_mask.to(a.dtype)
1208
- n_valid = (mask_f @ mask_f.transpose(-1, -2)).unsqueeze(-1).clamp(min=1.0)
 
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 = self.norm(x)
1287
- a = self.a_proj(x)
1288
- b = self.b_proj(x)
1289
- return self.out_proj(F.silu(a) * b)
 
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
- a_norm = F.layer_norm(a, (self.d_model,), None, None, self.eps)
1311
- s_norm = F.layer_norm(s, (self.d_cond,), self.s_scale, None, self.eps)
1312
- return torch.sigmoid(self.s_gate(s_norm)) * a_norm + self.s_shift(s_norm)
 
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
- t = torch.as_tensor(t_hat, device=self.w.device, dtype=self.w.dtype).reshape(-1)
1334
- return torch.cos(2.0 * torch.pi * (t[:, None] * self.w[None, :] + self.b[None, :]))
 
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
- x12 = self.w12(x)
1364
- x1, x2 = x12.split(self.hidden_features, dim=-1)
1365
- hidden = F.silu(x1)
 
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
- x1, x2 = x.chunk(2, dim=-1)
1390
- return torch.cat((-x2, x1), dim=-1)
 
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
- return F.rms_norm(x, (x.size(-1),)).to(x.dtype)
 
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
- x = x.to(self.w_up.weight.dtype)
1483
- x1, x2 = self.w_up(x).chunk(2, dim=-1)
1484
- return self.w_down(F.silu(x1) * x2)
 
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
- return F.rms_norm(x, (x.shape[-1],)) * (1 + scale) + shift
 
1597
 
1598
 
1599
  def _gated_residual_raw(x: Tensor, gate: Tensor, y: Tensor) -> Tensor:
1600
- return x + gate * y
 
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
- mod = self.adaln_modulation(c_l)
 
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
- mask_exp = atom_attention_mask.repeat_interleave(num_diffusion_samples, 0)
1725
- seqlens = mask_exp.sum(dim=-1, dtype=torch.int32)
1726
- indices = torch.nonzero(mask_exp.flatten(), as_tuple=False).flatten()
 
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
- atom_to_token_exp = atom_to_token.repeat_interleave(num_diffusion_samples, 0)
1939
- a_to_q = self.token_to_atom_linear(a_i)
1940
- a_to_q = gather_token_to_atom(a_to_q, atom_to_token_exp)
1941
- q_l = q_l + a_to_q
 
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
- ) # A has shape (b, h, q, k).
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
- x = self.adaln(a, s) if s is not None else self.pre_norm(a)
 
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
- s_inputs_normed = self.s_inputs_norm(s_inputs)
147
- z_base = self.z_norm(z)
 
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
- return self.ffn(self.norm(x))
 
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