lhallee commited on
Commit
14bc716
·
verified ·
1 Parent(s): 1f3f255

Update FastPLMs runtime and cards from 328deef (files only)

Browse files
fastplms/attention/__init__.py CHANGED
@@ -1,10 +1,17 @@
1
  """Shared attention backends, masks, and optional optimized kernels."""
2
 
 
 
 
 
 
 
3
  from ._core import (
4
  LEGACY_CHECKPOINT_ATTENTION_BACKENDS,
5
  VALID_ATTENTION_BACKENDS,
6
  AttentionBackend,
7
  BlockMask,
 
8
  _ensure_flash_kernels_loaded,
9
  _get_flex_attention_fn,
10
  _get_flex_block_mask,
@@ -18,6 +25,7 @@ from ._core import (
18
  flex_attention,
19
  get_attention_mask,
20
  get_attn_implementation,
 
21
  index_first_axis,
22
  index_put_first_axis,
23
  kernels_flash_attention_func,
@@ -36,13 +44,18 @@ from .interfaces import (
36
 
37
 
38
  __all__ = [
 
39
  "FASTPLMS_ATTENTION_FUNCTIONS",
40
  "FASTPLMS_ATTENTION_MASKS",
41
  "LEGACY_CHECKPOINT_ATTENTION_BACKENDS",
42
  "VALID_ATTENTION_BACKENDS",
43
  "AttentionBackend",
 
 
 
44
  "BlockMask",
45
  "FastPLMsAttentionMixin",
 
46
  "_ensure_flash_kernels_loaded",
47
  "_get_flex_attention_fn",
48
  "_get_flex_block_mask",
@@ -56,6 +69,7 @@ __all__ = [
56
  "flex_attention",
57
  "get_attention_mask",
58
  "get_attn_implementation",
 
59
  "index_first_axis",
60
  "index_put_first_axis",
61
  "kernels_flash_attention_func",
 
1
  """Shared attention backends, masks, and optional optimized kernels."""
2
 
3
+ from ._auto import (
4
+ AUTO_ATTENTION,
5
+ AttentionCandidate,
6
+ AttentionExecutionContext,
7
+ AttentionResolution,
8
+ )
9
  from ._core import (
10
  LEGACY_CHECKPOINT_ATTENTION_BACKENDS,
11
  VALID_ATTENTION_BACKENDS,
12
  AttentionBackend,
13
  BlockMask,
14
+ FlashPaddingLayout,
15
  _ensure_flash_kernels_loaded,
16
  _get_flex_attention_fn,
17
  _get_flex_block_mask,
 
25
  flex_attention,
26
  get_attention_mask,
27
  get_attn_implementation,
28
+ get_flash_padding_layout,
29
  index_first_axis,
30
  index_put_first_axis,
31
  kernels_flash_attention_func,
 
44
 
45
 
46
  __all__ = [
47
+ "AUTO_ATTENTION",
48
  "FASTPLMS_ATTENTION_FUNCTIONS",
49
  "FASTPLMS_ATTENTION_MASKS",
50
  "LEGACY_CHECKPOINT_ATTENTION_BACKENDS",
51
  "VALID_ATTENTION_BACKENDS",
52
  "AttentionBackend",
53
+ "AttentionCandidate",
54
+ "AttentionExecutionContext",
55
+ "AttentionResolution",
56
  "BlockMask",
57
  "FastPLMsAttentionMixin",
58
+ "FlashPaddingLayout",
59
  "_ensure_flash_kernels_loaded",
60
  "_get_flex_attention_fn",
61
  "_get_flex_block_mask",
 
69
  "flex_attention",
70
  "get_attention_mask",
71
  "get_attn_implementation",
72
+ "get_flash_padding_layout",
73
  "index_first_axis",
74
  "index_put_first_axis",
75
  "kernels_flash_attention_func",
fastplms/attention/_auto.py ADDED
@@ -0,0 +1,182 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Opt-in automatic selection of an attention implementation.
2
+
3
+ ``attn_implementation="auto"`` is a request, not a backend. Each family lists
4
+ its implementations in a preference order backed by measured evidence, and
5
+ FastPLMs configures the first one that this machine can execute. After that the
6
+ model holds a named implementation. Configuration files and embedding
7
+ fingerprints record that name and never the word ``auto``.
8
+
9
+ Eager attention, SDPA, and Flex attention can be judged when the model is built.
10
+ FlashAttention depends on the device and on the dtype that Q, K, and V will
11
+ have, so a preference order that contains it is resolved at the first forward or
12
+ by an explicit ``resolve_attn_implementation`` call.
13
+ """
14
+
15
+ from __future__ import annotations
16
+
17
+ import torch
18
+
19
+ from dataclasses import dataclass
20
+
21
+ from ._core import _ensure_flash_kernels_loaded, resolve_attention_backend
22
+ from ._kernel_lock import require_kernels_package
23
+
24
+
25
+ AUTO_ATTENTION = "auto"
26
+ _FLASH_IMPLEMENTATIONS = frozenset({"flash_attention_2", "flash_attention_3"})
27
+ _DTYPE_NAMES = {
28
+ torch.float32: "float32",
29
+ torch.bfloat16: "bfloat16",
30
+ torch.float16: "float16",
31
+ }
32
+
33
+
34
+ @dataclass(frozen=True)
35
+ class AttentionExecutionContext:
36
+ """The device and dtype that the attention inputs of a forward will have."""
37
+
38
+ device: torch.device
39
+ dtype: torch.dtype
40
+
41
+
42
+ @dataclass(frozen=True)
43
+ class AttentionCandidate:
44
+ """One implementation from the preference order and why it was or was not usable."""
45
+
46
+ implementation: str
47
+ usable: bool
48
+ reason: str
49
+
50
+
51
+ @dataclass(frozen=True)
52
+ class AttentionResolution:
53
+ """The outcome of one ``auto`` request.
54
+
55
+ While ``deferred`` is true the model runs ``resolved`` provisionally, and the
56
+ first forward or ``resolve_attn_implementation`` replaces this record.
57
+ """
58
+
59
+ requested: str
60
+ resolved: str
61
+ candidates: tuple[AttentionCandidate, ...]
62
+ context: AttentionExecutionContext | None
63
+ deferred: bool
64
+
65
+
66
+ def needs_execution_context(order: tuple[str, ...]) -> bool:
67
+ return any(implementation in _FLASH_IMPLEMENTATIONS for implementation in order)
68
+
69
+
70
+ def provisional_implementation(order: tuple[str, ...]) -> str:
71
+ """Return the implementation a model runs until its context is known."""
72
+ for implementation in order:
73
+ if implementation not in _FLASH_IMPLEMENTATIONS:
74
+ return implementation
75
+ raise ValueError(f"The automatic attention order {order} has no context-free implementation.")
76
+
77
+
78
+ def attention_execution_context(
79
+ module: torch.nn.Module,
80
+ device: torch.device | str | None = None,
81
+ dtype: torch.dtype | None = None,
82
+ ) -> AttentionExecutionContext:
83
+ """Describe where the next forward of ``module`` will run its attention.
84
+
85
+ Under CUDA autocast FP32 parameters produce autocast-dtype Q, K, and V, so
86
+ the autocast dtype is the one that decides kernel eligibility.
87
+ """
88
+ parameter = next(module.parameters(), None)
89
+ if parameter is None:
90
+ raise RuntimeError("Automatic attention selection requires a model with parameters.")
91
+ resolved_device = parameter.device if device is None else torch.device(device)
92
+ if dtype is not None:
93
+ return AttentionExecutionContext(resolved_device, dtype)
94
+ if resolved_device.type == "cuda" and torch.is_autocast_enabled("cuda"):
95
+ return AttentionExecutionContext(resolved_device, torch.get_autocast_dtype("cuda"))
96
+ return AttentionExecutionContext(resolved_device, parameter.dtype)
97
+
98
+
99
+ def _unusable(implementation: str, reason: str) -> AttentionCandidate:
100
+ return AttentionCandidate(implementation, False, reason)
101
+
102
+
103
+ def _flash_candidate(implementation: str, context: AttentionExecutionContext) -> AttentionCandidate:
104
+ """Judge a FlashAttention kernel, leaving the possible download for the last gate."""
105
+ from fastplms.registry import get_model_registry
106
+
107
+ kernel_spec = get_model_registry().attention_kernels[implementation]
108
+ if context.device.type != "cuda":
109
+ return _unusable(
110
+ implementation, f"It requires a CUDA device; the model is on {context.device}."
111
+ )
112
+ dtype_name = _DTYPE_NAMES.get(context.dtype, str(context.dtype))
113
+ if dtype_name not in kernel_spec.dtypes:
114
+ supported = ", ".join(kernel_spec.dtypes)
115
+ return _unusable(
116
+ implementation,
117
+ f"It supports only {supported}; the attention inputs would be {dtype_name}. "
118
+ "Use CUDA BF16 autocast or BF16 weights.",
119
+ )
120
+ capability = torch.cuda.get_device_capability(context.device)
121
+ if capability < kernel_spec.min_cuda_capability:
122
+ required = ".".join(str(part) for part in kernel_spec.min_cuda_capability)
123
+ observed = ".".join(str(part) for part in capability)
124
+ return _unusable(
125
+ implementation,
126
+ f"It requires CUDA compute capability {required} or newer; this GPU has {observed}.",
127
+ )
128
+ try:
129
+ require_kernels_package()
130
+ _ensure_flash_kernels_loaded(implementation)
131
+ except RuntimeError as error:
132
+ return _unusable(implementation, str(error))
133
+ return AttentionCandidate(implementation, True, "The manifest-locked kernel loaded.")
134
+
135
+
136
+ def _candidate(
137
+ implementation: str, context: AttentionExecutionContext | None
138
+ ) -> AttentionCandidate:
139
+ if implementation in _FLASH_IMPLEMENTATIONS:
140
+ if context is None:
141
+ return _unusable(implementation, "It needs a device and dtype, which are not known.")
142
+ return _flash_candidate(implementation, context)
143
+ try:
144
+ # Flex attention is the one device-independent backend a PyTorch build can lack.
145
+ resolve_attention_backend(implementation)
146
+ except RuntimeError as error:
147
+ return _unusable(implementation, str(error))
148
+ return AttentionCandidate(implementation, True, "It runs on every supported device and dtype.")
149
+
150
+
151
+ def resolve_auto_attention(
152
+ order: tuple[str, ...],
153
+ context: AttentionExecutionContext | None,
154
+ ) -> AttentionResolution:
155
+ """Select the first implementation in ``order`` that can execute in ``context``."""
156
+ candidates: list[AttentionCandidate] = []
157
+ for implementation in order:
158
+ candidate = _candidate(implementation, context)
159
+ candidates.append(candidate)
160
+ if candidate.usable:
161
+ return AttentionResolution(
162
+ requested=AUTO_ATTENTION,
163
+ resolved=implementation,
164
+ candidates=tuple(candidates),
165
+ context=context,
166
+ deferred=False,
167
+ )
168
+ reasons = "; ".join(
169
+ f"{candidate.implementation}: {candidate.reason}" for candidate in candidates
170
+ )
171
+ raise RuntimeError(f"No implementation in the automatic attention order is usable. {reasons}")
172
+
173
+
174
+ def deferred_resolution(order: tuple[str, ...]) -> AttentionResolution:
175
+ """Record an ``auto`` request that waits for its execution context."""
176
+ return AttentionResolution(
177
+ requested=AUTO_ATTENTION,
178
+ resolved=provisional_implementation(order),
179
+ candidates=(),
180
+ context=None,
181
+ deferred=True,
182
+ )
fastplms/attention/_core.py CHANGED
@@ -11,6 +11,7 @@ import warnings
11
  import torch
12
  from collections import OrderedDict
13
  from collections.abc import Callable
 
14
  from enum import Enum
15
  from threading import RLock
16
  from types import MappingProxyType
@@ -440,12 +441,53 @@ index_first_axis = IndexFirstAxis.apply
440
  index_put_first_axis = IndexPutFirstAxis.apply
441
 
442
 
 
 
 
 
 
 
 
 
 
 
 
 
 
443
  def pad_input(
444
  hidden_states: torch.Tensor, indices: torch.Tensor, batch: int, seqlen: int
445
  ) -> torch.Tensor:
446
  # hidden_states: (t, ...); indices: (t,)
447
- output = index_put_first_axis(hidden_states, indices, batch * seqlen) # (b * l, ...)
448
- return rearrange(output, "(b s) ... -> b s ...", b=batch) # (b, l, ...)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
449
 
450
 
451
  def _unpad_input(
@@ -453,6 +495,7 @@ def _unpad_input(
453
  key_layer: torch.Tensor,
454
  value_layer: torch.Tensor,
455
  attention_mask_2d: torch.Tensor,
 
456
  ) -> tuple[
457
  torch.Tensor,
458
  torch.Tensor,
@@ -463,17 +506,18 @@ def _unpad_input(
463
  ]:
464
  # query_layer, key_layer, value_layer: (b, l, h, d); attention_mask_2d: (b, l)
465
  batch_size, seq_len, num_heads, head_dim = query_layer.shape
466
- seqlens = attention_mask_2d.sum(dim=1).int() # (b,)
467
- cu_seqlens = F.pad(seqlens.cumsum(0, dtype=torch.int32), (1, 0)) # (b + 1,)
468
- max_seqlen = int(seqlens.max().item())
469
- indices = attention_mask_2d.flatten().nonzero(as_tuple=False).flatten() # (t,)
470
- query_layer = index_first_axis( # (t, h, d)
 
471
  query_layer.reshape(batch_size * seq_len, num_heads, head_dim), indices
472
  )
473
- key_layer = index_first_axis( # (t, h, d)
474
  key_layer.reshape(batch_size * seq_len, num_heads, head_dim), indices
475
  )
476
- value_layer = index_first_axis( # (t, h, d)
477
  value_layer.reshape(batch_size * seq_len, num_heads, head_dim), indices
478
  )
479
  return (
@@ -513,6 +557,25 @@ def _validate_flash_padding_mask(
513
  return attention_mask_2d.to(dtype=torch.bool) # (b, l)
514
 
515
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
516
  def kernels_flash_attention_func(
517
  query_states: torch.Tensor,
518
  key_states: torch.Tensor,
@@ -521,6 +584,7 @@ def kernels_flash_attention_func(
521
  causal: bool = False,
522
  softmax_scale: float | None = None,
523
  implementation: str = "flash_attention_3",
 
524
  ) -> torch.Tensor:
525
  """Public flash-attention entry point with optional padding handling.
526
 
@@ -533,6 +597,10 @@ def kernels_flash_attention_func(
533
  before calling this function (ESM2, DPLM, DPLM2, E1, and ESMFold do), pass
534
  `softmax_scale=1.0`. Otherwise the flash kernel applies its default scale
535
  again, yielding an effective `1/head_dim` scale that drifts across layers.
 
 
 
 
536
  """
537
  # query_states, key_states, value_states: (b, l, h, d)
538
  # attention_mask_2d: (b, l) or None
@@ -559,6 +627,8 @@ def kernels_flash_attention_func(
559
  value_states,
560
  attention_mask_2d,
561
  )
 
 
562
  _ensure_flash_kernels_loaded(implementation)
563
  if attention_mask_2d is not None:
564
  batch_size, q_len = query_states.shape[:2]
@@ -574,6 +644,7 @@ def kernels_flash_attention_func(
574
  key_states,
575
  value_states,
576
  attention_mask_2d,
 
577
  )
578
  attn_output_unpad = _kernels_flash_varlen_forward( # (t, h, d)
579
  query_states=query_states,
@@ -587,8 +658,8 @@ def kernels_flash_attention_func(
587
  softmax_scale=softmax_scale,
588
  implementation=implementation,
589
  )
590
- output = pad_input(attn_output_unpad, indices_q, batch_size, q_len) # (b, l, h, d)
591
- return output.masked_fill(~attention_mask_2d[:, :, None, None], 0) # (b, l, h, d)
592
  else:
593
  return _kernels_flash_forward( # (b, l, h, d)
594
  query_states=query_states,
@@ -803,6 +874,21 @@ def get_attention_mask(
803
  return attention_mask_2d, attention_mask_4d, None # (b, l), (b, 1, 1, l), None
804
 
805
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
806
  def bool_to_additive_mask(
807
  bool_mask: torch.Tensor,
808
  dtype: torch.dtype,
 
11
  import torch
12
  from collections import OrderedDict
13
  from collections.abc import Callable
14
+ from dataclasses import dataclass
15
  from enum import Enum
16
  from threading import RLock
17
  from types import MappingProxyType
 
441
  index_put_first_axis = IndexPutFirstAxis.apply
442
 
443
 
444
+ def _select_first_axis(states: torch.Tensor, indices: torch.Tensor) -> torch.Tensor:
445
+ """Copy the rows at ``indices``, through the autograd wrapper only when it is needed.
446
+
447
+ Without a gradient to route, the wrapper and its reshapes are host overhead
448
+ that every layer pays, and plain indexing copies the same values.
449
+ """
450
+ # states: (n, ...); indices: (m,)
451
+ if states.requires_grad:
452
+ selected: torch.Tensor = index_first_axis(states, indices) # (m, ...)
453
+ return selected
454
+ return states[indices] # (m, ...)
455
+
456
+
457
  def pad_input(
458
  hidden_states: torch.Tensor, indices: torch.Tensor, batch: int, seqlen: int
459
  ) -> torch.Tensor:
460
  # hidden_states: (t, ...); indices: (t,)
461
+ if hidden_states.requires_grad:
462
+ output = index_put_first_axis(hidden_states, indices, batch * seqlen) # (b * l, ...)
463
+ return rearrange(output, "(b s) ... -> b s ...", b=batch) # (b, l, ...)
464
+ output = hidden_states.new_zeros(batch * seqlen, *hidden_states.shape[1:]) # (b * l, ...)
465
+ output[indices] = hidden_states
466
+ return output.view(batch, seqlen, *hidden_states.shape[1:]) # (b, l, ...)
467
+
468
+
469
+ @dataclass(frozen=True)
470
+ class FlashPaddingLayout:
471
+ """Varlen metadata shared by every layer of one padded self-attention forward.
472
+
473
+ The token indices and the longest row each cost a host synchronization to
474
+ derive. An encoder builds this record once per forward so its layers do not
475
+ repeat that work. It is valid only for the mask it was built from.
476
+ """
477
+
478
+ attention_mask_2d: torch.Tensor # (b, l), bool; True marks a real token
479
+ indices: torch.Tensor # (t,), positions of the t real tokens in the flat (b * l) axis
480
+ cu_seqlens: torch.Tensor # (b + 1,), int32 cumulative row lengths
481
+ max_seqlen: int # longest row
482
+
483
+
484
+ def _flash_padding_layout(attention_mask_2d: torch.Tensor) -> FlashPaddingLayout:
485
+ # attention_mask_2d: (b, l), bool
486
+ seqlens = attention_mask_2d.sum(dim=1).int() # (b,)
487
+ cu_seqlens = F.pad(seqlens.cumsum(0, dtype=torch.int32), (1, 0)) # (b + 1,)
488
+ max_seqlen = int(seqlens.max().item())
489
+ indices = attention_mask_2d.flatten().nonzero(as_tuple=False).flatten() # (t,)
490
+ return FlashPaddingLayout(attention_mask_2d, indices, cu_seqlens, max_seqlen)
491
 
492
 
493
  def _unpad_input(
 
495
  key_layer: torch.Tensor,
496
  value_layer: torch.Tensor,
497
  attention_mask_2d: torch.Tensor,
498
+ padding_layout: FlashPaddingLayout | None = None,
499
  ) -> tuple[
500
  torch.Tensor,
501
  torch.Tensor,
 
506
  ]:
507
  # query_layer, key_layer, value_layer: (b, l, h, d); attention_mask_2d: (b, l)
508
  batch_size, seq_len, num_heads, head_dim = query_layer.shape
509
+ if padding_layout is None:
510
+ padding_layout = _flash_padding_layout(attention_mask_2d)
511
+ indices = padding_layout.indices # (t,)
512
+ cu_seqlens = padding_layout.cu_seqlens # (b + 1,)
513
+ max_seqlen = padding_layout.max_seqlen
514
+ query_layer = _select_first_axis( # (t, h, d)
515
  query_layer.reshape(batch_size * seq_len, num_heads, head_dim), indices
516
  )
517
+ key_layer = _select_first_axis( # (t, h, d)
518
  key_layer.reshape(batch_size * seq_len, num_heads, head_dim), indices
519
  )
520
+ value_layer = _select_first_axis( # (t, h, d)
521
  value_layer.reshape(batch_size * seq_len, num_heads, head_dim), indices
522
  )
523
  return (
 
557
  return attention_mask_2d.to(dtype=torch.bool) # (b, l)
558
 
559
 
560
+ def _validate_flash_padding_layout(
561
+ padding_layout: FlashPaddingLayout,
562
+ attention_mask_2d: torch.Tensor | None,
563
+ ) -> None:
564
+ """Reject a layout that cannot belong to this call without reading tensor values."""
565
+
566
+ # attention_mask_2d: (b, l) or None
567
+ if attention_mask_2d is None:
568
+ raise ValueError("A FlashAttention padding layout requires the mask it was built from.")
569
+ layout_mask = padding_layout.attention_mask_2d # (b, l)
570
+ if layout_mask.shape != attention_mask_2d.shape:
571
+ raise ValueError(
572
+ "FlashAttention padding layout was built for mask shape "
573
+ f"{tuple(layout_mask.shape)}, but this call uses {tuple(attention_mask_2d.shape)}."
574
+ )
575
+ if layout_mask.device != attention_mask_2d.device:
576
+ raise ValueError("FlashAttention padding layout and padding mask must share a device.")
577
+
578
+
579
  def kernels_flash_attention_func(
580
  query_states: torch.Tensor,
581
  key_states: torch.Tensor,
 
584
  causal: bool = False,
585
  softmax_scale: float | None = None,
586
  implementation: str = "flash_attention_3",
587
+ padding_layout: FlashPaddingLayout | None = None,
588
  ) -> torch.Tensor:
589
  """Public flash-attention entry point with optional padding handling.
590
 
 
597
  before calling this function (ESM2, DPLM, DPLM2, E1, and ESMFold do), pass
598
  `softmax_scale=1.0`. Otherwise the flash kernel applies its default scale
599
  again, yielding an effective `1/head_dim` scale that drifts across layers.
600
+
601
+ `padding_layout` is the record `get_flash_padding_layout` built from
602
+ `attention_mask_2d`. Pass both so that a multi-layer encoder derives the
603
+ varlen metadata once per forward. The output is identical without it.
604
  """
605
  # query_states, key_states, value_states: (b, l, h, d)
606
  # attention_mask_2d: (b, l) or None
 
627
  value_states,
628
  attention_mask_2d,
629
  )
630
+ if padding_layout is not None:
631
+ _validate_flash_padding_layout(padding_layout, attention_mask_2d)
632
  _ensure_flash_kernels_loaded(implementation)
633
  if attention_mask_2d is not None:
634
  batch_size, q_len = query_states.shape[:2]
 
644
  key_states,
645
  value_states,
646
  attention_mask_2d,
647
+ padding_layout,
648
  )
649
  attn_output_unpad = _kernels_flash_varlen_forward( # (t, h, d)
650
  query_states=query_states,
 
658
  softmax_scale=softmax_scale,
659
  implementation=implementation,
660
  )
661
+ # pad_input scatters the real tokens into zeros, so padded positions are already zero.
662
+ return pad_input(attn_output_unpad, indices_q, batch_size, q_len) # (b, l, h, d)
663
  else:
664
  return _kernels_flash_forward( # (b, l, h, d)
665
  query_states=query_states,
 
874
  return attention_mask_2d, attention_mask_4d, None # (b, l), (b, 1, 1, l), None
875
 
876
 
877
+ def get_flash_padding_layout(
878
+ effective_backend: AttentionBackend,
879
+ attention_mask_2d: torch.Tensor | None,
880
+ ) -> FlashPaddingLayout | None:
881
+ """Build the FlashAttention varlen metadata once for all encoder layers.
882
+
883
+ `attention_mask_2d` is the bool mask that `get_attention_mask` returned.
884
+ Returns None when the call does not run a padded FlashAttention kernel.
885
+ """
886
+ # attention_mask_2d: (b, l) or None
887
+ if attention_mask_2d is None or not resolve_attention_backend(effective_backend).is_flash:
888
+ return None
889
+ return _flash_padding_layout(attention_mask_2d)
890
+
891
+
892
  def bool_to_additive_mask(
893
  bool_mask: torch.Tensor,
894
  dtype: torch.dtype,
fastplms/attention/interfaces.py CHANGED
@@ -8,6 +8,15 @@ from functools import partial
8
  from typing import Any
9
  from transformers import AttentionInterface, AttentionMaskInterface
10
 
 
 
 
 
 
 
 
 
 
11
  from ._core import (
12
  AttentionBackend,
13
  canonical_checkpoint_attention_backend,
@@ -95,6 +104,13 @@ class FastPLMsAttentionMixin:
95
  "sdpa",
96
  "flex_attention",
97
  )
 
 
 
 
 
 
 
98
 
99
  def _validate_attention_name(self, implementation: str) -> None:
100
  if implementation not in self._fastplms_attention_implementations:
@@ -147,11 +163,18 @@ class FastPLMsAttentionMixin:
147
  def __init__(self, config, *args: Any, **kwargs: Any) -> None:
148
  sentinel = object()
149
  internal = getattr(config, "_attn_implementation_internal", sentinel)
150
- stored = (
151
- getattr(config, "_attn_implementation", None) if internal is sentinel else internal
152
- )
153
  legacy = getattr(config, "attn_backend", None)
154
  requested = stored if stored is not None else legacy
 
 
 
 
 
 
 
 
 
155
  if requested is not None:
156
  if not isinstance(requested, str):
157
  raise TypeError(
@@ -180,6 +203,93 @@ class FastPLMsAttentionMixin:
180
  resolved = get_attn_implementation(config)
181
  self._validate_attention_name(resolved)
182
  set_config_attn_implementation(config, resolved)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
183
 
184
  def set_attn_implementation(
185
  self,
@@ -194,6 +304,23 @@ class FastPLMsAttentionMixin:
194
  raise ValueError(
195
  "FastPLMs models have one attention backbone; pass a string or {'': name}."
196
  )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
197
  resolved_name = self._check_and_adjust_attn_implementation(
198
  attn_implementation,
199
  is_init_check=False,
@@ -213,6 +340,17 @@ class FastPLMsAttentionMixin:
213
  )
214
 
215
 
 
 
 
 
 
 
 
 
 
 
 
216
  def validate_transformers_attention_interfaces() -> None:
217
  """Verify that Transformers exposes functions and masks for every backend.
218
 
 
8
  from typing import Any
9
  from transformers import AttentionInterface, AttentionMaskInterface
10
 
11
+ from ._auto import (
12
+ AUTO_ATTENTION,
13
+ AttentionResolution,
14
+ attention_execution_context,
15
+ deferred_resolution,
16
+ needs_execution_context,
17
+ provisional_implementation,
18
+ resolve_auto_attention,
19
+ )
20
  from ._core import (
21
  AttentionBackend,
22
  canonical_checkpoint_attention_backend,
 
104
  "sdpa",
105
  "flex_attention",
106
  )
107
+ # Preference order for ``attn_implementation="auto"``, mirrored from the
108
+ # family's ``attention_auto_order`` in models.toml. Empty rejects the request.
109
+ _fastplms_attention_auto_order: tuple[str, ...] = ()
110
+ # Supplied by the Transformers ``PreTrainedModel`` that follows this mixin in every
111
+ # MRO. Each family's configuration class adds the ``attn_backend`` field this mixin
112
+ # reads, and they share no base that declares it, so the boundary is dynamic.
113
+ config: Any
114
 
115
  def _validate_attention_name(self, implementation: str) -> None:
116
  if implementation not in self._fastplms_attention_implementations:
 
163
  def __init__(self, config, *args: Any, **kwargs: Any) -> None:
164
  sentinel = object()
165
  internal = getattr(config, "_attn_implementation_internal", sentinel)
166
+ stored = getattr(config, "_attn_implementation", None) if internal is sentinel else internal
 
 
167
  legacy = getattr(config, "attn_backend", None)
168
  requested = stored if stored is not None else legacy
169
+ auto_requested = requested == AUTO_ATTENTION
170
+ serialized_backend: str | None = None
171
+ if auto_requested:
172
+ # The configuration never holds ``auto``. Family layers are built on the
173
+ # provisional implementation, and a saved copy keeps the backend that a
174
+ # named load of the same checkpoint would have stored.
175
+ requested = provisional_implementation(self._require_attention_auto_order())
176
+ serialized_backend = legacy if legacy not in (None, AUTO_ATTENTION) else requested
177
+ stored = None
178
  if requested is not None:
179
  if not isinstance(requested, str):
180
  raise TypeError(
 
203
  resolved = get_attn_implementation(config)
204
  self._validate_attention_name(resolved)
205
  set_config_attn_implementation(config, resolved)
206
+ if auto_requested:
207
+ self.__dict__["_fastplms_serialized_attn_backend"] = serialized_backend
208
+ self._begin_auto_attention()
209
+
210
+ def _require_attention_auto_order(self) -> tuple[str, ...]:
211
+ order = self._fastplms_attention_auto_order
212
+ if not order:
213
+ raise ValueError(
214
+ f"{type(self).__name__} does not support attn_implementation='auto'; "
215
+ f"request one of {self._fastplms_attention_implementations}."
216
+ )
217
+ return order
218
+
219
+ @property
220
+ def attention_resolution(self) -> AttentionResolution | None:
221
+ """The record of an ``auto`` request, or None when a backend was named."""
222
+ return self.__dict__.get("_fastplms_attention_resolution")
223
+
224
+ def _begin_auto_attention(self) -> None:
225
+ """Resolve now when no candidate needs a device, else at the first forward."""
226
+ order = self._require_attention_auto_order()
227
+ self._cancel_pending_auto_attention()
228
+ if not needs_execution_context(order):
229
+ resolution = resolve_auto_attention(order, None)
230
+ self._apply_attn_implementation(resolution.resolved)
231
+ self.__dict__["_fastplms_attention_resolution"] = resolution
232
+ return
233
+ resolution = deferred_resolution(order)
234
+ self._apply_attn_implementation(resolution.resolved)
235
+ self.__dict__["_fastplms_attention_resolution"] = resolution
236
+ # The first forward runs inside the caller's autocast context, which is
237
+ # what decides FlashAttention eligibility for FP32 parameters.
238
+ self.__dict__["_fastplms_auto_attention_hook"] = (
239
+ self._as_module().register_forward_pre_hook(_resolve_auto_attention_before_forward)
240
+ )
241
+
242
+ def _as_module(self) -> torch.nn.Module:
243
+ if not isinstance(self, torch.nn.Module):
244
+ raise TypeError(
245
+ f"{type(self).__name__} must be a torch.nn.Module to defer attention selection."
246
+ )
247
+ return self
248
+
249
+ def _cancel_pending_auto_attention(self) -> None:
250
+ hook = self.__dict__.pop("_fastplms_auto_attention_hook", None)
251
+ if hook is not None:
252
+ hook.remove()
253
+
254
+ def resolve_attn_implementation(
255
+ self,
256
+ device: torch.device | str | None = None,
257
+ dtype: torch.dtype | None = None,
258
+ ) -> AttentionResolution:
259
+ """Settle a pending ``auto`` request for the device and dtype of the next forward.
260
+
261
+ The first forward does this by itself. Call it earlier, for example before
262
+ ``torch.compile`` or before fingerprinting an embedding run, and pass
263
+ ``dtype`` when the forward will run under an autocast context that is not
264
+ active yet. A settled request returns its record unchanged.
265
+ """
266
+ resolution = self.attention_resolution
267
+ if resolution is None:
268
+ raise RuntimeError(
269
+ f"{type(self).__name__} was not configured with attn_implementation='auto'."
270
+ )
271
+ if not resolution.deferred:
272
+ return resolution
273
+ self._cancel_pending_auto_attention()
274
+ resolution = resolve_auto_attention(
275
+ self._require_attention_auto_order(),
276
+ attention_execution_context(self._as_module(), device=device, dtype=dtype),
277
+ )
278
+ self._apply_attn_implementation(resolution.resolved)
279
+ self.__dict__["_fastplms_attention_resolution"] = resolution
280
+ return resolution
281
+
282
+ def save_pretrained(self, *args: Any, **kwargs: Any) -> Any:
283
+ """Save without the machine-specific outcome of an ``auto`` request."""
284
+ # ``save_pretrained`` comes from the ``PreTrainedModel`` later in the MRO.
285
+ if self.attention_resolution is None:
286
+ return super().save_pretrained(*args, **kwargs) # type: ignore[misc]
287
+ selected_backend = self.config.attn_backend
288
+ self.config.attn_backend = self.__dict__["_fastplms_serialized_attn_backend"]
289
+ try:
290
+ return super().save_pretrained(*args, **kwargs) # type: ignore[misc]
291
+ finally:
292
+ self.config.attn_backend = selected_backend
293
 
294
  def set_attn_implementation(
295
  self,
 
304
  raise ValueError(
305
  "FastPLMs models have one attention backbone; pass a string or {'': name}."
306
  )
307
+ if attn_implementation == AUTO_ATTENTION:
308
+ if allow_all_kernels:
309
+ raise ValueError("FastPLMs does not load external attention kernels.")
310
+ self.__dict__.setdefault(
311
+ "_fastplms_serialized_attn_backend", getattr(self.config, "attn_backend", None)
312
+ )
313
+ self._begin_auto_attention()
314
+ return
315
+ # A named request replaces any earlier automatic selection.
316
+ self._cancel_pending_auto_attention()
317
+ self.__dict__.pop("_fastplms_attention_resolution", None)
318
+ self.__dict__.pop("_fastplms_serialized_attn_backend", None)
319
+ self._apply_attn_implementation(attn_implementation, allow_all_kernels)
320
+
321
+ def _apply_attn_implementation(
322
+ self, attn_implementation: str, allow_all_kernels: bool = False
323
+ ) -> None:
324
  resolved_name = self._check_and_adjust_attn_implementation(
325
  attn_implementation,
326
  is_init_check=False,
 
340
  )
341
 
342
 
343
+ # Selection reads the manifest, can load a kernel, and rewrites layer attributes.
344
+ # It runs eagerly so that a compiled model never traces it.
345
+ @torch.compiler.disable # type: ignore[untyped-decorator]
346
+ def _resolve_auto_attention_before_forward(
347
+ module: torch.nn.Module, _arguments: tuple[Any, ...]
348
+ ) -> None:
349
+ if not isinstance(module, FastPLMsAttentionMixin):
350
+ raise TypeError("The automatic attention hook belongs on a FastPLMs model.")
351
+ module.resolve_attn_implementation()
352
+
353
+
354
  def validate_transformers_attention_interfaces() -> None:
355
  """Verify that Transformers exposes functions and masks for every backend.
356
 
fastplms/embeddings/batches.py ADDED
@@ -0,0 +1,378 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
fastplms/embeddings/identity.py ADDED
@@ -0,0 +1,510 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
fastplms/embeddings/inputs.py ADDED
@@ -0,0 +1,263 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Normalize ordered inputs and plan bounded windows without retaining a full stream."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import hashlib
6
+ import sqlite3
7
+ import tempfile
8
+ 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
fastplms/embeddings/output.py ADDED
@@ -0,0 +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
fastplms/embeddings/runner.py CHANGED
@@ -1,1583 +1,425 @@
1
- """Model-independent dataset embedding orchestration."""
2
-
3
- from __future__ import annotations
4
-
5
- import hashlib
6
- import json
7
- import platform
8
- import sqlite3
9
- import tempfile
10
- import torch
11
- from collections.abc import Callable, Iterable, Iterator, Mapping, Sequence
12
- from contextlib import contextmanager
13
- from pathlib import Path
14
- from typing import Any, overload
15
  from torch import Tensor
16
 
17
- from .pooling import Pooler
18
- from .storage import (
19
- SafetensorsStreamWriter,
20
- append_sqlite_records,
21
- initialize_sqlite_run,
22
- load_result,
23
- load_sqlite_result,
24
- safetensors_result_exists,
25
- save_result,
26
- tensor_sha256,
27
- update_sqlite_run_metadata,
28
- )
29
- from .types import (
30
- EmbeddingBatch,
31
- EmbeddingInput,
32
- EmbeddingRecord,
33
- EmbeddingResult,
34
- LazyTensorReference,
35
- )
36
-
37
-
38
- _MAX_PARTI_RESIDUES = 2_048
39
- _RUN_FINGERPRINT_SCHEMA_VERSION = 3
40
- _MODEL_STATE_HASH_CHUNK_BYTES = 16 * 1024**2
41
- _DEFAULT_BATCH_WINDOW_MULTIPLIER = 16
42
- _SUPPORTED_STORAGE_FORMATS = frozenset({"safetensors", "sqlite"})
43
-
44
-
45
- def _validate_parti_length(M: Tensor) -> None:
46
- """Reject an oversized attention graph before model inference."""
47
-
48
- # M: (b, l)
49
- n_residues = int(M.to(dtype=torch.int64).sum(dim=1).max().item())
50
- if n_residues > _MAX_PARTI_RESIDUES:
51
- raise ValueError(f"parti supports at most {_MAX_PARTI_RESIDUES:,} biological residues.")
52
-
53
-
54
- def select_hidden_state_embeddings(
55
- last_hidden_state: Tensor,
56
- hidden_states: tuple[Tensor, ...] | None,
57
- *,
58
- hidden_state_index: int = -1,
59
- store_all_hidden_states: bool = False,
60
- ) -> Tensor:
61
- """Select one hidden state or stack every state without changing values."""
62
- # last_hidden_state and each hidden_states entry: (b, l, d)
63
- if store_all_hidden_states:
64
- if not hidden_states:
65
- raise ValueError("store_all_hidden_states requires model hidden states.")
66
- # H has shape (b, n, l, d), where n follows the model's output order.
67
- return torch.stack(hidden_states, dim=1) # (b, n, l, d)
68
- if hidden_state_index == -1:
69
- return last_hidden_state # (b, l, d)
70
- if not hidden_states:
71
- raise ValueError("hidden_state_index requires model hidden states.")
72
- return hidden_states[hidden_state_index] # (b, l, d)
73
-
74
-
75
- def iter_fasta(path: str | Path) -> Iterator[EmbeddingInput]:
76
- """Yield FASTA records in source order without reading the file into memory."""
77
-
78
- identifier: str | None = None
79
- sequence_parts: list[str] = []
80
- found_record = False
81
- with Path(path).open("r", encoding="utf-8") as handle:
82
- for line_number, raw_line in enumerate(handle, start=1):
83
- line = raw_line.strip()
84
- if not line:
85
- continue
86
- if line.startswith(">"):
87
- if identifier is not None:
88
- found_record = True
89
- yield EmbeddingInput(identifier, "".join(sequence_parts))
90
- identifier = line[1:].strip().split(maxsplit=1)[0]
91
- if not identifier:
92
- raise ValueError(f"Missing FASTA identifier on line {line_number}.")
93
- sequence_parts = []
94
- else:
95
- if identifier is None:
96
- raise ValueError(
97
- f"Sequence data precedes the first FASTA header on line {line_number}."
98
- )
99
- sequence_parts.append("".join(line.split()))
100
- if identifier is not None:
101
- found_record = True
102
- yield EmbeddingInput(identifier, "".join(sequence_parts))
103
- if not found_record:
104
- raise ValueError(f"No FASTA records found in {path}.")
105
-
106
-
107
- def parse_fasta(path: str | Path) -> list[EmbeddingInput]:
108
- """Parse FASTA records while preserving identifiers, order, and duplicates."""
109
-
110
- return list(iter_fasta(path))
111
-
112
-
113
- def _normalize_input_item(
114
- position: int,
115
- item: str | EmbeddingInput | tuple[str, str],
116
- ) -> EmbeddingInput:
117
- if isinstance(item, EmbeddingInput):
118
- return item
119
- if isinstance(item, str):
120
- return EmbeddingInput(str(position), item)
121
- if isinstance(item, tuple) and len(item) == 2:
122
- return EmbeddingInput(str(item[0]), str(item[1]))
123
- raise TypeError(
124
- "inputs must contain sequences, EmbeddingInput values, or (id, sequence) tuples."
125
- )
126
-
127
-
128
- class _InputSpool(Sequence[EmbeddingInput]):
129
- """Immutable disk-backed normalized inputs with an incremental digest."""
130
-
131
- def __init__(
132
- self,
133
- values: Iterable[str | EmbeddingInput | tuple[str, str]],
134
- ) -> None:
135
- self._temporary: tempfile.TemporaryDirectory[str] | None = tempfile.TemporaryDirectory(
136
- prefix="fastplms-inputs-"
137
- )
138
- self.path = Path(self._temporary.name) / "inputs.sqlite"
139
- self._connection: sqlite3.Connection | None = sqlite3.connect(self.path)
140
- self._connection.execute(
141
- "CREATE TABLE inputs ("
142
- "position INTEGER PRIMARY KEY, input_id TEXT NOT NULL, sequence TEXT NOT NULL)"
143
- )
144
- digest = hashlib.sha256()
145
- count = 0
146
- pending: list[tuple[int, str, str]] = []
147
- try:
148
- for position, item in enumerate(values):
149
- record = _normalize_input_item(position, item)
150
- for value in (record.id, record.sequence):
151
- encoded = value.encode("utf-8")
152
- digest.update(len(encoded).to_bytes(8, "big"))
153
- digest.update(encoded)
154
- pending.append((position, record.id, record.sequence))
155
- count += 1
156
- if len(pending) == 1_024:
157
- self._connection.executemany("INSERT INTO inputs VALUES (?, ?, ?)", pending)
158
- pending.clear()
159
- if pending:
160
- self._connection.executemany("INSERT INTO inputs VALUES (?, ?, ?)", pending)
161
- if count == 0:
162
- raise ValueError("inputs must contain at least one sequence.")
163
- self._connection.commit()
164
- self._connection.close()
165
- self._connection = sqlite3.connect(
166
- f"{self.path.resolve().as_uri()}?mode=ro",
167
- uri=True,
168
- )
169
- except BaseException:
170
- self.close()
171
- raise
172
- digest.update(count.to_bytes(8, "big"))
173
- self.input_fingerprint = digest.hexdigest()
174
- self._count = count
175
-
176
- def _require_connection(self) -> sqlite3.Connection:
177
- if self._connection is None:
178
- raise RuntimeError("Input spool is closed.")
179
- return self._connection
180
-
181
- def __len__(self) -> int:
182
- return self._count
183
-
184
- def __iter__(self) -> Iterator[EmbeddingInput]:
185
- cursor = self._require_connection().execute(
186
- "SELECT input_id, sequence FROM inputs ORDER BY position"
187
- )
188
- while rows := cursor.fetchmany(1_024):
189
- for input_id, sequence in rows:
190
- yield EmbeddingInput(input_id, sequence)
191
-
192
- @overload
193
- def __getitem__(self, index: int, /) -> EmbeddingInput: ...
194
-
195
- @overload
196
- def __getitem__(self, index: slice, /) -> list[EmbeddingInput]: ...
197
-
198
- def __getitem__(self, index: int | slice) -> EmbeddingInput | list[EmbeddingInput]:
199
- connection = self._require_connection()
200
-
201
- if isinstance(index, slice):
202
- start, stop, step = index.indices(self._count)
203
- if step != 1:
204
- return [self[position] for position in range(start, stop, step)]
205
- rows = connection.execute(
206
- "SELECT input_id, sequence FROM inputs "
207
- "WHERE position >= ? AND position < ? ORDER BY position",
208
- (start, stop),
209
- ).fetchall()
210
- return [EmbeddingInput(input_id, sequence) for input_id, sequence in rows]
211
- position = index + self._count if index < 0 else index
212
- if position < 0 or position >= self._count:
213
- raise IndexError(index)
214
- row = connection.execute(
215
- "SELECT input_id, sequence FROM inputs WHERE position = ?", (position,)
216
- ).fetchone()
217
- if row is None:
218
- raise IndexError(index)
219
- return EmbeddingInput(row[0], row[1])
220
-
221
- def close(self) -> None:
222
- connection = getattr(self, "_connection", None)
223
- if connection is not None:
224
- connection.close()
225
- self._connection = None
226
- temporary = getattr(self, "_temporary", None)
227
- if temporary is not None:
228
- temporary.cleanup()
229
- self._temporary = None
230
-
231
- def __del__(self) -> None:
232
- self.close()
233
-
234
-
235
- def _normalize_inputs(
236
- inputs: (Iterable[str | EmbeddingInput | tuple[str, str]] | Mapping[str, str] | str | Path),
237
- *,
238
- disk_backed: bool,
239
- ) -> Sequence[EmbeddingInput]:
240
- is_fasta_path = isinstance(inputs, Path)
241
- if isinstance(inputs, str):
242
- try:
243
- is_fasta_path = Path(inputs).is_file()
244
- except OSError:
245
- is_fasta_path = False
246
- should_spool = disk_backed or is_fasta_path or not isinstance(inputs, (str, Sequence, Mapping))
247
- values: Iterable[str | EmbeddingInput | tuple[str, str]]
248
- if isinstance(inputs, Path):
249
- values = iter_fasta(inputs)
250
- elif isinstance(inputs, str):
251
- values = iter_fasta(inputs) if is_fasta_path else [inputs]
252
- elif isinstance(inputs, Mapping):
253
- values = inputs.items()
254
- else:
255
- values = inputs
256
- if should_spool:
257
- return _InputSpool(values)
258
- records: list[EmbeddingInput] = []
259
- for position, item in enumerate(values):
260
- records.append(_normalize_input_item(position, item))
261
- if not records:
262
- raise ValueError("inputs must contain at least one sequence.")
263
- return records
264
-
265
-
266
- def _validate_untruncated_lengths(
267
- records: Sequence[EmbeddingInput],
268
- *,
269
- max_length: int | None,
270
- truncate: bool,
271
- ) -> None:
272
- """Fail before inference when a biological-residue limit would be exceeded."""
273
-
274
- if max_length is None or truncate:
275
- return
276
- for position, record in enumerate(records):
277
- residue_count = len(record.sequence)
278
- if residue_count > max_length:
279
- raise ValueError(
280
- f"Input at position {position} with id {record.id!r} has "
281
- f"{residue_count} biological residues, exceeding max_length={max_length} "
282
- "while truncate=False."
283
- )
284
-
285
-
286
- def _model_device(model: Any) -> torch.device:
287
- try:
288
- return torch.device(next(model.parameters()).device)
289
- except (AttributeError, StopIteration):
290
- return torch.device("cpu")
291
-
292
-
293
- def _attention_backend(model: Any) -> str | None:
294
- config = getattr(model, "config", None)
295
- for name in ("_attn_implementation", "attn_implementation", "attn_backend"):
296
- value = getattr(config, name, None)
297
- if value:
298
- return str(value)
299
- return None
300
-
301
-
302
- def _attention_kernel_metadata(backend: str | None) -> dict[str, Any] | None:
303
- if backend not in {"flash_attention_2", "flash_attention_3"}:
304
- return None
305
- from fastplms.registry import get_model_registry
306
-
307
- spec = get_model_registry().attention_kernels[backend]
308
- return {
309
- "repository": spec.repository,
310
- "revision": spec.revision,
311
- "version": spec.version,
312
- "expected_variant": spec.expected_variant,
313
- "dtypes": list(spec.dtypes),
314
- }
315
-
316
-
317
- def _fingerprint_jsonable(value: Any) -> Any:
318
- if isinstance(value, Mapping):
319
- return {str(key): _fingerprint_jsonable(item) for key, item in value.items()}
320
- if isinstance(value, (list, tuple)):
321
- return [_fingerprint_jsonable(item) for item in value]
322
- if isinstance(value, (set, frozenset)):
323
- return sorted((_fingerprint_jsonable(item) for item in value), key=repr)
324
- if isinstance(value, Path):
325
- return str(value)
326
- if isinstance(value, Tensor):
327
- return {
328
- "dtype": str(value.dtype).removeprefix("torch."),
329
- "shape": list(value.shape),
330
- "sha256": tensor_sha256(value),
331
- }
332
- if isinstance(value, torch.dtype):
333
- return str(value).removeprefix("torch.")
334
- if isinstance(value, torch.device):
335
- return str(value)
336
- if value is None or isinstance(value, (str, int, float, bool)):
337
- return value
338
- return {
339
- "class": f"{value.__class__.__module__}.{value.__class__.__qualname__}",
340
- "value": str(value),
341
- }
342
-
343
-
344
- def _tokenizer_content_sha256(tokenizer: Any) -> str:
345
- content: dict[str, Any] = {
346
- "init_kwargs": getattr(tokenizer, "init_kwargs", None),
347
- "special_tokens_map": getattr(tokenizer, "special_tokens_map", None),
348
- "model_max_length": getattr(tokenizer, "model_max_length", None),
349
- "padding_side": getattr(tokenizer, "padding_side", None),
350
- "truncation_side": getattr(tokenizer, "truncation_side", None),
351
- }
352
- get_vocab = getattr(tokenizer, "get_vocab", None)
353
- if callable(get_vocab):
354
- content["vocabulary"] = get_vocab()
355
- get_added_vocab = getattr(tokenizer, "get_added_vocab", None)
356
- if callable(get_added_vocab):
357
- content["added_vocabulary"] = get_added_vocab()
358
- backend = getattr(tokenizer, "backend_tokenizer", None)
359
- backend_to_str = getattr(backend, "to_str", None)
360
- if callable(backend_to_str):
361
- content["backend"] = backend_to_str()
362
- serialized = json.dumps(
363
- _fingerprint_jsonable(content),
364
- sort_keys=True,
365
- separators=(",", ":"),
366
- ensure_ascii=False,
367
- ).encode()
368
- return hashlib.sha256(serialized).hexdigest()
369
-
370
-
371
- def _tokenizer_metadata(model: Any, tokenizer: Any | None) -> dict[str, Any]:
372
- resolved = tokenizer if tokenizer is not None else getattr(model, "tokenizer", None)
373
- if resolved is None:
374
- # Raw-sequence families such as E1 retain their loader context on the
375
- # model/encoder rather than exposing a Transformers tokenizer. Bind the
376
- # non-secret source policy to resume identity without serializing a Hub
377
- # token or forcing lazy tokenizer initialization.
378
- for candidate in (model, getattr(model, "model", None)):
379
- settings = getattr(candidate, "__dict__", {}).get("_fastplms_tokenizer_kwargs")
380
- if isinstance(settings, Mapping):
381
- token_value = settings.get("token")
382
- return {
383
- "mode": "native-sequence",
384
- "source": (
385
- str(settings.get("tokenizer_source"))
386
- if settings.get("tokenizer_source") is not None
387
- else None
388
- ),
389
- "revision": settings.get("revision"),
390
- "cache_dir": (
391
- str(settings.get("cache_dir"))
392
- if settings.get("cache_dir") is not None
393
- else None
394
- ),
395
- "local_files_only": bool(settings.get("local_files_only", False)),
396
- "token_policy": (
397
- "disabled"
398
- if token_value is False
399
- else "provided"
400
- if token_value is not None
401
- else "default"
402
- ),
403
- }
404
- return {"mode": "native-sequence"}
405
- return {
406
- "mode": "tokenizer",
407
- "class": f"{resolved.__class__.__module__}.{resolved.__class__.__qualname__}",
408
- "name_or_path": getattr(resolved, "name_or_path", None),
409
- "vocab_size": getattr(resolved, "vocab_size", None),
410
- "special_token_ids": list(getattr(resolved, "all_special_ids", ())),
411
- "content_sha256": _tokenizer_content_sha256(resolved),
412
- }
413
-
414
-
415
- @contextmanager
416
- def _temporary_eval(model: Any) -> Iterator[None]:
417
- was_training = getattr(model, "training", None)
418
- eval_method = getattr(model, "eval", None)
419
- train_method = getattr(model, "train", None)
420
- if (
421
- not isinstance(was_training, bool)
422
- or not callable(eval_method)
423
- or not callable(train_method)
424
- ):
425
- yield
426
- return
427
- eval_method()
428
- try:
429
- yield
430
- finally:
431
- train_method(was_training)
432
-
433
-
434
- def _software_versions() -> dict[str, str | None]:
435
- try:
436
- import fastplms
437
-
438
- fastplms_version = fastplms.__version__
439
- except (AttributeError, ImportError):
440
- fastplms_version = None
441
- try:
442
- import safetensors
443
-
444
- safetensors_version = safetensors.__version__
445
- except ImportError:
446
- safetensors_version = None
447
- try:
448
- import transformers
449
-
450
- transformers_version = transformers.__version__
451
- except ImportError:
452
- transformers_version = None
453
- return {
454
- "fastplms": fastplms_version,
455
- "python": platform.python_version(),
456
- "safetensors": safetensors_version,
457
- "torch": torch.__version__,
458
- "torch_cuda": torch.version.cuda,
459
- "transformers": transformers_version,
460
- }
461
-
462
-
463
- def _adapter_identity_metadata(model: Any) -> dict[str, Any] | None:
464
- """Return deterministic PEFT/adapter identity without tensor payloads."""
465
-
466
- peft_config = getattr(model, "peft_config", None)
467
- if not isinstance(peft_config, Mapping) or not peft_config:
468
- return None
469
- configurations: dict[str, Any] = {}
470
- for name, config in sorted(peft_config.items(), key=lambda item: str(item[0])):
471
- to_dict = getattr(config, "to_dict", None)
472
- if callable(to_dict):
473
- value = to_dict()
474
- else:
475
- try:
476
- value = vars(config)
477
- except TypeError:
478
- value = config
479
- configurations[str(name)] = _fingerprint_jsonable(value)
480
- active_adapters = getattr(model, "active_adapters", None)
481
- if callable(active_adapters):
482
- active_adapters = active_adapters()
483
- return {
484
- "active": _fingerprint_jsonable(active_adapters),
485
- "configurations": configurations,
486
- }
487
-
488
-
489
- def _execution_identity_metadata(model: Any) -> dict[str, Any]:
490
- """Capture runtime policy that can change persisted numerical results."""
491
-
492
- parameter_dtypes = sorted(
493
- {
494
- str(parameter.dtype).removeprefix("torch.")
495
- for parameter in getattr(model, "parameters", lambda: ())()
496
- }
497
- )
498
- return {
499
- "device": _model_device(model).type,
500
- "hf_device_map": _fingerprint_jsonable(getattr(model, "hf_device_map", None)),
501
- "parameter_dtypes": parameter_dtypes,
502
- "software": _software_versions(),
503
- }
504
-
505
-
506
- def _biological_residue_mask(
507
- input_ids: Tensor,
508
- attention_mask: Tensor,
509
- tokenizer: Any,
510
- ) -> Tensor:
511
- """Remove padding and tokenizer-declared special tokens from M."""
512
-
513
- # input_ids, attention_mask: (b, l)
514
- M = attention_mask.to(dtype=torch.bool) # (b, l)
515
- special_ids = tuple(int(token_id) for token_id in getattr(tokenizer, "all_special_ids", ()))
516
- if special_ids:
517
- specials = torch.tensor( # (n_special,)
518
- special_ids,
519
- device=input_ids.device,
520
- dtype=input_ids.dtype,
521
- )
522
- M = M & ~torch.isin(input_ids, specials) # (b, l)
523
- return M # (b, l)
524
-
525
-
526
- def _generic_embedding_batch(
527
- model: Any,
528
- sequences: list[str],
529
- *,
530
- tokenizer: Any | None,
531
- max_length: int | None,
532
- truncate: bool,
533
- need_attentions: bool,
534
- model_kwargs: dict[str, Any],
535
- ) -> EmbeddingBatch:
536
- config = getattr(model, "config", None)
537
- model_type = str(getattr(config, "model_type", "")).lower()
538
- if tokenizer is None:
539
- tokenizer = getattr(model, "tokenizer", None)
540
-
541
- if tokenizer is None and model_type == "e1":
542
- output = model._embed(sequences, return_attention_mask=True, **model_kwargs)
543
- if not isinstance(output, tuple) or len(output) != 2:
544
- raise TypeError("E1 _embed must return (X, residue_mask).")
545
- X, M = output # (b, l, d), (b, l)
546
- preparer = getattr(model, "prep_tokens", None)
547
- if preparer is not None and hasattr(preparer, "get_batch_kwargs"):
548
- prepared = preparer.get_batch_kwargs(sequences, device=X.device)
549
- input_ids = prepared["input_ids"] # (b, l)
550
- boundary_ids = preparer.boundary_token_ids.to( # (n_boundary,)
551
- device=input_ids.device, dtype=input_ids.dtype
552
- )
553
- # E1 wraps each raw sequence in BOS, context-label, terminal-label,
554
- # and EOS tokens. Only amino-acid rows are biological residues.
555
- M = M.to(dtype=torch.bool) & ~torch.isin(input_ids, boundary_ids) # (b, l)
556
- if need_attentions:
557
- raise ValueError("parti is not available for tokenizer-free E1 embedding.")
558
- return EmbeddingBatch( # X: (b, l, d); residue_mask: (b, l)
559
- X=X,
560
- residue_mask=M.to(dtype=torch.bool),
561
- )
562
- if tokenizer is None:
563
- raise ValueError("A tokenizer is required for this model's embedding path.")
564
-
565
- tokenize_kwargs: dict[str, Any] = {
566
- "return_tensors": "pt",
567
- "padding": True,
568
- "truncation": truncate,
569
- }
570
- if max_length is not None and truncate:
571
- # ``max_length`` is a biological-residue limit. Tokenizer limits include
572
- # boundary tokens, so reserve their declared width instead of dropping
573
- # residues at the exact boundary.
574
- special_token_count = 0
575
- num_special_tokens_to_add = getattr(tokenizer, "num_special_tokens_to_add", None)
576
- if callable(num_special_tokens_to_add):
577
- special_token_count = int(num_special_tokens_to_add(pair=False))
578
- tokenize_kwargs["max_length"] = max_length + special_token_count
579
- sequence_tokenizer = getattr(model, "_tokenize_sequence_batch", None)
580
- if callable(sequence_tokenizer):
581
- encoded = sequence_tokenizer(sequences, tokenizer=tokenizer, **tokenize_kwargs)
582
- else:
583
- encoded = tokenizer(sequences, **tokenize_kwargs)
584
- device = _model_device(model)
585
- input_ids = encoded["input_ids"].to(device) # (b, l)
586
- attention_mask = encoded.get( # (b, l)
587
- "attention_mask",
588
- input_ids.new_ones(input_ids.shape),
589
- ).to(device)
590
- M = _biological_residue_mask(input_ids, attention_mask, tokenizer) # (b, l)
591
- if need_attentions:
592
- # Validate l before either the backbone or its quadratic attention graph
593
- # is materialized. M has shape (b, l).
594
- _validate_parti_length(M)
595
- X = model._embed(input_ids, attention_mask, **model_kwargs) # (b, l, d)
596
- attentions = None
597
- if need_attentions:
598
- output = model(
599
- input_ids=input_ids,
600
- attention_mask=attention_mask,
601
- output_attentions=True,
602
- return_dict=True,
603
- )
604
- attentions = getattr(output, "attentions", None) # each: (b, h, l, l)
605
- if attentions is None:
606
- raise ValueError("The model did not return attentions required by parti.")
607
- return EmbeddingBatch( # X: (b, l, d); M: (b, l)
608
- X=X,
609
- residue_mask=M,
610
- attentions=attentions,
611
- )
612
-
613
-
614
- def _first_metadata_value(*values: Any) -> Any:
615
- for value in values:
616
- if isinstance(value, str):
617
- if value.strip():
618
- return value
619
- elif value is not None:
620
- return value
621
- return None
622
-
623
-
624
- def _model_identity_metadata(model: Any) -> dict[str, Any]:
625
- """Resolve model and checkpoint identity, including local artifact fallbacks."""
626
-
627
- config = getattr(model, "config", None)
628
- checkpoint_revision = _first_metadata_value(
629
- getattr(config, "fastplms_checkpoint_revision", None),
630
- getattr(config, "_commit_hash", None),
631
- )
632
- return {
633
- "model_id": _first_metadata_value(
634
- getattr(config, "fastplms_model_id", None),
635
- getattr(config, "_name_or_path", None),
636
- ),
637
- "model_revision": _first_metadata_value(
638
- getattr(config, "_commit_hash", None),
639
- checkpoint_revision,
640
- ),
641
- "checkpoint_repo_id": getattr(config, "fastplms_checkpoint_repo_id", None),
642
- "checkpoint_revision": checkpoint_revision,
643
- "checkpoint_hash": _first_metadata_value(
644
- getattr(model, "checkpoint_hash", None),
645
- getattr(config, "checkpoint_hash", None),
646
- getattr(config, "fastplms_checkpoint_hash", None),
647
- ),
648
- "weights_revision": getattr(config, "fastplms_weights_revision", None),
649
- "runtime_revision": getattr(config, "fastplms_runtime_revision", None),
650
- "source_tree_sha256": getattr(config, "fastplms_source_tree_sha256", None),
651
- "runtime_bundle_sha256": getattr(config, "fastplms_runtime_bundle_sha256", None),
652
- }
653
-
654
-
655
- def _bounded_tensor_chunks(X: Tensor, max_elements: int) -> Iterable[Tensor]:
656
- """Yield X in logical row-major order without materializing a full copy."""
657
-
658
- # X: (...)
659
- if X.numel() == 0:
660
- return
661
- if X.ndim == 0:
662
- yield X
663
- return
664
- trailing_elements = 1
665
- for size in X.shape[1:]:
666
- trailing_elements *= int(size)
667
- if trailing_elements <= max_elements:
668
- rows_per_chunk = max(1, max_elements // trailing_elements)
669
- for start in range(0, X.shape[0], rows_per_chunk):
670
- yield X[start : start + rows_per_chunk] # (chunk_rows, ...)
671
- return
672
- for row in X:
673
- yield from _bounded_tensor_chunks(row, max_elements)
674
-
675
-
676
- def _model_state_sha256(model: Any) -> str:
677
- """Hash named parameters and persistent buffers using bounded CPU copies."""
678
-
679
- # Never cache this digest from tensor identity or ``Tensor._version``.
680
- # ``Parameter.data`` and independent tensor aliases can mutate shared storage
681
- # without changing either signal, while persisted resume identity must bind
682
- # the authoritative bytes visible at the start of this run.
683
- state = model.state_dict(keep_vars=True)
684
- digest = hashlib.sha256()
685
- for name, value in sorted(state.items()):
686
- if not isinstance(value, Tensor):
687
- raise TypeError(f"Model state entry {name!r} is not a tensor.")
688
- if value.is_meta:
689
- raise ValueError(
690
- f"Cannot fingerprint meta-device model state entry {name!r}; pass "
691
- "model_state_fingerprint with a caller-owned state identity."
692
- )
693
- header = json.dumps(
694
- {
695
- "name": name,
696
- "dtype": str(value.dtype).removeprefix("torch."),
697
- "shape": list(value.shape),
698
- },
699
- sort_keys=True,
700
- separators=(",", ":"),
701
- ).encode()
702
- digest.update(len(header).to_bytes(8, "big"))
703
- digest.update(header)
704
- max_elements = max(1, _MODEL_STATE_HASH_CHUNK_BYTES // value.element_size())
705
- for chunk in _bounded_tensor_chunks(value.detach(), max_elements):
706
- cpu_chunk = chunk.to(device="cpu").contiguous() # chunk.shape
707
- digest.update(cpu_chunk.reshape(-1).view(torch.uint8).numpy().tobytes())
708
- return digest.hexdigest()
709
-
710
-
711
- def _input_sha256(records: Iterable[EmbeddingInput]) -> str:
712
- """Hash an ordered input stream without constructing a duplicate JSON payload."""
713
-
714
- precomputed = getattr(records, "input_fingerprint", None)
715
- if isinstance(precomputed, str):
716
- return precomputed
717
- digest = hashlib.sha256()
718
- count = 0
719
- for record in records:
720
- count += 1
721
- for value in (record.id, record.sequence):
722
- encoded = value.encode("utf-8")
723
- digest.update(len(encoded).to_bytes(8, "big"))
724
- digest.update(encoded)
725
- digest.update(count.to_bytes(8, "big"))
726
- return digest.hexdigest()
727
-
728
-
729
- def _run_fingerprint(
730
- model: Any,
731
- records: Sequence[EmbeddingInput],
732
- *,
733
- pooling: Sequence[str],
734
- full_embeddings: bool,
735
- max_length: int | None,
736
- truncate: bool,
737
- dtype: torch.dtype | None,
738
- model_kwargs: dict[str, Any],
739
- tokenizer_metadata: dict[str, Any],
740
- model_state_fingerprint: str | None,
741
- persist_output: bool,
742
- embedding_context: Mapping[str, Any],
743
- batch_size: int,
744
- batch_window_size: int,
745
- max_tokens_per_batch: int | None,
746
- ) -> tuple[str, str, str | None, str]:
747
- input_fingerprint = _input_sha256(records)
748
- attention_backend = _attention_backend(model)
749
- model_identity = _model_identity_metadata(model)
750
- if model_state_fingerprint is None and persist_output:
751
- resolved_model_state_fingerprint = _model_state_sha256(model)
752
- model_state_fingerprint_source = "computed"
753
- elif model_state_fingerprint is not None:
754
- resolved_model_state_fingerprint = model_state_fingerprint.strip()
755
- if not resolved_model_state_fingerprint:
756
- raise ValueError("model_state_fingerprint must not be empty.")
757
- model_state_fingerprint_source = "caller"
758
- else:
759
- resolved_model_state_fingerprint = None
760
- model_state_fingerprint_source = "not-computed"
761
- payload = {
762
- "fingerprint_schema_version": _RUN_FINGERPRINT_SCHEMA_VERSION,
763
- "input_fingerprint": input_fingerprint,
764
- "model_state_fingerprint": resolved_model_state_fingerprint,
765
- "model_state_fingerprint_source": model_state_fingerprint_source,
766
- "model_class": f"{model.__class__.__module__}.{model.__class__.__qualname__}",
767
- **model_identity,
768
- "attention_backend": attention_backend,
769
- "attention_kernel": _attention_kernel_metadata(attention_backend),
770
- "layer": repr(
771
- getattr(model, "embedding_layer", model_kwargs.get("hidden_state_index", -1))
772
- ),
773
- "projection": getattr(model, "embedding_projection", None),
774
- "esmc_source": getattr(model, "_esmc_source", None),
775
- "esmc_revision": getattr(model, "_esmc_source_revision", None),
776
- "esmc_files": getattr(model, "_esmc_source_files", None),
777
- "token_policy": getattr(model, "embedding_token_policy", None),
778
- "tokenizer": tokenizer_metadata,
779
- "adapter": _adapter_identity_metadata(model),
780
- "execution": _execution_identity_metadata(model),
781
- "embedding_context": _fingerprint_jsonable(embedding_context),
782
- "pooling": list(pooling),
783
- "full_embeddings": full_embeddings,
784
- "max_length": max_length,
785
- "truncate": truncate,
786
- "dtype": str(dtype) if dtype is not None else None,
787
- "batching": {
788
- "batch_size": batch_size,
789
- "batch_window_size": batch_window_size,
790
- "max_tokens_per_batch": max_tokens_per_batch,
791
- "input_storage": ("disk-spool" if isinstance(records, _InputSpool) else "memory"),
792
- },
793
- "model_kwargs": {
794
- key: _fingerprint_jsonable(value) for key, value in sorted(model_kwargs.items())
795
- },
796
- "residue_mask_policy": "attention-mask-minus-special-tokens",
797
- }
798
- run_fingerprint = hashlib.sha256(
799
- json.dumps(payload, sort_keys=True, separators=(",", ":")).encode()
800
- ).hexdigest()
801
- return (
802
- input_fingerprint,
803
- run_fingerprint,
804
- resolved_model_state_fingerprint,
805
- model_state_fingerprint_source,
806
- )
807
-
808
-
809
- def _output_exists(path: str | Path, format: str) -> bool:
810
- path = Path(path)
811
- if format == "sqlite":
812
- return path.is_file()
813
- return safetensors_result_exists(path)
814
-
815
-
816
- def _output_descriptor(position: int, record: EmbeddingRecord) -> dict[str, Any]:
817
- tensor = record.tensor
818
- if isinstance(tensor, LazyTensorReference):
819
- dtype = tensor.dtype
820
- shape = tensor.shape
821
- digest = tensor.sha256
822
- else:
823
- dtype = str(tensor.dtype).removeprefix("torch.")
824
- shape = tuple(tensor.shape)
825
- digest = tensor_sha256(tensor)
826
- return {
827
- "position": position,
828
- "id": record.id,
829
- "dtype": dtype,
830
- "shape": shape,
831
- "sha256": digest,
832
- }
833
-
834
-
835
- def _ordered_string_sha256(values: Sequence[str]) -> str:
836
- digest = hashlib.sha256()
837
- for value in values:
838
- encoded = value.encode("utf-8")
839
- digest.update(len(encoded).to_bytes(8, "big"))
840
- digest.update(encoded)
841
- digest.update(len(values).to_bytes(8, "big"))
842
- return digest.hexdigest()
843
-
844
-
845
- def _embedding_context(
846
- model: Any,
847
- records: Sequence[EmbeddingInput],
848
- *,
849
- hidden_state_source: str,
850
- decoder_inputs: Sequence[str] | None,
851
- decoder_input_ids: Tensor | None,
852
- decoder_attention_mask: Tensor | None,
853
- model_kwargs: Mapping[str, Any],
854
- ) -> tuple[dict[str, Any], tuple[str, ...] | None]:
855
- if hidden_state_source not in {"encoder", "decoder"}:
856
- raise ValueError("hidden_state_source must be 'encoder' or 'decoder'.")
857
- hidden_state_index = model_kwargs.get("hidden_state_index", -1)
858
- if not isinstance(hidden_state_index, int) or isinstance(hidden_state_index, bool):
859
- raise TypeError("hidden_state_index must be an integer.")
860
- store_all_hidden_states = model_kwargs.get("store_all_hidden_states", False)
861
- if not isinstance(store_all_hidden_states, bool):
862
- raise TypeError("store_all_hidden_states must be a boolean.")
863
- normalized_decoder_inputs: tuple[str, ...] | None = None
864
- has_decoder_inputs = decoder_inputs is not None
865
- has_decoder_ids = decoder_input_ids is not None
866
- if hidden_state_source == "encoder":
867
- if has_decoder_inputs or has_decoder_ids or decoder_attention_mask is not None:
868
- raise ValueError("Decoder inputs are only valid when hidden_state_source='decoder'.")
869
- else:
870
- if has_decoder_inputs == has_decoder_ids:
871
- raise ValueError(
872
- "Decoder embedding requires exactly one of decoder_inputs or decoder_input_ids."
873
- )
874
- decoder_input_fingerprint: str | None = None
875
- if decoder_inputs is not None:
876
- if isinstance(decoder_inputs, (str, bytes)) or not isinstance(decoder_inputs, Sequence):
877
- raise TypeError("decoder_inputs must be an aligned sequence of strings.")
878
- normalized_decoder_inputs = tuple(decoder_inputs)
879
- if not all(isinstance(value, str) and value for value in normalized_decoder_inputs):
880
- raise ValueError("decoder_inputs must contain non-empty strings.")
881
- if len(normalized_decoder_inputs) != len(records):
882
- raise ValueError("decoder_inputs must align one-to-one with embedding inputs.")
883
- decoder_input_fingerprint = _ordered_string_sha256(normalized_decoder_inputs)
884
- if decoder_attention_mask is not None:
885
- raise ValueError("decoder_attention_mask requires decoder_input_ids.")
886
- if decoder_input_ids is not None:
887
- if not isinstance(decoder_input_ids, Tensor) or decoder_input_ids.ndim != 2:
888
- raise ValueError("decoder_input_ids must have shape (batch, sequence).")
889
- if decoder_input_ids.shape[0] != len(records):
890
- raise ValueError("decoder_input_ids must align one-to-one with embedding inputs.")
891
- if decoder_input_ids.dtype == torch.bool or decoder_input_ids.is_floating_point():
892
- raise TypeError("decoder_input_ids must use an integer token dtype.")
893
- decoder_input_fingerprint = tensor_sha256(decoder_input_ids)
894
- decoder_mask_fingerprint: str | None = None
895
- if decoder_attention_mask is not None:
896
- if not isinstance(decoder_attention_mask, Tensor):
897
- raise TypeError("decoder_attention_mask must be a tensor.")
898
- if decoder_input_ids is None or decoder_attention_mask.shape != decoder_input_ids.shape:
899
- raise ValueError("decoder_attention_mask must match decoder_input_ids shape.")
900
- decoder_mask_fingerprint = tensor_sha256(decoder_attention_mask)
901
-
902
- context: dict[str, Any] = {
903
- "hidden_state_source": hidden_state_source,
904
- "hidden_state_index": hidden_state_index,
905
- "store_all_hidden_states": store_all_hidden_states,
906
- "decoder_input_fingerprint": decoder_input_fingerprint,
907
- "decoder_attention_mask_fingerprint": decoder_mask_fingerprint,
908
- "decoder_alignment": "input-position" if hidden_state_source == "decoder" else None,
909
- }
910
- metadata_hook = getattr(model, "_embedding_metadata", None)
911
- model_metadata: Mapping[str, Any] | None = None
912
- if callable(metadata_hook):
913
- model_metadata = metadata_hook(**context)
914
- if not isinstance(model_metadata, Mapping):
915
- raise TypeError("_embedding_metadata must return a mapping.")
916
- context["model_embedding"] = _fingerprint_jsonable(model_metadata)
917
- if hidden_state_source == "decoder":
918
- has_decoder_batch = callable(getattr(model, "_embedding_batch", None))
919
- declares_decoder_stack = (
920
- model_metadata is not None and model_metadata.get("hidden_state_stack") == "decoder"
921
- )
922
- if not has_decoder_batch or not declares_decoder_stack:
923
- raise ValueError(
924
- f"{model.__class__.__name__} does not declare decoder embedding support."
925
- )
926
- return context, normalized_decoder_inputs
927
-
928
-
929
- def _planned_batches(
930
- records: Sequence[EmbeddingInput],
931
- positions: range,
932
- *,
933
- batch_size: int,
934
- max_tokens_per_batch: int | None,
935
- max_length: int | None,
936
- truncate: bool,
937
- ) -> Iterator[list[int]]:
938
- """Length-bucket one bounded window while retaining stable output positions."""
939
-
940
- def effective_length(position: int) -> int:
941
- length = len(records[position].sequence)
942
- return min(length, max_length) if truncate and max_length is not None else length
943
-
944
- ordered = sorted(positions, key=lambda position: (-effective_length(position), position))
945
- batch: list[int] = []
946
- longest = 0
947
- for position in ordered:
948
- length = effective_length(position)
949
- if max_tokens_per_batch is not None and length > max_tokens_per_batch:
950
- raise ValueError(
951
- f"Input at position {position} has {length} residues, exceeding "
952
- f"max_tokens_per_batch={max_tokens_per_batch}."
953
- )
954
- candidate_longest = max(longest, length)
955
- exceeds_tokens = (
956
- max_tokens_per_batch is not None
957
- and candidate_longest * (len(batch) + 1) > max_tokens_per_batch
958
- )
959
- if batch and (len(batch) >= batch_size or exceeds_tokens):
960
- yield batch
961
- batch = []
962
- longest = 0
963
- batch.append(position)
964
- longest = max(longest, length)
965
- if batch:
966
- yield batch
967
-
968
-
969
- def embed_dataset(
970
- model: Any,
971
- inputs: (Iterable[str | EmbeddingInput | tuple[str, str]] | Mapping[str, str] | str | Path),
972
- *,
973
- batch_size: int = 2,
974
- pooling: str | Sequence[str] | None = None,
975
- full_embeddings: bool = False,
976
- output: str | Path | None = None,
977
- format: str = "safetensors",
978
- resume: bool = True,
979
- tokenizer: Any | None = None,
980
- max_length: int | None = None,
981
- truncate: bool = True,
982
- dtype: torch.dtype | None = torch.float32,
983
- shard_size: int = 2 * 1024**3,
984
- model_state_fingerprint: str | None = None,
985
- batch_window_size: int | None = None,
986
- max_tokens_per_batch: int | None = None,
987
- hidden_state_source: str = "encoder",
988
- decoder_inputs: Sequence[str] | None = None,
989
- decoder_input_ids: Tensor | None = None,
990
- decoder_attention_mask: Tensor | None = None,
991
- _embedding_batch_fn: Callable[..., EmbeddingBatch] | None = None,
992
- _embedding_batch_identity: Mapping[str, Any] | None = None,
993
- _allowed_unsupported_pooling: Sequence[str] = (),
994
- **model_kwargs: Any,
995
- ) -> EmbeddingResult:
996
- """Embed protein sequences with stable ordering and residue-only pooling."""
997
-
998
- for name, value in (
999
- ("batch_size", batch_size),
1000
- ("shard_size", shard_size),
1001
- ):
1002
- if not isinstance(value, int) or isinstance(value, bool):
1003
- raise TypeError(f"{name} must be a positive integer.")
1004
- if value <= 0:
1005
- raise ValueError(f"{name} must be a positive integer.")
1006
- for optional_name, optional_value in (
1007
- ("max_length", max_length),
1008
- ("max_tokens_per_batch", max_tokens_per_batch),
1009
- ("batch_window_size", batch_window_size),
1010
- ):
1011
- if optional_value is not None and (
1012
- not isinstance(optional_value, int) or isinstance(optional_value, bool)
1013
- ):
1014
- raise TypeError(f"{optional_name} must be a positive integer when provided.")
1015
- if optional_value is not None and optional_value <= 0:
1016
- raise ValueError(f"{optional_name} must be a positive integer when provided.")
1017
- for name, value in (
1018
- ("full_embeddings", full_embeddings),
1019
- ("resume", resume),
1020
- ("truncate", truncate),
1021
- ):
1022
- if not isinstance(value, bool):
1023
- raise TypeError(f"{name} must be a boolean.")
1024
- if not isinstance(format, str):
1025
- raise TypeError("format must be a string.")
1026
- if output is not None and not isinstance(output, (str, Path)):
1027
- raise TypeError("output must be a path or None.")
1028
- if model_state_fingerprint is not None and (
1029
- not isinstance(model_state_fingerprint, str) or not model_state_fingerprint
1030
- ):
1031
- raise ValueError("model_state_fingerprint must be a non-empty string when provided.")
1032
- if hidden_state_source not in {"encoder", "decoder"}:
1033
- raise ValueError("hidden_state_source must be 'encoder' or 'decoder'.")
1034
- hidden_state_index = model_kwargs.get("hidden_state_index", -1)
1035
- if not isinstance(hidden_state_index, int) or isinstance(hidden_state_index, bool):
1036
- raise TypeError("hidden_state_index must be an integer.")
1037
- store_all_hidden_states = model_kwargs.get("store_all_hidden_states", False)
1038
- if not isinstance(store_all_hidden_states, bool):
1039
- raise TypeError("store_all_hidden_states must be a boolean.")
1040
- if decoder_input_ids is not None:
1041
- if not isinstance(decoder_input_ids, Tensor):
1042
- raise TypeError("decoder_input_ids must be a tensor.")
1043
- if decoder_input_ids.is_meta:
1044
- raise ValueError("decoder_input_ids cannot be a meta tensor.")
1045
- if decoder_input_ids.ndim != 2 or decoder_input_ids.shape[1] == 0:
1046
- raise ValueError("decoder_input_ids must have non-empty shape (batch, sequence).")
1047
- if decoder_input_ids.dtype not in {torch.int32, torch.int64}:
1048
- raise TypeError("decoder_input_ids must use torch.int32 or torch.int64.")
1049
- if decoder_attention_mask is not None:
1050
- if not isinstance(decoder_attention_mask, Tensor):
1051
- raise TypeError("decoder_attention_mask must be a tensor.")
1052
- if decoder_attention_mask.is_meta:
1053
- raise ValueError("decoder_attention_mask cannot be a meta tensor.")
1054
- if decoder_attention_mask.is_complex() or not bool(
1055
- torch.isfinite(decoder_attention_mask).all()
1056
- ):
1057
- raise ValueError("decoder_attention_mask must contain finite binary values.")
1058
- if not bool(((decoder_attention_mask == 0) | (decoder_attention_mask == 1)).all()):
1059
- raise ValueError("decoder_attention_mask must contain finite binary values.")
1060
- pooling_names = (
1061
- (("mean",) if not full_embeddings else ())
1062
- if pooling is None
1063
- else ((pooling,) if isinstance(pooling, str) else tuple(pooling))
1064
- )
1065
- if full_embeddings and pooling is not None:
1066
- raise ValueError("full_embeddings=True cannot be combined with pooling.")
1067
- if not full_embeddings and not pooling_names:
1068
- raise ValueError("pooling is required unless full_embeddings=True.")
1069
- pooler = Pooler(pooling_names) if pooling_names else None
1070
-
1071
- if batch_size <= 0:
1072
- raise ValueError("batch_size must be positive.")
1073
- if format == "pth" or (output is not None and Path(output).suffix.lower() == ".pth"):
1074
- raise ValueError("Writing pickle-based .pth embeddings is not supported.")
1075
- if format not in _SUPPORTED_STORAGE_FORMATS:
1076
- raise ValueError("format must be 'safetensors' or 'sqlite'.")
1077
- if max_length is not None and max_length <= 0:
1078
- raise ValueError("max_length must be positive when provided.")
1079
- if max_tokens_per_batch is not None and max_tokens_per_batch <= 0:
1080
- raise ValueError("max_tokens_per_batch must be positive when provided.")
1081
- if not isinstance(dtype, (torch.dtype, type(None))):
1082
- raise TypeError("dtype must be a torch.dtype or None.")
1083
- if batch_window_size is not None and batch_window_size <= 0:
1084
- raise ValueError("batch_window_size must be positive when provided.")
1085
- if _embedding_batch_fn is not None and not callable(_embedding_batch_fn):
1086
- raise TypeError("_embedding_batch_fn must be callable when provided.")
1087
- if _embedding_batch_fn is not None and _embedding_batch_identity is None:
1088
- raise ValueError(
1089
- "_embedding_batch_identity is required with _embedding_batch_fn so persisted "
1090
- "runs bind the family-specific embedding behavior."
1091
- )
1092
- if _embedding_batch_identity is not None and not isinstance(_embedding_batch_identity, Mapping):
1093
- raise TypeError("_embedding_batch_identity must be a mapping when provided.")
1094
- if isinstance(_allowed_unsupported_pooling, (str, bytes)) or not isinstance(
1095
- _allowed_unsupported_pooling, Sequence
1096
- ):
1097
- raise TypeError("_allowed_unsupported_pooling must be a sequence of pooler names.")
1098
- if not all(isinstance(name, str) for name in _allowed_unsupported_pooling):
1099
- raise TypeError("_allowed_unsupported_pooling must contain only strings.")
1100
- allowed_unsupported_pooling = frozenset(_allowed_unsupported_pooling)
1101
- if allowed_unsupported_pooling and _embedding_batch_fn is None:
1102
- raise ValueError(
1103
- "_allowed_unsupported_pooling is only valid with a family-specific _embedding_batch_fn."
1104
- )
1105
- resolved_batch_window_size = (
1106
- batch_size * _DEFAULT_BATCH_WINDOW_MULTIPLIER
1107
- if batch_window_size is None
1108
- else batch_window_size
1109
- )
1110
- if resolved_batch_window_size < batch_size:
1111
- raise ValueError("batch_window_size must be at least batch_size.")
1112
- records = _normalize_inputs(inputs, disk_backed=output is not None)
1113
- _validate_untruncated_lengths(
1114
- records,
1115
- max_length=max_length,
1116
- truncate=truncate,
1117
- )
1118
- pooling_names = (
1119
- (("mean",) if not full_embeddings else ())
1120
- if pooling is None
1121
- else ((pooling,) if isinstance(pooling, str) else tuple(pooling))
1122
- )
1123
- if full_embeddings:
1124
- if pooling is not None:
1125
- raise ValueError("full_embeddings=True cannot be combined with pooling.")
1126
- elif not pooling_names:
1127
- raise ValueError("pooling is required unless full_embeddings=True.")
1128
- store_all_hidden_states = bool(model_kwargs.get("store_all_hidden_states", False))
1129
- if store_all_hidden_states and not full_embeddings:
1130
- raise ValueError("store_all_hidden_states=True requires full_embeddings=True.")
1131
-
1132
- unsupported = set(getattr(model, "embedding_unsupported_pooling", ()))
1133
- unknown_pooling_overrides = allowed_unsupported_pooling.difference(unsupported)
1134
- if unknown_pooling_overrides:
1135
- raise ValueError(
1136
- "_allowed_unsupported_pooling may only override poolers declared unsupported "
1137
- f"by the model; unknown overrides: {sorted(unknown_pooling_overrides)}."
1138
- )
1139
- unsupported.difference_update(allowed_unsupported_pooling)
1140
- requested_unsupported = unsupported.intersection(pooling_names)
1141
- if requested_unsupported:
1142
- raise ValueError(
1143
- f"{model.__class__.__name__} does not support pooling operations "
1144
- f"{sorted(requested_unsupported)}."
1145
- )
1146
-
1147
- # Constructing the pooler validates names and duplicate operations before
1148
- # any checkpoint hashing, tokenization, or inference occurs.
1149
- pooler = Pooler(pooling_names) if pooling_names else None
1150
- embedding_context, normalized_decoder_inputs = _embedding_context(
1151
- model,
1152
- records,
1153
- hidden_state_source=hidden_state_source,
1154
- decoder_inputs=decoder_inputs,
1155
- decoder_input_ids=decoder_input_ids,
1156
- decoder_attention_mask=decoder_attention_mask,
1157
- model_kwargs=model_kwargs,
1158
- )
1159
- if _embedding_batch_identity is not None:
1160
- embedding_context["family_adapter"] = _fingerprint_jsonable(_embedding_batch_identity)
1161
- if allowed_unsupported_pooling:
1162
- embedding_context["family_adapter_pooling_override"] = sorted(
1163
- allowed_unsupported_pooling
1164
- )
1165
-
1166
- tokenizer_metadata = _tokenizer_metadata(model, tokenizer)
1167
- (
1168
- input_fingerprint,
1169
- run_fingerprint,
1170
- resolved_model_state_fingerprint,
1171
- model_state_fingerprint_source,
1172
- ) = _run_fingerprint(
1173
- model,
1174
- records,
1175
- pooling=pooling_names,
1176
- full_embeddings=full_embeddings,
1177
- max_length=max_length,
1178
- truncate=truncate,
1179
- dtype=dtype,
1180
- model_kwargs=model_kwargs,
1181
- tokenizer_metadata=tokenizer_metadata,
1182
- model_state_fingerprint=model_state_fingerprint,
1183
- persist_output=output is not None,
1184
- embedding_context=embedding_context,
1185
- batch_size=batch_size,
1186
- batch_window_size=resolved_batch_window_size,
1187
- max_tokens_per_batch=max_tokens_per_batch,
1188
- )
1189
- output_already_exists = output is not None and _output_exists(output, format)
1190
- existing: EmbeddingResult | None = None
1191
- start_position = 0
1192
- if output is not None and resume and output_already_exists:
1193
- if format == "sqlite":
1194
- try:
1195
- existing = load_sqlite_result(output, run_id=run_fingerprint)
1196
- except KeyError:
1197
- existing = load_result(output, format=format)
1198
- else:
1199
- existing = load_result(output, format=format)
1200
- if existing.metadata.get("fingerprint_schema_version") != (_RUN_FINGERPRINT_SCHEMA_VERSION):
1201
- raise ValueError(
1202
- "Existing embeddings use an incompatible run fingerprint schema; "
1203
- "choose another output or set resume=False."
1204
- )
1205
- if existing.metadata.get("run_fingerprint") != run_fingerprint:
1206
- raise ValueError(
1207
- "Existing embeddings were produced by a different run fingerprint; "
1208
- "choose another output or set resume=False."
1209
- )
1210
- if len(existing) > len(records):
1211
- raise ValueError(
1212
- "Existing embeddings are not an ordered prefix of the requested inputs."
1213
- )
1214
- prefix_matches = all(
1215
- (observed.id, observed.sequence) == (expected.id, expected.sequence)
1216
- for expected, observed in zip(records, existing, strict=False)
1217
- )
1218
- if not prefix_matches:
1219
- raise ValueError(
1220
- "Existing embeddings are not an ordered prefix of the requested inputs."
1221
- )
1222
- if len(existing) == len(records) and existing.metadata.get("complete", True):
1223
- return existing
1224
- start_position = len(existing)
1225
-
1226
- sqlite_run_id: str | None = None
1227
- sqlite_replace_on_first_commit = False
1228
- sqlite_initial_metadata: dict[str, Any] | None = None
1229
- if output is not None and format == "sqlite":
1230
- sqlite_initial_metadata = {
1231
- "format_version": 1,
1232
- "fingerprint_schema_version": _RUN_FINGERPRINT_SCHEMA_VERSION,
1233
- "run_fingerprint": run_fingerprint,
1234
- "input_fingerprint": input_fingerprint,
1235
- "model_state_fingerprint": resolved_model_state_fingerprint,
1236
- "model_state_fingerprint_source": model_state_fingerprint_source,
1237
- "complete": False,
1238
- }
1239
- sqlite_run_id = run_fingerprint
1240
- if not resume and output_already_exists:
1241
- try:
1242
- load_sqlite_result(output, run_id=run_fingerprint)
1243
- except KeyError:
1244
- pass
1245
- else:
1246
- # Keep an exact prior run readable until replacement inference
1247
- # has produced the first complete commit window.
1248
- sqlite_replace_on_first_commit = True
1249
- if not sqlite_replace_on_first_commit:
1250
- initialize_sqlite_run(
1251
- output,
1252
- sqlite_initial_metadata,
1253
- resume=resume,
1254
- )
1255
-
1256
- stream_safetensors = output is not None and format == "safetensors"
1257
- attention_backend = _attention_backend(model)
1258
- output_records: list[EmbeddingRecord] = (
1259
- [] if sqlite_run_id is not None or stream_safetensors else list(existing or ())
1260
- )
1261
- output_descriptors: list[dict[str, Any]] | None = [] if output is None else None
1262
- pool_slices: dict[str, tuple[int, int]] = {}
1263
- if existing and pooler is not None:
1264
- pooled_width = existing[0].load_tensor().shape[-1]
1265
- if pooled_width % len(pooling_names) != 0:
1266
- raise ValueError("Stored pooled width is inconsistent with pooling metadata.")
1267
- pool_slices = pooler.output_slices(pooled_width // len(pooling_names))
1268
-
1269
- safetensors_writer: SafetensorsStreamWriter | None = None
1270
- if stream_safetensors:
1271
- if output is None:
1272
- raise RuntimeError("Safetensors streaming was enabled without an output destination.")
1273
- transactional_overwrite = output_already_exists and not resume
1274
- safetensors_writer = SafetensorsStreamWriter(
1275
- output,
1276
- {
1277
- "format_version": 1,
1278
- "fingerprint_schema_version": _RUN_FINGERPRINT_SCHEMA_VERSION,
1279
- "run_fingerprint": run_fingerprint,
1280
- "input_fingerprint": input_fingerprint,
1281
- "model_state_fingerprint": resolved_model_state_fingerprint,
1282
- "model_state_fingerprint_source": model_state_fingerprint_source,
1283
- "complete": False,
1284
- },
1285
- shard_size=shard_size,
1286
- existing=existing or (),
1287
- reuse_existing=bool(resume and existing is not None),
1288
- publish_initial=not transactional_overwrite,
1289
- publish_incremental=not transactional_overwrite,
1290
- )
1291
- need_attentions = "parti" in pooling_names
1292
-
1293
- config = getattr(model, "config", None)
1294
- model_type = str(getattr(config, "model_type", "")).lower()
1295
- resolved_tokenizer = tokenizer if tokenizer is not None else getattr(model, "tokenizer", None)
1296
- with _temporary_eval(model), torch.inference_mode():
1297
- for window_start in range(start_position, len(records), resolved_batch_window_size):
1298
- window_stop = min(window_start + resolved_batch_window_size, len(records))
1299
- window_records = records[window_start:window_stop]
1300
- if not isinstance(window_records, Sequence):
1301
- raise RuntimeError("The immutable embedding spool returned a non-sequence window.")
1302
- window_results: dict[int, EmbeddingRecord] = {}
1303
- for local_positions in _planned_batches(
1304
- window_records,
1305
- range(len(window_records)),
1306
- batch_size=batch_size,
1307
- max_tokens_per_batch=max_tokens_per_batch,
1308
- max_length=max_length,
1309
- truncate=truncate,
1310
- ):
1311
- batch_positions = [window_start + position for position in local_positions]
1312
- batch_records = [window_records[position] for position in local_positions]
1313
- sequences = [
1314
- record.sequence[:max_length]
1315
- if truncate and max_length is not None
1316
- else record.sequence
1317
- for record in batch_records
1318
- ]
1319
- batch_model_kwargs = dict(model_kwargs)
1320
- if model_type == "fast_ankh" or hidden_state_source == "decoder":
1321
- batch_model_kwargs["hidden_state_source"] = hidden_state_source
1322
- if normalized_decoder_inputs is not None:
1323
- batch_model_kwargs["decoder_inputs"] = [
1324
- normalized_decoder_inputs[position] for position in batch_positions
1325
- ]
1326
- if decoder_input_ids is not None:
1327
- # decoder_input_ids: (n_records, l_decoder)
1328
- indices = torch.tensor( # (b,)
1329
- batch_positions,
1330
- device=decoder_input_ids.device,
1331
- dtype=torch.long,
1332
- )
1333
- batch_model_kwargs["decoder_input_ids"] = ( # (b, l_decoder)
1334
- decoder_input_ids.index_select(0, indices)
1335
- )
1336
- if decoder_attention_mask is not None:
1337
- # decoder_attention_mask: (n_records, l_decoder)
1338
- indices = torch.tensor( # (b,)
1339
- batch_positions,
1340
- device=decoder_attention_mask.device,
1341
- dtype=torch.long,
1342
- )
1343
- batch_model_kwargs["decoder_attention_mask"] = (
1344
- decoder_attention_mask.index_select(0, indices) # (b, l_decoder)
1345
- )
1346
- custom_batch = _embedding_batch_fn or getattr(model, "_embedding_batch", None)
1347
- if custom_batch is not None:
1348
- if model_type == "fast_ankh":
1349
- batch = custom_batch(
1350
- sequences,
1351
- tokenizer=resolved_tokenizer,
1352
- max_length=max_length,
1353
- truncate=truncate,
1354
- need_attentions=need_attentions,
1355
- **batch_model_kwargs,
1356
- )
1357
- else:
1358
- batch = custom_batch(sequences, **batch_model_kwargs)
1359
- if not isinstance(batch, EmbeddingBatch):
1360
- raise TypeError("_embedding_batch must return EmbeddingBatch.")
1361
- else:
1362
- batch = _generic_embedding_batch(
1363
- model,
1364
- sequences,
1365
- tokenizer=tokenizer,
1366
- max_length=max_length,
1367
- truncate=truncate,
1368
- need_attentions=need_attentions,
1369
- model_kwargs=batch_model_kwargs,
1370
- )
1371
- X = batch.X # (b, l, d) or (b, n_states, l, d)
1372
- raw_mask = batch.residue_mask # (b, l)
1373
- if not isinstance(X, Tensor) or not isinstance(raw_mask, Tensor):
1374
- raise TypeError("Embedding batches must provide Tensor X and residue_mask.")
1375
- if X.is_meta or raw_mask.is_meta:
1376
- raise ValueError("Embedding batches cannot contain meta tensors.")
1377
- if not X.is_floating_point():
1378
- raise TypeError("Embedding batches must use a floating-point X dtype.")
1379
- if raw_mask.is_complex() or not bool(torch.isfinite(raw_mask).all()):
1380
- raise ValueError("Embedding residue_mask must contain finite binary values.")
1381
- if not bool(((raw_mask == 0) | (raw_mask == 1)).all()):
1382
- raise ValueError("Embedding residue_mask must contain finite binary values.")
1383
- M = raw_mask.to(device=X.device, dtype=torch.bool) # (b, l)
1384
- valid_X_shape = (
1385
- X.ndim == 3
1386
- and X.shape[0] == len(batch_records)
1387
- and X.shape[-1] > 0
1388
- and M.shape == X.shape[:2]
1389
- )
1390
- valid_all_states_shape = (
1391
- X.ndim == 4
1392
- and store_all_hidden_states
1393
- and full_embeddings
1394
- and X.shape[0] == len(batch_records)
1395
- and X.shape[1] > 0
1396
- and X.shape[-1] > 0
1397
- and M.shape == (X.shape[0], X.shape[2])
1398
- )
1399
- if not (valid_X_shape or valid_all_states_shape):
1400
- raise ValueError(
1401
- "Embedding batches must provide X with shape (b, l, d), or "
1402
- "(b, states, l, d) when storing all hidden states, and "
1403
- "residue_mask with shape (b, l)."
1404
- )
1405
- if not bool(M.any(dim=1).all()):
1406
- raise ValueError("Every embedding sample must contain a biological residue.")
1407
- finite_selected = ( # X.shape
1408
- torch.isfinite(X) | ~M.unsqueeze(-1)
1409
- if X.ndim == 3
1410
- else torch.isfinite(X) | ~M[:, None, :, None]
1411
- )
1412
- if not bool(finite_selected.all()):
1413
- raise ValueError("Biological residue embeddings produced non-finite output.")
1414
- if need_attentions:
1415
- # Validate the biological graph only after mask integrity is established.
1416
- _validate_parti_length(M)
1417
- if dtype is not None:
1418
- X = X.to(dtype=dtype) # unchanged shape
1419
-
1420
- if full_embeddings:
1421
- if X.ndim == 4:
1422
- values = [
1423
- X_i[:, M_i, :].detach().cpu() # (n_states, r_i, d)
1424
- for X_i, M_i in zip(X, M, strict=True)
1425
- ]
1426
- else:
1427
- values = [
1428
- X_i[M_i].detach().cpu() # (r_i, d)
1429
- for X_i, M_i in zip(X, M, strict=True)
1430
- ]
1431
- else:
1432
- if pooler is None:
1433
- raise RuntimeError(
1434
- "Pooled embedding output was requested without an initialized pooler."
1435
- )
1436
- Y = pooler( # (b, n_poolers * d)
1437
- X,
1438
- M,
1439
- attentions=batch.attentions,
1440
- attention_backend=attention_backend,
1441
- )
1442
- pool_slices = pooler.output_slices(X.shape[-1])
1443
- values = list(Y.detach().cpu().unbind(0)) # each: (n_poolers * d,)
1444
- for position, record, value in zip(
1445
- batch_positions, batch_records, values, strict=True
1446
- ):
1447
- window_results[position] = EmbeddingRecord(record.id, record.sequence, value)
1448
-
1449
- new_records = [
1450
- window_results[position] for position in range(window_start, window_stop)
1451
- ]
1452
- if output_descriptors is not None:
1453
- output_descriptors.extend(
1454
- _output_descriptor(window_start + offset, record)
1455
- for offset, record in enumerate(new_records)
1456
- )
1457
- if output is not None and sqlite_run_id is not None:
1458
- append_sqlite_records(
1459
- output,
1460
- sqlite_run_id,
1461
- window_start,
1462
- new_records,
1463
- replace_metadata=(
1464
- sqlite_initial_metadata if sqlite_replace_on_first_commit else None
1465
- ),
1466
- )
1467
- sqlite_replace_on_first_commit = False
1468
- elif safetensors_writer is not None:
1469
- safetensors_writer.append(new_records)
1470
- else:
1471
- output_records.extend(new_records)
1472
-
1473
- software_versions = _software_versions()
1474
- projection = getattr(model, "embedding_projection", None)
1475
- resolved_layer = getattr(
1476
- model,
1477
- "embedding_layer",
1478
- model_kwargs.get("hidden_state_index", -1),
1479
- )
1480
- token_policy = getattr(
1481
- model,
1482
- "embedding_token_policy",
1483
- {
1484
- "unit": "residue",
1485
- "include": ["biological residues"],
1486
- "exclude": [
1487
- "BOS",
1488
- "EOS",
1489
- "padding",
1490
- "chain delimiters",
1491
- "non-protein tokens",
1492
- ],
1493
- },
1494
- )
1495
- model_identity = _model_identity_metadata(model)
1496
- metadata: dict[str, Any] = {
1497
- "format_version": 1,
1498
- "fingerprint_schema_version": _RUN_FINGERPRINT_SCHEMA_VERSION,
1499
- "run_fingerprint": run_fingerprint,
1500
- "input_fingerprint": input_fingerprint,
1501
- "model_state_fingerprint": resolved_model_state_fingerprint,
1502
- "model_state_fingerprint_source": model_state_fingerprint_source,
1503
- "model_class": f"{model.__class__.__module__}.{model.__class__.__qualname__}",
1504
- **model_identity,
1505
- "dtype": str(dtype).removeprefix("torch.") if dtype is not None else "model",
1506
- "attention_backend": attention_backend,
1507
- "attention_kernel": _attention_kernel_metadata(attention_backend),
1508
- "layer": resolved_layer,
1509
- "projection": projection,
1510
- "esmc_source": getattr(model, "_esmc_source", None),
1511
- "esmc_revision": getattr(model, "_esmc_source_revision", None),
1512
- "esmc_files": getattr(model, "_esmc_source_files", None),
1513
- "token_policy": token_policy,
1514
- "tokenizer": tokenizer_metadata,
1515
- **embedding_context,
1516
- "pooling": list(pooling_names),
1517
- "pool_slices": pool_slices,
1518
- "full_embeddings": full_embeddings,
1519
- "max_length": max_length,
1520
- "truncate": truncate,
1521
- "truncation": {"enabled": truncate, "max_length": max_length},
1522
- "batching": {
1523
- "batch_size": batch_size,
1524
- "batch_window_size": resolved_batch_window_size,
1525
- "max_tokens_per_batch": max_tokens_per_batch,
1526
- "input_storage": ("disk-spool" if isinstance(records, _InputSpool) else "memory"),
1527
- "ordering": "bounded-length-bucketed-stable-output",
1528
- "resume_commit_granularity": (
1529
- "not-applicable"
1530
- if output is None
1531
- else "batch-window"
1532
- if format == "sqlite"
1533
- else "shard-flush"
1534
- ),
1535
- },
1536
- "residue_mask_policy": "biological-residues-only",
1537
- "record_count": len(records),
1538
- "descriptor_index": (
1539
- "memory-metadata"
1540
- if output is None
1541
- else "sqlite-records"
1542
- if format == "sqlite"
1543
- else "safetensors-generation-index"
1544
- ),
1545
- "storage_format": format if output is not None else "memory",
1546
- "software": software_versions,
1547
- "execution": _execution_identity_metadata(model),
1548
- "adapter": _adapter_identity_metadata(model),
1549
- "torch_version": software_versions["torch"],
1550
- "transformers_version": software_versions["transformers"],
1551
- "complete": True,
1552
- }
1553
- if output_descriptors is not None:
1554
- metadata["outputs"] = output_descriptors
1555
- metadata["tensor_hashes"] = [item["sha256"] for item in output_descriptors]
1556
- status = getattr(model, "esmc_precision_status", None)
1557
- if status is not None:
1558
- metadata["esmc_precision"] = status.as_dict() if hasattr(status, "as_dict") else status
1559
- if output is not None and sqlite_run_id is not None:
1560
- update_sqlite_run_metadata(output, sqlite_run_id, metadata)
1561
- return load_sqlite_result(output, run_id=sqlite_run_id)
1562
- if safetensors_writer is not None:
1563
- return safetensors_writer.publish(complete=True, metadata=metadata)
1564
- result = EmbeddingResult(output_records, metadata)
1565
- if output is not None:
1566
- return save_result(result, output, format=format, shard_size=shard_size)
1567
- return result
1568
-
1569
-
1570
- class EmbeddingMixin:
1571
- """Small delegation mixin shared by FastPLMs model classes."""
1572
-
1573
- def embed_dataset(self, inputs: Any, **kwargs: Any) -> EmbeddingResult:
1574
- return embed_dataset(self, inputs, **kwargs)
1575
-
1576
-
1577
- __all__ = [
1578
- "EmbeddingMixin",
1579
- "embed_dataset",
1580
- "iter_fasta",
1581
- "parse_fasta",
1582
- "select_hidden_state_embeddings",
1583
- ]
 
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
+ ]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
fastplms/models.toml CHANGED
@@ -11,6 +11,7 @@ revision = "db6b51744f0cd7061386442c09df890fc6d9f47e"
11
  version = 2
12
  expected_variant = "flash_attn2"
13
  dtypes = ["bfloat16"]
 
14
 
15
  [[attention_kernels]]
16
  implementation = "flash_attention_3"
@@ -19,6 +20,7 @@ revision = "43f0bd269777115d94ff826e0d113ce9c1c9087b"
19
  version = 1
20
  expected_variant = "flash_attn3"
21
  dtypes = ["bfloat16"]
 
22
 
23
  [[runtime_assets]]
24
  id = "esmfold2_ccd"
@@ -161,6 +163,8 @@ extra = "core"
161
  reference_container = "reference-esm2"
162
  reference_adapter = "tests.parity.support.reference_adapters.esm2"
163
  attention = ["eager", "sdpa", "flex_attention", "flash_attention_2", "flash_attention_3"]
 
 
164
  dtypes = ["float32", "bfloat16"]
165
  bf16_execution = "fp32_parameters_autocast"
166
  precisions = ["default"]
@@ -185,6 +189,8 @@ extra = "core"
185
  reference_container = "reference-biohub-esm"
186
  reference_adapter = "tests.parity.support.reference_adapters.esm_plusplus"
187
  attention = ["eager", "sdpa", "flex_attention", "flash_attention_2", "flash_attention_3"]
 
 
188
  dtypes = ["float32", "bfloat16"]
189
  bf16_execution = "static_parameters"
190
  precisions = ["default", "fp8"]
@@ -210,6 +216,7 @@ extra = "core"
210
  reference_container = "reference-biohub-esm"
211
  reference_adapter = "tests.parity.support.reference_adapters.esm3"
212
  attention = ["eager", "sdpa", "flex_attention"]
 
213
  dtypes = ["float32", "bfloat16"]
214
  bf16_execution = "fp32_parameters_autocast"
215
  precisions = ["default"]
@@ -234,6 +241,7 @@ extra = "core"
234
  reference_container = "reference-e1"
235
  reference_adapter = "tests.parity.support.reference_adapters.e1"
236
  attention = ["sdpa", "flex_attention"]
 
237
  dtypes = ["float32", "bfloat16"]
238
  bf16_execution = "static_parameters"
239
  precisions = ["default"]
@@ -260,6 +268,8 @@ extra = "core"
260
  reference_container = "reference-dplm"
261
  reference_adapter = "tests.parity.support.reference_adapters.dplm"
262
  attention = ["eager", "sdpa", "flex_attention", "flash_attention_3"]
 
 
263
  dtypes = ["float32", "bfloat16"]
264
  bf16_execution = "fp32_parameters_autocast"
265
  precisions = ["default"]
@@ -284,6 +294,7 @@ extra = "core"
284
  reference_container = "reference-dplm"
285
  reference_adapter = "tests.parity.support.reference_adapters.dplm2"
286
  attention = ["sdpa"]
 
287
  dtypes = ["float32", "bfloat16"]
288
  bf16_execution = "fp32_parameters_autocast"
289
  precisions = ["default"]
@@ -309,6 +320,7 @@ extra = "core"
309
  reference_container = "reference-ankh"
310
  reference_adapter = "tests.parity.support.reference_adapters.ankh"
311
  attention = ["eager", "sdpa"]
 
312
  dtypes = ["float32", "bfloat16"]
313
  bf16_execution = "static_parameters"
314
  precisions = ["default"]
@@ -358,6 +370,7 @@ extra = "structure"
358
  reference_container = "reference-esmfold"
359
  reference_adapter = "tests.parity.support.reference_adapters.esmfold"
360
  attention = ["eager", "sdpa", "flex_attention"]
 
361
  dtypes = ["float32", "bfloat16"]
362
  bf16_execution = "fp32_parameters_autocast"
363
  precisions = ["default"]
@@ -383,6 +396,7 @@ extra = "structure"
383
  reference_container = "reference-esmfold2"
384
  reference_adapter = "tests.parity.support.reference_adapters.esmfold2"
385
  attention = ["eager", "sdpa", "flex_attention"]
 
386
  dtypes = ["float32", "bfloat16"]
387
  bf16_execution = "fp32_parameters_autocast"
388
  precisions = ["auto", "fp32", "bf16", "fp8"]
 
11
  version = 2
12
  expected_variant = "flash_attn2"
13
  dtypes = ["bfloat16"]
14
+ min_cuda_capability = [8, 0]
15
 
16
  [[attention_kernels]]
17
  implementation = "flash_attention_3"
 
20
  version = 1
21
  expected_variant = "flash_attn3"
22
  dtypes = ["bfloat16"]
23
+ min_cuda_capability = [8, 0]
24
 
25
  [[runtime_assets]]
26
  id = "esmfold2_ccd"
 
163
  reference_container = "reference-esm2"
164
  reference_adapter = "tests.parity.support.reference_adapters.esm2"
165
  attention = ["eager", "sdpa", "flex_attention", "flash_attention_2", "flash_attention_3"]
166
+ attention_auto_order = ["sdpa"]
167
+ attention_auto_evidence = "docs/evidence/attention/backend_latency.json"
168
  dtypes = ["float32", "bfloat16"]
169
  bf16_execution = "fp32_parameters_autocast"
170
  precisions = ["default"]
 
189
  reference_container = "reference-biohub-esm"
190
  reference_adapter = "tests.parity.support.reference_adapters.esm_plusplus"
191
  attention = ["eager", "sdpa", "flex_attention", "flash_attention_2", "flash_attention_3"]
192
+ attention_auto_order = ["sdpa"]
193
+ attention_auto_evidence = "docs/evidence/attention/backend_latency.json"
194
  dtypes = ["float32", "bfloat16"]
195
  bf16_execution = "static_parameters"
196
  precisions = ["default", "fp8"]
 
216
  reference_container = "reference-biohub-esm"
217
  reference_adapter = "tests.parity.support.reference_adapters.esm3"
218
  attention = ["eager", "sdpa", "flex_attention"]
219
+ attention_auto_order = ["sdpa"]
220
  dtypes = ["float32", "bfloat16"]
221
  bf16_execution = "fp32_parameters_autocast"
222
  precisions = ["default"]
 
241
  reference_container = "reference-e1"
242
  reference_adapter = "tests.parity.support.reference_adapters.e1"
243
  attention = ["sdpa", "flex_attention"]
244
+ attention_auto_order = ["sdpa"]
245
  dtypes = ["float32", "bfloat16"]
246
  bf16_execution = "static_parameters"
247
  precisions = ["default"]
 
268
  reference_container = "reference-dplm"
269
  reference_adapter = "tests.parity.support.reference_adapters.dplm"
270
  attention = ["eager", "sdpa", "flex_attention", "flash_attention_3"]
271
+ attention_auto_order = ["sdpa"]
272
+ attention_auto_evidence = "docs/evidence/attention/backend_latency.json"
273
  dtypes = ["float32", "bfloat16"]
274
  bf16_execution = "fp32_parameters_autocast"
275
  precisions = ["default"]
 
294
  reference_container = "reference-dplm"
295
  reference_adapter = "tests.parity.support.reference_adapters.dplm2"
296
  attention = ["sdpa"]
297
+ attention_auto_order = ["sdpa"]
298
  dtypes = ["float32", "bfloat16"]
299
  bf16_execution = "fp32_parameters_autocast"
300
  precisions = ["default"]
 
320
  reference_container = "reference-ankh"
321
  reference_adapter = "tests.parity.support.reference_adapters.ankh"
322
  attention = ["eager", "sdpa"]
323
+ attention_auto_order = ["sdpa", "eager"]
324
  dtypes = ["float32", "bfloat16"]
325
  bf16_execution = "static_parameters"
326
  precisions = ["default"]
 
370
  reference_container = "reference-esmfold"
371
  reference_adapter = "tests.parity.support.reference_adapters.esmfold"
372
  attention = ["eager", "sdpa", "flex_attention"]
373
+ attention_auto_order = ["sdpa"]
374
  dtypes = ["float32", "bfloat16"]
375
  bf16_execution = "fp32_parameters_autocast"
376
  precisions = ["default"]
 
396
  reference_container = "reference-esmfold2"
397
  reference_adapter = "tests.parity.support.reference_adapters.esmfold2"
398
  attention = ["eager", "sdpa", "flex_attention"]
399
+ attention_auto_order = ["sdpa"]
400
  dtypes = ["float32", "bfloat16"]
401
  bf16_execution = "fp32_parameters_autocast"
402
  precisions = ["auto", "fp32", "bf16", "fp8"]
fastplms/models/esm_plusplus/modeling_esm_plusplus.py CHANGED
@@ -35,10 +35,12 @@ try:
35
  AttentionBackend,
36
  BlockMask,
37
  FastPLMsAttentionMixin,
 
38
  _get_flex_attention_fn,
39
  _get_flex_block_mask,
40
  flex_attention,
41
  get_attention_mask,
 
42
  kernels_flash_attention_func,
43
  resolve_attention_backend,
44
  resolve_attention_backend_for_call,
@@ -52,11 +54,13 @@ except ModuleNotFoundError as error:
52
  "EmbeddingMixin",
53
  "FastPLMsAttentionMixin",
54
  "FastPLMTestTimeTrainingMixin",
 
55
  "Pooler",
56
  "_get_flex_attention_fn",
57
  "_get_flex_block_mask",
58
  "flex_attention",
59
  "get_attention_mask",
 
60
  "kernels_flash_attention_func",
61
  "resolve_attention_backend",
62
  "resolve_attention_backend_for_call",
@@ -296,13 +300,27 @@ def apply_rotary_emb_torch(
296
  raise AssertionError("rotary width exceeds the attention head dimension")
297
 
298
  token_count = x.shape[1]
299
- cos_full = torch.cat((cos[:token_count], cos[:token_count]), dim=-1).unsqueeze(1)
300
- sin_full = torch.cat((sin[:token_count], sin[:token_count]), dim=-1).unsqueeze(1)
301
- x_rotary = x[..., :rotary_width]
302
- y_rotary = x_rotary * cos_full + rotate_half(x_rotary, interleaved) * sin_full
 
 
 
 
 
 
 
 
 
 
 
 
 
 
303
  if rotary_width == x.shape[-1]:
304
  return y_rotary
305
- return torch.cat((y_rotary, x[..., rotary_width:]), dim=-1)
306
 
307
 
308
  class RotaryEmbedding(torch.nn.Module):
@@ -344,6 +362,10 @@ class RotaryEmbedding(torch.nn.Module):
344
  self._sin_cached: torch.Tensor | None = None
345
  self._cos_k_cached: torch.Tensor | None = None
346
  self._sin_k_cached: torch.Tensor | None = None
 
 
 
 
347
 
348
  def reset_parameters(self, device: torch.device | str | None = None) -> None:
349
  """Rebuild the non-persistent frequency buffers on ``device``."""
@@ -426,6 +448,8 @@ class RotaryEmbedding(torch.nn.Module):
426
  if self.scale is None:
427
  self._cos_cached = cos_angles.to(dtype)
428
  self._sin_cached = sin_angles.to(dtype)
 
 
429
  return
430
 
431
  centered_positions = (
@@ -459,11 +483,17 @@ class RotaryEmbedding(torch.nn.Module):
459
  if self.scale is not None:
460
  raise AssertionError("Scaled rotary embeddings are unsupported for ESMC.")
461
 
462
- cos_angles = self._cos_cached
463
- sin_angles = self._sin_cached
 
 
 
 
 
 
464
  return (
465
- apply_rotary_emb_torch(q, cos_angles, sin_angles, self.interleaved, True),
466
- apply_rotary_emb_torch(k, cos_angles, sin_angles, self.interleaved, True),
467
  )
468
 
469
 
@@ -540,6 +570,7 @@ class MultiHeadAttention(nn.Module):
540
  flex_block_mask: BlockMask | None = None,
541
  output_attentions: bool = False,
542
  output_s_max: bool = False,
 
543
  ) -> tuple[torch.Tensor, torch.Tensor | None, list[torch.Tensor] | None]:
544
  # x: (b, l, d)
545
  qkv = self.layernorm_qkv(x) # (b, l, 3 * d)
@@ -562,6 +593,7 @@ class MultiHeadAttention(nn.Module):
562
  flex_block_mask=flex_block_mask,
563
  output_attentions=output_attentions,
564
  output_s_max=output_s_max,
 
565
  )
566
 
567
  output = self.out_proj(attn_output)
@@ -577,6 +609,7 @@ class MultiHeadAttention(nn.Module):
577
  flex_block_mask: BlockMask | None = None,
578
  output_attentions: bool = False,
579
  output_s_max: bool = False,
 
580
  ) -> tuple[torch.Tensor, torch.Tensor | None, list[torch.Tensor] | None]:
581
  if output_attentions:
582
  return self._manual_attn(
@@ -590,7 +623,7 @@ class MultiHeadAttention(nn.Module):
590
  return attn_output, None, s_max
591
  if self.attn_backend.is_flash:
592
  attn_output, attn_weights = self._kernels_flash_attn(
593
- query_heads, key_heads, value_heads, attention_mask_2d
594
  )
595
  elif self.attn_backend == AttentionBackend.FLEX:
596
  attn_output, attn_weights = self._flex_attn(
@@ -647,6 +680,7 @@ class MultiHeadAttention(nn.Module):
647
  key_heads: torch.Tensor,
648
  value_heads: torch.Tensor,
649
  attention_mask_2d: torch.Tensor | None = None,
 
650
  ) -> tuple[torch.Tensor, None]:
651
  query_tokens = query_heads.transpose(1, 2).contiguous()
652
  key_tokens = key_heads.transpose(1, 2).contiguous()
@@ -658,6 +692,7 @@ class MultiHeadAttention(nn.Module):
658
  attention_mask_2d=attention_mask_2d,
659
  causal=False,
660
  implementation=self.attn_backend.value,
 
661
  )
662
  return rearrange(attn_output, "b s h d -> b s (h d)"), None
663
 
@@ -747,6 +782,7 @@ class UnifiedTransformerBlock(nn.Module):
747
  flex_block_mask: BlockMask | None = None,
748
  output_attentions: bool = False,
749
  output_s_max: bool = False,
 
750
  ) -> tuple[torch.Tensor, torch.Tensor | None, list[torch.Tensor] | None]:
751
  attn_output, attn_weights, s_max = self.attn(
752
  x,
@@ -755,6 +791,7 @@ class UnifiedTransformerBlock(nn.Module):
755
  flex_block_mask=flex_block_mask,
756
  output_attentions=output_attentions,
757
  output_s_max=output_s_max,
 
758
  )
759
  x = x + self.dropout(attn_output) / self.scaling_factor
760
  x = x + self.dropout(self.ffn(x)) / self.scaling_factor
@@ -857,16 +894,20 @@ class TransformerStack(nn.Module):
857
  # Match the pinned Biohub Transformers contract: a supplied sequence_id
858
  # is authoritative and must encode padding as -1. attention_mask is
859
  # ignored in that mode rather than intersected with the chain mask.
860
- attention_mask_2d, attention_mask_4d, flex_block_mask = (
861
- self._prepare_attention_masks(
862
- attention_mask=attention_mask,
863
- sequence_id=sequence_id,
864
- batch_size=x.shape[0],
865
- seq_len=x.shape[1],
866
- device=x.device,
867
- dtype=x.dtype,
868
- output_attentions=bool(output_attentions),
869
- )
 
 
 
 
870
  )
871
 
872
  for layer_index, block in enumerate(self.blocks):
@@ -890,6 +931,7 @@ class TransformerStack(nn.Module):
890
  flex_block_mask=flex_block_mask,
891
  output_attentions=output_attentions,
892
  output_s_max=output_s_max,
 
893
  )
894
  else:
895
  x, attn_weights, s_max = block(
@@ -899,6 +941,7 @@ class TransformerStack(nn.Module):
899
  flex_block_mask=flex_block_mask,
900
  output_attentions=output_attentions,
901
  output_s_max=output_s_max,
 
902
  )
903
 
904
  if attentions is not None:
@@ -1048,6 +1091,7 @@ class PreTrainedESMplusplusModel(FastPLMsAttentionMixin, PreTrainedModel):
1048
  "flash_attention_2",
1049
  "flash_attention_3",
1050
  )
 
1051
 
1052
  def __init__(self, config: ESMplusplusConfig, *args: object, **kwargs: object) -> None:
1053
  super().__init__(config, *args, **kwargs)
 
35
  AttentionBackend,
36
  BlockMask,
37
  FastPLMsAttentionMixin,
38
+ FlashPaddingLayout,
39
  _get_flex_attention_fn,
40
  _get_flex_block_mask,
41
  flex_attention,
42
  get_attention_mask,
43
+ get_flash_padding_layout,
44
  kernels_flash_attention_func,
45
  resolve_attention_backend,
46
  resolve_attention_backend_for_call,
 
54
  "EmbeddingMixin",
55
  "FastPLMsAttentionMixin",
56
  "FastPLMTestTimeTrainingMixin",
57
+ "FlashPaddingLayout",
58
  "Pooler",
59
  "_get_flex_attention_fn",
60
  "_get_flex_block_mask",
61
  "flex_attention",
62
  "get_attention_mask",
63
+ "get_flash_padding_layout",
64
  "kernels_flash_attention_func",
65
  "resolve_attention_backend",
66
  "resolve_attention_backend_for_call",
 
300
  raise AssertionError("rotary width exceeds the attention head dimension")
301
 
302
  token_count = x.shape[1]
303
+ cos_full = torch.cat((cos[:token_count], cos[:token_count]), dim=-1) # (l, d_r)
304
+ sin_full = torch.cat((sin[:token_count], sin[:token_count]), dim=-1) # (l, d_r)
305
+ return _rotate_with_full_tables(x, cos_full, sin_full, interleaved)
306
+
307
+
308
+ def _rotate_with_full_tables(
309
+ x: torch.Tensor,
310
+ cos_full: torch.Tensor,
311
+ sin_full: torch.Tensor,
312
+ interleaved: bool,
313
+ ) -> torch.Tensor:
314
+ # x: (b, l, h, d); cos_full, sin_full: (l, d_r) with rotary width d_r <= d
315
+ rotary_width = cos_full.shape[-1]
316
+ x_rotary = x[..., :rotary_width] # (b, l, h, d_r)
317
+ y_rotary = ( # (b, l, h, d_r)
318
+ x_rotary * cos_full.unsqueeze(1)
319
+ + rotate_half(x_rotary, interleaved) * sin_full.unsqueeze(1)
320
+ )
321
  if rotary_width == x.shape[-1]:
322
  return y_rotary
323
+ return torch.cat((y_rotary, x[..., rotary_width:]), dim=-1) # (b, l, h, d)
324
 
325
 
326
  class RotaryEmbedding(torch.nn.Module):
 
362
  self._sin_cached: torch.Tensor | None = None
363
  self._cos_k_cached: torch.Tensor | None = None
364
  self._sin_k_cached: torch.Tensor | None = None
365
+ # Both halves of a head rotate by the same angles. Q and K of every layer
366
+ # read these full-width tables instead of concatenating the halves again.
367
+ self._cos_full_cached: torch.Tensor | None = None
368
+ self._sin_full_cached: torch.Tensor | None = None
369
 
370
  def reset_parameters(self, device: torch.device | str | None = None) -> None:
371
  """Rebuild the non-persistent frequency buffers on ``device``."""
 
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 = (
 
483
  if self.scale is not None:
484
  raise AssertionError("Scaled rotary embeddings are unsupported for ESMC.")
485
 
486
+ if self._cos_full_cached is None or self._sin_full_cached is None:
487
+ raise RuntimeError("Rotary cache initialization did not produce full-width tables.")
488
+ if 2 * self._cos_cached.shape[-1] > q.shape[-1]:
489
+ raise AssertionError("rotary width exceeds the attention head dimension")
490
+
491
+ token_count = q.shape[1]
492
+ cos_full = self._cos_full_cached[:token_count] # (l, d_r)
493
+ sin_full = self._sin_full_cached[:token_count] # (l, d_r)
494
  return (
495
+ _rotate_with_full_tables(q, cos_full, sin_full, self.interleaved),
496
+ _rotate_with_full_tables(k, cos_full, sin_full, self.interleaved),
497
  )
498
 
499
 
 
570
  flex_block_mask: BlockMask | None = None,
571
  output_attentions: bool = False,
572
  output_s_max: bool = False,
573
+ flash_padding_layout: FlashPaddingLayout | None = None,
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)
 
593
  flex_block_mask=flex_block_mask,
594
  output_attentions=output_attentions,
595
  output_s_max=output_s_max,
596
+ flash_padding_layout=flash_padding_layout,
597
  )
598
 
599
  output = self.out_proj(attn_output)
 
609
  flex_block_mask: BlockMask | None = None,
610
  output_attentions: bool = False,
611
  output_s_max: bool = False,
612
+ flash_padding_layout: FlashPaddingLayout | None = None,
613
  ) -> tuple[torch.Tensor, torch.Tensor | None, list[torch.Tensor] | None]:
614
  if output_attentions:
615
  return self._manual_attn(
 
623
  return attn_output, None, s_max
624
  if self.attn_backend.is_flash:
625
  attn_output, attn_weights = self._kernels_flash_attn(
626
+ query_heads, key_heads, value_heads, attention_mask_2d, flash_padding_layout
627
  )
628
  elif self.attn_backend == AttentionBackend.FLEX:
629
  attn_output, attn_weights = self._flex_attn(
 
680
  key_heads: torch.Tensor,
681
  value_heads: torch.Tensor,
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()
 
692
  attention_mask_2d=attention_mask_2d,
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
 
 
782
  flex_block_mask: BlockMask | None = None,
783
  output_attentions: bool = False,
784
  output_s_max: bool = False,
785
+ flash_padding_layout: FlashPaddingLayout | None = None,
786
  ) -> tuple[torch.Tensor, torch.Tensor | None, list[torch.Tensor] | None]:
787
  attn_output, attn_weights, s_max = self.attn(
788
  x,
 
791
  flex_block_mask=flex_block_mask,
792
  output_attentions=output_attentions,
793
  output_s_max=output_s_max,
794
+ flash_padding_layout=flash_padding_layout,
795
  )
796
  x = x + self.dropout(attn_output) / self.scaling_factor
797
  x = x + self.dropout(self.ffn(x)) / self.scaling_factor
 
894
  # Match the pinned Biohub Transformers contract: a supplied sequence_id
895
  # is authoritative and must encode padding as -1. attention_mask is
896
  # ignored in that mode rather than intersected with the chain mask.
897
+ attention_mask_2d, attention_mask_4d, flex_block_mask = self._prepare_attention_masks(
898
+ attention_mask=attention_mask,
899
+ sequence_id=sequence_id,
900
+ batch_size=x.shape[0],
901
+ seq_len=x.shape[1],
902
+ device=x.device,
903
+ dtype=x.dtype,
904
+ output_attentions=bool(output_attentions),
905
+ )
906
+ # A call that returns attention weights runs eager attention and needs no layout.
907
+ flash_padding_layout = (
908
+ None
909
+ if output_attentions
910
+ else get_flash_padding_layout(self.attention_backend, attention_mask_2d)
911
  )
912
 
913
  for layer_index, block in enumerate(self.blocks):
 
931
  flex_block_mask=flex_block_mask,
932
  output_attentions=output_attentions,
933
  output_s_max=output_s_max,
934
+ flash_padding_layout=flash_padding_layout,
935
  )
936
  else:
937
  x, attn_weights, s_max = block(
 
941
  flex_block_mask=flex_block_mask,
942
  output_attentions=output_attentions,
943
  output_s_max=output_s_max,
944
+ flash_padding_layout=flash_padding_layout,
945
  )
946
 
947
  if attentions is not None:
 
1091
  "flash_attention_2",
1092
  "flash_attention_3",
1093
  )
1094
+ _fastplms_attention_auto_order = ("sdpa",)
1095
 
1096
  def __init__(self, config: ESMplusplusConfig, *args: object, **kwargs: object) -> None:
1097
  super().__init__(config, *args, **kwargs)
fastplms/registry.py CHANGED
The diff for this file is too large to render. See raw diff
 
fastplms_bundle.py CHANGED
The diff for this file is too large to render. See raw diff
 
modeling_fastplms.py CHANGED
@@ -13,7 +13,7 @@ from zipfile import ZIP_DEFLATED, ZipFile
13
 
14
  from .fastplms_bundle import RUNTIME_DATA, RUNTIME_HASH
15
 
16
- if RUNTIME_HASH != "18a1768d1a03557ecd0fbc7ccf8f12cdb8d18d23b190e9d56ccf7c0c0225112e":
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 != "1b702f30eb56fa1c6b436c423cd901ab7583868bdc38e15da835a315d3ce15b8":
17
  raise RuntimeError("FastPLMs runtime identity differs from the bridge.")
18
 
19
  _RUNTIME_TEMPORARIES = []