lhallee commited on
Commit
412308b
·
verified ·
1 Parent(s): 3d6499f

Apply coding standards from 1cb5747 (files only)

Browse files
README.md CHANGED
@@ -79,8 +79,10 @@ shows padding explicitly:
79
 
80
  ```python
81
  import torch
 
82
  from transformers import AutoTokenizer
83
 
 
84
  model_id = "Synthyra/ESMplusplus_large"
85
  tokenizer = AutoTokenizer.from_pretrained(
86
  model_id,
@@ -129,12 +131,14 @@ Residue labels have shape `(b, l)` and use `-100` outside biological positions.
129
 
130
  ```python
131
  import torch
 
132
  from transformers import AutoTokenizer
133
  from transformers import (
134
  AutoModelForSequenceClassification,
135
  AutoModelForTokenClassification,
136
  )
137
 
 
138
  model_id = "Synthyra/ESMplusplus_large"
139
  sequence_model = AutoModelForSequenceClassification.from_pretrained(
140
  model_id, num_labels=2, trust_remote_code=True
@@ -145,13 +149,13 @@ token_model = AutoModelForTokenClassification.from_pretrained(
145
  tokenizer = AutoTokenizer.from_pretrained(model_id, trust_remote_code=True)
146
  sequences = ["MSTNPKPQRKTKRNT", "MKTIIALSYIFCLVFA"]
147
  batch = tokenizer(sequences, padding=True, return_tensors="pt")
148
- biological = batch["attention_mask"].bool()
149
  for special_id in tokenizer.all_special_ids:
150
- biological &= batch["input_ids"].ne(special_id)
151
 
152
- sequence_labels = torch.zeros(len(sequences), dtype=torch.long)
153
- token_labels = torch.full_like(batch["input_ids"], -100)
154
- token_labels[biological] = 0
155
 
156
  with torch.inference_mode():
157
  sequence_output = sequence_model(**batch, labels=sequence_labels)
@@ -171,6 +175,7 @@ python -m pip install "datasets>=4.8,<5" "peft>=0.19,<0.20"
171
  ```python
172
  from peft import LoraConfig, TaskType, get_peft_model
173
 
 
174
  peft_model = get_peft_model(
175
  sequence_model,
176
  LoraConfig(
@@ -198,6 +203,7 @@ adapters. Base checkpoint weights stay frozen:
198
  ```python
199
  from transformers import AutoModelForMaskedLM
200
 
 
201
  ttt_model = AutoModelForMaskedLM.from_pretrained(
202
  "Synthyra/ESMplusplus_large",
203
  trust_remote_code=True,
@@ -243,12 +249,13 @@ Select an SAE for this ESMC scale, then load only the layers you need:
243
  ```python
244
  import torch
245
 
 
246
  model.load_sae_models("biohub/ESMC-600M-sae-layer27-k64-codebook65536", [27])
247
 
248
  with torch.inference_mode():
249
  output = model(**batch, normalize_sae=True)
250
 
251
- features = output.sae_outputs["layer27"]
252
  print(features.shape, features.layout) # (valid_token_count, codebook_dim), sparse COO
253
  ```
254
 
 
79
 
80
  ```python
81
  import torch
82
+
83
  from transformers import AutoTokenizer
84
 
85
+
86
  model_id = "Synthyra/ESMplusplus_large"
87
  tokenizer = AutoTokenizer.from_pretrained(
88
  model_id,
 
131
 
132
  ```python
133
  import torch
134
+
135
  from transformers import AutoTokenizer
136
  from transformers import (
137
  AutoModelForSequenceClassification,
138
  AutoModelForTokenClassification,
139
  )
140
 
141
+
142
  model_id = "Synthyra/ESMplusplus_large"
143
  sequence_model = AutoModelForSequenceClassification.from_pretrained(
144
  model_id, num_labels=2, trust_remote_code=True
 
149
  tokenizer = AutoTokenizer.from_pretrained(model_id, trust_remote_code=True)
150
  sequences = ["MSTNPKPQRKTKRNT", "MKTIIALSYIFCLVFA"]
151
  batch = tokenizer(sequences, padding=True, return_tensors="pt")
152
+ biological = batch["attention_mask"].bool() # (b, l)
153
  for special_id in tokenizer.all_special_ids:
154
+ biological &= batch["input_ids"].ne(special_id) # (b, l)
155
 
156
+ sequence_labels = torch.zeros(len(sequences), dtype=torch.long) # (b,)
157
+ token_labels = torch.full_like(batch["input_ids"], -100) # (b, l)
158
+ token_labels[biological] = 0 # selected biological positions; labels stay (b, l)
159
 
160
  with torch.inference_mode():
161
  sequence_output = sequence_model(**batch, labels=sequence_labels)
 
175
  ```python
176
  from peft import LoraConfig, TaskType, get_peft_model
177
 
178
+
179
  peft_model = get_peft_model(
180
  sequence_model,
181
  LoraConfig(
 
203
  ```python
204
  from transformers import AutoModelForMaskedLM
205
 
206
+
207
  ttt_model = AutoModelForMaskedLM.from_pretrained(
208
  "Synthyra/ESMplusplus_large",
209
  trust_remote_code=True,
 
249
  ```python
250
  import torch
251
 
252
+
253
  model.load_sae_models("biohub/ESMC-600M-sae-layer27-k64-codebook65536", [27])
254
 
255
  with torch.inference_mode():
256
  output = model(**batch, normalize_sae=True)
257
 
258
+ features = output.sae_outputs["layer27"] # (valid_tokens, codebook_dim), sparse COO
259
  print(features.shape, features.layout) # (valid_token_count, codebook_dim), sparse COO
260
  ```
261
 
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_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/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
fastplms/registry.py CHANGED
@@ -9,6 +9,7 @@ from __future__ import annotations
9
 
10
  import re
11
  import tomllib
 
12
  from collections.abc import Iterator, Mapping
13
  from dataclasses import dataclass
14
  from functools import lru_cache
 
9
 
10
  import re
11
  import tomllib
12
+
13
  from collections.abc import Iterator, Mapping
14
  from dataclasses import dataclass
15
  from functools import lru_cache
fastplms_bundle.py CHANGED
The diff for this file is too large to render. See raw diff
 
modeling_fastplms.py CHANGED
@@ -13,7 +13,7 @@ from zipfile import ZIP_DEFLATED, ZipFile
13
 
14
  from .fastplms_bundle import RUNTIME_DATA, RUNTIME_HASH
15
 
16
- if RUNTIME_HASH != "1b702f30eb56fa1c6b436c423cd901ab7583868bdc38e15da835a315d3ce15b8":
17
  raise RuntimeError("FastPLMs runtime identity differs from the bridge.")
18
 
19
  _RUNTIME_TEMPORARIES = []
 
13
 
14
  from .fastplms_bundle import RUNTIME_DATA, RUNTIME_HASH
15
 
16
+ if RUNTIME_HASH != "729de763a994eb7320c3876cb7f6a480c93059dd407aead6004991988217ab55":
17
  raise RuntimeError("FastPLMs runtime identity differs from the bridge.")
18
 
19
  _RUNTIME_TEMPORARIES = []