lhallee commited on
Commit
2ca20a1
·
verified ·
1 Parent(s): a6afadf

Update FastPLMs files from Synthyra/FastPLMs 32d9951 (PR 52: ESMFold2 noise cap, DPLM2 parity)

Browse files
README.md CHANGED
@@ -36,7 +36,7 @@ python -m pip install -r \
36
  The FastPLMs implementation itself is embedded in the model repository.
37
  Transformers loads it through `trust_remote_code=True`.
38
 
39
- This model requires Python 3.11-3.14, PyTorch 2.13, and Transformers 5.13.
40
 
41
  The artifact requirements include the structure dependencies.
42
 
@@ -117,7 +117,7 @@ print(token_output.logits.shape) # (b, l, 3)
117
  Install the training dependencies. Then attach LoRA to the loaded checkpoint:
118
 
119
  ```bash
120
- python -m pip install "datasets>=4.8,<5" "peft>=0.19,<0.20"
121
  ```
122
 
123
  ```python
 
36
  The FastPLMs implementation itself is embedded in the model repository.
37
  Transformers loads it through `trust_remote_code=True`.
38
 
39
+ This model requires Python 3.12-3.14, PyTorch 2.14, and Transformers 5.17.
40
 
41
  The artifact requirements include the structure dependencies.
42
 
 
117
  Install the training dependencies. Then attach LoRA to the loaded checkpoint:
118
 
119
  ```bash
120
+ python -m pip install "datasets>=5.0" "peft>=0.21"
121
  ```
122
 
123
  ```python
fastplms/atomic_files.py ADDED
@@ -0,0 +1,65 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Replace a file in one step: readers see the old bytes or the new bytes, never a partial write."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import os
6
+ import tempfile
7
+
8
+ from collections.abc import Iterator
9
+ from contextlib import contextmanager
10
+ from pathlib import Path
11
+ from typing import IO, Any
12
+
13
+
14
+ def write_bytes_atomically(path: Path, payload: bytes, *, create_parent: bool = False) -> None:
15
+ """Write ``payload`` to a temporary file beside ``path``, flush it to disk, then rename it over ``path``.
16
+
17
+ A failed write removes the temporary file and leaves ``path`` untouched. The parent directory must
18
+ exist unless ``create_parent`` is true. The temporary file sits in the same directory so the rename
19
+ never crosses a file system.
20
+ """
21
+
22
+ with _staged_handle(path, "wb", create_parent=create_parent) as handle:
23
+ handle.write(payload)
24
+
25
+
26
+ def write_text_atomically(
27
+ path: Path,
28
+ text: str,
29
+ *,
30
+ encoding: str = "utf-8",
31
+ newline: str | None = None,
32
+ create_parent: bool = False,
33
+ ) -> None:
34
+ """Write ``text`` as ``write_bytes_atomically`` does.
35
+
36
+ ``newline`` is the argument of ``open``: ``None`` translates ``"\\n"`` to the platform separator and
37
+ ``"\\n"`` writes it unchanged.
38
+ """
39
+
40
+ with _staged_handle(
41
+ path, "w", create_parent=create_parent, encoding=encoding, newline=newline
42
+ ) as handle:
43
+ handle.write(text)
44
+
45
+
46
+ @contextmanager
47
+ def _staged_handle(
48
+ path: Path, mode: str, *, create_parent: bool, **open_arguments: Any
49
+ ) -> Iterator[IO[Any]]:
50
+ """Yield a handle to a temporary file; on a clean exit flush it and rename it over ``path``."""
51
+
52
+ if create_parent:
53
+ path.parent.mkdir(parents=True, exist_ok=True)
54
+ descriptor, temporary_name = tempfile.mkstemp(
55
+ prefix=f".{path.name}.", suffix=".tmp", dir=path.parent
56
+ )
57
+ try:
58
+ with os.fdopen(descriptor, mode, **open_arguments) as handle:
59
+ yield handle
60
+ handle.flush()
61
+ os.fsync(handle.fileno())
62
+ os.replace(temporary_name, path)
63
+ except BaseException:
64
+ Path(temporary_name).unlink(missing_ok=True)
65
+ raise
fastplms/attention/__init__.py CHANGED
@@ -22,6 +22,7 @@ from ._core import (
22
  canonical_checkpoint_attention_backend,
23
  clear_flex_attention_caches,
24
  create_block_mask,
 
25
  flex_attention,
26
  get_attention_mask,
27
  get_attn_implementation,
@@ -66,6 +67,7 @@ __all__ = [
66
  "canonical_checkpoint_attention_backend",
67
  "clear_flex_attention_caches",
68
  "create_block_mask",
 
69
  "flex_attention",
70
  "get_attention_mask",
71
  "get_attn_implementation",
 
22
  canonical_checkpoint_attention_backend,
23
  clear_flex_attention_caches,
24
  create_block_mask,
25
+ flash_kernel_unsupported_reason,
26
  flex_attention,
27
  get_attention_mask,
28
  get_attn_implementation,
 
67
  "canonical_checkpoint_attention_backend",
68
  "clear_flex_attention_caches",
69
  "create_block_mask",
70
+ "flash_kernel_unsupported_reason",
71
  "flex_attention",
72
  "get_attention_mask",
73
  "get_attn_implementation",
fastplms/attention/_core.py CHANGED
@@ -16,10 +16,15 @@ from dataclasses import dataclass
16
  from enum import Enum
17
  from threading import RLock
18
  from types import MappingProxyType
 
19
  from einops import rearrange
20
  from torch.nn import functional as F
21
 
22
- from ._kernel_lock import load_locked_kernel
 
 
 
 
23
 
24
 
25
  try:
@@ -30,12 +35,16 @@ except ImportError:
30
  BlockMask = None
31
 
32
  _MAX_FLEX_CACHE_ENTRIES = 128
33
- _compiled_flex_attention: OrderedDict[tuple, object] = OrderedDict()
34
- _flex_block_masks: OrderedDict[tuple, BlockMask] = OrderedDict()
35
  _flex_cache_lock = RLock()
36
 
 
 
37
 
38
- def _remember(cache: OrderedDict, key: tuple, value):
 
 
39
  """Insert an item into a bounded least-recently-used cache."""
40
  cache[key] = value
41
  cache.move_to_end(key)
@@ -64,7 +73,7 @@ def _get_flex_attention_fn(
64
  shape: tuple[int, ...] | None = None,
65
  sequence_lengths: tuple[int, ...] | None = None,
66
  mask_semantics: str = "padding",
67
- ):
68
  """Return a compiled Flex callable for an explicit execution signature.
69
 
70
  Compilation depends on execution shape, device, dtype, and mask semantics.
@@ -116,15 +125,17 @@ def _get_flex_block_mask(
116
  because compiled Flex plans can specialize on it even though the pattern
117
  tensor itself is boolean or integer.
118
  """
 
 
119
  if create_block_mask is None:
120
  raise RuntimeError(
121
  "'flex_attention' was requested, but torch.create_block_mask is unavailable."
122
  )
123
- pattern = mask_pattern.detach().to(device=device).contiguous() # mask_pattern.shape
124
  # One device-to-host transfer is required for an exact cache identity. Use
125
  # the contiguous buffer directly instead of materializing one Python int
126
  # per byte, which is prohibitively expensive for long batched sequences.
127
- host_pattern = pattern.to(device="cpu").contiguous() # mask_pattern.shape
128
  pattern_bytes = host_pattern.view(torch.uint8).numpy().tobytes(order="C") # bytes
129
  cache_key = (
130
  str(device),
@@ -153,7 +164,7 @@ def _get_flex_block_mask(
153
 
154
  # Hugging Face `kernels` exposes slightly different APIs for FlashAttention 2
155
  # and 3. Detect the loaded variant once so every caller uses the same dispatch.
156
- def _infer_kernels_flash_variant(kernel) -> str | None:
157
  if hasattr(kernel, "fwd") and hasattr(kernel, "varlen_fwd"):
158
  return "flash_attn2"
159
  if hasattr(kernel, "flash_attn_func") and hasattr(kernel, "flash_attn_varlen_func"):
@@ -273,6 +284,40 @@ def _ensure_flash_kernels_loaded(implementation: str) -> tuple[object, str]:
273
  return loaded
274
 
275
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
276
  def _kernels_flash_forward(
277
  query_states: torch.Tensor,
278
  key_states: torch.Tensor,
@@ -371,7 +416,7 @@ def _kernels_flash_varlen_forward(
371
  # before the kernel call and restore the original padded batch shape afterward.
372
  class IndexFirstAxis(torch.autograd.Function):
373
  @staticmethod
374
- def forward(ctx, input, indices) -> torch.Tensor:
375
  # input: (n, ...); indices: (m,)
376
  ctx.save_for_backward(indices)
377
  if input.ndim < 2:
@@ -391,7 +436,7 @@ class IndexFirstAxis(torch.autograd.Function):
391
  ).reshape(-1, *other_shape)
392
 
393
  @staticmethod
394
- def backward(ctx, grad_output) -> tuple[torch.Tensor, None]:
395
  # grad_output: (m, ...)
396
  (indices,) = ctx.saved_tensors
397
  if grad_output.ndim < 2:
@@ -412,7 +457,9 @@ class IndexFirstAxis(torch.autograd.Function):
412
 
413
  class IndexPutFirstAxis(torch.autograd.Function):
414
  @staticmethod
415
- def forward(ctx, values, indices, first_axis_dim) -> torch.Tensor:
 
 
416
  # values: (m, ...); indices: (m,)
417
  ctx.save_for_backward(indices)
418
  if indices.ndim != 1:
@@ -432,7 +479,7 @@ class IndexPutFirstAxis(torch.autograd.Function):
432
  return output # (n, ...)
433
 
434
  @staticmethod
435
- def backward(ctx, grad_output) -> tuple[torch.Tensor, None, None]:
436
  # grad_output: (n, ...)
437
  (indices,) = ctx.saved_tensors
438
  return grad_output[indices], None, None # (m, ...), None, None
@@ -451,7 +498,7 @@ def _select_first_axis(states: torch.Tensor, indices: torch.Tensor) -> torch.Ten
451
  # states: (n, ...); indices: (m,)
452
  if states.requires_grad:
453
  selected: torch.Tensor = index_first_axis(states, indices) # (m, ...)
454
- return selected
455
  return states[indices] # (m, ...)
456
 
457
 
@@ -528,7 +575,7 @@ def _unpad_input(
528
  indices,
529
  (cu_seqlens, cu_seqlens),
530
  (max_seqlen, max_seqlen),
531
- )
532
 
533
 
534
  def _validate_flash_padding_mask(
@@ -786,7 +833,7 @@ def resolve_attention_backend(
786
  return resolved
787
 
788
 
789
- def get_attn_implementation(config) -> str:
790
  """Read the Transformers attention setting, defaulting to SDPA."""
791
  requested = getattr(config, "_attn_implementation", None)
792
  if requested is None:
@@ -794,7 +841,7 @@ def get_attn_implementation(config) -> str:
794
  return resolve_attention_backend(requested).value
795
 
796
 
797
- def set_config_attn_implementation(config, implementation: str) -> str:
798
  """Set both the Transformers field and the internal dispatch field."""
799
  resolved = resolve_attention_backend(implementation).value
800
  if hasattr(config, "_attn_implementation_internal"):
@@ -823,7 +870,7 @@ def get_attention_mask(
823
  """
824
  # attention_mask: (b, l) or None
825
  if attention_mask is None:
826
- return None, None, None
827
 
828
  if attention_mask.ndim != 2:
829
  raise ValueError(
@@ -850,12 +897,15 @@ def get_attention_mask(
850
  raise RuntimeError(
851
  "'flex_attention' was requested, but torch.create_block_mask is unavailable."
852
  )
853
- def mask_mod(batch_idx, head_idx, q_idx, kv_idx):
 
 
 
854
  del head_idx, q_idx
855
  # Match eager and SDPA: padding masks suppress invalid keys only.
856
  # Invalid queries still attend to real keys and therefore remain
857
  # finite; downstream residue masks exclude their outputs.
858
- return attention_mask_2d[batch_idx, kv_idx]
859
 
860
  flex_block_mask = _get_flex_block_mask(
861
  mask_pattern=attention_mask_2d,
 
16
  from enum import Enum
17
  from threading import RLock
18
  from types import MappingProxyType
19
+ from typing import TYPE_CHECKING, Any, TypeVar
20
  from einops import rearrange
21
  from torch.nn import functional as F
22
 
23
+ from ._kernel_lock import load_locked_kernel, locked_variant_for_this_system
24
+
25
+
26
+ if TYPE_CHECKING:
27
+ from transformers import PretrainedConfig
28
 
29
 
30
  try:
 
35
  BlockMask = None
36
 
37
  _MAX_FLEX_CACHE_ENTRIES = 128
38
+ _compiled_flex_attention: OrderedDict[tuple[Any, ...], object] = OrderedDict()
39
+ _flex_block_masks: OrderedDict[tuple[Any, ...], BlockMask] = OrderedDict()
40
  _flex_cache_lock = RLock()
41
 
42
+ _CacheValue = TypeVar("_CacheValue")
43
+
44
 
45
+ def _remember(
46
+ cache: OrderedDict[tuple[Any, ...], _CacheValue], key: tuple[Any, ...], value: _CacheValue
47
+ ) -> _CacheValue:
48
  """Insert an item into a bounded least-recently-used cache."""
49
  cache[key] = value
50
  cache.move_to_end(key)
 
73
  shape: tuple[int, ...] | None = None,
74
  sequence_lengths: tuple[int, ...] | None = None,
75
  mask_semantics: str = "padding",
76
+ ) -> Callable[..., Any] | None:
77
  """Return a compiled Flex callable for an explicit execution signature.
78
 
79
  Compilation depends on execution shape, device, dtype, and mask semantics.
 
125
  because compiled Flex plans can specialize on it even though the pattern
126
  tensor itself is boolean or integer.
127
  """
128
+ # mask_pattern: (b, l), the padding mask of the call
129
+ # mask_mod: ((), (), (), ()) -> (), taking the 0-d batch, head, query and key-value indices
130
  if create_block_mask is None:
131
  raise RuntimeError(
132
  "'flex_attention' was requested, but torch.create_block_mask is unavailable."
133
  )
134
+ pattern = mask_pattern.detach().to(device=device).contiguous() # (b, l)
135
  # One device-to-host transfer is required for an exact cache identity. Use
136
  # the contiguous buffer directly instead of materializing one Python int
137
  # per byte, which is prohibitively expensive for long batched sequences.
138
+ host_pattern = pattern.to(device="cpu").contiguous() # (b, l)
139
  pattern_bytes = host_pattern.view(torch.uint8).numpy().tobytes(order="C") # bytes
140
  cache_key = (
141
  str(device),
 
164
 
165
  # Hugging Face `kernels` exposes slightly different APIs for FlashAttention 2
166
  # and 3. Detect the loaded variant once so every caller uses the same dispatch.
167
+ def _infer_kernels_flash_variant(kernel: object) -> str | None:
168
  if hasattr(kernel, "fwd") and hasattr(kernel, "varlen_fwd"):
169
  return "flash_attn2"
170
  if hasattr(kernel, "flash_attn_func") and hasattr(kernel, "flash_attn_varlen_func"):
 
284
  return loaded
285
 
286
 
287
+ def flash_kernel_unsupported_reason(implementation: str) -> str | None:
288
+ """Return why this platform or GPU cannot run a manifest-locked kernel, or None if it can.
289
+
290
+ The answer comes from kernels.lock, the manifest, and the current CUDA device, so nothing
291
+ is downloaded or imported. A caller choosing between FlashAttention versions may skip a
292
+ kernel for this reason alone; any other loading failure must raise.
293
+ """
294
+ from fastplms.registry import get_model_registry
295
+
296
+ kernel_spec = get_model_registry().attention_kernels[implementation]
297
+ pinned = f"{kernel_spec.repository}@{kernel_spec.revision}"
298
+ if locked_variant_for_this_system(kernel_spec.repository, kernel_spec.revision) is None:
299
+ return f"kernels.lock pins no build of {pinned} for this platform."
300
+
301
+ # `kernels` matches a build to the PyTorch backend, so on a PyTorch without CUDA the
302
+ # build found above is a CPU or XPU one and needs no GPU.
303
+ if torch.version.cuda is None:
304
+ return None
305
+
306
+ if not torch.cuda.is_available():
307
+ return f"{pinned} runs on a CUDA device, and none is visible."
308
+
309
+ capability = torch.cuda.get_device_capability()
310
+ if capability < kernel_spec.min_cuda_capability:
311
+ required = ".".join(str(part) for part in kernel_spec.min_cuda_capability)
312
+ observed = ".".join(str(part) for part in capability)
313
+ return (
314
+ f"{pinned} requires CUDA compute capability {required} or newer; "
315
+ f"this GPU has {observed}."
316
+ )
317
+
318
+ return None
319
+
320
+
321
  def _kernels_flash_forward(
322
  query_states: torch.Tensor,
323
  key_states: torch.Tensor,
 
416
  # before the kernel call and restore the original padded batch shape afterward.
417
  class IndexFirstAxis(torch.autograd.Function):
418
  @staticmethod
419
+ def forward(ctx: Any, input: torch.Tensor, indices: torch.Tensor) -> torch.Tensor:
420
  # input: (n, ...); indices: (m,)
421
  ctx.save_for_backward(indices)
422
  if input.ndim < 2:
 
436
  ).reshape(-1, *other_shape)
437
 
438
  @staticmethod
439
+ def backward(ctx: Any, grad_output: torch.Tensor) -> tuple[torch.Tensor, None]:
440
  # grad_output: (m, ...)
441
  (indices,) = ctx.saved_tensors
442
  if grad_output.ndim < 2:
 
457
 
458
  class IndexPutFirstAxis(torch.autograd.Function):
459
  @staticmethod
460
+ def forward(
461
+ ctx: Any, values: torch.Tensor, indices: torch.Tensor, first_axis_dim: int
462
+ ) -> torch.Tensor:
463
  # values: (m, ...); indices: (m,)
464
  ctx.save_for_backward(indices)
465
  if indices.ndim != 1:
 
479
  return output # (n, ...)
480
 
481
  @staticmethod
482
+ def backward(ctx: Any, grad_output: torch.Tensor) -> tuple[torch.Tensor, None, None]:
483
  # grad_output: (n, ...)
484
  (indices,) = ctx.saved_tensors
485
  return grad_output[indices], None, None # (m, ...), None, None
 
498
  # states: (n, ...); indices: (m,)
499
  if states.requires_grad:
500
  selected: torch.Tensor = index_first_axis(states, indices) # (m, ...)
501
+ return selected # (m, ...)
502
  return states[indices] # (m, ...)
503
 
504
 
 
575
  indices,
576
  (cu_seqlens, cu_seqlens),
577
  (max_seqlen, max_seqlen),
578
+ ) # (t, h, d), (t, h, d), (t, h, d), (t,), ((b + 1,), (b + 1,)), (int, int)
579
 
580
 
581
  def _validate_flash_padding_mask(
 
833
  return resolved
834
 
835
 
836
+ def get_attn_implementation(config: PretrainedConfig) -> str:
837
  """Read the Transformers attention setting, defaulting to SDPA."""
838
  requested = getattr(config, "_attn_implementation", None)
839
  if requested is None:
 
841
  return resolve_attention_backend(requested).value
842
 
843
 
844
+ def set_config_attn_implementation(config: PretrainedConfig, implementation: str) -> str:
845
  """Set both the Transformers field and the internal dispatch field."""
846
  resolved = resolve_attention_backend(implementation).value
847
  if hasattr(config, "_attn_implementation_internal"):
 
870
  """
871
  # attention_mask: (b, l) or None
872
  if attention_mask is None:
873
+ return None, None, None # (None, None, None)
874
 
875
  if attention_mask.ndim != 2:
876
  raise ValueError(
 
897
  raise RuntimeError(
898
  "'flex_attention' was requested, but torch.create_block_mask is unavailable."
899
  )
900
+ def mask_mod(
901
+ batch_idx: torch.Tensor, head_idx: torch.Tensor, q_idx: torch.Tensor, kv_idx: torch.Tensor
902
+ ) -> torch.Tensor:
903
+ # batch_idx, head_idx, q_idx, kv_idx: ()
904
  del head_idx, q_idx
905
  # Match eager and SDPA: padding masks suppress invalid keys only.
906
  # Invalid queries still attend to real keys and therefore remain
907
  # finite; downstream residue masks exclude their outputs.
908
+ return attention_mask_2d[batch_idx, kv_idx] # ()
909
 
910
  flex_block_mask = _get_flex_block_mask(
911
  mask_pattern=attention_mask_2d,
fastplms/attention/_kernel_lock.py CHANGED
@@ -2,13 +2,26 @@
2
 
3
  from __future__ import annotations
4
 
 
5
  import json
6
  import os
 
7
 
8
  from pathlib import Path
9
  from typing import Any
10
 
11
 
 
 
 
 
 
 
 
 
 
 
 
12
  def require_kernels_package() -> None:
13
  """Fail early when the precompiled-kernel runtime is not installed."""
14
  try:
@@ -51,6 +64,68 @@ def _locked_entry(lock_path: Path, repository: str) -> dict[str, Any]:
51
  return matches[0]
52
 
53
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
54
  def _offline_mode() -> bool:
55
  """Return whether Hub access was explicitly disabled for this process."""
56
 
@@ -78,7 +153,7 @@ def _offline_snapshot_path(repository: str, revision: str) -> Path:
78
  if not snapshot.is_dir():
79
  raise RuntimeError(
80
  f"The exact offline kernel snapshot {repository}@{revision} is not cached under "
81
- f"{cache_root}. Run `kernels download` before enabling offline mode."
82
  )
83
  if repository_root not in snapshot.resolve().parents:
84
  raise RuntimeError(f"Refusing kernel snapshot outside its cache repository: {snapshot}")
@@ -88,7 +163,7 @@ def _offline_snapshot_path(repository: str, revision: str) -> Path:
88
  def _load_offline_locked_kernel(
89
  repository: str,
90
  revision: str,
91
- variant_locks: dict[str, object],
92
  ) -> object:
93
  """Validate and import the one compatible variant from a sparse Hub snapshot."""
94
  snapshot = _offline_snapshot_path(repository, revision)
@@ -97,7 +172,7 @@ def _load_offline_locked_kernel(
97
  raise RuntimeError(f"The cached kernel snapshot has no build directory: {snapshot}")
98
 
99
  cached_names = sorted(entry.name for entry in build_root.iterdir() if entry.is_dir())
100
- unexpected = sorted(set(cached_names).difference(variant_locks))
101
  if unexpected:
102
  raise RuntimeError(
103
  f"The cached {repository}@{revision} snapshot contains unlocked variants: "
@@ -106,7 +181,6 @@ def _load_offline_locked_kernel(
106
 
107
  try:
108
  from kernels import get_local_kernel
109
- from kernels.utils import validate_kernel
110
  from kernels.variants import get_variants_local, resolve_variants
111
  except ImportError as error:
112
  raise RuntimeError(
@@ -130,47 +204,91 @@ def _load_offline_locked_kernel(
130
  f"found {names}."
131
  )
132
  variant_name = compatible[0].variant_str
133
- variant_lock = variant_locks.get(variant_name)
134
- expected_hash = getattr(variant_lock, "hash", None)
135
- if not isinstance(expected_hash, str) or not expected_hash.startswith("sha256-"):
136
- raise RuntimeError(f"The kernel lock for {variant_name} has no valid SHA-256 digest.")
137
-
138
- # Hash validation deliberately happens before import. This operates on the
139
- # sparse snapshot produced by `kernels download` and avoids Hub 1.23's
140
- # full-snapshot completeness check in offline mode.
141
- validate_kernel(repo_path=snapshot, variant=variant_name, hash=expected_hash)
142
  return get_local_kernel(build_root / variant_name)
143
 
144
 
145
- def load_locked_kernel(repository: str, revision: str) -> object:
146
- """Download, hash-validate, then import one immutable precompiled kernel."""
147
- require_kernels_package()
148
  try:
149
- from kernels import get_local_kernel, install_kernel
150
- from kernels.lockfile import KernelLock
151
  except ImportError as error:
152
  raise RuntimeError(
153
  "Precompiled FlashAttention requires requirements/features/flash.in."
154
  ) from error
155
 
156
- lock_path = _kernel_lock_path()
157
- kernel_lock = KernelLock.from_json(_locked_entry(lock_path, repository))
158
- if kernel_lock.sha != revision:
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
159
  raise RuntimeError(
160
  f"The typed manifest pins {repository}@{revision}, but kernels.lock pins "
161
- f"{kernel_lock.sha}."
162
  )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
163
 
164
  if _offline_mode():
165
- return _load_offline_locked_kernel(repository, revision, kernel_lock.variants)
166
-
167
- # `install_kernel` downloads data without importing it and validates the
168
- # selected build against the tracked variant hash. Only then is the exact
169
- # validated path imported directly. Offline mode uses the sparse-cache
170
- # resolver above because Hub 1.23 rejects partial snapshots as incomplete.
171
- validated_path = install_kernel(
172
- repository,
173
- revision=kernel_lock.sha,
174
- variant_locks=kernel_lock.variants,
 
 
 
175
  )
176
- return get_local_kernel(validated_path)
 
 
2
 
3
  from __future__ import annotations
4
 
5
+ import hashlib
6
  import json
7
  import os
8
+ import re
9
 
10
  from pathlib import Path
11
  from typing import Any
12
 
13
 
14
+ # kernels.lock records, per locked build variant, one SHA-256 over the variant's relative
15
+ # file paths and the Git or Git LFS object IDs of their contents. kernels 0.17 no longer
16
+ # writes or checks these digests, so this module checks them before anything is imported.
17
+ _VARIANT_HASH_TYPE = "git_lfs_concat"
18
+ _VARIANT_DIGEST = re.compile(r"sha256-[0-9a-f]{64}")
19
+ # The Hub cache names a Git blob by its 40-character SHA-1 object ID and Git LFS content
20
+ # by its 64-character SHA-256.
21
+ _GIT_OBJECT_ID_LENGTH = 40
22
+ _LFS_OBJECT_ID_LENGTH = 64
23
+
24
+
25
  def require_kernels_package() -> None:
26
  """Fail early when the precompiled-kernel runtime is not installed."""
27
  try:
 
64
  return matches[0]
65
 
66
 
67
+ def _locked_variant_digests(entry: dict[str, Any]) -> dict[str, str]:
68
+ """Each locked build variant of one kernels.lock entry, mapped to its digest."""
69
+ variants = entry.get("variants")
70
+ if not isinstance(variants, dict) or not variants:
71
+ raise RuntimeError(f"kernels.lock locks no build variants for {entry.get('repo_id')!r}.")
72
+ digests: dict[str, str] = {}
73
+ for variant_name, variant_lock in variants.items():
74
+ expected_hash = variant_lock.get("hash") if isinstance(variant_lock, dict) else None
75
+ if (
76
+ not isinstance(expected_hash, str)
77
+ or _VARIANT_DIGEST.fullmatch(expected_hash) is None
78
+ or variant_lock.get("hash_type") != _VARIANT_HASH_TYPE
79
+ ):
80
+ raise RuntimeError(f"The kernel lock for {variant_name} has no valid SHA-256 digest.")
81
+ digests[variant_name] = expected_hash
82
+ return digests
83
+
84
+
85
+ def _git_blob_object_id(contents: bytes) -> bytes:
86
+ """Return the SHA-1 object ID Git assigns to a blob with these contents."""
87
+ return hashlib.sha1(b"blob %d\0" % len(contents) + contents).digest()
88
+
89
+
90
+ def validate_variant_digest(snapshot: Path, variant_name: str, expected_hash: str) -> None:
91
+ """Check one cached build variant against its kernels.lock digest before import.
92
+
93
+ Snapshot files link into the Hub cache's content-addressed blobs. Each linked file
94
+ contributes its path relative to the variant, then the object ID its contents hash to:
95
+ a Git blob ID when the cache names the blob by SHA-1, a Git LFS SHA-256 otherwise.
96
+ Files that are not links are skipped, because importing a kernel writes bytecode
97
+ beside it.
98
+ """
99
+ variant_root = snapshot / "build" / variant_name
100
+ linked_files: list[tuple[bytes, Path]] = []
101
+ for directory, _, file_names in os.walk(variant_root):
102
+ for file_name in file_names:
103
+ path = Path(directory) / file_name
104
+ if path.is_symlink():
105
+ relative_name = path.relative_to(variant_root).as_posix().encode("utf-8")
106
+ linked_files.append((relative_name, path))
107
+
108
+ digest = hashlib.sha256()
109
+ for relative_name, path in sorted(linked_files):
110
+ contents = path.read_bytes()
111
+ object_id_length = len(path.resolve().name)
112
+ if object_id_length == _GIT_OBJECT_ID_LENGTH:
113
+ object_id = _git_blob_object_id(contents)
114
+ elif object_id_length == _LFS_OBJECT_ID_LENGTH:
115
+ object_id = hashlib.sha256(contents).digest()
116
+ else:
117
+ raise RuntimeError(f"Unexpected Hub cache blob name behind {path}.")
118
+ digest.update(relative_name)
119
+ digest.update(object_id)
120
+
121
+ received_hash = f"sha256-{digest.hexdigest()}"
122
+ if received_hash != expected_hash:
123
+ raise RuntimeError(
124
+ f"The cached kernel variant {variant_name} hashes to {received_hash}, but "
125
+ f"kernels.lock records {expected_hash}."
126
+ )
127
+
128
+
129
  def _offline_mode() -> bool:
130
  """Return whether Hub access was explicitly disabled for this process."""
131
 
 
153
  if not snapshot.is_dir():
154
  raise RuntimeError(
155
  f"The exact offline kernel snapshot {repository}@{revision} is not cached under "
156
+ f"{cache_root}. Load the kernel once with Hub access before enabling offline mode."
157
  )
158
  if repository_root not in snapshot.resolve().parents:
159
  raise RuntimeError(f"Refusing kernel snapshot outside its cache repository: {snapshot}")
 
163
  def _load_offline_locked_kernel(
164
  repository: str,
165
  revision: str,
166
+ variant_digests: dict[str, str],
167
  ) -> object:
168
  """Validate and import the one compatible variant from a sparse Hub snapshot."""
169
  snapshot = _offline_snapshot_path(repository, revision)
 
172
  raise RuntimeError(f"The cached kernel snapshot has no build directory: {snapshot}")
173
 
174
  cached_names = sorted(entry.name for entry in build_root.iterdir() if entry.is_dir())
175
+ unexpected = sorted(set(cached_names).difference(variant_digests))
176
  if unexpected:
177
  raise RuntimeError(
178
  f"The cached {repository}@{revision} snapshot contains unlocked variants: "
 
181
 
182
  try:
183
  from kernels import get_local_kernel
 
184
  from kernels.variants import get_variants_local, resolve_variants
185
  except ImportError as error:
186
  raise RuntimeError(
 
204
  f"found {names}."
205
  )
206
  variant_name = compatible[0].variant_str
207
+
208
+ # Hash validation deliberately happens before import. It reads the sparse snapshot
209
+ # directly, so no Hub API judges whether the snapshot is complete.
210
+ validate_variant_digest(snapshot, variant_name, variant_digests[variant_name])
 
 
 
 
 
211
  return get_local_kernel(build_root / variant_name)
212
 
213
 
214
+ def _compatible_locked_variants(repository: str, variant_names: list[str]) -> list[str]:
215
+ """Return the locked build variants `kernels` can load on this system, preferred first."""
 
216
  try:
217
+ from kernels.variants import parse_variant, resolve_variants
 
218
  except ImportError as error:
219
  raise RuntimeError(
220
  "Precompiled FlashAttention requires requirements/features/flash.in."
221
  ) from error
222
 
223
+ try:
224
+ locked_variants = [parse_variant(variant_name) for variant_name in variant_names]
225
+ except ValueError as error:
226
+ raise RuntimeError(f"kernels.lock contains an invalid {repository} variant.") from error
227
+ compatible, _ = resolve_variants(locked_variants)
228
+ return [variant.variant_str for variant in compatible]
229
+
230
+
231
+ def _preferred_locked_variant(repository: str, revision: str, variant_names: list[str]) -> str:
232
+ """Return the locked build variant `kernels` prefers on this system."""
233
+ compatible = _compatible_locked_variants(repository, variant_names)
234
+ if not compatible:
235
+ raise RuntimeError(
236
+ f"kernels.lock locks no build of {repository}@{revision} for this system; "
237
+ f"locked variants: {', '.join(sorted(variant_names))}."
238
+ )
239
+ return compatible[0]
240
+
241
+
242
+ def _pinned_variant_digests(repository: str, revision: str) -> dict[str, str]:
243
+ """Return the locked build digests of a kernel whose kernels.lock entry pins `revision`."""
244
+ entry = _locked_entry(_kernel_lock_path(), repository)
245
+ locked_revision = entry.get("sha")
246
+ if locked_revision != revision:
247
  raise RuntimeError(
248
  f"The typed manifest pins {repository}@{revision}, but kernels.lock pins "
249
+ f"{locked_revision}."
250
  )
251
+ return _locked_variant_digests(entry)
252
+
253
+
254
+ def locked_variant_for_this_system(repository: str, revision: str) -> str | None:
255
+ """Return the locked build `kernels` would load here, or None when kernels.lock pins none.
256
+
257
+ Only kernels.lock is read, so a caller can tell a platform the lock does not cover
258
+ apart from a download, digest, or import failure.
259
+ """
260
+ variant_digests = _pinned_variant_digests(repository, revision)
261
+ compatible = _compatible_locked_variants(repository, list(variant_digests))
262
+ return compatible[0] if compatible else None
263
+
264
+
265
+ def load_locked_kernel(repository: str, revision: str) -> object:
266
+ """Download, hash-validate, then import one immutable precompiled kernel."""
267
+ require_kernels_package()
268
+ try:
269
+ from huggingface_hub import snapshot_download
270
+ from kernels import get_local_kernel
271
+ except ImportError as error:
272
+ raise RuntimeError(
273
+ "Precompiled FlashAttention requires requirements/features/flash.in."
274
+ ) from error
275
+
276
+ variant_digests = _pinned_variant_digests(repository, revision)
277
 
278
  if _offline_mode():
279
+ return _load_offline_locked_kernel(repository, revision, variant_digests)
280
+
281
+ # Only a locked build can be selected, and only its files are downloaded, at the
282
+ # immutable revision. The download imports nothing; the digest check runs first.
283
+ variant_name = _preferred_locked_variant(repository, revision, list(variant_digests))
284
+ snapshot = Path(
285
+ snapshot_download(
286
+ repository,
287
+ repo_type="kernel",
288
+ revision=revision,
289
+ allow_patterns=[f"build/{variant_name}/*"],
290
+ cache_dir=os.environ.get("KERNELS_CACHE") or None,
291
+ )
292
  )
293
+ validate_variant_digest(snapshot, variant_name, variant_digests[variant_name])
294
+ return get_local_kernel(snapshot / "build" / variant_name)
fastplms/attention/interfaces.py CHANGED
@@ -7,7 +7,7 @@ import torch
7
  from collections.abc import Mapping
8
  from functools import partial
9
  from typing import Any
10
- from transformers import AttentionInterface, AttentionMaskInterface
11
 
12
  from ._auto import (
13
  AUTO_ATTENTION,
@@ -94,7 +94,7 @@ class FastPLMsAttentionMixin:
94
 
95
  _supports_sdpa = True
96
  _supports_flex_attn = True
97
- # Transformers 5.13 uses the singular flag during model construction. A
98
  # family opts in only when its manifest entry advertises at least one of
99
  # the two FastPLMs kernels-only FlashAttention implementations.
100
  _supports_flash_attn = False
@@ -161,7 +161,7 @@ class FastPLMsAttentionMixin:
161
  allow_all_kernels=False,
162
  )
163
 
164
- def __init__(self, config, *args: Any, **kwargs: Any) -> None:
165
  sentinel = object()
166
  internal = getattr(config, "_attn_implementation_internal", sentinel)
167
  stored = getattr(config, "_attn_implementation", None) if internal is sentinel else internal
@@ -355,7 +355,7 @@ def _resolve_auto_attention_before_forward(
355
  def validate_transformers_attention_interfaces() -> None:
356
  """Verify that Transformers exposes functions and masks for every backend.
357
 
358
- Transformers 5.13 registers these canonical names. The FastPLMs function
359
  overrides remain instance-local and do not replace process-global handlers.
360
  """
361
  function_registry = FASTPLMS_ATTENTION_FUNCTIONS
 
7
  from collections.abc import Mapping
8
  from functools import partial
9
  from typing import Any
10
+ from transformers import AttentionInterface, AttentionMaskInterface, PretrainedConfig
11
 
12
  from ._auto import (
13
  AUTO_ATTENTION,
 
94
 
95
  _supports_sdpa = True
96
  _supports_flex_attn = True
97
+ # Transformers checks the singular flag during model construction. A
98
  # family opts in only when its manifest entry advertises at least one of
99
  # the two FastPLMs kernels-only FlashAttention implementations.
100
  _supports_flash_attn = False
 
161
  allow_all_kernels=False,
162
  )
163
 
164
+ def __init__(self, config: PretrainedConfig, *args: Any, **kwargs: Any) -> None:
165
  sentinel = object()
166
  internal = getattr(config, "_attn_implementation_internal", sentinel)
167
  stored = getattr(config, "_attn_implementation", None) if internal is sentinel else internal
 
355
  def validate_transformers_attention_interfaces() -> None:
356
  """Verify that Transformers exposes functions and masks for every backend.
357
 
358
+ The validated Transformers registers these canonical names. The FastPLMs function
359
  overrides remain instance-local and do not replace process-global handlers.
360
  """
361
  function_registry = FASTPLMS_ATTENTION_FUNCTIONS
fastplms/digests.py ADDED
@@ -0,0 +1,35 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """SHA-256 digests of files and JSON values, the identities FastPLMs records and compares.
2
+
3
+ This file exists twice, byte for byte: here and as ``features/digests.py``. ``features`` loads as a standalone
4
+ package (``foundry.embedding.private_store``) and imports nothing outside itself, so it cannot reach this
5
+ module. ``tests/tier1_unit/test_features_package_is_self_contained.py`` fails when the two differ.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import hashlib
11
+
12
+ from pathlib import Path
13
+ from typing import Any
14
+
15
+ from .json_files import compact_json
16
+
17
+
18
+ FILE_READ_BYTES = 1024 * 1024
19
+
20
+
21
+ def file_sha256(path: str | Path) -> str:
22
+ """Return the SHA-256 of a file's bytes, read in blocks so a checkpoint never sits in memory."""
23
+
24
+ digest = hashlib.sha256()
25
+ with Path(path).open("rb") as handle:
26
+ while block := handle.read(FILE_READ_BYTES):
27
+ digest.update(block)
28
+ return digest.hexdigest()
29
+
30
+
31
+ def json_sha256(value: Any, *, ensure_ascii: bool = True, allow_nan: bool = True) -> str:
32
+ """Return the SHA-256 of ``value`` in its compact, key-sorted JSON form (``compact_json``)."""
33
+
34
+ encoded = compact_json(value, ensure_ascii=ensure_ascii, allow_nan=allow_nan).encode("utf-8")
35
+ return hashlib.sha256(encoded).hexdigest()
fastplms/embeddings/__init__.py CHANGED
@@ -1,6 +1,9 @@
1
  """Ordered, residue-aware protein embedding utilities."""
2
 
3
- from .pooling import POOLING_NAMES, Pooler, pagerank_weights
 
 
 
4
  from .runner import (
5
  EmbeddingMixin,
6
  embed_dataset,
@@ -24,12 +27,21 @@ from .storage import (
24
  tensor_sha256,
25
  update_sqlite_run_metadata,
26
  )
 
 
 
 
 
 
27
  from .types import (
28
  EmbeddingBatch,
29
  EmbeddingInput,
30
  EmbeddingRecord,
31
  EmbeddingResult,
32
  LazyTensorReference,
 
 
 
33
  TensorValue,
34
  )
35
 
@@ -37,17 +49,34 @@ from .types import (
37
  __all__ = [
38
  "DEFAULT_SHARD_SIZE",
39
  "POOLING_NAMES",
 
 
 
40
  "EmbeddingBatch",
41
  "EmbeddingInput",
42
  "EmbeddingMixin",
43
  "EmbeddingRecord",
44
  "EmbeddingResult",
 
 
45
  "LazyTensorReference",
46
  "Pooler",
 
 
 
 
 
 
 
 
 
47
  "TensorValue",
 
48
  "append_sqlite_records",
49
  "convert_legacy_sqlite",
50
  "embed_dataset",
 
 
51
  "garbage_collect_safetensors_generations",
52
  "initialize_sqlite_run",
53
  "iter_fasta",
@@ -57,6 +86,9 @@ __all__ = [
57
  "load_sqlite_result",
58
  "pagerank_weights",
59
  "parse_fasta",
 
 
 
60
  "save_result",
61
  "save_safetensors_result",
62
  "save_sqlite_result",
 
1
  """Ordered, residue-aware protein embedding utilities."""
2
 
3
+ from .feature_runs import embed_into_features
4
+ from .pooling import (
5
+ POOLING_NAMES, POOLING_SEMANTICS_TOKENS, TOKEN_POOLING_NAMES, Pooler, pagerank_weights, pool_token_rows,
6
+ )
7
  from .runner import (
8
  EmbeddingMixin,
9
  embed_dataset,
 
27
  tensor_sha256,
28
  update_sqlite_run_metadata,
29
  )
30
+ from .taps import (
31
+ HiddenTap, LayerAccumulator, ReducedTap, RowSelection, SparseResidueTap, StreamingTap, TapBatch,
32
+ )
33
+ from .token_batches import BatchGeometry, TokenTapExecutor, plan_geometry_batches, plan_token_batches
34
+ from .token_runs import embed_token_features
35
+ from .tokens import ResidueVocabulary
36
  from .types import (
37
  EmbeddingBatch,
38
  EmbeddingInput,
39
  EmbeddingRecord,
40
  EmbeddingResult,
41
  LazyTensorReference,
42
+ TapRecord,
43
+ TapResult,
44
+ TapRunReceipt,
45
  TensorValue,
46
  )
47
 
 
49
  __all__ = [
50
  "DEFAULT_SHARD_SIZE",
51
  "POOLING_NAMES",
52
+ "POOLING_SEMANTICS_TOKENS",
53
+ "TOKEN_POOLING_NAMES",
54
+ "BatchGeometry",
55
  "EmbeddingBatch",
56
  "EmbeddingInput",
57
  "EmbeddingMixin",
58
  "EmbeddingRecord",
59
  "EmbeddingResult",
60
+ "HiddenTap",
61
+ "LayerAccumulator",
62
  "LazyTensorReference",
63
  "Pooler",
64
+ "ReducedTap",
65
+ "ResidueVocabulary",
66
+ "RowSelection",
67
+ "SparseResidueTap",
68
+ "StreamingTap",
69
+ "TapBatch",
70
+ "TapRecord",
71
+ "TapResult",
72
+ "TapRunReceipt",
73
  "TensorValue",
74
+ "TokenTapExecutor",
75
  "append_sqlite_records",
76
  "convert_legacy_sqlite",
77
  "embed_dataset",
78
+ "embed_into_features",
79
+ "embed_token_features",
80
  "garbage_collect_safetensors_generations",
81
  "initialize_sqlite_run",
82
  "iter_fasta",
 
86
  "load_sqlite_result",
87
  "pagerank_weights",
88
  "parse_fasta",
89
+ "plan_geometry_batches",
90
+ "plan_token_batches",
91
+ "pool_token_rows",
92
  "save_result",
93
  "save_safetensors_result",
94
  "save_sqlite_result",
fastplms/embeddings/batches.py CHANGED
@@ -13,7 +13,9 @@ from torch import Tensor
13
  from .identity import _model_device
14
  from .inputs import _planned_batches
15
  from .pooling import Pooler
16
- from .types import EmbeddingBatch, EmbeddingInput, EmbeddingRecord
 
 
17
 
18
 
19
  _MAX_PARTI_RESIDUES = 2_048
@@ -36,7 +38,7 @@ def select_hidden_state_embeddings(
36
  store_all_hidden_states: bool = False,
37
  ) -> Tensor:
38
  """Select one hidden state or stack every state without changing values."""
39
- # last_hidden_state and each hidden_states entry: (b, l, d)
40
  if store_all_hidden_states:
41
  if not hidden_states:
42
  raise ValueError("store_all_hidden_states requires model hidden states.")
@@ -101,6 +103,61 @@ def _biological_residue_mask(
101
  return M # (b, l)
102
 
103
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
104
  def _generic_embedding_batch(
105
  model: Any,
106
  sequences: list[str],
@@ -140,32 +197,13 @@ def _generic_embedding_batch(
140
  if tokenizer is None:
141
  raise ValueError("A tokenizer is required for this model's embedding path.")
142
 
143
- tokenize_kwargs: dict[str, Any] = {
144
- "return_tensors": "pt",
145
- "padding": True,
146
- "truncation": truncate,
147
- }
148
- if max_length is not None and truncate:
149
- # ``max_length`` is a biological-residue limit. Tokenizer limits include
150
- # boundary tokens, so reserve their declared width instead of dropping
151
- # residues at the exact boundary.
152
- special_token_count = 0
153
- num_special_tokens_to_add = getattr(tokenizer, "num_special_tokens_to_add", None)
154
- if callable(num_special_tokens_to_add):
155
- special_token_count = int(num_special_tokens_to_add(pair=False))
156
- tokenize_kwargs["max_length"] = max_length + special_token_count
157
- sequence_tokenizer = getattr(model, "_tokenize_sequence_batch", None)
158
- if callable(sequence_tokenizer):
159
- encoded = sequence_tokenizer(sequences, tokenizer=tokenizer, **tokenize_kwargs)
160
- else:
161
- encoded = tokenizer(sequences, **tokenize_kwargs)
162
- device = _model_device(model)
163
- input_ids = encoded["input_ids"].to(device) # (b, l)
164
- attention_mask = encoded.get( # (b, l)
165
- "attention_mask",
166
- input_ids.new_ones(input_ids.shape),
167
- ).to(device)
168
- M = _biological_residue_mask(input_ids, attention_mask, tokenizer) # (b, l)
169
  if need_attentions:
170
  # Validate l before either the backbone or its quadratic attention graph
171
  # is materialized. M has shape (b, l).
@@ -189,6 +227,24 @@ def _generic_embedding_batch(
189
  )
190
 
191
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
192
  @dataclass(eq=False)
193
  class BatchExecutor:
194
  """Model and batch policy for one bounded embedding window at a time."""
@@ -312,7 +368,7 @@ class BatchExecutor:
312
  if not bool(((raw_mask == 0) | (raw_mask == 1)).all()):
313
  raise ValueError("Embedding residue_mask must contain finite binary values.")
314
  M = raw_mask.to(device=X.device, dtype=torch.bool) # (b, l)
315
- valid_X_shape = (
316
  X.ndim == 3
317
  and X.shape[0] == len(batch_records)
318
  and X.shape[-1] > 0
@@ -327,7 +383,7 @@ class BatchExecutor:
327
  and X.shape[-1] > 0
328
  and M.shape == (X.shape[0], X.shape[2])
329
  )
330
- if not (valid_X_shape or valid_all_states_shape):
331
  raise ValueError(
332
  "Embedding batches must provide X with shape (b, l, d), or "
333
  "(b, states, l, d) when storing all hidden states, and "
@@ -335,7 +391,7 @@ class BatchExecutor:
335
  )
336
  if not bool(M.any(dim=1).all()):
337
  raise ValueError("Every embedding sample must contain a biological residue.")
338
- finite_selected = ( # X.shape
339
  torch.isfinite(X) | ~M.unsqueeze(-1)
340
  if X.ndim == 3
341
  else torch.isfinite(X) | ~M[:, None, :, None]
@@ -347,6 +403,11 @@ class BatchExecutor:
347
  _validate_parti_length(M)
348
  if self.dtype is not None:
349
  X = X.to(dtype=self.dtype) # unchanged shape
 
 
 
 
 
350
 
351
  if self.full_embeddings:
352
  if X.ndim == 4:
@@ -377,3 +438,224 @@ class BatchExecutor:
377
  for position in range(window_start, window_start + len(window_records))
378
  ]
379
  return new_records, pool_slices
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
13
  from .identity import _model_device
14
  from .inputs import _planned_batches
15
  from .pooling import Pooler
16
+ from .taps import HiddenTap, ReducedTap, SparseResidueTap, StreamingTap, TapBatch, TapPlan
17
+ from .types import EmbeddingBatch, EmbeddingInput, EmbeddingRecord, TapRecord
18
+ from ..features.layouts import TopKRow, validate_topk
19
 
20
 
21
  _MAX_PARTI_RESIDUES = 2_048
 
38
  store_all_hidden_states: bool = False,
39
  ) -> Tensor:
40
  """Select one hidden state or stack every state without changing values."""
41
+ # last_hidden_state: (b, l, d); hidden_states: (b, l, d) per entry
42
  if store_all_hidden_states:
43
  if not hidden_states:
44
  raise ValueError("store_all_hidden_states requires model hidden states.")
 
103
  return M # (b, l)
104
 
105
 
106
+ def canonical_residue_ids(sequence: str, tokenizer: Any) -> list[int]:
107
+ """Validate an already normalized protein's one-token-per-residue representation.
108
+
109
+ This does not normalize input or create sequence identity. Canonical callers supply their
110
+ verified inventory text; unknown residues, including J in ESMC, fail before inference.
111
+ """
112
+ if not sequence or not sequence.isascii() or not sequence.isalpha() or not sequence.isupper():
113
+ raise ValueError("Canonical feature input must be an already normalized uppercase protein.")
114
+ ids = tokenizer.convert_tokens_to_ids(list(sequence))
115
+ special = set(tokenizer.all_special_ids)
116
+ if (not isinstance(ids, list) or len(ids) != len(sequence)
117
+ or any(type(token) is not int or token in special for token in ids)):
118
+ raise ValueError("A canonical residue has no non-special tokenizer representation.")
119
+ return ids
120
+
121
+
122
+ def _tokenized_batch(
123
+ model: Any,
124
+ sequences: list[str],
125
+ *,
126
+ tokenizer: Any,
127
+ max_length: int | None,
128
+ truncate: bool,
129
+ ) -> tuple[Tensor, Tensor, Tensor]:
130
+ """Tokenize one batch on the model device: input IDs, attention mask, and residue mask M."""
131
+
132
+ tokenize_kwargs: dict[str, Any] = {
133
+ "return_tensors": "pt",
134
+ "padding": True,
135
+ "truncation": truncate,
136
+ }
137
+ if max_length is not None and truncate:
138
+ # ``max_length`` is a biological-residue limit. Tokenizer limits include
139
+ # boundary tokens, so reserve their declared width instead of dropping
140
+ # residues at the exact boundary.
141
+ special_token_count = 0
142
+ num_special_tokens_to_add = getattr(tokenizer, "num_special_tokens_to_add", None)
143
+ if callable(num_special_tokens_to_add):
144
+ special_token_count = int(num_special_tokens_to_add(pair=False))
145
+ tokenize_kwargs["max_length"] = max_length + special_token_count
146
+ sequence_tokenizer = getattr(model, "_tokenize_sequence_batch", None)
147
+ if callable(sequence_tokenizer):
148
+ encoded = sequence_tokenizer(sequences, tokenizer=tokenizer, **tokenize_kwargs)
149
+ else:
150
+ encoded = tokenizer(sequences, **tokenize_kwargs)
151
+ device = _model_device(model)
152
+ input_ids = encoded["input_ids"].to(device) # (b, l)
153
+ attention_mask = encoded.get( # (b, l)
154
+ "attention_mask",
155
+ input_ids.new_ones(input_ids.shape),
156
+ ).to(device)
157
+ M = _biological_residue_mask(input_ids, attention_mask, tokenizer) # (b, l)
158
+ return input_ids, attention_mask, M # each: (b, l)
159
+
160
+
161
  def _generic_embedding_batch(
162
  model: Any,
163
  sequences: list[str],
 
197
  if tokenizer is None:
198
  raise ValueError("A tokenizer is required for this model's embedding path.")
199
 
200
+ input_ids, attention_mask, M = _tokenized_batch(
201
+ model,
202
+ sequences,
203
+ tokenizer=tokenizer,
204
+ max_length=max_length,
205
+ truncate=truncate,
206
+ ) # each: (b, l)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
207
  if need_attentions:
208
  # Validate l before either the backbone or its quadratic attention graph
209
  # is materialized. M has shape (b, l).
 
227
  )
228
 
229
 
230
+ def _sparse_residue_rows(tap: SparseResidueTap, batch: TapBatch) -> list[TopKRow]:
231
+ """Validate and own each sequence's sparse biological-residue outputs on the CPU."""
232
+ rows = tap.reduce(batch)
233
+ if not isinstance(rows, Sequence) or len(rows) != batch.X.shape[0]:
234
+ raise ValueError("A sparse residue reducer must return one TopKRow per sequence.")
235
+ output = []
236
+ for row, length in zip(rows, batch.residue_mask.sum(dim=1).tolist(), strict=True):
237
+ if not isinstance(row, TopKRow):
238
+ raise TypeError("A sparse residue reducer must return TopKRow values.")
239
+ validate_topk(row.indices, row.values, tap.codebook_size, tap.sparse_count)
240
+ if row.values.shape[0] != length:
241
+ raise ValueError("Sparse residue output must retain every biological residue in order.")
242
+ output.append(TopKRow(
243
+ row.indices.detach().cpu().clone(), row.values.detach().cpu().clone(),
244
+ ))
245
+ return output
246
+
247
+
248
  @dataclass(eq=False)
249
  class BatchExecutor:
250
  """Model and batch policy for one bounded embedding window at a time."""
 
368
  if not bool(((raw_mask == 0) | (raw_mask == 1)).all()):
369
  raise ValueError("Embedding residue_mask must contain finite binary values.")
370
  M = raw_mask.to(device=X.device, dtype=torch.bool) # (b, l)
371
+ valid_embedding_shape = (
372
  X.ndim == 3
373
  and X.shape[0] == len(batch_records)
374
  and X.shape[-1] > 0
 
383
  and X.shape[-1] > 0
384
  and M.shape == (X.shape[0], X.shape[2])
385
  )
386
+ if not (valid_embedding_shape or valid_all_states_shape):
387
  raise ValueError(
388
  "Embedding batches must provide X with shape (b, l, d), or "
389
  "(b, states, l, d) when storing all hidden states, and "
 
391
  )
392
  if not bool(M.any(dim=1).all()):
393
  raise ValueError("Every embedding sample must contain a biological residue.")
394
+ finite_selected = ( # (b, l, d) or (b, n_states, l, d)
395
  torch.isfinite(X) | ~M.unsqueeze(-1)
396
  if X.ndim == 3
397
  else torch.isfinite(X) | ~M[:, None, :, None]
 
403
  _validate_parti_length(M)
404
  if self.dtype is not None:
405
  X = X.to(dtype=self.dtype) # unchanged shape
406
+ selected = M.unsqueeze(-1) if X.ndim == 3 else M[:, None, :, None] # (b, l, 1) or (b, 1, l, 1)
407
+ if not bool((torch.isfinite(X) | ~selected).all()):
408
+ raise ValueError(
409
+ "Embedding dtype conversion produced non-finite biological residues."
410
+ )
411
 
412
  if self.full_embeddings:
413
  if X.ndim == 4:
 
438
  for position in range(window_start, window_start + len(window_records))
439
  ]
440
  return new_records, pool_slices
441
+
442
+
443
+ def _tap_states(
444
+ model: Any,
445
+ sequences: list[str],
446
+ *,
447
+ tokenizer: Any | None,
448
+ max_length: int | None,
449
+ truncate: bool,
450
+ layers: tuple[int, ...],
451
+ streaming: tuple[StreamingTap, ...] = (),
452
+ require_residue_identity: bool = False,
453
+ ) -> tuple[dict[int, Tensor], Tensor, Tensor, dict[str, Tensor]]:
454
+ """Record ``layers`` in one forward pass; return the states, the token mask, and M."""
455
+
456
+ resolved_tokenizer = tokenizer if tokenizer is not None else getattr(model, "tokenizer", None)
457
+ if resolved_tokenizer is None:
458
+ raise ValueError("A tokenizer is required for this model's embedding path.")
459
+ expected_ids = (
460
+ [canonical_residue_ids(sequence, resolved_tokenizer) for sequence in sequences]
461
+ if require_residue_identity else None
462
+ )
463
+ input_ids, attention_mask, M = _tokenized_batch(
464
+ model,
465
+ sequences,
466
+ tokenizer=resolved_tokenizer,
467
+ max_length=max_length,
468
+ truncate=truncate,
469
+ ) # each: (b, l)
470
+ token_mask = attention_mask.to(dtype=torch.bool) # (b, l)
471
+ if expected_ids is not None:
472
+ for index, expected in enumerate(expected_ids):
473
+ # The actual biological mask must select each original residue exactly once,
474
+ # in order. This checks tokenization, padding and cropping before the encoder.
475
+ if input_ids[index, M[index]].tolist() != expected:
476
+ raise ValueError(
477
+ "Tokenizer and biological mask do not preserve original residue positions."
478
+ )
479
+ if not bool(M.any(dim=1).all()):
480
+ raise ValueError("Every embedding sample must contain a biological residue.")
481
+ streamed: dict[str, Tensor] = {}
482
+ if streaming:
483
+ if getattr(model, "embedding_streaming_tap_support", False) is not True:
484
+ raise ValueError("This model does not support streaming hidden-state taps.")
485
+ accumulators = [(tap, tap.begin()) for tap in streaming]
486
+ stream_layers = tuple(sorted({layer for tap in streaming for layer in tap.layers}))
487
+ seen: list[int] = []
488
+
489
+ def consume(layer: int, X: Tensor) -> None:
490
+ # X: (b, l, d), borrowed until this callback returns.
491
+ if len(seen) >= len(stream_layers) or layer != stream_layers[len(seen)]:
492
+ raise RuntimeError("Streaming hidden states arrived out of plan order.")
493
+ if X.ndim != 3 or X.shape[:2] != M.shape:
494
+ raise ValueError("Streaming hidden states are not token-aligned.")
495
+ if not bool((torch.isfinite(X) | ~M.unsqueeze(-1)).all()):
496
+ raise ValueError("Biological residue embeddings produced non-finite output.")
497
+ seen.append(layer)
498
+ batch = TapBatch(X, token_mask, M)
499
+ for tap, accumulator in accumulators:
500
+ if layer in tap.layers:
501
+ accumulator.update(layer, batch)
502
+
503
+ states = model._embed_taps(
504
+ input_ids, attention_mask, layers, stream_layers=stream_layers, state_consumer=consume,
505
+ )
506
+ if tuple(seen) != stream_layers:
507
+ raise RuntimeError("The encoder did not deliver every streaming hidden state.")
508
+ for tap, accumulator in accumulators:
509
+ Y = accumulator.finish() # (b, l, c)
510
+ if (not isinstance(Y, Tensor) or Y.ndim != 3
511
+ or Y.shape[:2] != M.shape or Y.shape[2] == 0):
512
+ raise ValueError(
513
+ f"Streaming tap {tap.name!r} must return token-aligned (b, l, c) features."
514
+ )
515
+ if Y.device != M.device or not Y.is_floating_point():
516
+ raise ValueError("Streaming features must be floating tensors on the input device.")
517
+ if not bool((torch.isfinite(Y) | ~M.unsqueeze(-1)).all()):
518
+ raise ValueError(
519
+ f"Streaming tap {tap.name!r} produced non-finite biological residues."
520
+ )
521
+ streamed[tap.name] = Y
522
+ else:
523
+ states = model._embed_taps(input_ids, attention_mask, layers) # each: (b, l, d)
524
+ # states: (b, l, d) per requested layer; token_mask, M: (b, l); streamed: (b, l, c) per streaming tap
525
+ return states, token_mask, M, streamed # ((b, l, d), ...), (b, l), (b, l), {tap: (b, l, c)}
526
+
527
+
528
+ def _reduced_rows(tap: ReducedTap, batch: TapBatch) -> list[Tensor]:
529
+ """Apply a reducer and split its output into one CPU tensor per sequence."""
530
+
531
+ # batch.X: (b, l, d)
532
+ Y = tap.reduce(batch) # (b, ...)
533
+ if not isinstance(Y, Tensor):
534
+ raise TypeError(f"The reducer of tap {tap.name!r} must return a Tensor.")
535
+ if Y.ndim == 0 or Y.shape[0] != batch.X.shape[0]:
536
+ raise ValueError(
537
+ f"The reducer of tap {tap.name!r} must return one row per sequence, shape "
538
+ f"(b, ...) with b={batch.X.shape[0]}; it returned {tuple(Y.shape)}."
539
+ )
540
+ if Y.is_floating_point() and not bool(torch.isfinite(Y).all()):
541
+ raise ValueError(f"The reducer of tap {tap.name!r} produced non-finite output.")
542
+ return list(Y.detach().cpu().unbind(0)) # b tensors of shape (...), Y without its batch axis
543
+
544
+
545
+ @dataclass(eq=False)
546
+ class TapExecutor:
547
+ """Model and batch policy for one bounded window of a tap plan.
548
+
549
+ Each batch runs one forward pass that records every tapped state and stops after the
550
+ deepest. A hidden tap's dtype overrides the run dtype. Conversion starts from the original
551
+ state for each tap, so a low-precision residue output cannot quantize a pooled sibling.
552
+ A reducer receives the run dtype and its output keeps the dtype it returns.
553
+ """
554
+
555
+ model: Any
556
+ plan: TapPlan
557
+ batch_size: int
558
+ max_tokens_per_batch: int | None
559
+ max_length: int | None
560
+ truncate: bool
561
+ tokenizer: Any | None
562
+ dtype: torch.dtype | None
563
+ attention_backend: str | None
564
+ require_residue_identity: bool = False
565
+ poolers: dict[str, Pooler] = field(init=False)
566
+
567
+ def __post_init__(self) -> None:
568
+ pooled_streams = [tap.name for tap in self.plan.taps if isinstance(tap, StreamingTap) and tap.pooling is not None]
569
+ if pooled_streams:
570
+ raise ValueError(f"Pooled streaming taps {pooled_streams} run in a token run (TokenTapExecutor) only.")
571
+ self.poolers = {
572
+ tap.name: Pooler(tap.pooling)
573
+ for tap in self.plan.taps
574
+ if isinstance(tap, HiddenTap) and tap.pooling is not None
575
+ }
576
+
577
+ def run_window(
578
+ self,
579
+ window_records: Sequence[EmbeddingInput],
580
+ *,
581
+ window_start: int,
582
+ ) -> tuple[list[TapRecord], dict[str, dict[str, tuple[int, int]]]]:
583
+ """Restore source order after length-bucketed inference; return each pooled tap's slices."""
584
+
585
+ pool_slices: dict[str, dict[str, tuple[int, int]]] = {}
586
+ window_results: dict[int, TapRecord] = {}
587
+ for local_positions in _planned_batches(
588
+ window_records,
589
+ range(len(window_records)),
590
+ batch_size=self.batch_size,
591
+ max_tokens_per_batch=self.max_tokens_per_batch,
592
+ max_length=self.max_length,
593
+ truncate=self.truncate,
594
+ ):
595
+ batch_records = [window_records[position] for position in local_positions]
596
+ sequences = [
597
+ record.sequence[: self.max_length]
598
+ if self.truncate and self.max_length is not None
599
+ else record.sequence
600
+ for record in batch_records
601
+ ]
602
+ states, token_mask, M, streamed = _tap_states( # states: each (b, l, d); masks: (b, l)
603
+ self.model,
604
+ sequences,
605
+ tokenizer=self.tokenizer,
606
+ max_length=self.max_length,
607
+ truncate=self.truncate,
608
+ layers=self.plan.captured_layers,
609
+ streaming=tuple(tap for tap in self.plan.taps if isinstance(tap, StreamingTap)),
610
+ require_residue_identity=self.require_residue_identity,
611
+ )
612
+ if not bool(M.any(dim=1).all()):
613
+ raise ValueError("Every embedding sample must contain a biological residue.")
614
+ for X in states.values(): # each: (b, l, d)
615
+ if not bool((torch.isfinite(X) | ~M.unsqueeze(-1)).all()):
616
+ raise ValueError("Biological residue embeddings produced non-finite output.")
617
+ outputs: dict[str, list[Tensor] | list[TopKRow]] = {}
618
+ for tap, layer in zip(self.plan.taps, self.plan.layers, strict=True):
619
+ if isinstance(tap, StreamingTap):
620
+ outputs[tap.name] = _residue_embeddings(streamed[tap.name], M) # each: (r_i, c)
621
+ continue
622
+ X = states[layer] # (b, l, d)
623
+ dtype = self.dtype
624
+ if isinstance(tap, HiddenTap) and tap.dtype is not None:
625
+ dtype = tap.dtype
626
+ if dtype is not None:
627
+ X = X.to(dtype=dtype) # (b, l, d), from the original captured state
628
+ if not bool((torch.isfinite(X) | ~M.unsqueeze(-1)).all()):
629
+ raise ValueError(
630
+ "Tap dtype conversion produced non-finite biological residues."
631
+ )
632
+ if isinstance(tap, SparseResidueTap):
633
+ batch = TapBatch(X=X, token_mask=token_mask, residue_mask=M)
634
+ outputs[tap.name] = _sparse_residue_rows(tap, batch) # each pair: (r_i,k)
635
+ elif isinstance(tap, ReducedTap):
636
+ batch = TapBatch(X=X, token_mask=token_mask, residue_mask=M)
637
+ outputs[tap.name] = _reduced_rows(tap, batch) # each: (...)
638
+ elif tap.pooling is None:
639
+ outputs[tap.name] = _residue_embeddings(X, M) # each: (r_i, d)
640
+ else:
641
+ pooler = self.poolers[tap.name]
642
+ Y = pooler(X, M, attention_backend=self.attention_backend) # (b, n_poolers * d)
643
+ pool_slices[tap.name] = pooler.output_slices(X.shape[-1])
644
+ outputs[tap.name] = list(Y.detach().cpu().unbind(0)) # each: (n_poolers * d,)
645
+ for offset, (position, record) in enumerate(
646
+ zip(local_positions, batch_records, strict=True)
647
+ ):
648
+ window_results[window_start + position] = TapRecord(
649
+ record.id,
650
+ record.sequence,
651
+ {name: values[offset] for name, values in outputs.items()},
652
+ # This correspondence was proven against actual token IDs and M before
653
+ # inference, not inferred from a tensor's row count.
654
+ tuple(range(len(sequences[offset]))) if self.require_residue_identity else None,
655
+ )
656
+
657
+ new_records = [
658
+ window_results[position]
659
+ for position in range(window_start, window_start + len(window_records))
660
+ ]
661
+ return new_records, pool_slices
fastplms/embeddings/feature_runs.py ADDED
@@ -0,0 +1,232 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Run a tap plan and persist each tap into its own feature store.
2
+
3
+ This is the one place a model runs to fill the store. ``embed_into_features`` embeds only the
4
+ sequences the stores lack, takes one forward pass per batch for every tap, and writes each tap's
5
+ rows into the store of its key. A second call with the same sequences runs no model at all.
6
+
7
+ Each tap becomes one feature, so a run that taps the last hidden state, a mean-pooled vector, and
8
+ max-pooled sparse-autoencoder codes fills three stores from one pass. The store's layout decides
9
+ how a tap's tensor is stored: a pooled vector goes in dense, per-residue rows go in ragged, and a
10
+ sparse-autoencoder vector goes in csr, compressed to its exactly non-zero codes.
11
+
12
+ The caller owns the keys, because the key composes the model, its revision, the autoencoder, the
13
+ layer, the pooling, the dtype, and the residue limit, and only the caller knows the pinned
14
+ revisions. ``foundry.embedding.feature_key`` computes them.
15
+ """
16
+
17
+ from __future__ import annotations
18
+
19
+ from collections.abc import Mapping, Sequence
20
+ from contextlib import ExitStack
21
+ from pathlib import Path
22
+ from typing import Any, Protocol
23
+ from torch import Tensor
24
+
25
+ from .runner import embed_dataset
26
+ from .taps import Tap
27
+ from .types import TapRecord, TapRunReceipt
28
+ from ..features.layouts import DENSE, RAGGED, RAGGED_TOPK, SparseRow, TopKRow
29
+ from ..features.store import FeatureStore, SegmentReceipt, SegmentWriter, StoredFeature
30
+
31
+
32
+ class FeatureRunContract(Protocol):
33
+ """Scientific identity policy supplied by the caller, independent of the storage format."""
34
+
35
+ def validate(
36
+ self, model: Any, sequences: Sequence[str], features: Mapping[str, StoredFeature],
37
+ taps: Sequence[Tap], options: Mapping[str, Any],
38
+ ) -> None: ...
39
+
40
+ def validate_cached(self, name: str, store: FeatureStore, sequences: Sequence[str]) -> None: ...
41
+
42
+ def bind_rows(self, name: str, records: Sequence[TapRecord]) -> Sequence[Mapping[str, Any]]: ...
43
+
44
+ def before_commit(self) -> None: ...
45
+
46
+
47
+ def embed_into_features(
48
+ model: Any,
49
+ sequences: Sequence[str],
50
+ root: str | Path,
51
+ features: Mapping[str, StoredFeature],
52
+ *,
53
+ taps: Sequence[Tap],
54
+ metadata: Mapping[str, Any] | None = None,
55
+ contract: FeatureRunContract | None = None,
56
+ max_part_bytes: int = 256 * 1024**2,
57
+ keep_special_tokens: bool = False,
58
+ **embed_kwargs: Any,
59
+ ) -> dict[str, SegmentReceipt]:
60
+ """Fill each named feature under ``root`` from one pass over the sequences it lacks.
61
+
62
+ ``features`` maps a tap name to the spec of the feature it fills, and must name every tap.
63
+ ``metadata`` is recorded on every segment this run commits, beside the run fingerprint.
64
+ ``max_part_bytes`` bounds each part's encoded tensor payload, excluding its small safetensors
65
+ header and row-identity sidecar. A single row exceeding it fails without being committed.
66
+ Remaining keyword arguments go to ``embed_dataset``. Tensor outputs are streamed one bounded
67
+ batch window at a time; input identities and part metadata still scale with inventory size.
68
+
69
+ A contract that keeps the special tokens (``contract.keep_special_tokens`` is True) selects the canonical
70
+ token path by itself, so a caller cannot forget the argument; a contract that keeps residues only (False)
71
+ refuses ``keep_special_tokens=True``.
72
+
73
+ ``keep_special_tokens=True`` is the canonical token path: every per-token stream holds l + 2 rows per
74
+ protein (row 0 CLS, rows 1..l residues, row l + 1 EOS, shape (n, d) with n = sum(l_i + 2) overall), every
75
+ pooled stream covers those same rows, and the run goes through ``embed_token_features`` and its
76
+ asynchronous writer. Its ``embed_kwargs`` are the contract's ``embedding_options`` (``max_length`` is the
77
+ crop in residues), and it returns the last committed segment of each stream; call ``embed_token_features``
78
+ for every segment.
79
+
80
+ Returns the committed segment of each feature that gained rows. A feature whose sequences were
81
+ all present is absent from the result, and an empty result means the model never ran.
82
+ """
83
+
84
+ # The contract names the stored layout: l + 2 rows (True), l rows (False), or no claim (no attribute).
85
+ contract_keeps = getattr(contract, "keep_special_tokens", None)
86
+ if contract_keeps is True:
87
+ keep_special_tokens = True
88
+ elif contract_keeps is False and keep_special_tokens:
89
+ raise ValueError("keep_special_tokens=True needs a contract captured with the special tokens kept.")
90
+
91
+ if keep_special_tokens:
92
+ from .token_batches import CANONICAL_MAX_RESIDUES
93
+ from .token_runs import embed_token_features
94
+
95
+ options = dict(embed_kwargs)
96
+ settings: dict[str, Any] = {
97
+ "max_residues": options.pop("max_length", CANONICAL_MAX_RESIDUES),
98
+ "max_sequences": options.pop("batch_size", 256),
99
+ "max_tokens": options.pop("max_tokens_per_batch", None) or 32768,
100
+ "window": options.pop("batch_window_size", 65536),
101
+ "dtype": options.pop("dtype", None),
102
+ }
103
+ options.pop("truncate", None) # the canonical crop is always a prefix crop
104
+ if "fixed_batch_size" in options:
105
+ settings["fixed_batch_size"] = options.pop("fixed_batch_size")
106
+ if options:
107
+ raise ValueError(f"keep_special_tokens takes no other extraction options; received {sorted(options)}.")
108
+ received = embed_token_features(
109
+ model, sequences, root, features, taps=taps, contract=contract, metadata=metadata,
110
+ part_bytes=max_part_bytes, **settings,
111
+ )
112
+ return {name: group[-1] for name, group in received.items()}
113
+
114
+ if type(max_part_bytes) is not int or max_part_bytes <= 0:
115
+ raise ValueError("max_part_bytes must be a positive integer.")
116
+ if "tap_sink" in embed_kwargs:
117
+ raise ValueError("embed_into_features owns its tap_sink destination.")
118
+ tap_names = [tap.name for tap in taps]
119
+ if set(features) != set(tap_names):
120
+ raise ValueError(
121
+ "features must name exactly the taps this run takes.\n"
122
+ f" taps: {sorted(tap_names)}\n"
123
+ f" features: {sorted(features)}"
124
+ )
125
+ for name, spec in features.items():
126
+ if spec.positions:
127
+ raise ValueError(
128
+ f"Feature {spec.key!r} stores the argmax residue of each code, and no tap carries "
129
+ f"them, so this run cannot fill it from tap {name!r}. A pipeline that computes "
130
+ "positions itself writes them through the store's own segment writer."
131
+ )
132
+
133
+ ordered = _distinct(sequences)
134
+ if not ordered:
135
+ raise ValueError("embed_into_features needs at least one sequence.")
136
+ if contract is None and any(
137
+ spec.descriptor.get("schema") == "feature_spec_v1" for spec in features.values()
138
+ ):
139
+ raise ValueError("Complete feature descriptors require a FeatureRunContract.")
140
+ if contract is not None:
141
+ contract.validate(model, ordered, features, taps, embed_kwargs)
142
+ stores = {name: FeatureStore.open(root, spec) for name, spec in features.items()}
143
+ wanted = {name: frozenset(store.missing(ordered)) for name, store in stores.items()}
144
+ if contract is not None:
145
+ for name, store in stores.items():
146
+ contract.validate_cached(name, store, [s for s in ordered if s not in wanted[name]])
147
+ to_embed = [sequence for sequence in ordered if any(sequence in group for group in wanted.values())]
148
+ if not to_embed:
149
+ return {}
150
+
151
+ receipts: dict[str, SegmentReceipt] = {}
152
+ with ExitStack() as stack:
153
+ writers: dict[str, SegmentWriter] = {}
154
+ # Recheck live dependencies and caller descriptors after all windows have been staged.
155
+ check = None if contract is None else lambda: contract.validate(
156
+ model, (), features, taps, embed_kwargs,
157
+ )
158
+
159
+ def append_window(records: Sequence[TapRecord], identity: Mapping[str, str]) -> None:
160
+ for name, store in stores.items():
161
+ embedded = [record for record in records if record.sequence in wanted[name]]
162
+ if not embedded:
163
+ continue
164
+ rows = [_row_for(store.spec, record, name) for record in embedded]
165
+ identities = None if contract is None else contract.bind_rows(name, embedded)
166
+ if name not in writers:
167
+ run_metadata = {
168
+ **dict(metadata or {}), "input_fingerprint": identity["input_fingerprint"],
169
+ "storage_policy": {"max_part_tensor_bytes": max_part_bytes},
170
+ }
171
+ writers[name] = stack.enter_context(store.segment(
172
+ identity["run_fingerprint"], run_metadata, before_commit=check,
173
+ ))
174
+ writers[name].append_bounded(
175
+ [record.sequence for record in embedded], rows, row_metadata=identities,
176
+ max_tensor_bytes=max_part_bytes,
177
+ )
178
+
179
+ run_receipt = embed_dataset(
180
+ model, to_embed, taps=list(taps), require_residue_identity=contract is not None,
181
+ tap_sink=append_window, **embed_kwargs,
182
+ )
183
+ if not isinstance(run_receipt, TapRunReceipt) or run_receipt.record_count != len(to_embed):
184
+ raise TypeError("Feature extraction did not complete delivery of every requested row.")
185
+ for name, writer in writers.items():
186
+ receipts[name] = writer.commit()
187
+ return receipts
188
+
189
+
190
+ def _row_for(spec: StoredFeature, record: TapRecord, tap: str) -> Tensor | SparseRow | TopKRow:
191
+ """One tap's output for one sequence, in the shape its store stores."""
192
+
193
+ value = record.tensors[tap] # (w,) pooled, (r_i, d) per residue, or a TopKRow of (r_i, k) tensors
194
+ if spec.layout == RAGGED_TOPK:
195
+ if not isinstance(value, TopKRow):
196
+ raise ValueError(f"Feature {spec.key!r} needs sparse TopKRow output from tap {tap!r}.")
197
+ return value # TopKRow of (r_i, k) indices and values
198
+ if not isinstance(value, Tensor):
199
+ raise ValueError("Sparse residue tap output requires a ragged_topk feature.")
200
+ if spec.layout == DENSE:
201
+ if value.ndim != 1:
202
+ raise ValueError(
203
+ f"Feature {spec.key!r} is dense, so tap {tap!r} must give one vector per sequence; "
204
+ f"received shape {tuple(value.shape)}. A per-residue tap needs a ragged feature."
205
+ )
206
+ return value # (w,)
207
+ if spec.layout == RAGGED:
208
+ if value.ndim != 2:
209
+ raise ValueError(
210
+ f"Feature {spec.key!r} is ragged, so tap {tap!r} must give (r_i, d) residue rows; "
211
+ f"received shape {tuple(value.shape)}."
212
+ )
213
+ return value # (r_i, d)
214
+ if value.ndim != 1:
215
+ raise ValueError(
216
+ f"Feature {spec.key!r} is csr, so tap {tap!r} must give one vector per sequence; "
217
+ f"received shape {tuple(value.shape)}."
218
+ )
219
+ return SparseRow.from_dense(value) # SparseRow of (nnz,) tensors, from a (w,) vector
220
+
221
+
222
+ def _distinct(sequences: Sequence[str]) -> list[str]:
223
+ seen: set[str] = set()
224
+ ordered: list[str] = []
225
+ for sequence in sequences:
226
+ if sequence not in seen:
227
+ seen.add(sequence)
228
+ ordered.append(sequence)
229
+ return ordered
230
+
231
+
232
+ __all__ = ["embed_into_features"]
fastplms/embeddings/identity.py CHANGED
@@ -3,7 +3,6 @@
3
  from __future__ import annotations
4
 
5
  import hashlib
6
- import json
7
  import platform
8
  import torch
9
 
@@ -12,12 +11,14 @@ from pathlib import Path
12
  from typing import Any
13
  from torch import Tensor
14
 
15
- from .inputs import _InputSpool
16
  from .storage import tensor_sha256
17
  from .types import EmbeddingInput
 
 
18
 
19
 
20
- _RUN_FINGERPRINT_SCHEMA_VERSION = 3
21
  _MODEL_STATE_HASH_CHUNK_BYTES = 16 * 1024**2
22
 
23
 
@@ -79,6 +80,33 @@ def _fingerprint_jsonable(value: Any) -> Any:
79
  }
80
 
81
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
82
  def _tokenizer_content_sha256(tokenizer: Any) -> str:
83
  content: dict[str, Any] = {
84
  "init_kwargs": getattr(tokenizer, "init_kwargs", None),
@@ -93,17 +121,10 @@ def _tokenizer_content_sha256(tokenizer: Any) -> str:
93
  get_added_vocab = getattr(tokenizer, "get_added_vocab", None)
94
  if callable(get_added_vocab):
95
  content["added_vocabulary"] = get_added_vocab()
96
- backend = getattr(tokenizer, "backend_tokenizer", None)
97
- backend_to_str = getattr(backend, "to_str", None)
98
- if callable(backend_to_str):
99
- content["backend"] = backend_to_str()
100
- serialized = json.dumps(
101
- _fingerprint_jsonable(content),
102
- sort_keys=True,
103
- separators=(",", ":"),
104
- ensure_ascii=False,
105
- ).encode()
106
- return hashlib.sha256(serialized).hexdigest()
107
 
108
 
109
  def _tokenizer_metadata(model: Any, tokenizer: Any | None) -> dict[str, Any]:
@@ -301,14 +322,12 @@ def _model_state_sha256(model: Any) -> str:
301
  f"Cannot fingerprint meta-device model state entry {name!r}; pass "
302
  "model_state_fingerprint with a caller-owned state identity."
303
  )
304
- header = json.dumps(
305
  {
306
  "name": name,
307
  "dtype": str(value.dtype).removeprefix("torch."),
308
  "shape": list(value.shape),
309
- },
310
- sort_keys=True,
311
- separators=(",", ":"),
312
  ).encode()
313
  digest.update(len(header).to_bytes(8, "big"))
314
  digest.update(header)
@@ -354,6 +373,7 @@ def _run_fingerprint(
354
  batch_size: int,
355
  batch_window_size: int,
356
  max_tokens_per_batch: int | None,
 
357
  ) -> tuple[str, str, str | None, str]:
358
  input_fingerprint = _input_sha256(records)
359
  attention_backend = _attention_backend(model)
@@ -391,6 +411,7 @@ def _run_fingerprint(
391
  "execution": _execution_identity_metadata(model),
392
  "embedding_context": _fingerprint_jsonable(embedding_context),
393
  "pooling": list(pooling),
 
394
  "full_embeddings": full_embeddings,
395
  "max_length": max_length,
396
  "truncate": truncate,
@@ -399,16 +420,16 @@ def _run_fingerprint(
399
  "batch_size": batch_size,
400
  "batch_window_size": batch_window_size,
401
  "max_tokens_per_batch": max_tokens_per_batch,
402
- "input_storage": ("disk-spool" if isinstance(records, _InputSpool) else "memory"),
403
  },
404
  "model_kwargs": {
405
  key: _fingerprint_jsonable(value) for key, value in sorted(model_kwargs.items())
406
  },
407
  "residue_mask_policy": "attention-mask-minus-special-tokens",
408
  }
409
- run_fingerprint = hashlib.sha256(
410
- json.dumps(payload, sort_keys=True, separators=(",", ":")).encode()
411
- ).hexdigest()
 
412
  return (
413
  input_fingerprint,
414
  run_fingerprint,
@@ -437,6 +458,7 @@ def _embedding_context(
437
  decoder_attention_mask: Tensor | None,
438
  model_kwargs: Mapping[str, Any],
439
  ) -> tuple[dict[str, Any], tuple[str, ...] | None]:
 
440
  if hidden_state_source not in {"encoder", "decoder"}:
441
  raise ValueError("hidden_state_source must be 'encoder' or 'decoder'.")
442
  hidden_state_index = model_kwargs.get("hidden_state_index", -1)
 
3
  from __future__ import annotations
4
 
5
  import hashlib
 
6
  import platform
7
  import torch
8
 
 
11
  from typing import Any
12
  from torch import Tensor
13
 
14
+ from .pooling import POOLING_SEMANTICS
15
  from .storage import tensor_sha256
16
  from .types import EmbeddingInput
17
+ from ..digests import json_sha256
18
+ from ..json_files import compact_json
19
 
20
 
21
+ _RUN_FINGERPRINT_SCHEMA_VERSION = 5
22
  _MODEL_STATE_HASH_CHUNK_BYTES = 16 * 1024**2
23
 
24
 
 
80
  }
81
 
82
 
83
+ def _backend_tokenizer_content(backend: Any) -> str | None:
84
+ """The backend tokenizer's serialization, without the padding and truncation of its last call.
85
+
86
+ Transformers sets padding and truncation on the Rust tokenizer at the start of every encode
87
+ call, from that call's own arguments, and leaves them set. What the backend holds is
88
+ therefore the previous call's settings, which never change the next encoding. Hashing them
89
+ gave the first run in a process a different fingerprint from every identical run after it.
90
+ This run's own padding and truncation are fingerprinted through ``max_length``,
91
+ ``truncate``, and the batching policy.
92
+
93
+ The settings are cleared on a copy, so the caller's tokenizer keeps its state. A tokenizer
94
+ that was never called already serializes without them, so this normalization itself
95
+ does not change its identity. The enclosing run schema also versions pooling semantics.
96
+ """
97
+ to_str = getattr(backend, "to_str", None)
98
+ if not callable(to_str):
99
+ return None
100
+ serialized = to_str()
101
+ from_str = getattr(type(backend), "from_str", None)
102
+ if not callable(from_str):
103
+ return serialized
104
+ call_free = from_str(serialized)
105
+ call_free.no_truncation()
106
+ call_free.no_padding()
107
+ return call_free.to_str()
108
+
109
+
110
  def _tokenizer_content_sha256(tokenizer: Any) -> str:
111
  content: dict[str, Any] = {
112
  "init_kwargs": getattr(tokenizer, "init_kwargs", None),
 
121
  get_added_vocab = getattr(tokenizer, "get_added_vocab", None)
122
  if callable(get_added_vocab):
123
  content["added_vocabulary"] = get_added_vocab()
124
+ backend_content = _backend_tokenizer_content(getattr(tokenizer, "backend_tokenizer", None))
125
+ if backend_content is not None:
126
+ content["backend"] = backend_content
127
+ return json_sha256(_fingerprint_jsonable(content), ensure_ascii=False)
 
 
 
 
 
 
 
128
 
129
 
130
  def _tokenizer_metadata(model: Any, tokenizer: Any | None) -> dict[str, Any]:
 
322
  f"Cannot fingerprint meta-device model state entry {name!r}; pass "
323
  "model_state_fingerprint with a caller-owned state identity."
324
  )
325
+ header = compact_json(
326
  {
327
  "name": name,
328
  "dtype": str(value.dtype).removeprefix("torch."),
329
  "shape": list(value.shape),
330
+ }
 
 
331
  ).encode()
332
  digest.update(len(header).to_bytes(8, "big"))
333
  digest.update(header)
 
373
  batch_size: int,
374
  batch_window_size: int,
375
  max_tokens_per_batch: int | None,
376
+ taps: Sequence[Mapping[str, Any]] | None = None,
377
  ) -> tuple[str, str, str | None, str]:
378
  input_fingerprint = _input_sha256(records)
379
  attention_backend = _attention_backend(model)
 
411
  "execution": _execution_identity_metadata(model),
412
  "embedding_context": _fingerprint_jsonable(embedding_context),
413
  "pooling": list(pooling),
414
+ "pooling_semantics": dict(POOLING_SEMANTICS),
415
  "full_embeddings": full_embeddings,
416
  "max_length": max_length,
417
  "truncate": truncate,
 
420
  "batch_size": batch_size,
421
  "batch_window_size": batch_window_size,
422
  "max_tokens_per_batch": max_tokens_per_batch,
 
423
  },
424
  "model_kwargs": {
425
  key: _fingerprint_jsonable(value) for key, value in sorted(model_kwargs.items())
426
  },
427
  "residue_mask_policy": "attention-mask-minus-special-tokens",
428
  }
429
+ if taps is not None:
430
+ # Each hidden tap binds its requested dtype; each reduced tap binds its own contract.
431
+ payload["taps"] = _fingerprint_jsonable(taps)
432
+ run_fingerprint = json_sha256(payload)
433
  return (
434
  input_fingerprint,
435
  run_fingerprint,
 
458
  decoder_attention_mask: Tensor | None,
459
  model_kwargs: Mapping[str, Any],
460
  ) -> tuple[dict[str, Any], tuple[str, ...] | None]:
461
+ # decoder_input_ids, decoder_attention_mask: (n_records, l_decoder), aligned with records
462
  if hidden_state_source not in {"encoder", "decoder"}:
463
  raise ValueError("hidden_state_source must be 'encoder' or 'decoder'.")
464
  hidden_state_index = model_kwargs.get("hidden_state_index", -1)
fastplms/embeddings/inputs.py CHANGED
@@ -3,46 +3,104 @@
3
  from __future__ import annotations
4
 
5
  import hashlib
 
6
  import sqlite3
7
  import tempfile
 
8
 
9
  from collections.abc import Iterable, Iterator, Mapping, Sequence
 
10
  from pathlib import Path
11
- from typing import overload
12
 
13
  from .types import EmbeddingInput
14
 
15
 
16
- def iter_fasta(path: str | Path) -> Iterator[EmbeddingInput]:
17
- """Yield FASTA records in source order without reading the file into memory."""
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
18
 
19
- identifier: str | None = None
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
20
  sequence_parts: list[str] = []
21
  found_record = False
22
- with Path(path).open("r", encoding="utf-8") as handle:
23
- for line_number, raw_line in enumerate(handle, start=1):
24
- line = raw_line.strip()
25
- if not line:
26
- continue
27
- if line.startswith(">"):
28
- if identifier is not None:
29
- found_record = True
30
- yield EmbeddingInput(identifier, "".join(sequence_parts))
31
- identifier = line[1:].strip().split(maxsplit=1)[0]
32
- if not identifier:
33
- raise ValueError(f"Missing FASTA identifier on line {line_number}.")
34
- sequence_parts = []
35
- else:
36
- if identifier is None:
37
- raise ValueError(
38
- f"Sequence data precedes the first FASTA header on line {line_number}."
39
- )
40
- sequence_parts.append("".join(line.split()))
41
- if identifier is not None:
42
  found_record = True
43
- yield EmbeddingInput(identifier, "".join(sequence_parts))
44
- if not found_record:
45
- raise ValueError(f"No FASTA records found in {path}.")
 
 
 
 
 
 
 
 
46
 
47
 
48
  def parse_fasta(path: str | Path) -> list[EmbeddingInput]:
@@ -66,6 +124,27 @@ def _normalize_input_item(
66
  )
67
 
68
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
69
  class _InputSpool(Sequence[EmbeddingInput]):
70
  """Immutable disk-backed normalized inputs with an incremental digest."""
71
 
@@ -73,19 +152,19 @@ class _InputSpool(Sequence[EmbeddingInput]):
73
  self,
74
  values: Iterable[str | EmbeddingInput | tuple[str, str]],
75
  ) -> None:
76
- self._temporary: tempfile.TemporaryDirectory[str] | None = tempfile.TemporaryDirectory(
77
- prefix="fastplms-inputs-"
78
- )
79
- self.path = Path(self._temporary.name) / "inputs.sqlite"
80
- self._connection: sqlite3.Connection | None = sqlite3.connect(self.path)
81
- self._connection.execute(
82
- "CREATE TABLE inputs ("
83
- "position INTEGER PRIMARY KEY, input_id TEXT NOT NULL, sequence TEXT NOT NULL)"
84
- )
85
  digest = hashlib.sha256()
86
  count = 0
87
  pending: list[tuple[int, str, str]] = []
88
  try:
 
 
 
 
 
89
  for position, item in enumerate(values):
90
  record = _normalize_input_item(position, item)
91
  for value in (record.id, record.sequence):
@@ -95,15 +174,15 @@ class _InputSpool(Sequence[EmbeddingInput]):
95
  pending.append((position, record.id, record.sequence))
96
  count += 1
97
  if len(pending) == 1_024:
98
- self._connection.executemany("INSERT INTO inputs VALUES (?, ?, ?)", pending)
99
  pending.clear()
100
  if pending:
101
- self._connection.executemany("INSERT INTO inputs VALUES (?, ?, ?)", pending)
102
  if count == 0:
103
  raise ValueError("inputs must contain at least one sequence.")
104
- self._connection.commit()
105
- self._connection.close()
106
- self._connection = sqlite3.connect(
107
  f"{self.path.resolve().as_uri()}?mode=ro",
108
  uri=True,
109
  )
@@ -115,9 +194,10 @@ class _InputSpool(Sequence[EmbeddingInput]):
115
  self._count = count
116
 
117
  def _require_connection(self) -> sqlite3.Connection:
118
- if self._connection is None:
 
119
  raise RuntimeError("Input spool is closed.")
120
- return self._connection
121
 
122
  def __len__(self) -> int:
123
  return self._count
@@ -160,17 +240,7 @@ class _InputSpool(Sequence[EmbeddingInput]):
160
  return EmbeddingInput(row[0], row[1])
161
 
162
  def close(self) -> None:
163
- connection = getattr(self, "_connection", None)
164
- if connection is not None:
165
- connection.close()
166
- self._connection = None
167
- temporary = getattr(self, "_temporary", None)
168
- if temporary is not None:
169
- temporary.cleanup()
170
- self._temporary = None
171
-
172
- def __del__(self) -> None:
173
- self.close()
174
 
175
 
176
  def _normalize_inputs(
 
3
  from __future__ import annotations
4
 
5
  import hashlib
6
+ import shutil
7
  import sqlite3
8
  import tempfile
9
+ import weakref
10
 
11
  from collections.abc import Iterable, Iterator, Mapping, Sequence
12
+ from dataclasses import dataclass
13
  from pathlib import Path
14
+ from typing import NamedTuple, overload
15
 
16
  from .types import EmbeddingInput
17
 
18
 
19
+ class FastaRecord(NamedTuple):
20
+ """One FASTA record: its header text and its sequence lines joined."""
21
+
22
+ header: str
23
+ sequence: str
24
+
25
+
26
+ @dataclass(frozen=True)
27
+ class FastaDialect:
28
+ """How one caller reads FASTA lines. FastPLMs has three readers, and each keeps its rules here.
29
+
30
+ ``strip_lines``: strip each line before use; otherwise lines are used as given, so a space-only line
31
+ is sequence data.
32
+ ``comment_prefix``: lines that start with it are skipped, or ``None`` for no comments.
33
+ ``first_word_header``: the header is its first whitespace-delimited word, not the whole header text.
34
+ ``squeeze_sequence_whitespace``: remove all whitespace inside a sequence line.
35
+ ``orphan_message``: raised as ``ValueError`` for sequence data before the first header, formatted
36
+ with ``line_number`` and ``source``; ``None`` skips such lines.
37
+ ``empty_message``: raised as ``ValueError`` when the input holds no record, formatted with
38
+ ``source``; ``None`` allows empty input.
39
+ """
40
+
41
+ strip_lines: bool
42
+ comment_prefix: str | None
43
+ first_word_header: bool
44
+ squeeze_sequence_whitespace: bool
45
+ orphan_message: str | None
46
+ empty_message: str | None
47
+
48
 
49
+ EMBEDDING_FASTA = FastaDialect(
50
+ strip_lines=True,
51
+ comment_prefix=None,
52
+ first_word_header=True,
53
+ squeeze_sequence_whitespace=True,
54
+ orphan_message="Sequence data precedes the first FASTA header on line {line_number}.",
55
+ empty_message="No FASTA records found in {source}.",
56
+ )
57
+
58
+
59
+ def scan_fasta_lines(
60
+ lines: Iterable[str], dialect: FastaDialect, *, source: str
61
+ ) -> Iterator[FastaRecord]:
62
+ """Yield the records of FASTA ``lines`` in order, one record at a time.
63
+
64
+ ``source`` names the input in the dialect's messages. A record is yielded when the next header or
65
+ the end of input shows it is complete, so an error raised later in the input surfaces after the
66
+ records before it.
67
+ """
68
+
69
+ header: str | None = None
70
  sequence_parts: list[str] = []
71
  found_record = False
72
+ for line_number, raw_line in enumerate(lines, start=1):
73
+ line = raw_line.strip() if dialect.strip_lines else raw_line
74
+ if not line or (dialect.comment_prefix is not None and line.startswith(dialect.comment_prefix)):
75
+ continue
76
+
77
+ if line.startswith(">"):
78
+ if header is not None:
79
+ found_record = True
80
+ yield FastaRecord(header, "".join(sequence_parts))
81
+ header = line[1:].strip()
82
+ if dialect.first_word_header:
83
+ # An empty header has no first word and raises IndexError here.
84
+ header = header.split(maxsplit=1)[0]
85
+ sequence_parts = []
86
+ elif header is not None:
87
+ sequence_parts.append("".join(line.split()) if dialect.squeeze_sequence_whitespace else line)
88
+ elif dialect.orphan_message is not None:
89
+ raise ValueError(dialect.orphan_message.format(line_number=line_number, source=source))
90
+
91
+ if header is not None:
92
  found_record = True
93
+ yield FastaRecord(header, "".join(sequence_parts))
94
+ if not found_record and dialect.empty_message is not None:
95
+ raise ValueError(dialect.empty_message.format(source=source))
96
+
97
+
98
+ def iter_fasta(path: str | Path) -> Iterator[EmbeddingInput]:
99
+ """Yield FASTA records in source order without reading the file into memory."""
100
+
101
+ with Path(path).open("r", encoding="utf-8") as handle:
102
+ for record in scan_fasta_lines(handle, EMBEDDING_FASTA, source=str(path)):
103
+ yield EmbeddingInput(record.header, record.sequence)
104
 
105
 
106
  def parse_fasta(path: str | Path) -> list[EmbeddingInput]:
 
124
  )
125
 
126
 
127
+ class _SpoolFiles:
128
+ """A spool's directory and SQLite connection, released together: the connection first.
129
+
130
+ Windows cannot delete a file that an open connection holds. The cycle collector runs weakref
131
+ finalizers before any ``__del__``, so a spool freed as cyclic garbage, as one referenced from
132
+ a raised exception's traceback is, would have ``tempfile.TemporaryDirectory``'s own finalizer
133
+ remove the directory while the connection was still open. One finalizer owns both instead.
134
+ """
135
+
136
+ def __init__(self) -> None:
137
+ self.directory = Path(tempfile.mkdtemp(prefix="fastplms-inputs-"))
138
+ self.connection: sqlite3.Connection | None = None
139
+
140
+ def release(self) -> None:
141
+ if self.connection is not None:
142
+ self.connection.close()
143
+ self.connection = None
144
+ if self.directory.exists():
145
+ shutil.rmtree(self.directory)
146
+
147
+
148
  class _InputSpool(Sequence[EmbeddingInput]):
149
  """Immutable disk-backed normalized inputs with an incremental digest."""
150
 
 
152
  self,
153
  values: Iterable[str | EmbeddingInput | tuple[str, str]],
154
  ) -> None:
155
+ self._files = _SpoolFiles()
156
+ # Called by close(), or by garbage collection however the spool is freed; at most once.
157
+ self._release = weakref.finalize(self, self._files.release)
158
+ self.path = self._files.directory / "inputs.sqlite"
 
 
 
 
 
159
  digest = hashlib.sha256()
160
  count = 0
161
  pending: list[tuple[int, str, str]] = []
162
  try:
163
+ connection = self._files.connection = sqlite3.connect(self.path)
164
+ connection.execute(
165
+ "CREATE TABLE inputs ("
166
+ "position INTEGER PRIMARY KEY, input_id TEXT NOT NULL, sequence TEXT NOT NULL)"
167
+ )
168
  for position, item in enumerate(values):
169
  record = _normalize_input_item(position, item)
170
  for value in (record.id, record.sequence):
 
174
  pending.append((position, record.id, record.sequence))
175
  count += 1
176
  if len(pending) == 1_024:
177
+ connection.executemany("INSERT INTO inputs VALUES (?, ?, ?)", pending)
178
  pending.clear()
179
  if pending:
180
+ connection.executemany("INSERT INTO inputs VALUES (?, ?, ?)", pending)
181
  if count == 0:
182
  raise ValueError("inputs must contain at least one sequence.")
183
+ connection.commit()
184
+ connection.close()
185
+ self._files.connection = sqlite3.connect(
186
  f"{self.path.resolve().as_uri()}?mode=ro",
187
  uri=True,
188
  )
 
194
  self._count = count
195
 
196
  def _require_connection(self) -> sqlite3.Connection:
197
+ connection = self._files.connection
198
+ if connection is None:
199
  raise RuntimeError("Input spool is closed.")
200
+ return connection
201
 
202
  def __len__(self) -> int:
203
  return self._count
 
240
  return EmbeddingInput(row[0], row[1])
241
 
242
  def close(self) -> None:
243
+ self._release()
 
 
 
 
 
 
 
 
 
 
244
 
245
 
246
  def _normalize_inputs(
fastplms/embeddings/output.py CHANGED
@@ -125,11 +125,12 @@ class EmbeddingOutput:
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
@@ -209,7 +210,9 @@ class EmbeddingOutput:
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
 
 
 
125
  }
126
  self.sqlite_run_id = run_fingerprint
127
  if not resume and output_already_exists:
128
+ prior_run_readable = True
129
  try:
130
  load_sqlite_result(output, run_id=run_fingerprint)
131
  except KeyError:
132
+ prior_run_readable = False
133
+ if prior_run_readable:
134
  # Keep an exact prior run readable until replacement inference
135
  # has produced the first complete commit window.
136
  self.sqlite_replace_on_first_commit = True
 
210
  return load_sqlite_result(self.output, run_id=self.sqlite_run_id)
211
  if self.safetensors_writer is not None:
212
  return self.safetensors_writer.publish(complete=True, metadata=metadata)
213
+ embedding_result = EmbeddingResult(self.output_records, metadata)
214
  if self.output is not None:
215
+ return save_result(
216
+ embedding_result, self.output, format=self.format, shard_size=self.shard_size
217
+ )
218
+ return embedding_result
fastplms/embeddings/pooling.py CHANGED
@@ -10,6 +10,63 @@ from torch import Tensor
10
 
11
 
12
  POOLING_NAMES = frozenset({"mean", "max", "norm", "median", "std", "var", "cls", "parti"})
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
13
 
14
 
15
  def _validate_inputs(X: Tensor, M: Tensor) -> Tensor:
@@ -42,6 +99,7 @@ def _pooled_attention(attentions: Tensor | Sequence[Tensor], *, batch_size: int)
42
  must not change that reduction.
43
  """
44
 
 
45
  if isinstance(attentions, Sequence):
46
  if not attentions:
47
  raise ValueError("parti received an empty attention sequence.")
@@ -164,8 +222,12 @@ class Pooler:
164
  attentions: Tensor | Sequence[Tensor] | None = None,
165
  attention_backend: str | None = None,
166
  ) -> Tensor:
167
- # X: (b, l, d); residue_mask: (b, l)
168
  M = _validate_inputs(X, residue_mask) # (b, l)
 
 
 
 
169
  M_expanded = M.unsqueeze(-1) # (b, l, 1)
170
  count = M_expanded.sum(dim=1).clamp_min(1) # (b, 1)
171
  X_residues = X.masked_fill(~M_expanded, 0) # (b, l, d)
@@ -206,6 +268,7 @@ class Pooler:
206
  w = pagerank_weights(A_residue).to(dtype=X.dtype) # (r,)
207
  pooled.append(w @ X_i.index_select(0, indices)) # (d,)
208
  Y = torch.stack(pooled) # (b, d)
 
209
  if not bool(torch.isfinite(Y).all()):
210
  raise ValueError(
211
  f"Pooling operation {name!r} produced non-finite output from "
@@ -216,4 +279,7 @@ class Pooler:
216
  return torch.cat(outputs, dim=-1) # (b, len(self.names) * d)
217
 
218
 
219
- __all__ = ["POOLING_NAMES", "Pooler", "pagerank_weights"]
 
 
 
 
10
 
11
 
12
  POOLING_NAMES = frozenset({"mean", "max", "norm", "median", "std", "var", "cls", "parti"})
13
+ POOLING_SEMANTICS = {
14
+ "version": 2,
15
+ "mask": "biological_residues_only",
16
+ "accumulator": "float64_for_float64_input_else_float32",
17
+ "output_dtype": "input_dtype_after_reduction",
18
+ "variance_correction": 0,
19
+ "empty_residues": "reject",
20
+ "singleton_variance": 0,
21
+ }
22
+
23
+
24
+ # Pooling over every attended token, CLS and EOS included: the canonical policy of a store that
25
+ # keeps the special tokens. Version 3 supersedes the residue-only version 2 for those stores.
26
+ POOLING_SEMANTICS_TOKENS = {
27
+ "version": 3,
28
+ "mask": "all_attended_tokens_including_cls_and_eos",
29
+ "accumulator": "float64_for_float64_input_else_float32",
30
+ "output_dtype": "input_dtype_after_reduction",
31
+ "variance_correction": 0,
32
+ "empty_residues": "reject",
33
+ "singleton_variance": 0,
34
+ }
35
+ TOKEN_POOLING_NAMES = ("mean", "var", "std", "max", "norm")
36
+
37
+
38
+ def pool_token_rows(X: Tensor, token_mask: Tensor, names: Sequence[str]) -> Tensor:
39
+ """Pool every attended token of each sequence, CLS and EOS included, without any device read.
40
+
41
+ The reductions equal ``Pooler`` over a mask that is true on every token: float32 accumulation
42
+ (float64 for float64 input), population variance from the two-pass mean, and a cast to the input
43
+ dtype after reduction. Unlike ``Pooler`` this skips its input validation, which reads device
44
+ values and stalls the host; the caller checks finiteness on the device after the copy lands.
45
+ """
46
+ # X: (b, n, d) padded, n = the longest l + 2; token_mask: (b, n) true on CLS, residues and EOS
47
+ unsupported = [name for name in names if name not in TOKEN_POOLING_NAMES]
48
+ if unsupported or not names or len(set(names)) != len(names):
49
+ raise ValueError(f"Token pooling supports {TOKEN_POOLING_NAMES}, each once; received {list(names)}.")
50
+ work = torch.float64 if X.dtype == torch.float64 else torch.float32
51
+ values = X.to(work) # (b, n, d)
52
+ kept = token_mask.unsqueeze(-1) # (b, n, 1)
53
+ count = kept.sum(dim=1).clamp_min(1).to(work) # (b, 1), attended rows per sequence: l + 2
54
+ masked = values.masked_fill(~kept, 0) # (b, n, d), padding zeroed
55
+ mean = masked.sum(dim=1) / count # (b, d)
56
+ outputs: list[Tensor] = []
57
+ for name in names:
58
+ if name == "mean":
59
+ pooled = mean # (b, d)
60
+ elif name == "max":
61
+ pooled = values.masked_fill(~kept, -torch.inf).amax(dim=1) # (b, d)
62
+ elif name == "norm":
63
+ pooled = torch.linalg.vector_norm(masked, ord=2, dim=1) # (b, d)
64
+ else:
65
+ centered = (values - mean.unsqueeze(1)).masked_fill(~kept, 0) # (b, n, d)
66
+ variance = (centered * centered).sum(dim=1) / count # (b, d)
67
+ pooled = variance.sqrt() if name == "std" else variance # (b, d)
68
+ outputs.append(pooled.to(X.dtype))
69
+ return torch.cat(outputs, dim=-1) # (b, len(names) * d)
70
 
71
 
72
  def _validate_inputs(X: Tensor, M: Tensor) -> Tensor:
 
99
  must not change that reduction.
100
  """
101
 
102
+ # attentions: (b, ..., l, l), or a sequence of (b, h, l, l) layer maps
103
  if isinstance(attentions, Sequence):
104
  if not attentions:
105
  raise ValueError("parti received an empty attention sequence.")
 
222
  attentions: Tensor | Sequence[Tensor] | None = None,
223
  attention_backend: str | None = None,
224
  ) -> Tensor:
225
+ # X: (b, l, d); residue_mask: (b, l); attentions: (b, ..., l, l), or a sequence of (b, h, l, l) layer maps
226
  M = _validate_inputs(X, residue_mask) # (b, l)
227
+ output_dtype = X.dtype
228
+ # Sum/variance in FP16 can overflow even when the final answer is representable.
229
+ # Retain FP64 precision, otherwise accumulate in FP32 and cast only the result.
230
+ X = X.to(dtype=torch.float64 if X.dtype == torch.float64 else torch.float32) # (b, l, d)
231
  M_expanded = M.unsqueeze(-1) # (b, l, 1)
232
  count = M_expanded.sum(dim=1).clamp_min(1) # (b, 1)
233
  X_residues = X.masked_fill(~M_expanded, 0) # (b, l, d)
 
268
  w = pagerank_weights(A_residue).to(dtype=X.dtype) # (r,)
269
  pooled.append(w @ X_i.index_select(0, indices)) # (d,)
270
  Y = torch.stack(pooled) # (b, d)
271
+ Y = Y.to(dtype=output_dtype) # (b, d)
272
  if not bool(torch.isfinite(Y).all()):
273
  raise ValueError(
274
  f"Pooling operation {name!r} produced non-finite output from "
 
279
  return torch.cat(outputs, dim=-1) # (b, len(self.names) * d)
280
 
281
 
282
+ __all__ = [
283
+ "POOLING_NAMES", "POOLING_SEMANTICS_TOKENS", "TOKEN_POOLING_NAMES", "Pooler", "pagerank_weights",
284
+ "pool_token_rows",
285
+ ]
fastplms/embeddings/runner.py CHANGED
@@ -12,6 +12,7 @@ from torch import Tensor
12
  from . import identity
13
  from .batches import (
14
  BatchExecutor,
 
15
  _residue_embeddings as _residue_embeddings,
16
  _temporary_eval,
17
  select_hidden_state_embeddings as select_hidden_state_embeddings,
@@ -36,8 +37,11 @@ from .inputs import (
36
  parse_fasta as parse_fasta,
37
  )
38
  from .output import EmbeddingOutput
39
- from .pooling import Pooler
40
- from .types import EmbeddingBatch, EmbeddingInput, EmbeddingResult
 
 
 
41
 
42
 
43
  _DEFAULT_BATCH_WINDOW_MULTIPLIER = 16
@@ -51,6 +55,9 @@ def embed_dataset(
51
  batch_size: int = 2,
52
  pooling: str | Sequence[str] | None = None,
53
  full_embeddings: bool = False,
 
 
 
54
  output: str | Path | None = None,
55
  format: str = "safetensors",
56
  resume: bool = True,
@@ -70,9 +77,16 @@ def embed_dataset(
70
  _embedding_batch_identity: Mapping[str, Any] | None = None,
71
  _allowed_unsupported_pooling: Sequence[str] = (),
72
  **model_kwargs: Any,
73
- ) -> EmbeddingResult:
74
- """Embed protein sequences with stable ordering and residue-only pooling."""
 
 
 
 
 
 
75
 
 
76
  for name, value in (
77
  ("batch_size", batch_size),
78
  ("shard_size", shard_size),
@@ -94,11 +108,16 @@ def embed_dataset(
94
  raise ValueError(f"{optional_name} must be a positive integer when provided.")
95
  for name, value in (
96
  ("full_embeddings", full_embeddings),
 
97
  ("resume", resume),
98
  ("truncate", truncate),
99
  ):
100
  if not isinstance(value, bool):
101
  raise TypeError(f"{name} must be a boolean.")
 
 
 
 
102
  if not isinstance(format, str):
103
  raise TypeError("format must be a string.")
104
  if output is not None and not isinstance(output, (str, Path)):
@@ -135,15 +154,26 @@ def embed_dataset(
135
  raise ValueError("decoder_attention_mask must contain finite binary values.")
136
  if not bool(((decoder_attention_mask == 0) | (decoder_attention_mask == 1)).all()):
137
  raise ValueError("decoder_attention_mask must contain finite binary values.")
138
- pooling_names = (
139
- (("mean",) if not full_embeddings else ())
140
- if pooling is None
141
- else ((pooling,) if isinstance(pooling, str) else tuple(pooling))
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
142
  )
143
- if full_embeddings and pooling is not None:
144
- raise ValueError("full_embeddings=True cannot be combined with pooling.")
145
- if not full_embeddings and not pooling_names:
146
- raise ValueError("pooling is required unless full_embeddings=True.")
147
  pooler = Pooler(pooling_names) if pooling_names else None
148
 
149
  if batch_size <= 0:
@@ -187,22 +217,12 @@ def embed_dataset(
187
  )
188
  if resolved_batch_window_size < batch_size:
189
  raise ValueError("batch_window_size must be at least batch_size.")
190
- records = _normalize_inputs(inputs, disk_backed=output is not None)
191
  _validate_untruncated_lengths(
192
  records,
193
  max_length=max_length,
194
  truncate=truncate,
195
  )
196
- pooling_names = (
197
- (("mean",) if not full_embeddings else ())
198
- if pooling is None
199
- else ((pooling,) if isinstance(pooling, str) else tuple(pooling))
200
- )
201
- if full_embeddings:
202
- if pooling is not None:
203
- raise ValueError("full_embeddings=True cannot be combined with pooling.")
204
- elif not pooling_names:
205
- raise ValueError("pooling is required unless full_embeddings=True.")
206
  store_all_hidden_states = bool(model_kwargs.get("store_all_hidden_states", False))
207
  if store_all_hidden_states and not full_embeddings:
208
  raise ValueError("store_all_hidden_states=True requires full_embeddings=True.")
@@ -215,7 +235,10 @@ def embed_dataset(
215
  f"by the model; unknown overrides: {sorted(unknown_pooling_overrides)}."
216
  )
217
  unsupported.difference_update(allowed_unsupported_pooling)
218
- requested_unsupported = unsupported.intersection(pooling_names)
 
 
 
219
  if requested_unsupported:
220
  raise ValueError(
221
  f"{model.__class__.__name__} does not support pooling operations "
@@ -248,6 +271,7 @@ def embed_dataset(
248
  model.resolve_attn_implementation()
249
 
250
  tokenizer_metadata = _tokenizer_metadata(model, tokenizer)
 
251
  (
252
  input_fingerprint,
253
  run_fingerprint,
@@ -269,7 +293,64 @@ def embed_dataset(
269
  batch_size=batch_size,
270
  batch_window_size=resolved_batch_window_size,
271
  max_tokens_per_batch=max_tokens_per_batch,
 
272
  )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
273
  destination = EmbeddingOutput(
274
  records,
275
  output=output,
@@ -286,7 +367,6 @@ def embed_dataset(
286
  if destination.completed is not None:
287
  return destination.completed
288
 
289
- attention_backend = _attention_backend(model)
290
  executor = BatchExecutor(
291
  model=model,
292
  batch_size=batch_size,
@@ -321,13 +401,181 @@ def embed_dataset(
321
  )
322
  destination.append(window_start, new_records)
323
 
324
- software_versions = identity._software_versions()
325
- projection = getattr(model, "embedding_projection", None)
326
- resolved_layer = getattr(
327
  model,
328
- "embedding_layer",
329
- model_kwargs.get("hidden_state_index", -1),
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
330
  )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
331
  token_policy = getattr(
332
  model,
333
  "embedding_token_policy",
@@ -349,14 +597,14 @@ def embed_dataset(
349
  "fingerprint_schema_version": _RUN_FINGERPRINT_SCHEMA_VERSION,
350
  "run_fingerprint": run_fingerprint,
351
  "input_fingerprint": input_fingerprint,
352
- "model_state_fingerprint": resolved_model_state_fingerprint,
353
  "model_state_fingerprint_source": model_state_fingerprint_source,
354
  "model_class": f"{model.__class__.__module__}.{model.__class__.__qualname__}",
355
  **model_identity,
356
  "dtype": str(dtype).removeprefix("torch.") if dtype is not None else "model",
357
  "attention_backend": attention_backend,
358
  "attention_kernel": _attention_kernel_metadata(attention_backend),
359
- "layer": resolved_layer,
360
  "projection": projection,
361
  "esmc_source": getattr(model, "_esmc_source", None),
362
  "esmc_revision": getattr(model, "_esmc_source_revision", None),
@@ -365,14 +613,18 @@ def embed_dataset(
365
  "tokenizer": tokenizer_metadata,
366
  **embedding_context,
367
  "pooling": list(pooling_names),
 
368
  "pool_slices": pool_slices,
369
  "full_embeddings": full_embeddings,
370
  "max_length": max_length,
371
  "truncate": truncate,
372
  "truncation": {"enabled": truncate, "max_length": max_length},
 
 
 
373
  "batching": {
374
  "batch_size": batch_size,
375
- "batch_window_size": resolved_batch_window_size,
376
  "max_tokens_per_batch": max_tokens_per_batch,
377
  "input_storage": ("disk-spool" if isinstance(records, _InputSpool) else "memory"),
378
  "ordering": "bounded-length-bucketed-stable-output",
@@ -386,13 +638,7 @@ def embed_dataset(
386
  },
387
  "residue_mask_policy": "biological-residues-only",
388
  "record_count": len(records),
389
- "descriptor_index": (
390
- "memory-metadata"
391
- if output is None
392
- else "sqlite-records"
393
- if format == "sqlite"
394
- else "safetensors-generation-index"
395
- ),
396
  "storage_format": format if output is not None else "memory",
397
  "software": software_versions,
398
  "execution": _execution_identity_metadata(model),
@@ -401,19 +647,20 @@ def embed_dataset(
401
  "transformers_version": software_versions["transformers"],
402
  "complete": True,
403
  }
404
- if destination.output_descriptors is not None:
405
- metadata["outputs"] = destination.output_descriptors
406
- metadata["tensor_hashes"] = [item["sha256"] for item in destination.output_descriptors]
407
  status = getattr(model, "esmc_precision_status", None)
408
  if status is not None:
409
  metadata["esmc_precision"] = status.as_dict() if hasattr(status, "as_dict") else status
410
- return destination.finish(metadata)
411
 
412
 
413
  class EmbeddingMixin:
414
  """Small delegation mixin shared by FastPLMs model classes."""
415
 
416
- def embed_dataset(self, inputs: Any, **kwargs: Any) -> EmbeddingResult:
 
 
417
  return embed_dataset(self, inputs, **kwargs)
418
 
419
 
 
12
  from . import identity
13
  from .batches import (
14
  BatchExecutor,
15
+ TapExecutor,
16
  _residue_embeddings as _residue_embeddings,
17
  _temporary_eval,
18
  select_hidden_state_embeddings as select_hidden_state_embeddings,
 
37
  parse_fasta as parse_fasta,
38
  )
39
  from .output import EmbeddingOutput
40
+ from .pooling import POOLING_SEMANTICS, Pooler
41
+ from .taps import Tap, TapPlan, plan_taps
42
+ from .types import (
43
+ EmbeddingBatch, EmbeddingInput, EmbeddingResult, TapRecord, TapResult, TapRunReceipt,
44
+ )
45
 
46
 
47
  _DEFAULT_BATCH_WINDOW_MULTIPLIER = 16
 
55
  batch_size: int = 2,
56
  pooling: str | Sequence[str] | None = None,
57
  full_embeddings: bool = False,
58
+ taps: Sequence[Tap] | None = None,
59
+ tap_sink: Callable[[Sequence[TapRecord], Mapping[str, str]], None] | None = None,
60
+ require_residue_identity: bool = False,
61
  output: str | Path | None = None,
62
  format: str = "safetensors",
63
  resume: bool = True,
 
77
  _embedding_batch_identity: Mapping[str, Any] | None = None,
78
  _allowed_unsupported_pooling: Sequence[str] = (),
79
  **model_kwargs: Any,
80
+ ) -> EmbeddingResult | TapResult | TapRunReceipt:
81
+ """Embed protein sequences with stable ordering and residue-only pooling.
82
+
83
+ ``taps`` instead returns a ``TapResult``: each tap's output from one forward pass per
84
+ batch, kept in memory. With ``tap_sink``, deliver bounded windows to the callback and return
85
+ a ``TapRunReceipt`` without retaining their tensors. The callback receives ordered records
86
+ and the run/input fingerprints; it must release the records to preserve bounded memory.
87
+ """
88
 
89
+ # decoder_input_ids, decoder_attention_mask: (n_records, l_decoder), aligned with the inputs
90
  for name, value in (
91
  ("batch_size", batch_size),
92
  ("shard_size", shard_size),
 
108
  raise ValueError(f"{optional_name} must be a positive integer when provided.")
109
  for name, value in (
110
  ("full_embeddings", full_embeddings),
111
+ ("require_residue_identity", require_residue_identity),
112
  ("resume", resume),
113
  ("truncate", truncate),
114
  ):
115
  if not isinstance(value, bool):
116
  raise TypeError(f"{name} must be a boolean.")
117
+ if require_residue_identity and taps is None:
118
+ raise ValueError("Residue identity validation requires a tap plan.")
119
+ if tap_sink is not None and (taps is None or not callable(tap_sink)):
120
+ raise ValueError("tap_sink requires a tap plan and a callable destination.")
121
  if not isinstance(format, str):
122
  raise TypeError("format must be a string.")
123
  if output is not None and not isinstance(output, (str, Path)):
 
154
  raise ValueError("decoder_attention_mask must contain finite binary values.")
155
  if not bool(((decoder_attention_mask == 0) | (decoder_attention_mask == 1)).all()):
156
  raise ValueError("decoder_attention_mask must contain finite binary values.")
157
+ tap_plan = (
158
+ _tap_plan(
159
+ model,
160
+ taps,
161
+ pooling=pooling,
162
+ full_embeddings=full_embeddings,
163
+ output=output,
164
+ model_kwargs=model_kwargs,
165
+ family_adapter=(
166
+ _embedding_batch_fn is not None or _embedding_batch_identity is not None
167
+ ),
168
+ )
169
+ if taps is not None
170
+ else None
171
+ )
172
+ pooling_names = _requested_pooling(
173
+ pooling,
174
+ full_embeddings=full_embeddings,
175
+ taps_requested=tap_plan is not None,
176
  )
 
 
 
 
177
  pooler = Pooler(pooling_names) if pooling_names else None
178
 
179
  if batch_size <= 0:
 
217
  )
218
  if resolved_batch_window_size < batch_size:
219
  raise ValueError("batch_window_size must be at least batch_size.")
220
+ records = _normalize_inputs(inputs, disk_backed=output is not None or tap_sink is not None)
221
  _validate_untruncated_lengths(
222
  records,
223
  max_length=max_length,
224
  truncate=truncate,
225
  )
 
 
 
 
 
 
 
 
 
 
226
  store_all_hidden_states = bool(model_kwargs.get("store_all_hidden_states", False))
227
  if store_all_hidden_states and not full_embeddings:
228
  raise ValueError("store_all_hidden_states=True requires full_embeddings=True.")
 
235
  f"by the model; unknown overrides: {sorted(unknown_pooling_overrides)}."
236
  )
237
  unsupported.difference_update(allowed_unsupported_pooling)
238
+ requested_pooling = set(pooling_names)
239
+ if tap_plan is not None:
240
+ requested_pooling.update(tap_plan.pooling_names)
241
+ requested_unsupported = unsupported.intersection(requested_pooling)
242
  if requested_unsupported:
243
  raise ValueError(
244
  f"{model.__class__.__name__} does not support pooling operations "
 
271
  model.resolve_attn_implementation()
272
 
273
  tokenizer_metadata = _tokenizer_metadata(model, tokenizer)
274
+ tap_identity = tap_plan.identity() if tap_plan is not None else None
275
  (
276
  input_fingerprint,
277
  run_fingerprint,
 
293
  batch_size=batch_size,
294
  batch_window_size=resolved_batch_window_size,
295
  max_tokens_per_batch=max_tokens_per_batch,
296
+ taps=tap_identity,
297
  )
298
+ attention_backend = _attention_backend(model)
299
+ if tap_plan is not None:
300
+ tap_records, tap_pool_slices = _embed_tap_windows(
301
+ model,
302
+ records,
303
+ TapExecutor(
304
+ model=model,
305
+ plan=tap_plan,
306
+ batch_size=batch_size,
307
+ max_tokens_per_batch=max_tokens_per_batch,
308
+ max_length=max_length,
309
+ truncate=truncate,
310
+ tokenizer=tokenizer,
311
+ dtype=dtype,
312
+ attention_backend=attention_backend,
313
+ require_residue_identity=require_residue_identity,
314
+ ),
315
+ window_size=resolved_batch_window_size,
316
+ sink=tap_sink,
317
+ run_identity={
318
+ "run_fingerprint": run_fingerprint, "input_fingerprint": input_fingerprint,
319
+ },
320
+ )
321
+ metadata = _run_metadata(
322
+ model,
323
+ records,
324
+ run_fingerprint=run_fingerprint,
325
+ input_fingerprint=input_fingerprint,
326
+ model_state_fingerprint=resolved_model_state_fingerprint,
327
+ model_state_fingerprint_source=model_state_fingerprint_source,
328
+ dtype=dtype,
329
+ attention_backend=attention_backend,
330
+ layer=None,
331
+ tokenizer_metadata=tokenizer_metadata,
332
+ embedding_context=embedding_context,
333
+ pooling_names=pooling_names,
334
+ pool_slices={},
335
+ full_embeddings=full_embeddings,
336
+ max_length=max_length,
337
+ truncate=truncate,
338
+ batch_size=batch_size,
339
+ batch_window_size=resolved_batch_window_size,
340
+ max_tokens_per_batch=max_tokens_per_batch,
341
+ output=output,
342
+ format=format,
343
+ descriptor_index="not-recorded",
344
+ taps={
345
+ "plan": tap_identity,
346
+ "stop_after_layer": tap_plan.deepest_layer,
347
+ "pool_slices": tap_pool_slices,
348
+ },
349
+ )
350
+ if tap_sink is not None:
351
+ metadata["storage_format"] = "tap-sink"
352
+ return TapRunReceipt(len(records), metadata)
353
+ return TapResult(tap_records, metadata)
354
  destination = EmbeddingOutput(
355
  records,
356
  output=output,
 
367
  if destination.completed is not None:
368
  return destination.completed
369
 
 
370
  executor = BatchExecutor(
371
  model=model,
372
  batch_size=batch_size,
 
401
  )
402
  destination.append(window_start, new_records)
403
 
404
+ metadata = _run_metadata(
 
 
405
  model,
406
+ records,
407
+ run_fingerprint=run_fingerprint,
408
+ input_fingerprint=input_fingerprint,
409
+ model_state_fingerprint=resolved_model_state_fingerprint,
410
+ model_state_fingerprint_source=model_state_fingerprint_source,
411
+ dtype=dtype,
412
+ attention_backend=attention_backend,
413
+ layer=getattr(
414
+ model,
415
+ "embedding_layer",
416
+ model_kwargs.get("hidden_state_index", -1),
417
+ ),
418
+ tokenizer_metadata=tokenizer_metadata,
419
+ embedding_context=embedding_context,
420
+ pooling_names=pooling_names,
421
+ pool_slices=pool_slices,
422
+ full_embeddings=full_embeddings,
423
+ max_length=max_length,
424
+ truncate=truncate,
425
+ batch_size=batch_size,
426
+ batch_window_size=resolved_batch_window_size,
427
+ max_tokens_per_batch=max_tokens_per_batch,
428
+ output=output,
429
+ format=format,
430
+ descriptor_index=(
431
+ "memory-metadata"
432
+ if output is None
433
+ else "sqlite-records"
434
+ if format == "sqlite"
435
+ else "safetensors-generation-index"
436
+ ),
437
  )
438
+ if destination.output_descriptors is not None:
439
+ metadata["outputs"] = destination.output_descriptors
440
+ metadata["tensor_hashes"] = [item["sha256"] for item in destination.output_descriptors]
441
+ return destination.finish(metadata)
442
+
443
+
444
+ def _tap_plan(
445
+ model: Any,
446
+ taps: Sequence[Tap],
447
+ *,
448
+ pooling: str | Sequence[str] | None,
449
+ full_embeddings: bool,
450
+ output: str | Path | None,
451
+ model_kwargs: Mapping[str, Any],
452
+ family_adapter: bool,
453
+ ) -> TapPlan:
454
+ """Validate a tap request against the arguments it excludes and the model's hidden states."""
455
+
456
+ excluded = [
457
+ name
458
+ for name, requested in (
459
+ ("pooling", pooling is not None),
460
+ ("full_embeddings", full_embeddings),
461
+ ("hidden_state_index", "hidden_state_index" in model_kwargs),
462
+ ("store_all_hidden_states", "store_all_hidden_states" in model_kwargs),
463
+ )
464
+ if requested
465
+ ]
466
+ if excluded:
467
+ raise ValueError(
468
+ f"taps= cannot be combined with {', '.join(excluded)}; each tap names its own "
469
+ "layer and pooling."
470
+ )
471
+ if model_kwargs:
472
+ raise ValueError(
473
+ f"taps= takes no model keyword arguments; received {sorted(model_kwargs)}."
474
+ )
475
+ if family_adapter:
476
+ raise ValueError(
477
+ "taps= runs the model's own one-pass path; _embedding_batch_fn and "
478
+ "_embedding_batch_identity do not apply."
479
+ )
480
+ if output is not None:
481
+ raise ValueError(
482
+ "taps= returns its records in memory; omit output=. To persist them, call "
483
+ "embed_into_features, which writes each tap into the feature store of its key and "
484
+ "embeds only the sequences that store lacks."
485
+ )
486
+ if getattr(model, "embedding_tap_support", False) is not True:
487
+ raise ValueError(
488
+ f"{model.__class__.__name__} does not support taps=. One-pass taps need a model "
489
+ "family that implements them, such as ESM++ (ESMC)."
490
+ )
491
+ return plan_taps(taps, int(model.embedding_tap_state_count))
492
+
493
+
494
+ def _requested_pooling(
495
+ pooling: str | Sequence[str] | None,
496
+ *,
497
+ full_embeddings: bool,
498
+ taps_requested: bool,
499
+ ) -> tuple[str, ...]:
500
+ """Pooler names of a single-output run: mean by default, none for residues or taps."""
501
+
502
+ if taps_requested:
503
+ return ()
504
+ if full_embeddings:
505
+ if pooling is not None:
506
+ raise ValueError("full_embeddings=True cannot be combined with pooling.")
507
+ return ()
508
+ names = (
509
+ ("mean",)
510
+ if pooling is None
511
+ else ((pooling,) if isinstance(pooling, str) else tuple(pooling))
512
+ )
513
+ if not names:
514
+ raise ValueError("pooling is required unless full_embeddings=True.")
515
+ return names
516
+
517
+
518
+ def _embed_tap_windows(
519
+ model: Any,
520
+ records: Sequence[EmbeddingInput],
521
+ executor: TapExecutor,
522
+ *,
523
+ window_size: int,
524
+ sink: Callable[[Sequence[TapRecord], Mapping[str, str]], None] | None = None,
525
+ run_identity: Mapping[str, str] | None = None,
526
+ ) -> tuple[list[TapRecord], dict[str, dict[str, tuple[int, int]]]]:
527
+ """Run bounded windows in source order, retaining tensors only without a sink."""
528
+
529
+ tap_records: list[TapRecord] = []
530
+ pool_slices: dict[str, dict[str, tuple[int, int]]] = {}
531
+ with _temporary_eval(model), torch.inference_mode():
532
+ for window_start in range(0, len(records), window_size):
533
+ window_stop = min(window_start + window_size, len(records))
534
+ window_records = records[window_start:window_stop]
535
+ if not isinstance(window_records, Sequence):
536
+ raise RuntimeError("The immutable embedding spool returned a non-sequence window.")
537
+ new_records, pool_slices = executor.run_window(
538
+ window_records, window_start=window_start
539
+ )
540
+ if sink is None:
541
+ tap_records.extend(new_records)
542
+ else:
543
+ sink(new_records, dict(run_identity or {}))
544
+ # Release the previous window before allocating the next, including on CPU.
545
+ del new_records
546
+ return tap_records, pool_slices
547
+
548
+
549
+ def _run_metadata(
550
+ model: Any,
551
+ records: Sequence[EmbeddingInput],
552
+ *,
553
+ run_fingerprint: str,
554
+ input_fingerprint: str,
555
+ model_state_fingerprint: str | None,
556
+ model_state_fingerprint_source: str,
557
+ dtype: torch.dtype | None,
558
+ attention_backend: str | None,
559
+ layer: Any,
560
+ tokenizer_metadata: dict[str, Any],
561
+ embedding_context: Mapping[str, Any],
562
+ pooling_names: Sequence[str],
563
+ pool_slices: Mapping[str, tuple[int, int]],
564
+ full_embeddings: bool,
565
+ max_length: int | None,
566
+ truncate: bool,
567
+ batch_size: int,
568
+ batch_window_size: int,
569
+ max_tokens_per_batch: int | None,
570
+ output: str | Path | None,
571
+ format: str,
572
+ descriptor_index: str,
573
+ taps: Mapping[str, Any] | None = None,
574
+ ) -> dict[str, Any]:
575
+ """Everything a finished run records so that it can be reproduced and resumed."""
576
+
577
+ software_versions = identity._software_versions()
578
+ projection = getattr(model, "embedding_projection", None)
579
  token_policy = getattr(
580
  model,
581
  "embedding_token_policy",
 
597
  "fingerprint_schema_version": _RUN_FINGERPRINT_SCHEMA_VERSION,
598
  "run_fingerprint": run_fingerprint,
599
  "input_fingerprint": input_fingerprint,
600
+ "model_state_fingerprint": model_state_fingerprint,
601
  "model_state_fingerprint_source": model_state_fingerprint_source,
602
  "model_class": f"{model.__class__.__module__}.{model.__class__.__qualname__}",
603
  **model_identity,
604
  "dtype": str(dtype).removeprefix("torch.") if dtype is not None else "model",
605
  "attention_backend": attention_backend,
606
  "attention_kernel": _attention_kernel_metadata(attention_backend),
607
+ "layer": layer,
608
  "projection": projection,
609
  "esmc_source": getattr(model, "_esmc_source", None),
610
  "esmc_revision": getattr(model, "_esmc_source_revision", None),
 
613
  "tokenizer": tokenizer_metadata,
614
  **embedding_context,
615
  "pooling": list(pooling_names),
616
+ "pooling_semantics": dict(POOLING_SEMANTICS),
617
  "pool_slices": pool_slices,
618
  "full_embeddings": full_embeddings,
619
  "max_length": max_length,
620
  "truncate": truncate,
621
  "truncation": {"enabled": truncate, "max_length": max_length},
622
+ "retained_positions": (
623
+ "biological_residues_in_input_order_after_optional_prefix_crop_before_forward"
624
+ ),
625
  "batching": {
626
  "batch_size": batch_size,
627
+ "batch_window_size": batch_window_size,
628
  "max_tokens_per_batch": max_tokens_per_batch,
629
  "input_storage": ("disk-spool" if isinstance(records, _InputSpool) else "memory"),
630
  "ordering": "bounded-length-bucketed-stable-output",
 
638
  },
639
  "residue_mask_policy": "biological-residues-only",
640
  "record_count": len(records),
641
+ "descriptor_index": descriptor_index,
 
 
 
 
 
 
642
  "storage_format": format if output is not None else "memory",
643
  "software": software_versions,
644
  "execution": _execution_identity_metadata(model),
 
647
  "transformers_version": software_versions["transformers"],
648
  "complete": True,
649
  }
650
+ if taps is not None:
651
+ metadata["taps"] = taps
 
652
  status = getattr(model, "esmc_precision_status", None)
653
  if status is not None:
654
  metadata["esmc_precision"] = status.as_dict() if hasattr(status, "as_dict") else status
655
+ return metadata
656
 
657
 
658
  class EmbeddingMixin:
659
  """Small delegation mixin shared by FastPLMs model classes."""
660
 
661
+ def embed_dataset(
662
+ self, inputs: Any, **kwargs: Any,
663
+ ) -> EmbeddingResult | TapResult | TapRunReceipt:
664
  return embed_dataset(self, inputs, **kwargs)
665
 
666
 
fastplms/embeddings/storage.py CHANGED
@@ -22,6 +22,7 @@ from .types import (
22
  EmbeddingResult,
23
  LazyTensorReference,
24
  )
 
25
 
26
 
27
  _DTYPE_NAMES: dict[torch.dtype, str] = {
@@ -102,13 +103,15 @@ def _bounded_tensor_chunks(X: Tensor, max_bytes: int) -> Iterator[Tensor]:
102
 
103
 
104
  def _tensor_hash_chunks(X: Tensor) -> Iterator[bytes]:
105
- for chunk in _bounded_tensor_chunks(X, _TENSOR_HASH_CHUNK_BYTES):
 
106
  yield chunk.view(torch.uint8).numpy().tobytes()
107
 
108
 
109
  def tensor_sha256(X: Tensor) -> str:
110
  """Hash dtype, shape, and exact tensor bytes."""
111
 
 
112
  if not isinstance(X, Tensor):
113
  raise TypeError("X must be a tensor.")
114
  if X.dtype not in _DTYPE_NAMES:
@@ -126,22 +129,23 @@ def tensor_sha256(X: Tensor) -> str:
126
 
127
 
128
  def _encode_tensor(X: Tensor) -> tuple[str, str, bytes]:
 
129
  if X.dtype not in _DTYPE_NAMES:
130
  raise TypeError(f"Unsupported tensor dtype {X.dtype}.")
131
  shape = json.dumps(tuple(X.shape), separators=(",", ":"))
132
  return _DTYPE_NAMES[X.dtype], shape, _tensor_bytes(X)
133
 
134
 
135
- def _decode_tensor(dtype_name: str, shape_json: str, data: bytes) -> Tensor:
136
  try:
137
  dtype = _NAME_DTYPES[dtype_name]
138
  except KeyError as error:
139
  raise ValueError(f"Unsupported stored dtype {dtype_name!r}.") from error
140
  shape = tuple(json.loads(shape_json))
141
  # uint8 is used only as a byte-level carrier, preserving BF16 bits exactly.
142
- byte_array = np.frombuffer(data, dtype=np.uint8).copy() # (n_bytes,)
143
  X = torch.from_numpy(byte_array).view(dtype) # (n_elements,)
144
- return X.reshape(shape).clone() # shape
145
 
146
 
147
  def _index_path(path: str | Path) -> Path:
@@ -172,10 +176,6 @@ def _resolve_index_child(root: Path, relative: str, *, label: str) -> Path:
172
  return candidate
173
 
174
 
175
- def _canonical_json_bytes(payload: dict[str, Any]) -> bytes:
176
- return (json.dumps(payload, indent=2, sort_keys=True) + "\n").encode("utf-8")
177
-
178
-
179
  def _load_authoritative_index(
180
  path: str | Path,
181
  ) -> tuple[dict[str, Any], Path, dict[str, Any]]:
@@ -198,7 +198,7 @@ def _load_authoritative_index(
198
  snapshot = run_manifest.get("index_payload")
199
  if isinstance(snapshot, dict):
200
  payload = snapshot
201
- index_bytes = _canonical_json_bytes(payload)
202
  elif snapshot is None:
203
  index_bytes = stable_index_path.read_bytes()
204
  payload = json.loads(index_bytes.decode("utf-8"))
@@ -268,7 +268,7 @@ def _load_safetensor(path: Path, key: str) -> Tensor:
268
  except ImportError as error:
269
  raise ImportError("Loading embeddings requires the 'safetensors' package.") from error
270
  with safe_open(path, framework="pt", device="cpu") as handle:
271
- return cast(Tensor, handle.get_tensor(key))
272
 
273
 
274
  def _safetensors_shard_prefix(path: str | Path) -> str:
@@ -361,7 +361,7 @@ def _record_from_safetensors_descriptor(root: Path, item: dict[str, Any]) -> Emb
361
  raise ValueError(f"Safetensors tensor shard is missing: {relative}.")
362
 
363
  def load_tensor() -> Tensor:
364
- return _load_safetensor(tensor_path, key)
365
 
366
  reference = LazyTensorReference(
367
  source=str(tensor_path),
@@ -720,7 +720,7 @@ class SafetensorsStreamWriter:
720
  raise FileExistsError(
721
  f"Refusing to reuse immutable safetensors generation index {generation_index_path}."
722
  )
723
- encoded_index = _canonical_json_bytes(payload)
724
  temporary_generation_index.write_bytes(encoded_index)
725
  temporary_generation_index.replace(generation_index_path)
726
 
@@ -739,7 +739,7 @@ class SafetensorsStreamWriter:
739
  temporary_manifest = self.run_manifest_path.with_name(
740
  f".{self.run_manifest_path.name}.{pointer_identity}.tmp"
741
  )
742
- temporary_manifest.write_bytes(_canonical_json_bytes(run_manifest))
743
  temporary_manifest.replace(self.run_manifest_path)
744
 
745
  # ``index.json`` is a non-authoritative convenience pointer. The run
@@ -753,7 +753,7 @@ class SafetensorsStreamWriter:
753
  temporary_index = self.index_path.with_name(
754
  f".{self.index_path.name}.{pointer_identity}.tmp"
755
  )
756
- temporary_index.write_bytes(_canonical_json_bytes(stable_pointer))
757
  temporary_index.replace(self.index_path)
758
 
759
  return load_safetensors_result(self.index_path)
@@ -771,7 +771,7 @@ class SafetensorsStreamWriter:
771
 
772
 
773
  def save_safetensors_result(
774
- result: EmbeddingResult,
775
  path: str | Path,
776
  *,
777
  shard_size: int = DEFAULT_SHARD_SIZE,
@@ -780,13 +780,13 @@ def save_safetensors_result(
780
 
781
  writer = SafetensorsStreamWriter(
782
  path,
783
- result.metadata,
784
  shard_size=shard_size,
785
  publish_initial=False,
786
  publish_incremental=False,
787
  )
788
- writer.append(result, publish=False)
789
- return writer.publish(complete=bool(result.metadata.get("complete", True)))
790
 
791
 
792
  def load_safetensors_result(path: str | Path) -> EmbeddingResult:
@@ -925,19 +925,19 @@ def _ensure_sqlite_schema(connection: sqlite3.Connection) -> None:
925
  connection.commit()
926
 
927
 
928
- def save_sqlite_result(result: EmbeddingResult, path: str | Path) -> EmbeddingResult:
929
  """Transactionally store an ordered result in normalized SQLite tables."""
930
 
931
  path = Path(path)
932
  path.parent.mkdir(parents=True, exist_ok=True)
933
- run_id = str(result.metadata.get("run_fingerprint", ""))
934
  if not run_id:
935
  raise ValueError("SQLite results require metadata['run_fingerprint'].")
936
  metadata_json = json.dumps(
937
  _persistent_metadata(
938
- result.metadata,
939
  descriptor_index="sqlite-records",
940
- record_count=len(result),
941
  ),
942
  sort_keys=True,
943
  )
@@ -951,13 +951,13 @@ def save_sqlite_result(result: EmbeddingResult, path: str | Path) -> EmbeddingRe
951
  "SELECT ?, ?, COALESCE(MAX(published_order), 0) + 1 FROM runs",
952
  (run_id, metadata_json),
953
  )
954
- for position, record in enumerate(result):
955
  X = record.load_tensor().detach().cpu().contiguous() # (...)
956
- dtype_name, shape_json, data = _encode_tensor(X)
957
  digest = tensor_sha256(X)
958
  connection.execute(
959
  "INSERT INTO tensors VALUES (?, ?, ?, ?, ?, ?)",
960
- (run_id, position, dtype_name, shape_json, data, digest),
961
  )
962
  connection.execute(
963
  "INSERT INTO records VALUES (?, ?, ?, ?)",
@@ -1057,11 +1057,11 @@ def append_sqlite_records(
1057
  for offset, record in enumerate(records):
1058
  position = start_position + offset
1059
  X = record.load_tensor().detach().cpu().contiguous() # (...)
1060
- dtype_name, shape_json, data = _encode_tensor(X)
1061
  digest = tensor_sha256(X)
1062
  connection.execute(
1063
  "INSERT INTO tensors VALUES (?, ?, ?, ?, ?, ?)",
1064
- (run_id, position, dtype_name, shape_json, data, digest),
1065
  )
1066
  connection.execute(
1067
  "INSERT INTO records VALUES (?, ?, ?, ?)",
@@ -1142,7 +1142,7 @@ def _load_sqlite_tensor(path: Path, run_id: str, position: int) -> Tensor:
1142
  ).fetchone()
1143
  if row is None:
1144
  raise KeyError(f"Missing SQLite tensor {run_id}:{position}.")
1145
- return _decode_tensor(*row)
1146
 
1147
 
1148
  def _validate_sqlite_descriptor_row(
@@ -1180,7 +1180,7 @@ def _sqlite_record_from_row(path: Path, run_id: str, row: Sequence[Any]) -> Embe
1180
  )
1181
 
1182
  def load_tensor() -> Tensor:
1183
- return _load_sqlite_tensor(path, run_id, position)
1184
 
1185
  reference = LazyTensorReference(
1186
  source=str(path),
@@ -1284,7 +1284,7 @@ def load_sqlite_result(
1284
  _validate_sqlite_result_schema(connection, path)
1285
  if run_id is None:
1286
  run_columns = {
1287
- str(info[1]) for info in connection.execute("PRAGMA table_info(runs)").fetchall()
1288
  }
1289
  if "published_order" in run_columns:
1290
  row = connection.execute(
@@ -1434,36 +1434,36 @@ _LEGACY_CODE_DTYPES: dict[int, tuple[np.dtype[Any], torch.dtype]] = {
1434
 
1435
 
1436
  def _decode_legacy_sqlite_blob(
1437
- data: bytes,
1438
  *,
1439
  fallback_shape: tuple[int, ...] | None,
1440
  allow_unsafe_pickle: bool,
1441
  ) -> Tensor:
1442
- if len(data) >= 6 and data[0] == _LEGACY_COMPACT_VERSION:
1443
- dtype_code = int(data[1])
1444
  if dtype_code not in _LEGACY_CODE_DTYPES:
1445
  raise ValueError(f"Unsupported legacy compact dtype code {dtype_code}.")
1446
- (ndim,) = struct.unpack_from("<i", data, 2)
1447
- if ndim < 0 or ndim > 16 or len(data) < 6 + 4 * ndim:
1448
  raise ValueError("Malformed legacy compact embedding header.")
1449
- shape = tuple(int(value) for value in struct.unpack_from(f"<{ndim}i", data, 6))
1450
  if any(size < 0 for size in shape):
1451
  raise ValueError("Malformed negative legacy embedding dimension.")
1452
  numpy_dtype, target_dtype = _LEGACY_CODE_DTYPES[dtype_code]
1453
  offset = 6 + 4 * ndim
1454
  expected = int(np.prod(shape, dtype=np.int64)) * numpy_dtype.itemsize
1455
- if len(data) - offset != expected:
1456
  raise ValueError("Legacy compact embedding payload length does not match shape.")
1457
- array = ( # shape
1458
- np.frombuffer(data, dtype=numpy_dtype, offset=offset).copy().reshape(shape)
1459
  )
1460
- return torch.from_numpy(array).to(dtype=target_dtype) # shape
1461
 
1462
  try:
1463
- loaded = torch.load(io.BytesIO(data), map_location="cpu", weights_only=True)
1464
  except Exception as safe_error:
1465
  if allow_unsafe_pickle:
1466
- loaded = torch.load(io.BytesIO(data), map_location="cpu", weights_only=False)
1467
  elif fallback_shape is None:
1468
  raise ValueError(
1469
  "Legacy embedding blob is neither compact nor safely loadable. "
@@ -1472,14 +1472,14 @@ def _decode_legacy_sqlite_blob(
1472
  ) from safe_error
1473
  else:
1474
  expected = int(np.prod(fallback_shape, dtype=np.int64)) * 4
1475
- if len(data) != expected:
1476
  raise ValueError(
1477
  "Legacy raw FP32 payload length does not match fallback_shape."
1478
  ) from safe_error
1479
- array = np.frombuffer(data, dtype=np.float32).copy().reshape( # fallback_shape
1480
  fallback_shape
1481
  )
1482
- return torch.from_numpy(array) # fallback_shape
1483
  if not isinstance(loaded, Tensor):
1484
  raise ValueError("Legacy serialized embedding payload must contain one tensor.")
1485
  return loaded.detach().cpu() # (...)
@@ -1522,13 +1522,13 @@ def convert_legacy_sqlite(
1522
 
1523
  records: list[EmbeddingRecord] = []
1524
  content_digest = hashlib.sha256()
1525
- for position, (sequence, data) in enumerate(rows):
1526
  if not isinstance(sequence, str) or not sequence:
1527
  raise ValueError("Legacy embedding sequences must be non-empty strings.")
1528
- if not isinstance(data, bytes):
1529
- data = bytes(data)
1530
  tensor = _decode_legacy_sqlite_blob(
1531
- data,
1532
  fallback_shape=fallback_shape,
1533
  allow_unsafe_pickle=allow_unsafe_pickle,
1534
  )
@@ -1559,16 +1559,16 @@ def convert_legacy_sqlite(
1559
 
1560
 
1561
  def save_result(
1562
- result: EmbeddingResult,
1563
  path: str | Path,
1564
  *,
1565
  format: str = "safetensors",
1566
  shard_size: int = DEFAULT_SHARD_SIZE,
1567
  ) -> EmbeddingResult:
1568
  if format == "safetensors":
1569
- return save_safetensors_result(result, path, shard_size=shard_size)
1570
  if format == "sqlite":
1571
- return save_sqlite_result(result, path)
1572
  if format == "pth":
1573
  raise ValueError("Writing pickle-based .pth embeddings is not supported.")
1574
  raise ValueError("format must be 'safetensors' or 'sqlite'.")
 
22
  EmbeddingResult,
23
  LazyTensorReference,
24
  )
25
+ from ..json_files import indented_json
26
 
27
 
28
  _DTYPE_NAMES: dict[torch.dtype, str] = {
 
103
 
104
 
105
  def _tensor_hash_chunks(X: Tensor) -> Iterator[bytes]:
106
+ # X: (...)
107
+ for chunk in _bounded_tensor_chunks(X, _TENSOR_HASH_CHUNK_BYTES): # (n_chunk,)
108
  yield chunk.view(torch.uint8).numpy().tobytes()
109
 
110
 
111
  def tensor_sha256(X: Tensor) -> str:
112
  """Hash dtype, shape, and exact tensor bytes."""
113
 
114
+ # X: (...)
115
  if not isinstance(X, Tensor):
116
  raise TypeError("X must be a tensor.")
117
  if X.dtype not in _DTYPE_NAMES:
 
129
 
130
 
131
  def _encode_tensor(X: Tensor) -> tuple[str, str, bytes]:
132
+ # X: (...)
133
  if X.dtype not in _DTYPE_NAMES:
134
  raise TypeError(f"Unsupported tensor dtype {X.dtype}.")
135
  shape = json.dumps(tuple(X.shape), separators=(",", ":"))
136
  return _DTYPE_NAMES[X.dtype], shape, _tensor_bytes(X)
137
 
138
 
139
+ def _decode_tensor(dtype_name: str, shape_json: str, raw_bytes: bytes) -> Tensor:
140
  try:
141
  dtype = _NAME_DTYPES[dtype_name]
142
  except KeyError as error:
143
  raise ValueError(f"Unsupported stored dtype {dtype_name!r}.") from error
144
  shape = tuple(json.loads(shape_json))
145
  # uint8 is used only as a byte-level carrier, preserving BF16 bits exactly.
146
+ byte_array = np.frombuffer(raw_bytes, dtype=np.uint8).copy() # (n_bytes,)
147
  X = torch.from_numpy(byte_array).view(dtype) # (n_elements,)
148
+ return X.reshape(shape).clone() # (...), the stored shape
149
 
150
 
151
  def _index_path(path: str | Path) -> Path:
 
176
  return candidate
177
 
178
 
 
 
 
 
179
  def _load_authoritative_index(
180
  path: str | Path,
181
  ) -> tuple[dict[str, Any], Path, dict[str, Any]]:
 
198
  snapshot = run_manifest.get("index_payload")
199
  if isinstance(snapshot, dict):
200
  payload = snapshot
201
+ index_bytes = indented_json(payload).encode("utf-8")
202
  elif snapshot is None:
203
  index_bytes = stable_index_path.read_bytes()
204
  payload = json.loads(index_bytes.decode("utf-8"))
 
268
  except ImportError as error:
269
  raise ImportError("Loading embeddings requires the 'safetensors' package.") from error
270
  with safe_open(path, framework="pt", device="cpu") as handle:
271
+ return cast(Tensor, handle.get_tensor(key)) # (...), as stored under key
272
 
273
 
274
  def _safetensors_shard_prefix(path: str | Path) -> str:
 
361
  raise ValueError(f"Safetensors tensor shard is missing: {relative}.")
362
 
363
  def load_tensor() -> Tensor:
364
+ return _load_safetensor(tensor_path, key) # (...), the descriptor's shape
365
 
366
  reference = LazyTensorReference(
367
  source=str(tensor_path),
 
720
  raise FileExistsError(
721
  f"Refusing to reuse immutable safetensors generation index {generation_index_path}."
722
  )
723
+ encoded_index = indented_json(payload).encode("utf-8")
724
  temporary_generation_index.write_bytes(encoded_index)
725
  temporary_generation_index.replace(generation_index_path)
726
 
 
739
  temporary_manifest = self.run_manifest_path.with_name(
740
  f".{self.run_manifest_path.name}.{pointer_identity}.tmp"
741
  )
742
+ temporary_manifest.write_bytes(indented_json(run_manifest).encode("utf-8"))
743
  temporary_manifest.replace(self.run_manifest_path)
744
 
745
  # ``index.json`` is a non-authoritative convenience pointer. The run
 
753
  temporary_index = self.index_path.with_name(
754
  f".{self.index_path.name}.{pointer_identity}.tmp"
755
  )
756
+ temporary_index.write_bytes(indented_json(stable_pointer).encode("utf-8"))
757
  temporary_index.replace(self.index_path)
758
 
759
  return load_safetensors_result(self.index_path)
 
771
 
772
 
773
  def save_safetensors_result(
774
+ embedding_result: EmbeddingResult,
775
  path: str | Path,
776
  *,
777
  shard_size: int = DEFAULT_SHARD_SIZE,
 
780
 
781
  writer = SafetensorsStreamWriter(
782
  path,
783
+ embedding_result.metadata,
784
  shard_size=shard_size,
785
  publish_initial=False,
786
  publish_incremental=False,
787
  )
788
+ writer.append(embedding_result, publish=False)
789
+ return writer.publish(complete=bool(embedding_result.metadata.get("complete", True)))
790
 
791
 
792
  def load_safetensors_result(path: str | Path) -> EmbeddingResult:
 
925
  connection.commit()
926
 
927
 
928
+ def save_sqlite_result(embedding_result: EmbeddingResult, path: str | Path) -> EmbeddingResult:
929
  """Transactionally store an ordered result in normalized SQLite tables."""
930
 
931
  path = Path(path)
932
  path.parent.mkdir(parents=True, exist_ok=True)
933
+ run_id = str(embedding_result.metadata.get("run_fingerprint", ""))
934
  if not run_id:
935
  raise ValueError("SQLite results require metadata['run_fingerprint'].")
936
  metadata_json = json.dumps(
937
  _persistent_metadata(
938
+ embedding_result.metadata,
939
  descriptor_index="sqlite-records",
940
+ record_count=len(embedding_result),
941
  ),
942
  sort_keys=True,
943
  )
 
951
  "SELECT ?, ?, COALESCE(MAX(published_order), 0) + 1 FROM runs",
952
  (run_id, metadata_json),
953
  )
954
+ for position, record in enumerate(embedding_result):
955
  X = record.load_tensor().detach().cpu().contiguous() # (...)
956
+ dtype_name, shape_json, raw_bytes = _encode_tensor(X)
957
  digest = tensor_sha256(X)
958
  connection.execute(
959
  "INSERT INTO tensors VALUES (?, ?, ?, ?, ?, ?)",
960
+ (run_id, position, dtype_name, shape_json, raw_bytes, digest),
961
  )
962
  connection.execute(
963
  "INSERT INTO records VALUES (?, ?, ?, ?)",
 
1057
  for offset, record in enumerate(records):
1058
  position = start_position + offset
1059
  X = record.load_tensor().detach().cpu().contiguous() # (...)
1060
+ dtype_name, shape_json, raw_bytes = _encode_tensor(X)
1061
  digest = tensor_sha256(X)
1062
  connection.execute(
1063
  "INSERT INTO tensors VALUES (?, ?, ?, ?, ?, ?)",
1064
+ (run_id, position, dtype_name, shape_json, raw_bytes, digest),
1065
  )
1066
  connection.execute(
1067
  "INSERT INTO records VALUES (?, ?, ?, ?)",
 
1142
  ).fetchone()
1143
  if row is None:
1144
  raise KeyError(f"Missing SQLite tensor {run_id}:{position}.")
1145
+ return _decode_tensor(*row) # (...), the stored shape
1146
 
1147
 
1148
  def _validate_sqlite_descriptor_row(
 
1180
  )
1181
 
1182
  def load_tensor() -> Tensor:
1183
+ return _load_sqlite_tensor(path, run_id, position) # (...), the stored shape
1184
 
1185
  reference = LazyTensorReference(
1186
  source=str(path),
 
1284
  _validate_sqlite_result_schema(connection, path)
1285
  if run_id is None:
1286
  run_columns = {
1287
+ str(column[1]) for column in connection.execute("PRAGMA table_info(runs)").fetchall()
1288
  }
1289
  if "published_order" in run_columns:
1290
  row = connection.execute(
 
1434
 
1435
 
1436
  def _decode_legacy_sqlite_blob(
1437
+ blob: bytes,
1438
  *,
1439
  fallback_shape: tuple[int, ...] | None,
1440
  allow_unsafe_pickle: bool,
1441
  ) -> Tensor:
1442
+ if len(blob) >= 6 and blob[0] == _LEGACY_COMPACT_VERSION:
1443
+ dtype_code = int(blob[1])
1444
  if dtype_code not in _LEGACY_CODE_DTYPES:
1445
  raise ValueError(f"Unsupported legacy compact dtype code {dtype_code}.")
1446
+ (ndim,) = struct.unpack_from("<i", blob, 2)
1447
+ if ndim < 0 or ndim > 16 or len(blob) < 6 + 4 * ndim:
1448
  raise ValueError("Malformed legacy compact embedding header.")
1449
+ shape = tuple(int(value) for value in struct.unpack_from(f"<{ndim}i", blob, 6))
1450
  if any(size < 0 for size in shape):
1451
  raise ValueError("Malformed negative legacy embedding dimension.")
1452
  numpy_dtype, target_dtype = _LEGACY_CODE_DTYPES[dtype_code]
1453
  offset = 6 + 4 * ndim
1454
  expected = int(np.prod(shape, dtype=np.int64)) * numpy_dtype.itemsize
1455
+ if len(blob) - offset != expected:
1456
  raise ValueError("Legacy compact embedding payload length does not match shape.")
1457
+ array = ( # (...), the stored shape
1458
+ np.frombuffer(blob, dtype=numpy_dtype, offset=offset).copy().reshape(shape)
1459
  )
1460
+ return torch.from_numpy(array).to(dtype=target_dtype) # (...), the stored shape
1461
 
1462
  try:
1463
+ loaded = torch.load(io.BytesIO(blob), map_location="cpu", weights_only=True)
1464
  except Exception as safe_error:
1465
  if allow_unsafe_pickle:
1466
+ loaded = torch.load(io.BytesIO(blob), map_location="cpu", weights_only=False)
1467
  elif fallback_shape is None:
1468
  raise ValueError(
1469
  "Legacy embedding blob is neither compact nor safely loadable. "
 
1472
  ) from safe_error
1473
  else:
1474
  expected = int(np.prod(fallback_shape, dtype=np.int64)) * 4
1475
+ if len(blob) != expected:
1476
  raise ValueError(
1477
  "Legacy raw FP32 payload length does not match fallback_shape."
1478
  ) from safe_error
1479
+ array = np.frombuffer(blob, dtype=np.float32).copy().reshape( # (...), fallback_shape
1480
  fallback_shape
1481
  )
1482
+ return torch.from_numpy(array) # (...), fallback_shape
1483
  if not isinstance(loaded, Tensor):
1484
  raise ValueError("Legacy serialized embedding payload must contain one tensor.")
1485
  return loaded.detach().cpu() # (...)
 
1522
 
1523
  records: list[EmbeddingRecord] = []
1524
  content_digest = hashlib.sha256()
1525
+ for position, (sequence, blob) in enumerate(rows):
1526
  if not isinstance(sequence, str) or not sequence:
1527
  raise ValueError("Legacy embedding sequences must be non-empty strings.")
1528
+ if not isinstance(blob, bytes):
1529
+ blob = bytes(blob)
1530
  tensor = _decode_legacy_sqlite_blob(
1531
+ blob,
1532
  fallback_shape=fallback_shape,
1533
  allow_unsafe_pickle=allow_unsafe_pickle,
1534
  )
 
1559
 
1560
 
1561
  def save_result(
1562
+ embedding_result: EmbeddingResult,
1563
  path: str | Path,
1564
  *,
1565
  format: str = "safetensors",
1566
  shard_size: int = DEFAULT_SHARD_SIZE,
1567
  ) -> EmbeddingResult:
1568
  if format == "safetensors":
1569
+ return save_safetensors_result(embedding_result, path, shard_size=shard_size)
1570
  if format == "sqlite":
1571
+ return save_sqlite_result(embedding_result, path)
1572
  if format == "pth":
1573
  raise ValueError("Writing pickle-based .pth embeddings is not supported.")
1574
  raise ValueError("format must be 'safetensors' or 'sqlite'.")
fastplms/embeddings/taps.py ADDED
@@ -0,0 +1,413 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Tap plans: several hidden-state outputs from one forward pass per batch.
2
+
3
+ A tap names one hidden state and what to keep from it. ``HiddenTap`` keeps the rows the mask selects
4
+ (biological residues in a residue run; CLS, residues and EOS in a token run) or pools them
5
+ (``Pooler`` in a residue run, ``pool_token_rows`` in a token run). ``ReducedTap`` hands the state to a
6
+ caller's reducer, such as a sparse-autoencoder encoder and its pooling. ``StreamingTap`` reduces layers
7
+ as they arrive without saving their full hidden states. A plan runs one forward pass per
8
+ batch that stops once the deepest tapped state exists.
9
+
10
+ Layer indices follow the FastPLMs hidden-state order: index ``i`` is the input to block ``i``,
11
+ index ``n`` (the block count) is the final normalized state, and negative indices count back
12
+ from it, so ``-1`` is the final state.
13
+
14
+ Symbols: b sequences of a batch; l token columns of the padded batch (CLS, residues, EOS, padding); d hidden
15
+ width; n attended rows of a batch; r residues of one sequence. In a residue run the mask is false on CLS, EOS and
16
+ padding, so a sequence is r rows of an ``(n, d)`` output. In a token run (canonical) it is false on padding only, so a
17
+ sequence is r + 2 rows (row 0 CLS, rows 1..r residues, row r + 1 EOS) in every ``(n, d)`` output and every pooling.
18
+ """
19
+
20
+ from __future__ import annotations
21
+
22
+ import json
23
+ import math
24
+ import torch
25
+
26
+ from collections.abc import Callable, Mapping, Sequence
27
+ from dataclasses import dataclass, field
28
+ from typing import Any, Protocol
29
+ from torch import Tensor
30
+
31
+ from types import MappingProxyType
32
+ from .pooling import Pooler
33
+ from ..features.layouts import TopKRow
34
+
35
+
36
+ @dataclass(frozen=True, slots=True)
37
+ class RowSelection:
38
+ """The rows of a ``(b, l)`` grid that a tap keeps, known on the host so no device value is read.
39
+
40
+ ``flat_index`` holds each kept row's position in the row-major ``(b * l)`` grid, ``owner`` the
41
+ sequence it belongs to, and ``counts`` how many rows each sequence keeps. Gathering with an index
42
+ built from host-known lengths never stalls the host the way boolean indexing does.
43
+ """
44
+
45
+ flat_index: Tensor # (n,) int64 on the device, n = sum of counts
46
+ owner: Tensor # (n,) int64 on the device, the sequence of each kept row
47
+ counts: tuple[int, ...] # (b,) kept rows per sequence, on the host
48
+ sizes: Tensor # (b,) int64 on the device, the same counts
49
+
50
+ def gather(self, X: Tensor) -> Tensor:
51
+ """The kept rows of ``X`` in sequence order, without padding."""
52
+ # X: (b, l, w) -> (n, w)
53
+ return X.reshape(-1, X.shape[-1]).index_select(0, self.flat_index) # (n, w)
54
+
55
+
56
+ @dataclass(frozen=True, slots=True)
57
+ class TapBatch:
58
+ """One batch of a tapped hidden state, as a ``ReducedTap`` reducer receives it.
59
+
60
+ ``X`` has shape ``(b, l, d)``. ``token_mask`` has shape ``(b, l)`` and marks every attended
61
+ token, BOS and EOS included. ``residue_mask`` has shape ``(b, l)`` and marks the rows that
62
+ ``HiddenTap`` keeps and pools: the biological residues, or, when the run keeps the special
63
+ tokens, every attended token (``l`` then counts CLS and EOS). All three share X's device.
64
+ ``rows`` is the host-known selection of ``residue_mask`` when the executor has one, and
65
+ ``cache`` lets taps of one batch share an intermediate such as an SAE encoding.
66
+ """
67
+
68
+ X: Tensor
69
+ token_mask: Tensor
70
+ residue_mask: Tensor
71
+ rows: RowSelection | None = None
72
+ cache: dict[Any, Any] = field(default_factory=dict)
73
+
74
+ def selection(self) -> RowSelection:
75
+ """The kept rows, taken from the executor when it supplied them, else read from the mask."""
76
+ if self.rows is not None:
77
+ return self.rows
78
+ kept = self.residue_mask.bool() # (b, l)
79
+ owner, position = kept.nonzero(as_tuple=True) # (n,), (n,)
80
+ sizes = kept.sum(dim=1) # (b,)
81
+ return RowSelection(owner * kept.shape[1] + position, owner, tuple(sizes.tolist()), sizes)
82
+
83
+
84
+ @dataclass(frozen=True, slots=True)
85
+ class HiddenTap:
86
+ """One hidden state, kept as ragged rows or pooled.
87
+
88
+ ``pooling=None`` keeps one ``(r_i, d)`` tensor of mask-selected rows per sequence: ``r_i`` biological
89
+ residues in a residue run, ``l_i + 2`` rows (CLS, residues, EOS) in a token run. Pooler
90
+ names give one pooled vector per sequence over those same rows, concatenated in request order. ``parti`` needs the
91
+ attention graph of a full pass, so a tap rejects it. ``dtype`` overrides the run's dtype
92
+ for this output, starting from the original captured state; None inherits the run dtype.
93
+ """
94
+
95
+ name: str
96
+ layer: int
97
+ pooling: str | Sequence[str] | None = None
98
+ dtype: torch.dtype | None = None
99
+
100
+ def __post_init__(self) -> None:
101
+ _require_name_and_layer(self.name, self.layer)
102
+ if self.dtype is not None and self.dtype not in (
103
+ torch.float16, torch.bfloat16, torch.float32, torch.float64
104
+ ):
105
+ raise ValueError(
106
+ "A hidden tap dtype must be float16, bfloat16, float32, float64, "
107
+ "or None to inherit the run dtype."
108
+ )
109
+ if self.pooling is None:
110
+ return
111
+ names = Pooler(self.pooling).names # validates names and rejects duplicates
112
+ if "parti" in names:
113
+ raise ValueError(
114
+ f"Tap {self.name!r} cannot pool with 'parti', which needs the attention graph "
115
+ "of a full forward pass."
116
+ )
117
+ object.__setattr__(self, "pooling", names)
118
+
119
+
120
+ @dataclass(frozen=True, slots=True)
121
+ class ReducedTap:
122
+ """One hidden state reduced by a caller's function.
123
+
124
+ ``reduce`` maps a ``TapBatch`` to a tensor with one row per sequence, shape ``(b, ...)``.
125
+ ``identity`` describes the reducer in plain data: strings, numbers, booleans, None, lists,
126
+ and string-keyed mappings. The run fingerprint records it, so two reducers with equal
127
+ identities must return equal outputs.
128
+ """
129
+
130
+ name: str
131
+ layer: int
132
+ reduce: Callable[[TapBatch], Tensor]
133
+ identity: Mapping[str, Any]
134
+
135
+ def __post_init__(self) -> None:
136
+ _require_name_and_layer(self.name, self.layer)
137
+ if not callable(self.reduce):
138
+ raise TypeError(f"Tap {self.name!r} needs a callable reduce.")
139
+ if not isinstance(self.identity, Mapping) or not self.identity:
140
+ raise TypeError(f"Tap {self.name!r} needs a non-empty identity mapping.")
141
+ _require_plain_data(self.identity, f"identity of tap {self.name!r}")
142
+ # A private canonical copy, so later changes to the caller's mapping cannot change the
143
+ # fingerprint of this tap.
144
+ canonical = json.loads(json.dumps(dict(self.identity), sort_keys=True, allow_nan=False))
145
+ object.__setattr__(self, "identity", MappingProxyType(canonical))
146
+
147
+
148
+ @dataclass(frozen=True, slots=True)
149
+ class SparseResidueTap:
150
+ """Reduce one state to sparse codes per kept row, in sequence and row order.
151
+
152
+ The reducer receives the same masks as dense taps and returns one ``TopKRow`` per sequence.
153
+ Its output must retain every row the mask keeps: the biological residues in a residue run, all
154
+ ``l_i + 2`` attended tokens in a token run. Both integer indices and floating
155
+ values remain sparse through extraction and persistence.
156
+ """
157
+
158
+ name: str
159
+ layer: int
160
+ reduce: Callable[[TapBatch], Sequence[TopKRow]]
161
+ identity: Mapping[str, Any]
162
+ codebook_size: int
163
+ sparse_count: int
164
+ # A run that keeps CLS and EOS reads every sequence's codes as one packed (n, k) pair, n = sum(l_i + 2),
165
+ # and splits nothing: the token executor calls this instead of ``reduce``. Not part of the identity.
166
+ reduce_packed: Callable[[TapBatch], TopKRow] | None = None
167
+
168
+ def __post_init__(self) -> None:
169
+ checked = ReducedTap(self.name, self.layer, self.reduce, self.identity)
170
+ object.__setattr__(self, "identity", checked.identity)
171
+ if (type(self.codebook_size) is not int or not 1 <= self.codebook_size <= 2**31
172
+ or type(self.sparse_count) is not int
173
+ or not 1 <= self.sparse_count <= self.codebook_size):
174
+ raise ValueError(
175
+ "Sparse residue taps require integer 1 <= sparse_count <= codebook_size <= 2**31."
176
+ )
177
+
178
+
179
+ class LayerAccumulator(Protocol):
180
+ """Batch-local state for a streaming reduction. Never mutate the borrowed hidden state."""
181
+
182
+ def update(self, layer: int, batch: TapBatch) -> None: ...
183
+
184
+ def finish(self) -> Tensor:
185
+ """Return a token-aligned tensor of shape (b, l, c)."""
186
+ ...
187
+
188
+
189
+ @dataclass(frozen=True, slots=True)
190
+ class StreamingTap:
191
+ """Reduce selected layers as they arrive, retaining only the accumulator's own state.
192
+
193
+ ``begin`` creates a fresh accumulator for each batch. ``update`` borrows each original
194
+ hidden state, before the run's output dtype conversion, in ascending layer order. ``finish``
195
+ returns token-aligned residue features; the engine applies its biological mask and restores
196
+ input order. Reducers own their arithmetic and describe it in ``identity``. They must not
197
+ retain or mutate borrowed states. No callback is installed on the model between calls.
198
+
199
+ ``pooling`` (token runs only) pools the finished rows of each sequence over every attended token, as a
200
+ pooled ``HiddenTap`` does, so one value per sequence is kept instead of one per token; ``dtype`` converts
201
+ the finished rows first (float32 for float32 moments of a 16-bit reducer).
202
+ """
203
+
204
+ name: str
205
+ layers: tuple[int, ...]
206
+ begin: Callable[[], LayerAccumulator]
207
+ identity: Mapping[str, Any]
208
+ required_state_count: int | None = None
209
+ pooling: str | Sequence[str] | None = None
210
+ dtype: torch.dtype | None = None
211
+
212
+ def __post_init__(self) -> None:
213
+ if self.dtype is not None and self.dtype not in (torch.float16, torch.bfloat16, torch.float32, torch.float64):
214
+ raise ValueError("A streaming tap dtype must be float16, bfloat16, float32, float64, or None.")
215
+ if self.pooling is not None:
216
+ names = Pooler(self.pooling).names # validates names and rejects duplicates
217
+ if "parti" in names:
218
+ raise ValueError(f"Tap {self.name!r} cannot pool with 'parti', which needs the attention graph.")
219
+ object.__setattr__(self, "pooling", names)
220
+ layers = tuple(self.layers)
221
+ if not layers or any(type(layer) is not int or layer < 0 for layer in layers):
222
+ raise ValueError("Streaming layers must be nonempty nonnegative integer indices.")
223
+ if tuple(sorted(set(layers))) != layers:
224
+ raise ValueError("Streaming layers must be distinct and ascending.")
225
+ object.__setattr__(self, "layers", layers)
226
+ if self.required_state_count is not None and (
227
+ type(self.required_state_count) is not int or self.required_state_count <= layers[-1]
228
+ ):
229
+ raise ValueError("required_state_count must include every streamed layer.")
230
+ # Reuse the reducer identity validation and detached canonical copy.
231
+ checked = ReducedTap(self.name, layers[-1], self.begin, self.identity)
232
+ object.__setattr__(self, "identity", checked.identity)
233
+
234
+ @property
235
+ def layer(self) -> int:
236
+ """The deepest required state, for the existing early-stop plan."""
237
+ return self.layers[-1]
238
+
239
+
240
+ Tap = HiddenTap | ReducedTap | StreamingTap | SparseResidueTap
241
+
242
+
243
+ @dataclass(frozen=True, slots=True)
244
+ class TapPlan:
245
+ """Validated taps, each layer resolved to a hidden-state index in ``0..n``."""
246
+
247
+ taps: tuple[Tap, ...]
248
+ layers: tuple[int, ...]
249
+
250
+ @property
251
+ def captured_layers(self) -> tuple[int, ...]:
252
+ """The distinct hidden states the forward pass must record, in ascending order."""
253
+
254
+ return tuple(sorted({
255
+ layer for tap, layer in zip(self.taps, self.layers, strict=True)
256
+ if not isinstance(tap, StreamingTap)
257
+ }))
258
+
259
+ @property
260
+ def streamed_layers(self) -> tuple[int, ...]:
261
+ return tuple(sorted({
262
+ layer for tap in self.taps if isinstance(tap, StreamingTap) for layer in tap.layers
263
+ }))
264
+
265
+ @property
266
+ def deepest_layer(self) -> int:
267
+ """The hidden state after which the forward pass stops."""
268
+
269
+ return max(self.layers)
270
+
271
+ @property
272
+ def pooling_names(self) -> frozenset[str]:
273
+ """Every pooler name the plan's hidden taps request."""
274
+
275
+ return frozenset(
276
+ name
277
+ for tap in self.taps
278
+ if isinstance(tap, HiddenTap) and tap.pooling is not None
279
+ for name in tap.pooling
280
+ )
281
+
282
+ def identity(self) -> list[dict[str, Any]]:
283
+ """Every tap in plan order, as the run fingerprint and metadata record it."""
284
+
285
+ described: list[dict[str, Any]] = []
286
+ for tap, layer in zip(self.taps, self.layers, strict=True):
287
+ if isinstance(tap, HiddenTap):
288
+ pooling = None if tap.pooling is None else list(tap.pooling)
289
+ described.append(
290
+ {
291
+ "name": tap.name,
292
+ "kind": "hidden",
293
+ "layer": layer,
294
+ "pooling": pooling,
295
+ "dtype": (
296
+ str(tap.dtype).removeprefix("torch.") if tap.dtype is not None else None
297
+ ),
298
+ }
299
+ )
300
+ elif isinstance(tap, StreamingTap):
301
+ record = {
302
+ "name": tap.name, "kind": "streaming", "layers": list(tap.layers),
303
+ "identity": dict(tap.identity),
304
+ "required_state_count": tap.required_state_count,
305
+ }
306
+ if tap.pooling is not None: # absent for a per-token tap, so its existing fingerprint holds
307
+ record.update(pooling=list(tap.pooling),
308
+ dtype=None if tap.dtype is None else str(tap.dtype).removeprefix("torch."))
309
+ described.append(record)
310
+ elif isinstance(tap, SparseResidueTap):
311
+ described.append({
312
+ "name": tap.name, "kind": "sparse_residue", "layer": layer,
313
+ "identity": dict(tap.identity), "codebook_size": tap.codebook_size,
314
+ "sparse_count": tap.sparse_count,
315
+ })
316
+ else:
317
+ described.append(
318
+ {
319
+ "name": tap.name,
320
+ "kind": "reduced",
321
+ "layer": layer,
322
+ "identity": dict(tap.identity),
323
+ }
324
+ )
325
+ return described
326
+
327
+
328
+ def plan_taps(taps: object, state_count: int) -> TapPlan:
329
+ """Validate ``taps`` against a model that exposes ``state_count`` hidden states.
330
+
331
+ ``taps`` is the caller's ``embed_dataset`` argument, checked here rather than trusted.
332
+ """
333
+
334
+ if isinstance(taps, (str, bytes)) or not isinstance(taps, Sequence):
335
+ raise TypeError(
336
+ "taps must be a sequence of HiddenTap, ReducedTap, StreamingTap "
337
+ "or SparseResidueTap values."
338
+ )
339
+ if not taps:
340
+ raise ValueError("taps must contain at least one tap.")
341
+ checked: list[Tap] = []
342
+ for tap in taps:
343
+ if not isinstance(tap, (HiddenTap, ReducedTap, StreamingTap, SparseResidueTap)):
344
+ raise TypeError(
345
+ "taps must contain HiddenTap, ReducedTap, StreamingTap or SparseResidueTap values; "
346
+ f"found {type(tap).__name__}."
347
+ )
348
+ checked.append(tap)
349
+ names = [tap.name for tap in checked]
350
+ repeated = sorted({name for name in names if names.count(name) > 1})
351
+ if repeated:
352
+ raise ValueError(f"Tap names must be unique; repeated: {repeated}.")
353
+ layers: list[int] = []
354
+ for tap in checked:
355
+ if isinstance(tap, StreamingTap) and tap.required_state_count not in (None, state_count):
356
+ raise ValueError(
357
+ f"Tap {tap.name!r} requires {tap.required_state_count} hidden states, "
358
+ f"not {state_count}."
359
+ )
360
+ if not -state_count <= tap.layer < state_count:
361
+ raise ValueError(
362
+ f"Tap {tap.name!r} names layer {tap.layer}, outside this model's hidden states "
363
+ f"{-state_count}..{state_count - 1}. Index i is the input to block i, and "
364
+ f"{state_count - 1} or -1 is the final normalized state."
365
+ )
366
+ layers.append(tap.layer % state_count)
367
+ return TapPlan(taps=tuple(checked), layers=tuple(layers))
368
+
369
+
370
+ def _require_name_and_layer(name: object, layer: object) -> None:
371
+ if not isinstance(name, str) or not name:
372
+ raise ValueError("A tap name must be a non-empty string.")
373
+ if not isinstance(layer, int) or isinstance(layer, bool):
374
+ raise TypeError(f"Tap {name!r} layer must be an integer hidden-state index.")
375
+
376
+
377
+ def _require_plain_data(value: object, where: str) -> None:
378
+ """Reject content whose serialized form could differ between runs of one reducer."""
379
+
380
+ if value is None or isinstance(value, (str, bool, int)):
381
+ return
382
+ if isinstance(value, float):
383
+ if not math.isfinite(value):
384
+ raise ValueError(f"The {where} holds a non-finite number.")
385
+ return
386
+ if isinstance(value, Mapping):
387
+ for key, item in value.items():
388
+ if not isinstance(key, str):
389
+ raise TypeError(f"The {where} has a non-string key {key!r}.")
390
+ _require_plain_data(item, where)
391
+ return
392
+ if isinstance(value, (list, tuple)):
393
+ for item in value:
394
+ _require_plain_data(item, where)
395
+ return
396
+ raise TypeError(
397
+ f"The {where} holds a {type(value).__name__}. An identity holds only strings, numbers, "
398
+ "booleans, None, lists, and string-keyed mappings, so its fingerprint is stable."
399
+ )
400
+
401
+
402
+ __all__ = [
403
+ "HiddenTap",
404
+ "LayerAccumulator",
405
+ "ReducedTap",
406
+ "RowSelection",
407
+ "SparseResidueTap",
408
+ "StreamingTap",
409
+ "Tap",
410
+ "TapBatch",
411
+ "TapPlan",
412
+ "plan_taps",
413
+ ]
fastplms/embeddings/token_batches.py ADDED
@@ -0,0 +1,391 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Embed canonical proteins with CLS and EOS kept: one pass per token-budget batch, no host stall, pinned outputs.
2
+
3
+ A canonical run keeps the special tokens. A protein of l residues, after the N-terminal crop, is l + 2 rows in
4
+ every per-token stream: row 0 CLS, rows 1..l residues, row l + 1 EOS. Padding is the only masked position, so a
5
+ pooled vector averages all l + 2 rows. A legacy residue-only store has l rows (b, l, d) instead; the two never
6
+ mix, because a v2 descriptor says ``special_tokens: kept``.
7
+
8
+ The executor never reads a device value while it builds or runs a batch. Lengths are known on the host, so the
9
+ attention mask, the row selection, and every offset come from host arrays, and the finite check is one device
10
+ flag read after the outputs land. That keeps the device queue full while the previous batch is written.
11
+
12
+ A run with a ``BatchGeometry`` gives every sequence one batch shape, whatever its companions: its l + 2 tokens round
13
+ up to a bucket of T columns, and every batch of that bucket holds exactly ``rows(T)`` sequences padded to T. A
14
+ GEMM's reduction order follows its shape, so a fixed shape per sequence is what makes a stored row independent of
15
+ the batch that made it. Without a geometry, batches follow the token budget and pad to their longest member.
16
+
17
+ Symbols: b sequences of a batch; l residues of one sequence after the crop; m = max(l) + 2 padded token columns,
18
+ or the bucket T of a geometry batch; n = sum(l_i + 2) attended token rows; d hidden width; c SAE codebook; k SAE
19
+ codes kept per token; w stored columns of a stream.
20
+ """
21
+
22
+ from __future__ import annotations
23
+
24
+ import numpy as np
25
+ import torch
26
+
27
+ from collections.abc import Iterator, Mapping, Sequence
28
+ from dataclasses import dataclass
29
+ from typing import Any
30
+ from torch import Tensor
31
+
32
+ from .pooling import pool_token_rows
33
+ from .taps import HiddenTap, ReducedTap, RowSelection, SparseResidueTap, StreamingTap, TapBatch, TapPlan
34
+ from .tokens import ResidueVocabulary
35
+ from ..features.async_writer import PackedBatch
36
+ from ..features.layouts import INDEX_DTYPE
37
+
38
+
39
+ # CLS before the residues and EOS after them.
40
+ SPECIAL_TOKEN_ROWS = 2
41
+ # The N-terminal crop in residues: with CLS and EOS it fills a 2,048-token context. The one place the engine names it.
42
+ CANONICAL_MAX_RESIDUES = 2046
43
+ # A geometry run's batch algorithm and its rule for a bucket's last, partial batch.
44
+ GEOMETRY_ALGORITHM = "bucketed_fixed_shape_v1"
45
+ PARTIAL_BATCH_POLICY = "repeat_last_sequence_discard_duplicate_outputs_v1"
46
+
47
+
48
+ @dataclass(frozen=True, slots=True)
49
+ class BatchGeometry:
50
+ """Fixed batch shapes, so a sequence meets the same kernels in whatever batch it runs.
51
+
52
+ A sequence of l residues after the crop needs l + 2 token columns and runs in the bucket
53
+ T = ``bucket_tokens`` * ceil((l + 2) / ``bucket_tokens``), at most ``max_columns``. Every batch of bucket T holds
54
+ exactly ``rows(T)`` sequences, as many as ``token_budget`` padded token rows hold and at most ``max_rows``, each
55
+ padded to T columns; a bucket's last batch repeats its last sequence to fill and discards the copies' outputs.
56
+ """
57
+
58
+ bucket_tokens: int
59
+ token_budget: int # padded token rows one batch may hold
60
+ max_rows: int # sequences one batch may hold
61
+ max_columns: int = CANONICAL_MAX_RESIDUES + SPECIAL_TOKEN_ROWS
62
+
63
+ def __post_init__(self) -> None:
64
+ for name in ("bucket_tokens", "token_budget", "max_rows", "max_columns"):
65
+ value = getattr(self, name)
66
+ if type(value) is not int or value < 1:
67
+ raise ValueError(f"BatchGeometry.{name} must be a positive integer; received {value!r}.")
68
+ if self.max_columns % self.bucket_tokens:
69
+ raise ValueError("BatchGeometry.max_columns must be a multiple of bucket_tokens.")
70
+ if self.token_budget < self.max_columns:
71
+ raise ValueError("BatchGeometry.token_budget must hold one sequence of max_columns tokens.")
72
+
73
+ def columns(self, residues: int) -> int:
74
+ """The bucket T of a sequence of ``residues`` after the crop: its l + 2 tokens rounded up to the bucket width."""
75
+ tokens = residues + SPECIAL_TOKEN_ROWS
76
+ if residues < 1 or tokens > self.max_columns:
77
+ raise ValueError(
78
+ f"A sequence of {residues} residues needs {tokens} token columns; this geometry holds 3 to {self.max_columns}."
79
+ )
80
+ return -(-tokens // self.bucket_tokens) * self.bucket_tokens
81
+
82
+ def rows(self, columns: int) -> int:
83
+ """Sequences in every batch of bucket ``columns``: as many as the token budget holds, at most ``max_rows``."""
84
+ if columns % self.bucket_tokens or not 0 < columns <= self.max_columns:
85
+ raise ValueError(f"{columns} is not a bucket of this geometry.")
86
+ return min(self.max_rows, self.token_budget // columns)
87
+
88
+ def shapes(self) -> dict[int, int]:
89
+ """Every bucket T with the row count rows(T) of its batches."""
90
+ return {
91
+ columns: self.rows(columns)
92
+ for columns in range(self.bucket_tokens, self.max_columns + 1, self.bucket_tokens)
93
+ }
94
+
95
+ def describe(self) -> dict[str, int | str]:
96
+ """The geometry as a run record and a feature contract state it."""
97
+ return {
98
+ "algorithm": GEOMETRY_ALGORITHM, "bucket_tokens": self.bucket_tokens, "token_budget": self.token_budget,
99
+ "max_rows": self.max_rows, "max_columns": self.max_columns, "partial_batch": PARTIAL_BATCH_POLICY,
100
+ }
101
+
102
+
103
+ def plan_geometry_batches(
104
+ lengths: Sequence[int], digests: Sequence[str], geometry: BatchGeometry,
105
+ ) -> Iterator[tuple[int, ...]]:
106
+ """Batches of one bucket each: the longest bucket first, members in row-key order, ``rows(T)`` to a batch.
107
+
108
+ ``lengths`` holds the residues l of each sequence after the crop and ``digests`` their row keys. A bucket's last
109
+ batch may hold fewer sequences; the executor fills it to ``rows(T)``. The longest bucket first makes an
110
+ out-of-memory failure show on the first batch, and key order makes the plan independent of the input's order.
111
+ Yields indices into ``lengths``.
112
+ """
113
+ if len(lengths) != len(digests):
114
+ raise ValueError("plan_geometry_batches needs one row key per length.")
115
+ buckets: dict[int, list[int]] = {}
116
+ for index, residues in enumerate(lengths):
117
+ buckets.setdefault(geometry.columns(residues), []).append(index)
118
+ for columns in sorted(buckets, reverse=True):
119
+ members = sorted(buckets[columns], key=lambda index: digests[index])
120
+ size = geometry.rows(columns)
121
+ for start in range(0, len(members), size):
122
+ yield tuple(members[start:start + size])
123
+
124
+
125
+ def plan_token_batches(
126
+ lengths: Sequence[int], *, max_sequences: int, max_tokens: int, window: int,
127
+ ) -> Iterator[tuple[int, ...]]:
128
+ """Group sequences into batches of similar length under a padded-token budget.
129
+
130
+ ``lengths`` holds the residues l of each sequence after the crop. Sequences are sorted longest first
131
+ inside each window of ``window`` consecutive sequences, so a batch pads to its first member: it holds
132
+ at most ``max_sequences`` sequences and ``b * (l_first + 2) <= max_tokens`` padded token rows. Longest
133
+ first also makes an out-of-memory failure show on the first batch. Yields indices into ``lengths``.
134
+ """
135
+ if min(max_sequences, max_tokens, window) < 1:
136
+ raise ValueError("max_sequences, max_tokens and window must be positive.")
137
+ for start in range(0, len(lengths), window):
138
+ order = sorted(range(start, min(start + window, len(lengths))), key=lambda index: (-lengths[index], index))
139
+ batch: list[int] = []
140
+ for index in order:
141
+ if lengths[index] + SPECIAL_TOKEN_ROWS > max_tokens:
142
+ raise ValueError(
143
+ f"A sequence of {lengths[index]} residues needs {lengths[index] + SPECIAL_TOKEN_ROWS} token rows, "
144
+ f"more than max_tokens={max_tokens}."
145
+ )
146
+ padded_rows = (len(batch) + 1) * (lengths[batch[0]] + SPECIAL_TOKEN_ROWS) if batch else 0
147
+ if batch and (len(batch) + 1 > max_sequences or padded_rows > max_tokens):
148
+ yield tuple(batch)
149
+ batch = []
150
+ batch.append(index)
151
+ if batch:
152
+ yield tuple(batch)
153
+
154
+
155
+ @dataclass(frozen=True, slots=True)
156
+ class HostBatch:
157
+ """The integer arrays of one batch, built on the host from known lengths."""
158
+
159
+ input_ids: np.ndarray # (b, m) int64, CLS, residue ids, EOS, then padding
160
+ rows: np.ndarray # (b,) int64, l_i + 2 attended tokens of each sequence
161
+ flat_index: np.ndarray # (n,) int64, each attended token's position in the row-major (b * m) grid
162
+ owner: np.ndarray # (n,) int64, the sequence of each attended token
163
+
164
+
165
+ def build_host_batch(vocabulary: ResidueVocabulary, texts: Sequence[str], *, columns: int | None = None) -> HostBatch:
166
+ """Token ids, lengths, and the row selection of ``texts`` (already cropped), with no device value.
167
+
168
+ ``columns`` pads every sequence to that many token columns (a geometry batch's bucket T); None pads to the
169
+ longest sequence.
170
+ """
171
+ encoded = [vocabulary.encode(text) for text in texts] # b arrays of (l_i + 2,)
172
+ rows = np.fromiter((len(ids) for ids in encoded), dtype=np.int64, count=len(encoded)) # (b,)
173
+ longest = int(rows.max())
174
+ if columns is not None and columns < longest:
175
+ raise ValueError(f"A batch padded to {columns} columns cannot hold a sequence of {longest} tokens.")
176
+ m = longest if columns is None else columns
177
+ input_ids = np.full((len(encoded), m), vocabulary.pad_id, dtype=np.int64) # (b, m)
178
+ for index, ids in enumerate(encoded):
179
+ input_ids[index, : len(ids)] = ids
180
+ owner = np.repeat(np.arange(len(encoded), dtype=np.int64), rows) # (n,)
181
+ starts = np.cumsum(rows) - rows # (b,) first packed row of each sequence
182
+ within = np.arange(int(rows.sum()), dtype=np.int64) - np.repeat(starts, rows) # (n,) token index in its sequence
183
+ return HostBatch(input_ids, rows, owner * m + within, owner)
184
+
185
+
186
+ def _to_device(arrays: Sequence[np.ndarray], device: torch.device) -> list[Tensor]:
187
+ """Copy several host arrays to the device in one transfer from pinned memory, without blocking the host."""
188
+ # arrays: (b, m), (b,), (n,), (n,) int64 for the executor's batch; each comes back with its own shape.
189
+ flat = np.concatenate([array.reshape(-1) for array in arrays]) # (total,) int64
190
+ if device.type == "cuda":
191
+ staging = torch.empty(flat.shape[0], dtype=torch.int64, pin_memory=True) # (total,)
192
+ staging.numpy()[:] = flat
193
+ moved = staging.to(device, non_blocking=True) # (total,)
194
+ else:
195
+ moved = torch.from_numpy(flat).to(device) # (total,)
196
+ pieces, cursor = [], 0
197
+ for array in arrays:
198
+ pieces.append(moved[cursor : cursor + array.size].view(array.shape))
199
+ cursor += array.size
200
+ return pieces # views shaped like arrays, e.g. (b, m), (b,), (n,), (n,), of one device buffer
201
+
202
+
203
+ class TokenTapExecutor:
204
+ """Run a tap plan over canonical proteins and hand each batch to the writer as pinned host tensors.
205
+
206
+ ``plan`` holds the taps. Every tap sees all attended tokens, CLS and EOS included, so a hidden tap keeps
207
+ ``(n, d)`` rows and a pooled tap averages l + 2 rows. ``max_residues`` is the N-terminal crop. ``dtype`` is
208
+ the run dtype a tap converts to unless it names its own, and None keeps the model's dtype. ``geometry`` runs
209
+ every batch at its bucket's fixed shape (``BatchGeometry``); ``fixed_batch_size`` fills every batch to that
210
+ many sequences but pads it to its longest member. A run uses one of the two at most.
211
+ """
212
+
213
+ def __init__(
214
+ self, model: Any, plan: TapPlan, *, vocabulary: ResidueVocabulary, max_residues: int | None,
215
+ dtype: torch.dtype | None,
216
+ fixed_batch_size: int | None = None,
217
+ geometry: BatchGeometry | None = None,
218
+ ) -> None:
219
+ if getattr(model, "embedding_tap_support", False) is not True:
220
+ raise ValueError(f"{type(model).__name__} does not support one-pass taps.")
221
+ if any(isinstance(tap, StreamingTap) for tap in plan.taps) and getattr(
222
+ model, "embedding_streaming_tap_support", False
223
+ ) is not True:
224
+ raise ValueError("This model does not support streaming hidden-state taps.")
225
+ if max_residues is not None and max_residues < 1:
226
+ raise ValueError("max_residues must be positive.")
227
+ if fixed_batch_size is not None and (type(fixed_batch_size) is not int or fixed_batch_size < 1):
228
+ raise ValueError("fixed_batch_size must be a positive integer.")
229
+ if geometry is not None:
230
+ if fixed_batch_size is not None:
231
+ raise ValueError("A run takes its batch shapes from a geometry or a fixed batch size, not both.")
232
+ if max_residues is None or max_residues + SPECIAL_TOKEN_ROWS > geometry.max_columns:
233
+ raise ValueError("A geometry run needs a crop whose l + 2 tokens fit the geometry's widest bucket.")
234
+ self.model = model
235
+ self.plan = plan
236
+ self.vocabulary = vocabulary
237
+ self.max_residues = max_residues
238
+ self.dtype = dtype
239
+ self.fixed_batch_size = fixed_batch_size
240
+ self.geometry = geometry
241
+ self.device = next(model.parameters()).device
242
+
243
+ def crop(self, sequence: str) -> str:
244
+ """The N-terminal crop: the first ``max_residues`` residues, so l <= max_residues and l + 2 tokens."""
245
+ return sequence if self.max_residues is None else sequence[: self.max_residues]
246
+
247
+ def batch_shape(self, texts: Sequence[str]) -> tuple[int, int] | None:
248
+ """The (rows, columns) a geometry runs ``texts`` (already cropped) at; None without a geometry."""
249
+ if self.geometry is None:
250
+ return None
251
+ buckets = {self.geometry.columns(len(text)) for text in texts}
252
+ if len(buckets) != 1:
253
+ raise ValueError(f"A geometry batch holds sequences of one bucket; these fall in {sorted(buckets)}.")
254
+ (columns,) = buckets
255
+ rows = self.geometry.rows(columns)
256
+ if len(texts) > rows:
257
+ raise ValueError(f"Bucket {columns} runs {rows} sequences to a batch; received {len(texts)}.")
258
+ return rows, columns
259
+
260
+ def run_batch(self, sequences: Sequence[str], digests: Sequence[str]) -> PackedBatch:
261
+ """Embed one batch and start the copies of its outputs to the host; returns without waiting for them."""
262
+ count = len(sequences)
263
+ if not count or len(digests) != count:
264
+ raise ValueError("A batch needs at least one sequence and one digest per sequence.")
265
+ if self.fixed_batch_size is not None and count > self.fixed_batch_size:
266
+ raise ValueError("The input exceeds the fixed physical batch size.")
267
+ texts = [self.crop(sequence) for sequence in sequences]
268
+ shape = self.batch_shape(texts)
269
+ columns = None
270
+ if shape is not None:
271
+ rows, columns = shape
272
+ texts += [texts[-1]] * (rows - count)
273
+ elif self.fixed_batch_size is not None:
274
+ texts += [texts[-1]] * (self.fixed_batch_size - count)
275
+ host = build_host_batch(self.vocabulary, texts, columns=columns)
276
+ input_ids, rows, flat_index, owner = _to_device(
277
+ (host.input_ids, host.rows, host.flat_index, host.owner), self.device,
278
+ ) # (b, m), (b,), (n,), (n,)
279
+ m = input_ids.shape[1]
280
+ # (b, m): true on CLS, residues and EOS; padding is the only masked position.
281
+ token_mask = torch.arange(m, device=self.device).unsqueeze(0) < rows.unsqueeze(1) # (b, m)
282
+ selection = RowSelection(flat_index, owner, tuple(int(count) for count in host.rows), rows)
283
+ cache: dict[Any, Any] = {}
284
+
285
+ def batch_for(X: Tensor) -> TapBatch:
286
+ # X: (b, m, d) one layer's hidden state; padding columns stay in X and are masked by token_mask (b, m).
287
+ return TapBatch(X=X, token_mask=token_mask, residue_mask=token_mask, rows=selection, cache=cache)
288
+
289
+ streaming = tuple(tap for tap in self.plan.taps if isinstance(tap, StreamingTap))
290
+ accumulators = {tap.name: tap.begin() for tap in streaming}
291
+ stream_layers = self.plan.streamed_layers
292
+
293
+ def consume(layer: int, X: Tensor) -> None:
294
+ # X: (b, m, d), borrowed until this callback returns
295
+ borrowed = batch_for(X)
296
+ for tap in streaming:
297
+ if layer in tap.layers:
298
+ accumulators[tap.name].update(layer, borrowed)
299
+
300
+ states = self.model._embed_taps(
301
+ input_ids, token_mask, self.plan.captured_layers, stream_layers=stream_layers,
302
+ state_consumer=consume if streaming else None, assume_valid_mask=True,
303
+ ) # {layer: (b, m, d)}
304
+ outputs: dict[str, dict[str, Tensor]] = {}
305
+ for tap, layer in zip(self.plan.taps, self.plan.layers, strict=True):
306
+ if isinstance(tap, StreamingTap):
307
+ Y = accumulators[tap.name].finish() # (b, m, w)
308
+ if tap.dtype is not None:
309
+ Y = Y.to(tap.dtype) # (b, m, w)
310
+ if tap.pooling is None:
311
+ outputs[tap.name] = {"values": selection.gather(Y)} # (n, w)
312
+ else:
313
+ outputs[tap.name] = {"values": pool_token_rows(Y, token_mask, tap.pooling)} # (b, p * w)
314
+ continue
315
+ dtype = tap.dtype if isinstance(tap, HiddenTap) and tap.dtype is not None else self.dtype
316
+ X = states[layer] # (b, m, d)
317
+ if dtype is not None:
318
+ X = X.to(dtype) # (b, m, d)
319
+ if isinstance(tap, SparseResidueTap):
320
+ if tap.reduce_packed is None:
321
+ raise ValueError(f"Tap {tap.name!r} has no packed reducer; it cannot keep special tokens.")
322
+ packed = tap.reduce_packed(batch_for(X)) # indices, values (n, k)
323
+ outputs[tap.name] = { # indices and values: (n, k)
324
+ "indices": packed.indices.to(INDEX_DTYPE),
325
+ "values": packed.values,
326
+ }
327
+ elif isinstance(tap, ReducedTap):
328
+ outputs[tap.name] = {"values": tap.reduce(batch_for(X))} # (b, w)
329
+ elif tap.pooling is None:
330
+ outputs[tap.name] = {"values": selection.gather(X)} # (n, d)
331
+ else:
332
+ outputs[tap.name] = {"values": pool_token_rows(X, token_mask, tap.pooling)} # (b, p * d)
333
+ if len(texts) != count:
334
+ token_rows = int(host.rows[:count].sum())
335
+ for tap in self.plan.taps:
336
+ per_token = isinstance(tap, SparseResidueTap) or (
337
+ isinstance(tap, (HiddenTap, StreamingTap)) and tap.pooling is None)
338
+ retained = token_rows if per_token else count
339
+ outputs[tap.name] = {name: tensor[:retained] for name, tensor in outputs[tap.name].items()}
340
+ return self._deliver(outputs, sequences, digests, tuple(int(rows) for rows in host.rows[:count]))
341
+
342
+ def _deliver(
343
+ self, outputs: Mapping[str, Mapping[str, Tensor]], sequences: Sequence[str], digests: Sequence[str],
344
+ rows: tuple[int, ...],
345
+ ) -> PackedBatch:
346
+ """Start the device-to-host copies, record one event after them, and check finiteness on the device."""
347
+ # outputs: (b, w) pooled, (n, w) per token row, (n, k) top-k indices, per tap; n = sum(l_i + 2).
348
+ finite = [torch.isfinite(tensor).all() for group in outputs.values() for tensor in group.values()
349
+ if tensor.is_floating_point()] # () per floating tensor
350
+ everything_finite = torch.stack(finite).all() if finite else torch.ones((), dtype=torch.bool, device=self.device) # ()
351
+ if self.device.type != "cuda":
352
+ hosts = {name: dict(group) for name, group in outputs.items()}
353
+ verdict = everything_finite
354
+
355
+ def wait() -> None:
356
+ if not bool(verdict):
357
+ raise ValueError("A tap produced a non-finite value.")
358
+ else:
359
+ hosts = {name: {key: _pinned_copy(tensor) for key, tensor in group.items()} for name, group in outputs.items()}
360
+ flag = _pinned_copy(everything_finite)
361
+ event = torch.cuda.Event()
362
+ event.record()
363
+
364
+ def wait() -> None:
365
+ event.synchronize()
366
+ if not bool(flag):
367
+ raise ValueError("A tap produced a non-finite value.")
368
+
369
+ return PackedBatch(tuple(sequences), tuple(digests), rows, hosts, wait)
370
+
371
+
372
+ def _pinned_copy(tensor: Tensor) -> Tensor:
373
+ """A pinned host tensor the copy of ``tensor`` is queued into on the current stream, without waiting."""
374
+ # tensor: (...) any shape; the host copy has the same shape and dtype.
375
+ host = torch.empty(tensor.shape, dtype=tensor.dtype, pin_memory=True) # (...)
376
+ host.copy_(tensor.detach(), non_blocking=True)
377
+ return host # (...) the shape of tensor
378
+
379
+
380
+ __all__ = [
381
+ "CANONICAL_MAX_RESIDUES",
382
+ "GEOMETRY_ALGORITHM",
383
+ "PARTIAL_BATCH_POLICY",
384
+ "SPECIAL_TOKEN_ROWS",
385
+ "BatchGeometry",
386
+ "HostBatch",
387
+ "TokenTapExecutor",
388
+ "build_host_batch",
389
+ "plan_geometry_batches",
390
+ "plan_token_batches",
391
+ ]
fastplms/embeddings/token_runs.py ADDED
@@ -0,0 +1,252 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Fill feature stores from canonical proteins with CLS and EOS kept, through the asynchronous writer.
2
+
3
+ ``embed_token_features`` is the fast path of ``embed_into_features(keep_special_tokens=True)``. It takes a
4
+ protein inventory, embeds only the rows a stream lacks, and writes each tap's rows into the store of its key.
5
+ The device loop (``TokenTapExecutor``) and the disk loop (``AsyncFeatureWriter``) run concurrently, joined by a
6
+ bounded queue of pinned host buffers, so the model is not paused for hashing, compression, or fsync.
7
+
8
+ Per batch, every per-token stream holds l + 2 rows per sequence (row 0 CLS, rows 1..l residues, row l + 1
9
+ EOS) and every pooled stream averages or maximizes over those same l + 2 rows. A legacy residue-only store
10
+ holds l rows; this path never writes one.
11
+
12
+ Symbols: b sequences of a batch; l residues of a sequence after the N-terminal crop; n = sum(l_i + 2) token rows.
13
+ """
14
+
15
+ from __future__ import annotations
16
+
17
+ import hashlib
18
+ import json
19
+ import torch
20
+
21
+ from collections.abc import Callable, Iterator, Mapping, Sequence
22
+ from contextlib import contextmanager
23
+ from pathlib import Path
24
+ from typing import Any, Protocol
25
+
26
+ from .batches import _temporary_eval
27
+ from .pooling import POOLING_SEMANTICS_TOKENS
28
+ from .taps import Tap, plan_taps
29
+ from .token_batches import (
30
+ CANONICAL_MAX_RESIDUES, SPECIAL_TOKEN_ROWS, BatchGeometry, TokenTapExecutor, plan_geometry_batches, plan_token_batches,
31
+ )
32
+ from .tokens import ResidueVocabulary, check_canonical_text
33
+ from ..features.async_writer import AsyncFeatureWriter
34
+ from ..features.store import FeatureStore, SegmentReceipt, StoredFeature, sequence_digest
35
+
36
+
37
+ GIB = 1024**3
38
+ TOKEN_RUN_SCHEMA = "token_features_v1"
39
+ BATCH_ALGORITHM = "token_budget_length_sorted_v1"
40
+ # The token budget's defaults, for a run without a geometry.
41
+ DEFAULT_MAX_SEQUENCES = 256
42
+ DEFAULT_MAX_TOKENS = 32768
43
+ DEFAULT_WINDOW = 65536
44
+
45
+
46
+ class TokenFeatureContract(Protocol):
47
+ """What a token run asks of a scientific contract: validate it, name each row, and recheck before a commit."""
48
+
49
+ def validate(
50
+ self, model: Any, sequences: Sequence[str], features: Mapping[str, StoredFeature],
51
+ taps: Sequence[Tap], options: Mapping[str, Any],
52
+ ) -> None: ...
53
+
54
+ def validate_cached(self, name: str, store: FeatureStore, sequences: Sequence[str]) -> None: ...
55
+
56
+ def row_identities(
57
+ self, name: str, sequences: Sequence[str], digests: Sequence[str],
58
+ ) -> Sequence[Mapping[str, Any]]: ...
59
+
60
+ def check_row_keys(self, sequences: Sequence[str], digests: Sequence[str]) -> None:
61
+ """Raise unless each digest is the key the contract holds for its sequence."""
62
+
63
+ def before_commit(self) -> None: ...
64
+
65
+
66
+ def distinct_with_digests(
67
+ sequences: Sequence[str], digests: Sequence[str] | None,
68
+ ) -> tuple[list[str], list[str]]:
69
+ """Distinct sequences in first-seen order and their row keys, hashing each exact text once.
70
+
71
+ ``digests`` may supply the keys when the caller already verified them against the text.
72
+ """
73
+ if digests is not None and len(digests) != len(sequences):
74
+ raise ValueError("digests must hold one key per sequence.")
75
+ seen: set[str] = set()
76
+ ordered: list[str] = []
77
+ keys: list[str] = []
78
+ for position, sequence in enumerate(sequences):
79
+ if sequence in seen:
80
+ continue
81
+ seen.add(sequence)
82
+ ordered.append(sequence)
83
+ keys.append(sequence_digest(sequence) if digests is None else digests[position])
84
+ return ordered, keys
85
+
86
+
87
+ def embed_token_features(
88
+ model: Any,
89
+ sequences: Sequence[str],
90
+ root: str | Path,
91
+ features: Mapping[str, StoredFeature],
92
+ *,
93
+ taps: Sequence[Tap],
94
+ contract: TokenFeatureContract | None = None,
95
+ digests: Sequence[str] | None = None,
96
+ metadata: Mapping[str, Any] | None = None,
97
+ max_residues: int | None = CANONICAL_MAX_RESIDUES,
98
+ max_sequences: int | None = None,
99
+ max_tokens: int | None = None,
100
+ window: int | None = None,
101
+ dtype: torch.dtype | None = None,
102
+ fixed_batch_size: int | None = None,
103
+ geometry: BatchGeometry | None = None,
104
+ part_bytes: int = GIB,
105
+ segment_bytes: int = 8 * GIB,
106
+ queue_bytes: int = 4 * GIB,
107
+ workers: int = 4,
108
+ verify_cached: bool = True,
109
+ model_state_fingerprint: str | None = None,
110
+ progress: Callable[[int], None] | None = None,
111
+ on_plan: Callable[[int], None] | None = None,
112
+ batch_watch: Callable[[Iterator[tuple[int, ...]]], Iterator[tuple[int, ...]]] | None = None,
113
+ ) -> dict[str, tuple[SegmentReceipt, ...]]:
114
+ """Fill each named feature under ``root`` from one pass over the sequences it lacks, special tokens kept.
115
+
116
+ ``features`` maps a tap name to its feature and names every tap. ``max_residues`` is the N-terminal
117
+ crop (``CANONICAL_MAX_RESIDUES`` keeps a sequence within 2048 tokens). Batches hold at most ``max_sequences`` sequences and
118
+ ``max_tokens`` padded token rows, sorted by length inside windows of ``window`` sequences (256, 32768 and 65536 when
119
+ None). A ``geometry`` instead runs every sequence at its bucket's fixed shape (``BatchGeometry``), whatever batch it
120
+ falls in, and takes none of those three nor ``fixed_batch_size``. Parts are
121
+ about ``part_bytes``, a segment commits every ``segment_bytes`` across all streams, and at most
122
+ ``queue_bytes`` of finished rows wait for the writer. ``digests`` are the rows' SHA-256 keys when the
123
+ caller already has them; otherwise each text is hashed once here. ``verify_cached=False`` skips the
124
+ contract's re-read of rows already stored, which a resumed run of a large store does separately.
125
+ ``progress`` receives the sequences of each batch once its rows are packed for writing, and ``on_plan`` the number of
126
+ sequences this run will embed (fewer than ``sequences`` on a resume), before the first batch.
127
+
128
+ Returns the committed segments of each stream that gained rows, oldest first; an empty result means the
129
+ model never ran. A killed run commits whole segments only, so a rerun resumes from the last one.
130
+ """
131
+ names = {tap.name for tap in taps}
132
+ if geometry is not None:
133
+ if any(value is not None for value in (max_sequences, max_tokens, window, fixed_batch_size)):
134
+ raise ValueError(
135
+ "A geometry run takes its batch shapes from the geometry; pass no max_sequences, max_tokens, window "
136
+ "or fixed_batch_size."
137
+ )
138
+ if max_residues is None or max_residues + SPECIAL_TOKEN_ROWS > geometry.max_columns:
139
+ raise ValueError("A geometry run needs a crop whose l + 2 tokens fit the geometry's widest bucket.")
140
+ else:
141
+ max_sequences = DEFAULT_MAX_SEQUENCES if max_sequences is None else max_sequences
142
+ max_tokens = DEFAULT_MAX_TOKENS if max_tokens is None else max_tokens
143
+ window = DEFAULT_WINDOW if window is None else window
144
+ if fixed_batch_size is not None:
145
+ if fixed_batch_size != max_sequences or max_residues is None or max_tokens < fixed_batch_size * (max_residues + 2):
146
+ raise ValueError("Fixed batches require matching max_sequences and a token budget covering the full cropped context.")
147
+ if set(features) != names:
148
+ raise ValueError(
149
+ "features must name exactly the taps this run takes.\n"
150
+ f" taps: {sorted(names)}\n features: {sorted(features)}"
151
+ )
152
+ if any(spec.positions for spec in features.values()):
153
+ raise ValueError("A token run stores no argmax positions.")
154
+ if contract is None and any(spec.descriptor.get("schema") is not None for spec in features.values()):
155
+ raise ValueError("A descriptor that carries a schema (feature_spec_v1, v2 or v3) requires its contract.")
156
+ if getattr(contract, "keep_special_tokens", None) is False:
157
+ raise ValueError("A token run needs a contract captured with the special tokens kept.")
158
+ ordered, keys = distinct_with_digests(sequences, digests)
159
+ if not ordered:
160
+ raise ValueError("embed_token_features needs at least one sequence.")
161
+ # A row is keyed by the hash of its text, so the text must be the normalized one (uppercase, no whitespace):
162
+ # a second spelling of a protein would otherwise become a second row. Checked for all before any forward.
163
+ for sequence in ordered:
164
+ check_canonical_text(sequence)
165
+ if contract is not None:
166
+ contract.check_row_keys(ordered, keys) # a dict lookup per row: a caller's digest is never trusted
167
+ # The contract compares these with the options it measured; its per-sequence checks are the caller's.
168
+ extraction_options: dict[str, Any] = {"max_length": max_residues, "truncate": True, "dtype": dtype}
169
+ if geometry is not None:
170
+ extraction_options["geometry"] = geometry.describe()
171
+ else:
172
+ extraction_options.update(batch_size=max_sequences, batch_window_size=window, max_tokens_per_batch=max_tokens)
173
+ if fixed_batch_size is not None:
174
+ extraction_options["fixed_batch_size"] = fixed_batch_size
175
+ contract.validate(model, (), features, taps, extraction_options)
176
+ stores = {name: FeatureStore.open(root, spec, deep_verify=False) for name, spec in features.items()}
177
+ # Each stream's missing rows come from one index pass over the keys, never a second hash of the text.
178
+ wanted = {name: frozenset(set(keys) - store.present_digests(keys)) for name, store in stores.items()}
179
+ if verify_cached and contract is not None:
180
+ for name, store in stores.items():
181
+ cached = [sequence for sequence, key in zip(ordered, keys, strict=True) if key not in wanted[name]]
182
+ if cached:
183
+ contract.validate_cached(name, store, cached)
184
+ chosen = [position for position, key in enumerate(keys) if any(key in group for group in wanted.values())]
185
+ if not chosen:
186
+ return {}
187
+ texts = [ordered[position] for position in chosen]
188
+ text_keys = [keys[position] for position in chosen]
189
+ if on_plan is not None:
190
+ on_plan(len(texts))
191
+ plan = plan_taps(list(taps), int(model.embedding_tap_state_count))
192
+ executor = TokenTapExecutor(
193
+ model, plan, vocabulary=ResidueVocabulary(model.tokenizer), max_residues=max_residues, dtype=dtype,
194
+ fixed_batch_size=fixed_batch_size, geometry=geometry,
195
+ )
196
+ lengths = [len(executor.crop(text)) for text in texts] # (count,) residues l after the crop
197
+ batch_policy: dict[str, Any] = dict(geometry.describe()) if geometry is not None else {
198
+ "algorithm": BATCH_ALGORITHM, "max_sequences": max_sequences, "max_tokens": max_tokens, "window": window,
199
+ }
200
+ run = {
201
+ "schema": TOKEN_RUN_SCHEMA, "special_tokens": "kept", "max_residues": max_residues,
202
+ "batch_policy": batch_policy,
203
+ "pooling_semantics": dict(POOLING_SEMANTICS_TOKENS), "model_state_fingerprint": model_state_fingerprint,
204
+ "storage_policy": {"max_part_bytes": part_bytes, "segment_bytes": segment_bytes},
205
+ }
206
+ if fixed_batch_size is not None:
207
+ run["batch_policy"].update(algorithm="fixed_rows_duplicate_pad_v1", fixed_batch_size=fixed_batch_size)
208
+ # Each stream's own missing rows name the segment. A kill after one stream committed leaves that stream
209
+ # with fewer missing rows, so the rerun gets new segment names and never collides with the committed one.
210
+ wanted_digests = {
211
+ features[name].key: hashlib.sha256("".join(sorted(group)).encode("ascii")).hexdigest()
212
+ for name, group in wanted.items()
213
+ }
214
+ fingerprint = hashlib.sha256(json.dumps({"run": run, "wanted": wanted_digests}, sort_keys=True)
215
+ .encode("utf-8")).hexdigest()[:32]
216
+
217
+ def identities(stream: str, batch_texts: Sequence[str], batch_keys: Sequence[str]) -> Sequence[Mapping[str, Any]]:
218
+ if contract is None:
219
+ return [{} for _ in batch_texts]
220
+ # Identities bind the original text, whose length the crop policy turns into l; never the cropped text.
221
+ return contract.row_identities(stream, batch_texts, batch_keys)
222
+
223
+ writer = AsyncFeatureWriter(
224
+ stores, fingerprint=fingerprint, metadata={**dict(metadata or {}), **run}, wanted=wanted,
225
+ row_records=identities, part_bytes=part_bytes, segment_bytes=segment_bytes, queue_bytes=queue_bytes,
226
+ before_commit=None if contract is None else contract.before_commit, workers=workers, progress=progress,
227
+ )
228
+ try:
229
+ with _inference(model):
230
+ if geometry is not None:
231
+ batches = plan_geometry_batches(lengths, text_keys, geometry)
232
+ else:
233
+ batches = plan_token_batches(lengths, max_sequences=max_sequences, max_tokens=max_tokens, window=window)
234
+ if batch_watch is not None:
235
+ batches = batch_watch(batches)
236
+ for members in batches:
237
+ writer.submit(executor.run_batch([texts[i] for i in members], [text_keys[i] for i in members]))
238
+ except BaseException:
239
+ writer.abort()
240
+ raise
241
+ receipts = writer.close()
242
+ return {name: tuple(group) for name, group in receipts.items() if group}
243
+
244
+
245
+ @contextmanager
246
+ def _inference(model: Any) -> Iterator[None]:
247
+ """Evaluation mode and inference mode for the whole loop, restored afterwards."""
248
+ with _temporary_eval(model), torch.inference_mode():
249
+ yield
250
+
251
+
252
+ __all__ = ["BATCH_ALGORITHM", "TOKEN_RUN_SCHEMA", "TokenFeatureContract", "distinct_with_digests", "embed_token_features"]
fastplms/embeddings/tokens.py ADDED
@@ -0,0 +1,74 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Tokenize canonical proteins once, as lookup-table rows, and check the table against the tokenizer.
2
+
3
+ A canonical sequence is already normalized (uppercase ASCII letters), and an ESM tokenizer maps each
4
+ residue letter to one id, so a protein of `l` residues is `l + 2` ids: CLS, the residues, EOS. The
5
+ vocabulary builds a 256-entry table from the tokenizer once, then encodes a sequence by one array
6
+ lookup instead of a Python call per residue. `verify` holds the table to the tokenizer itself, so
7
+ the ids the model sees are the ids the tokenizer would have produced.
8
+
9
+ Symbols: `l` residues of one sequence, so `l + 2` token ids (row 0 CLS, rows 1..l residues, row `l + 1` EOS).
10
+ """
11
+
12
+ from __future__ import annotations
13
+
14
+ import string
15
+ import numpy as np
16
+
17
+ from typing import Any
18
+ from numpy.typing import NDArray
19
+
20
+
21
+ UNKNOWN_ID = -1
22
+ _LETTERS = string.ascii_uppercase
23
+
24
+
25
+ def check_canonical_text(sequence: str) -> None:
26
+ """Reject text that is not an already normalized uppercase protein; this never normalizes."""
27
+ if not sequence or not sequence.isascii() or not sequence.isalpha() or not sequence.isupper():
28
+ raise ValueError("Canonical feature input must be an already normalized uppercase protein.")
29
+
30
+
31
+ class ResidueVocabulary:
32
+ """Residue letter to token id, built once from a tokenizer and verified against it."""
33
+
34
+ def __init__(self, tokenizer: Any) -> None:
35
+ special = set(tokenizer.all_special_ids)
36
+ vocabulary = tokenizer.get_vocab()
37
+ table = np.full(256, UNKNOWN_ID, dtype=np.int64) # (256,) ASCII code to token id
38
+ for letter in _LETTERS:
39
+ token_id = vocabulary.get(letter)
40
+ if token_id is not None and token_id not in special:
41
+ table[ord(letter)] = token_id
42
+ self.table = table
43
+ self.cls_id = int(tokenizer.cls_token_id)
44
+ self.eos_id = int(tokenizer.eos_token_id)
45
+ self.pad_id = int(tokenizer.pad_token_id)
46
+ self.verify(tokenizer)
47
+
48
+ def verify(self, tokenizer: Any) -> None:
49
+ """Hold every mapped letter to the tokenizer: `CLS, letter, EOS` must be its encoding."""
50
+ for letter in _LETTERS:
51
+ token_id = int(self.table[ord(letter)])
52
+ if token_id == UNKNOWN_ID:
53
+ continue
54
+ encoded = tokenizer([letter], add_special_tokens=True)["input_ids"][0]
55
+ if list(encoded) != [self.cls_id, token_id, self.eos_id]:
56
+ raise ValueError(
57
+ f"Residue {letter!r} encodes as {list(encoded)}, not the table's "
58
+ f"{[self.cls_id, token_id, self.eos_id]}; the tokenizer is not one token per residue."
59
+ )
60
+
61
+ def encode(self, sequence: str) -> NDArray[np.int64]:
62
+ """Token ids `(l + 2,)` of one canonical sequence: CLS, one id per residue, EOS."""
63
+ check_canonical_text(sequence)
64
+ letters = np.frombuffer(sequence.encode("ascii"), dtype=np.uint8) # (l,)
65
+ residues = self.table[letters] # (l,) token ids, UNKNOWN_ID where the tokenizer lacks the letter
66
+ if bool((residues == UNKNOWN_ID).any()):
67
+ raise ValueError("A canonical residue has no non-special tokenizer representation.")
68
+ ids = np.empty(len(letters) + 2, dtype=np.int64) # (l+2,)
69
+ ids[0], ids[-1] = self.cls_id, self.eos_id
70
+ ids[1:-1] = residues
71
+ return ids # (l+2,)
72
+
73
+
74
+ __all__ = ["UNKNOWN_ID", "ResidueVocabulary", "check_canonical_text"]
fastplms/embeddings/types.py CHANGED
@@ -7,6 +7,9 @@ from dataclasses import dataclass, field
7
  from typing import Any, Literal, overload
8
  from torch import Tensor
9
 
 
 
 
10
 
11
  @dataclass(frozen=True, slots=True)
12
  class EmbeddingInput:
@@ -38,7 +41,7 @@ class LazyTensorReference:
38
 
39
  if not isinstance(verify, bool):
40
  raise TypeError("verify must be a boolean.")
41
- X = self._loader() # self.shape
42
  if not isinstance(X, Tensor):
43
  raise TypeError(f"Stored tensor loader for {self.key!r} must return a Tensor.")
44
  if tuple(X.shape) != self.shape:
@@ -56,7 +59,7 @@ class LazyTensorReference:
56
  digest = tensor_sha256(X)
57
  if digest != self.sha256:
58
  raise ValueError(f"Stored tensor {self.key!r} failed SHA-256 verification.")
59
- return X # self.shape
60
 
61
 
62
  TensorValue = Tensor | LazyTensorReference
@@ -176,11 +179,75 @@ class EmbeddingBatch:
176
  attentions: Tensor | tuple[Tensor, ...] | None = None
177
 
178
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
179
  __all__ = [
180
  "EmbeddingBatch",
181
  "EmbeddingInput",
182
  "EmbeddingRecord",
183
  "EmbeddingResult",
184
  "LazyTensorReference",
 
 
 
185
  "TensorValue",
186
  ]
 
7
  from typing import Any, Literal, overload
8
  from torch import Tensor
9
 
10
+ from types import MappingProxyType
11
+ from ..features.layouts import TopKRow
12
+
13
 
14
  @dataclass(frozen=True, slots=True)
15
  class EmbeddingInput:
 
41
 
42
  if not isinstance(verify, bool):
43
  raise TypeError("verify must be a boolean.")
44
+ X = self._loader() # (...), equal to self.shape
45
  if not isinstance(X, Tensor):
46
  raise TypeError(f"Stored tensor loader for {self.key!r} must return a Tensor.")
47
  if tuple(X.shape) != self.shape:
 
59
  digest = tensor_sha256(X)
60
  if digest != self.sha256:
61
  raise ValueError(f"Stored tensor {self.key!r} failed SHA-256 verification.")
62
+ return X # (...), equal to self.shape
63
 
64
 
65
  TensorValue = Tensor | LazyTensorReference
 
179
  attentions: Tensor | tuple[Tensor, ...] | None = None
180
 
181
 
182
+ @dataclass(frozen=True, slots=True)
183
+ class TapRecord:
184
+ """One sequence's outputs from a tap plan, keyed by tap name."""
185
+
186
+ id: str
187
+ sequence: str
188
+ tensors: Mapping[str, Tensor | TopKRow]
189
+ retained_positions: tuple[int, ...] | None = None
190
+
191
+ def __post_init__(self) -> None:
192
+ if not isinstance(self.id, str) or not self.id:
193
+ raise ValueError("TapRecord.id must be a non-empty string.")
194
+ if not isinstance(self.sequence, str) or not self.sequence:
195
+ raise ValueError("TapRecord.sequence must be a non-empty string.")
196
+ if not isinstance(self.tensors, Mapping) or not self.tensors:
197
+ raise TypeError("TapRecord.tensors must be a non-empty mapping of tap name to Tensor.")
198
+ if not all(
199
+ isinstance(name, str) and isinstance(value, (Tensor, TopKRow))
200
+ for name, value in self.tensors.items()
201
+ ):
202
+ raise TypeError("TapRecord.tensors must map tap names to Tensor or TopKRow values.")
203
+ object.__setattr__(self, "tensors", MappingProxyType(dict(self.tensors)))
204
+ if self.retained_positions is not None:
205
+ positions = self.retained_positions
206
+ if (type(positions) is not tuple or not positions
207
+ or any(type(p) is not int or not 0 <= p < len(self.sequence) for p in positions)
208
+ or tuple(sorted(set(positions))) != positions):
209
+ raise ValueError(
210
+ "TapRecord retained positions must be ordered original-sequence indices."
211
+ )
212
+
213
+
214
+ @dataclass(frozen=True, slots=True)
215
+ class TapRunReceipt:
216
+ """Completed sink delivery, retaining run metadata but no output tensors."""
217
+
218
+ record_count: int
219
+ metadata: Mapping[str, Any]
220
+
221
+
222
+ class TapResult:
223
+ """Ordered tap records and the metadata needed to reproduce them."""
224
+
225
+ def __init__(
226
+ self,
227
+ records: Sequence[TapRecord],
228
+ metadata: Mapping[str, Any] | None = None,
229
+ ) -> None:
230
+ self.records: tuple[TapRecord, ...] = tuple(records)
231
+ self.metadata = dict(metadata or {})
232
+
233
+ def __len__(self) -> int:
234
+ return len(self.records)
235
+
236
+ def __iter__(self) -> Iterator[TapRecord]:
237
+ return iter(self.records)
238
+
239
+ def __getitem__(self, index: int) -> TapRecord:
240
+ return self.records[index]
241
+
242
+
243
  __all__ = [
244
  "EmbeddingBatch",
245
  "EmbeddingInput",
246
  "EmbeddingRecord",
247
  "EmbeddingResult",
248
  "LazyTensorReference",
249
+ "TapRecord",
250
+ "TapResult",
251
+ "TapRunReceipt",
252
  "TensorValue",
253
  ]
fastplms/features/__init__.py ADDED
@@ -0,0 +1,77 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """The FastPLMs feature store: one storage format for every embedding this workspace keeps.
2
+
3
+ A feature is one value per sequence, addressed by the SHA-256 of the sequence and by a key that
4
+ names the model, its revision, the sparse autoencoder, the layer, the pooling, the dtype, and the
5
+ residue limit. `store` holds the directory format and `layouts` the row layouts: dense vectors,
6
+ compressed-sparse pooled rows, ragged hidden states, and ragged top-k residue codes. `reader`
7
+ serves rows by random access to loops that read a batch at a time, `writing` commits a window of
8
+ rows as one segment, and `conversion` moves a cache another format holds into a store and proves the
9
+ rows survived.
10
+ """
11
+
12
+ from .async_writer import AsyncFeatureWriter, PackedBatch
13
+ from .conversion import (
14
+ ConversionMismatch,
15
+ ConversionReceipt,
16
+ conversion_fingerprint,
17
+ convert_rows,
18
+ describe_file,
19
+ )
20
+ from .layouts import (
21
+ CSR,
22
+ DENSE,
23
+ LAYOUT_NAMES,
24
+ RAGGED,
25
+ RAGGED_TOPK,
26
+ SparseRow,
27
+ TopKRow,
28
+ dtype_name,
29
+ value_dtype,
30
+ )
31
+ from .reader import CsrRows, FeatureReader
32
+ from .store import (
33
+ FORMAT,
34
+ FeatureStore,
35
+ RowAddress,
36
+ SegmentReceipt,
37
+ SegmentWriter,
38
+ StoredFeature,
39
+ features_in,
40
+ open_feature,
41
+ partition_sequences,
42
+ sequence_digest,
43
+ )
44
+ from .writing import write_rows
45
+
46
+
47
+ __all__ = [
48
+ "CSR",
49
+ "DENSE",
50
+ "FORMAT",
51
+ "LAYOUT_NAMES",
52
+ "RAGGED",
53
+ "RAGGED_TOPK",
54
+ "AsyncFeatureWriter",
55
+ "ConversionMismatch",
56
+ "ConversionReceipt",
57
+ "CsrRows",
58
+ "FeatureReader",
59
+ "FeatureStore",
60
+ "PackedBatch",
61
+ "RowAddress",
62
+ "SegmentReceipt",
63
+ "SegmentWriter",
64
+ "SparseRow",
65
+ "StoredFeature",
66
+ "TopKRow",
67
+ "conversion_fingerprint",
68
+ "convert_rows",
69
+ "describe_file",
70
+ "dtype_name",
71
+ "features_in",
72
+ "open_feature",
73
+ "partition_sequences",
74
+ "sequence_digest",
75
+ "value_dtype",
76
+ "write_rows",
77
+ ]
fastplms/features/async_writer.py ADDED
@@ -0,0 +1,375 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Write the rows of many streams behind the model: one bounded queue, one ingest thread, parallel part writes.
2
+
3
+ The embedding loop runs on the device. This writer takes each finished batch from pinned host buffers,
4
+ packs rows into parts of about ``part_bytes``, writes and hashes each part in one pass on a pool thread,
5
+ and commits every stream together once ``segment_bytes`` of parts exist, so a killed run loses at most
6
+ one segment. ``submit`` blocks only when the queued batches exceed ``queue_bytes``, which is what keeps
7
+ the device from outrunning the disk.
8
+
9
+ Symbols: b sequences of a batch; n token rows of a batch (sum of l_i + 2, with l_i the residues of
10
+ sequence i after the crop); w stored columns of a stream; k SAE codes kept per token; c SAE codebook.
11
+ Layouts of one batch, by stream:
12
+
13
+ - ragged values (n, w): row 0 CLS, rows 1..l_i residues, row l_i + 1 EOS of each sequence, in order.
14
+ - ragged_topk indices, values (n, k): the top-k SAE codes of every token row.
15
+ - dense values (b, w): one pooled vector per sequence.
16
+ - csr values (b, c) dense on arrival, compressed here to (nnz,) indices and values.
17
+ """
18
+
19
+ from __future__ import annotations
20
+
21
+ import threading
22
+ import torch
23
+
24
+ from collections import deque
25
+ from collections.abc import Callable, Mapping, Sequence
26
+ from concurrent.futures import Future, ThreadPoolExecutor
27
+ from contextlib import ExitStack
28
+ from dataclasses import dataclass, field
29
+ from typing import Any
30
+ from torch import Tensor
31
+
32
+ from .layouts import CSR, DENSE, INDEX_DTYPE, OFFSET_DTYPE, RAGGED, RAGGED_TOPK
33
+ from .store import FeatureStore, SegmentReceipt, SegmentWriter
34
+
35
+
36
+ @dataclass(eq=False)
37
+ class PackedBatch:
38
+ """One batch of finished rows for every stream, on the host, as the executor hands them over.
39
+
40
+ ``streams`` maps a stream to its host tensors by layout: ragged ``{"values": (n, w)}``, ragged
41
+ top-k ``{"values": (n, k), "indices": (n, k)}``, dense ``{"values": (b, w)}``, csr
42
+ ``{"values": (b, c)}`` dense, which the writer compresses. ``rows`` counts the stored rows of each
43
+ sequence in a ragged stream (l_i + 2 when the special tokens are kept). ``wait`` blocks until the
44
+ device copies landed and the batch passed its finite check, and raises if it did not.
45
+ """
46
+
47
+ sequences: tuple[str, ...] # (b,) the exact text of each row, for the row identities
48
+ digests: tuple[str, ...] # (b,) SHA-256 row keys, hashed once by the caller
49
+ rows: tuple[int, ...] # (b,) stored rows per sequence for ragged streams
50
+ streams: Mapping[str, Mapping[str, Tensor]] # per stream: values (n, w) | (n, k) with indices (n, k) | (b, w) | (b, c)
51
+ wait: Callable[[], None]
52
+ nbytes: int = field(init=False)
53
+
54
+ def __post_init__(self) -> None:
55
+ self.nbytes = sum(
56
+ tensor.numel() * tensor.element_size() for group in self.streams.values() for tensor in group.values()
57
+ )
58
+
59
+
60
+ @dataclass(eq=False)
61
+ class _StreamBuffer:
62
+ """Rows of one stream waiting to fill a part."""
63
+
64
+ sequences: list[str] = field(default_factory=list)
65
+ digests: list[str] = field(default_factory=list)
66
+ rows: list[int] = field(default_factory=list) # stored rows per sequence, as in PackedBatch
67
+ nnz: list[int] = field(default_factory=list) # csr entries per sequence, zero for other layouts
68
+ tensors: dict[str, list[Tensor]] = field(default_factory=dict) # per name: slices (n_i, w) or (b_i, w), joined at flush
69
+ nbytes: int = 0
70
+
71
+
72
+ def _row_bytes(layout: str, tensors: Mapping[str, Tensor], rows: Sequence[int], nnz: Sequence[int]) -> list[int]:
73
+ """Encoded payload of each sequence's row, its offset included, which the part budget counts."""
74
+ # tensors: (b, w) per stream when dense, (nnz,) when csr, (n, w) when ragged; n = sum(rows) token rows.
75
+ if layout == DENSE:
76
+ per_row = sum(tensor.shape[1] * tensor.element_size() for tensor in tensors.values())
77
+ return [per_row] * len(rows)
78
+ if layout == CSR:
79
+ per_entry = sum(tensor.element_size() for tensor in tensors.values())
80
+ return [8 + count * per_entry for count in nnz]
81
+ per_token = sum(tensor.shape[1] * tensor.element_size() for tensor in tensors.values())
82
+ return [8 + count * per_token for count in rows]
83
+
84
+
85
+ class AsyncFeatureWriter:
86
+ """Take finished batches from the embedding loop and commit them as rolling segments of every stream.
87
+
88
+ ``stores`` maps a stream to its opened store, whose descriptor carries the layout. ``wanted`` maps a
89
+ stream to the digests it lacks, so a rerun fills a lagging stream without rewriting the others.
90
+ ``row_records(stream, sequences, digests)`` returns each row's persisted identity. ``before_commit`` runs once
91
+ per committed group, before the first stream's marker, and may raise to refuse the commit.
92
+ """
93
+
94
+ def __init__(
95
+ self,
96
+ stores: Mapping[str, FeatureStore],
97
+ *,
98
+ fingerprint: str,
99
+ metadata: Mapping[str, Any],
100
+ wanted: Mapping[str, frozenset[str]],
101
+ row_records: Callable[[str, Sequence[str], Sequence[str]], Sequence[Mapping[str, Any]]],
102
+ part_bytes: int,
103
+ segment_bytes: int,
104
+ queue_bytes: int,
105
+ before_commit: Callable[[], None] | None = None,
106
+ workers: int = 4,
107
+ progress: Callable[[int], None] | None = None,
108
+ verify_staged: bool = False,
109
+ ) -> None:
110
+ if min(part_bytes, segment_bytes, queue_bytes, workers) < 1:
111
+ raise ValueError("part_bytes, segment_bytes, queue_bytes and workers must be positive.")
112
+ self._stores = dict(stores)
113
+ self._fingerprint = fingerprint
114
+ self._metadata = dict(metadata)
115
+ self._wanted = dict(wanted)
116
+ self._row_records = row_records
117
+ self._part_bytes = part_bytes
118
+ self._segment_bytes = segment_bytes
119
+ self._queue_bytes = queue_bytes
120
+ self._before_commit = before_commit
121
+ self._progress = progress
122
+ self._verify_staged = verify_staged
123
+ self._buffers = {name: _StreamBuffer() for name in self._stores}
124
+ self._pool = ThreadPoolExecutor(max_workers=workers, thread_name_prefix="feature-part")
125
+ # Parts in flight, written or waiting for a pool thread; more would only queue host memory.
126
+ self._slots = threading.Semaphore(workers + 2)
127
+ self._futures: deque[Future[None]] = deque()
128
+ self._condition = threading.Condition()
129
+ self._queue: deque[PackedBatch | None] = deque()
130
+ self._queued_bytes = 0
131
+ self._error: BaseException | None = None
132
+ # Held around each commit, so an abort lands before a commit starts or after it ends, never during one.
133
+ self._commit_lock = threading.Lock()
134
+ self._aborted = False
135
+ self._segment_index = 0
136
+ self._segment_written = 0
137
+ self._receipts: dict[str, list[SegmentReceipt]] = {name: [] for name in self._stores}
138
+ self._stack = ExitStack()
139
+ self._writers: dict[str, SegmentWriter] = {}
140
+ self._checked = False
141
+ self._open_segments()
142
+ self._thread = threading.Thread(target=self._ingest, name="feature-ingest", daemon=True)
143
+ self._thread.start()
144
+
145
+ def submit(self, batch: PackedBatch) -> None:
146
+ """Queue one batch, blocking while the queue holds more than ``queue_bytes`` of finished rows."""
147
+ with self._condition:
148
+ while self._error is None and self._queued_bytes and self._queued_bytes + batch.nbytes > self._queue_bytes:
149
+ self._condition.wait()
150
+ self._raise_error()
151
+ self._queue.append(batch)
152
+ self._queued_bytes += batch.nbytes
153
+ self._condition.notify_all()
154
+
155
+ def close(self) -> dict[str, list[SegmentReceipt]]:
156
+ """Write what is queued, commit the last segment, and return every committed segment by stream."""
157
+ with self._condition:
158
+ self._queue.append(None)
159
+ self._condition.notify_all()
160
+ self._thread.join()
161
+ self._pool.shutdown(wait=True)
162
+ if self._error is not None:
163
+ self._discard()
164
+ raise self._error
165
+ return self._receipts
166
+
167
+ def abort(self) -> None:
168
+ """Stop without committing: queued rows are dropped and open segments stay uncommitted for ``sweep``."""
169
+ with self._commit_lock: # a commit in flight finishes; none starts after this
170
+ self._aborted = True
171
+ with self._condition:
172
+ self._error = self._error or RuntimeError("The feature writer was aborted.")
173
+ self._queue.clear()
174
+ self._condition.notify_all()
175
+ self._thread.join()
176
+ self._pool.shutdown(wait=True)
177
+ self._discard()
178
+
179
+ def _discard(self) -> None:
180
+ """Close the open segments without committing them; a closed writer is skipped by its context."""
181
+ with self._commit_lock: # one closer at a time, and never during a commit
182
+ for writer in self._writers.values():
183
+ writer.closed = True
184
+ self._stack.close()
185
+
186
+ def _raise_error(self) -> None:
187
+ if self._error is not None:
188
+ raise self._error
189
+
190
+ def _open_segments(self) -> None:
191
+ name = f"{self._fingerprint}-{self._segment_index:05d}"
192
+ for stream, store in self._stores.items():
193
+ self._writers[stream] = self._stack.enter_context(store.segment(
194
+ name, {**self._metadata, "segment_index": self._segment_index},
195
+ before_commit=self._check_once, verify_staged=self._verify_staged,
196
+ ))
197
+
198
+ def _check_once(self) -> None:
199
+ """Run the caller's pre-commit check for the first stream of a group; its siblings reuse the verdict."""
200
+ if not self._checked:
201
+ if self._before_commit is not None:
202
+ self._before_commit()
203
+ self._checked = True
204
+
205
+ def _ingest(self) -> None:
206
+ try:
207
+ while True:
208
+ with self._condition:
209
+ while not self._queue and self._error is None:
210
+ self._condition.wait()
211
+ if self._error is not None:
212
+ return
213
+ batch = self._queue.popleft()
214
+ if batch is None:
215
+ self._commit_group()
216
+ return
217
+ batch.wait() # the device copies landed and the finite check passed
218
+ for stream in self._stores:
219
+ self._add(stream, batch)
220
+ with self._condition:
221
+ self._queued_bytes -= batch.nbytes
222
+ self._condition.notify_all()
223
+ if self._progress is not None:
224
+ self._progress(len(batch.sequences))
225
+ if self._segment_written >= self._segment_bytes:
226
+ self._commit_group(reopen=True)
227
+ # A worker thread has no caller to raise to: keep the error for `_raise_error` on the caller's thread.
228
+ except BaseException as error: # noqa: broad-except
229
+ with self._condition:
230
+ self._error = self._error or error
231
+ self._condition.notify_all()
232
+
233
+ def _add(self, stream: str, batch: PackedBatch) -> None:
234
+ layout = self._stores[stream].spec.layout
235
+ keep = [index for index, digest in enumerate(batch.digests) if digest in self._wanted[stream]]
236
+ if not keep:
237
+ return
238
+ tensors = batch.streams[stream] # values (n, w) | (n, k) with indices (n, k) | (b, w) | (b, c), by layout
239
+ rows = list(batch.rows) # (b,) stored rows per sequence: l_i + 2 when CLS and EOS are kept
240
+ if layout == CSR:
241
+ tensors, nnz = _compress_csr(tensors["values"]) # indices, values (nnz,); nnz per sequence (b,)
242
+ else:
243
+ nnz = [0] * len(rows)
244
+ if len(keep) != len(rows):
245
+ tensors, rows, nnz = _select_rows(layout, tensors, rows, nnz, keep)
246
+ sequences = [batch.sequences[index] for index in keep] # (m,) m sequences this stream lacks
247
+ digests = [batch.digests[index] for index in keep] # (m,)
248
+ sizes = _row_bytes(layout, tensors, rows, nnz) # (m,) payload bytes per sequence
249
+ if any(size > self._part_bytes for size in sizes):
250
+ raise ValueError("A feature row exceeds max_part_bytes; increase the explicit part budget.")
251
+ buffer = self._buffers[stream]
252
+ start = 0
253
+ while start < len(sequences):
254
+ # Take rows until the next would overflow the part, flush, and go on; a flushed buffer fits any row.
255
+ room = self._part_bytes - buffer.nbytes
256
+ stop, used = start, 0
257
+ while stop < len(sequences) and used + sizes[stop] <= room:
258
+ used += sizes[stop]
259
+ stop += 1
260
+ if stop == start:
261
+ self._flush(stream)
262
+ buffer = self._buffers[stream]
263
+ continue
264
+ _take(buffer, layout, tensors, sequences, digests, rows, nnz, start, stop, used)
265
+ self._segment_written += used # buffered rows count toward the segment, so a small segment commits per batch
266
+ if buffer.nbytes >= self._part_bytes:
267
+ self._flush(stream)
268
+ buffer = self._buffers[stream]
269
+ start = stop
270
+
271
+ def _flush(self, stream: str) -> None:
272
+ buffer = self._buffers[stream]
273
+ if not buffer.sequences:
274
+ return
275
+ self._buffers[stream] = _StreamBuffer()
276
+ # A failed part surfaces at the next flush, not only at the end of the segment.
277
+ while self._futures and self._futures[0].done():
278
+ self._futures.popleft().result()
279
+ writer = self._writers[stream]
280
+ part = writer.reserve_part()
281
+ self._slots.acquire()
282
+ future = self._pool.submit(self._write_part, stream, writer, part, buffer)
283
+ future.add_done_callback(lambda _: self._slots.release())
284
+ self._futures.append(future)
285
+
286
+ def _write_part(self, stream: str, writer: SegmentWriter, part: int, buffer: _StreamBuffer) -> None:
287
+ spec = self._stores[stream].spec
288
+ tensors = _pack(spec.layout, buffer) # offsets (b + 1,) and values (n, w), or (b, w), or indptr (b + 1,) and (nnz,)
289
+ identities = self._row_records(stream, buffer.sequences, buffer.digests) # (b,) one identity per sequence
290
+ residues = buffer.rows if spec.layout in (RAGGED, RAGGED_TOPK) else [0] * len(buffer.sequences) # (b,) stored rows
291
+ writer.append_packed(part, buffer.digests, tensors, residues, row_metadata=identities)
292
+
293
+ def _drain(self) -> None:
294
+ """Wait for every part in flight, surfacing the first write error."""
295
+ while self._futures:
296
+ self._futures.popleft().result()
297
+
298
+ def _commit_group(self, *, reopen: bool = False) -> None:
299
+ """Flush every stream, commit its open segment, and optionally open the next segment group."""
300
+ for stream in self._stores:
301
+ self._flush(stream)
302
+ self._drain()
303
+ self._checked = False
304
+ for stream, writer in self._writers.items():
305
+ with self._commit_lock:
306
+ if self._aborted:
307
+ raise RuntimeError("The feature writer was aborted before its segment committed.")
308
+ if writer.parts:
309
+ self._receipts[stream].append(writer.commit())
310
+ else:
311
+ writer.abandon()
312
+ self._stack.close()
313
+ if reopen:
314
+ self._stack = ExitStack()
315
+ self._segment_index += 1
316
+ self._segment_written = 0
317
+ self._open_segments()
318
+
319
+
320
+ def _compress_csr(dense: Tensor) -> tuple[dict[str, Tensor], list[int]]:
321
+ """Keep the nonzero codes of each sequence's row, in row-major order, as indices and values."""
322
+ # dense: (b, c) float32, zero where no code fired anywhere in the sequence
323
+ present = dense != 0 # (b, c)
324
+ counts = present.sum(dim=1).tolist() # (b,) nnz per sequence
325
+ columns = present.nonzero()[:, 1].to(INDEX_DTYPE) # (nnz,)
326
+ return {"indices": columns, "values": dense[present]}, [int(count) for count in counts] # values (nnz,)
327
+
328
+
329
+ def _select_rows(
330
+ layout: str, tensors: Mapping[str, Tensor], rows: Sequence[int], nnz: Sequence[int], keep: Sequence[int],
331
+ ) -> tuple[dict[str, Tensor], list[int], list[int]]:
332
+ """Gather the sequences at ``keep`` from packed batch tensors, for a rerun that fills only some rows."""
333
+ # tensors: (b, w) when dense, (n, w) when ragged, (nnz,) when csr; the kept rows keep the layout.
334
+ if layout == DENSE:
335
+ index = torch.tensor(list(keep), dtype=torch.int64) # (m,) m kept sequences
336
+ picked = {name: tensor[index] for name, tensor in tensors.items()}
337
+ return picked, [rows[i] for i in keep], [0] * len(keep) # (m, w) per stream, m kept sequences
338
+ spans = rows if layout in (RAGGED, RAGGED_TOPK) else nnz
339
+ pieces = {name: torch.split(tensor, list(spans)) for name, tensor in tensors.items()} # (span_i, ...) per sequence
340
+ selected = {name: torch.cat([parts[i] for i in keep]) for name, parts in pieces.items()}
341
+ return selected, [rows[i] for i in keep], [nnz[i] for i in keep] # (n_kept, w) or (nnz_kept,) per stream
342
+
343
+
344
+ def _take(
345
+ buffer: _StreamBuffer, layout: str, tensors: Mapping[str, Tensor], sequences: Sequence[str],
346
+ digests: Sequence[str], rows: Sequence[int], nnz: Sequence[int], start: int, stop: int, used: int,
347
+ ) -> None:
348
+ """Move sequences ``[start, stop)`` of a batch into the stream's part buffer, slicing the packed tensors."""
349
+ # tensors: (b, w) when dense, (n, w) when ragged, (nnz,) when csr; rows low:high of each go to the buffer.
350
+ buffer.sequences.extend(sequences[start:stop])
351
+ buffer.digests.extend(digests[start:stop])
352
+ buffer.rows.extend(rows[start:stop])
353
+ buffer.nnz.extend(nnz[start:stop])
354
+ if layout == DENSE:
355
+ low, high = start, stop # one row per sequence
356
+ else:
357
+ spans = rows if layout in (RAGGED, RAGGED_TOPK) else nnz
358
+ low, high = sum(spans[:start]), sum(spans[:stop]) # first and last packed row of the slice
359
+ for name, tensor in tensors.items():
360
+ buffer.tensors.setdefault(name, []).append(tensor[low:high])
361
+ buffer.nbytes += used
362
+
363
+
364
+ def _pack(layout: str, buffer: _StreamBuffer) -> dict[str, Tensor]:
365
+ """Concatenate a part's batches into the layout's tensors, with offsets naming each sequence's span."""
366
+ joined = {name: torch.cat(parts) for name, parts in buffer.tensors.items()} # (n, ...) one copy per tensor
367
+ if layout == DENSE:
368
+ return joined # values (b, w)
369
+ spans = torch.tensor(buffer.rows if layout in (RAGGED, RAGGED_TOPK) else buffer.nnz, dtype=OFFSET_DTYPE) # (b,)
370
+ offsets = torch.zeros(len(spans) + 1, dtype=OFFSET_DTYPE) # (b + 1,)
371
+ offsets[1:] = torch.cumsum(spans, dim=0)
372
+ return {("indptr" if layout == CSR else "offsets"): offsets, **joined} # offsets (b + 1,); values (n, w) or (nnz,)
373
+
374
+
375
+ __all__ = ["AsyncFeatureWriter", "PackedBatch"]
fastplms/features/conversion.py ADDED
@@ -0,0 +1,272 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Checked conversion of a cache another format holds into a feature store.
2
+
3
+ A project that already holds embeddings in its own format keeps them by converting them, not by
4
+ re-embedding. ``convert_rows`` takes the old cache as a function that decodes it, writes the rows
5
+ through the store's ordinary segment writer, then decodes the old cache a second time and compares
6
+ every row it produced with what the store now returns. Comparison is bit-exact on the stored
7
+ representation, so a converter never has to argue that two floats are close enough.
8
+
9
+ Nothing is repaired or skipped quietly. A row the store's dtype cannot hold exactly, a sequence
10
+ absent after the commit, a row that differs, a source that repeats a sequence with different rows,
11
+ and a source that changes between its two decodes each raise ``ConversionMismatch``. The old cache is
12
+ never modified or deleted: the move is undone by removing the segments this call names, and the
13
+ origin recorded in each commit marker says which files those rows came from.
14
+
15
+ What a project supplies is only the decoder of its old format. The decoder yields
16
+ ``(sequence, row)`` pairs, where a row is a tensor for ``dense`` and ``ragged`` features, a
17
+ ``SparseRow`` for ``csr``, and a ``TopKRow`` for ``ragged_topk``. It takes no arguments, because it
18
+ is called twice.
19
+ """
20
+
21
+ from __future__ import annotations
22
+
23
+ import json
24
+ import torch
25
+
26
+ from collections.abc import Callable, Iterable, Iterator, Mapping
27
+ from dataclasses import dataclass
28
+ from pathlib import Path
29
+ from typing import Any
30
+ from torch import Tensor
31
+
32
+ from .digests import file_sha256, json_sha256
33
+ from .layouts import CSR, RAGGED_TOPK, SparseRow, TopKRow
34
+ from .reader import FeatureReader
35
+ from .store import (
36
+ COMMIT_FILE,
37
+ SEGMENTS_DIRECTORY,
38
+ FeatureStore,
39
+ StoredFeature,
40
+ sequence_digest,
41
+ )
42
+
43
+
44
+ Row = Tensor | SparseRow | TopKRow
45
+
46
+
47
+ class ConversionMismatch(ValueError):
48
+ """The store does not hold exactly what the old cache held."""
49
+
50
+
51
+ @dataclass(frozen=True, slots=True)
52
+ class ConversionReceipt:
53
+ """What one call did: the segments it committed and how many rows it compared."""
54
+
55
+ segments: tuple[str, ...]
56
+ source_rows: int
57
+ written: int
58
+ skipped: int
59
+ verified: int
60
+
61
+
62
+ def describe_file(path: str | Path) -> dict[str, Any]:
63
+ """A file as an origin records it: its name, size, and content digest."""
64
+
65
+ location = Path(path)
66
+ return {
67
+ "name": location.name,
68
+ "bytes": location.stat().st_size,
69
+ "sha256": file_sha256(location),
70
+ }
71
+
72
+
73
+ def conversion_fingerprint(origin: Mapping[str, Any]) -> str:
74
+ """The stable name of the conversion of this origin, from which segment names derive."""
75
+
76
+ return "convert-" + json_sha256(dict(origin), allow_nan=False)[:16]
77
+
78
+
79
+ def convert_rows(
80
+ store: FeatureStore,
81
+ source: Callable[[], Iterable[tuple[str, Row]]],
82
+ *,
83
+ origin: Mapping[str, Any],
84
+ rows_per_segment: int = 100_000,
85
+ window_rows: int = 1_024,
86
+ max_tensor_bytes: int = 256 * 1024**2,
87
+ ) -> ConversionReceipt:
88
+ """Write the decoded rows into ``store``, then prove the store holds them.
89
+
90
+ ``origin`` is plain data naming where the rows came from, normally the old cache's format, the
91
+ decoder, and ``describe_file`` of each file. It is recorded in every segment's commit marker and
92
+ names the conversion, so calling this again with the same origin resumes an interrupted
93
+ conversion: committed segments are kept, sequences already in the store are not written twice,
94
+ and the comparison runs over the whole source. A long conversion commits a segment once it holds
95
+ ``rows_per_segment`` new rows, rounded up to a window, so an interruption loses at most one
96
+ segment.
97
+
98
+ Rows the store held before the call are compared too. A store that mixes converted rows with
99
+ rows embedded afresh therefore fails here, which is the point: the old numbers are not the
100
+ store's numbers.
101
+ """
102
+
103
+ if rows_per_segment < 1 or window_rows < 1:
104
+ raise ValueError("rows_per_segment and window_rows must be positive.")
105
+ fingerprint = conversion_fingerprint(origin)
106
+ ordinal = _committed_segments(store, fingerprint)
107
+ segments: list[str] = []
108
+ source_rows = 0
109
+ written = 0
110
+
111
+ def counted() -> Iterator[tuple[str, Row]]:
112
+ nonlocal source_rows
113
+ for pair in source():
114
+ source_rows += 1
115
+ yield pair
116
+
117
+ windows = _unique_windows(store.spec, counted(), window_rows)
118
+ pending = next(windows, None)
119
+ while pending is not None:
120
+ name = f"{fingerprint}-{ordinal:05d}"
121
+ metadata = {"conversion": {"origin": dict(origin), "ordinal": ordinal}}
122
+ staged: set[str] = set() # digests written into this segment, which is not yet committed
123
+ with store.segment(name, metadata) as writer:
124
+ while pending is not None and len(staged) < rows_per_segment:
125
+ absent = set(store.missing([sequence for sequence, _ in pending]))
126
+ fresh = [
127
+ pair for pair in pending
128
+ if pair[0] in absent and sequence_digest(pair[0]) not in staged
129
+ ]
130
+ if fresh:
131
+ writer.append_bounded(
132
+ [sequence for sequence, _ in fresh], [row for _, row in fresh],
133
+ max_tensor_bytes=max_tensor_bytes,
134
+ )
135
+ staged.update(sequence_digest(sequence) for sequence, _ in fresh)
136
+ pending = next(windows, None)
137
+ if staged:
138
+ segments.append(name)
139
+ ordinal += 1
140
+ written += len(staged)
141
+
142
+ if source_rows == 0:
143
+ raise ConversionMismatch("The source decoded no rows; check the path and the decoder.")
144
+ verified = _verify(store, source, window_rows)
145
+ if verified != source_rows:
146
+ raise ConversionMismatch(
147
+ f"The source decoded {source_rows} rows the first time and {verified} the second."
148
+ )
149
+ return ConversionReceipt(
150
+ segments=tuple(segments), source_rows=source_rows, written=written,
151
+ skipped=source_rows - written, verified=verified,
152
+ )
153
+
154
+
155
+ def _committed_segments(store: FeatureStore, fingerprint: str) -> int:
156
+ directory = store.directory / SEGMENTS_DIRECTORY
157
+ return sum(
158
+ 1 for marker in directory.glob(f"{fingerprint}-*/{COMMIT_FILE}") if marker.is_file()
159
+ )
160
+
161
+
162
+ def _windows(rows: Iterable[tuple[str, Row]], size: int) -> Iterator[list[tuple[str, Row]]]:
163
+ window: list[tuple[str, Row]] = []
164
+ for pair in rows:
165
+ window.append(pair)
166
+ if len(window) == size:
167
+ yield window
168
+ window = []
169
+ if window:
170
+ yield window
171
+
172
+
173
+ def _unique_windows(
174
+ spec: StoredFeature, rows: Iterable[tuple[str, Row]], size: int,
175
+ ) -> Iterator[list[tuple[str, Row]]]:
176
+ """Windows in which each sequence appears once.
177
+
178
+ A sequence repeated inside a window keeps its first row, and is an error if the rows differ.
179
+ A repeat in a later window is skipped when its sequence is already written, and the comparison
180
+ after the commit raises if its row differs from the one that was.
181
+ """
182
+
183
+ for window in _windows(rows, size):
184
+ unique: dict[str, tuple[str, Row]] = {}
185
+ for sequence, row in window:
186
+ digest = sequence_digest(sequence)
187
+ first = unique.setdefault(digest, (sequence, row))
188
+ if first[1] is not row and _row_bytes(first[1], spec) != _row_bytes(row, spec):
189
+ raise ConversionMismatch(
190
+ f"The source repeats a sequence (sha256 {digest[:12]}) with different rows."
191
+ )
192
+ yield list(unique.values())
193
+
194
+
195
+ def _verify(
196
+ store: FeatureStore, source: Callable[[], Iterable[tuple[str, Row]]], window_rows: int,
197
+ ) -> int:
198
+ """Decode the source again and compare each row with a fresh read of the committed store."""
199
+
200
+ spec = store.spec
201
+ verified = 0
202
+ with FeatureReader.open(store.directory) as reader:
203
+ for window in _windows(source(), window_rows):
204
+ sequences = [sequence for sequence, _ in window]
205
+ absent = reader.missing(sequences)
206
+ if absent:
207
+ raise ConversionMismatch(
208
+ f"{len(absent)} converted sequences are absent from {spec.key!r} after commit."
209
+ )
210
+ for (sequence, row), stored in zip(window, _read(reader, sequences), strict=True):
211
+ if _row_bytes(row, spec) != _row_bytes(stored, spec):
212
+ raise ConversionMismatch(
213
+ f"Feature {spec.key!r} differs from the source for the sequence with "
214
+ f"sha256 {sequence_digest(sequence)[:12]}."
215
+ )
216
+ verified += len(window)
217
+ return verified
218
+
219
+
220
+ def _read(reader: FeatureReader, sequences: list[str]) -> list[Row]:
221
+ layout = reader.spec.layout
222
+ if layout == CSR:
223
+ return list(reader.read_sparse(sequences))
224
+ if layout == RAGGED_TOPK:
225
+ return list(reader.read_topk(sequences))
226
+ return list(reader.read(sequences))
227
+
228
+
229
+ def _row_bytes(row: Row, spec: StoredFeature) -> bytes:
230
+ """A row as the store holds it, as bytes: shapes, then each tensor's raw storage.
231
+
232
+ Values are cast to the feature's dtype first, and a cast that changes any value raises, so
233
+ equal bytes mean equal stored rows.
234
+ """
235
+
236
+ # row: (w,) dense, (r_i, d) ragged; a SparseRow holds (nnz,) tensors and a TopKRow (r_i, k) tensors
237
+ if isinstance(row, SparseRow):
238
+ tensors = [row.indices.to(torch.int32), _stored_values(row.values, spec.dtype)] # (nnz,), (nnz,)
239
+ if row.positions is not None:
240
+ tensors.append(row.positions.to(torch.int16)) # (nnz,)
241
+ elif isinstance(row, TopKRow):
242
+ tensors = [row.indices.to(torch.int32), _stored_values(row.values, spec.dtype)] # (r_i, k), (r_i, k)
243
+ else:
244
+ tensors = [_stored_values(row, spec.dtype)] # (w,) or (r_i, d)
245
+ shapes = json.dumps([list(tensor.shape) for tensor in tensors]).encode("utf-8")
246
+ return shapes + b"".join(_raw_bytes(tensor) for tensor in tensors)
247
+
248
+
249
+ def _stored_values(values: Tensor, dtype: torch.dtype) -> Tensor:
250
+ # values: (...), the shape of one stored tensor; the cast keeps it
251
+ cast = values.detach().to("cpu").to(dtype) # (...)
252
+ if values.dtype != dtype and _raw_bytes(cast.to(values.dtype)) != _raw_bytes(values.detach().cpu()):
253
+ raise ConversionMismatch(
254
+ f"Storing {values.dtype} values as {dtype} would change them; choose a wider dtype."
255
+ )
256
+ return cast # (...)
257
+
258
+
259
+ def _raw_bytes(tensor: Tensor) -> bytes:
260
+ # tensor: (...)
261
+ flat = tensor.detach().to("cpu").contiguous().reshape(-1) # (n,), n = tensor.numel()
262
+ # An empty tensor can carry a zero stride, which the byte view refuses, and has no bytes to compare.
263
+ return b"" if flat.numel() == 0 else flat.view(torch.uint8).numpy().tobytes()
264
+
265
+
266
+ __all__ = [
267
+ "ConversionMismatch",
268
+ "ConversionReceipt",
269
+ "convert_rows",
270
+ "conversion_fingerprint",
271
+ "describe_file",
272
+ ]
fastplms/features/digests.py ADDED
@@ -0,0 +1,35 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """SHA-256 digests of files and JSON values, the identities FastPLMs records and compares.
2
+
3
+ This file exists twice, byte for byte: here and as ``features/digests.py``. ``features`` loads as a standalone
4
+ package (``foundry.embedding.private_store``) and imports nothing outside itself, so it cannot reach this
5
+ module. ``tests/tier1_unit/test_features_package_is_self_contained.py`` fails when the two differ.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import hashlib
11
+
12
+ from pathlib import Path
13
+ from typing import Any
14
+
15
+ from .json_files import compact_json
16
+
17
+
18
+ FILE_READ_BYTES = 1024 * 1024
19
+
20
+
21
+ def file_sha256(path: str | Path) -> str:
22
+ """Return the SHA-256 of a file's bytes, read in blocks so a checkpoint never sits in memory."""
23
+
24
+ digest = hashlib.sha256()
25
+ with Path(path).open("rb") as handle:
26
+ while block := handle.read(FILE_READ_BYTES):
27
+ digest.update(block)
28
+ return digest.hexdigest()
29
+
30
+
31
+ def json_sha256(value: Any, *, ensure_ascii: bool = True, allow_nan: bool = True) -> str:
32
+ """Return the SHA-256 of ``value`` in its compact, key-sorted JSON form (``compact_json``)."""
33
+
34
+ encoded = compact_json(value, ensure_ascii=ensure_ascii, allow_nan=allow_nan).encode("utf-8")
35
+ return hashlib.sha256(encoded).hexdigest()
fastplms/features/json_files.py ADDED
@@ -0,0 +1,45 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """The two JSON text forms FastPLMs writes: compact for hashing, indented for files people read.
2
+
3
+ This file exists twice, byte for byte: here and as ``features/json_files.py``. ``features`` loads as a
4
+ standalone package (``foundry.embedding.private_store``) and imports nothing outside itself, so it cannot
5
+ reach this module. ``tests/tier1_unit/test_features_package_is_self_contained.py`` fails when the two differ.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import json
11
+
12
+ from typing import Any
13
+
14
+
15
+ def compact_json(value: Any, *, ensure_ascii: bool = True, allow_nan: bool = True) -> str:
16
+ """Serialize with sorted keys and no whitespace, the form a digest or an identity is taken over."""
17
+
18
+ return json.dumps(
19
+ value,
20
+ sort_keys=True,
21
+ separators=(",", ":"),
22
+ ensure_ascii=ensure_ascii,
23
+ allow_nan=allow_nan,
24
+ )
25
+
26
+
27
+ def indented_json(
28
+ value: Any,
29
+ *,
30
+ ensure_ascii: bool = True,
31
+ allow_nan: bool = True,
32
+ sort_keys: bool = True,
33
+ ) -> str:
34
+ """Serialize with two-space indentation and one trailing newline, the form of a stored JSON file."""
35
+
36
+ return (
37
+ json.dumps(
38
+ value,
39
+ indent=2,
40
+ sort_keys=sort_keys,
41
+ ensure_ascii=ensure_ascii,
42
+ allow_nan=allow_nan,
43
+ )
44
+ + "\n"
45
+ )
fastplms/features/layouts.py ADDED
@@ -0,0 +1,354 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """The row layouts a feature segment stores, as memory-mappable safetensors tensors.
2
+
3
+ A feature is one value per sequence, and the layouts differ in what that value is:
4
+
5
+ - ``dense``: a fixed-width vector, ``values`` of shape ``(n, w)``. Pooled embeddings.
6
+ - ``csr``: a sparse fixed-width vector, compressed by row. Max-pooled sparse-autoencoder codes,
7
+ where a row holds at most one entry per code that fired anywhere in the sequence. ``positions``
8
+ carries the argmax residue of each entry, which is what makes a code interpretable as "here",
9
+ and is optional because a pooling that discards it has nothing to store.
10
+ - ``ragged``: per-row values, ``values`` of shape ``(sum r_i, d)`` with ``offsets`` naming each
11
+ sequence's span. Hidden states, where ``r_i`` is the sequence's stored row count: its biological residue
12
+ count ``l`` in a residue-only (``feature_spec_v1``) store, and ``l + 2`` in a canonical store, whose rows are
13
+ CLS, the ``l`` residues, then EOS.
14
+ - ``ragged_topk``: per-row sparse codes, ``indices`` and ``values`` of shape
15
+ ``(sum r_i, k)`` with sequence ``offsets``, ``r_i`` as above. All k entries, including zeros, retain their
16
+ order. Unlike pooled csr positions, these rows retain every stored row and support crop-local pooling.
17
+
18
+ Rows are addressed individually and never sliced by column, so compressed-sparse-row is the right
19
+ sparse layout: a batch materializes with one gather and no search. The index dtypes follow
20
+ enzyme_loop's measured store: int32 code indices, int16 residue positions, and int64 row offsets,
21
+ which keep an entry to eight bytes against four bytes per column dense.
22
+
23
+ Every layout stores its values in the feature's declared dtype, so a segment is a lossless record
24
+ of what the model produced at that precision, and no reader has to guess.
25
+ """
26
+
27
+ from __future__ import annotations
28
+
29
+ import torch
30
+
31
+ from collections.abc import Sequence
32
+ from dataclasses import dataclass
33
+ from torch import Tensor
34
+
35
+
36
+ DENSE = "dense"
37
+ CSR = "csr"
38
+ RAGGED = "ragged"
39
+ RAGGED_TOPK = "ragged_topk"
40
+ LAYOUT_NAMES = (DENSE, CSR, RAGGED, RAGGED_TOPK)
41
+
42
+ INDEX_DTYPE = torch.int32
43
+ POSITION_DTYPE = torch.int16
44
+ OFFSET_DTYPE = torch.int64
45
+
46
+ VALUE_DTYPES = (torch.float64, torch.float32, torch.float16, torch.bfloat16)
47
+ DTYPE_NAMES = {dtype: str(dtype).removeprefix("torch.") for dtype in VALUE_DTYPES}
48
+ NAME_DTYPES = {name: dtype for dtype, name in DTYPE_NAMES.items()}
49
+
50
+
51
+ @dataclass(frozen=True, slots=True)
52
+ class TopKRow:
53
+ """One sequence's sparse residue codes: integer indices and values, both ``(r_i, k)``.
54
+
55
+ The store validates the declared codebook and k. Zero-valued entries and code order are
56
+ retained exactly; there is deliberately no implicit residue-by-codebook densification.
57
+ """
58
+
59
+ indices: Tensor
60
+ values: Tensor
61
+
62
+
63
+ @dataclass(frozen=True, slots=True)
64
+ class SparseRow:
65
+ """One compressed row: the codes that fired, their values, and where each peaked.
66
+
67
+ ``indices`` and ``values`` have shape ``(nnz_i,)``. ``positions`` has the same shape when the
68
+ segment stores them, and is None when it does not.
69
+ """
70
+
71
+ indices: Tensor
72
+ values: Tensor
73
+ positions: Tensor | None
74
+
75
+ def to_dense(self, width: int) -> Tensor:
76
+ """The row as a ``(width,)`` vector, zero where no code fired."""
77
+
78
+ # self.indices, self.values: (nnz,), the entries that fired
79
+ dense = torch.zeros(width, dtype=self.values.dtype) # (w,)
80
+ dense[self.indices.to(torch.int64)] = self.values
81
+ return dense # (w,)
82
+
83
+ @classmethod
84
+ def from_dense(cls, vector: Tensor, positions: Tensor | None = None) -> SparseRow:
85
+ """Compress a ``(w,)`` vector by keeping its exactly non-zero entries.
86
+
87
+ Top-k sparse-autoencoder pooling leaves exact zeros where a code never fired, so this
88
+ drops nothing a reader could want. It is not a threshold: a code that fired weakly is
89
+ kept. ``positions`` is the full ``(w,)`` argmax residue vector, gathered at the same
90
+ entries.
91
+ """
92
+
93
+ # vector: (w,); positions: (w,) or None
94
+ if vector.ndim != 1:
95
+ raise ValueError(f"A sparse row comes from a one-dimensional vector; received {tuple(vector.shape)}.")
96
+ kept = torch.nonzero(vector, as_tuple=False).flatten() # (nnz,)
97
+ if positions is not None and positions.shape != vector.shape:
98
+ raise ValueError(
99
+ f"positions must have the vector's shape {tuple(vector.shape)}; received "
100
+ f"{tuple(positions.shape)}."
101
+ )
102
+ return cls(
103
+ indices=kept.to(INDEX_DTYPE), # (nnz,)
104
+ values=vector[kept], # (nnz,)
105
+ positions=None if positions is None else positions[kept], # (nnz,)
106
+ )
107
+
108
+
109
+ def dtype_name(dtype: torch.dtype) -> str:
110
+ """The stored name of a value dtype, rejecting one no layout stores."""
111
+
112
+ if dtype not in DTYPE_NAMES:
113
+ raise ValueError(
114
+ f"A feature stores float64, float32, float16, or bfloat16 values; received {dtype}."
115
+ )
116
+ return DTYPE_NAMES[dtype]
117
+
118
+
119
+ def value_dtype(name: str) -> torch.dtype:
120
+ """The dtype a stored name means."""
121
+
122
+ if name not in NAME_DTYPES:
123
+ raise ValueError(f"Unknown feature value dtype {name!r}; expected one of {list(NAME_DTYPES)}.")
124
+ return NAME_DTYPES[name]
125
+
126
+
127
+ def encode_dense(rows: Sequence[Tensor], width: int, dtype: torch.dtype) -> dict[str, Tensor]:
128
+ """Stack ``(w,)`` rows into one ``values`` tensor of shape ``(n, w)``."""
129
+
130
+ # rows: (w,) each, n tensors; w = width
131
+ stacked = torch.empty((len(rows), width), dtype=dtype) # (n, w)
132
+ for position, row in enumerate(rows):
133
+ vector = row.detach().to("cpu") # (w,)
134
+ if vector.ndim != 1 or vector.shape[0] != width:
135
+ raise ValueError(
136
+ f"A dense feature row must have shape ({width},); row {position} has "
137
+ f"{tuple(vector.shape)}."
138
+ )
139
+ stacked[position] = vector.to(dtype)
140
+ return {"values": stacked} # {"values": (n, w)}
141
+
142
+
143
+ def row_tensor_bytes(
144
+ row: Tensor | SparseRow | TopKRow, layout: str, width: int, dtype: torch.dtype,
145
+ *, positions: bool,
146
+ ) -> int:
147
+ """Exact per-row encoded payload, including its offset but not the part's initial offset."""
148
+ # row: (w,) dense, (r_i, d) ragged, or a SparseRow or TopKRow
149
+ value_size = torch.empty(0, dtype=dtype).element_size()
150
+ if layout == RAGGED_TOPK:
151
+ if not isinstance(row, TopKRow):
152
+ raise TypeError("Ragged top-k payload sizing requires a TopKRow.")
153
+ return 8 + row.values.numel() * (4 + value_size)
154
+ if layout == CSR:
155
+ if not isinstance(row, SparseRow):
156
+ raise TypeError("CSR payload sizing requires a SparseRow.")
157
+ return 8 + row.values.numel() * (4 + value_size + (2 if positions else 0))
158
+ if not isinstance(row, Tensor):
159
+ raise TypeError("Dense and ragged payload sizing requires tensor rows.")
160
+ return width * value_size if layout == DENSE else 8 + row.numel() * value_size
161
+
162
+
163
+ def encode_csr(
164
+ rows: Sequence[SparseRow], width: int, dtype: torch.dtype
165
+ ) -> dict[str, Tensor]:
166
+ """Concatenate sparse rows, with ``indptr`` naming each row's span.
167
+
168
+ Every row must agree about positions: either all carry them or none does, because one segment
169
+ stores one tensor set.
170
+ """
171
+
172
+ # rows[i].indices, .values, .positions: (nnz_i,); the segment holds n rows and nnz = sum(nnz_i) entries
173
+ with_positions = [row.positions is not None for row in rows]
174
+ if any(with_positions) and not all(with_positions):
175
+ raise ValueError(
176
+ "Either every sparse row carries argmax positions or none does; this batch mixes both."
177
+ )
178
+ indptr = torch.zeros(len(rows) + 1, dtype=OFFSET_DTYPE) # (n + 1,)
179
+ indices: list[Tensor] = []
180
+ values: list[Tensor] = []
181
+ positions: list[Tensor] = []
182
+ for position, row in enumerate(rows):
183
+ if row.indices.dtype not in (torch.int16, torch.int32, torch.int64):
184
+ raise ValueError("Sparse code indices must be signed integer tensors.")
185
+ row_indices = row.indices.detach().to("cpu").to(torch.int64) # (nnz_i,)
186
+ row_values = row.values.detach().to("cpu") # (nnz_i,)
187
+ if row_indices.ndim != 1 or row_values.shape != row_indices.shape:
188
+ raise ValueError(
189
+ f"Sparse row {position} needs one-dimensional indices and values of equal length; "
190
+ f"received {tuple(row_indices.shape)} and {tuple(row_values.shape)}."
191
+ )
192
+ if row_indices.numel() and (int(row_indices.min()) < 0 or int(row_indices.max()) >= width):
193
+ raise ValueError(
194
+ f"Sparse row {position} names a code outside 0..{width - 1}."
195
+ )
196
+ if row_indices.numel() and int(row_indices.max()) > torch.iinfo(INDEX_DTYPE).max:
197
+ raise ValueError("Sparse code index exceeds the stored integer range.")
198
+ if row_indices.unique().numel() != row_indices.numel():
199
+ raise ValueError("Sparse code indices must be unique within each row.")
200
+ indptr[position + 1] = int(indptr[position]) + row_indices.numel()
201
+ indices.append(row_indices.to(INDEX_DTYPE))
202
+ values.append(row_values.to(dtype))
203
+ if row.positions is not None:
204
+ if row.positions.dtype not in (torch.int16, torch.int32, torch.int64):
205
+ raise ValueError("Sparse residue positions must be signed integer tensors.")
206
+ row_positions = row.positions.detach().to("cpu") # (nnz_i,)
207
+ if row_positions.shape != row_indices.shape:
208
+ raise ValueError(
209
+ f"Sparse row {position} has {row_positions.numel()} positions for "
210
+ f"{row_indices.numel()} entries."
211
+ )
212
+ if row_positions.numel() and (
213
+ int(row_positions.min()) < 0
214
+ or int(row_positions.max()) > torch.iinfo(POSITION_DTYPE).max
215
+ ):
216
+ raise ValueError("Sparse residue position exceeds the stored integer range.")
217
+ positions.append(row_positions.to(POSITION_DTYPE))
218
+
219
+ tensors = {
220
+ "indptr": indptr, # (n + 1,)
221
+ "indices": _concatenate(indices, INDEX_DTYPE), # (nnz,)
222
+ "values": _concatenate(values, dtype), # (nnz,)
223
+ }
224
+ if positions:
225
+ tensors["positions"] = _concatenate(positions, POSITION_DTYPE) # (nnz,)
226
+ return tensors # indptr (n + 1,); indices, values and positions (nnz,)
227
+
228
+
229
+ def encode_ragged(rows: Sequence[Tensor], width: int, dtype: torch.dtype) -> dict[str, Tensor]:
230
+ """Concatenate ``(r_i, d)`` residue blocks, with ``offsets`` naming each sequence's span."""
231
+
232
+ # rows: (r_i, d) each, n tensors; d = width
233
+ offsets = torch.zeros(len(rows) + 1, dtype=OFFSET_DTYPE) # (n + 1,)
234
+ blocks: list[Tensor] = []
235
+ for position, row in enumerate(rows):
236
+ block = row.detach().to("cpu") # (r_i, d)
237
+ if block.ndim != 2 or block.shape[1] != width:
238
+ raise ValueError(
239
+ f"A ragged feature row must have shape (r_i, {width}); row {position} has "
240
+ f"{tuple(block.shape)}."
241
+ )
242
+ offsets[position + 1] = int(offsets[position]) + block.shape[0]
243
+ blocks.append(block.to(dtype))
244
+ values = ( # (sum r_i, d)
245
+ torch.cat(blocks, dim=0) if blocks else torch.empty((0, width), dtype=dtype)
246
+ )
247
+ return {"offsets": offsets, "values": values} # {"offsets": (n + 1,), "values": (sum r_i, d)}
248
+
249
+
250
+ def validate_topk(indices: Tensor, values: Tensor, width: int, count: int) -> None:
251
+ """Reject malformed residue codes before writing and when reading committed tensors."""
252
+ # indices, values: (r, k), with k = count; r residues of one sequence, or of every sequence of a part
253
+ if indices.dtype not in (torch.int16, torch.int32, torch.int64):
254
+ raise ValueError("Top-k code indices must be signed integer tensors.")
255
+ if (indices.ndim != 2 or indices.shape[1] != count or values.shape != indices.shape):
256
+ raise ValueError(f"Top-k indices and values must both have shape (residues, {count}).")
257
+ if not values.is_floating_point() or not bool(torch.isfinite(values).all()):
258
+ raise ValueError("Top-k values must be finite floating point tensors.")
259
+ # Python integer bounds avoid narrowing 2**31 to int32 (or 16384 to int16).
260
+ if indices.numel() and (int(indices.min()) < 0 or int(indices.max()) >= width):
261
+ raise ValueError(f"Top-k code index is outside 0..{width - 1}.")
262
+ # Sorting validates uniqueness without changing the stored (residues,k) order.
263
+ ordered = indices.sort(dim=1).values # (r, k)
264
+ if bool((ordered[:, 1:] == ordered[:, :-1]).any()):
265
+ raise ValueError("Top-k indices must be unique within each residue.")
266
+
267
+
268
+ def encode_topk_rows(
269
+ rows: Sequence[TopKRow], width: int, count: int, dtype: torch.dtype,
270
+ ) -> dict[str, Tensor]:
271
+ """Encode sparse residue rows without allocating a residue-by-codebook tensor."""
272
+ # rows[i].indices, .values: (r_i, k), with k = count
273
+ offsets = torch.zeros(len(rows) + 1, dtype=OFFSET_DTYPE) # (n + 1,)
274
+ indices, values = [], []
275
+ for position, row in enumerate(rows):
276
+ if not isinstance(row, TopKRow):
277
+ raise TypeError("A ragged top-k feature requires TopKRow values.")
278
+ row_indices = row.indices.detach().to("cpu") # (r_i, k)
279
+ row_values = row.values.detach().to("cpu") # (r_i, k)
280
+ validate_topk(row_indices, row_values, width, count)
281
+ converted = row_values.to(dtype) # (r_i, k)
282
+ if not bool(torch.isfinite(converted).all()):
283
+ raise ValueError("Top-k value conversion exceeded the stored dtype range.")
284
+ offsets[position + 1] = int(offsets[position]) + row_values.shape[0]
285
+ indices.append(row_indices.to(INDEX_DTYPE))
286
+ values.append(converted)
287
+ return {
288
+ "offsets": offsets, # (n + 1,)
289
+ "indices": torch.cat(indices) if indices else torch.empty((0, count), dtype=INDEX_DTYPE),
290
+ "values": torch.cat(values) if values else torch.empty((0, count), dtype=dtype),
291
+ } # offsets (n + 1,); indices and values (sum r_i, k)
292
+
293
+
294
+ def row_count(layout: str, tensors: dict[str, Tensor]) -> int:
295
+ """How many sequences a segment's tensors hold."""
296
+
297
+ # tensors: (n, w) values for dense, (n + 1,) indptr or offsets otherwise
298
+ if layout == DENSE:
299
+ return int(tensors["values"].shape[0])
300
+ if layout == CSR:
301
+ return int(tensors["indptr"].shape[0]) - 1
302
+ if layout in (RAGGED, RAGGED_TOPK):
303
+ return int(tensors["offsets"].shape[0]) - 1
304
+ raise ValueError(f"Unknown feature layout {layout!r}; expected one of {list(LAYOUT_NAMES)}.")
305
+
306
+
307
+ def tensor_names(layout: str, *, positions: bool) -> tuple[str, ...]:
308
+ """The tensors a segment of this layout holds, in a stable order."""
309
+
310
+ if layout == DENSE:
311
+ return ("values",)
312
+ if layout == CSR:
313
+ return ("indptr", "indices", "values", "positions") if positions else (
314
+ "indptr", "indices", "values",
315
+ )
316
+ if layout == RAGGED:
317
+ return ("offsets", "values")
318
+ if layout == RAGGED_TOPK:
319
+ return ("offsets", "indices", "values")
320
+ raise ValueError(f"Unknown feature layout {layout!r}; expected one of {list(LAYOUT_NAMES)}.")
321
+
322
+
323
+ def _concatenate(parts: Sequence[Tensor], dtype: torch.dtype) -> Tensor:
324
+ # parts: (n_i, ...) each, with one shared trailing shape
325
+ if not parts:
326
+ return torch.empty(0, dtype=dtype) # (0,)
327
+ return torch.cat(list(parts), dim=0) # (sum n_i, ...)
328
+
329
+
330
+ __all__ = [
331
+ "CSR",
332
+ "DENSE",
333
+ "DTYPE_NAMES",
334
+ "INDEX_DTYPE",
335
+ "LAYOUT_NAMES",
336
+ "NAME_DTYPES",
337
+ "OFFSET_DTYPE",
338
+ "POSITION_DTYPE",
339
+ "RAGGED",
340
+ "RAGGED_TOPK",
341
+ "VALUE_DTYPES",
342
+ "SparseRow",
343
+ "TopKRow",
344
+ "dtype_name",
345
+ "encode_csr",
346
+ "encode_dense",
347
+ "encode_ragged",
348
+ "encode_topk_rows",
349
+ "row_count",
350
+ "row_tensor_bytes",
351
+ "tensor_names",
352
+ "validate_topk",
353
+ "value_dtype",
354
+ ]
fastplms/features/reader.py ADDED
@@ -0,0 +1,546 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Random-access reads of one feature, for loops that ask for a few rows at a time.
2
+
3
+ ``FeatureStore.read`` is the verified one-shot read. Each call opens a fresh index connection,
4
+ hashes every part the request touches, and loads each of those parts whole, so a training loop that
5
+ reads one batch per step pays for the whole store on every step. ``FeatureReader`` pays the same
6
+ verification once per part, then serves each row by a memory-mapped slice of the part file.
7
+
8
+ A reader gives the same rows as the store, in the same types, and refuses what the store refuses:
9
+ a missing sequence raises, changed part bytes raise on first access, and independent content pins
10
+ are honored. What it changes is the cost:
11
+ one read-only index connection per thread, one open handle per verified part, and one slice per row.
12
+
13
+ A reader belongs to a single feature directory and may be shared by threads. It pickles as the store
14
+ it reads, so a spawned worker reopens and re-verifies only the parts it touches. Call ``verify``
15
+ before forking to pay for verification once instead of once per worker.
16
+
17
+ A reader given a ``receipt`` path remembers its verifications across processes (see ``receipts``):
18
+ a part whose marker digests, size and modification time match the receipt is opened without hashing
19
+ or loading it, and every part the reader does verify is recorded. A pickled reader keeps its
20
+ receipt, so spawned workers skip the verification too. A store opened with content pins never trusts
21
+ a receipt.
22
+ """
23
+
24
+ from __future__ import annotations
25
+
26
+ import gzip
27
+ import json
28
+ import os
29
+ import sqlite3
30
+ import sys
31
+ import threading
32
+ import torch
33
+
34
+ from collections.abc import Iterable, Sequence
35
+ from concurrent.futures import ThreadPoolExecutor
36
+ from dataclasses import dataclass
37
+ from functools import partial
38
+ from pathlib import Path
39
+ from typing import Any, cast
40
+ from torch import Tensor
41
+ from tqdm import tqdm
42
+
43
+ from .digests import file_sha256
44
+ from .layouts import CSR, DENSE, RAGGED_TOPK, SparseRow, TopKRow
45
+ from .receipts import PartReceipt
46
+ from .store import (
47
+ COMMIT_FILE,
48
+ FEATURE_FILE,
49
+ INDEX_FILE,
50
+ PART_TEMPLATE,
51
+ SEGMENTS_DIRECTORY,
52
+ FeatureStore,
53
+ RowAddress,
54
+ StoredFeature,
55
+ sequence_digest,
56
+ )
57
+
58
+
59
+ _LOOKUP_CHUNK = 512
60
+
61
+
62
+ @dataclass(frozen=True, slots=True)
63
+ class CsrRows:
64
+ """Rows of a csr feature gathered into one compressed-sparse-row block, in the order asked.
65
+
66
+ ``indptr`` is ``(n + 1,)`` int64 and names each row's span of ``indices``, ``values`` and
67
+ ``positions``, which are ``(nnz,)``. Wrap them in whatever sparse type the caller uses: no
68
+ array library is imported here.
69
+ """
70
+
71
+ indptr: Tensor
72
+ indices: Tensor
73
+ values: Tensor
74
+ positions: Tensor | None
75
+
76
+
77
+ @dataclass(slots=True)
78
+ class _VerifiedPart:
79
+ """A part that passed verification: its open handle, and its row offsets in memory.
80
+
81
+ Offsets are ``(n + 1,)`` int64 and small next to the values they index, so they stay resident.
82
+ Values, indices and positions are read by slice from ``handle``.
83
+ """
84
+
85
+ handle: Any
86
+ offsets: Tensor | None
87
+
88
+
89
+ class FeatureReader:
90
+ """Verified random access to one feature's rows."""
91
+
92
+ def __init__(
93
+ self, store: FeatureStore, *, receipt: str | Path | None = None, trust_receipt: bool = True,
94
+ ) -> None:
95
+ self.store = store
96
+ self._receipt_path = None if receipt is None else Path(receipt)
97
+ self._trust_receipt = trust_receipt
98
+ # Pinned stores check each file against its pin, which a receipt cannot stand in for.
99
+ self._receipt = None
100
+ if self._receipt_path is not None and store._content_pins is None:
101
+ self._receipt = PartReceipt(
102
+ self._receipt_path, store.directory, store.spec.payload(), trust=trust_receipt,
103
+ )
104
+ self._lock = threading.RLock()
105
+ self._connections: dict[tuple[int, int], sqlite3.Connection] = {}
106
+ self._markers: dict[str, dict[str, Any]] = {}
107
+ self._parts: dict[tuple[str, int], _VerifiedPart] = {}
108
+ # Different parts verify concurrently; one part verifies only once.
109
+ self._part_locks: dict[tuple[str, int], threading.Lock] = {}
110
+ self._closed = False
111
+
112
+ @classmethod
113
+ def open(
114
+ cls, directory: str | Path, *, content_pins: dict[str, str] | None = None,
115
+ ) -> FeatureReader:
116
+ """A reader over an existing feature directory, which it never creates or repairs."""
117
+
118
+ return cls(FeatureStore.read_only(directory, content_pins=content_pins))
119
+
120
+ def __reduce__(self) -> tuple[Any, tuple[FeatureStore]]:
121
+ restore = partial(
122
+ FeatureReader, receipt=self._receipt_path, trust_receipt=self._trust_receipt,
123
+ )
124
+ return (restore, (self.store,))
125
+
126
+ def __enter__(self) -> FeatureReader:
127
+ return self
128
+
129
+ def __exit__(self, *exception: object) -> None:
130
+ self.close()
131
+
132
+ @property
133
+ def spec(self) -> StoredFeature:
134
+ return self.store.spec
135
+
136
+ def close(self) -> None:
137
+ """Save the receipt, then release every connection and part handle; reads then raise."""
138
+
139
+ if self._receipt is not None:
140
+ self._receipt.save()
141
+ with self._lock:
142
+ self._closed = True
143
+ for connection in self._connections.values():
144
+ connection.close()
145
+ self._connections.clear()
146
+ self._parts.clear()
147
+
148
+ # Membership --------------------------------------------------------------
149
+
150
+ def __len__(self) -> int:
151
+ return int(self._connection().execute("SELECT count(*) FROM rows").fetchone()[0])
152
+
153
+ def __contains__(self, sequence: str) -> bool:
154
+ digest = sequence_digest(sequence)
155
+ return digest in self._found([digest])
156
+
157
+ def missing(self, sequences: Iterable[str]) -> tuple[str, ...]:
158
+ """The sequences this feature lacks, in the order given, without repeats."""
159
+
160
+ wanted: dict[str, str] = {}
161
+ for sequence in sequences:
162
+ wanted.setdefault(sequence_digest(sequence), sequence)
163
+ present = self._found(list(wanted))
164
+ return tuple(sequence for digest, sequence in wanted.items() if digest not in present)
165
+
166
+ # Reading -----------------------------------------------------------------
167
+
168
+ def verify(
169
+ self, sequences: Sequence[str] | None = None, *, workers: int = 1,
170
+ progress: str | None = None,
171
+ ) -> int:
172
+ """Verify the parts holding these sequences, or every committed part, and count them.
173
+
174
+ Reading verifies lazily, so this is only for paying the cost at a moment of the caller's
175
+ choosing, such as before a data loader forks its workers. Verifying a part hashes its
176
+ bytes and loads its tensors, so ``workers`` threads verify that many parts at a time
177
+ (hashing and the loads release the interpreter lock). The first part to fail raises,
178
+ and parts not yet started are not verified. A part the receipt vouches for is opened
179
+ without hashing; the receipt is saved when the pass ends. ``progress`` names a progress
180
+ bar on stderr, one tick per part.
181
+ """
182
+
183
+ if workers < 1:
184
+ raise ValueError("verify needs at least one worker.")
185
+ if sequences is not None:
186
+ wanted = {(address.segment, address.part) for _, address in self._resolved(sequences)}
187
+ else:
188
+ wanted = {
189
+ (fingerprint, int(part["part"]))
190
+ for fingerprint in self._segment_names()
191
+ for part in self._marker(fingerprint)["parts"]
192
+ }
193
+ ordered = sorted(wanted)
194
+ bar = tqdm(
195
+ total=len(ordered), desc=progress, unit="part", file=sys.stderr, mininterval=10.0,
196
+ disable=progress is None,
197
+ )
198
+
199
+ def verified(segment: str, number: int) -> None:
200
+ self._part(segment, number)
201
+ bar.update()
202
+
203
+ try:
204
+ if workers == 1 or len(ordered) < 2:
205
+ for segment, number in ordered:
206
+ verified(segment, number)
207
+ return len(wanted)
208
+ with ThreadPoolExecutor(max_workers=workers) as pool:
209
+ futures = [pool.submit(verified, segment, number) for segment, number in ordered]
210
+ try:
211
+ for future in futures:
212
+ future.result()
213
+ except BaseException:
214
+ for future in futures:
215
+ future.cancel()
216
+ raise
217
+ return len(wanted)
218
+ finally:
219
+ bar.close()
220
+ if self._receipt is not None:
221
+ self._receipt.save()
222
+
223
+ def addresses(self, sequences: Sequence[str]) -> list[RowAddress]:
224
+ """Each sequence's address, in the order asked, from the index and its commit markers alone.
225
+
226
+ This reads no part, so it costs index lookups, not bytes. It raises ``KeyError`` for a
227
+ sequence the feature lacks and ``ValueError`` when the index and the commit marker disagree
228
+ about a row. Nothing here verifies the part holding the row: ``verify`` or a read does.
229
+ """
230
+
231
+ return [address for _, address in self._resolved(sequences)]
232
+
233
+ def content_pins(self, sequences: Sequence[str], *, workers: int = 1) -> dict[str, str]:
234
+ """Pin freshly hashed selection bytes after reusing this reader's layout verification.
235
+
236
+ Every selected file is hashed again, including when a verification receipt was trusted.
237
+ Cached part handles avoid a second whole-part tensor load. Marker changes and even
238
+ data rewrites preserving size and modification time fail before these pins are returned.
239
+ """
240
+ self.verify(sequences, workers=workers)
241
+ feature = self.store.directory / FEATURE_FILE
242
+ if StoredFeature.from_payload(json.loads(feature.read_text(encoding="utf-8"))) != self.spec:
243
+ raise ValueError("Feature descriptor changed while pinning a selection.")
244
+ pins = {FEATURE_FILE: file_sha256(feature)}
245
+ selected = {(address.segment, address.part) for _, address in self._resolved(sequences)}
246
+ markers: set[str] = set()
247
+ for segment, number in sorted(selected):
248
+ payload = self._marker(segment)
249
+ prefix = f"{SEGMENTS_DIRECTORY}/{segment}"
250
+ if segment not in markers:
251
+ marker = self.store.directory / prefix / COMMIT_FILE
252
+ if json.loads(marker.read_text(encoding="utf-8")) != payload:
253
+ raise ValueError("Feature commit marker changed while pinning a selection.")
254
+ pins[f"{prefix}/{COMMIT_FILE}"] = file_sha256(marker)
255
+ markers.add(segment)
256
+ part = payload["parts"][number]
257
+ pins[f"{prefix}/{PART_TEMPLATE.format(number)}"] = part["sha256"]
258
+ sidecar = part.get("row_metadata")
259
+ if sidecar is not None:
260
+ pins[f"{prefix}/{sidecar['file']}"] = sidecar["sha256"]
261
+
262
+ def check(relative: str) -> None:
263
+ actual = file_sha256(self.store.directory / relative)
264
+ if actual != pins[relative]:
265
+ raise ValueError("Feature bytes changed while pinning a selection.")
266
+ self.store._check_pin(relative, actual)
267
+
268
+ if workers == 1:
269
+ for relative in pins:
270
+ check(relative)
271
+ else:
272
+ with ThreadPoolExecutor(max_workers=workers) as pool:
273
+ list(pool.map(check, pins))
274
+ return pins
275
+
276
+ def residue_counts(self, sequences: Sequence[str]) -> list[int]:
277
+ """Each sequence's stored residue count, which a length-bucketed reader batches by."""
278
+
279
+ return [address.residues for _, address in self._resolved(sequences, verify=True)]
280
+
281
+ def row_metadata(self, sequences: Sequence[str]) -> list[dict[str, Any]]:
282
+ """Return row identities using the same immutable-part verification as feature reads."""
283
+ identities = []
284
+ parts: dict[tuple[str, int], list[dict[str, Any]]] = {}
285
+ for _, address in self._resolved(sequences):
286
+ self._part(address.segment, address.part)
287
+ key = (address.segment, address.part)
288
+ if key not in parts:
289
+ committed = self._marker(address.segment)["parts"][address.part]
290
+ sidecar = committed.get("row_metadata")
291
+ if not isinstance(sidecar, dict):
292
+ raise ValueError(
293
+ "Stored row identities are unavailable; recompute this legacy feature."
294
+ )
295
+ path = self.store.directory / SEGMENTS_DIRECTORY / address.segment / sidecar["file"]
296
+ records = json.loads(gzip.decompress(path.read_bytes()))
297
+ if (not isinstance(records, list) or len(records) != len(committed["digests"])
298
+ or not all(isinstance(record, dict) for record in records)):
299
+ raise ValueError("Stored row metadata does not match the committed part rows.")
300
+ parts[key] = records
301
+ identities.append(dict(parts[key][address.row]))
302
+ return identities
303
+
304
+ def read(self, sequences: Sequence[str]) -> list[Tensor]:
305
+ """Each sequence's feature as a tensor, in the order asked, as ``FeatureStore.read`` gives.
306
+
307
+ A dense feature gives ``(w,)``, a ragged one ``(r_i, d)``, and a csr one a densified
308
+ ``(w,)`` row. ``ragged_topk`` requires ``read_topk``.
309
+ """
310
+
311
+ spec = self.spec
312
+ if spec.layout == CSR:
313
+ # Returns b vectors, each (w,).
314
+ return [row.to_dense(spec.width) for row in self.read_sparse(sequences)] # (w,) per row
315
+ if spec.layout == RAGGED_TOPK:
316
+ raise ValueError(
317
+ "Use read_topk for sparse residue codes; implicit densification is refused."
318
+ )
319
+ rows: list[Tensor] = []
320
+ parts: dict[tuple[str, int], tuple[_VerifiedPart, Tensor]] = {}
321
+ for _, address in self._resolved(sequences):
322
+ key = (address.segment, address.part)
323
+ if key not in parts:
324
+ part = self._part(*key)
325
+ # One mapped tensor per touched part, not one safetensors wrapper per row.
326
+ parts[key] = (part, part.handle.get_tensor("values"))
327
+ part, values = parts[key] # (n_part, w) dense or (r_part, d) ragged
328
+ if spec.layout == DENSE:
329
+ rows.append(values[address.row].clone()) # (w,)
330
+ else:
331
+ start, stop = self._span(part, address)
332
+ rows.append(values[start:stop].clone()) # (r_i, d)
333
+ return rows # b tensors, each (w,) for dense or (r_i, d) for ragged
334
+
335
+ def read_sparse(self, sequences: Sequence[str]) -> list[SparseRow]:
336
+ """Each sequence's compressed row, in the order asked. Only for a csr feature."""
337
+
338
+ if self.spec.layout != CSR:
339
+ raise ValueError(f"read_sparse needs a csr feature; this one is {self.spec.layout}.")
340
+ rows: list[SparseRow] = []
341
+ parts: dict[tuple[str, int], tuple[_VerifiedPart, Tensor, Tensor, Tensor | None]] = {}
342
+ for _, address in self._resolved(sequences):
343
+ key = (address.segment, address.part)
344
+ if key not in parts:
345
+ part = self._part(*key)
346
+ parts[key] = (
347
+ part, part.handle.get_tensor("indices"), part.handle.get_tensor("values"),
348
+ part.handle.get_tensor("positions") if self.spec.positions else None,
349
+ )
350
+ part, indices, values, positions = parts[key] # (nnz_part,) per tensor
351
+ start, stop = self._span(part, address)
352
+ rows.append(SparseRow(
353
+ indices=indices[start:stop].clone(), # (nnz_i,)
354
+ values=values[start:stop].clone(), # (nnz_i,)
355
+ positions=(
356
+ positions[start:stop].clone() if positions is not None else None # (nnz_i,)
357
+ ),
358
+ ))
359
+ return rows
360
+
361
+ def read_csr(self, sequences: Sequence[str]) -> CsrRows:
362
+ """These sequences' rows as one compressed-sparse-row block, in the order asked.
363
+
364
+ A sequence asked for twice appears twice. This is the read for a caller that wants a matrix,
365
+ such as a design matrix for a gradient-boosted model, instead of one row at a time.
366
+ """
367
+
368
+ spec = self.spec
369
+ rows = self.read_sparse(sequences)
370
+ counts = torch.tensor([row.indices.numel() for row in rows], dtype=torch.int64) # (n,)
371
+ indptr = torch.zeros(len(rows) + 1, dtype=torch.int64) # (n + 1,)
372
+ torch.cumsum(counts, dim=0, out=indptr[1:])
373
+
374
+ def joined(tensors: list[Tensor], dtype: torch.dtype) -> Tensor:
375
+ # tensors: (nnz_i,) per input vector.
376
+ return torch.cat(tensors) if tensors else torch.empty(0, dtype=dtype) # (nnz,)
377
+
378
+ return CsrRows(
379
+ indptr=indptr,
380
+ indices=joined([row.indices for row in rows], torch.int32), # (nnz,)
381
+ values=joined([row.values for row in rows], spec.dtype), # (nnz,)
382
+ positions=(
383
+ joined([cast(Tensor, row.positions) for row in rows], torch.int16) # (nnz,)
384
+ if spec.positions else None
385
+ ),
386
+ )
387
+
388
+ def read_topk(self, sequences: Sequence[str]) -> list[TopKRow]:
389
+ """Each sequence's sparse ``(r_i, k)`` residue codes, in the order asked."""
390
+
391
+ if self.spec.layout != RAGGED_TOPK:
392
+ raise ValueError(
393
+ f"read_topk needs a ragged_topk feature; this one is {self.spec.layout}."
394
+ )
395
+ rows: list[TopKRow] = []
396
+ for _, address in self._resolved(sequences):
397
+ part = self._part(address.segment, address.part)
398
+ start, stop = self._span(part, address)
399
+ rows.append(TopKRow(
400
+ cast(Tensor, part.handle.get_slice("indices")[start:stop]), # (r_i, k)
401
+ cast(Tensor, part.handle.get_slice("values")[start:stop]), # (r_i, k)
402
+ ))
403
+ return rows
404
+
405
+ # Internals ---------------------------------------------------------------
406
+
407
+ def _connection(self) -> sqlite3.Connection:
408
+ """This thread's read-only index connection, opened on first use.
409
+
410
+ Connections are keyed by process and thread, so a forked worker or a prefetch thread opens
411
+ its own and never touches one that another owner holds.
412
+ """
413
+
414
+ owner = (os.getpid(), threading.get_ident())
415
+ with self._lock:
416
+ if self._closed:
417
+ raise ValueError("This feature reader is closed.")
418
+ connection = self._connections.get(owner)
419
+ if connection is None:
420
+ database = (self.store.directory / INDEX_FILE).resolve()
421
+ connection = sqlite3.connect(
422
+ database.as_uri() + "?mode=ro", uri=True, timeout=30,
423
+ check_same_thread=False,
424
+ )
425
+ self._connections[owner] = connection
426
+ return connection
427
+
428
+ def _found(self, digests: Sequence[str]) -> dict[str, RowAddress]:
429
+ found: dict[str, RowAddress] = {}
430
+ connection = self._connection()
431
+ for start in range(0, len(digests), _LOOKUP_CHUNK):
432
+ chunk = list(digests[start:start + _LOOKUP_CHUNK])
433
+ marks = ",".join("?" * len(chunk))
434
+ # fetchall ends the read transaction so writers can commit.
435
+ for digest, segment, part, row, residues in connection.execute(
436
+ f"SELECT digest, segment, part, row, residues FROM rows WHERE digest IN ({marks})",
437
+ chunk,
438
+ ).fetchall():
439
+ found[digest] = RowAddress(segment, part, row, residues)
440
+ return found
441
+
442
+ def _resolved(
443
+ self, sequences: Sequence[str], *, verify: bool = False,
444
+ ) -> list[tuple[str, RowAddress]]:
445
+ """Each sequence with its address, checked against the commit marker that owns it."""
446
+
447
+ self.store._check_pin(FEATURE_FILE)
448
+ digests = [sequence_digest(sequence) for sequence in sequences]
449
+ found = self._found(list(dict.fromkeys(digests)))
450
+ resolved: list[tuple[str, RowAddress]] = []
451
+ for sequence, digest in zip(sequences, digests, strict=True):
452
+ address = found.get(digest)
453
+ if address is None:
454
+ raise KeyError(
455
+ f"Feature {self.spec.key!r} has no row for a sequence of "
456
+ f"{len(sequence)} residues (sha256 {digest[:12]}). "
457
+ "Embed it first, or call missing() before reading."
458
+ )
459
+ payload = self._marker(address.segment)
460
+ if not 0 <= address.part < len(payload["parts"]):
461
+ raise ValueError("Index part does not name one committed feature part.")
462
+ part = payload["parts"][address.part]
463
+ if (not 0 <= address.row < len(part["digests"])
464
+ or part["digests"][address.row] != digest
465
+ or part["residues"][address.row] != address.residues):
466
+ raise ValueError("Feature index and committed sequence row disagree.")
467
+ if verify:
468
+ self._part(address.segment, address.part)
469
+ resolved.append((sequence, address))
470
+ return resolved
471
+
472
+ def _marker(self, segment: str) -> dict[str, Any]:
473
+ """A validated commit marker, read once per reader."""
474
+
475
+ with self._lock:
476
+ payload = self._markers.get(segment)
477
+ if payload is None:
478
+ payload = self.store._segment_payload(segment)
479
+ self._markers[segment] = payload
480
+ return payload
481
+
482
+ def _segment_names(self) -> list[str]:
483
+ directory = self.store.directory / SEGMENTS_DIRECTORY
484
+ return sorted(marker.parent.name for marker in directory.glob(f"*/{COMMIT_FILE}"))
485
+
486
+ def _part(self, segment: str, number: int) -> _VerifiedPart:
487
+ """A part's open handle, verified the first time and kept.
488
+
489
+ Verification is the store's own: checksum against the marker, physical layout, and the row
490
+ identity sidecar. The tensors it loads are dropped afterwards except the offsets.
491
+ """
492
+
493
+ key = (segment, number)
494
+ with self._lock:
495
+ if self._closed:
496
+ raise ValueError("This feature reader is closed.")
497
+ opened = self._parts.get(key)
498
+ if opened is not None:
499
+ return opened
500
+ part_lock = self._part_locks.setdefault(key, threading.Lock())
501
+ with part_lock:
502
+ # Verify outside the reader lock so other parts can progress.
503
+ with self._lock:
504
+ opened = self._parts.get(key)
505
+ if opened is not None:
506
+ return opened
507
+ from safetensors import safe_open
508
+
509
+ committed = self._marker(segment)["parts"][number]
510
+ path = str(self.store._part_path(segment, number))
511
+ offsets_name = "indptr" if self.spec.layout == CSR else "offsets"
512
+ receipt = self._receipt
513
+ described = None if receipt is None else receipt.describe(segment, number, committed)
514
+ if (receipt is not None and described is not None
515
+ and receipt.trusts(segment, number, described)):
516
+ # An earlier full verification passed these digests on files of this size and time.
517
+ handle = safe_open(path, framework="pt", device="cpu")
518
+ names = set(handle.keys())
519
+ offsets = handle.get_tensor(offsets_name) if offsets_name in names else None
520
+ else:
521
+ tensors = self.store._verified_part(segment, committed)
522
+ offsets = tensors.get(offsets_name)
523
+ del tensors
524
+ handle = safe_open(path, framework="pt", device="cpu")
525
+ if receipt is not None and described is not None:
526
+ if receipt.describe(segment, number, committed) != described:
527
+ raise ValueError(
528
+ f"A feature part changed during verification: {segment}/{number}."
529
+ )
530
+ receipt.record(segment, number, described)
531
+ opened = _VerifiedPart(handle=handle, offsets=offsets)
532
+ with self._lock:
533
+ if self._closed:
534
+ raise ValueError("This feature reader is closed.")
535
+ self._parts[key] = opened
536
+ return opened
537
+
538
+ @staticmethod
539
+ def _span(part: _VerifiedPart, address: RowAddress) -> tuple[int, int]:
540
+ """A row's ``[start, stop)`` span in its part's flattened values."""
541
+
542
+ assert part.offsets is not None, "Only sparse and ragged layouts carry row offsets."
543
+ return int(part.offsets[address.row]), int(part.offsets[address.row + 1])
544
+
545
+
546
+ __all__ = ["FeatureReader"]
fastplms/features/receipts.py ADDED
@@ -0,0 +1,125 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """A record of which feature parts a full verification has passed, so a later reader can skip it.
2
+
3
+ Verifying a part hashes its bytes and loads its tensors, which costs one pass over the file. A store
4
+ of hundreds of gigabytes read by many processes would otherwise pay that pass in every process, on
5
+ every open. A receipt keeps the result for one feature directory: for each part, the digests the
6
+ commit marker states, and the size and modification time of the part file and of its row-identity
7
+ sidecar when they were verified. A reader trusts a part only when the marker states the same digests
8
+ and both files still have the recorded size and modification time. Any other part is verified in
9
+ full.
10
+
11
+ A receipt is a cache of one machine's verification, never a proof: a file rewritten to the same size
12
+ and the same modification time passes it. ``FeatureReader.verify`` on a reader built with
13
+ ``trust_receipt=False`` hashes every byte and refreshes the receipt, and is the check for that case.
14
+ """
15
+
16
+ from __future__ import annotations
17
+
18
+ import json
19
+ import os
20
+ import threading
21
+ import warnings
22
+
23
+ from collections.abc import Mapping
24
+ from contextlib import suppress
25
+ from pathlib import Path
26
+ from typing import Any
27
+
28
+ from .digests import json_sha256
29
+ from .store import PART_TEMPLATE, SEGMENTS_DIRECTORY
30
+
31
+
32
+ SCHEMA = "feature_part_receipt_v1"
33
+
34
+
35
+ class PartReceipt:
36
+ """The verified parts of one feature directory, read from and saved to one JSON file."""
37
+
38
+ def __init__(
39
+ self, path: str | Path, directory: Path, spec_payload: Mapping[str, Any],
40
+ *, trust: bool = True,
41
+ ) -> None:
42
+ self.path = Path(path)
43
+ self._directory = directory
44
+ self._feature = spec_payload.get("key")
45
+ self._descriptor_sha256 = json_sha256(spec_payload)
46
+ self._trusted = self._load() if trust else {}
47
+ self._pending: dict[str, dict[str, Any]] = {}
48
+ self._lock = threading.Lock()
49
+
50
+ def _load(self) -> dict[str, dict[str, Any]]:
51
+ """The recorded parts, or none without a receipt or when it names another descriptor."""
52
+ try:
53
+ document = json.loads(self.path.read_text(encoding="utf-8"))
54
+ except FileNotFoundError:
55
+ return {}
56
+ except (OSError, ValueError) as error:
57
+ warnings.warn(
58
+ f"Ignoring the unreadable verification receipt {self.path}: {error}",
59
+ RuntimeWarning, stacklevel=3,
60
+ )
61
+ return {}
62
+ if (not isinstance(document, dict) or document.get("schema") != SCHEMA
63
+ or document.get("descriptor_sha256") != self._descriptor_sha256
64
+ or not isinstance(document.get("parts"), dict)):
65
+ return {}
66
+ return document["parts"]
67
+
68
+ def describe(self, segment: str, number: int, committed: Mapping[str, Any]) -> dict[str, Any]:
69
+ """What verifying this part vouches for: the marker's digests and the files' stat."""
70
+ base = self._directory / SEGMENTS_DIRECTORY / segment
71
+ part = (base / PART_TEMPLATE.format(number)).stat()
72
+ sidecar = committed.get("row_metadata")
73
+ rows = None
74
+ if isinstance(sidecar, dict):
75
+ stat = (base / str(sidecar["file"])).stat()
76
+ rows = {
77
+ "file": sidecar["file"], "sha256": sidecar["sha256"],
78
+ "size": stat.st_size, "mtime_ns": stat.st_mtime_ns,
79
+ }
80
+ return {
81
+ "sha256": committed["sha256"], "size": part.st_size, "mtime_ns": part.st_mtime_ns,
82
+ "rows": rows,
83
+ }
84
+
85
+ def trusts(self, segment: str, number: int, described: Mapping[str, Any]) -> bool:
86
+ """Whether an earlier verification passed these digests, sizes and modification times."""
87
+ return self._trusted.get(f"{segment}/{number}") == described
88
+
89
+ def record(self, segment: str, number: int, described: Mapping[str, Any]) -> None:
90
+ with self._lock:
91
+ self._pending[f"{segment}/{number}"] = dict(described)
92
+
93
+ def save(self) -> None:
94
+ """Merge the parts verified since the last save into the file, atomically.
95
+
96
+ Concurrent savers can drop each other's entries, which only costs a later verification. A
97
+ location that cannot be written warns, and this process keeps what it verified in memory.
98
+ """
99
+ with self._lock:
100
+ if not self._pending:
101
+ return
102
+ document = {
103
+ "schema": SCHEMA, "feature": self._feature,
104
+ "descriptor_sha256": self._descriptor_sha256,
105
+ "parts": {**self._load(), **self._pending},
106
+ }
107
+ owner = f"{os.getpid()}.{threading.get_ident()}"
108
+ temporary = self.path.with_name(f"{self.path.name}.{owner}.writing")
109
+ try:
110
+ self.path.parent.mkdir(parents=True, exist_ok=True)
111
+ temporary.write_text(json.dumps(document, sort_keys=True), encoding="utf-8")
112
+ temporary.replace(self.path)
113
+ except OSError as error:
114
+ with suppress(OSError):
115
+ temporary.unlink(missing_ok=True)
116
+ warnings.warn(
117
+ f"Could not save the verification receipt {self.path}: {error}. "
118
+ "Every open of this feature verifies its parts again.",
119
+ RuntimeWarning, stacklevel=2,
120
+ )
121
+ self._trusted = {**self._trusted, **self._pending}
122
+ self._pending.clear()
123
+
124
+
125
+ __all__ = ["SCHEMA", "PartReceipt"]
fastplms/features/store.py ADDED
@@ -0,0 +1,1567 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """The FastPLMs feature store: one directory per feature key, read by sequence.
2
+
3
+ A store answers one question: what is this feature of this sequence, and does it already exist?
4
+ That makes a run embed only what is missing, and lets one embedding pass serve every head trained
5
+ afterwards.
6
+
7
+ Layout on disk, under one root that holds many features::
8
+
9
+ <root>/<key>/
10
+ feature.json the key's descriptor, layout, width, and dtype
11
+ index.sqlite sequence sha256 -> segment, part, row, residue count
12
+ segments/<fingerprint>/
13
+ part-00000.safetensors immutable, memory-mappable, one per append
14
+ part-00001.safetensors
15
+ run.json the commit marker, written last
16
+
17
+ **A segment is immutable and committed once.** Parts appear as a run streams, and nothing reads
18
+ them until ``run.json`` names them; a run that dies leaves a directory the index ignores and
19
+ ``sweep`` removes. Nothing is ever rewritten, so a reader never sees a half-written feature and two
20
+ runs never race over one file.
21
+
22
+ **The index is a cache of the commit markers, not the record.** Every indexed row can be rebuilt
23
+ from the committed segments, which ``reindex`` does, so a lost or corrupt ``index.sqlite`` costs a
24
+ scan rather than the features.
25
+
26
+ **A sequence is identified by the SHA-256 of its exact UTF-8 bytes**, so two callers agree without
27
+ coordinating, and a sequence that differs by one residue is a different row. The store keeps the
28
+ digest and the residue count, never the sequence text: a caller that has the sequences can always
29
+ recompute the digest, and storing millions of them again would cost more than the features.
30
+
31
+ The key belongs to ``foundry.embedding.feature_key``, which composes the model, its revision, the
32
+ sparse autoencoder, the layer, the pooling, the dtype, and the residue limit into one filename-safe
33
+ name. This module never invents a key; it stores the name and the descriptor it is given and
34
+ refuses a second, different descriptor under the same name. FastPLMs ships without foundry, so the
35
+ store takes the name as a string rather than importing the key.
36
+ """
37
+
38
+ from __future__ import annotations
39
+
40
+ import gzip
41
+ import hashlib
42
+ import json
43
+ import os
44
+ import re
45
+ import sqlite3
46
+ import threading
47
+ import torch
48
+
49
+ from collections.abc import Callable, Iterable, Iterator, Mapping, Sequence
50
+ from contextlib import AbstractContextManager, ExitStack, contextmanager
51
+ from dataclasses import dataclass, field
52
+ from datetime import UTC, datetime
53
+ from itertools import pairwise
54
+ from pathlib import Path
55
+ from typing import Any, cast
56
+ from torch import Tensor
57
+
58
+ from .digests import file_sha256, json_sha256
59
+ from .json_files import indented_json
60
+ from .layouts import (
61
+ CSR,
62
+ DENSE,
63
+ LAYOUT_NAMES,
64
+ RAGGED,
65
+ RAGGED_TOPK,
66
+ SparseRow,
67
+ TopKRow,
68
+ dtype_name,
69
+ encode_csr,
70
+ encode_dense,
71
+ encode_ragged,
72
+ encode_topk_rows,
73
+ row_count,
74
+ row_tensor_bytes,
75
+ tensor_names,
76
+ validate_topk,
77
+ value_dtype,
78
+ )
79
+ from .transactions import file_lock, flush_and_evict, publish_file, sync_directory
80
+
81
+
82
+ FORMAT = "fastplms-feature-store-v1"
83
+ FEATURE_FILE = "feature.json"
84
+ INDEX_FILE = "index.sqlite"
85
+ SEGMENTS_DIRECTORY = "segments"
86
+ COMMIT_FILE = "run.json"
87
+ PART_TEMPLATE = "part-{:05d}.safetensors"
88
+
89
+ # Descriptor schemas that carry a complete scientific contract: v1 keeps residues only, v2 keeps CLS and EOS too,
90
+ # and v3 keeps v2's rows computed under a pinned embedding profile.
91
+ COMPLETE_SCHEMAS = frozenset({"feature_spec_v1", "feature_spec_v2", "feature_spec_v3"})
92
+
93
+ _NAME = re.compile(r"[A-Za-z0-9][A-Za-z0-9._-]{0,255}")
94
+ _PART = re.compile(r"part-(\d{5})\.safetensors")
95
+ _SHA256 = re.compile(r"[0-9a-f]{64}")
96
+ _WRITER_LEASE = object()
97
+
98
+
99
+ def sequence_digest(sequence: str) -> str:
100
+ """The SHA-256 of a sequence's exact UTF-8 bytes, which is its row identity."""
101
+
102
+ if not isinstance(sequence, str) or not sequence:
103
+ raise ValueError("A sequence must be a non-empty string.")
104
+ return hashlib.sha256(sequence.encode("utf-8")).hexdigest()
105
+
106
+
107
+ def partition_sequences(sequences: Iterable[str], *, shard: int, shards: int) -> tuple[str, ...]:
108
+ """Assign exact unique sequences by SHA-256 modulo shard count, preserving input order.
109
+
110
+ Every worker must receive the same input inventory and shard count. Assignment does not
111
+ depend on Python's randomized hash or process/device identity.
112
+ """
113
+ if type(shards) is not int or shards < 1 or type(shard) is not int or not 0 <= shard < shards:
114
+ raise ValueError("Require a positive shard count and 0 <= shard < shards.")
115
+ unique = dict.fromkeys(sequences)
116
+ return tuple(
117
+ sequence for sequence in unique if int(sequence_digest(sequence), 16) % shards == shard
118
+ )
119
+
120
+
121
+ @dataclass(frozen=True, slots=True)
122
+ class StoredFeature:
123
+ """What one feature is: its key, how its rows are laid out, and what they hold.
124
+
125
+ ``width`` is the vector width for ``dense``, the codebook size for ``csr``, and the hidden
126
+ width for ``ragged``. For ``ragged_topk``, ``width`` is the codebook size and ``sparse_count``
127
+ is the number of retained codes per residue. ``positions`` says whether ``csr`` rows carry
128
+ each entry's argmax residue. ``descriptor`` is the plain data the key was computed from,
129
+ so a store can say what it holds without the caller that made it.
130
+ """
131
+
132
+ key: str
133
+ layout: str
134
+ width: int
135
+ dtype: torch.dtype
136
+ positions: bool = False
137
+ descriptor: Mapping[str, Any] = field(default_factory=dict)
138
+ sparse_count: int | None = None
139
+
140
+ def __post_init__(self) -> None:
141
+ if not isinstance(self.key, str) or not _NAME.fullmatch(self.key):
142
+ raise ValueError(
143
+ "A feature key must be a filename-safe name, as "
144
+ "foundry.embedding.feature_key(...).name returns; received "
145
+ f"{self.key!r}."
146
+ )
147
+ if self.layout not in LAYOUT_NAMES:
148
+ raise ValueError(
149
+ f"layout must be one of {list(LAYOUT_NAMES)}; received {self.layout!r}."
150
+ )
151
+ if not isinstance(self.width, int) or isinstance(self.width, bool) or self.width <= 0:
152
+ raise ValueError(f"width must be a positive integer; received {self.width!r}.")
153
+ dtype_name(self.dtype)
154
+ if self.positions and self.layout != CSR:
155
+ raise ValueError("Only a csr feature stores argmax positions.")
156
+ if self.layout == RAGGED_TOPK:
157
+ if type(self.sparse_count) is not int or not 1 <= self.sparse_count <= self.width:
158
+ raise ValueError(
159
+ "A ragged top-k feature requires integer sparse_count in 1..width."
160
+ )
161
+ if self.width > torch.iinfo(torch.int32).max + 1:
162
+ raise ValueError("Top-k codebook exceeds the stored int32 index range.")
163
+ elif self.sparse_count is not None:
164
+ raise ValueError("Only a ragged top-k feature declares sparse_count.")
165
+ if not isinstance(self.descriptor, Mapping):
166
+ raise TypeError("descriptor must be a mapping of plain data.")
167
+ object.__setattr__(self, "descriptor", json.loads(json.dumps(dict(self.descriptor))))
168
+
169
+ def payload(self) -> dict[str, Any]:
170
+ """``feature.json``'s content."""
171
+
172
+ payload = {
173
+ "format": FORMAT,
174
+ "key": self.key,
175
+ "layout": self.layout,
176
+ "width": self.width,
177
+ "dtype": dtype_name(self.dtype),
178
+ "positions": self.positions,
179
+ "descriptor": dict(self.descriptor),
180
+ }
181
+ if self.sparse_count is not None:
182
+ payload["sparse_count"] = self.sparse_count
183
+ return payload
184
+
185
+ @classmethod
186
+ def from_payload(cls, payload: Mapping[str, Any]) -> StoredFeature:
187
+ """The spec a ``feature.json`` describes."""
188
+
189
+ if payload.get("format") != FORMAT:
190
+ raise ValueError(
191
+ f"Not a {FORMAT} feature directory; its format is {payload.get('format')!r}."
192
+ )
193
+ return cls(
194
+ key=str(payload["key"]),
195
+ layout=str(payload["layout"]),
196
+ width=int(payload["width"]),
197
+ dtype=value_dtype(str(payload["dtype"])),
198
+ positions=bool(payload.get("positions", False)),
199
+ descriptor=payload.get("descriptor") or {},
200
+ sparse_count=payload.get("sparse_count"),
201
+ )
202
+
203
+
204
+ @dataclass(frozen=True, slots=True)
205
+ class RowAddress:
206
+ """Where one sequence's row sits: which segment, which part, which row of it."""
207
+
208
+ segment: str
209
+ part: int
210
+ row: int
211
+ residues: int
212
+
213
+
214
+ @dataclass(frozen=True, slots=True)
215
+ class SegmentReceipt:
216
+ """What one committed segment holds, as ``run.json`` records it."""
217
+
218
+ fingerprint: str
219
+ parts: tuple[int, ...]
220
+ rows: int
221
+ committed_at: str
222
+ metadata: Mapping[str, Any]
223
+
224
+
225
+ class FeatureStore:
226
+ """One feature's directory, opened for reading and appending."""
227
+
228
+ def __init__(
229
+ self, directory: Path, spec: StoredFeature, *, read_only: bool = False,
230
+ content_pins: Mapping[str, str] | None = None, deep_verify: bool = True,
231
+ partial: bool = False,
232
+ ) -> None:
233
+ self.directory = directory.resolve()
234
+ self.spec = spec
235
+ self._read_only = read_only
236
+ # Deep verification re-reads every committed part whenever the index is recovered, which opening
237
+ # and every commit do. A store of terabytes opens with deep_verify=False: only segments the index
238
+ # lacks are read and indexed, and each part is still verified when a reader first touches it.
239
+ self._deep_verify = deep_verify
240
+ # A partial copy holds some of each segment's committed parts, as a fetch of a few rows from a
241
+ # published store leaves it. Its index lists the rows of the parts present, so a row of an absent
242
+ # part reads as missing, a part fetched later is indexed on the next open, and it takes no new
243
+ # segment. Absent parts cannot be verified, so it opens without deep verification.
244
+ if partial and deep_verify:
245
+ raise ValueError("A partial copy opens with deep_verify=False; its absent parts cannot be read.")
246
+ self._partial = partial
247
+ self._content_pins = None if content_pins is None else dict(content_pins)
248
+ if self._content_pins is not None:
249
+ if not read_only or FEATURE_FILE not in self._content_pins:
250
+ raise ValueError("Pinned feature handles must be read-only and pin feature.json.")
251
+ for relative, digest in self._content_pins.items():
252
+ components = relative.split("/") if isinstance(relative, str) else []
253
+ valid_path = relative == FEATURE_FILE or (
254
+ len(components) == 3 and components[0] == SEGMENTS_DIRECTORY
255
+ and _NAME.fullmatch(components[1])
256
+ and (components[2] == COMMIT_FILE or re.fullmatch(
257
+ r"part-\d{5,}\.(safetensors|rows\.json\.gz)", components[2],
258
+ ))
259
+ )
260
+ if not valid_path or not isinstance(digest, str) or not _SHA256.fullmatch(digest):
261
+ raise ValueError(
262
+ "Feature content pins must name immutable store files and SHA-256 digests."
263
+ )
264
+
265
+ # Opening -----------------------------------------------------------------
266
+
267
+ @classmethod
268
+ def open(
269
+ cls, root: str | Path, spec: StoredFeature, *, deep_verify: bool = True, partial: bool = False,
270
+ ) -> FeatureStore:
271
+ """Open the store for ``spec`` under ``root``, creating it when it does not exist.
272
+
273
+ An existing directory must describe exactly this spec. A mismatch is a caller using one key
274
+ for two different features, which would serve one project another project's numbers, so it
275
+ raises rather than migrating. ``deep_verify=False`` is for stores too large to re-read on
276
+ every open and commit, and ``partial=True`` for a copy holding only some committed parts,
277
+ which must already exist: see the constructor.
278
+ """
279
+
280
+ directory = Path(root) / spec.key
281
+ recorded = directory / FEATURE_FILE
282
+ if partial and not recorded.exists():
283
+ raise FileNotFoundError(f"{directory} holds no {FEATURE_FILE}, so it is no copy of a store.")
284
+ store = cls(directory, spec, deep_verify=deep_verify, partial=partial)
285
+ with store._write_lock():
286
+ if recorded.exists():
287
+ existing = StoredFeature.from_payload(
288
+ json.loads(recorded.read_text(encoding="utf-8")))
289
+ if existing != spec:
290
+ raise ValueError(
291
+ f"{directory} already holds a different feature under key {spec.key!r}.\n"
292
+ f" stored: {existing.payload()}\n"
293
+ f" requested: {spec.payload()}"
294
+ )
295
+ else:
296
+ (directory / SEGMENTS_DIRECTORY).mkdir(parents=True, exist_ok=True)
297
+ _write_json_atomically(recorded, spec.payload())
298
+ store._ensure_index_locked()
299
+ return store
300
+
301
+ @classmethod
302
+ def read_only(
303
+ cls, directory: str | Path, *, content_pins: Mapping[str, str] | None = None,
304
+ ) -> FeatureStore:
305
+ """Open an existing descriptor without creating or repairing any file.
306
+
307
+ Each query opens its own SQLite connection in read-only mode and closes it before
308
+ returning. A missing or corrupt index fails when queried, without being repaired.
309
+ """
310
+
311
+ path = Path(directory)
312
+ payload = json.loads((path / FEATURE_FILE).read_text(encoding="utf-8"))
313
+ store = cls(
314
+ path, StoredFeature.from_payload(payload), read_only=True, content_pins=content_pins,
315
+ )
316
+ store._check_pin(FEATURE_FILE)
317
+ return store
318
+
319
+ # Membership --------------------------------------------------------------
320
+
321
+ def __len__(self) -> int:
322
+ with self._connect() as connection:
323
+ return int(connection.execute("SELECT count(*) FROM rows").fetchone()[0])
324
+
325
+ def address(self, sequence: str) -> RowAddress | None:
326
+ """Where this sequence's row is, or None when the store lacks it."""
327
+
328
+ with self._connect() as connection:
329
+ found = connection.execute(
330
+ "SELECT segment, part, row, residues FROM rows WHERE digest = ?",
331
+ (sequence_digest(sequence),),
332
+ ).fetchone()
333
+ return None if found is None else RowAddress(*found)
334
+
335
+ def missing(self, sequences: Iterable[str]) -> tuple[str, ...]:
336
+ """The sequences this store lacks, in the order given, without repeats.
337
+
338
+ This is the call that makes a run embed only what is new.
339
+ """
340
+
341
+ wanted: dict[str, str] = {}
342
+ for sequence in sequences:
343
+ wanted.setdefault(sequence_digest(sequence), sequence)
344
+ if not wanted:
345
+ return ()
346
+ present = self.present_digests(list(wanted))
347
+ return tuple(sequence for digest, sequence in wanted.items() if digest not in present)
348
+
349
+ def present_digests(self, digests: Sequence[str]) -> set[str]:
350
+ """The digests among ``digests`` that this store already holds, from one index connection."""
351
+
352
+ present: set[str] = set()
353
+ with self._connect() as connection:
354
+ for chunk in _chunks(list(digests), 512):
355
+ marks = ",".join("?" * len(chunk))
356
+ present.update(
357
+ row[0]
358
+ for row in connection.execute(
359
+ f"SELECT digest FROM rows WHERE digest IN ({marks})", chunk
360
+ )
361
+ )
362
+ return present
363
+
364
+ def segments(self) -> tuple[SegmentReceipt, ...]:
365
+ """Every committed segment, oldest commit first, then by fingerprint.
366
+
367
+ Commit times have one-second resolution, so the fingerprint breaks the tie and the order is
368
+ total.
369
+ """
370
+
371
+ receipts = [_receipt(payload) for payload in self._verified_segments()]
372
+ return tuple(sorted(receipts, key=lambda receipt: (receipt.committed_at, receipt.fingerprint)))
373
+
374
+ # Writing -----------------------------------------------------------------
375
+
376
+ @contextmanager
377
+ def segment(
378
+ self, fingerprint: str, metadata: Mapping[str, Any] | None = None,
379
+ *, before_commit: Callable[[], None] | None = None, verify_staged: bool = True,
380
+ ) -> Iterator[SegmentWriter]:
381
+ """Open a new segment for one embedding run, committed on a clean exit.
382
+
383
+ ``verify_staged`` makes the commit re-read and re-hash every staged file. A caller that wrote
384
+ large parts through ``append_packed`` and hashed them as it wrote may pass False: the commit
385
+ then compares each file's size and modification time with what the write recorded, which
386
+ catches an appended or rewritten file without reading the payload a second time.
387
+
388
+ ``fingerprint`` identifies the run that produced these rows, and is FastPLMs' own run
389
+ fingerprint in a real run. A committed segment of that name already existing is an error:
390
+ the same run producing the same rows twice means one of the two is not what it claims.
391
+
392
+ A body that writes nothing leaves nothing behind, and a body that raises leaves the segment
393
+ uncommitted for ``sweep``.
394
+ """
395
+
396
+ self._require_writable()
397
+ if self._partial:
398
+ raise PermissionError("A partial copy takes no new segment; embed into a store of its own.")
399
+ if not isinstance(fingerprint, str) or not _NAME.fullmatch(fingerprint):
400
+ raise ValueError(f"A segment fingerprint must be a filename-safe name; received {fingerprint!r}.")
401
+ if self.spec.descriptor.get("schema") in COMPLETE_SCHEMAS and before_commit is None:
402
+ raise ValueError(
403
+ "A complete feature contract requires a pre-commit validation callback."
404
+ )
405
+ path = self.directory / SEGMENTS_DIRECTORY / fingerprint
406
+ with file_lock(self._segment_lock_path(fingerprint), wait=False):
407
+ with self._write_lock():
408
+ if (path / COMMIT_FILE).exists():
409
+ raise FileExistsError(
410
+ f"Segment {fingerprint!r} is already committed in {self.directory}."
411
+ )
412
+ # Only the owner of this segment lock can discard a crashed attempt.
413
+ if path.exists():
414
+ _discard_segment(path)
415
+ path.mkdir(parents=True)
416
+ sync_directory(path.parent)
417
+ writer = SegmentWriter(
418
+ self, path, fingerprint, dict(metadata or {}), before_commit, lease=_WRITER_LEASE,
419
+ verify_staged=verify_staged,
420
+ )
421
+ try:
422
+ yield writer
423
+ if not writer.committed and not writer.closed:
424
+ if writer.parts:
425
+ writer.commit()
426
+ else:
427
+ writer.abandon()
428
+ finally:
429
+ writer.closed = True
430
+
431
+ def sweep(self) -> tuple[str, ...]:
432
+ """Delete every uncommitted segment directory, and name what was deleted.
433
+
434
+ An uncommitted segment is the remains of a run that died. Nothing reads it.
435
+ """
436
+
437
+ self._require_writable()
438
+ removed: list[str] = []
439
+ for path in sorted((self.directory / SEGMENTS_DIRECTORY).iterdir()):
440
+ if not path.is_dir() or not _NAME.fullmatch(path.name):
441
+ continue
442
+ with ExitStack() as stack:
443
+ try:
444
+ stack.enter_context(file_lock(self._segment_lock_path(path.name), wait=False))
445
+ except BlockingIOError:
446
+ continue
447
+ with self._write_lock():
448
+ if path.exists() and not (path / COMMIT_FILE).exists():
449
+ _discard_segment(path)
450
+ removed.append(path.name)
451
+ return tuple(removed)
452
+
453
+ def reindex(self) -> int:
454
+ """Rebuild the index from the committed segments, and return the row count.
455
+
456
+ The index is a cache, so this is the repair when it is lost or doubted. It reads each
457
+ part's row count from its tensors and each part's digests from the commit marker.
458
+ """
459
+
460
+ self._require_writable()
461
+ with self._write_lock():
462
+ self._ensure_index_locked(repair=True)
463
+ return len(self)
464
+
465
+ # Reading -----------------------------------------------------------------
466
+
467
+ def content_pins(self, sequences: Sequence[str]) -> dict[str, str]:
468
+ """Pin selected immutable files, excluding the rebuildable index.
469
+
470
+ Adding an unrelated committed segment does not invalidate this selection. A pinned
471
+ reader checks the requested rows' actual index associations, markers and part bytes.
472
+ """
473
+ payload = json.loads((self.directory / FEATURE_FILE).read_text(encoding="utf-8"))
474
+ recorded = StoredFeature.from_payload(payload)
475
+ if recorded != self.spec:
476
+ raise ValueError("Feature descriptor changed while pinning a selection.")
477
+ pins = {FEATURE_FILE: file_sha256(self.directory / FEATURE_FILE)}
478
+ markers = {}
479
+ for _, address, _ in self._located(sequences, load=False):
480
+ prefix = f"{SEGMENTS_DIRECTORY}/{address.segment}"
481
+ if address.segment not in markers:
482
+ markers[address.segment] = self._segment_payload(address.segment)
483
+ marker = self.directory / prefix / COMMIT_FILE
484
+ pins[f"{prefix}/{COMMIT_FILE}"] = file_sha256(marker)
485
+ part = markers[address.segment]["parts"][address.part]
486
+ pins[f"{prefix}/{PART_TEMPLATE.format(address.part)}"] = part["sha256"]
487
+ if part.get("row_metadata") is not None:
488
+ pins[f"{prefix}/{part['row_metadata']['file']}"] = part["row_metadata"]["sha256"]
489
+ for relative, expected in pins.items():
490
+ actual = file_sha256(self.directory / relative)
491
+ if actual != expected:
492
+ raise ValueError("Feature bytes changed while pinning a selection.")
493
+ self._check_pin(relative, actual)
494
+ return pins
495
+
496
+ def read(self, sequences: Sequence[str]) -> list[Tensor]:
497
+ """Each sequence's feature as a tensor, in the order asked.
498
+
499
+ A dense feature gives ``(w,)``, a ragged one ``(r_i, d)``, and a csr one a densified
500
+ ``(w,)`` row; ``read_sparse`` keeps a csr row compressed. ``ragged_topk`` requires
501
+ ``read_topk`` to keep residue codes sparse. A sequence the store lacks raises,
502
+ because a silent zero row is indistinguishable from a real one.
503
+ """
504
+
505
+ if self.spec.layout == CSR:
506
+ return [row.to_dense(self.spec.width) for row in self.read_sparse(sequences)] # (w,) each
507
+ if self.spec.layout == RAGGED_TOPK:
508
+ raise ValueError(
509
+ "Use read_topk for sparse residue codes; implicit densification is refused."
510
+ )
511
+ rows: list[Tensor] = []
512
+ for sequence, address, tensors in self._located(sequences):
513
+ # tensors["values"]: (n, w) dense or (sum r_i, d) ragged, for the n rows of one part
514
+ if self.spec.layout == DENSE:
515
+ rows.append(tensors["values"][address.row].clone()) # (w,)
516
+ else:
517
+ offsets = tensors["offsets"] # (n + 1,)
518
+ start, stop = int(offsets[address.row]), int(offsets[address.row + 1])
519
+ rows.append(tensors["values"][start:stop].clone()) # (r_i, d)
520
+ del sequence
521
+ return rows # (w,) or (r_i, d) per sequence
522
+
523
+ def read_sparse(self, sequences: Sequence[str]) -> list[SparseRow]:
524
+ """Each sequence's compressed row, in the order asked. Only for a csr feature."""
525
+
526
+ if self.spec.layout != CSR:
527
+ raise ValueError(f"read_sparse needs a csr feature; this one is {self.spec.layout}.")
528
+ rows: list[SparseRow] = []
529
+ for _, address, tensors in self._located(sequences):
530
+ indptr = tensors["indptr"] # (n + 1,)
531
+ start, stop = int(indptr[address.row]), int(indptr[address.row + 1])
532
+ positions = tensors.get("positions") # (nnz,) or None, for the part's nnz entries
533
+ rows.append(
534
+ SparseRow(
535
+ indices=tensors["indices"][start:stop].clone(), # (nnz_i,)
536
+ values=tensors["values"][start:stop].clone(), # (nnz_i,)
537
+ positions=None if positions is None else positions[start:stop].clone(), # (nnz_i,)
538
+ )
539
+ )
540
+ return rows
541
+
542
+ def read_topk(self, sequences: Sequence[str]) -> list[TopKRow]:
543
+ """Return sparse ``(residues,k)`` codes in request order, including duplicate requests."""
544
+ if self.spec.layout != RAGGED_TOPK:
545
+ raise ValueError(
546
+ f"read_topk needs a ragged_topk feature; this one is {self.spec.layout}."
547
+ )
548
+ rows = []
549
+ for _, address, tensors in self._located(sequences):
550
+ offsets = tensors["offsets"] # (n + 1,)
551
+ start, stop = int(offsets[address.row]), int(offsets[address.row + 1])
552
+ rows.append(TopKRow(
553
+ tensors["indices"][start:stop].clone(), # (r_i, k)
554
+ tensors["values"][start:stop].clone(), # (r_i, k)
555
+ ))
556
+ return rows
557
+
558
+ def residue_counts(self, sequences: Sequence[str]) -> list[int]:
559
+ """Each sequence's stored residue count, which a length-bucketed reader batches by."""
560
+
561
+ return [address.residues for _, address, _ in self._located(sequences, load=False)]
562
+
563
+ def row_metadata(
564
+ self, sequences: Sequence[str], *, verify_data: bool = True,
565
+ ) -> list[dict[str, Any]]:
566
+ """Read opaque row identities in request order and verify committed file digests.
567
+
568
+ The contract provider owns the scientific schema. This layer verifies the physical
569
+ association between sequence, index address, commit marker, data and metadata file.
570
+ Legacy rows without identities raise rather than becoming canonical cache hits.
571
+ """
572
+ markers: dict[str, dict[str, Any]] = {}
573
+ parts: dict[tuple[str, int], list[dict[str, Any]]] = {}
574
+ identities = []
575
+ for sequence, address, _ in self._located(sequences, load=False):
576
+ if not _NAME.fullmatch(address.segment) or address.part < 0:
577
+ raise ValueError("Invalid committed feature address.")
578
+ if address.segment not in markers:
579
+ marker = self.directory / SEGMENTS_DIRECTORY / address.segment / COMMIT_FILE
580
+ payload = json.loads(marker.read_text(encoding="utf-8"))
581
+ if payload["key"] != self.spec.key or payload["fingerprint"] != address.segment:
582
+ raise ValueError("Committed feature identity does not match the index.")
583
+ markers[address.segment] = payload
584
+ payload = markers[address.segment]
585
+ matching = [part for part in payload["parts"] if part["part"] == address.part]
586
+ if len(matching) != 1:
587
+ raise ValueError("Index part does not name one committed feature part.")
588
+ part = matching[0]
589
+ if (not 0 <= address.row < len(part["digests"])
590
+ or part["digests"][address.row] != sequence_digest(sequence)
591
+ or part["residues"][address.row] != address.residues):
592
+ raise ValueError("Feature index and committed sequence row disagree.")
593
+ where = (address.segment, address.part)
594
+ if where not in parts:
595
+ identity = part.get("row_metadata")
596
+ if not isinstance(identity, dict):
597
+ raise ValueError(
598
+ "Stored row identities are unavailable; recompute this legacy feature."
599
+ )
600
+ filename = f"part-{address.part:05d}.rows.json.gz"
601
+ if identity.get("file") != filename:
602
+ raise ValueError("Invalid row metadata filename.")
603
+ path = self.directory / SEGMENTS_DIRECTORY / address.segment / filename
604
+ if file_sha256(path) != identity.get("sha256"):
605
+ raise ValueError("Stored row metadata digest does not match its commit marker.")
606
+ if verify_data and file_sha256(self._part_path(*where)) != part.get("sha256"):
607
+ raise ValueError("Stored feature data digest does not match its commit marker.")
608
+ rows = json.loads(gzip.decompress(path.read_bytes()))
609
+ if (not isinstance(rows, list) or len(rows) != len(part["digests"])
610
+ or not all(isinstance(row, dict) for row in rows)):
611
+ raise ValueError("Stored row metadata does not match the committed part rows.")
612
+ parts[where] = rows
613
+ identities.append(parts[where][address.row])
614
+ return identities
615
+
616
+ # Internals ---------------------------------------------------------------
617
+
618
+ def _located(
619
+ self, sequences: Sequence[str], *, load: bool = True
620
+ ) -> list[tuple[str, RowAddress, dict[str, Tensor]]]:
621
+ """Resolve each sequence to its address, reading each part file at most once."""
622
+
623
+ self._check_pin(FEATURE_FILE)
624
+ addresses = list(zip(sequences, self._addresses(sequences), strict=True))
625
+ cache: dict[tuple[str, int], dict[str, Tensor]] = {}
626
+ markers: dict[str, dict[str, Any]] = {}
627
+ located: list[tuple[str, RowAddress, dict[str, Tensor]]] = []
628
+ for sequence, address in addresses:
629
+ if address.segment not in markers:
630
+ markers[address.segment] = self._segment_payload(address.segment)
631
+ payload = markers[address.segment]
632
+ if not 0 <= address.part < len(payload["parts"]):
633
+ raise ValueError("Index part does not name one committed feature part.")
634
+ part = payload["parts"][address.part]
635
+ if (not 0 <= address.row < len(part["digests"])
636
+ or part["digests"][address.row] != sequence_digest(sequence)
637
+ or part["residues"][address.row] != address.residues):
638
+ raise ValueError("Feature index and committed sequence row disagree.")
639
+ where = (address.segment, address.part)
640
+ if where not in cache:
641
+ tensors = self._verified_part(address.segment, part)
642
+ cache[where] = tensors if load else {}
643
+ tensors = cache[where]
644
+ located.append((sequence, address, tensors))
645
+ return located # (sequence, address, part tensors) per sequence; the tensors are those of _verified_part
646
+
647
+ def _addresses(self, sequences: Sequence[str]) -> list[RowAddress]:
648
+ """Each sequence's address in the order given, from one index connection.
649
+
650
+ Opening the index once per sequence cost minutes for a run of a hundred thousand sequences on a
651
+ network volume. The first sequence the store lacks raises, as ``read`` always has.
652
+ """
653
+
654
+ digests = [sequence_digest(sequence) for sequence in sequences]
655
+ found: dict[str, RowAddress] = {}
656
+ with self._connect() as connection:
657
+ for chunk in _chunks(list(dict.fromkeys(digests)), 512):
658
+ marks = ",".join("?" * len(chunk))
659
+ for digest, *where in connection.execute(
660
+ f"SELECT digest, segment, part, row, residues FROM rows WHERE digest IN ({marks})",
661
+ chunk,
662
+ ):
663
+ found[digest] = RowAddress(*where)
664
+ addresses: list[RowAddress] = []
665
+ for sequence, digest in zip(sequences, digests, strict=True):
666
+ address = found.get(digest)
667
+ if address is None:
668
+ raise KeyError(
669
+ f"Feature {self.spec.key!r} has no row for a sequence of "
670
+ f"{len(sequence)} residues (sha256 {digest[:12]}). "
671
+ "Embed it first, or call missing() before reading."
672
+ )
673
+ addresses.append(address)
674
+ return addresses
675
+
676
+ def _segment_payload(self, fingerprint: str) -> dict[str, Any]:
677
+ """Validate a commit marker before using any filename or row it supplies."""
678
+ if not isinstance(fingerprint, str) or not _NAME.fullmatch(fingerprint):
679
+ raise ValueError("Invalid committed segment fingerprint.")
680
+ marker = self.directory / SEGMENTS_DIRECTORY / fingerprint / COMMIT_FILE
681
+ encoded = marker.read_bytes()
682
+ self._check_pin(
683
+ f"{SEGMENTS_DIRECTORY}/{fingerprint}/{COMMIT_FILE}",
684
+ hashlib.sha256(encoded).hexdigest(),
685
+ )
686
+ payload = json.loads(encoded.decode("utf-8"))
687
+ if not isinstance(payload, dict):
688
+ raise ValueError("A feature commit marker must be an object.")
689
+ if (payload.get("format") != FORMAT or payload.get("key") != self.spec.key
690
+ or payload.get("fingerprint") != fingerprint):
691
+ raise ValueError("Committed feature identity does not match its directory.")
692
+ if any(key in payload for key in (
693
+ "transaction_schema", "manifest_sha256", "descriptor_sha256",
694
+ )):
695
+ if (type(payload.get("transaction_schema")) is not int
696
+ or payload["transaction_schema"] != 2):
697
+ raise ValueError("Unsupported feature transaction schema.")
698
+ unsigned = {key: value for key, value in payload.items() if key != "manifest_sha256"}
699
+ if payload.get("manifest_sha256") != json_sha256(unsigned, allow_nan=False):
700
+ raise ValueError("Feature commit marker checksum mismatch.")
701
+ if payload.get("descriptor_sha256") != json_sha256(self.spec.payload(), allow_nan=False):
702
+ raise ValueError("Feature descriptor does not match its committed digest.")
703
+ parts = payload.get("parts")
704
+ if not isinstance(parts, list) or not parts:
705
+ raise ValueError("A committed segment must contain parts.")
706
+ seen: set[str] = set()
707
+ for number, part in enumerate(parts):
708
+ if (not isinstance(part, dict) or type(part.get("part")) is not int
709
+ or part["part"] != number):
710
+ raise ValueError("Committed part numbers must be consecutive and unique.")
711
+ digests, residues = part.get("digests"), part.get("residues")
712
+ if (not isinstance(digests, list) or not digests or not isinstance(residues, list)
713
+ or len(digests) != len(residues)):
714
+ raise ValueError("Committed sequence and residue counts disagree.")
715
+ if any(not isinstance(value, str) or not _SHA256.fullmatch(value) for value in digests):
716
+ raise ValueError("Invalid committed sequence digest.")
717
+ if len(set(digests)) != len(digests) or seen.intersection(digests):
718
+ raise ValueError("A committed segment repeats sequence rows.")
719
+ seen.update(digests)
720
+ if any(type(value) is not int or value < 0 for value in residues):
721
+ raise ValueError("Invalid committed residue count.")
722
+ if not isinstance(part.get("sha256"), str) or not _SHA256.fullmatch(part["sha256"]):
723
+ raise ValueError("Committed data checksum missing; recompute this legacy segment.")
724
+ if type(payload.get("rows")) is not int or payload["rows"] != len(seen):
725
+ raise ValueError("Committed segment row count disagrees with its parts.")
726
+ if (not isinstance(payload.get("committed_at"), str)
727
+ or not isinstance(payload.get("metadata"), dict)):
728
+ raise ValueError("Invalid committed segment metadata.")
729
+ return payload
730
+
731
+ def _verified_part(self, fingerprint: str, part: Mapping[str, Any]) -> dict[str, Tensor]:
732
+ """Check immutable bytes, physical layout, and optional row sidecars."""
733
+ number = part["part"]
734
+ data_digest = file_sha256(self._part_path(fingerprint, number))
735
+ if data_digest != part["sha256"]:
736
+ raise ValueError("Stored feature data digest does not match its commit marker.")
737
+ self._check_pin(
738
+ f"{SEGMENTS_DIRECTORY}/{fingerprint}/{PART_TEMPLATE.format(number)}", data_digest,
739
+ )
740
+ tensors = self._load_part(fingerprint, number)
741
+ # The part holds count rows. Dense: values (count, w). Otherwise offsets or indptr (count + 1,), with
742
+ # values (sum r_i, d) ragged, (sum r_i, k) top-k plus indices of the same shape, or (nnz,) csr.
743
+ spec, count = self.spec, len(part["digests"])
744
+ if set(tensors) != set(tensor_names(spec.layout, positions=spec.positions)):
745
+ raise ValueError("Committed tensor names do not match the feature layout.")
746
+ values = tensors["values"] # (count, w) dense, (sum r_i, d) ragged, (sum r_i, k) top-k, (nnz,) csr
747
+ if values.dtype != spec.dtype:
748
+ raise ValueError("Committed tensor dtype does not match the feature.")
749
+ if spec.layout == DENSE:
750
+ if values.shape != (count, spec.width) or any(part["residues"]):
751
+ raise ValueError("Committed dense shape or residue counts disagree.")
752
+ else:
753
+ offsets = tensors["indptr" if spec.layout == CSR else "offsets"] # (count + 1,)
754
+ if (offsets.dtype != torch.int64 or offsets.shape != (count + 1,)
755
+ or int(offsets[0]) != 0 or int(offsets[-1]) != len(values)
756
+ or bool((offsets[1:] < offsets[:-1]).any())):
757
+ raise ValueError("Committed row offsets are invalid.")
758
+ if spec.layout == RAGGED:
759
+ if (values.ndim != 2 or values.shape[1] != spec.width
760
+ or (offsets[1:] - offsets[:-1]).tolist() != part["residues"]):
761
+ raise ValueError("Committed ragged shape or residue counts disagree.")
762
+ elif spec.layout == RAGGED_TOPK:
763
+ indices = tensors["indices"] # (sum r_i, k)
764
+ if (indices.dtype != torch.int32
765
+ or (offsets[1:] - offsets[:-1]).tolist() != part["residues"]):
766
+ raise ValueError("Committed top-k index dtype or residue counts disagree.")
767
+ validate_topk(indices, values, spec.width, cast(int, spec.sparse_count))
768
+ else:
769
+ indices = tensors["indices"] # (nnz,)
770
+ if (values.ndim != 1 or indices.dtype != torch.int32
771
+ or indices.shape != values.shape or any(part["residues"])
772
+ or bool(((indices < 0) | (indices >= spec.width)).any())):
773
+ raise ValueError("Committed sparse shape or indices are invalid.")
774
+ for start, stop in pairwise(offsets):
775
+ codes = indices[int(start):int(stop)] # (nnz_i,)
776
+ if len(torch.unique(codes)) != len(codes):
777
+ raise ValueError("Committed sparse row repeats indices.")
778
+ if spec.positions:
779
+ positions = tensors["positions"] # (nnz,)
780
+ if (positions.dtype != torch.int16 or positions.shape != values.shape
781
+ or bool((positions < 0).any())):
782
+ raise ValueError("Committed sparse positions are invalid.")
783
+ if "tensor_bytes" in part and part["tensor_bytes"] != sum(
784
+ value.numel() * value.element_size() for value in tensors.values()
785
+ ):
786
+ raise ValueError("Committed tensor byte count disagrees with its payload.")
787
+ identity = part.get("row_metadata")
788
+ if identity is None and spec.descriptor.get("schema") in COMPLETE_SCHEMAS:
789
+ raise ValueError("Complete feature contracts require committed row identities.")
790
+ if identity is not None:
791
+ filename = f"part-{number:05d}.rows.json.gz"
792
+ if not isinstance(identity, dict) or identity.get("file") != filename:
793
+ raise ValueError("Invalid row metadata filename.")
794
+ path = self.directory / SEGMENTS_DIRECTORY / fingerprint / filename
795
+ metadata_digest = file_sha256(path)
796
+ if metadata_digest != identity.get("sha256"):
797
+ raise ValueError("Stored row metadata digest does not match its commit marker.")
798
+ self._check_pin(f"{SEGMENTS_DIRECTORY}/{fingerprint}/{filename}", metadata_digest)
799
+ rows = json.loads(gzip.decompress(path.read_bytes()))
800
+ if (not isinstance(rows, list) or len(rows) != count
801
+ or not all(isinstance(row, dict) for row in rows)):
802
+ raise ValueError("Stored row metadata does not match the committed part rows.")
803
+ return tensors # (count, w) values for dense; offsets or indptr (count + 1,) with values as above otherwise
804
+
805
+ def _check_pin(self, relative: str, actual: str | None = None) -> None:
806
+ """Check independent selection pins, in addition to a marker's internal checksums."""
807
+ if self._content_pins is not None:
808
+ if relative not in self._content_pins:
809
+ raise ValueError(
810
+ f"Requested feature file is outside the pinned selection: {relative}."
811
+ )
812
+ actual = file_sha256(self.directory / relative) if actual is None else actual
813
+ if actual != self._content_pins[relative]:
814
+ raise ValueError(
815
+ f"Feature file differs from its independent content pin: {relative}."
816
+ )
817
+
818
+ def _verified_segments(self) -> list[dict[str, Any]]:
819
+ payloads, seen = [], set()
820
+ for marker in sorted((self.directory / SEGMENTS_DIRECTORY).glob(f"*/{COMMIT_FILE}")):
821
+ payload = self._segment_payload(marker.parent.name)
822
+ for part in payload["parts"]:
823
+ if seen.intersection(part["digests"]):
824
+ raise ValueError("Committed segments contain conflicting sequence rows.")
825
+ seen.update(part["digests"])
826
+ self._verified_part(payload["fingerprint"], part)
827
+ payloads.append(payload)
828
+ return payloads
829
+
830
+ def _write_lock(self) -> AbstractContextManager[None]:
831
+ self._require_writable()
832
+ return file_lock(self.directory / ".locks" / "store.lock")
833
+
834
+ def _segment_lock_path(self, fingerprint: str) -> Path:
835
+ # 128 bits keep distinct segments on distinct locks, and a short name keeps the path under
836
+ # Windows' 260-character limit, which a store in a nested directory would otherwise exceed.
837
+ name = hashlib.sha256(fingerprint.encode("utf-8")).hexdigest()[:32]
838
+ return self.directory / ".locks" / f"segment-{name}.lock"
839
+
840
+ def _load_part(self, segment: str, part: int) -> dict[str, Tensor]:
841
+ path = self._part_path(segment, part)
842
+ try:
843
+ from safetensors import safe_open
844
+ except ImportError as error:
845
+ raise ImportError("Reading a feature store requires the 'safetensors' package.") from error
846
+ with safe_open(path, framework="pt", device="cpu") as handle:
847
+ # `safe_open` is a handle with keys(), not a mapping: iterating it directly does not work.
848
+ return {name: cast(Tensor, handle.get_tensor(name)) for name in handle.keys()} # noqa: dict-idiom # (...) each, as stored
849
+
850
+ def _part_path(self, segment: str, part: int) -> Path:
851
+ return self.directory / SEGMENTS_DIRECTORY / segment / PART_TEMPLATE.format(part)
852
+
853
+ @contextmanager
854
+ def _connect(self) -> Iterator[sqlite3.Connection]:
855
+ """A connection that is always closed.
856
+
857
+ `sqlite3.Connection` as a context manager ends the transaction but leaves the handle open,
858
+ which on Windows keeps the index file locked against the next writer.
859
+ """
860
+
861
+ database = self.directory / INDEX_FILE
862
+ connection = (
863
+ sqlite3.connect(database.resolve().as_uri() + "?mode=ro", uri=True)
864
+ if self._read_only else sqlite3.connect(database)
865
+ )
866
+ try:
867
+ yield connection
868
+ finally:
869
+ connection.close()
870
+
871
+ def _require_writable(self) -> None:
872
+ if self._read_only:
873
+ raise PermissionError("This feature store handle is read-only.")
874
+
875
+ def _ensure_index(self) -> None:
876
+ self._require_writable()
877
+ with self._write_lock():
878
+ self._ensure_index_locked()
879
+
880
+ def _ensure_index_locked(self, *, repair: bool = False) -> None:
881
+ if not self._deep_verify and not repair:
882
+ self._index_new_segments_locked()
883
+ return
884
+ payloads = self._verified_segments()
885
+ try:
886
+ with self._connect() as connection:
887
+ _rebuild_rows(connection, payloads, repair=repair)
888
+ except sqlite3.DatabaseError as error:
889
+ if getattr(error, "sqlite_errorcode", None) not in (
890
+ sqlite3.SQLITE_CORRUPT, sqlite3.SQLITE_NOTADB,
891
+ ):
892
+ raise
893
+ # A derived corrupt index can be replaced only after every source segment verifies.
894
+ temporary = self.directory / (INDEX_FILE + ".writing")
895
+ temporary.unlink(missing_ok=True)
896
+ connection = sqlite3.connect(temporary)
897
+ try:
898
+ _rebuild_rows(connection, payloads)
899
+ finally:
900
+ connection.close()
901
+ publish_file(temporary, self.directory / INDEX_FILE)
902
+
903
+ def _index_uncounted_segments(self) -> None:
904
+ """Index any committed segment the index does not hold, which a crash can leave behind."""
905
+
906
+ self._ensure_index()
907
+
908
+ def _index_new_segments_locked(self) -> None:
909
+ """Index only the committed segments the index lacks, reading their markers and nothing else.
910
+
911
+ This is the recovery of a store opened with ``deep_verify=False``. A segment already in the
912
+ index is trusted until a reader touches its parts, so the cost of a commit grows with the new
913
+ segment and the count of segments, never with the bytes already stored. A corrupt index is
914
+ rebuilt from the markers alone.
915
+ """
916
+
917
+ try:
918
+ with self._connect() as connection:
919
+ connection.execute("BEGIN IMMEDIATE")
920
+ _create_index_tables(connection)
921
+ self._index_missing_segments(connection)
922
+ connection.commit()
923
+ except sqlite3.DatabaseError as error:
924
+ if getattr(error, "sqlite_errorcode", None) not in (
925
+ sqlite3.SQLITE_CORRUPT, sqlite3.SQLITE_NOTADB,
926
+ ):
927
+ raise
928
+ temporary = self.directory / (INDEX_FILE + ".writing")
929
+ temporary.unlink(missing_ok=True)
930
+ connection = sqlite3.connect(temporary)
931
+ try:
932
+ connection.execute("BEGIN IMMEDIATE")
933
+ _create_index_tables(connection)
934
+ self._index_missing_segments(connection)
935
+ connection.commit()
936
+ finally:
937
+ connection.close()
938
+ publish_file(temporary, self.directory / INDEX_FILE)
939
+
940
+ def _index_missing_segments(self, connection: sqlite3.Connection) -> None:
941
+ known = {row[0] for row in connection.execute("SELECT segment FROM segments")}
942
+ # A partial copy's index is always this code's, and lists a segment once all its parts are in.
943
+ if not known and not self._partial:
944
+ # An index written before the segment table existed lists its segments only through its rows.
945
+ connection.execute("INSERT OR IGNORE INTO segments (segment) SELECT DISTINCT segment FROM rows")
946
+ known = {row[0] for row in connection.execute("SELECT segment FROM segments")}
947
+ for marker in sorted((self.directory / SEGMENTS_DIRECTORY).glob(f"*/{COMMIT_FILE}")):
948
+ name = marker.parent.name
949
+ if name in known:
950
+ continue
951
+ payload = self._segment_payload(name)
952
+ complete = True
953
+ for part in payload["parts"]:
954
+ number = int(part["part"])
955
+ if not self._part_path(name, number).is_file():
956
+ if not self._partial:
957
+ raise ValueError(f"Committed part {part['part']} of segment {name!r} is missing.")
958
+ complete = False # not fetched: its rows read as missing until a fetch brings it
959
+ continue
960
+ if self._partial and connection.execute(
961
+ "SELECT 1 FROM rows WHERE segment = ? AND part = ? LIMIT 1", (name, number),
962
+ ).fetchone():
963
+ continue # an earlier open of this copy indexed it
964
+ _insert_rows(connection, name, number, part)
965
+ if complete:
966
+ connection.execute("INSERT OR IGNORE INTO segments (segment) VALUES (?)", (name,))
967
+
968
+
969
+ class SegmentWriter:
970
+ """One run's segment: parts as it streams, then one commit marker and the index rows."""
971
+
972
+ def __init__(
973
+ self, store: FeatureStore, path: Path, fingerprint: str, metadata: dict[str, Any],
974
+ before_commit: Callable[[], None] | None = None,
975
+ *, lease: object | None = None, verify_staged: bool = True,
976
+ ) -> None:
977
+ store._require_writable()
978
+ if lease is not _WRITER_LEASE:
979
+ raise RuntimeError("Open a segment writer through FeatureStore.segment().")
980
+ self.store = store
981
+ self.path = path
982
+ self.fingerprint = fingerprint
983
+ self.metadata = json.loads(json.dumps(metadata, allow_nan=False))
984
+ self.before_commit = before_commit
985
+ self.committed = False
986
+ self.closed = False
987
+ self.verify_staged = verify_staged
988
+ self._owner_pid = os.getpid()
989
+ self._failed = False
990
+ self._seen_digests: set[str] = set()
991
+ self._parts: list[dict[str, Any]] = []
992
+ # Parts are numbered when reserved and recorded when written, which may be on another thread.
993
+ self._lock = threading.Lock()
994
+ self._reserved = 0
995
+ self._staged_stats: dict[str, tuple[int, int]] = {}
996
+
997
+ def __getstate__(self) -> dict[str, Any]:
998
+ # A lock does not pickle. A copy sent to another process gets a fresh one and still refuses every write,
999
+ # commit and abandon, because its owner pid is not that process's.
1000
+ state = dict(self.__dict__)
1001
+ del state["_lock"]
1002
+ return state
1003
+
1004
+ def __setstate__(self, state: dict[str, Any]) -> None:
1005
+ self.__dict__.update(state)
1006
+ self._lock = threading.Lock()
1007
+
1008
+ def _require_open(self) -> None:
1009
+ if self._owner_pid != os.getpid():
1010
+ raise RuntimeError("A segment writer belongs to the process that opened its context.")
1011
+ if self.committed or (self.path / COMMIT_FILE).exists():
1012
+ raise RuntimeError(f"Segment {self.fingerprint!r} is already committed.")
1013
+ if self.closed or self._failed:
1014
+ raise RuntimeError("This segment writer is closed or failed; start a fresh attempt.")
1015
+
1016
+ @property
1017
+ def parts(self) -> tuple[int, ...]:
1018
+ """The part numbers written so far, which are visible only once committed."""
1019
+
1020
+ return tuple(int(part["part"]) for part in self._parts)
1021
+
1022
+ def append_bounded(
1023
+ self, sequences: Sequence[str],
1024
+ rows: Sequence[Tensor] | Sequence[SparseRow] | Sequence[TopKRow],
1025
+ *, max_tensor_bytes: int, row_metadata: Sequence[Mapping[str, Any]] | None = None,
1026
+ ) -> tuple[int, ...]:
1027
+ """Split a bounded window into lossless parts capped by their encoded tensor payload.
1028
+
1029
+ Metadata sidecars and safetensors headers are separate. A row cannot span parts; reject
1030
+ an oversized row instead of silently writing an oversized part or changing its values.
1031
+ """
1032
+ # rows: (w,) dense or (r_i, d) ragged tensors; a SparseRow holds (nnz_i,) and a TopKRow (r_i, k) tensors
1033
+ if type(max_tensor_bytes) is not int or max_tensor_bytes <= 0:
1034
+ raise ValueError("max_tensor_bytes must be a positive integer.")
1035
+ if not sequences or len(sequences) != len(rows):
1036
+ raise ValueError("append_bounded needs one row per sequence and at least one row.")
1037
+ if len(set(sequences)) != len(sequences):
1038
+ raise ValueError("This batch repeats sequences; each feature has one row.")
1039
+ if row_metadata is not None and len(row_metadata) != len(sequences):
1040
+ raise ValueError("Row metadata must contain one identity per sequence.")
1041
+ spec = self.store.spec
1042
+ initial = 0 if spec.layout == DENSE else 8
1043
+ sizes = [
1044
+ row_tensor_bytes(row, spec.layout, spec.width, spec.dtype, positions=spec.positions)
1045
+ for row in rows
1046
+ ]
1047
+ if any(initial + size > max_tensor_bytes for size in sizes):
1048
+ raise ValueError(
1049
+ "A feature row exceeds max_tensor_bytes; increase the explicit part budget."
1050
+ )
1051
+ written, start, used = [], 0, initial
1052
+ for stop in range(len(rows) + 1):
1053
+ size = sizes[stop] if stop < len(rows) else 0
1054
+ if stop == len(rows) or used + size > max_tensor_bytes:
1055
+ part = self.append(
1056
+ sequences[start:stop], rows[start:stop],
1057
+ row_metadata=None if row_metadata is None else row_metadata[start:stop],
1058
+ )
1059
+ if self._parts[part]["tensor_bytes"] > max_tensor_bytes:
1060
+ raise RuntimeError("Encoded tensor payload exceeded the planned part budget.")
1061
+ written.append(part)
1062
+ start, used = stop, initial
1063
+ used += size
1064
+ return tuple(written)
1065
+
1066
+ def append(
1067
+ self, sequences: Sequence[str],
1068
+ rows: Sequence[Tensor] | Sequence[SparseRow] | Sequence[TopKRow],
1069
+ *, row_metadata: Sequence[Mapping[str, Any]] | None = None,
1070
+ ) -> int:
1071
+ """Write one part holding these rows, and return the part number.
1072
+
1073
+ Sequences the store already holds, or that repeat inside this batch, are an error: a
1074
+ feature has one row, and writing it twice makes two answers to one question. Call
1075
+ ``missing`` first.
1076
+ """
1077
+
1078
+ # rows: (w,) dense or (r_i, d) ragged tensors; a SparseRow holds (nnz_i,) and a TopKRow (r_i, k) tensors
1079
+ self._require_open()
1080
+ if len(sequences) != len(rows):
1081
+ raise ValueError(
1082
+ f"append needs one row per sequence; received {len(sequences)} sequences and "
1083
+ f"{len(rows)} rows."
1084
+ )
1085
+ if not sequences:
1086
+ raise ValueError("append needs at least one sequence.")
1087
+ digests = [sequence_digest(sequence) for sequence in sequences]
1088
+ repeated = sorted({digest for digest in digests if digests.count(digest) > 1})
1089
+ if repeated or self._seen_digests.intersection(digests):
1090
+ raise ValueError(f"This batch repeats {len(repeated)} sequences; each feature has one row.")
1091
+ already = self.store.missing(sequences)
1092
+ if len(already) != len(sequences):
1093
+ raise ValueError(
1094
+ f"{len(sequences) - len(already)} of these sequences already have a row in "
1095
+ f"{self.store.spec.key!r}; call missing() and embed only what it returns."
1096
+ )
1097
+
1098
+ spec = self.store.spec
1099
+ if spec.descriptor.get("schema") in COMPLETE_SCHEMAS and row_metadata is None:
1100
+ raise ValueError(
1101
+ "Complete feature contracts require a persisted identity for each row."
1102
+ )
1103
+ if row_metadata is not None and len(row_metadata) != len(sequences):
1104
+ raise ValueError("Row metadata must contain one identity per sequence.")
1105
+ encoded_metadata = None if row_metadata is None else gzip.compress(
1106
+ json.dumps([dict(row) for row in row_metadata], sort_keys=True, allow_nan=False,
1107
+ separators=(",", ":")).encode("utf-8"), mtime=0,
1108
+ )
1109
+ if spec.layout == DENSE:
1110
+ tensors = encode_dense(cast(Sequence[Tensor], rows), spec.width, spec.dtype)
1111
+ residues = [0] * len(sequences)
1112
+ elif spec.layout == CSR:
1113
+ tensors = encode_csr(cast(Sequence[SparseRow], rows), spec.width, spec.dtype)
1114
+ if spec.positions and "positions" not in tensors:
1115
+ raise ValueError(
1116
+ f"Feature {spec.key!r} stores argmax positions; these rows carry none."
1117
+ )
1118
+ if not spec.positions and "positions" in tensors:
1119
+ raise ValueError(
1120
+ f"Feature {spec.key!r} stores no argmax positions; these rows carry them."
1121
+ )
1122
+ residues = [0] * len(sequences)
1123
+ else:
1124
+ if spec.layout == RAGGED_TOPK:
1125
+ tensors = encode_topk_rows(
1126
+ cast(Sequence[TopKRow], rows), spec.width,
1127
+ cast(int, spec.sparse_count), spec.dtype,
1128
+ )
1129
+ else:
1130
+ tensors = encode_ragged(cast(Sequence[Tensor], rows), spec.width, spec.dtype)
1131
+ offsets = tensors["offsets"] # (n + 1,)
1132
+ residues = [
1133
+ int(offsets[position + 1]) - int(offsets[position])
1134
+ for position in range(len(sequences))
1135
+ ]
1136
+ written = row_count(spec.layout, tensors)
1137
+ if written != len(sequences):
1138
+ raise ValueError(
1139
+ f"Encoded {written} rows for {len(sequences)} sequences; the layout and the rows "
1140
+ "disagree."
1141
+ )
1142
+
1143
+ part = self.reserve_part()
1144
+ try:
1145
+ _transaction_event("before_part_write", self.path)
1146
+ _save_safetensors_atomically(
1147
+ self.path / PART_TEMPLATE.format(part),
1148
+ {name: tensors[name]
1149
+ for name in tensor_names(spec.layout, positions=spec.positions)},
1150
+ )
1151
+ _transaction_event("after_part_write", self.path)
1152
+ record = {
1153
+ "part": part, "digests": digests, "residues": residues,
1154
+ "sha256": file_sha256(self.path / PART_TEMPLATE.format(part)),
1155
+ "tensor_bytes": sum(
1156
+ value.numel() * value.element_size() for value in tensors.values()),
1157
+ }
1158
+ if encoded_metadata is not None:
1159
+ target = self.path / f"part-{part:05d}.rows.json.gz"
1160
+ temporary = target.with_name(target.name + ".writing")
1161
+ temporary.write_bytes(encoded_metadata)
1162
+ publish_file(temporary, target)
1163
+ record["row_metadata"] = {"file": target.name, "sha256": file_sha256(target)}
1164
+ _transaction_event("after_metadata_write", self.path)
1165
+ except BaseException:
1166
+ self._failed = True
1167
+ raise
1168
+ self._parts.append(record)
1169
+ self._seen_digests.update(digests)
1170
+ return part
1171
+
1172
+ def reserve_part(self) -> int:
1173
+ """The next part number, taken before a part is written so that writers on several threads never collide."""
1174
+
1175
+ with self._lock:
1176
+ number = self._reserved
1177
+ self._reserved += 1
1178
+ return number
1179
+
1180
+ def seen(self, digests: Sequence[str]) -> bool:
1181
+ """Whether any digest was already written by this segment, a cheap guard before a large write."""
1182
+
1183
+ with self._lock:
1184
+ return not self._seen_digests.isdisjoint(digests)
1185
+
1186
+ def append_packed(
1187
+ self, part: int, digests: Sequence[str], tensors: Mapping[str, Tensor], residues: Sequence[int],
1188
+ *, row_metadata: Sequence[Mapping[str, Any]] | None = None,
1189
+ ) -> dict[str, Any]:
1190
+ """Write one large part from tensors already packed in the layout, hashing as it writes.
1191
+
1192
+ ``tensors`` holds exactly the layout's tensors, as ``encode_*`` would build them. This skips
1193
+ the per-row copies, finite scans and index re-reads of ``append``: the caller proved the values
1194
+ finite on the device and packed them in order. Shapes, dtypes, offsets and the residue counts
1195
+ are still checked, because they decide whether a reader can address the part. The file and
1196
+ its digest come from one pass, and one flush covers the part. Safe to call from several
1197
+ threads with parts from ``reserve_part``.
1198
+
1199
+ Shapes, with b rows (sequences) in the part and n = sum(residues) stored rows: dense ``values`` (b, w);
1200
+ ragged ``offsets`` (b + 1,) and ``values`` (n, w); ragged top-k adds ``indices`` (n, k); csr ``indptr``
1201
+ (b + 1,) with ``indices`` and ``values`` (nnz,). ``residues[i]`` is row i's stored rows, which is l + 2
1202
+ for a stream that keeps CLS and EOS.
1203
+ """
1204
+ # tensors: (b, w) dense; (b + 1,) offsets and (n, w) values ragged; (b + 1,) indptr and (nnz,) csr.
1205
+ spec = self.store.spec
1206
+ if self._owner_pid != os.getpid():
1207
+ raise RuntimeError("A segment writer belongs to the process that opened its context.")
1208
+ if self.committed or (self.path / COMMIT_FILE).exists():
1209
+ raise RuntimeError(f"Segment {self.fingerprint!r} is already committed.")
1210
+ if self.closed or self._failed:
1211
+ raise RuntimeError("This segment writer is closed or failed; start a fresh attempt.")
1212
+ count = len(digests)
1213
+ if not count or len(residues) != count:
1214
+ raise ValueError("append_packed needs one residue count and one digest per row, and at least one row.")
1215
+ if spec.descriptor.get("schema") in COMPLETE_SCHEMAS and row_metadata is None:
1216
+ raise ValueError("Complete feature contracts require a persisted identity for each row.")
1217
+ if row_metadata is not None and len(row_metadata) != count:
1218
+ raise ValueError("Row metadata must contain one identity per sequence.")
1219
+ if set(tensors) != set(tensor_names(spec.layout, positions=spec.positions)):
1220
+ raise ValueError("Packed tensors do not match the feature layout.")
1221
+ if len(set(digests)) != count or self.seen(digests):
1222
+ raise ValueError("A packed part repeats sequences; each feature has one row.")
1223
+ _check_packed(spec, tensors, count, residues)
1224
+
1225
+ try:
1226
+ _transaction_event("before_part_write", self.path)
1227
+ target = self.path / PART_TEMPLATE.format(part)
1228
+ digest, size = _write_safetensors_streaming(target, tensors)
1229
+ _transaction_event("after_part_write", self.path)
1230
+ record: dict[str, Any] = {
1231
+ "part": part, "digests": list(digests), "residues": [int(value) for value in residues],
1232
+ "sha256": digest,
1233
+ "tensor_bytes": sum(value.numel() * value.element_size() for value in tensors.values()),
1234
+ }
1235
+ stats = {target.name: (size, target.stat().st_mtime_ns)}
1236
+ if row_metadata is not None:
1237
+ encoded = gzip.compress(
1238
+ json.dumps([dict(row) for row in row_metadata], sort_keys=True, allow_nan=False,
1239
+ separators=(",", ":")).encode("utf-8"), mtime=0,
1240
+ )
1241
+ sidecar = self.path / f"part-{part:05d}.rows.json.gz"
1242
+ temporary = sidecar.with_name(sidecar.name + ".writing")
1243
+ temporary.write_bytes(encoded)
1244
+ publish_file(temporary, sidecar, sync_parent=False)
1245
+ record["row_metadata"] = {"file": sidecar.name, "sha256": hashlib.sha256(encoded).hexdigest()}
1246
+ stats[sidecar.name] = (len(encoded), sidecar.stat().st_mtime_ns)
1247
+ _transaction_event("after_metadata_write", self.path)
1248
+ except BaseException:
1249
+ self._failed = True
1250
+ raise
1251
+ with self._lock:
1252
+ self._parts.append(record)
1253
+ self._seen_digests.update(digests)
1254
+ self._staged_stats.update(stats)
1255
+ return record
1256
+
1257
+ def commit(self) -> SegmentReceipt:
1258
+ """Write the commit marker, then index the parts. Nothing reads a segment until this."""
1259
+
1260
+ self._require_open()
1261
+ with self._lock:
1262
+ # Parts written on several threads finish in any order; the marker lists them by number.
1263
+ self._parts.sort(key=lambda part: part["part"])
1264
+ if not self._parts:
1265
+ raise RuntimeError(f"Segment {self.fingerprint!r} holds no parts to commit.")
1266
+ payload = {
1267
+ "format": FORMAT,
1268
+ "key": self.store.spec.key,
1269
+ "fingerprint": self.fingerprint,
1270
+ "committed_at": datetime.now(UTC).isoformat(timespec="seconds"),
1271
+ "rows": sum(len(part["digests"]) for part in self._parts),
1272
+ "metadata": self.metadata,
1273
+ "parts": self._parts,
1274
+ "transaction_schema": 2,
1275
+ "descriptor_sha256": json_sha256(self.store.spec.payload(), allow_nan=False),
1276
+ }
1277
+ recorded = json.loads((self.store.directory / FEATURE_FILE).read_text(encoding="utf-8"))
1278
+ if recorded != self.store.spec.payload():
1279
+ raise ValueError("Feature descriptor changed while the segment was staged.")
1280
+ for part in self._parts:
1281
+ self._check_staged(self.path / PART_TEMPLATE.format(part["part"]), part["sha256"])
1282
+ identity = part.get("row_metadata")
1283
+ if identity:
1284
+ self._check_staged(self.path / identity["file"], identity["sha256"])
1285
+ if self.before_commit is not None:
1286
+ self.before_commit()
1287
+ # Parts and sidecars were renamed without a directory flush; one flush covers them all.
1288
+ sync_directory(self.path)
1289
+ payload["manifest_sha256"] = json_sha256(payload, allow_nan=False)
1290
+ with self.store._write_lock():
1291
+ # Recover prior durable markers before deciding whether any row conflicts.
1292
+ self.store._ensure_index_locked()
1293
+ with self.store._connect() as connection:
1294
+ connection.execute("BEGIN IMMEDIATE")
1295
+ _transaction_event("before_index_update", self.path)
1296
+ _create_index_tables(connection) # idempotent: an index from before the segment list gains it here
1297
+ for part in self._parts:
1298
+ _insert_rows(connection, self.fingerprint, int(part["part"]), part)
1299
+ connection.execute("INSERT OR IGNORE INTO segments (segment) VALUES (?)", (self.fingerprint,))
1300
+ _transaction_event("after_index_update", self.path)
1301
+ _transaction_event("before_commit_marker", self.path)
1302
+ _write_json_atomically(self.path / COMMIT_FILE, payload)
1303
+ self.committed = True
1304
+ _transaction_event("after_commit_marker", self.path)
1305
+ connection.commit()
1306
+ _transaction_event("after_index_commit", self.path)
1307
+ return _receipt(payload)
1308
+
1309
+ def _check_staged(self, path: Path, digest: str) -> None:
1310
+ """Refuse a staged file that changed since it was written.
1311
+
1312
+ A writer that hashed as it wrote (``verify_staged=False``) compares size and modification
1313
+ time, which catches an appended or rewritten file without a second read of the payload.
1314
+ """
1315
+ recorded = self._staged_stats.get(path.name)
1316
+ if self.verify_staged or recorded is None:
1317
+ if file_sha256(path) != digest:
1318
+ raise ValueError(
1319
+ "Staged feature data changed before commit." if path.suffix == ".safetensors"
1320
+ else "Staged row metadata changed before commit."
1321
+ )
1322
+ return
1323
+ observed = path.stat()
1324
+ if (observed.st_size, observed.st_mtime_ns) != recorded:
1325
+ raise ValueError(
1326
+ "Staged feature data changed before commit." if path.suffix == ".safetensors"
1327
+ else "Staged row metadata changed before commit."
1328
+ )
1329
+
1330
+ def abandon(self) -> None:
1331
+ """Delete this segment's parts, for a run that decides not to keep them."""
1332
+
1333
+ if self._owner_pid != os.getpid():
1334
+ raise RuntimeError("A segment writer belongs to the process that opened its context.")
1335
+ if self.committed or (self.path / COMMIT_FILE).exists():
1336
+ raise RuntimeError(f"Segment {self.fingerprint!r} is committed and immutable.")
1337
+ if self.closed:
1338
+ raise RuntimeError("This segment writer is closed.")
1339
+ _discard_segment(self.path)
1340
+ self._parts.clear()
1341
+ self._seen_digests.clear()
1342
+ self.closed = True
1343
+
1344
+
1345
+ # The documented public entry point (docs/feature_store.md, `features.__all__`), kept under its name.
1346
+ def open_feature(root: str | Path, spec: StoredFeature) -> FeatureStore: # noqa: renaming-wrapper
1347
+ """Open, or create, the store for one feature under ``root``."""
1348
+
1349
+ return FeatureStore.open(root, spec)
1350
+
1351
+
1352
+ def features_in(root: str | Path) -> tuple[FeatureStore, ...]:
1353
+ """Every feature store under ``root``, by key."""
1354
+
1355
+ return tuple(
1356
+ FeatureStore.read_only(recorded.parent)
1357
+ for recorded in sorted(Path(root).glob(f"*/{FEATURE_FILE}"))
1358
+ )
1359
+
1360
+
1361
+ def _insert_rows(
1362
+ connection: sqlite3.Connection, fingerprint: str, part: int, payload: Mapping[str, Any]
1363
+ ) -> None:
1364
+ try:
1365
+ connection.executemany(
1366
+ "INSERT INTO rows (digest, segment, part, row, residues) VALUES (?, ?, ?, ?, ?)",
1367
+ [(digest, fingerprint, part, row, int(residues))
1368
+ for row, (digest, residues) in enumerate(
1369
+ zip(payload["digests"], payload["residues"], strict=True)
1370
+ )],
1371
+ )
1372
+ except sqlite3.IntegrityError as error:
1373
+ raise ValueError(
1374
+ f"Segment {fingerprint!r} conflicts with an existing sequence row."
1375
+ ) from error
1376
+
1377
+
1378
+ def _receipt(payload: Mapping[str, Any]) -> SegmentReceipt:
1379
+ return SegmentReceipt(
1380
+ fingerprint=str(payload["fingerprint"]),
1381
+ parts=tuple(int(part["part"]) for part in payload["parts"]),
1382
+ rows=int(payload["rows"]),
1383
+ committed_at=str(payload["committed_at"]),
1384
+ metadata=payload.get("metadata") or {},
1385
+ )
1386
+
1387
+
1388
+ def _write_json_atomically(path: Path, payload: Mapping[str, Any]) -> None:
1389
+ temporary = path.with_name(f"{path.name}.writing")
1390
+ temporary.write_text(indented_json(payload, allow_nan=False), encoding="utf-8")
1391
+ publish_file(temporary, path)
1392
+
1393
+
1394
+ def _save_safetensors_atomically(path: Path, tensors: Mapping[str, Tensor]) -> None:
1395
+ # tensors: (n, w), (n + 1,), (nnz,) or (sum r_i, d) by name and layout; each is written as it is
1396
+ try:
1397
+ from safetensors.torch import save_file
1398
+ except ImportError as error:
1399
+ raise ImportError("Writing a feature store requires the 'safetensors' package.") from error
1400
+ temporary = path.with_name(f"{path.name}.writing")
1401
+ save_file({name: tensor.contiguous() for name, tensor in tensors.items()}, str(temporary))
1402
+ publish_file(temporary, path)
1403
+
1404
+
1405
+ _SAFETENSORS_DTYPES = {
1406
+ torch.float64: "F64", torch.float32: "F32", torch.float16: "F16", torch.bfloat16: "BF16",
1407
+ torch.int64: "I64", torch.int32: "I32", torch.int16: "I16",
1408
+ }
1409
+ _WRITE_BLOCK_BYTES = 64 * 1024**2
1410
+
1411
+
1412
+ def _write_safetensors_streaming(path: Path, tensors: Mapping[str, Tensor]) -> tuple[str, int]:
1413
+ """Write ``tensors`` as a safetensors file in one pass, hashing the bytes as they are written.
1414
+
1415
+ The file is the safetensors format that ``safe_open`` reads: an 8-byte little-endian header
1416
+ length, a JSON header padded with spaces to a multiple of 8, then the raw tensor bytes. Tensors go
1417
+ in decreasing element size, so every tensor starts aligned to its own element size. The tensor
1418
+ memory is written from its own buffer in blocks, with no second serialization copy, and the
1419
+ SHA-256 of the file comes from those same blocks. One flush makes the part durable. Returns the
1420
+ digest and the size in bytes.
1421
+ """
1422
+ # tensors: (b, w), (n, w), (n, k), (b + 1,) or (nnz,) by layout; any shape, the header records each as is.
1423
+ ordered = sorted(tensors.items(), key=lambda item: (-item[1].element_size(), item[0]))
1424
+ header: dict[str, Any] = {}
1425
+ cursor = 0
1426
+ for name, tensor in ordered:
1427
+ size = tensor.numel() * tensor.element_size()
1428
+ header[name] = {
1429
+ "dtype": _SAFETENSORS_DTYPES[tensor.dtype], "shape": list(tensor.shape),
1430
+ "data_offsets": [cursor, cursor + size],
1431
+ }
1432
+ cursor += size
1433
+ encoded = json.dumps(header, separators=(",", ":")).encode("utf-8")
1434
+ encoded += b" " * (-len(encoded) % 8)
1435
+ prefix = len(encoded).to_bytes(8, "little")
1436
+ temporary = path.with_name(f"{path.name}.writing")
1437
+ digest = hashlib.sha256()
1438
+ with temporary.open("wb") as handle:
1439
+ for block in (prefix, encoded):
1440
+ handle.write(block)
1441
+ digest.update(block)
1442
+ for _, tensor in ordered:
1443
+ raw = tensor.detach().contiguous().reshape(-1).view(torch.uint8) # (bytes,) over the tensor's own memory
1444
+ buffer = memoryview(raw.numpy())
1445
+ for start in range(0, len(buffer), _WRITE_BLOCK_BYTES):
1446
+ block = buffer[start : start + _WRITE_BLOCK_BYTES]
1447
+ handle.write(block)
1448
+ digest.update(block)
1449
+ handle.flush()
1450
+ flush_and_evict(handle)
1451
+ temporary.replace(path)
1452
+ return digest.hexdigest(), 8 + len(encoded) + cursor
1453
+
1454
+
1455
+ def _check_packed(
1456
+ spec: StoredFeature, tensors: Mapping[str, Tensor], count: int, residues: Sequence[int],
1457
+ ) -> None:
1458
+ """The structural checks `append_packed` keeps: whatever decides whether a reader can address the part."""
1459
+ # tensors: (b, w) dense; (b + 1,) offsets and (n, w) values ragged; (b + 1,) indptr and (nnz,) csr.
1460
+ values = tensors["values"] # (b, w) dense, (n, w) ragged, (nnz,) csr
1461
+ if values.dtype != spec.dtype:
1462
+ raise ValueError("Packed tensor dtype does not match the feature.")
1463
+ if spec.layout == DENSE:
1464
+ if tuple(values.shape) != (count, spec.width) or any(residues):
1465
+ raise ValueError("Packed dense shape or residue counts disagree.")
1466
+ return
1467
+ offsets = tensors["indptr" if spec.layout == CSR else "offsets"] # (b + 1,) int64 row boundaries
1468
+ if (offsets.dtype != torch.int64 or tuple(offsets.shape) != (count + 1,)
1469
+ or int(offsets[0]) != 0 or int(offsets[-1]) != len(values)
1470
+ or bool((offsets[1:] < offsets[:-1]).any())):
1471
+ raise ValueError("Packed row offsets are invalid.")
1472
+ spans = (offsets[1:] - offsets[:-1]).tolist() # (b,) stored rows (ragged) or entries (csr) per row
1473
+ if spec.layout == RAGGED:
1474
+ if values.ndim != 2 or values.shape[1] != spec.width or spans != list(residues):
1475
+ raise ValueError("Packed ragged shape or residue counts disagree.")
1476
+ elif spec.layout == RAGGED_TOPK:
1477
+ indices = tensors["indices"] # (n, k) int32 codes, the shape of values
1478
+ if (indices.dtype != torch.int32 or tuple(indices.shape) != tuple(values.shape)
1479
+ or values.ndim != 2 or values.shape[1] != spec.sparse_count or spans != list(residues)):
1480
+ raise ValueError("Packed top-k shape, index dtype or residue counts disagree.")
1481
+ else:
1482
+ indices = tensors["indices"]
1483
+ if (values.ndim != 1 or indices.dtype != torch.int32 or indices.shape != values.shape
1484
+ or any(residues)):
1485
+ raise ValueError("Packed sparse shape or residue counts disagree.")
1486
+ if spec.positions and (tensors["positions"].dtype != torch.int16
1487
+ or tensors["positions"].shape != values.shape):
1488
+ raise ValueError("Packed sparse positions are invalid.")
1489
+
1490
+
1491
+ def _discard_segment(path: Path) -> None:
1492
+ if path.is_symlink() or any(not part.is_file() or part.is_symlink() for part in path.iterdir()):
1493
+ raise ValueError("Refusing to discard an uncommitted segment with unexpected entries.")
1494
+ for part in path.iterdir():
1495
+ part.unlink()
1496
+ path.rmdir()
1497
+ sync_directory(path.parent)
1498
+
1499
+
1500
+ def _rebuild_rows(
1501
+ connection: sqlite3.Connection, payloads: Sequence[Mapping[str, Any]], *, repair: bool = False,
1502
+ ) -> None:
1503
+ connection.execute("BEGIN IMMEDIATE")
1504
+ _create_index_tables(connection)
1505
+ expected = {
1506
+ digest: (payload["fingerprint"], part["part"], row, residues)
1507
+ for payload in payloads for part in payload["parts"]
1508
+ for row, (digest, residues) in enumerate(
1509
+ zip(part["digests"], part["residues"], strict=True))
1510
+ }
1511
+ if repair:
1512
+ connection.execute("DELETE FROM rows")
1513
+ indexed = {
1514
+ row[0]: tuple(row[1:])
1515
+ for row in connection.execute("SELECT digest, segment, part, row, residues FROM rows")
1516
+ }
1517
+ if any(expected.get(digest) != address for digest, address in indexed.items()):
1518
+ raise ValueError(
1519
+ "Feature index disagrees with committed rows; inspect it and call reindex()."
1520
+ )
1521
+ # Recover missing rows only. Changed addresses must fail, not silently become cache hits.
1522
+ connection.executemany(
1523
+ "INSERT INTO rows (digest, segment, part, row, residues) VALUES (?, ?, ?, ?, ?)",
1524
+ [(digest, *address) for digest, address in expected.items() if digest not in indexed],
1525
+ )
1526
+ connection.executemany(
1527
+ "INSERT OR IGNORE INTO segments (segment) VALUES (?)", [(payload["fingerprint"],) for payload in payloads],
1528
+ )
1529
+ connection.commit()
1530
+
1531
+
1532
+ def _create_index_tables(connection: sqlite3.Connection) -> None:
1533
+ """The row index, and the list of segments it already holds, so recovery can skip them."""
1534
+
1535
+ connection.execute(
1536
+ "CREATE TABLE IF NOT EXISTS rows (digest TEXT PRIMARY KEY, segment TEXT NOT NULL, "
1537
+ "part INTEGER NOT NULL, row INTEGER NOT NULL, residues INTEGER NOT NULL)"
1538
+ )
1539
+ connection.execute("CREATE INDEX IF NOT EXISTS rows_by_segment ON rows (segment, part)")
1540
+ connection.execute("CREATE TABLE IF NOT EXISTS segments (segment TEXT PRIMARY KEY)")
1541
+
1542
+
1543
+ def _transaction_event(stage: str, path: Path) -> None:
1544
+ """A test observation point at a real I/O boundary; production performs no action."""
1545
+
1546
+
1547
+ def _chunks(values: Sequence[str], size: int) -> Iterator[list[str]]:
1548
+ for start in range(0, len(values), size):
1549
+ yield list(values[start : start + size])
1550
+
1551
+
1552
+ __all__ = [
1553
+ "COMMIT_FILE",
1554
+ "FEATURE_FILE",
1555
+ "FORMAT",
1556
+ "INDEX_FILE",
1557
+ "SEGMENTS_DIRECTORY",
1558
+ "FeatureStore",
1559
+ "RowAddress",
1560
+ "SegmentReceipt",
1561
+ "SegmentWriter",
1562
+ "StoredFeature",
1563
+ "features_in",
1564
+ "open_feature",
1565
+ "partition_sequences",
1566
+ "sequence_digest",
1567
+ ]
fastplms/features/transactions.py ADDED
@@ -0,0 +1,100 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Process ownership and durable publication for the local feature store.
2
+
3
+ Locks are kernel-owned, not PID files, and release when a process exits. Lock files stay
4
+ in place: unlinking one would let another process lock a different inode at the same path.
5
+ See https://docs.python.org/3/library/fcntl.html and
6
+ https://www.sqlite.org/atomiccommit.html for the locking and flush assumptions.
7
+ """
8
+
9
+ from __future__ import annotations
10
+
11
+ import errno
12
+ import os
13
+ import time
14
+
15
+ from collections.abc import Iterator
16
+ from contextlib import contextmanager
17
+ from pathlib import Path
18
+ from typing import BinaryIO
19
+
20
+
21
+ @contextmanager
22
+ def file_lock(path: Path, *, wait: bool = True) -> Iterator[None]:
23
+ """Hold a process lock on a stable file, optionally refusing a competing owner."""
24
+ path.parent.mkdir(parents=True, exist_ok=True)
25
+ owner_pid = os.getpid()
26
+ with path.open("a+b") as handle:
27
+ if os.name == "nt":
28
+ import msvcrt
29
+
30
+ while True:
31
+ try:
32
+ handle.seek(0)
33
+ msvcrt.locking(handle.fileno(), msvcrt.LK_NBLCK, 1)
34
+ break
35
+ except OSError as error:
36
+ if error.errno not in (errno.EACCES, errno.EAGAIN, errno.EDEADLK):
37
+ raise
38
+ if not wait:
39
+ raise BlockingIOError(
40
+ "Feature segment already has an active writer."
41
+ ) from error
42
+ time.sleep(0.05)
43
+ try:
44
+ yield
45
+ finally:
46
+ if os.getpid() == owner_pid:
47
+ handle.seek(0)
48
+ msvcrt.locking(handle.fileno(), msvcrt.LK_UNLCK, 1)
49
+ else:
50
+ import fcntl
51
+
52
+ flags = fcntl.LOCK_EX | (0 if wait else fcntl.LOCK_NB)
53
+ try:
54
+ fcntl.flock(handle.fileno(), flags)
55
+ except BlockingIOError as error:
56
+ raise BlockingIOError("Feature segment already has an active writer.") from error
57
+ try:
58
+ yield
59
+ finally:
60
+ if os.getpid() == owner_pid:
61
+ fcntl.flock(handle.fileno(), fcntl.LOCK_UN)
62
+
63
+
64
+ def sync_directory(path: Path) -> None:
65
+ """Flush directory entries on POSIX; Windows has no equivalent directory fsync."""
66
+ if os.name != "nt":
67
+ descriptor = os.open(path, os.O_RDONLY | os.O_DIRECTORY)
68
+ try:
69
+ os.fsync(descriptor)
70
+ finally:
71
+ os.close(descriptor)
72
+
73
+
74
+ def flush_and_evict(handle: BinaryIO) -> None:
75
+ """Make a written file durable, then drop its pages from the page cache.
76
+
77
+ A feature store writes far more than it reads back soon, and written pages stay cached. On a GH200 the
78
+ kernel fills the GPU's HBM, which Linux exposes as a NUMA node, with that cache, so ``nvidia-smi`` reads
79
+ the device as full and the next large allocation waits while the kernel evicts. The pages are clean
80
+ after ``fsync``, so ``POSIX_FADV_DONTNEED`` drops all of them. Windows has no ``posix_fadvise`` and
81
+ keeps only the flush. Callers flush Python buffers first.
82
+ """
83
+ descriptor = handle.fileno()
84
+ os.fsync(descriptor)
85
+ if hasattr(os, "posix_fadvise"):
86
+ os.posix_fadvise(descriptor, 0, 0, os.POSIX_FADV_DONTNEED) # offset 0, length 0: the whole file
87
+
88
+
89
+ def publish_file(temporary: Path, destination: Path, *, sync_parent: bool = True) -> None:
90
+ """Flush a complete staged file, rename it, then flush its containing directory.
91
+
92
+ A caller that publishes many files into one directory passes ``sync_parent=False`` and flushes the
93
+ directory once, before the commit marker, so a part costs one flush and not two. The staged file
94
+ leaves the page cache once it is durable (``flush_and_evict``).
95
+ """
96
+ with temporary.open("r+b") as handle:
97
+ flush_and_evict(handle)
98
+ temporary.replace(destination)
99
+ if sync_parent:
100
+ sync_directory(destination.parent)
fastplms/features/writing.py ADDED
@@ -0,0 +1,52 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Commit a window of rows to a feature store as one segment.
2
+
3
+ A pipeline that embeds a window of sequences at a time, and wants a run that dies to lose at most
4
+ that window, does the same thing after every window: keep the sequences the store lacks, write them
5
+ as a segment, and name the segment so a restart cannot commit the same rows under two names.
6
+ ``write_rows`` is that step. The segment is named by the digest of its sequences, so the name
7
+ depends only on what the segment holds, and a window a restart repeats is skipped by calling
8
+ ``FeatureStore.missing`` first, exactly as a run does before it embeds.
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ import hashlib
14
+
15
+ from collections.abc import Mapping
16
+ from typing import Any
17
+
18
+ from .conversion import Row
19
+ from .store import FeatureStore, sequence_digest
20
+
21
+
22
+ DEFAULT_MAX_TENSOR_BYTES = 256 * 1024**2
23
+
24
+
25
+ def write_rows(
26
+ store: FeatureStore,
27
+ rows: Mapping[str, Row],
28
+ *,
29
+ metadata: Mapping[str, Any] | None = None,
30
+ max_tensor_bytes: int = DEFAULT_MAX_TENSOR_BYTES,
31
+ ) -> int:
32
+ """Commit ``rows`` (sequence to row) as one segment and return how many rows it holds.
33
+
34
+ A tensor for a dense or ragged feature, a ``SparseRow`` for csr, and a ``TopKRow`` for ragged
35
+ top-k, as ``SegmentWriter.append`` takes them, in the order the mapping yields. An empty mapping
36
+ writes nothing and returns zero. A sequence the store already holds raises, because a feature
37
+ has one row per sequence; call ``store.missing`` first.
38
+
39
+ ``metadata`` is plain data recorded in the segment's commit marker, for what the run wants to
40
+ remember about these rows. A window too large for ``max_tensor_bytes`` is split into parts of a
41
+ segment that commits or vanishes as one.
42
+ """
43
+
44
+ if not rows:
45
+ return 0
46
+ sequences = list(rows)
47
+ digest = hashlib.sha256("\n".join(sequence_digest(sequence) for sequence in sequences).encode("utf-8"))
48
+ with store.segment("rows-" + digest.hexdigest()[:16], metadata) as writer:
49
+ writer.append_bounded(
50
+ sequences, [rows[sequence] for sequence in sequences], max_tensor_bytes=max_tensor_bytes,
51
+ )
52
+ return len(sequences)
fastplms/json_files.py ADDED
@@ -0,0 +1,45 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """The two JSON text forms FastPLMs writes: compact for hashing, indented for files people read.
2
+
3
+ This file exists twice, byte for byte: here and as ``features/json_files.py``. ``features`` loads as a
4
+ standalone package (``foundry.embedding.private_store``) and imports nothing outside itself, so it cannot
5
+ reach this module. ``tests/tier1_unit/test_features_package_is_self_contained.py`` fails when the two differ.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import json
11
+
12
+ from typing import Any
13
+
14
+
15
+ def compact_json(value: Any, *, ensure_ascii: bool = True, allow_nan: bool = True) -> str:
16
+ """Serialize with sorted keys and no whitespace, the form a digest or an identity is taken over."""
17
+
18
+ return json.dumps(
19
+ value,
20
+ sort_keys=True,
21
+ separators=(",", ":"),
22
+ ensure_ascii=ensure_ascii,
23
+ allow_nan=allow_nan,
24
+ )
25
+
26
+
27
+ def indented_json(
28
+ value: Any,
29
+ *,
30
+ ensure_ascii: bool = True,
31
+ allow_nan: bool = True,
32
+ sort_keys: bool = True,
33
+ ) -> str:
34
+ """Serialize with two-space indentation and one trailing newline, the form of a stored JSON file."""
35
+
36
+ return (
37
+ json.dumps(
38
+ value,
39
+ indent=2,
40
+ sort_keys=sort_keys,
41
+ ensure_ascii=ensure_ascii,
42
+ allow_nan=allow_nan,
43
+ )
44
+ + "\n"
45
+ )
fastplms/models.toml CHANGED
@@ -4,11 +4,58 @@ legal_files = [
4
  "THIRD_PARTY_NOTICES.md=sha256:d35e506b728868b52290672d54a89c6aade09529af99dbc8a4db72d9e9ca3460",
5
  ]
6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7
  [[attention_kernels]]
8
  implementation = "flash_attention_2"
9
  repository = "kernels-community/flash-attn2"
10
- revision = "db6b51744f0cd7061386442c09df890fc6d9f47e"
11
- version = 2
12
  expected_variant = "flash_attn2"
13
  dtypes = ["bfloat16"]
14
  min_cuda_capability = [8, 0]
@@ -64,6 +111,9 @@ distribution_files = [
64
  id = "biohub-transformers"
65
  path = "vendor/upstream/biohub-transformers"
66
  url = "https://github.com/Biohub/transformers.git"
 
 
 
67
  revision = "3a8956fb4d4ea16b0ec8e71deef2c2909b6a5cbf"
68
  license = "Apache-2.0"
69
  license_files = ["LICENSE"]
@@ -177,7 +227,7 @@ conversion_provenance = "Input: the pinned official ESM2 state dictionary. Trans
177
  representative = "esm2_8m"
178
  documentation = "docs/models.md#esm2"
179
  test_tiers = ["check", "compliance", "feature", "artifact", "benchmark"]
180
- runtime_paths = ["__init__.py", "registry.py", "runtime.py", "models.toml", "models/__init__.py", "attention", "embeddings", "models/_esm_rotary.py", "models/esm2", "models/ttt.py"]
181
  auto_map = { AutoConfig = "fastplms.models.esm2.modeling_fastesm.FastEsmConfig", AutoModel = "fastplms.models.esm2.modeling_fastesm.FastEsmModel", AutoModelForMaskedLM = "fastplms.models.esm2.modeling_fastesm.FastEsmForMaskedLM", AutoModelForSequenceClassification = "fastplms.models.esm2.modeling_fastesm.FastEsmForSequenceClassification", AutoModelForTokenClassification = "fastplms.models.esm2.modeling_fastesm.FastEsmForTokenClassification" }
182
 
183
  [families.esm_plusplus]
@@ -204,7 +254,7 @@ conversion_provenance = "Input: the pinned Biohub ESMC checkpoint. Transformatio
204
  representative = "esmc_small"
205
  documentation = "docs/models.md#esm-and-esmc"
206
  test_tiers = ["check", "compliance", "feature", "artifact", "benchmark"]
207
- runtime_paths = ["__init__.py", "registry.py", "runtime.py", "models.toml", "models/__init__.py", "attention", "embeddings", "models/esm_plusplus", "models/ttt.py"]
208
  auto_map = { AutoConfig = "fastplms.models.esm_plusplus.modeling_esm_plusplus.ESMplusplusConfig", AutoModel = "fastplms.models.esm_plusplus.modeling_esm_plusplus.ESMplusplusModel", AutoModelForMaskedLM = "fastplms.models.esm_plusplus.modeling_esm_plusplus.ESMplusplusForMaskedLM", AutoModelForSequenceClassification = "fastplms.models.esm_plusplus.modeling_esm_plusplus.ESMplusplusForSequenceClassification", AutoModelForTokenClassification = "fastplms.models.esm_plusplus.modeling_esm_plusplus.ESMplusplusForTokenClassification" }
209
 
210
  [families.esm3]
@@ -229,7 +279,7 @@ conversion_provenance = "Input: the pinned Biohub ESM3 checkpoint. Transformatio
229
  representative = "esm3_small"
230
  documentation = "docs/models.md#esm3"
231
  test_tiers = ["check", "compliance", "feature", "artifact", "benchmark"]
232
- runtime_paths = ["__init__.py", "registry.py", "runtime.py", "models.toml", "models/__init__.py", "attention", "embeddings", "models/esm3", "models/ttt.py"]
233
  auto_map = { AutoConfig = "fastplms.models.esm3.modeling_esm3.FastESM3Config", AutoModel = "fastplms.models.esm3.modeling_esm3.FastESM3Model", AutoModelForSequenceClassification = "fastplms.models.esm3.modeling_esm3.FastESM3ForSequenceClassification", AutoModelForTokenClassification = "fastplms.models.esm3.modeling_esm3.FastESM3ForTokenClassification" }
234
 
235
  [families.e1]
@@ -256,7 +306,7 @@ conversion_provenance = "Input: the pinned Profluent-E1 checkpoint and tokenizer
256
  representative = "e1_150m"
257
  documentation = "docs/models.md#e1"
258
  test_tiers = ["check", "compliance", "feature", "artifact", "benchmark"]
259
- runtime_paths = ["__init__.py", "registry.py", "runtime.py", "models.toml", "models/__init__.py", "attention", "embeddings", "models/e1", "models/ttt.py"]
260
  auto_map = { AutoConfig = "fastplms.models.e1.modeling_e1.E1Config", AutoModel = "fastplms.models.e1.modeling_e1.E1Model", AutoModelForMaskedLM = "fastplms.models.e1.modeling_e1.E1ForMaskedLM", AutoModelForSequenceClassification = "fastplms.models.e1.modeling_e1.E1ForSequenceClassification", AutoModelForTokenClassification = "fastplms.models.e1.modeling_e1.E1ForTokenClassification" }
261
 
262
  [families.dplm]
@@ -282,7 +332,7 @@ conversion_provenance = "Input: the pinned official DPLM1 checkpoint. Transforma
282
  representative = "dplm_150m"
283
  documentation = "docs/models.md#dplm"
284
  test_tiers = ["check", "compliance", "feature", "artifact", "benchmark"]
285
- runtime_paths = ["__init__.py", "registry.py", "runtime.py", "models.toml", "models/__init__.py", "attention", "embeddings", "models/_diffusion_generation.py", "models/_esm_rotary.py", "models/dplm", "models/ttt.py"]
286
  auto_map = { AutoConfig = "fastplms.models.dplm.modeling_dplm.DPLMConfig", AutoModel = "fastplms.models.dplm.modeling_dplm.DPLMModel", AutoModelForMaskedLM = "fastplms.models.dplm.modeling_dplm.DPLMForMaskedLM", AutoModelForSequenceClassification = "fastplms.models.dplm.modeling_dplm.DPLMForSequenceClassification", AutoModelForTokenClassification = "fastplms.models.dplm.modeling_dplm.DPLMForTokenClassification" }
287
 
288
  [families.dplm2]
@@ -307,7 +357,7 @@ conversion_provenance = "Input: the pinned official DPLM2 checkpoint. Transforma
307
  representative = "dplm2_150m"
308
  documentation = "docs/models.md#dplm2"
309
  test_tiers = ["check", "compliance", "feature", "artifact", "benchmark"]
310
- runtime_paths = ["__init__.py", "registry.py", "runtime.py", "models.toml", "models/__init__.py", "attention", "embeddings", "models/_diffusion_generation.py", "models/_esm_rotary.py", "models/dplm2", "models/ttt.py"]
311
  auto_map = { AutoConfig = "fastplms.models.dplm2.modeling_dplm2.DPLM2Config", AutoModel = "fastplms.models.dplm2.modeling_dplm2.DPLM2Model", AutoModelForMaskedLM = "fastplms.models.dplm2.modeling_dplm2.DPLM2ForMaskedLM", AutoModelForSequenceClassification = "fastplms.models.dplm2.modeling_dplm2.DPLM2ForSequenceClassification", AutoModelForTokenClassification = "fastplms.models.dplm2.modeling_dplm2.DPLM2ForTokenClassification" }
312
  tokenizer_class = "fastplms.models.dplm2.tokenization_dplm2.DPLM2Tokenizer"
313
 
@@ -334,7 +384,7 @@ representative = "ankh_base"
334
  documentation = "docs/models.md#ankh"
335
  test_tiers = ["check", "compliance", "feature", "artifact", "benchmark"]
336
  requires_complete_weight_publication = false
337
- runtime_paths = ["__init__.py", "registry.py", "runtime.py", "models.toml", "models/__init__.py", "attention", "embeddings", "models/ankh", "models/ttt.py"]
338
  auto_map = { AutoConfig = "fastplms.models.ankh.modeling_ankh.FastAnkhConfig", AutoModel = "fastplms.models.ankh.modeling_ankh.FastAnkhModel", AutoModelForMaskedLM = "fastplms.models.ankh.modeling_ankh.FastAnkhForMaskedLMExtension", AutoModelForSeq2SeqLM = "fastplms.models.ankh.modeling_ankh.FastAnkhForConditionalGeneration", AutoModelForSequenceClassification = "fastplms.models.ankh.modeling_ankh.FastAnkhForSequenceClassification", AutoModelForTokenClassification = "fastplms.models.ankh.modeling_ankh.FastAnkhForTokenClassification" }
339
 
340
  [families.boltz2]
@@ -358,7 +408,7 @@ conversion_provenance = "Input: the pinned official Boltz2 checkpoint. Transform
358
  representative = "boltz2"
359
  documentation = "docs/models.md#boltz2"
360
  test_tiers = ["structure", "artifact", "benchmark"]
361
- runtime_paths = ["__init__.py", "registry.py", "runtime.py", "models.toml", "models/__init__.py", "models/boltz"]
362
  auto_map = { AutoConfig = "fastplms.models.boltz.modeling_boltz2.Boltz2Config", AutoModel = "fastplms.models.boltz.modeling_boltz2.Boltz2Model" }
363
 
364
  [families.esmfold]
@@ -383,7 +433,7 @@ conversion_provenance = "Input: the pinned native Meta ESMFold checkpoint plus i
383
  representative = "esmfold"
384
  documentation = "docs/models.md#esmfold"
385
  test_tiers = ["check", "compliance", "structure", "feature", "artifact", "benchmark"]
386
- runtime_paths = ["__init__.py", "registry.py", "runtime.py", "models.toml", "models/__init__.py", "attention", "embeddings", "models/_esm_rotary.py", "models/classification_probe.py", "models/esmfold"]
387
  auto_map = { AutoConfig = "fastplms.models.esmfold.modeling_fast_esmfold.FastEsmFoldConfig", AutoModel = "fastplms.models.esmfold.modeling_fast_esmfold.FastEsmForProteinFolding", AutoModelForSequenceClassification = "fastplms.models.esmfold.modeling_fast_esmfold.FastEsmForSequenceClassification", AutoModelForTokenClassification = "fastplms.models.esmfold.modeling_fast_esmfold.FastEsmForTokenClassification" }
388
 
389
  [families.esmfold2]
@@ -410,7 +460,7 @@ conversion_provenance = "Input: each pinned Biohub ESMFold2 checkpoint and its s
410
  representative = "esmfold2"
411
  documentation = "docs/esmfold2.md"
412
  test_tiers = ["check", "compliance", "structure", "feature", "artifact", "benchmark"]
413
- runtime_paths = ["__init__.py", "registry.py", "runtime.py", "models.toml", "models/__init__.py", "attention", "embeddings", "models/classification_probe.py", "models/_esm_rotary.py", "models/esmfold2", "models/esm_plusplus", "models/ttt.py"]
414
  auto_map = { AutoConfig = "fastplms.models.esmfold2.configuration_esmfold2.ESMFold2Config", AutoModel = "fastplms.models.esmfold2.modeling_esmfold2.ESMFold2Model", AutoModelForSequenceClassification = "fastplms.models.esmfold2.modeling_esmfold2_classification.ESMFold2ForSequenceClassification", AutoModelForTokenClassification = "fastplms.models.esmfold2.modeling_esmfold2_classification.ESMFold2ForTokenClassification" }
415
 
416
  [[models]]
@@ -418,7 +468,7 @@ id = "esm2_8m"
418
  family = "esm2"
419
  size_category = "small"
420
  generation_contract = "not_applicable"
421
- official_golden = { metadata = "tests/goldens/esm2_8m.json=sha256:6975e86d1d8f27488bf2a676551feaa48cc19254c9d24b6acb09198122745609", tensors = "tests/goldens/esm2_8m.safetensors=sha256:b40217566c33c71988d28869de353be54a3b3ebfc21fdfd29056e88cf7e99f4c" }
422
  fast_repo = "Synthyra/ESM2-8M"
423
  fast_revision = "185ecbd45665d050a8dae326d91886d330c5f9d0"
424
  fast_files = [
@@ -457,7 +507,7 @@ id = "esm2_35m"
457
  family = "esm2"
458
  size_category = "small"
459
  generation_contract = "not_applicable"
460
- official_golden = { metadata = "tests/goldens/esm2_35m.json=sha256:e919d3ce6d20b6a942d27d92323814ae7594a0129dc9c4de27c5053e96675bcd", tensors = "tests/goldens/esm2_35m.safetensors=sha256:c9b8bb616cf884fb7744521a2fcc6eed23586342d11241e6c9ef16454ec31e17" }
461
  fast_repo = "Synthyra/ESM2-35M"
462
  fast_revision = "37ab9f56b41e365b3bd9e25d6fefe9150fd910f0"
463
  fast_files = [
@@ -496,7 +546,7 @@ id = "esm2_150m"
496
  family = "esm2"
497
  size_category = "medium"
498
  generation_contract = "not_applicable"
499
- official_golden = { metadata = "tests/goldens/esm2_150m.json=sha256:c04c93486024ba0fa1c81fbfbe92ee79d1d4c7f1cfcc2c9886728522f752feab", tensors = "tests/goldens/esm2_150m.safetensors=sha256:c03fe9916dba137b452a6bbe944c7dc414db4019a6f0921e87b92d4bb6a8a42f" }
500
  fast_repo = "Synthyra/ESM2-150M"
501
  fast_revision = "979e0880dfc9e0c0080839b83d9d2dc05b92786a"
502
  fast_files = [
@@ -535,7 +585,7 @@ id = "esm2_650m"
535
  family = "esm2"
536
  size_category = "large"
537
  generation_contract = "not_applicable"
538
- official_golden = { metadata = "tests/goldens/esm2_650m.json=sha256:f18332172fcb3abf5dd2485fd55f5b0d193ad3b93a44cc744e0d02817c927477", tensors = "tests/goldens/esm2_650m.safetensors=sha256:c3a66b75add03628e62e238cb63da6a9e4d321f8160e84bdf2a131c096977f86" }
539
  fast_repo = "Synthyra/ESM2-650M"
540
  fast_revision = "ca0718a5d52b80d5c60dd76860e55e061a95fb0a"
541
  fast_files = [
@@ -574,7 +624,7 @@ id = "esm2_3b"
574
  family = "esm2"
575
  size_category = "xlarge"
576
  generation_contract = "not_applicable"
577
- official_golden = { metadata = "tests/goldens/esm2_3b.json=sha256:5043b2333c57a34d54fac53916722d1acb4b6fd50395b9abafa805435b184a48", tensors = "tests/goldens/esm2_3b.safetensors=sha256:dfd5a8cb05d3e814a080185c4808c8e7ec2277f070f395562fcfbe4376789e4e" }
578
  notes = "The pinned default SDPA BF16 path uses a checkpoint-specific numeric calibration: relative L2 target/hard limit 0.06/0.07, relative Q99.9 0.15/0.18, first-percentile residue cosine 0.994/0.992, and pooled cosine 0.998/0.997. Exact state identity and the global logits-distribution contract remain required."
579
  fast_repo = "Synthyra/ESM2-3B"
580
  fast_revision = "ff89d0180f414ab9c677219a25da79bf09185456"
@@ -617,7 +667,7 @@ id = "esmc_small"
617
  family = "esm_plusplus"
618
  size_category = "medium"
619
  generation_contract = "not_applicable"
620
- official_golden = { metadata = "tests/goldens/esmc_small.json=sha256:bb02652cf3cc484756b98ffa4ba55ed4c55870d2cea3342adb1d920ba9dfe10a", tensors = "tests/goldens/esmc_small.safetensors=sha256:03378d0f0fdd8161178ebb2c1f0da1b9776a726c8e8d3a10c009808a24de5654" }
621
  notes = "Release contract: SDPA must match the pinned Biohub implementation bit-for-bit across every hidden state, last hidden state, logits, special token, and padding position. Eager and FlashAttention 2 are release-gated in BF16 against the pinned boundary-length and biological panels with a relative-L2 engineering target of 0.029, hard limit of 0.03, relative-Q99.9 target of 0.049, first-percentile residue-cosine target of 0.997, and Jensen-Shannon target of 0.0004. The global pooled-cosine and top-1 thresholds remain unchanged. Flex Attention and FlashAttention 3 remain selectable as opt-in alternatives, but they are not strict-parity choices: on the locked H100 BF16 generated-boundary panel, ESMC-6B Flex Attention exceeds the 0.03 relative-L2 hard limit and FlashAttention 3 falls below the 0.995 residue-cosine hard limit. The deviation is consistent with backend-specific BF16 kernel arithmetic; it is not a weight-conversion difference or silent fallback. Use SDPA for exact Biohub parity or FlashAttention 2 for release-gated acceleration."
622
  fast_repo = "Synthyra/ESMplusplus_small"
623
  fast_revision = "46c5f7d562e47d4c14165b424c71ab7db008e6fb"
@@ -643,7 +693,7 @@ id = "esmc_large"
643
  family = "esm_plusplus"
644
  size_category = "large"
645
  generation_contract = "not_applicable"
646
- official_golden = { metadata = "tests/goldens/esmc_large.json=sha256:7a4d614f67b6fde417f3fd89f61e7ec442ae284769734b2b73e14945a816a8fd", tensors = "tests/goldens/esmc_large.safetensors=sha256:e13302df4cf7e8381552f1043a8fd0f31f3e0d50b2ab6009fb86b7940ae8ff79" }
647
  notes = "Release contract: SDPA must match the pinned Biohub implementation bit-for-bit across every hidden state, last hidden state, logits, special token, and padding position. Eager and FlashAttention 2 are release-gated in BF16 against the pinned boundary-length and biological panels with a relative-L2 engineering target of 0.029, hard limit of 0.03, relative-Q99.9 target of 0.049, first-percentile residue-cosine target of 0.997, and Jensen-Shannon target of 0.0004. The global pooled-cosine and top-1 thresholds remain unchanged. Flex Attention and FlashAttention 3 remain selectable as opt-in alternatives, but they are not strict-parity choices: on the locked H100 BF16 generated-boundary panel, ESMC-6B Flex Attention exceeds the 0.03 relative-L2 hard limit and FlashAttention 3 falls below the 0.995 residue-cosine hard limit. The deviation is consistent with backend-specific BF16 kernel arithmetic; it is not a weight-conversion difference or silent fallback. Use SDPA for exact Biohub parity or FlashAttention 2 for release-gated acceleration."
648
  fast_repo = "Synthyra/ESMplusplus_large"
649
  fast_revision = "f813401638b3fddab09748aec1ad2bf537aa4208"
@@ -669,7 +719,7 @@ id = "esmc_6b"
669
  family = "esm_plusplus"
670
  size_category = "xlarge"
671
  generation_contract = "not_applicable"
672
- official_golden = { metadata = "tests/goldens/esmc_6b.json=sha256:e229d938719782f280fab22dfc4c43e86109fdb0cc523631168c5a491afaace3", tensors = "tests/goldens/esmc_6b.safetensors=sha256:a948945e985c7deaca7be8b7eed09c0a9521a2af3f2b10fc2ec7a7d2a0f99ada" }
673
  notes = "Release contract: SDPA must match the pinned Biohub implementation bit-for-bit across every hidden state, last hidden state, logits, special token, and padding position. Eager and FlashAttention 2 are release-gated in BF16 against the pinned boundary-length and biological panels with a relative-L2 engineering target of 0.029, hard limit of 0.03, relative-Q99.9 target of 0.049, first-percentile residue-cosine target of 0.997, and Jensen-Shannon target of 0.0004. The global pooled-cosine and top-1 thresholds remain unchanged. Flex Attention and FlashAttention 3 remain selectable as opt-in alternatives, but they are not strict-parity choices: on the locked H100 BF16 generated-boundary panel, ESMC-6B Flex Attention exceeds the 0.03 relative-L2 hard limit and FlashAttention 3 falls below the 0.995 residue-cosine hard limit. The deviation is consistent with backend-specific BF16 kernel arithmetic; it is not a weight-conversion difference or silent fallback. Use SDPA for exact Biohub parity or FlashAttention 2 for release-gated acceleration."
674
  fast_repo = "Synthyra/ESMplusplus_6B"
675
  fast_revision = "0d579cce3b0f09efa6b3baddf6cc3fd8c9b616c8"
@@ -681,6 +731,7 @@ fast_files = [
681
  "model-00004-of-00006.safetensors=sha256:e46c6113c89c6f3e9b072c1bef02d763a625c37bcd8f9da2ed9363891c9a0758",
682
  "model-00005-of-00006.safetensors=sha256:6d92cb2bf9791de644de2ae86f8523d802ac3b4aaabfff0716ab6c2b97f6fb14",
683
  "model-00006-of-00006.safetensors=sha256:5fc1a8632490bb34162823c35d0d591337b9e4195b22cc0560741397a6e9d0b3",
 
684
  "special_tokens_map.json=git-sha1:c907ee1dc19b24241749b32d665c291c7e6e8e4b",
685
  "tokenizer.json=git-sha1:f49735e56cebab0e791aeaae777757b7fd114f71",
686
  "tokenizer_config.json=git-sha1:2985ed2b8aa8ecfb1d12f53f47d2b8a44cc21756",
@@ -706,7 +757,7 @@ family = "esm3"
706
  tokenizer_source = "esmc_small"
707
  size_category = "large"
708
  generation_contract = "not_applicable"
709
- official_golden = { metadata = "tests/goldens/esm3_small.json=sha256:5470e8596cbba0e2882647eccbc53c36d8b48b0f3947d1fe0bcea68da1078c32", tensors = "tests/goldens/esm3_small.safetensors=sha256:d957922f810c9ab4c557d80d5aaaf6a3aab79a5a45e4638012a634a4134803b1" }
710
  fast_repo = "Synthyra/ESM3_small"
711
  fast_revision = "7ddb5a740f9e5f93933eb6410c0ee8684bc63ec1"
712
  fast_files = [
@@ -732,7 +783,7 @@ id = "e1_150m"
732
  family = "e1"
733
  size_category = "small"
734
  generation_contract = "not_applicable"
735
- official_golden = { metadata = "tests/goldens/e1_150m.json=sha256:701a64a6ab1a2fec5a427555b6af96232526c15cb3d5b4dc7fb253ac8f20b922", tensors = "tests/goldens/e1_150m.safetensors=sha256:6558bc8f1a7b20629eaaaa6f72601d0c2cdb859a5dc13595549b1773b6e2de41" }
736
  fast_repo = "Synthyra/Profluent-E1-150M"
737
  fast_revision = "7c5f3bbf697226a2e0900db7a100f9201774a907"
738
  fast_files = [
@@ -751,7 +802,7 @@ id = "e1_300m"
751
  family = "e1"
752
  size_category = "medium"
753
  generation_contract = "not_applicable"
754
- official_golden = { metadata = "tests/goldens/e1_300m.json=sha256:d3478f3f5957a0e0377864074dde0107de890019f96cb63548ee17ffb8f3ec3a", tensors = "tests/goldens/e1_300m.safetensors=sha256:92778b9ef95a803ddc84b3e3ca764c59e045872a94bcff0eb0cd47647732c188" }
755
  fast_repo = "Synthyra/Profluent-E1-300M"
756
  fast_revision = "5ef52c0ad2ae2578f40622696b763523810e8e26"
757
  fast_files = [
@@ -770,7 +821,7 @@ id = "e1_600m"
770
  family = "e1"
771
  size_category = "large"
772
  generation_contract = "not_applicable"
773
- official_golden = { metadata = "tests/goldens/e1_600m.json=sha256:914be191c28141c1f84535cdb69ead0588a2057bb19d46c5bc7f3891a3d6739e", tensors = "tests/goldens/e1_600m.safetensors=sha256:22ed8417a4651ded255099f6d15c63c2c40552e700d2b0470d1adfde3a39c513" }
774
  fast_repo = "Synthyra/Profluent-E1-600M"
775
  fast_revision = "6c8bf0ec83b0e0178677c528b101efffd0677742"
776
  fast_files = [
@@ -789,7 +840,7 @@ id = "dplm_150m"
789
  family = "dplm"
790
  size_category = "small"
791
  generation_contract = "required"
792
- official_golden = { metadata = "tests/goldens/dplm_150m.json=sha256:3228551fe3bed951db9ec97347143ec4462ce7c221ac240b7ce7730948c1dc1f", tensors = "tests/goldens/dplm_150m.safetensors=sha256:392992235195beed97ab8359b90a2e11e52f4326606f99a471447bed81d146bd" }
793
  fast_repo = "Synthyra/DPLM-150M"
794
  fast_revision = "90ba742754151a774f3b7ed580170d0a76b3e69d"
795
  fast_files = [
@@ -814,7 +865,7 @@ id = "dplm_650m"
814
  family = "dplm"
815
  size_category = "large"
816
  generation_contract = "required"
817
- official_golden = { metadata = "tests/goldens/dplm_650m.json=sha256:bf58d0ce73aaac7e6fb1923ef3d9adad67122df2a3dd414c3229488ef9587a6d", tensors = "tests/goldens/dplm_650m.safetensors=sha256:073f0a6abea7e48f28c2d921ff8329a28e22627f01979277cb324908a01b3378" }
818
  fast_repo = "Synthyra/DPLM-650M"
819
  fast_revision = "05dc16d97c5c028aed924c9ed681cee4ab609760"
820
  fast_files = [
@@ -839,7 +890,7 @@ id = "dplm_3b"
839
  family = "dplm"
840
  size_category = "xlarge"
841
  generation_contract = "required"
842
- official_golden = { metadata = "tests/goldens/dplm_3b.json=sha256:a5b6df8b9c7b371976892ec1d6c45581a32ad3a6325c6c0a0b3267012848c8ed", tensors = "tests/goldens/dplm_3b.safetensors=sha256:75b0a0854fc391133920b0feaaeb8f69ab7568a88b3759627aca1556c4338c1e" }
843
  fast_repo = "Synthyra/DPLM-3B"
844
  fast_revision = "7d764dd3d70ecf1ac0e64693de64a0064aacac65"
845
  fast_files = [
@@ -869,7 +920,7 @@ id = "dplm2_150m"
869
  family = "dplm2"
870
  size_category = "small"
871
  generation_contract = "required"
872
- official_golden = { metadata = "tests/goldens/dplm2_150m.json=sha256:d269de779ea1503de72c77e7b2e6224afc9797bd945b40c571ff6faec782e4aa", tensors = "tests/goldens/dplm2_150m.safetensors=sha256:17fc26600938ba5364b8ecb96750786d33e9f92bcd4ea4df3e12a389340748eb" }
873
  artifact_source = "official"
874
  canonical_state_sha256 = "82e1751f59052b8de72b082517557db47947e8d9b4ac2f11278369e6c0cbf001"
875
  fast_repo = "Synthyra/DPLM2-150M"
@@ -896,7 +947,7 @@ id = "dplm2_650m"
896
  family = "dplm2"
897
  size_category = "large"
898
  generation_contract = "required"
899
- official_golden = { metadata = "tests/goldens/dplm2_650m.json=sha256:d9a7548f9af657a72d441ca70f27379863724fcce8ddd3da4f672104b7bfb772", tensors = "tests/goldens/dplm2_650m.safetensors=sha256:c4e0e467c252c3ac813363d2d4b17a5e3bd99e75fad315e76d97689b4655ddac" }
900
  artifact_source = "official"
901
  canonical_state_sha256 = "cba76b6602d2258de9fffff953b608d93cb8ef4a9e89b0bbd27e160c81e78bb4"
902
  fast_repo = "Synthyra/DPLM2-650M"
@@ -925,7 +976,7 @@ size_category = "xlarge"
925
  # The pinned public sampler fails before generation because cls_token_id is None.
926
  # State, tokenizer, and inference parity remain required for this checkpoint.
927
  generation_contract = "official_unavailable"
928
- official_golden = { metadata = "tests/goldens/dplm2_3b.json=sha256:d6e0e02af53b13cb129192f06e264758aa21c9ebf4ee82411cf67037082d2329", tensors = "tests/goldens/dplm2_3b.safetensors=sha256:838b11824d08f83bcb0c0b3268e579f3a87dbfb965370cfe5c3f8793b96b1964" }
929
  notes = "The pinned official DPLM2-3B sampler fails before generation, so live generation equivalence cannot be established for this checkpoint. State, tokenizer, and inference parity remain required."
930
  artifact_source = "official"
931
  canonical_state_sha256 = "8c46ec09115dbe6cbfb91d94ab5e906369d57e27fe620a7741c6f8cb1b6ca890"
@@ -958,7 +1009,7 @@ id = "ankh_base"
958
  family = "ankh"
959
  size_category = "medium"
960
  generation_contract = "required"
961
- official_golden = { metadata = "tests/goldens/ankh_base.json=sha256:ebce8d7de821827ee995789c9b38d79252d3b2f76888130b0a8a7eedafaefe2b", tensors = "tests/goldens/ankh_base.safetensors=sha256:f0e78aa15d11749e0c64ff57f9e88c51cec6538a0adf8951f839df70cc708b65" }
962
  notes = "ANKH parity covers the official encoder and sequence-to-sequence heads. AutoModelForMaskedLM exposes the separately named FastPLMs synthesized masked-LM extension and is not an official ANKH head."
963
  artifact_source = "official"
964
  canonical_state_sha256 = "cdd8d30d88e5bf41f44e1eef4470d8e46607aba5f7c7c805b06c035b89c8c16f"
@@ -987,7 +1038,7 @@ id = "ankh_large"
987
  family = "ankh"
988
  size_category = "large"
989
  generation_contract = "required"
990
- official_golden = { metadata = "tests/goldens/ankh_large.json=sha256:59492518b021de5cfaea87d672c9448c8558e99a3443ba2cc7ab544963196ecb", tensors = "tests/goldens/ankh_large.safetensors=sha256:3fb8d3ac27716d15a9ea92aeef6acf2b977bcc887d9b535000539e523673459b" }
991
  notes = "ANKH parity covers the official encoder and sequence-to-sequence heads. AutoModelForMaskedLM exposes the separately named FastPLMs synthesized masked-LM extension and is not an official ANKH head."
992
  artifact_source = "official"
993
  canonical_state_sha256 = "e498a2e9aea76ef784cbe3e596c6b3f5e9a40e209ad837f7e3207099e4d74483"
@@ -1017,7 +1068,7 @@ id = "ankh2_large"
1017
  family = "ankh"
1018
  size_category = "large"
1019
  generation_contract = "required"
1020
- official_golden = { metadata = "tests/goldens/ankh2_large.json=sha256:e8df38994ca1a1e0c598ace34a0b257b264937e4fdbb01bc41544985116b02a4", tensors = "tests/goldens/ankh2_large.safetensors=sha256:25fe1569f55c635fab8fa49c1d62a889a35a2a738bad921f5764a85b58fd4b5d" }
1021
  notes = "ANKH parity covers the official encoder and sequence-to-sequence heads. AutoModelForMaskedLM exposes the separately named FastPLMs synthesized masked-LM extension and is not an official ANKH head."
1022
  artifact_source = "official"
1023
  canonical_state_sha256 = "597c4fe2fa8711f11a25317905f1d62fa92905e55fdd5c0a79614cd9c9d2bca3"
@@ -1049,7 +1100,7 @@ id = "ankh3_large"
1049
  family = "ankh"
1050
  size_category = "large"
1051
  generation_contract = "required"
1052
- official_golden = { metadata = "tests/goldens/ankh3_large.json=sha256:2e5bb05b3baa5baa78f61fef7d2a2c669b0da5dbfaf6b50b12abd3e17253a961", tensors = "tests/goldens/ankh3_large.safetensors=sha256:e5c494ac418e0a2fe7bdad1376676d48960d58ec9e044d19bfffccb8c3288513" }
1053
  notes = "ANKH parity covers the official encoder and sequence-to-sequence heads. AutoModelForMaskedLM exposes the separately named FastPLMs synthesized masked-LM extension and is not an official ANKH head."
1054
  artifact_source = "official"
1055
  canonical_state_sha256 = "60acb7ef86e85dc0c51fc1edf4c8e69a0480049723b6b2c95e6e9faa720c112a"
@@ -1083,7 +1134,7 @@ id = "ankh3_xl"
1083
  family = "ankh"
1084
  size_category = "xlarge"
1085
  generation_contract = "required"
1086
- official_golden = { metadata = "tests/goldens/ankh3_xl.json=sha256:66bb12e033e4163be225d636108a479393228a4f5061015c8af114e766c3c486", tensors = "tests/goldens/ankh3_xl.safetensors=sha256:72d34567d0228cb6f1ee701c578ed4039fead4346e3f161a52e0e74df28dc8ae" }
1087
  notes = "ANKH parity covers the official encoder and sequence-to-sequence heads. AutoModelForMaskedLM exposes the separately named FastPLMs synthesized masked-LM extension and is not an official ANKH head. The official PyTorch shard index is deliberately excluded: the builder verifies every declared source shard directly and writes a new canonical safetensors index."
1088
  artifact_source = "official"
1089
  canonical_state_sha256 = "dd2188e0d2ca65232135714eef6de394239734d843ddae4928c7398685d858e7"
@@ -1176,7 +1227,7 @@ family = "esmfold2"
1176
  size_category = "structure"
1177
  generation_contract = "not_applicable"
1178
  msa_conditioning = true
1179
- official_golden = { metadata = "tests/goldens/esmfold2.json=sha256:f6e0ed1ec400b9a0fcc817db51774be968dc454b7a32645a07c479e42423ab20", tensors = "tests/goldens/esmfold2.safetensors=sha256:e4d6be4344c528e26b13f79a9303549e3de7e582da195c0078db3ce957fad420" }
1180
  fast_repo = "Synthyra/ESMFold2"
1181
  fast_revision = "cd5a0927cec585a778d983b99a8db23d2e9b281e"
1182
  fast_files = [
@@ -1196,7 +1247,7 @@ family = "esmfold2"
1196
  size_category = "structure"
1197
  generation_contract = "not_applicable"
1198
  msa_conditioning = false
1199
- official_golden = { metadata = "tests/goldens/esmfold2_fast.json=sha256:091b004c0b330217b59c12acd6da3d6edaf91e48d95f6d5f40fc20399cef9478", tensors = "tests/goldens/esmfold2_fast.safetensors=sha256:6e2e1cd07401538b4d9df994f82abe7a5b38a01e8d1ee26681e1216d44a81990" }
1200
  fast_repo = "Synthyra/ESMFold2-Fast"
1201
  fast_revision = "407875bfcaa42552bfcb25acd67ee1888b790170"
1202
  fast_files = [
@@ -1216,7 +1267,7 @@ family = "esmfold2"
1216
  size_category = "structure"
1217
  generation_contract = "not_applicable"
1218
  msa_conditioning = true
1219
- official_golden = { metadata = "tests/goldens/esmfold2_experimental_cutoff2025.json=sha256:cfd0e35b2bc468a0dc4f614d3acfa2fce004f96e9ae2433256ed095b829d55cc", tensors = "tests/goldens/esmfold2_experimental_cutoff2025.safetensors=sha256:9347466bbe803b6f5dc82e3356ca6cbbf2c2edd8765f9fd273385bda255019f6" }
1220
  fast_repo = "Synthyra/ESMFold2-Experimental-Cutoff2025"
1221
  fast_revision = "632ff4a9e68f1de78ee956a613267bdcdb5b354d"
1222
  fast_files = [
@@ -1237,7 +1288,7 @@ family = "esmfold2"
1237
  size_category = "structure"
1238
  generation_contract = "not_applicable"
1239
  msa_conditioning = false
1240
- official_golden = { metadata = "tests/goldens/esmfold2_experimental_fast_cutoff2025.json=sha256:1d0b2da4f1579243f37ae04bd4b834b747005cd8e8e7665e00d088123c43afd9", tensors = "tests/goldens/esmfold2_experimental_fast_cutoff2025.safetensors=sha256:516e216d05d7e6bee59e77126d3e595e2bb7821929433f00c259c5d5241964bb" }
1241
  fast_repo = "Synthyra/ESMFold2-Experimental-Fast-Cutoff2025"
1242
  fast_revision = "8f022c2514a6c32692aaca078a8391d6bc6c4bac"
1243
  fast_files = [
@@ -1254,6 +1305,7 @@ auto_map = { AutoConfig = "fastplms.models.esmfold2.configuration_esmfold2.ESMFo
1254
 
1255
  [[models]]
1256
  id = "esmfold2_300"
 
1257
  confidence_adaptation = { release = "v1", head_sha256 = "40fd7f3d82fcefe8ad20ab2b32a37a68a84b54a527a4bce6eb9437bad2b77e31", base_weight_sha256 = "44d6797c5efebf24753d502b40950e0874871c96ceea14f2d7f7e39cebac67fd", donor_repo = "biohub/ESMFold2-Experimental-Fast-Cutoff2025", donor_revision = "74b88548bf19688b8727432db0d698cb2e1d8783", donor_weight_sha256 = "4e903b740ad6ad704ec60881bfd593e0d6c874a630ffa0f0838276e0b665088f", training_url = "https://wandb.ai/lhallee/fastplms-confidence/runs/9558b6d23daf", evaluation_url = "https://huggingface.co/datasets/Synthyra/FastPLMs-artifacts/tree/309f353b07e0e46de4d77a5266d4eddd695538e3/confidence-v2/v2-reproduction-20260922/public/evaluation/esmfold2_300", evidence_path = "docs/evidence/confidence/esmfold2_300-v1.json", frozen_base = { repo = "Synthyra/ESMFold2-300", revision = "a38a62ae930d157484b331c2bf4241684573adba", files = ["config.json=git-sha1:47ec20cf8b234c3b41d6f3ae1bdfe95d4eb4849e", "model.safetensors=sha256:44d6797c5efebf24753d502b40950e0874871c96ceea14f2d7f7e39cebac67fd"] } }
1258
  family = "esmfold2"
1259
  size_category = "structure"
@@ -1273,6 +1325,7 @@ backbone = { repo = "biohub/ESMC-300M-1500000", revision = "56803b6378b82e16c3b2
1273
 
1274
  [[models]]
1275
  id = "esmfold2_600"
 
1276
  confidence_adaptation = { release = "v1", head_sha256 = "e84726a050722e1b722712c87d17a5388bd3699e2520e4a59abb1d828dfb8de7", base_weight_sha256 = "11a53c1b4700b6c62a5a464fc3ec7076160c19e8584157f88e13275116cfd602", donor_repo = "biohub/ESMFold2-Experimental-Fast-Cutoff2025", donor_revision = "74b88548bf19688b8727432db0d698cb2e1d8783", donor_weight_sha256 = "4e903b740ad6ad704ec60881bfd593e0d6c874a630ffa0f0838276e0b665088f", training_url = "https://wandb.ai/lhallee/fastplms-confidence/runs/820d2cfa56c0", evaluation_url = "https://huggingface.co/datasets/Synthyra/FastPLMs-artifacts/tree/6e62186cd36b9047cc4691980076be9f76482192/confidence-v2/v2-reproduction-20260922/public/evaluation/esmfold2_600", evidence_path = "docs/evidence/confidence/esmfold2_600-v1.json", frozen_base = { repo = "Synthyra/ESMFold2-600", revision = "71c67d0b2b73dc245ea7c3cc0d0476439a882d08", files = ["config.json=git-sha1:8e271837cbdada96c4974c8e543f84065e0f06f1", "model.safetensors=sha256:11a53c1b4700b6c62a5a464fc3ec7076160c19e8584157f88e13275116cfd602"] } }
1277
  family = "esmfold2"
1278
  size_category = "structure"
@@ -1289,3 +1342,277 @@ notes = "Experimental Fast model with a frozen 600M ESM++ backbone, 24 folding b
1289
  auto_map = { AutoConfig = "fastplms.models.esmfold2.configuration_esmfold2.ESMFold2Config", AutoModel = "fastplms.models.esmfold2.modeling_esmfold2_experimental.ESMFold2ExperimentalModel", AutoModelForSequenceClassification = "fastplms.models.esmfold2.modeling_esmfold2_classification.ESMFold2ExperimentalForSequenceClassification", AutoModelForTokenClassification = "fastplms.models.esmfold2.modeling_esmfold2_classification.ESMFold2ExperimentalForTokenClassification" }
1290
  backbone_model = "esmc_large"
1291
  backbone = { repo = "biohub/ESMC-600M-1500000", revision = "21af9cc429af76ebda6c48074fb624db4735aaaf", files = ["config.json=git-sha1:ec29f6009b21d710f64bf1c058f3a9710833d692", "model.safetensors=sha256:d6869f5ae0f11e5dc829b195e062e87cfcc2f851a08a5edbaf5d1083ae7f76cc", "tokenizer.json=git-sha1:81c797f56768b22dec0301fa771f018b7e43e98c", "tokenizer_config.json=git-sha1:f49f57b24a1c93bd544974811e8ecbd61b7fae89", "special_tokens_map.json=git-sha1:c907ee1dc19b24241749b32d665c291c7e6e8e4b"] }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4
  "THIRD_PARTY_NOTICES.md=sha256:d35e506b728868b52290672d54a89c6aade09529af99dbc8a4db72d9e9ca3460",
5
  ]
6
 
7
+ # Released hidden-state SAEs used by the canonical suite. The training-backbone revision
8
+ # is not reported upstream; base_model names the selected inference base, not that history.
9
+ [[sparse_autoencoders]]
10
+ id = "esmc_small_sae_depth"
11
+ base_model = "esmc_small"
12
+ layer = 23
13
+ input_width = 960
14
+ k = 64
15
+ codebook_dim = 16384
16
+ input_kind = "hidden_state"
17
+ checkpoint_repo = "biohub/ESMC-300M-sae-layer23-k64-codebook16384"
18
+ checkpoint_revision = "71b054fc726a9f11153a601118f03657ac6e6e80"
19
+ checkpoint_files = [
20
+ "config.json=sha256:e2864c5f9756c052f343310bafeddafc3edc9514ad85981a51a1a3411186434d",
21
+ "layer_23.safetensors=sha256:825257f785600a7a9d462882d5cf3431ebbbba94761434e33b0c4814a7cc6270",
22
+ ]
23
+
24
+ [[sparse_autoencoders]]
25
+ id = "esmc_large_sae_depth"
26
+ base_model = "esmc_large"
27
+ layer = 27
28
+ input_width = 1152
29
+ k = 64
30
+ codebook_dim = 16384
31
+ input_kind = "hidden_state"
32
+ checkpoint_repo = "biohub/ESMC-600M-sae-layer27-k64-codebook16384"
33
+ checkpoint_revision = "b96480999aca49b7684e093f1adfcb09b55a1720"
34
+ checkpoint_files = [
35
+ "config.json=sha256:0e74c8af3ff4ac644f26be308e0fe6da6c75a4a52f0d3f05b0cbbdda27efc9ce",
36
+ "layer_27.safetensors=sha256:5e93c230aa200876cb300377c8e22a141a7e2e316c6a3b962b2a78c4518535db",
37
+ ]
38
+
39
+ [[sparse_autoencoders]]
40
+ id = "esmc_6b_sae_depth"
41
+ base_model = "esmc_6b"
42
+ layer = 60
43
+ input_width = 2560
44
+ k = 64
45
+ codebook_dim = 16384
46
+ input_kind = "hidden_state"
47
+ checkpoint_repo = "biohub/ESMC-6B-sae-layer60-k64-codebook16384"
48
+ checkpoint_revision = "99752fe6e4d25fbb26f887db2d9225e7a577da73"
49
+ checkpoint_files = [
50
+ "config.json=sha256:42567d3f757bea5c6b618abf55ed905ce4304d1eed891a2c1feae8e1ef393fd0",
51
+ "layer_60.safetensors=sha256:ddc1417b42cffe2d7fc2ec31783a04ac88800e6a8f8c8542f1f92d8f4f16094c",
52
+ ]
53
+
54
  [[attention_kernels]]
55
  implementation = "flash_attention_2"
56
  repository = "kernels-community/flash-attn2"
57
+ revision = "81fb77c12b2ad5d69380669b46739d5868614502"
58
+ version = 3
59
  expected_variant = "flash_attn2"
60
  dtypes = ["bfloat16"]
61
  min_cuda_capability = [8, 0]
 
111
  id = "biohub-transformers"
112
  path = "vendor/upstream/biohub-transformers"
113
  url = "https://github.com/Biohub/transformers.git"
114
+ # Biohub/transformers left GitHub by 2026-10-08. Its pinned commit stays reachable through the
115
+ # fork network of huggingface/transformers, which serves the clone and the source archive.
116
+ fetch_url = "https://github.com/huggingface/transformers.git"
117
  revision = "3a8956fb4d4ea16b0ec8e71deef2c2909b6a5cbf"
118
  license = "Apache-2.0"
119
  license_files = ["LICENSE"]
 
227
  representative = "esm2_8m"
228
  documentation = "docs/models.md#esm2"
229
  test_tiers = ["check", "compliance", "feature", "artifact", "benchmark"]
230
+ runtime_paths = ["__init__.py", "registry.py", "runtime.py", "atomic_files.py", "digests.py", "json_files.py", "models.toml", "models/__init__.py", "attention", "embeddings", "features","models/_esm_rotary.py", "models/esm2", "models/ttt.py"]
231
  auto_map = { AutoConfig = "fastplms.models.esm2.modeling_fastesm.FastEsmConfig", AutoModel = "fastplms.models.esm2.modeling_fastesm.FastEsmModel", AutoModelForMaskedLM = "fastplms.models.esm2.modeling_fastesm.FastEsmForMaskedLM", AutoModelForSequenceClassification = "fastplms.models.esm2.modeling_fastesm.FastEsmForSequenceClassification", AutoModelForTokenClassification = "fastplms.models.esm2.modeling_fastesm.FastEsmForTokenClassification" }
232
 
233
  [families.esm_plusplus]
 
254
  representative = "esmc_small"
255
  documentation = "docs/models.md#esm-and-esmc"
256
  test_tiers = ["check", "compliance", "feature", "artifact", "benchmark"]
257
+ runtime_paths = ["__init__.py", "registry.py", "runtime.py", "atomic_files.py", "digests.py", "json_files.py", "models.toml", "models/__init__.py", "attention", "embeddings", "features","models/esm_plusplus", "models/ttt.py"]
258
  auto_map = { AutoConfig = "fastplms.models.esm_plusplus.modeling_esm_plusplus.ESMplusplusConfig", AutoModel = "fastplms.models.esm_plusplus.modeling_esm_plusplus.ESMplusplusModel", AutoModelForMaskedLM = "fastplms.models.esm_plusplus.modeling_esm_plusplus.ESMplusplusForMaskedLM", AutoModelForSequenceClassification = "fastplms.models.esm_plusplus.modeling_esm_plusplus.ESMplusplusForSequenceClassification", AutoModelForTokenClassification = "fastplms.models.esm_plusplus.modeling_esm_plusplus.ESMplusplusForTokenClassification" }
259
 
260
  [families.esm3]
 
279
  representative = "esm3_small"
280
  documentation = "docs/models.md#esm3"
281
  test_tiers = ["check", "compliance", "feature", "artifact", "benchmark"]
282
+ runtime_paths = ["__init__.py", "registry.py", "runtime.py", "atomic_files.py", "digests.py", "json_files.py", "models.toml", "models/__init__.py", "attention", "embeddings", "features","models/esm3", "models/ttt.py"]
283
  auto_map = { AutoConfig = "fastplms.models.esm3.modeling_esm3.FastESM3Config", AutoModel = "fastplms.models.esm3.modeling_esm3.FastESM3Model", AutoModelForSequenceClassification = "fastplms.models.esm3.modeling_esm3.FastESM3ForSequenceClassification", AutoModelForTokenClassification = "fastplms.models.esm3.modeling_esm3.FastESM3ForTokenClassification" }
284
 
285
  [families.e1]
 
306
  representative = "e1_150m"
307
  documentation = "docs/models.md#e1"
308
  test_tiers = ["check", "compliance", "feature", "artifact", "benchmark"]
309
+ runtime_paths = ["__init__.py", "registry.py", "runtime.py", "atomic_files.py", "digests.py", "json_files.py", "models.toml", "models/__init__.py", "attention", "embeddings", "features","models/e1", "models/ttt.py"]
310
  auto_map = { AutoConfig = "fastplms.models.e1.modeling_e1.E1Config", AutoModel = "fastplms.models.e1.modeling_e1.E1Model", AutoModelForMaskedLM = "fastplms.models.e1.modeling_e1.E1ForMaskedLM", AutoModelForSequenceClassification = "fastplms.models.e1.modeling_e1.E1ForSequenceClassification", AutoModelForTokenClassification = "fastplms.models.e1.modeling_e1.E1ForTokenClassification" }
311
 
312
  [families.dplm]
 
332
  representative = "dplm_150m"
333
  documentation = "docs/models.md#dplm"
334
  test_tiers = ["check", "compliance", "feature", "artifact", "benchmark"]
335
+ runtime_paths = ["__init__.py", "registry.py", "runtime.py", "atomic_files.py", "digests.py", "json_files.py", "models.toml", "models/__init__.py", "attention", "embeddings", "features","models/_diffusion_generation.py", "models/_esm_rotary.py", "models/dplm", "models/ttt.py"]
336
  auto_map = { AutoConfig = "fastplms.models.dplm.modeling_dplm.DPLMConfig", AutoModel = "fastplms.models.dplm.modeling_dplm.DPLMModel", AutoModelForMaskedLM = "fastplms.models.dplm.modeling_dplm.DPLMForMaskedLM", AutoModelForSequenceClassification = "fastplms.models.dplm.modeling_dplm.DPLMForSequenceClassification", AutoModelForTokenClassification = "fastplms.models.dplm.modeling_dplm.DPLMForTokenClassification" }
337
 
338
  [families.dplm2]
 
357
  representative = "dplm2_150m"
358
  documentation = "docs/models.md#dplm2"
359
  test_tiers = ["check", "compliance", "feature", "artifact", "benchmark"]
360
+ runtime_paths = ["__init__.py", "registry.py", "runtime.py", "atomic_files.py", "digests.py", "json_files.py", "models.toml", "models/__init__.py", "attention", "embeddings", "features","models/_diffusion_generation.py", "models/_esm_rotary.py", "models/dplm2", "models/ttt.py"]
361
  auto_map = { AutoConfig = "fastplms.models.dplm2.modeling_dplm2.DPLM2Config", AutoModel = "fastplms.models.dplm2.modeling_dplm2.DPLM2Model", AutoModelForMaskedLM = "fastplms.models.dplm2.modeling_dplm2.DPLM2ForMaskedLM", AutoModelForSequenceClassification = "fastplms.models.dplm2.modeling_dplm2.DPLM2ForSequenceClassification", AutoModelForTokenClassification = "fastplms.models.dplm2.modeling_dplm2.DPLM2ForTokenClassification" }
362
  tokenizer_class = "fastplms.models.dplm2.tokenization_dplm2.DPLM2Tokenizer"
363
 
 
384
  documentation = "docs/models.md#ankh"
385
  test_tiers = ["check", "compliance", "feature", "artifact", "benchmark"]
386
  requires_complete_weight_publication = false
387
+ runtime_paths = ["__init__.py", "registry.py", "runtime.py", "atomic_files.py", "digests.py", "json_files.py", "models.toml", "models/__init__.py", "attention", "embeddings", "features","models/ankh", "models/ttt.py"]
388
  auto_map = { AutoConfig = "fastplms.models.ankh.modeling_ankh.FastAnkhConfig", AutoModel = "fastplms.models.ankh.modeling_ankh.FastAnkhModel", AutoModelForMaskedLM = "fastplms.models.ankh.modeling_ankh.FastAnkhForMaskedLMExtension", AutoModelForSeq2SeqLM = "fastplms.models.ankh.modeling_ankh.FastAnkhForConditionalGeneration", AutoModelForSequenceClassification = "fastplms.models.ankh.modeling_ankh.FastAnkhForSequenceClassification", AutoModelForTokenClassification = "fastplms.models.ankh.modeling_ankh.FastAnkhForTokenClassification" }
389
 
390
  [families.boltz2]
 
408
  representative = "boltz2"
409
  documentation = "docs/models.md#boltz2"
410
  test_tiers = ["structure", "artifact", "benchmark"]
411
+ runtime_paths = ["__init__.py", "registry.py", "runtime.py", "atomic_files.py", "digests.py", "json_files.py", "models.toml", "models/__init__.py", "models/boltz"]
412
  auto_map = { AutoConfig = "fastplms.models.boltz.modeling_boltz2.Boltz2Config", AutoModel = "fastplms.models.boltz.modeling_boltz2.Boltz2Model" }
413
 
414
  [families.esmfold]
 
433
  representative = "esmfold"
434
  documentation = "docs/models.md#esmfold"
435
  test_tiers = ["check", "compliance", "structure", "feature", "artifact", "benchmark"]
436
+ runtime_paths = ["__init__.py", "registry.py", "runtime.py", "atomic_files.py", "digests.py", "json_files.py", "models.toml", "models/__init__.py", "attention", "embeddings", "features","models/_esm_rotary.py", "models/classification_probe.py", "models/esmfold"]
437
  auto_map = { AutoConfig = "fastplms.models.esmfold.modeling_fast_esmfold.FastEsmFoldConfig", AutoModel = "fastplms.models.esmfold.modeling_fast_esmfold.FastEsmForProteinFolding", AutoModelForSequenceClassification = "fastplms.models.esmfold.modeling_fast_esmfold.FastEsmForSequenceClassification", AutoModelForTokenClassification = "fastplms.models.esmfold.modeling_fast_esmfold.FastEsmForTokenClassification" }
438
 
439
  [families.esmfold2]
 
460
  representative = "esmfold2"
461
  documentation = "docs/esmfold2.md"
462
  test_tiers = ["check", "compliance", "structure", "feature", "artifact", "benchmark"]
463
+ runtime_paths = ["__init__.py", "registry.py", "runtime.py", "atomic_files.py", "digests.py", "json_files.py", "models.toml", "models/__init__.py", "attention", "embeddings", "features","models/classification_probe.py", "models/_esm_rotary.py", "models/esmfold2", "models/esm_plusplus", "models/ttt.py"]
464
  auto_map = { AutoConfig = "fastplms.models.esmfold2.configuration_esmfold2.ESMFold2Config", AutoModel = "fastplms.models.esmfold2.modeling_esmfold2.ESMFold2Model", AutoModelForSequenceClassification = "fastplms.models.esmfold2.modeling_esmfold2_classification.ESMFold2ForSequenceClassification", AutoModelForTokenClassification = "fastplms.models.esmfold2.modeling_esmfold2_classification.ESMFold2ForTokenClassification" }
465
 
466
  [[models]]
 
468
  family = "esm2"
469
  size_category = "small"
470
  generation_contract = "not_applicable"
471
+ official_golden = { metadata = "tests/goldens/esm2_8m.json=sha256:06b3ccbf45a46a7aed3833b49503567f7534ffe60e6b3e49fcfcf17d04b7237e", tensors = "tests/goldens/esm2_8m.safetensors=sha256:d08a7572cbef20b8b19b545bcb0427b7e9ae19b986d015c559cd5a3a2cfc8aa4" }
472
  fast_repo = "Synthyra/ESM2-8M"
473
  fast_revision = "185ecbd45665d050a8dae326d91886d330c5f9d0"
474
  fast_files = [
 
507
  family = "esm2"
508
  size_category = "small"
509
  generation_contract = "not_applicable"
510
+ official_golden = { metadata = "tests/goldens/esm2_35m.json=sha256:61cd24fd91ef2dc49cbc0f97b14c0d5c97849eca6f36ed8c294295147909cb38", tensors = "tests/goldens/esm2_35m.safetensors=sha256:4f82d10286e16041c2f23365dfdd8508b911864633287dd76580425527c2d922" }
511
  fast_repo = "Synthyra/ESM2-35M"
512
  fast_revision = "37ab9f56b41e365b3bd9e25d6fefe9150fd910f0"
513
  fast_files = [
 
546
  family = "esm2"
547
  size_category = "medium"
548
  generation_contract = "not_applicable"
549
+ official_golden = { metadata = "tests/goldens/esm2_150m.json=sha256:864e1f9c1d8e939e8da050b8617b8a7d6983fbb11a5911275d1571d08e460a79", tensors = "tests/goldens/esm2_150m.safetensors=sha256:20ad986b8e6e09f36d0158f5939d1914e0809b84329695028d767aab93338621" }
550
  fast_repo = "Synthyra/ESM2-150M"
551
  fast_revision = "979e0880dfc9e0c0080839b83d9d2dc05b92786a"
552
  fast_files = [
 
585
  family = "esm2"
586
  size_category = "large"
587
  generation_contract = "not_applicable"
588
+ official_golden = { metadata = "tests/goldens/esm2_650m.json=sha256:b57f18c8803e68e6c621f479fd7a70596fbc416fc4a17d2ab91c639b480006c4", tensors = "tests/goldens/esm2_650m.safetensors=sha256:261a8b71c4b7c90b1f294031558ee0bf00d77a98352021eb78771200c8d40548" }
589
  fast_repo = "Synthyra/ESM2-650M"
590
  fast_revision = "ca0718a5d52b80d5c60dd76860e55e061a95fb0a"
591
  fast_files = [
 
624
  family = "esm2"
625
  size_category = "xlarge"
626
  generation_contract = "not_applicable"
627
+ official_golden = { metadata = "tests/goldens/esm2_3b.json=sha256:aa441786329889d43983811cac155cabd37b38cfd47bc90d415732a295c33c6d", tensors = "tests/goldens/esm2_3b.safetensors=sha256:61080e55e07db4a562a19e9b6b662e71a9fd3d8fee357aa3df00697d2fa58e33" }
628
  notes = "The pinned default SDPA BF16 path uses a checkpoint-specific numeric calibration: relative L2 target/hard limit 0.06/0.07, relative Q99.9 0.15/0.18, first-percentile residue cosine 0.994/0.992, and pooled cosine 0.998/0.997. Exact state identity and the global logits-distribution contract remain required."
629
  fast_repo = "Synthyra/ESM2-3B"
630
  fast_revision = "ff89d0180f414ab9c677219a25da79bf09185456"
 
667
  family = "esm_plusplus"
668
  size_category = "medium"
669
  generation_contract = "not_applicable"
670
+ official_golden = { metadata = "tests/goldens/esmc_small.json=sha256:15cbbf909a12b5b81ff47b5f142babb75ebf9b3d85673905f2320c2a119ae347", tensors = "tests/goldens/esmc_small.safetensors=sha256:98219356e2845cbd10a80715b19c95632956d40567a7a993409c426b56040b06" }
671
  notes = "Release contract: SDPA must match the pinned Biohub implementation bit-for-bit across every hidden state, last hidden state, logits, special token, and padding position. Eager and FlashAttention 2 are release-gated in BF16 against the pinned boundary-length and biological panels with a relative-L2 engineering target of 0.029, hard limit of 0.03, relative-Q99.9 target of 0.049, first-percentile residue-cosine target of 0.997, and Jensen-Shannon target of 0.0004. The global pooled-cosine and top-1 thresholds remain unchanged. Flex Attention and FlashAttention 3 remain selectable as opt-in alternatives, but they are not strict-parity choices: on the locked H100 BF16 generated-boundary panel, ESMC-6B Flex Attention exceeds the 0.03 relative-L2 hard limit and FlashAttention 3 falls below the 0.995 residue-cosine hard limit. The deviation is consistent with backend-specific BF16 kernel arithmetic; it is not a weight-conversion difference or silent fallback. Use SDPA for exact Biohub parity or FlashAttention 2 for release-gated acceleration."
672
  fast_repo = "Synthyra/ESMplusplus_small"
673
  fast_revision = "46c5f7d562e47d4c14165b424c71ab7db008e6fb"
 
693
  family = "esm_plusplus"
694
  size_category = "large"
695
  generation_contract = "not_applicable"
696
+ official_golden = { metadata = "tests/goldens/esmc_large.json=sha256:86ffe179a9feafd067793f440d25e3606f184f0d32063aeaeb1842b14aa55d3b", tensors = "tests/goldens/esmc_large.safetensors=sha256:2b1ce3de5a7a27f171055f5c7aec4e7e50f0e809a1dc46508aaa0b472365015d" }
697
  notes = "Release contract: SDPA must match the pinned Biohub implementation bit-for-bit across every hidden state, last hidden state, logits, special token, and padding position. Eager and FlashAttention 2 are release-gated in BF16 against the pinned boundary-length and biological panels with a relative-L2 engineering target of 0.029, hard limit of 0.03, relative-Q99.9 target of 0.049, first-percentile residue-cosine target of 0.997, and Jensen-Shannon target of 0.0004. The global pooled-cosine and top-1 thresholds remain unchanged. Flex Attention and FlashAttention 3 remain selectable as opt-in alternatives, but they are not strict-parity choices: on the locked H100 BF16 generated-boundary panel, ESMC-6B Flex Attention exceeds the 0.03 relative-L2 hard limit and FlashAttention 3 falls below the 0.995 residue-cosine hard limit. The deviation is consistent with backend-specific BF16 kernel arithmetic; it is not a weight-conversion difference or silent fallback. Use SDPA for exact Biohub parity or FlashAttention 2 for release-gated acceleration."
698
  fast_repo = "Synthyra/ESMplusplus_large"
699
  fast_revision = "f813401638b3fddab09748aec1ad2bf537aa4208"
 
719
  family = "esm_plusplus"
720
  size_category = "xlarge"
721
  generation_contract = "not_applicable"
722
+ official_golden = { metadata = "tests/goldens/esmc_6b.json=sha256:7e044e7e8d106d43169c4148521c54d0d09bb989624543d648247a110f1907d4", tensors = "tests/goldens/esmc_6b.safetensors=sha256:46139de6b28ab4f9e3244fb6ef3f259fc5e6518a3825f0fb4fd4a4fece6da8e6" }
723
  notes = "Release contract: SDPA must match the pinned Biohub implementation bit-for-bit across every hidden state, last hidden state, logits, special token, and padding position. Eager and FlashAttention 2 are release-gated in BF16 against the pinned boundary-length and biological panels with a relative-L2 engineering target of 0.029, hard limit of 0.03, relative-Q99.9 target of 0.049, first-percentile residue-cosine target of 0.997, and Jensen-Shannon target of 0.0004. The global pooled-cosine and top-1 thresholds remain unchanged. Flex Attention and FlashAttention 3 remain selectable as opt-in alternatives, but they are not strict-parity choices: on the locked H100 BF16 generated-boundary panel, ESMC-6B Flex Attention exceeds the 0.03 relative-L2 hard limit and FlashAttention 3 falls below the 0.995 residue-cosine hard limit. The deviation is consistent with backend-specific BF16 kernel arithmetic; it is not a weight-conversion difference or silent fallback. Use SDPA for exact Biohub parity or FlashAttention 2 for release-gated acceleration."
724
  fast_repo = "Synthyra/ESMplusplus_6B"
725
  fast_revision = "0d579cce3b0f09efa6b3baddf6cc3fd8c9b616c8"
 
731
  "model-00004-of-00006.safetensors=sha256:e46c6113c89c6f3e9b072c1bef02d763a625c37bcd8f9da2ed9363891c9a0758",
732
  "model-00005-of-00006.safetensors=sha256:6d92cb2bf9791de644de2ae86f8523d802ac3b4aaabfff0716ab6c2b97f6fb14",
733
  "model-00006-of-00006.safetensors=sha256:5fc1a8632490bb34162823c35d0d591337b9e4195b22cc0560741397a6e9d0b3",
734
+ "model.safetensors.index.json=git-sha1:f30f0b6b35e11d09a516c078b0847ee023bd1817",
735
  "special_tokens_map.json=git-sha1:c907ee1dc19b24241749b32d665c291c7e6e8e4b",
736
  "tokenizer.json=git-sha1:f49735e56cebab0e791aeaae777757b7fd114f71",
737
  "tokenizer_config.json=git-sha1:2985ed2b8aa8ecfb1d12f53f47d2b8a44cc21756",
 
757
  tokenizer_source = "esmc_small"
758
  size_category = "large"
759
  generation_contract = "not_applicable"
760
+ official_golden = { metadata = "tests/goldens/esm3_small.json=sha256:e476bad7ccdfa2f908261a03e20116c0854eece292a553196f940b594088204d", tensors = "tests/goldens/esm3_small.safetensors=sha256:251e050926a1f2426401bac6b93b4cb00041d0b0a5c8a458d0ef0e6a5d0d87c3" }
761
  fast_repo = "Synthyra/ESM3_small"
762
  fast_revision = "7ddb5a740f9e5f93933eb6410c0ee8684bc63ec1"
763
  fast_files = [
 
783
  family = "e1"
784
  size_category = "small"
785
  generation_contract = "not_applicable"
786
+ official_golden = { metadata = "tests/goldens/e1_150m.json=sha256:e7520331135fd4c98f52f21c470dbba5cb5f4defb418494c63199fb000621e65", tensors = "tests/goldens/e1_150m.safetensors=sha256:03a2e93e7b3e54b12f99eea7678ff80828922f92c3683d4909346a05d74815bc" }
787
  fast_repo = "Synthyra/Profluent-E1-150M"
788
  fast_revision = "7c5f3bbf697226a2e0900db7a100f9201774a907"
789
  fast_files = [
 
802
  family = "e1"
803
  size_category = "medium"
804
  generation_contract = "not_applicable"
805
+ official_golden = { metadata = "tests/goldens/e1_300m.json=sha256:d97637bf789fed61b4ad575f904f37befcea42e9436e8d841e8f8f2cff1d1b17", tensors = "tests/goldens/e1_300m.safetensors=sha256:d125ea866899d2788330d6877d77c48c69febbf7f01d629f4107f9647cfa13ff" }
806
  fast_repo = "Synthyra/Profluent-E1-300M"
807
  fast_revision = "5ef52c0ad2ae2578f40622696b763523810e8e26"
808
  fast_files = [
 
821
  family = "e1"
822
  size_category = "large"
823
  generation_contract = "not_applicable"
824
+ official_golden = { metadata = "tests/goldens/e1_600m.json=sha256:0c9afc83b96ed8ad7f339a26df4ba5853e7994ad54dd3cd837deb9869c0cd2a9", tensors = "tests/goldens/e1_600m.safetensors=sha256:d470d6833f0b38bec070aa37e9f3c49717ea82d31c794af843550bb7e56253a6" }
825
  fast_repo = "Synthyra/Profluent-E1-600M"
826
  fast_revision = "6c8bf0ec83b0e0178677c528b101efffd0677742"
827
  fast_files = [
 
840
  family = "dplm"
841
  size_category = "small"
842
  generation_contract = "required"
843
+ official_golden = { metadata = "tests/goldens/dplm_150m.json=sha256:541167ffcd12b2d7e101046b2d17e800c70ba5db1b561212575301f2d1c407f8", tensors = "tests/goldens/dplm_150m.safetensors=sha256:f8e4ddd580d4708dad2d7da002735b2a01d69ece478c46c4932eb2f3933dc836" }
844
  fast_repo = "Synthyra/DPLM-150M"
845
  fast_revision = "90ba742754151a774f3b7ed580170d0a76b3e69d"
846
  fast_files = [
 
865
  family = "dplm"
866
  size_category = "large"
867
  generation_contract = "required"
868
+ official_golden = { metadata = "tests/goldens/dplm_650m.json=sha256:2bd482c94b1b6b3ac901b5ff9b06e2d189cda1a8fe2f8ff32e5d52b64c2c446b", tensors = "tests/goldens/dplm_650m.safetensors=sha256:70e2c46ed722942d3b628b556b92c9994688cb4853b278eb1e40a3dfbbd9dfc4" }
869
  fast_repo = "Synthyra/DPLM-650M"
870
  fast_revision = "05dc16d97c5c028aed924c9ed681cee4ab609760"
871
  fast_files = [
 
890
  family = "dplm"
891
  size_category = "xlarge"
892
  generation_contract = "required"
893
+ official_golden = { metadata = "tests/goldens/dplm_3b.json=sha256:0f8d8fb7df562c30abf730abe55db4a8c13d40877dd21772375fe1f34030159a", tensors = "tests/goldens/dplm_3b.safetensors=sha256:6942218079232a8185b1ae6e1578474c7784722e28f528b76ec89ca77777e2b9" }
894
  fast_repo = "Synthyra/DPLM-3B"
895
  fast_revision = "7d764dd3d70ecf1ac0e64693de64a0064aacac65"
896
  fast_files = [
 
920
  family = "dplm2"
921
  size_category = "small"
922
  generation_contract = "required"
923
+ official_golden = { metadata = "tests/goldens/dplm2_150m.json=sha256:5f2c496e4b557faf70bdde03b531f753fc7f20c5868bab79dcf378ee358781a1", tensors = "tests/goldens/dplm2_150m.safetensors=sha256:7390559bc54f27f08972c8f0adcf17954065450d937b3984303c1ae1ec003855" }
924
  artifact_source = "official"
925
  canonical_state_sha256 = "82e1751f59052b8de72b082517557db47947e8d9b4ac2f11278369e6c0cbf001"
926
  fast_repo = "Synthyra/DPLM2-150M"
 
947
  family = "dplm2"
948
  size_category = "large"
949
  generation_contract = "required"
950
+ official_golden = { metadata = "tests/goldens/dplm2_650m.json=sha256:af7b8c4fab48773e1ddd2cfd941cf3d518ea27ed8519f467ddf0da788a0aa6be", tensors = "tests/goldens/dplm2_650m.safetensors=sha256:c45f4ca267183b11dd06244ed3c3f39cf5b6139f778733fa63a6e2b5a46b0678" }
951
  artifact_source = "official"
952
  canonical_state_sha256 = "cba76b6602d2258de9fffff953b608d93cb8ef4a9e89b0bbd27e160c81e78bb4"
953
  fast_repo = "Synthyra/DPLM2-650M"
 
976
  # The pinned public sampler fails before generation because cls_token_id is None.
977
  # State, tokenizer, and inference parity remain required for this checkpoint.
978
  generation_contract = "official_unavailable"
979
+ official_golden = { metadata = "tests/goldens/dplm2_3b.json=sha256:ded7978cf84da1d56d8b4693418030ec3b0086bf6466be2242fb0ce424296b9f", tensors = "tests/goldens/dplm2_3b.safetensors=sha256:7490f771f8b8f8137b333f98d6b7dbd2cb48708c8649d4459279a6711b9e5870" }
980
  notes = "The pinned official DPLM2-3B sampler fails before generation, so live generation equivalence cannot be established for this checkpoint. State, tokenizer, and inference parity remain required."
981
  artifact_source = "official"
982
  canonical_state_sha256 = "8c46ec09115dbe6cbfb91d94ab5e906369d57e27fe620a7741c6f8cb1b6ca890"
 
1009
  family = "ankh"
1010
  size_category = "medium"
1011
  generation_contract = "required"
1012
+ official_golden = { metadata = "tests/goldens/ankh_base.json=sha256:21f129b7b71c026cbd3b3fdb712af06d5b6f256b487ef8b01fcc506d2bc6c7b6", tensors = "tests/goldens/ankh_base.safetensors=sha256:449db3638117d5b6027380fdf4b05698bb156024b2b3f7e0ceb295a86e5ec756" }
1013
  notes = "ANKH parity covers the official encoder and sequence-to-sequence heads. AutoModelForMaskedLM exposes the separately named FastPLMs synthesized masked-LM extension and is not an official ANKH head."
1014
  artifact_source = "official"
1015
  canonical_state_sha256 = "cdd8d30d88e5bf41f44e1eef4470d8e46607aba5f7c7c805b06c035b89c8c16f"
 
1038
  family = "ankh"
1039
  size_category = "large"
1040
  generation_contract = "required"
1041
+ official_golden = { metadata = "tests/goldens/ankh_large.json=sha256:8f269215ee29887661abb6c139139dd1782f5112b94eb63be968282a7f050e1b", tensors = "tests/goldens/ankh_large.safetensors=sha256:4a7aea66e91cf880b0b839495d7704831725b672fe169f8e7258b5643d77aff9" }
1042
  notes = "ANKH parity covers the official encoder and sequence-to-sequence heads. AutoModelForMaskedLM exposes the separately named FastPLMs synthesized masked-LM extension and is not an official ANKH head."
1043
  artifact_source = "official"
1044
  canonical_state_sha256 = "e498a2e9aea76ef784cbe3e596c6b3f5e9a40e209ad837f7e3207099e4d74483"
 
1068
  family = "ankh"
1069
  size_category = "large"
1070
  generation_contract = "required"
1071
+ official_golden = { metadata = "tests/goldens/ankh2_large.json=sha256:bdafcd179bc055ea223de43228801b1b6c35196741ac4cc4cca0571e3291e6dc", tensors = "tests/goldens/ankh2_large.safetensors=sha256:44d5fc5c74a0d166ceac54d12505a109f8338cb7090af2f37e39bd78545a3c52" }
1072
  notes = "ANKH parity covers the official encoder and sequence-to-sequence heads. AutoModelForMaskedLM exposes the separately named FastPLMs synthesized masked-LM extension and is not an official ANKH head."
1073
  artifact_source = "official"
1074
  canonical_state_sha256 = "597c4fe2fa8711f11a25317905f1d62fa92905e55fdd5c0a79614cd9c9d2bca3"
 
1100
  family = "ankh"
1101
  size_category = "large"
1102
  generation_contract = "required"
1103
+ official_golden = { metadata = "tests/goldens/ankh3_large.json=sha256:ab3a671260bdad481635d4a4be1b8072d5a2e02e3df1178202713c9a7e56af49", tensors = "tests/goldens/ankh3_large.safetensors=sha256:cc55274281e87d22cf7828117678df2e37a1ece543f5b401b65a0cd6f8b6fa27" }
1104
  notes = "ANKH parity covers the official encoder and sequence-to-sequence heads. AutoModelForMaskedLM exposes the separately named FastPLMs synthesized masked-LM extension and is not an official ANKH head."
1105
  artifact_source = "official"
1106
  canonical_state_sha256 = "60acb7ef86e85dc0c51fc1edf4c8e69a0480049723b6b2c95e6e9faa720c112a"
 
1134
  family = "ankh"
1135
  size_category = "xlarge"
1136
  generation_contract = "required"
1137
+ official_golden = { metadata = "tests/goldens/ankh3_xl.json=sha256:a2d302227ee616698af18502d03e2b6b136d589134058d82b018520f18498431", tensors = "tests/goldens/ankh3_xl.safetensors=sha256:d18c086a4849b220761019f219590fe0c1c6f17698c365ccb0a5a8804446286c" }
1138
  notes = "ANKH parity covers the official encoder and sequence-to-sequence heads. AutoModelForMaskedLM exposes the separately named FastPLMs synthesized masked-LM extension and is not an official ANKH head. The official PyTorch shard index is deliberately excluded: the builder verifies every declared source shard directly and writes a new canonical safetensors index."
1139
  artifact_source = "official"
1140
  canonical_state_sha256 = "dd2188e0d2ca65232135714eef6de394239734d843ddae4928c7398685d858e7"
 
1227
  size_category = "structure"
1228
  generation_contract = "not_applicable"
1229
  msa_conditioning = true
1230
+ official_golden = { metadata = "tests/goldens/esmfold2.json=sha256:60b9e8827615a73acd89ac23f13e3d0a412c17879fbf3308444546e6bbb4ff03", tensors = "tests/goldens/esmfold2.safetensors=sha256:a1a9f3b7a9f7e36ef6a7077d76fa1b1c8b5529ed42ac2cfa485797162ce866fb" }
1231
  fast_repo = "Synthyra/ESMFold2"
1232
  fast_revision = "cd5a0927cec585a778d983b99a8db23d2e9b281e"
1233
  fast_files = [
 
1247
  size_category = "structure"
1248
  generation_contract = "not_applicable"
1249
  msa_conditioning = false
1250
+ official_golden = { metadata = "tests/goldens/esmfold2_fast.json=sha256:659bb338be4584787faff619d1cbcb8766aece7b5facf8a84470e850299b662f", tensors = "tests/goldens/esmfold2_fast.safetensors=sha256:8da91024a2cc63984585267a18f595892ebe489b2a05c4f79a994c3ac5de2f40" }
1251
  fast_repo = "Synthyra/ESMFold2-Fast"
1252
  fast_revision = "407875bfcaa42552bfcb25acd67ee1888b790170"
1253
  fast_files = [
 
1267
  size_category = "structure"
1268
  generation_contract = "not_applicable"
1269
  msa_conditioning = true
1270
+ official_golden = { metadata = "tests/goldens/esmfold2_experimental_cutoff2025.json=sha256:2e42d958f7a99ca8edaf0d9fac8e1a58b78200df652116b7bc102ac799123633", tensors = "tests/goldens/esmfold2_experimental_cutoff2025.safetensors=sha256:8168f8692c4932e1160e733b15210836120223615f6190fb0ecfa2f940c2420f" }
1271
  fast_repo = "Synthyra/ESMFold2-Experimental-Cutoff2025"
1272
  fast_revision = "632ff4a9e68f1de78ee956a613267bdcdb5b354d"
1273
  fast_files = [
 
1288
  size_category = "structure"
1289
  generation_contract = "not_applicable"
1290
  msa_conditioning = false
1291
+ official_golden = { metadata = "tests/goldens/esmfold2_experimental_fast_cutoff2025.json=sha256:d41c54ca28270fc1c1b95485c5009990bd2c215d35ec5cc2c2fa611d096bad21", tensors = "tests/goldens/esmfold2_experimental_fast_cutoff2025.safetensors=sha256:97b58525aefd21c23ec6cb2defa3b3b41982c872dfe70d94306648edbe40a5b9" }
1292
  fast_repo = "Synthyra/ESMFold2-Experimental-Fast-Cutoff2025"
1293
  fast_revision = "8f022c2514a6c32692aaca078a8391d6bc6c4bac"
1294
  fast_files = [
 
1305
 
1306
  [[models]]
1307
  id = "esmfold2_300"
1308
+ official_golden = { metadata = "tests/goldens/esmfold2_300.json=sha256:d4fd12f27352c53582bb40a8d85eab76d6d98f647ce7660a6f891a1dfe68c039", tensors = "tests/goldens/esmfold2_300.safetensors=sha256:40a954ff14edc7c2bff95a252241f08db4cae3821ca9b92ae216775cceed8e8b" }
1309
  confidence_adaptation = { release = "v1", head_sha256 = "40fd7f3d82fcefe8ad20ab2b32a37a68a84b54a527a4bce6eb9437bad2b77e31", base_weight_sha256 = "44d6797c5efebf24753d502b40950e0874871c96ceea14f2d7f7e39cebac67fd", donor_repo = "biohub/ESMFold2-Experimental-Fast-Cutoff2025", donor_revision = "74b88548bf19688b8727432db0d698cb2e1d8783", donor_weight_sha256 = "4e903b740ad6ad704ec60881bfd593e0d6c874a630ffa0f0838276e0b665088f", training_url = "https://wandb.ai/lhallee/fastplms-confidence/runs/9558b6d23daf", evaluation_url = "https://huggingface.co/datasets/Synthyra/FastPLMs-artifacts/tree/309f353b07e0e46de4d77a5266d4eddd695538e3/confidence-v2/v2-reproduction-20260922/public/evaluation/esmfold2_300", evidence_path = "docs/evidence/confidence/esmfold2_300-v1.json", frozen_base = { repo = "Synthyra/ESMFold2-300", revision = "a38a62ae930d157484b331c2bf4241684573adba", files = ["config.json=git-sha1:47ec20cf8b234c3b41d6f3ae1bdfe95d4eb4849e", "model.safetensors=sha256:44d6797c5efebf24753d502b40950e0874871c96ceea14f2d7f7e39cebac67fd"] } }
1310
  family = "esmfold2"
1311
  size_category = "structure"
 
1325
 
1326
  [[models]]
1327
  id = "esmfold2_600"
1328
+ official_golden = { metadata = "tests/goldens/esmfold2_600.json=sha256:ea249ff118975f29979143ea081cfa1c0fb60cb1fa393dd99a12caf413524230", tensors = "tests/goldens/esmfold2_600.safetensors=sha256:69bbb8d53a816b29e979e4469728df0a6b662e4d0be81a76068b36974d2b909e" }
1329
  confidence_adaptation = { release = "v1", head_sha256 = "e84726a050722e1b722712c87d17a5388bd3699e2520e4a59abb1d828dfb8de7", base_weight_sha256 = "11a53c1b4700b6c62a5a464fc3ec7076160c19e8584157f88e13275116cfd602", donor_repo = "biohub/ESMFold2-Experimental-Fast-Cutoff2025", donor_revision = "74b88548bf19688b8727432db0d698cb2e1d8783", donor_weight_sha256 = "4e903b740ad6ad704ec60881bfd593e0d6c874a630ffa0f0838276e0b665088f", training_url = "https://wandb.ai/lhallee/fastplms-confidence/runs/820d2cfa56c0", evaluation_url = "https://huggingface.co/datasets/Synthyra/FastPLMs-artifacts/tree/6e62186cd36b9047cc4691980076be9f76482192/confidence-v2/v2-reproduction-20260922/public/evaluation/esmfold2_600", evidence_path = "docs/evidence/confidence/esmfold2_600-v1.json", frozen_base = { repo = "Synthyra/ESMFold2-600", revision = "71c67d0b2b73dc245ea7c3cc0d0476439a882d08", files = ["config.json=git-sha1:8e271837cbdada96c4974c8e543f84065e0f06f1", "model.safetensors=sha256:11a53c1b4700b6c62a5a464fc3ec7076160c19e8584157f88e13275116cfd602"] } }
1330
  family = "esmfold2"
1331
  size_category = "structure"
 
1342
  auto_map = { AutoConfig = "fastplms.models.esmfold2.configuration_esmfold2.ESMFold2Config", AutoModel = "fastplms.models.esmfold2.modeling_esmfold2_experimental.ESMFold2ExperimentalModel", AutoModelForSequenceClassification = "fastplms.models.esmfold2.modeling_esmfold2_classification.ESMFold2ExperimentalForSequenceClassification", AutoModelForTokenClassification = "fastplms.models.esmfold2.modeling_esmfold2_classification.ESMFold2ExperimentalForTokenClassification" }
1343
  backbone_model = "esmc_large"
1344
  backbone = { repo = "biohub/ESMC-600M-1500000", revision = "21af9cc429af76ebda6c48074fb624db4735aaaf", files = ["config.json=git-sha1:ec29f6009b21d710f64bf1c058f3a9710833d692", "model.safetensors=sha256:d6869f5ae0f11e5dc829b195e062e87cfcc2f851a08a5edbaf5d1083ae7f76cc", "tokenizer.json=git-sha1:81c797f56768b22dec0301fa771f018b7e43e98c", "tokenizer_config.json=git-sha1:f49f57b24a1c93bd544974811e8ecbd61b7fae89", "special_tokens_map.json=git-sha1:c907ee1dc19b24241749b32d665c291c7e6e8e4b"] }
1345
+
1346
+ # Parity goldens, uploaded 2026-09-22 so the reference tensors live somewhere durable
1347
+ # rather than only inside the GitHub repository. Each sha256 is the one already recorded
1348
+ # in that model's official_golden.tensors, so the existing integrity check is unchanged.
1349
+
1350
+ [[golden_artifacts]]
1351
+ id = "esm2_8m"
1352
+ repository = "Synthyra/fastplms_parity_goldens"
1353
+ revision = "50b93a0fc2be21a36c2522118754c74c4a9631ad"
1354
+ path = "goldens/esm2_8m.safetensors"
1355
+ sha256 = "d08a7572cbef20b8b19b545bcb0427b7e9ae19b986d015c559cd5a3a2cfc8aa4"
1356
+ size = 537589
1357
+ offline_behavior = "requires_cached_verified_file"
1358
+
1359
+ [[golden_artifacts]]
1360
+ id = "esm2_35m"
1361
+ repository = "Synthyra/fastplms_parity_goldens"
1362
+ revision = "50b93a0fc2be21a36c2522118754c74c4a9631ad"
1363
+ path = "goldens/esm2_35m.safetensors"
1364
+ sha256 = "4f82d10286e16041c2f23365dfdd8508b911864633287dd76580425527c2d922"
1365
+ size = 779509
1366
+ offline_behavior = "requires_cached_verified_file"
1367
+
1368
+ [[golden_artifacts]]
1369
+ id = "esm2_150m"
1370
+ repository = "Synthyra/fastplms_parity_goldens"
1371
+ revision = "50b93a0fc2be21a36c2522118754c74c4a9631ad"
1372
+ path = "goldens/esm2_150m.safetensors"
1373
+ sha256 = "20ad986b8e6e09f36d0158f5939d1914e0809b84329695028d767aab93338621"
1374
+ size = 1021437
1375
+ offline_behavior = "requires_cached_verified_file"
1376
+
1377
+ [[golden_artifacts]]
1378
+ id = "esm2_650m"
1379
+ repository = "Synthyra/fastplms_parity_goldens"
1380
+ revision = "50b93a0fc2be21a36c2522118754c74c4a9631ad"
1381
+ path = "goldens/esm2_650m.safetensors"
1382
+ sha256 = "261a8b71c4b7c90b1f294031558ee0bf00d77a98352021eb78771200c8d40548"
1383
+ size = 1989117
1384
+ offline_behavior = "requires_cached_verified_file"
1385
+
1386
+ [[golden_artifacts]]
1387
+ id = "esm2_3b"
1388
+ repository = "Synthyra/fastplms_parity_goldens"
1389
+ revision = "50b93a0fc2be21a36c2522118754c74c4a9631ad"
1390
+ path = "goldens/esm2_3b.safetensors"
1391
+ sha256 = "61080e55e07db4a562a19e9b6b662e71a9fd3d8fee357aa3df00697d2fa58e33"
1392
+ size = 3924485
1393
+ offline_behavior = "requires_cached_verified_file"
1394
+
1395
+ [[golden_artifacts]]
1396
+ id = "esmc_small"
1397
+ repository = "Synthyra/fastplms_parity_goldens"
1398
+ revision = "50b93a0fc2be21a36c2522118754c74c4a9631ad"
1399
+ path = "goldens/esmc_small.safetensors"
1400
+ sha256 = "98219356e2845cbd10a80715b19c95632956d40567a7a993409c426b56040b06"
1401
+ size = 1165354
1402
+ offline_behavior = "requires_cached_verified_file"
1403
+
1404
+ [[golden_artifacts]]
1405
+ id = "esmc_large"
1406
+ repository = "Synthyra/fastplms_parity_goldens"
1407
+ revision = "50b93a0fc2be21a36c2522118754c74c4a9631ad"
1408
+ path = "goldens/esmc_large.safetensors"
1409
+ sha256 = "2b1ce3de5a7a27f171055f5c7aec4e7e50f0e809a1dc46508aaa0b472365015d"
1410
+ size = 1383082
1411
+ offline_behavior = "requires_cached_verified_file"
1412
+
1413
+ [[golden_artifacts]]
1414
+ id = "esmc_6b"
1415
+ repository = "Synthyra/fastplms_parity_goldens"
1416
+ revision = "50b93a0fc2be21a36c2522118754c74c4a9631ad"
1417
+ path = "goldens/esmc_6b.safetensors"
1418
+ sha256 = "46139de6b28ab4f9e3244fb6ef3f259fc5e6518a3825f0fb4fd4a4fece6da8e6"
1419
+ size = 2979762
1420
+ offline_behavior = "requires_cached_verified_file"
1421
+
1422
+ [[golden_artifacts]]
1423
+ id = "esm3_small"
1424
+ repository = "Synthyra/fastplms_parity_goldens"
1425
+ revision = "50b93a0fc2be21a36c2522118754c74c4a9631ad"
1426
+ path = "goldens/esm3_small.safetensors"
1427
+ sha256 = "251e050926a1f2426401bac6b93b4cb00041d0b0a5c8a458d0ef0e6a5d0d87c3"
1428
+ size = 2398877
1429
+ offline_behavior = "requires_cached_verified_file"
1430
+
1431
+ [[golden_artifacts]]
1432
+ id = "e1_150m"
1433
+ repository = "Synthyra/fastplms_parity_goldens"
1434
+ revision = "50b93a0fc2be21a36c2522118754c74c4a9631ad"
1435
+ path = "goldens/e1_150m.safetensors"
1436
+ sha256 = "03a2e93e7b3e54b12f99eea7678ff80828922f92c3683d4909346a05d74815bc"
1437
+ size = 958851
1438
+ offline_behavior = "requires_cached_verified_file"
1439
+
1440
+ [[golden_artifacts]]
1441
+ id = "e1_300m"
1442
+ repository = "Synthyra/fastplms_parity_goldens"
1443
+ revision = "50b93a0fc2be21a36c2522118754c74c4a9631ad"
1444
+ path = "goldens/e1_300m.safetensors"
1445
+ sha256 = "d125ea866899d2788330d6877d77c48c69febbf7f01d629f4107f9647cfa13ff"
1446
+ size = 1258379
1447
+ offline_behavior = "requires_cached_verified_file"
1448
+
1449
+ [[golden_artifacts]]
1450
+ id = "e1_600m"
1451
+ repository = "Synthyra/fastplms_parity_goldens"
1452
+ revision = "50b93a0fc2be21a36c2522118754c74c4a9631ad"
1453
+ path = "goldens/e1_600m.safetensors"
1454
+ sha256 = "d470d6833f0b38bec070aa37e9f3c49717ea82d31c794af843550bb7e56253a6"
1455
+ size = 1557899
1456
+ offline_behavior = "requires_cached_verified_file"
1457
+
1458
+ [[golden_artifacts]]
1459
+ id = "dplm_150m"
1460
+ repository = "Synthyra/fastplms_parity_goldens"
1461
+ revision = "50b93a0fc2be21a36c2522118754c74c4a9631ad"
1462
+ path = "goldens/dplm_150m.safetensors"
1463
+ sha256 = "f8e4ddd580d4708dad2d7da002735b2a01d69ece478c46c4932eb2f3933dc836"
1464
+ size = 1021437
1465
+ offline_behavior = "requires_cached_verified_file"
1466
+
1467
+ [[golden_artifacts]]
1468
+ id = "dplm_650m"
1469
+ repository = "Synthyra/fastplms_parity_goldens"
1470
+ revision = "50b93a0fc2be21a36c2522118754c74c4a9631ad"
1471
+ path = "goldens/dplm_650m.safetensors"
1472
+ sha256 = "70e2c46ed722942d3b628b556b92c9994688cb4853b278eb1e40a3dfbbd9dfc4"
1473
+ size = 1989117
1474
+ offline_behavior = "requires_cached_verified_file"
1475
+
1476
+ [[golden_artifacts]]
1477
+ id = "dplm_3b"
1478
+ repository = "Synthyra/fastplms_parity_goldens"
1479
+ revision = "50b93a0fc2be21a36c2522118754c74c4a9631ad"
1480
+ path = "goldens/dplm_3b.safetensors"
1481
+ sha256 = "6942218079232a8185b1ae6e1578474c7784722e28f528b76ec89ca77777e2b9"
1482
+ size = 3924485
1483
+ offline_behavior = "requires_cached_verified_file"
1484
+
1485
+ [[golden_artifacts]]
1486
+ id = "dplm2_150m"
1487
+ repository = "Synthyra/fastplms_parity_goldens"
1488
+ revision = "50b93a0fc2be21a36c2522118754c74c4a9631ad"
1489
+ path = "goldens/dplm2_150m.safetensors"
1490
+ sha256 = "7390559bc54f27f08972c8f0adcf17954065450d937b3984303c1ae1ec003855"
1491
+ size = 26826946
1492
+ offline_behavior = "requires_cached_verified_file"
1493
+
1494
+ [[golden_artifacts]]
1495
+ id = "dplm2_650m"
1496
+ repository = "Synthyra/fastplms_parity_goldens"
1497
+ revision = "50b93a0fc2be21a36c2522118754c74c4a9631ad"
1498
+ path = "goldens/dplm2_650m.safetensors"
1499
+ sha256 = "c45f4ca267183b11dd06244ed3c3f39cf5b6139f778733fa63a6e2b5a46b0678"
1500
+ size = 28762314
1501
+ offline_behavior = "requires_cached_verified_file"
1502
+
1503
+ [[golden_artifacts]]
1504
+ id = "dplm2_3b"
1505
+ repository = "Synthyra/fastplms_parity_goldens"
1506
+ revision = "50b93a0fc2be21a36c2522118754c74c4a9631ad"
1507
+ path = "goldens/dplm2_3b.safetensors"
1508
+ sha256 = "7490f771f8b8f8137b333f98d6b7dbd2cb48708c8649d4459279a6711b9e5870"
1509
+ size = 32633034
1510
+ offline_behavior = "requires_cached_verified_file"
1511
+
1512
+ [[golden_artifacts]]
1513
+ id = "ankh_base"
1514
+ repository = "Synthyra/fastplms_parity_goldens"
1515
+ revision = "50b93a0fc2be21a36c2522118754c74c4a9631ad"
1516
+ path = "goldens/ankh_base.safetensors"
1517
+ sha256 = "449db3638117d5b6027380fdf4b05698bb156024b2b3f7e0ceb295a86e5ec756"
1518
+ size = 860722
1519
+ offline_behavior = "requires_cached_verified_file"
1520
+
1521
+ [[golden_artifacts]]
1522
+ id = "ankh_large"
1523
+ repository = "Synthyra/fastplms_parity_goldens"
1524
+ revision = "50b93a0fc2be21a36c2522118754c74c4a9631ad"
1525
+ path = "goldens/ankh_large.safetensors"
1526
+ sha256 = "4a7aea66e91cf880b0b839495d7704831725b672fe169f8e7258b5643d77aff9"
1527
+ size = 1717818
1528
+ offline_behavior = "requires_cached_verified_file"
1529
+
1530
+ [[golden_artifacts]]
1531
+ id = "ankh2_large"
1532
+ repository = "Synthyra/fastplms_parity_goldens"
1533
+ revision = "50b93a0fc2be21a36c2522118754c74c4a9631ad"
1534
+ path = "goldens/ankh2_large.safetensors"
1535
+ sha256 = "44d5fc5c74a0d166ceac54d12505a109f8338cb7090af2f37e39bd78545a3c52"
1536
+ size = 1717818
1537
+ offline_behavior = "requires_cached_verified_file"
1538
+
1539
+ [[golden_artifacts]]
1540
+ id = "ankh3_large"
1541
+ repository = "Synthyra/fastplms_parity_goldens"
1542
+ revision = "50b93a0fc2be21a36c2522118754c74c4a9631ad"
1543
+ path = "goldens/ankh3_large.safetensors"
1544
+ sha256 = "cc55274281e87d22cf7828117678df2e37a1ece543f5b401b65a0cd6f8b6fa27"
1545
+ size = 1717818
1546
+ offline_behavior = "requires_cached_verified_file"
1547
+
1548
+ [[golden_artifacts]]
1549
+ id = "ankh3_xl"
1550
+ repository = "Synthyra/fastplms_parity_goldens"
1551
+ revision = "50b93a0fc2be21a36c2522118754c74c4a9631ad"
1552
+ path = "goldens/ankh3_xl.safetensors"
1553
+ sha256 = "d18c086a4849b220761019f219590fe0c1c6f17698c365ccb0a5a8804446286c"
1554
+ size = 2860602
1555
+ offline_behavior = "requires_cached_verified_file"
1556
+
1557
+ [[golden_artifacts]]
1558
+ id = "esmfold"
1559
+ repository = "Synthyra/fastplms_parity_goldens"
1560
+ revision = "50b93a0fc2be21a36c2522118754c74c4a9631ad"
1561
+ path = "goldens/esmfold.safetensors"
1562
+ sha256 = "873b1b325a43d8e0f35f355c8914a2a9fe611cc48763875e9e6a22e09ec9ebcb"
1563
+ size = 179144
1564
+ offline_behavior = "requires_cached_verified_file"
1565
+
1566
+ [[golden_artifacts]]
1567
+ id = "esmfold2"
1568
+ repository = "Synthyra/fastplms_parity_goldens"
1569
+ revision = "50b93a0fc2be21a36c2522118754c74c4a9631ad"
1570
+ path = "goldens/esmfold2.safetensors"
1571
+ sha256 = "a1a9f3b7a9f7e36ef6a7077d76fa1b1c8b5529ed42ac2cfa485797162ce866fb"
1572
+ size = 2573976
1573
+ offline_behavior = "requires_cached_verified_file"
1574
+
1575
+ [[golden_artifacts]]
1576
+ id = "esmfold2_fast"
1577
+ repository = "Synthyra/fastplms_parity_goldens"
1578
+ revision = "50b93a0fc2be21a36c2522118754c74c4a9631ad"
1579
+ path = "goldens/esmfold2_fast.safetensors"
1580
+ sha256 = "8da91024a2cc63984585267a18f595892ebe489b2a05c4f79a994c3ac5de2f40"
1581
+ size = 2573976
1582
+ offline_behavior = "requires_cached_verified_file"
1583
+
1584
+ [[golden_artifacts]]
1585
+ id = "esmfold2_experimental_cutoff2025"
1586
+ repository = "Synthyra/fastplms_parity_goldens"
1587
+ revision = "50b93a0fc2be21a36c2522118754c74c4a9631ad"
1588
+ path = "goldens/esmfold2_experimental_cutoff2025.safetensors"
1589
+ sha256 = "8168f8692c4932e1160e733b15210836120223615f6190fb0ecfa2f940c2420f"
1590
+ size = 1771064
1591
+ offline_behavior = "requires_cached_verified_file"
1592
+
1593
+ [[golden_artifacts]]
1594
+ id = "esmfold2_experimental_fast_cutoff2025"
1595
+ repository = "Synthyra/fastplms_parity_goldens"
1596
+ revision = "50b93a0fc2be21a36c2522118754c74c4a9631ad"
1597
+ path = "goldens/esmfold2_experimental_fast_cutoff2025.safetensors"
1598
+ sha256 = "97b58525aefd21c23ec6cb2defa3b3b41982c872dfe70d94306648edbe40a5b9"
1599
+ size = 1771064
1600
+ offline_behavior = "requires_cached_verified_file"
1601
+
1602
+ [[golden_artifacts]]
1603
+ id = "esmfold2_300"
1604
+ repository = "Synthyra/fastplms_parity_goldens"
1605
+ revision = "50b93a0fc2be21a36c2522118754c74c4a9631ad"
1606
+ path = "goldens/esmfold2_300.safetensors"
1607
+ sha256 = "40a954ff14edc7c2bff95a252241f08db4cae3821ca9b92ae216775cceed8e8b"
1608
+ size = 865376
1609
+ offline_behavior = "requires_cached_verified_file"
1610
+
1611
+ [[golden_artifacts]]
1612
+ id = "esmfold2_600"
1613
+ repository = "Synthyra/fastplms_parity_goldens"
1614
+ revision = "50b93a0fc2be21a36c2522118754c74c4a9631ad"
1615
+ path = "goldens/esmfold2_600.safetensors"
1616
+ sha256 = "69bbb8d53a816b29e979e4469728df0a6b662e4d0be81a76068b36974d2b909e"
1617
+ size = 865376
1618
+ offline_behavior = "requires_cached_verified_file"
fastplms/models/esmfold/modeling_fast_esmfold.py CHANGED
@@ -798,7 +798,7 @@ class FastEsmForProteinFolding(FastPLMsAttentionMixin, EsmForProteinFolding):
798
  finally:
799
  _ESMFOLD_CAPTURED_ATTENTIONS.reset(capture_token)
800
  _ESMFOLD_OUTPUT_ATTENTIONS.reset(request_token)
801
- # Transformers 5.13 returns categorical lDDT probabilities on [0, 1],
802
  # while Meta ESMFold's public forward output reports pLDDT on [0, 100].
803
  output["plddt"] = output["plddt"] * 100
804
  payload = dict(output)
 
798
  finally:
799
  _ESMFOLD_CAPTURED_ATTENTIONS.reset(capture_token)
800
  _ESMFOLD_OUTPUT_ATTENTIONS.reset(request_token)
801
+ # Transformers returns categorical lDDT probabilities on [0, 1],
802
  # while Meta ESMFold's public forward output reports pLDDT on [0, 100].
803
  output["plddt"] = output["plddt"] * 100
804
  payload = dict(output)
fastplms/registry.py CHANGED
@@ -87,6 +87,8 @@ _ROOT_FIELDS = frozenset(
87
  "families",
88
  "models",
89
  "runtime_assets",
 
 
90
  }
91
  )
92
  _UPSTREAM_FIELDS = frozenset(
@@ -94,6 +96,7 @@ _UPSTREAM_FIELDS = frozenset(
94
  "id",
95
  "path",
96
  "url",
 
97
  "revision",
98
  "license",
99
  "license_files",
@@ -177,6 +180,17 @@ _RUNTIME_ASSET_FIELDS = frozenset(
177
  "offline_behavior",
178
  }
179
  )
 
 
 
 
 
 
 
 
 
 
 
180
 
181
 
182
  class RegistryError(ValueError):
@@ -262,6 +276,20 @@ class CheckpointSource:
262
  return MappingProxyType({item.path: item for item in self.files})
263
 
264
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
265
  @dataclass(frozen=True, slots=True)
266
  class OracleAsset:
267
  """Hash-pinned external file required by a native parity oracle."""
@@ -297,6 +325,19 @@ class OfficialGolden:
297
  tensors: FileDigest
298
 
299
 
 
 
 
 
 
 
 
 
 
 
 
 
 
300
  @dataclass(frozen=True, slots=True)
301
  class UpstreamSource:
302
  """Pinned official implementation used as a parity oracle."""
@@ -309,6 +350,13 @@ class UpstreamSource:
309
  license_files: tuple[str, ...]
310
  license_digests: tuple[FileDigest, ...] = ()
311
  distribution_files: tuple[FileDigest, ...] = ()
 
 
 
 
 
 
 
312
 
313
 
314
  @dataclass(frozen=True, slots=True)
@@ -478,6 +526,8 @@ class ModelRegistry(Mapping[str, ModelSpec]):
478
  runtime_assets: Mapping[str, RuntimeAsset] = MappingProxyType({}),
479
  attention_kernels: Mapping[str, AttentionKernelSpec] = MappingProxyType({}),
480
  legal_files: tuple[FileDigest, ...] = (),
 
 
481
  ) -> None:
482
  self.schema_version = schema_version
483
  self.upstreams = MappingProxyType(dict(upstreams))
@@ -486,6 +536,8 @@ class ModelRegistry(Mapping[str, ModelSpec]):
486
  self._models = MappingProxyType(dict(models))
487
  self.runtime_assets = MappingProxyType(dict(runtime_assets))
488
  self.legal_files = legal_files
 
 
489
 
490
  def __getitem__(self, key: str) -> ModelSpec:
491
  return self._models[key]
@@ -501,6 +553,22 @@ class ModelRegistry(Mapping[str, ModelSpec]):
501
  raise KeyError(family_id)
502
  return tuple(model for model in self._models.values() if model.family.id == family_id)
503
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
504
  def supported_attention_dtypes(
505
  self,
506
  family_id: str,
@@ -988,8 +1056,11 @@ def _parse_upstreams(raw: object) -> dict[str, UpstreamSource]:
988
  raise RegistryError(f"Duplicate upstream path: {path!r}")
989
  paths.add(path)
990
  url = _require_str(value, "url", context)
991
- if not url.startswith("https://github.com/") or not url.endswith(".git"):
992
- raise RegistryError(f"{context}.url must be an HTTPS GitHub clone URL.")
 
 
 
993
  license_files = _require_str_list(value, "license_files", context)
994
  license_digests = _require_digest_list(value, "license_digests", context)
995
  if tuple(item.path for item in license_digests) != license_files:
@@ -1026,6 +1097,7 @@ def _parse_upstreams(raw: object) -> dict[str, UpstreamSource]:
1026
  license_files=license_files,
1027
  license_digests=license_digests,
1028
  distribution_files=distribution_files,
 
1029
  )
1030
  return upstreams
1031
 
@@ -1288,6 +1360,67 @@ def _parse_runtime_assets(
1288
  return runtime_assets
1289
 
1290
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1291
  def _parse_confidence_adaptation(
1292
  table: Mapping[str, Any], context: str, model_id: str | None = None
1293
  ) -> ConfidenceAdaptation | None:
@@ -1667,6 +1800,57 @@ def _validate_registry(
1667
  )
1668
 
1669
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1670
  def _load_manifest_bytes(raw_bytes: bytes) -> ModelRegistry:
1671
  try:
1672
  manifest = tomllib.loads(raw_bytes.decode("utf-8"))
@@ -1685,6 +1869,8 @@ def _load_manifest_bytes(raw_bytes: bytes) -> ModelRegistry:
1685
  runtime_assets = _parse_runtime_assets(manifest.get("runtime_assets"), families)
1686
  models = _parse_models(manifest.get("models"), families)
1687
  _validate_registry(upstreams, attention_kernels, families, models)
 
 
1688
  return ModelRegistry(
1689
  schema_version=1,
1690
  upstreams=upstreams,
@@ -1693,6 +1879,8 @@ def _load_manifest_bytes(raw_bytes: bytes) -> ModelRegistry:
1693
  models=models,
1694
  runtime_assets=runtime_assets,
1695
  legal_files=legal_files,
 
 
1696
  )
1697
 
1698
 
@@ -1729,6 +1917,7 @@ __all__ = [
1729
  "CheckpointSource",
1730
  "FileDigest",
1731
  "GenerationContract",
 
1732
  "ModelFamily",
1733
  "ModelRegistry",
1734
  "ModelSpec",
@@ -1737,6 +1926,7 @@ __all__ = [
1737
  "RuntimeAsset",
1738
  "RuntimeAssetTrustKind",
1739
  "RuntimeExtra",
 
1740
  "TestTier",
1741
  "UpstreamSource",
1742
  "VramTier",
 
87
  "families",
88
  "models",
89
  "runtime_assets",
90
+ "golden_artifacts",
91
+ "sparse_autoencoders",
92
  }
93
  )
94
  _UPSTREAM_FIELDS = frozenset(
 
96
  "id",
97
  "path",
98
  "url",
99
+ "fetch_url",
100
  "revision",
101
  "license",
102
  "license_files",
 
180
  "offline_behavior",
181
  }
182
  )
183
+ _GOLDEN_ARTIFACT_FIELDS = frozenset(
184
+ {
185
+ "id",
186
+ "repository",
187
+ "revision",
188
+ "path",
189
+ "sha256",
190
+ "size",
191
+ "offline_behavior",
192
+ }
193
+ )
194
 
195
 
196
  class RegistryError(ValueError):
 
276
  return MappingProxyType({item.path: item for item in self.files})
277
 
278
 
279
+ @dataclass(frozen=True, slots=True)
280
+ class SparseAutoencoderSpec:
281
+ """Pinned released SAE and its input contract, not its unknown training revision."""
282
+
283
+ id: str
284
+ base_model: str
285
+ checkpoint: CheckpointSource
286
+ layer: int
287
+ input_width: int
288
+ k: int
289
+ codebook_dim: int
290
+ input_kind: Literal["hidden_state"] = "hidden_state"
291
+
292
+
293
  @dataclass(frozen=True, slots=True)
294
  class OracleAsset:
295
  """Hash-pinned external file required by a native parity oracle."""
 
325
  tensors: FileDigest
326
 
327
 
328
+ @dataclass(frozen=True, slots=True)
329
+ class GoldenArtifact:
330
+ """Durable Hub copy of one model's official golden tensors at an immutable revision."""
331
+
332
+ model_id: str
333
+ repository: str
334
+ revision: str
335
+ path: str
336
+ sha256: str
337
+ size: int
338
+ offline_behavior: str
339
+
340
+
341
  @dataclass(frozen=True, slots=True)
342
  class UpstreamSource:
343
  """Pinned official implementation used as a parity oracle."""
 
350
  license_files: tuple[str, ...]
351
  license_digests: tuple[FileDigest, ...] = ()
352
  distribution_files: tuple[FileDigest, ...] = ()
353
+ # Where the pinned revision is fetched when `url`, the source of record, no longer serves it.
354
+ fetch_url: str = ""
355
+
356
+ @property
357
+ def clone_url(self) -> str:
358
+ """The URL that `.gitmodules` and source archives use."""
359
+ return self.fetch_url or self.url
360
 
361
 
362
  @dataclass(frozen=True, slots=True)
 
526
  runtime_assets: Mapping[str, RuntimeAsset] = MappingProxyType({}),
527
  attention_kernels: Mapping[str, AttentionKernelSpec] = MappingProxyType({}),
528
  legal_files: tuple[FileDigest, ...] = (),
529
+ golden_artifacts: Mapping[str, GoldenArtifact] = MappingProxyType({}),
530
+ sparse_autoencoders: Mapping[str, SparseAutoencoderSpec] = MappingProxyType({}),
531
  ) -> None:
532
  self.schema_version = schema_version
533
  self.upstreams = MappingProxyType(dict(upstreams))
 
536
  self._models = MappingProxyType(dict(models))
537
  self.runtime_assets = MappingProxyType(dict(runtime_assets))
538
  self.legal_files = legal_files
539
+ self.golden_artifacts = MappingProxyType(dict(golden_artifacts))
540
+ self.sparse_autoencoders = MappingProxyType(dict(sparse_autoencoders))
541
 
542
  def __getitem__(self, key: str) -> ModelSpec:
543
  return self._models[key]
 
553
  raise KeyError(family_id)
554
  return tuple(model for model in self._models.values() if model.family.id == family_id)
555
 
556
+ def sae_for_base(
557
+ self, base_model: str, *, layer: int, k: int, codebook_dim: int
558
+ ) -> SparseAutoencoderSpec:
559
+ """Resolve an exact registered SAE selection without a Hub lookup or fallback."""
560
+ matches = tuple(
561
+ spec for spec in self.sparse_autoencoders.values()
562
+ if (spec.base_model, spec.layer, spec.k, spec.codebook_dim)
563
+ == (base_model, layer, k, codebook_dim)
564
+ )
565
+ if len(matches) != 1:
566
+ raise RegistryError(
567
+ f"Expected one pinned SAE for {base_model}, layer={layer}, k={k}, "
568
+ f"codebook_dim={codebook_dim}; found {len(matches)}."
569
+ )
570
+ return matches[0]
571
+
572
  def supported_attention_dtypes(
573
  self,
574
  family_id: str,
 
1056
  raise RegistryError(f"Duplicate upstream path: {path!r}")
1057
  paths.add(path)
1058
  url = _require_str(value, "url", context)
1059
+ fetch_url = _optional_str(value, "fetch_url", context) or ""
1060
+ for field, candidate in (("url", url), ("fetch_url", fetch_url)):
1061
+ is_github = candidate.startswith("https://github.com/")
1062
+ if candidate and not (is_github and candidate.endswith(".git")):
1063
+ raise RegistryError(f"{context}.{field} must be an HTTPS GitHub clone URL.")
1064
  license_files = _require_str_list(value, "license_files", context)
1065
  license_digests = _require_digest_list(value, "license_digests", context)
1066
  if tuple(item.path for item in license_digests) != license_files:
 
1097
  license_files=license_files,
1098
  license_digests=license_digests,
1099
  distribution_files=distribution_files,
1100
+ fetch_url=fetch_url,
1101
  )
1102
  return upstreams
1103
 
 
1360
  return runtime_assets
1361
 
1362
 
1363
+ def _parse_golden_artifacts(
1364
+ raw: object,
1365
+ models: Mapping[str, ModelSpec],
1366
+ ) -> dict[str, GoldenArtifact]:
1367
+ """Parse the optional Hub copies of official golden tensors, keyed by model ID.
1368
+
1369
+ A copy must be byte-identical to the tensors its model's official_golden pins, so its
1370
+ SHA-256 must equal that digest and the existing integrity check covers both files.
1371
+ """
1372
+ if raw is None:
1373
+ return {}
1374
+ if not isinstance(raw, list):
1375
+ raise RegistryError("golden_artifacts must be an array of tables.")
1376
+ golden_artifacts: dict[str, GoldenArtifact] = {}
1377
+ for index, value in enumerate(raw):
1378
+ context = f"golden_artifacts[{index}]"
1379
+ if not isinstance(value, dict):
1380
+ raise RegistryError(f"{context} must be a table.")
1381
+ _reject_unknown_fields(value, _GOLDEN_ARTIFACT_FIELDS, context)
1382
+ model_id = _require_str(value, "id", context)
1383
+ model = models.get(model_id)
1384
+ if model is None:
1385
+ raise RegistryError(f"{context}.id references unknown model {model_id!r}.")
1386
+ if model.official_golden is None:
1387
+ raise RegistryError(f"{context}.id names {model_id!r}, which has no official_golden.")
1388
+ if model_id in golden_artifacts:
1389
+ raise RegistryError(f"Duplicate golden artifact for model {model_id!r}.")
1390
+ repository = _require_str(value, "repository", context)
1391
+ if _REPOSITORY_ID_RE.fullmatch(repository) is None:
1392
+ raise RegistryError(f"{context}.repository must be a Hugging Face repository ID.")
1393
+ revision = _require_str(value, "revision", context)
1394
+ _validate_revision(revision, f"{context}.revision")
1395
+ path = _require_str(value, "path", context)
1396
+ expected_path = f"goldens/{model_id}.safetensors"
1397
+ if path != expected_path:
1398
+ raise RegistryError(f"{context}.path must be {expected_path!r}.")
1399
+ sha256 = _require_str(value, "sha256", context)
1400
+ if sha256 != model.official_golden.tensors.digest:
1401
+ raise RegistryError(
1402
+ f"{context}.sha256 must equal the official_golden tensors digest of {model_id!r}."
1403
+ )
1404
+ size = value.get("size")
1405
+ if isinstance(size, bool) or not isinstance(size, int) or size <= 0:
1406
+ raise RegistryError(f"{context}.size must be a positive byte count.")
1407
+ offline_behavior = _require_str(value, "offline_behavior", context)
1408
+ if offline_behavior not in _ALLOWED_RUNTIME_ASSET_OFFLINE_BEHAVIORS:
1409
+ raise RegistryError(
1410
+ f"{context}.offline_behavior is unsupported: {offline_behavior!r}."
1411
+ )
1412
+ golden_artifacts[model_id] = GoldenArtifact(
1413
+ model_id=model_id,
1414
+ repository=repository,
1415
+ revision=revision,
1416
+ path=path,
1417
+ sha256=sha256,
1418
+ size=size,
1419
+ offline_behavior=offline_behavior,
1420
+ )
1421
+ return golden_artifacts
1422
+
1423
+
1424
  def _parse_confidence_adaptation(
1425
  table: Mapping[str, Any], context: str, model_id: str | None = None
1426
  ) -> ConfidenceAdaptation | None:
 
1800
  )
1801
 
1802
 
1803
+ def _parse_sparse_autoencoders(
1804
+ raw: object, models: Mapping[str, ModelSpec]
1805
+ ) -> dict[str, SparseAutoencoderSpec]:
1806
+ """Validate optional SAE records independently of the base-model mapping."""
1807
+ if raw is None:
1808
+ return {}
1809
+ if not isinstance(raw, list):
1810
+ raise RegistryError("sparse_autoencoders must be an array of tables.")
1811
+ allowed = frozenset({
1812
+ "id", "base_model", "layer", "input_width", "k", "codebook_dim", "input_kind",
1813
+ "checkpoint_repo", "checkpoint_revision", "checkpoint_files",
1814
+ })
1815
+ records: dict[str, SparseAutoencoderSpec] = {}
1816
+ selections: set[tuple[str, int, int, int]] = set()
1817
+ for index, table in enumerate(raw):
1818
+ context = f"sparse_autoencoders[{index}]"
1819
+ if not isinstance(table, dict):
1820
+ raise RegistryError(f"{context} must be a table.")
1821
+ _reject_unknown_fields(table, allowed, context)
1822
+ identifier = _require_str(table, "id", context)
1823
+ if _IDENTIFIER_RE.fullmatch(identifier) is None or identifier in records:
1824
+ raise RegistryError(f"{context} has an invalid or duplicate SAE id: {identifier!r}.")
1825
+ base = _require_str(table, "base_model", context)
1826
+ if base not in models or models[base].family.id != "esm_plusplus":
1827
+ raise RegistryError(f"{context}.base_model must name a registered ESMC base.")
1828
+ numbers: dict[str, int] = {}
1829
+ for field in ("layer", "input_width", "k", "codebook_dim"):
1830
+ value = table.get(field)
1831
+ if type(value) is not int or value < (0 if field == "layer" else 1):
1832
+ raise RegistryError(f"{context}.{field} must be a valid integer dimension.")
1833
+ numbers[field] = value
1834
+ if numbers["k"] > numbers["codebook_dim"]:
1835
+ raise RegistryError(f"{context}.k exceeds codebook_dim.")
1836
+ if _require_str(table, "input_kind", context) != "hidden_state":
1837
+ raise RegistryError(f"{context} requires hidden_state input, not residual updates.")
1838
+ checkpoint = _parse_checkpoint(table, "checkpoint", context)
1839
+ expected_files = {"config.json", f"layer_{numbers['layer']}.safetensors"}
1840
+ if set(checkpoint.file_map) != expected_files or any(
1841
+ item.algorithm != "sha256" for item in checkpoint.files
1842
+ ):
1843
+ raise RegistryError(f"{context} must SHA-256 pin exactly {sorted(expected_files)}.")
1844
+ selection = (base, numbers["layer"], numbers["k"], numbers["codebook_dim"])
1845
+ if selection in selections:
1846
+ raise RegistryError(f"{context} duplicates an SAE selection: {selection}.")
1847
+ selections.add(selection)
1848
+ records[identifier] = SparseAutoencoderSpec(
1849
+ id=identifier, base_model=base, checkpoint=checkpoint, **numbers
1850
+ )
1851
+ return records
1852
+
1853
+
1854
  def _load_manifest_bytes(raw_bytes: bytes) -> ModelRegistry:
1855
  try:
1856
  manifest = tomllib.loads(raw_bytes.decode("utf-8"))
 
1869
  runtime_assets = _parse_runtime_assets(manifest.get("runtime_assets"), families)
1870
  models = _parse_models(manifest.get("models"), families)
1871
  _validate_registry(upstreams, attention_kernels, families, models)
1872
+ golden_artifacts = _parse_golden_artifacts(manifest.get("golden_artifacts"), models)
1873
+ sparse_autoencoders = _parse_sparse_autoencoders(manifest.get("sparse_autoencoders"), models)
1874
  return ModelRegistry(
1875
  schema_version=1,
1876
  upstreams=upstreams,
 
1879
  models=models,
1880
  runtime_assets=runtime_assets,
1881
  legal_files=legal_files,
1882
+ golden_artifacts=golden_artifacts,
1883
+ sparse_autoencoders=sparse_autoencoders,
1884
  )
1885
 
1886
 
 
1917
  "CheckpointSource",
1918
  "FileDigest",
1919
  "GenerationContract",
1920
+ "GoldenArtifact",
1921
  "ModelFamily",
1922
  "ModelRegistry",
1923
  "ModelSpec",
 
1926
  "RuntimeAsset",
1927
  "RuntimeAssetTrustKind",
1928
  "RuntimeExtra",
1929
+ "SparseAutoencoderSpec",
1930
  "TestTier",
1931
  "UpstreamSource",
1932
  "VramTier",
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 != "76819a3e9cb9afb1c218ca7c60eda76c75b21d9d65d5fd611b689e2edf097e19":
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 != "06a0952e7281f265e6307f7ad76411852fea6ea4c1c56839317d7f039fd1963c":
17
  raise RuntimeError("FastPLMs runtime identity differs from the bridge.")
18
 
19
  _RUNTIME_TEMPORARIES = []
requirements.txt CHANGED
@@ -1,20 +1,20 @@
1
  # Direct runtime dependencies for Synthyra/FastESMFold.
2
  # FastPLMs source is embedded in this model repository.
3
- torch>=2.13,<2.14
4
- transformers>=5.13,<5.14
5
- huggingface-hub>=0.34,<2
6
- tokenizers>=0.22,<0.23
7
- safetensors>=0.5,<1
8
- numpy>=1.26,<3
9
- einops>=0.8,<1
10
- tqdm>=4.67,<5
11
- accelerate>=1.10,<2
12
- biopython>=1.85,<2
13
- biotite>=1.4,<2
14
- brotli>=1.1,<2
15
- msgpack>=1.1,<2
16
- msgpack-numpy>=0.4.8,<1
17
- omegaconf>=2.3,<3
18
- rdkit>=2025.9,<2027
19
- scipy>=1.15,<2
20
- zstandard>=0.23,<1
 
1
  # Direct runtime dependencies for Synthyra/FastESMFold.
2
  # FastPLMs source is embedded in this model repository.
3
+ torch>=2.14
4
+ transformers>=5.17
5
+ huggingface-hub>=1.32
6
+ tokenizers>=0.23
7
+ safetensors>=0.8
8
+ numpy>=2.5
9
+ einops>=0.8
10
+ tqdm>=4.70
11
+ accelerate>=1.15
12
+ biopython>=1.88
13
+ biotite>=1.7
14
+ brotli>=1.2
15
+ msgpack>=1.2
16
+ msgpack-numpy>=0.4.8
17
+ omegaconf>=2.3
18
+ rdkit>=2026.3
19
+ scipy>=1.18
20
+ zstandard>=0.25