Instructions to use Synthyra/ESMFold2-Fast with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use Synthyra/ESMFold2-Fast with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("feature-extraction", model="Synthyra/ESMFold2-Fast", trust_remote_code=True)# pip install -U transformers accelerate # Load model directly from transformers import AutoTokenizer, AutoModel tokenizer = AutoTokenizer.from_pretrained("Synthyra/ESMFold2-Fast", trust_remote_code=True) model = AutoModel.from_pretrained("Synthyra/ESMFold2-Fast", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
Update FastPLMs files from Synthyra/FastPLMs 32d9951 (PR 52: ESMFold2 noise cap, DPLM2 parity)
Browse files- README.md +2 -2
- fastplms/atomic_files.py +65 -0
- fastplms/attention/__init__.py +2 -0
- fastplms/attention/_core.py +69 -19
- fastplms/attention/_kernel_lock.py +151 -33
- fastplms/attention/interfaces.py +4 -4
- fastplms/digests.py +35 -0
- fastplms/embeddings/__init__.py +33 -1
- fastplms/embeddings/batches.py +313 -31
- fastplms/embeddings/feature_runs.py +232 -0
- fastplms/embeddings/identity.py +44 -22
- fastplms/embeddings/inputs.py +124 -54
- fastplms/embeddings/output.py +8 -5
- fastplms/embeddings/pooling.py +68 -2
- fastplms/embeddings/runner.py +292 -45
- fastplms/embeddings/storage.py +52 -52
- fastplms/embeddings/taps.py +413 -0
- fastplms/embeddings/token_batches.py +391 -0
- fastplms/embeddings/token_runs.py +252 -0
- fastplms/embeddings/tokens.py +74 -0
- fastplms/embeddings/types.py +69 -2
- fastplms/features/__init__.py +77 -0
- fastplms/features/async_writer.py +375 -0
- fastplms/features/conversion.py +272 -0
- fastplms/features/digests.py +35 -0
- fastplms/features/json_files.py +45 -0
- fastplms/features/layouts.py +354 -0
- fastplms/features/reader.py +546 -0
- fastplms/features/receipts.py +125 -0
- fastplms/features/store.py +1567 -0
- fastplms/features/transactions.py +100 -0
- fastplms/features/writing.py +52 -0
- fastplms/json_files.py +45 -0
- fastplms/models.toml +366 -39
- fastplms/models/esm_plusplus/modeling_esm_plusplus.py +137 -12
- fastplms/models/esmfold2/embedding.py +4 -1
- fastplms/models/esmfold2/esmfold2_conformers.py +3 -0
- fastplms/models/esmfold2/esmfold2_molecular_complex.py +11 -4
- fastplms/models/esmfold2/esmfold2_parsing.py +15 -21
- fastplms/models/esmfold2/modeling_esmfold2.py +3 -2
- fastplms/models/esmfold2/modeling_esmfold2_common.py +4 -1
- fastplms/models/esmfold2/modeling_esmfold2_experimental.py +4 -3
- fastplms/registry.py +192 -2
- fastplms_bundle.py +0 -0
- modeling_fastplms.py +1 -1
- requirements.txt +18 -18
README.md
CHANGED
|
@@ -79,7 +79,7 @@ python -m pip install -r \
|
|
| 79 |
The FastPLMs implementation itself is embedded in the model repository.
|
| 80 |
Transformers loads it through `trust_remote_code=True`.
|
| 81 |
|
| 82 |
-
This model requires Python 3.
|
| 83 |
|
| 84 |
The artifact requirements include the structure dependencies.
|
| 85 |
|
|
@@ -144,7 +144,7 @@ print(token_output.logits.shape) # (b, l, 3)
|
|
| 144 |
Install the training dependencies. Then attach LoRA to the loaded checkpoint:
|
| 145 |
|
| 146 |
```bash
|
| 147 |
-
python -m pip install "datasets>=
|
| 148 |
```
|
| 149 |
|
| 150 |
```python
|
|
|
|
| 79 |
The FastPLMs implementation itself is embedded in the model repository.
|
| 80 |
Transformers loads it through `trust_remote_code=True`.
|
| 81 |
|
| 82 |
+
This model requires Python 3.12-3.14, PyTorch 2.14, and Transformers 5.17.
|
| 83 |
|
| 84 |
The artifact requirements include the structure dependencies.
|
| 85 |
|
|
|
|
| 144 |
Install the training dependencies. Then attach LoRA to the loaded checkpoint:
|
| 145 |
|
| 146 |
```bash
|
| 147 |
+
python -m pip install "datasets>=5.0" "peft>=0.21"
|
| 148 |
```
|
| 149 |
|
| 150 |
```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(
|
|
|
|
|
|
|
| 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() #
|
| 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() #
|
| 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(
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
|
|
|
| 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}.
|
| 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 |
-
|
| 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(
|
| 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 |
-
|
| 134 |
-
|
| 135 |
-
|
| 136 |
-
|
| 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
|
| 146 |
-
"""
|
| 147 |
-
require_kernels_package()
|
| 148 |
try:
|
| 149 |
-
from kernels import
|
| 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 |
-
|
| 157 |
-
|
| 158 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 159 |
raise RuntimeError(
|
| 160 |
f"The typed manifest pins {repository}@{revision}, but kernels.lock pins "
|
| 161 |
-
f"{
|
| 162 |
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 163 |
|
| 164 |
if _offline_mode():
|
| 165 |
-
return _load_offline_locked_kernel(repository, revision,
|
| 166 |
-
|
| 167 |
-
#
|
| 168 |
-
#
|
| 169 |
-
|
| 170 |
-
|
| 171 |
-
|
| 172 |
-
|
| 173 |
-
|
| 174 |
-
|
|
|
|
|
|
|
|
|
|
| 175 |
)
|
| 176 |
-
|
|
|
|
|
|
| 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
|
| 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 |
-
|
| 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 .
|
|
|
|
|
|
|
|
|
|
| 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 .
|
|
|
|
|
|
|
| 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
|
| 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 |
-
|
| 144 |
-
|
| 145 |
-
|
| 146 |
-
|
| 147 |
-
|
| 148 |
-
|
| 149 |
-
|
| 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 |
-
|
| 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 (
|
| 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 = ( #
|
| 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 .
|
| 16 |
from .storage import tensor_sha256
|
| 17 |
from .types import EmbeddingInput
|
|
|
|
|
|
|
| 18 |
|
| 19 |
|
| 20 |
-
_RUN_FINGERPRINT_SCHEMA_VERSION =
|
| 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 |
-
|
| 97 |
-
|
| 98 |
-
|
| 99 |
-
|
| 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 =
|
| 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 |
-
|
| 410 |
-
|
| 411 |
-
|
|
|
|
| 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 |
-
|
| 17 |
-
"""
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 18 |
|
| 19 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 20 |
sequence_parts: list[str] = []
|
| 21 |
found_record = False
|
| 22 |
-
|
| 23 |
-
|
| 24 |
-
|
| 25 |
-
|
| 26 |
-
|
| 27 |
-
|
| 28 |
-
|
| 29 |
-
|
| 30 |
-
|
| 31 |
-
|
| 32 |
-
|
| 33 |
-
|
| 34 |
-
|
| 35 |
-
|
| 36 |
-
|
| 37 |
-
|
| 38 |
-
|
| 39 |
-
|
| 40 |
-
|
| 41 |
-
if
|
| 42 |
found_record = True
|
| 43 |
-
yield
|
| 44 |
-
if not found_record:
|
| 45 |
-
raise ValueError(
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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.
|
| 77 |
-
|
| 78 |
-
)
|
| 79 |
-
self.path =
|
| 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 |
-
|
| 99 |
pending.clear()
|
| 100 |
if pending:
|
| 101 |
-
|
| 102 |
if count == 0:
|
| 103 |
raise ValueError("inputs must contain at least one sequence.")
|
| 104 |
-
|
| 105 |
-
|
| 106 |
-
self.
|
| 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 |
-
|
|
|
|
| 119 |
raise RuntimeError("Input spool is closed.")
|
| 120 |
-
return
|
| 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 |
-
|
| 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 |
-
|
| 132 |
-
|
| 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 |
-
|
| 213 |
if self.output is not None:
|
| 214 |
-
return save_result(
|
| 215 |
-
|
|
|
|
|
|
|
|
|
| 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__ = [
|
|
|
|
|
|
|
|
|
|
|
|
| 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 .
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 139 |
-
(
|
| 140 |
-
|
| 141 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 325 |
-
projection = getattr(model, "embedding_projection", None)
|
| 326 |
-
resolved_layer = getattr(
|
| 327 |
model,
|
| 328 |
-
|
| 329 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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":
|
| 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":
|
| 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":
|
| 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
|
| 405 |
-
metadata["
|
| 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
|
| 411 |
|
| 412 |
|
| 413 |
class EmbeddingMixin:
|
| 414 |
"""Small delegation mixin shared by FastPLMs model classes."""
|
| 415 |
|
| 416 |
-
def embed_dataset(
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
| 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,
|
| 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(
|
| 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 =
|
| 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 =
|
| 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(
|
| 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(
|
| 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 |
-
|
| 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 |
-
|
| 784 |
shard_size=shard_size,
|
| 785 |
publish_initial=False,
|
| 786 |
publish_incremental=False,
|
| 787 |
)
|
| 788 |
-
writer.append(
|
| 789 |
-
return writer.publish(complete=bool(
|
| 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(
|
| 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(
|
| 934 |
if not run_id:
|
| 935 |
raise ValueError("SQLite results require metadata['run_fingerprint'].")
|
| 936 |
metadata_json = json.dumps(
|
| 937 |
_persistent_metadata(
|
| 938 |
-
|
| 939 |
descriptor_index="sqlite-records",
|
| 940 |
-
record_count=len(
|
| 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(
|
| 955 |
X = record.load_tensor().detach().cpu().contiguous() # (...)
|
| 956 |
-
dtype_name, shape_json,
|
| 957 |
digest = tensor_sha256(X)
|
| 958 |
connection.execute(
|
| 959 |
"INSERT INTO tensors VALUES (?, ?, ?, ?, ?, ?)",
|
| 960 |
-
(run_id, position, dtype_name, shape_json,
|
| 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,
|
| 1061 |
digest = tensor_sha256(X)
|
| 1062 |
connection.execute(
|
| 1063 |
"INSERT INTO tensors VALUES (?, ?, ?, ?, ?, ?)",
|
| 1064 |
-
(run_id, position, dtype_name, shape_json,
|
| 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(
|
| 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 |
-
|
| 1438 |
*,
|
| 1439 |
fallback_shape: tuple[int, ...] | None,
|
| 1440 |
allow_unsafe_pickle: bool,
|
| 1441 |
) -> Tensor:
|
| 1442 |
-
if len(
|
| 1443 |
-
dtype_code = int(
|
| 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",
|
| 1447 |
-
if ndim < 0 or ndim > 16 or len(
|
| 1448 |
raise ValueError("Malformed legacy compact embedding header.")
|
| 1449 |
-
shape = tuple(int(value) for value in struct.unpack_from(f"<{ndim}i",
|
| 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(
|
| 1456 |
raise ValueError("Legacy compact embedding payload length does not match shape.")
|
| 1457 |
-
array = ( # shape
|
| 1458 |
-
np.frombuffer(
|
| 1459 |
)
|
| 1460 |
-
return torch.from_numpy(array).to(dtype=target_dtype) # shape
|
| 1461 |
|
| 1462 |
try:
|
| 1463 |
-
loaded = torch.load(io.BytesIO(
|
| 1464 |
except Exception as safe_error:
|
| 1465 |
if allow_unsafe_pickle:
|
| 1466 |
-
loaded = torch.load(io.BytesIO(
|
| 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(
|
| 1476 |
raise ValueError(
|
| 1477 |
"Legacy raw FP32 payload length does not match fallback_shape."
|
| 1478 |
) from safe_error
|
| 1479 |
-
array = np.frombuffer(
|
| 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,
|
| 1526 |
if not isinstance(sequence, str) or not sequence:
|
| 1527 |
raise ValueError("Legacy embedding sequences must be non-empty strings.")
|
| 1528 |
-
if not isinstance(
|
| 1529 |
-
|
| 1530 |
tensor = _decode_legacy_sqlite_blob(
|
| 1531 |
-
|
| 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 |
-
|
| 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(
|
| 1570 |
if format == "sqlite":
|
| 1571 |
-
return save_sqlite_result(
|
| 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 = "
|
| 11 |
-
version =
|
| 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:
|
| 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:
|
| 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:
|
| 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:
|
| 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:
|
| 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:
|
| 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:
|
| 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:
|
| 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:
|
| 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:
|
| 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:
|
| 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:
|
| 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:
|
| 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:
|
| 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:
|
| 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:
|
| 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:
|
| 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:
|
| 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:
|
| 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:
|
| 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:
|
| 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:
|
| 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:
|
| 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:
|
| 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:
|
| 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:
|
| 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:
|
| 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/esm_plusplus/modeling_esm_plusplus.py
CHANGED
|
@@ -10,7 +10,7 @@ import torch
|
|
| 10 |
import torch.nn as nn
|
| 11 |
import torch.nn.functional as F
|
| 12 |
|
| 13 |
-
from collections.abc import Sequence
|
| 14 |
from contextlib import contextmanager
|
| 15 |
from dataclasses import asdict, dataclass
|
| 16 |
from functools import partial
|
|
@@ -416,7 +416,7 @@ class RotaryEmbedding(torch.nn.Module):
|
|
| 416 |
and self._seq_len_cached >= token_count
|
| 417 |
and cached.device == device
|
| 418 |
and cached.dtype == dtype
|
| 419 |
-
and not (
|
| 420 |
)
|
| 421 |
|
| 422 |
def _rotary_angles(
|
|
@@ -808,6 +808,7 @@ class TransformerOutput(ModelOutput):
|
|
| 808 |
s_max: tuple[list[torch.Tensor], ...] | None = None
|
| 809 |
sae_outputs: dict[str, torch.Tensor] | None = None
|
| 810 |
sae_hidden_states: dict[int, torch.Tensor] | None = None
|
|
|
|
| 811 |
|
| 812 |
|
| 813 |
@dataclass
|
|
@@ -884,13 +885,53 @@ class TransformerStack(nn.Module):
|
|
| 884 |
output_s_max: bool | None = False,
|
| 885 |
esmfold2_hidden_states: bool = False,
|
| 886 |
sae_layers: tuple[int, ...] = (),
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 887 |
) -> TransformerOutput:
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 888 |
# x: (b, l, d); attention_mask, sequence_id: (b, l)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 889 |
hidden_states = () if output_hidden_states else None
|
| 890 |
attentions = () if output_attentions else None
|
| 891 |
full_s_max = () if output_s_max else None
|
| 892 |
-
|
| 893 |
-
|
| 894 |
# Match the pinned Biohub Transformers contract: a supplied sequence_id
|
| 895 |
# is authoritative and must encode padding as -1. attention_mask is
|
| 896 |
# ignored in that mode rather than intersected with the chain mask.
|
|
@@ -902,6 +943,7 @@ class TransformerStack(nn.Module):
|
|
| 902 |
device=x.device,
|
| 903 |
dtype=x.dtype,
|
| 904 |
output_attentions=bool(output_attentions),
|
|
|
|
| 905 |
)
|
| 906 |
# A call that returns attention weights runs eager attention and needs no layout.
|
| 907 |
flash_padding_layout = (
|
|
@@ -911,6 +953,8 @@ class TransformerStack(nn.Module):
|
|
| 911 |
)
|
| 912 |
|
| 913 |
for layer_index, block in enumerate(self.blocks):
|
|
|
|
|
|
|
| 914 |
if output_hidden_states:
|
| 915 |
if hidden_states is None:
|
| 916 |
raise RuntimeError(
|
|
@@ -922,6 +966,10 @@ class TransformerStack(nn.Module):
|
|
| 922 |
hidden_states += (x,)
|
| 923 |
if sae_hidden_states is not None and layer_index in sae_layer_set:
|
| 924 |
sae_hidden_states[layer_index] = x
|
|
|
|
|
|
|
|
|
|
|
|
|
| 925 |
if self.gradient_checkpointing and self.training:
|
| 926 |
x, attn_weights, s_max = self._gradient_checkpointing_func(
|
| 927 |
block.__call__,
|
|
@@ -949,12 +997,20 @@ class TransformerStack(nn.Module):
|
|
| 949 |
if full_s_max is not None:
|
| 950 |
full_s_max += (s_max,)
|
| 951 |
|
| 952 |
-
last_hidden_state =
|
| 953 |
-
if
|
| 954 |
-
|
| 955 |
-
|
| 956 |
-
|
| 957 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 958 |
|
| 959 |
return TransformerOutput(
|
| 960 |
last_hidden_state=last_hidden_state,
|
|
@@ -962,6 +1018,7 @@ class TransformerStack(nn.Module):
|
|
| 962 |
attentions=attentions,
|
| 963 |
s_max=full_s_max,
|
| 964 |
sae_hidden_states=sae_hidden_states,
|
|
|
|
| 965 |
)
|
| 966 |
|
| 967 |
@torch.compiler.disable
|
|
@@ -975,6 +1032,7 @@ class TransformerStack(nn.Module):
|
|
| 975 |
dtype: torch.dtype | None = None,
|
| 976 |
effective_backend: AttentionBackend | None = None,
|
| 977 |
output_attentions: bool = False,
|
|
|
|
| 978 |
) -> tuple[torch.Tensor | None, torch.Tensor | None, BlockMask | None]:
|
| 979 |
mask_name = "sequence_id" if sequence_id is not None else "attention_mask"
|
| 980 |
mask_pattern = sequence_id if sequence_id is not None else attention_mask
|
|
@@ -1006,7 +1064,7 @@ class TransformerStack(nn.Module):
|
|
| 1006 |
attention_mask_2d = (
|
| 1007 |
mask_pattern if mask_pattern.dtype == torch.bool else mask_pattern != -1
|
| 1008 |
)
|
| 1009 |
-
if not bool(attention_mask_2d.any(dim=1).all()):
|
| 1010 |
raise ValueError("attention_mask must keep at least one valid key per batch row.")
|
| 1011 |
|
| 1012 |
if mask_pattern.dtype == torch.bool:
|
|
@@ -1092,6 +1150,9 @@ class PreTrainedESMplusplusModel(FastPLMsAttentionMixin, PreTrainedModel):
|
|
| 1092 |
"flash_attention_3",
|
| 1093 |
)
|
| 1094 |
_fastplms_attention_auto_order = ("sdpa",)
|
|
|
|
|
|
|
|
|
|
| 1095 |
|
| 1096 |
def __init__(self, config: ESMplusplusConfig, *args: object, **kwargs: object) -> None:
|
| 1097 |
super().__init__(config, *args, **kwargs)
|
|
@@ -1367,14 +1428,78 @@ class PreTrainedESMplusplusModel(FastPLMsAttentionMixin, PreTrainedModel):
|
|
| 1367 |
if output.attentions is not None
|
| 1368 |
else None
|
| 1369 |
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1370 |
return TransformerOutput(
|
| 1371 |
-
last_hidden_state=
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1372 |
hidden_states=hidden_states,
|
| 1373 |
attentions=attentions,
|
| 1374 |
s_max=output.s_max,
|
| 1375 |
sae_hidden_states=output.sae_hidden_states,
|
|
|
|
| 1376 |
)
|
| 1377 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1378 |
@property
|
| 1379 |
def tokenizer(self) -> EsmSequenceTokenizer:
|
| 1380 |
"""Construct the sequence tokenizer only when a raw-sequence API needs it."""
|
|
|
|
| 10 |
import torch.nn as nn
|
| 11 |
import torch.nn.functional as F
|
| 12 |
|
| 13 |
+
from collections.abc import Callable, Sequence
|
| 14 |
from contextlib import contextmanager
|
| 15 |
from dataclasses import asdict, dataclass
|
| 16 |
from functools import partial
|
|
|
|
| 416 |
and self._seq_len_cached >= token_count
|
| 417 |
and cached.device == device
|
| 418 |
and cached.dtype == dtype
|
| 419 |
+
and not (torch.is_grad_enabled() and cached.is_inference()) # autograd cannot save inference-created tables, in eval mode too
|
| 420 |
)
|
| 421 |
|
| 422 |
def _rotary_angles(
|
|
|
|
| 808 |
s_max: tuple[list[torch.Tensor], ...] | None = None
|
| 809 |
sae_outputs: dict[str, torch.Tensor] | None = None
|
| 810 |
sae_hidden_states: dict[int, torch.Tensor] | None = None
|
| 811 |
+
captured_hidden_states: dict[int, torch.Tensor] | None = None
|
| 812 |
|
| 813 |
|
| 814 |
@dataclass
|
|
|
|
| 885 |
output_s_max: bool | None = False,
|
| 886 |
esmfold2_hidden_states: bool = False,
|
| 887 |
sae_layers: tuple[int, ...] = (),
|
| 888 |
+
capture_layers: tuple[int, ...] = (),
|
| 889 |
+
stop_after_layer: int | None = None,
|
| 890 |
+
stream_layers: tuple[int, ...] = (),
|
| 891 |
+
state_consumer: Callable[[int, torch.Tensor], None] | None = None,
|
| 892 |
+
assume_valid_mask: bool = False,
|
| 893 |
) -> TransformerOutput:
|
| 894 |
+
"""Run the blocks and record the requested hidden states.
|
| 895 |
+
|
| 896 |
+
Hidden state ``i`` is the input to block ``i``, and state ``n_layers`` is the final
|
| 897 |
+
normalized state. ``capture_layers`` and ``sae_layers`` record states by that index.
|
| 898 |
+
``stop_after_layer`` ends the pass once its state exists: no later block runs, and the
|
| 899 |
+
final norm runs only when that state is ``n_layers``, the default.
|
| 900 |
+
``assume_valid_mask`` skips the check that every row keeps a valid key. That check reads a
|
| 901 |
+
device value and so stalls the host until the device is idle; a caller that builds its mask
|
| 902 |
+
from known lengths, as the canonical token executor does, sets it.
|
| 903 |
+
"""
|
| 904 |
# x: (b, l, d); attention_mask, sequence_id: (b, l)
|
| 905 |
+
n_layers = len(self.blocks)
|
| 906 |
+
depth = n_layers if stop_after_layer is None else stop_after_layer
|
| 907 |
+
if not 0 <= depth <= n_layers:
|
| 908 |
+
raise ValueError(
|
| 909 |
+
f"stop_after_layer must be a hidden-state index in 0..{n_layers}; "
|
| 910 |
+
f"received {stop_after_layer!r}."
|
| 911 |
+
)
|
| 912 |
+
capture_set = set(capture_layers)
|
| 913 |
+
stream_set = set(stream_layers)
|
| 914 |
+
if bool(stream_set) != (state_consumer is not None):
|
| 915 |
+
raise ValueError("Streaming layers and their state consumer must be supplied together.")
|
| 916 |
+
sae_layer_set = set(sae_layers)
|
| 917 |
+
# A full pass keeps the SAE contract unchanged; an early stop must reach every SAE state.
|
| 918 |
+
requested = capture_set | sae_layer_set if depth < n_layers else capture_set
|
| 919 |
+
requested = requested | stream_set
|
| 920 |
+
uncomputed = sorted(index for index in requested if not 0 <= index <= depth)
|
| 921 |
+
if uncomputed:
|
| 922 |
+
raise ValueError(
|
| 923 |
+
f"Hidden states {uncomputed} are outside the states 0..{depth} this call computes."
|
| 924 |
+
)
|
| 925 |
+
if depth < n_layers and (output_hidden_states or output_attentions or output_s_max):
|
| 926 |
+
raise ValueError(
|
| 927 |
+
"A stack that stops before its final state returns only captured hidden states; "
|
| 928 |
+
"output_hidden_states, output_attentions, and output_s_max need every block."
|
| 929 |
+
)
|
| 930 |
hidden_states = () if output_hidden_states else None
|
| 931 |
attentions = () if output_attentions else None
|
| 932 |
full_s_max = () if output_s_max else None
|
| 933 |
+
sae_hidden_states: dict[int, torch.Tensor] | None = {} if sae_layer_set else None
|
| 934 |
+
captured_hidden_states: dict[int, torch.Tensor] | None = {} if capture_set else None
|
| 935 |
# Match the pinned Biohub Transformers contract: a supplied sequence_id
|
| 936 |
# is authoritative and must encode padding as -1. attention_mask is
|
| 937 |
# ignored in that mode rather than intersected with the chain mask.
|
|
|
|
| 943 |
device=x.device,
|
| 944 |
dtype=x.dtype,
|
| 945 |
output_attentions=bool(output_attentions),
|
| 946 |
+
validate_mask=not assume_valid_mask,
|
| 947 |
)
|
| 948 |
# A call that returns attention weights runs eager attention and needs no layout.
|
| 949 |
flash_padding_layout = (
|
|
|
|
| 953 |
)
|
| 954 |
|
| 955 |
for layer_index, block in enumerate(self.blocks):
|
| 956 |
+
if layer_index == depth:
|
| 957 |
+
break
|
| 958 |
if output_hidden_states:
|
| 959 |
if hidden_states is None:
|
| 960 |
raise RuntimeError(
|
|
|
|
| 966 |
hidden_states += (x,)
|
| 967 |
if sae_hidden_states is not None and layer_index in sae_layer_set:
|
| 968 |
sae_hidden_states[layer_index] = x
|
| 969 |
+
if captured_hidden_states is not None and layer_index in capture_set:
|
| 970 |
+
captured_hidden_states[layer_index] = x
|
| 971 |
+
if layer_index in stream_set:
|
| 972 |
+
state_consumer(layer_index, x)
|
| 973 |
if self.gradient_checkpointing and self.training:
|
| 974 |
x, attn_weights, s_max = self._gradient_checkpointing_func(
|
| 975 |
block.__call__,
|
|
|
|
| 997 |
if full_s_max is not None:
|
| 998 |
full_s_max += (s_max,)
|
| 999 |
|
| 1000 |
+
last_hidden_state: torch.Tensor | None = None
|
| 1001 |
+
if depth == n_layers:
|
| 1002 |
+
last_hidden_state = self.norm(x) # (b, l, d)
|
| 1003 |
+
if output_hidden_states:
|
| 1004 |
+
hidden_states += (last_hidden_state,)
|
| 1005 |
+
# State `depth` is the final normalized state, or, after an early stop, the input to
|
| 1006 |
+
# block `depth`, which the final norm never touches.
|
| 1007 |
+
deepest_state = x if last_hidden_state is None else last_hidden_state # (b, l, d)
|
| 1008 |
+
if sae_hidden_states is not None and depth in sae_layer_set:
|
| 1009 |
+
sae_hidden_states[depth] = deepest_state
|
| 1010 |
+
if captured_hidden_states is not None and depth in capture_set:
|
| 1011 |
+
captured_hidden_states[depth] = deepest_state
|
| 1012 |
+
if depth in stream_set:
|
| 1013 |
+
state_consumer(depth, deepest_state)
|
| 1014 |
|
| 1015 |
return TransformerOutput(
|
| 1016 |
last_hidden_state=last_hidden_state,
|
|
|
|
| 1018 |
attentions=attentions,
|
| 1019 |
s_max=full_s_max,
|
| 1020 |
sae_hidden_states=sae_hidden_states,
|
| 1021 |
+
captured_hidden_states=captured_hidden_states,
|
| 1022 |
)
|
| 1023 |
|
| 1024 |
@torch.compiler.disable
|
|
|
|
| 1032 |
dtype: torch.dtype | None = None,
|
| 1033 |
effective_backend: AttentionBackend | None = None,
|
| 1034 |
output_attentions: bool = False,
|
| 1035 |
+
validate_mask: bool = True,
|
| 1036 |
) -> tuple[torch.Tensor | None, torch.Tensor | None, BlockMask | None]:
|
| 1037 |
mask_name = "sequence_id" if sequence_id is not None else "attention_mask"
|
| 1038 |
mask_pattern = sequence_id if sequence_id is not None else attention_mask
|
|
|
|
| 1064 |
attention_mask_2d = (
|
| 1065 |
mask_pattern if mask_pattern.dtype == torch.bool else mask_pattern != -1
|
| 1066 |
)
|
| 1067 |
+
if validate_mask and not bool(attention_mask_2d.any(dim=1).all()):
|
| 1068 |
raise ValueError("attention_mask must keep at least one valid key per batch row.")
|
| 1069 |
|
| 1070 |
if mask_pattern.dtype == torch.bool:
|
|
|
|
| 1150 |
"flash_attention_3",
|
| 1151 |
)
|
| 1152 |
_fastplms_attention_auto_order = ("sdpa",)
|
| 1153 |
+
# embed_dataset(taps=...) runs _embed_taps; families without this marker reject taps.
|
| 1154 |
+
embedding_tap_support: ClassVar[bool] = True
|
| 1155 |
+
embedding_streaming_tap_support: ClassVar[bool] = True
|
| 1156 |
|
| 1157 |
def __init__(self, config: ESMplusplusConfig, *args: object, **kwargs: object) -> None:
|
| 1158 |
super().__init__(config, *args, **kwargs)
|
|
|
|
| 1428 |
if output.attentions is not None
|
| 1429 |
else None
|
| 1430 |
)
|
| 1431 |
+
captured_hidden_states = (
|
| 1432 |
+
{
|
| 1433 |
+
index: state[:, :sequence_length] # (b, l, d)
|
| 1434 |
+
for index, state in output.captured_hidden_states.items()
|
| 1435 |
+
}
|
| 1436 |
+
if output.captured_hidden_states is not None
|
| 1437 |
+
else None
|
| 1438 |
+
)
|
| 1439 |
return TransformerOutput(
|
| 1440 |
+
last_hidden_state=(
|
| 1441 |
+
output.last_hidden_state[:, :sequence_length]
|
| 1442 |
+
if output.last_hidden_state is not None
|
| 1443 |
+
else None
|
| 1444 |
+
),
|
| 1445 |
hidden_states=hidden_states,
|
| 1446 |
attentions=attentions,
|
| 1447 |
s_max=output.s_max,
|
| 1448 |
sae_hidden_states=output.sae_hidden_states,
|
| 1449 |
+
captured_hidden_states=captured_hidden_states,
|
| 1450 |
)
|
| 1451 |
|
| 1452 |
+
@property
|
| 1453 |
+
def embedding_tap_state_count(self) -> int:
|
| 1454 |
+
"""Hidden states a tap can name: each block's input, then the final normalized state."""
|
| 1455 |
+
|
| 1456 |
+
return int(self.config.num_hidden_layers) + 1
|
| 1457 |
+
|
| 1458 |
+
def _embed_taps(
|
| 1459 |
+
self,
|
| 1460 |
+
input_ids: torch.Tensor,
|
| 1461 |
+
attention_mask: torch.Tensor,
|
| 1462 |
+
layers: tuple[int, ...],
|
| 1463 |
+
*,
|
| 1464 |
+
stream_layers: tuple[int, ...] = (),
|
| 1465 |
+
state_consumer: Callable[[int, torch.Tensor], None] | None = None,
|
| 1466 |
+
assume_valid_mask: bool = False,
|
| 1467 |
+
) -> dict[int, torch.Tensor]:
|
| 1468 |
+
"""Hidden states at the FastPLMs indices ``layers`` from one pass.
|
| 1469 |
+
|
| 1470 |
+
The pass stops once the deepest requested state exists. It pads, enters, and trims the
|
| 1471 |
+
FP8 context exactly as ``forward`` does, so each state equals the one
|
| 1472 |
+
``forward(output_hidden_states=True)`` returns for the same batch. ``assume_valid_mask``
|
| 1473 |
+
skips the device-side mask check, which would stall the host on every batch.
|
| 1474 |
+
"""
|
| 1475 |
+
# input_ids, attention_mask: (b, l)
|
| 1476 |
+
if not layers and not stream_layers:
|
| 1477 |
+
raise ValueError("_embed_taps needs at least one hidden-state index.")
|
| 1478 |
+
input_ids, attention_mask, _, _, original_length = self._pad_fp8_inputs(
|
| 1479 |
+
input_ids, attention_mask, None, None
|
| 1480 |
+
) # input_ids, attention_mask: (b, l_fp8), l_fp8 = l unless FP8 pads to 16
|
| 1481 |
+
x = self.embed(input_ids) # (b, l_fp8, d)
|
| 1482 |
+
|
| 1483 |
+
def consume(index: int, state: torch.Tensor) -> None:
|
| 1484 |
+
# Reducer masks describe original tokens, never the alignment-only FP8 padding.
|
| 1485 |
+
trimmed = state[:, :original_length] if original_length is not None else state
|
| 1486 |
+
state_consumer(index, trimmed) # (b, l, d)
|
| 1487 |
+
|
| 1488 |
+
with _esmplusplus_fp8_context(self._esmc_fp8, self.device):
|
| 1489 |
+
output = self.transformer(
|
| 1490 |
+
x=x,
|
| 1491 |
+
attention_mask=attention_mask,
|
| 1492 |
+
capture_layers=layers,
|
| 1493 |
+
stop_after_layer=max((*layers, *stream_layers)),
|
| 1494 |
+
stream_layers=stream_layers,
|
| 1495 |
+
state_consumer=consume if state_consumer is not None else None,
|
| 1496 |
+
assume_valid_mask=assume_valid_mask,
|
| 1497 |
+
)
|
| 1498 |
+
captured = self._trim_transformer_output(output, original_length).captured_hidden_states
|
| 1499 |
+
if captured is None and layers:
|
| 1500 |
+
raise RuntimeError("ESM++ did not capture the requested hidden states.")
|
| 1501 |
+
return captured or {} # each: (b, l, d); streaming-only runs retain no states
|
| 1502 |
+
|
| 1503 |
@property
|
| 1504 |
def tokenizer(self) -> EsmSequenceTokenizer:
|
| 1505 |
"""Construct the sequence tokenizer only when a raw-sequence API needs it."""
|
fastplms/models/esmfold2/embedding.py
CHANGED
|
@@ -110,7 +110,10 @@ class ESMFold2EmbeddingMixin:
|
|
| 110 |
def embed_dataset(self, inputs: Any, **kwargs: Any) -> EmbeddingResult:
|
| 111 |
"""Embed single-chain proteins using the learned 256-wide ESMFold2 summary."""
|
| 112 |
|
| 113 |
-
|
|
|
|
|
|
|
|
|
|
| 114 |
|
| 115 |
|
| 116 |
__all__ = ["ESMFold2EmbeddingMixin"]
|
|
|
|
| 110 |
def embed_dataset(self, inputs: Any, **kwargs: Any) -> EmbeddingResult:
|
| 111 |
"""Embed single-chain proteins using the learned 256-wide ESMFold2 summary."""
|
| 112 |
|
| 113 |
+
embeddings = embed_dataset(self, inputs, **kwargs)
|
| 114 |
+
# ESMFold2 serves no one-pass taps, and embed_dataset refuses taps= before inference.
|
| 115 |
+
assert isinstance(embeddings, EmbeddingResult)
|
| 116 |
+
return embeddings
|
| 117 |
|
| 118 |
|
| 119 |
__all__ = ["ESMFold2EmbeddingMixin"]
|
fastplms/models/esmfold2/esmfold2_conformers.py
CHANGED
|
@@ -116,6 +116,9 @@ def _open_verified_asset(
|
|
| 116 |
remaining -= len(chunk)
|
| 117 |
copied_size = snapshot.tell()
|
| 118 |
extra_byte = source.read(1)
|
|
|
|
|
|
|
|
|
|
| 119 |
if remaining or extra_byte:
|
| 120 |
observed_size = copied_size if remaining else copied_size + len(extra_byte)
|
| 121 |
raise ValueError(
|
|
|
|
| 116 |
remaining -= len(chunk)
|
| 117 |
copied_size = snapshot.tell()
|
| 118 |
extra_byte = source.read(1)
|
| 119 |
+
# Release the asset once the snapshot holds its bytes. An open handle
|
| 120 |
+
# would block another process from replacing the path on Windows.
|
| 121 |
+
source.close()
|
| 122 |
if remaining or extra_byte:
|
| 123 |
observed_size = copied_size if remaining else copied_size + len(extra_byte)
|
| 124 |
raise ValueError(
|
fastplms/models/esmfold2/esmfold2_molecular_complex.py
CHANGED
|
@@ -242,14 +242,21 @@ def _protein_sequence_and_indices(
|
|
| 242 |
return protein_indices, "".join(sequence), chain_ids, entity_ids, sym_ids, confidences
|
| 243 |
|
| 244 |
|
| 245 |
-
def _protein_entity_metadata_value(value: int | str) -> int
|
| 246 |
-
"""Restore the numeric entity
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 247 |
if isinstance(value, str):
|
| 248 |
try:
|
| 249 |
return int(value)
|
| 250 |
except ValueError:
|
| 251 |
pass
|
| 252 |
-
return
|
| 253 |
|
| 254 |
|
| 255 |
def _atom37_from_flat(
|
|
@@ -787,7 +794,7 @@ class MolecularComplex:
|
|
| 787 |
metadata = ProteinComplexMetadata(
|
| 788 |
entity_lookup={
|
| 789 |
int(entity): _protein_entity_metadata_value(
|
| 790 |
-
self.metadata.entity_lookup.get(int(entity), int(entity))
|
| 791 |
)
|
| 792 |
for entity in unique_entities
|
| 793 |
},
|
|
|
|
| 242 |
return protein_indices, "".join(sequence), chain_ids, entity_ids, sym_ids, confidences
|
| 243 |
|
| 244 |
|
| 245 |
+
def _protein_entity_metadata_value(value: int | str, entity: int) -> int:
|
| 246 |
+
"""Restore the numeric entity label ProteinComplex metadata carries for `entity`.
|
| 247 |
+
|
| 248 |
+
A complex converted from a ProteinComplex keeps its labels as numeric strings. A folded
|
| 249 |
+
complex records the entity type ("polymer", "non-polymer") under the same key instead, which
|
| 250 |
+
is no label, so the entity keeps its own number. ProteinChain accepts only an integer label.
|
| 251 |
+
"""
|
| 252 |
+
if isinstance(value, int) and not isinstance(value, bool):
|
| 253 |
+
return value
|
| 254 |
if isinstance(value, str):
|
| 255 |
try:
|
| 256 |
return int(value)
|
| 257 |
except ValueError:
|
| 258 |
pass
|
| 259 |
+
return entity
|
| 260 |
|
| 261 |
|
| 262 |
def _atom37_from_flat(
|
|
|
|
| 794 |
metadata = ProteinComplexMetadata(
|
| 795 |
entity_lookup={
|
| 796 |
int(entity): _protein_entity_metadata_value(
|
| 797 |
+
self.metadata.entity_lookup.get(int(entity), int(entity)), int(entity)
|
| 798 |
)
|
| 799 |
for entity in unique_entities
|
| 800 |
},
|
fastplms/models/esmfold2/esmfold2_parsing.py
CHANGED
|
@@ -10,9 +10,22 @@ from contextlib import AbstractContextManager, nullcontext
|
|
| 10 |
from pathlib import Path
|
| 11 |
from typing import NamedTuple, TextIO
|
| 12 |
|
|
|
|
| 13 |
from .esmfold2_utils_types import PathOrBuffer
|
| 14 |
|
| 15 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 16 |
class FastaEntry(NamedTuple):
|
| 17 |
"""One FASTA record in source order."""
|
| 18 |
|
|
@@ -23,27 +36,8 @@ class FastaEntry(NamedTuple):
|
|
| 23 |
def parse_fasta(text: str) -> Generator[FastaEntry, None, None]:
|
| 24 |
"""Yield records from FASTA text without normalizing sequence symbols."""
|
| 25 |
|
| 26 |
-
|
| 27 |
-
|
| 28 |
-
found_record = False
|
| 29 |
-
|
| 30 |
-
for line in text.splitlines():
|
| 31 |
-
if not line or line.startswith("#"):
|
| 32 |
-
continue
|
| 33 |
-
if line.startswith(">"):
|
| 34 |
-
if header is not None:
|
| 35 |
-
found_record = True
|
| 36 |
-
yield FastaEntry(header, "".join(sequence_lines))
|
| 37 |
-
header = line[1:].strip()
|
| 38 |
-
sequence_lines.clear()
|
| 39 |
-
elif header is not None:
|
| 40 |
-
sequence_lines.append(line)
|
| 41 |
-
|
| 42 |
-
if header is not None:
|
| 43 |
-
found_record = True
|
| 44 |
-
yield FastaEntry(header, "".join(sequence_lines))
|
| 45 |
-
if not found_record:
|
| 46 |
-
raise ValueError("Found no sequences in input")
|
| 47 |
|
| 48 |
|
| 49 |
def _open_reader(source: PathOrBuffer) -> AbstractContextManager[TextIO]:
|
|
|
|
| 10 |
from pathlib import Path
|
| 11 |
from typing import NamedTuple, TextIO
|
| 12 |
|
| 13 |
+
from ...embeddings.inputs import FastaDialect, scan_fasta_lines
|
| 14 |
from .esmfold2_utils_types import PathOrBuffer
|
| 15 |
|
| 16 |
|
| 17 |
+
# Lines are used as given and `#` lines are comments. A header keeps its full text, and sequence data
|
| 18 |
+
# before the first header is skipped.
|
| 19 |
+
ESMFOLD2_FASTA = FastaDialect(
|
| 20 |
+
strip_lines=False,
|
| 21 |
+
comment_prefix="#",
|
| 22 |
+
first_word_header=False,
|
| 23 |
+
squeeze_sequence_whitespace=False,
|
| 24 |
+
orphan_message=None,
|
| 25 |
+
empty_message="Found no sequences in input",
|
| 26 |
+
)
|
| 27 |
+
|
| 28 |
+
|
| 29 |
class FastaEntry(NamedTuple):
|
| 30 |
"""One FASTA record in source order."""
|
| 31 |
|
|
|
|
| 36 |
def parse_fasta(text: str) -> Generator[FastaEntry, None, None]:
|
| 37 |
"""Yield records from FASTA text without normalizing sequence symbols."""
|
| 38 |
|
| 39 |
+
for record in scan_fasta_lines(text.splitlines(), ESMFOLD2_FASTA, source="input"):
|
| 40 |
+
yield FastaEntry(record.header, record.sequence)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 41 |
|
| 42 |
|
| 43 |
def _open_reader(source: PathOrBuffer) -> AbstractContextManager[TextIO]:
|
fastplms/models/esmfold2/modeling_esmfold2.py
CHANGED
|
@@ -57,6 +57,7 @@ from .esmfold2_constants_esm3 import (
|
|
| 57 |
)
|
| 58 |
from .modeling_esmfold2_common import (
|
| 59 |
CHAR_VOCAB_SIZE,
|
|
|
|
| 60 |
MAX_ATOMIC_NUMBER,
|
| 61 |
MSA_CONDITIONING_INPUT_NAMES,
|
| 62 |
NUM_RES_TYPES,
|
|
@@ -1490,7 +1491,7 @@ class ESMFold2Model(
|
|
| 1490 |
early_exit: bool = False,
|
| 1491 |
noise_scale: float | None = None,
|
| 1492 |
step_scale: float | None = None,
|
| 1493 |
-
max_inference_sigma: float | None =
|
| 1494 |
output_attentions: bool | None = None,
|
| 1495 |
output_hidden_states: bool | None = None,
|
| 1496 |
return_dict: bool | None = None,
|
|
@@ -1793,7 +1794,7 @@ class ESMFold2Model(
|
|
| 1793 |
seed: int | None = None,
|
| 1794 |
noise_scale: float | None = None,
|
| 1795 |
step_scale: float | None = None,
|
| 1796 |
-
max_inference_sigma:
|
| 1797 |
early_exit: bool = False,
|
| 1798 |
complex_id: str = "pred",
|
| 1799 |
verbose: bool = False,
|
|
|
|
| 57 |
)
|
| 58 |
from .modeling_esmfold2_common import (
|
| 59 |
CHAR_VOCAB_SIZE,
|
| 60 |
+
DEFAULT_MAX_INFERENCE_SIGMA,
|
| 61 |
MAX_ATOMIC_NUMBER,
|
| 62 |
MSA_CONDITIONING_INPUT_NAMES,
|
| 63 |
NUM_RES_TYPES,
|
|
|
|
| 1491 |
early_exit: bool = False,
|
| 1492 |
noise_scale: float | None = None,
|
| 1493 |
step_scale: float | None = None,
|
| 1494 |
+
max_inference_sigma: float | None = DEFAULT_MAX_INFERENCE_SIGMA,
|
| 1495 |
output_attentions: bool | None = None,
|
| 1496 |
output_hidden_states: bool | None = None,
|
| 1497 |
return_dict: bool | None = None,
|
|
|
|
| 1794 |
seed: int | None = None,
|
| 1795 |
noise_scale: float | None = None,
|
| 1796 |
step_scale: float | None = None,
|
| 1797 |
+
max_inference_sigma: float | None = DEFAULT_MAX_INFERENCE_SIGMA,
|
| 1798 |
early_exit: bool = False,
|
| 1799 |
complex_id: str = "pred",
|
| 1800 |
verbose: bool = False,
|
fastplms/models/esmfold2/modeling_esmfold2_common.py
CHANGED
|
@@ -62,6 +62,9 @@ _VALID_BACKENDS = (None, BACKEND_FUSED, BACKEND_CUEQ)
|
|
| 62 |
ATOM_ATTENTION_DENSE = "dense"
|
| 63 |
ATOM_ATTENTION_WINDOWED = "windowed"
|
| 64 |
_VALID_ATOM_ATTENTION = (ATOM_ATTENTION_DENSE, ATOM_ATTENTION_WINDOWED)
|
|
|
|
|
|
|
|
|
|
| 65 |
MSA_CONDITIONING_INPUT_NAMES = (
|
| 66 |
"msa",
|
| 67 |
"msa_attention_mask",
|
|
@@ -2793,7 +2796,7 @@ class DiffusionStructureHead(nn.Module):
|
|
| 2793 |
token_attention_mask: Tensor | None = None,
|
| 2794 |
num_diffusion_samples: int = 1,
|
| 2795 |
num_sampling_steps: int | None = None,
|
| 2796 |
-
max_inference_sigma: float | None =
|
| 2797 |
noise_scale: float | None = None,
|
| 2798 |
step_scale: float | None = None,
|
| 2799 |
return_atom_repr: bool = False,
|
|
|
|
| 62 |
ATOM_ATTENTION_DENSE = "dense"
|
| 63 |
ATOM_ATTENTION_WINDOWED = "windowed"
|
| 64 |
_VALID_ATOM_ATTENTION = (ATOM_ATTENTION_DENSE, ATOM_ATTENTION_WINDOWED)
|
| 65 |
+
# The official sampler starts every fold at this noise level unless a caller overrides it;
|
| 66 |
+
# the public forwards and fold() default to it so an omitted argument keeps the cap.
|
| 67 |
+
DEFAULT_MAX_INFERENCE_SIGMA = 256.0
|
| 68 |
MSA_CONDITIONING_INPUT_NAMES = (
|
| 69 |
"msa",
|
| 70 |
"msa_attention_mask",
|
|
|
|
| 2796 |
token_attention_mask: Tensor | None = None,
|
| 2797 |
num_diffusion_samples: int = 1,
|
| 2798 |
num_sampling_steps: int | None = None,
|
| 2799 |
+
max_inference_sigma: float | None = DEFAULT_MAX_INFERENCE_SIGMA,
|
| 2800 |
noise_scale: float | None = None,
|
| 2801 |
step_scale: float | None = None,
|
| 2802 |
return_atom_repr: bool = False,
|
fastplms/models/esmfold2/modeling_esmfold2_experimental.py
CHANGED
|
@@ -21,8 +21,8 @@ from tqdm.auto import tqdm
|
|
| 21 |
from transformers.modeling_utils import PreTrainedModel
|
| 22 |
|
| 23 |
from .attention import ESMFold2AttentionMixin
|
| 24 |
-
from .configuration_esmfold2 import ESMFold2Config
|
| 25 |
from .confidence_checkpoint import install_confidence_checkpoint
|
|
|
|
| 26 |
from .embedding import ESMFold2EmbeddingMixin
|
| 27 |
from .modeling_esmfold2 import (
|
| 28 |
ESMCPrecision,
|
|
@@ -38,6 +38,7 @@ from .modeling_esmfold2 import (
|
|
| 38 |
)
|
| 39 |
from .modeling_esmfold2_common import (
|
| 40 |
CHAR_VOCAB_SIZE,
|
|
|
|
| 41 |
MAX_ATOMIC_NUMBER,
|
| 42 |
MSA_CONDITIONING_INPUT_NAMES,
|
| 43 |
NUM_RES_TYPES,
|
|
@@ -731,7 +732,7 @@ class ESMFold2ExperimentalModel(ESMFold2EmbeddingMixin, ESMFold2AttentionMixin,
|
|
| 731 |
provide_soft_sequence_to_msa_and_profile: bool = True,
|
| 732 |
noise_scale: float | None = None,
|
| 733 |
step_scale: float | None = None,
|
| 734 |
-
max_inference_sigma: float | None =
|
| 735 |
output_attentions: bool | None = None,
|
| 736 |
output_hidden_states: bool | None = None,
|
| 737 |
return_dict: bool | None = None,
|
|
@@ -1094,7 +1095,7 @@ class ESMFold2ExperimentalModel(ESMFold2EmbeddingMixin, ESMFold2AttentionMixin,
|
|
| 1094 |
seed: int | None = None,
|
| 1095 |
noise_scale: float | None = None,
|
| 1096 |
step_scale: float | None = None,
|
| 1097 |
-
max_inference_sigma:
|
| 1098 |
early_exit: bool = False,
|
| 1099 |
complex_id: str = "pred",
|
| 1100 |
verbose: bool = False,
|
|
|
|
| 21 |
from transformers.modeling_utils import PreTrainedModel
|
| 22 |
|
| 23 |
from .attention import ESMFold2AttentionMixin
|
|
|
|
| 24 |
from .confidence_checkpoint import install_confidence_checkpoint
|
| 25 |
+
from .configuration_esmfold2 import ESMFold2Config
|
| 26 |
from .embedding import ESMFold2EmbeddingMixin
|
| 27 |
from .modeling_esmfold2 import (
|
| 28 |
ESMCPrecision,
|
|
|
|
| 38 |
)
|
| 39 |
from .modeling_esmfold2_common import (
|
| 40 |
CHAR_VOCAB_SIZE,
|
| 41 |
+
DEFAULT_MAX_INFERENCE_SIGMA,
|
| 42 |
MAX_ATOMIC_NUMBER,
|
| 43 |
MSA_CONDITIONING_INPUT_NAMES,
|
| 44 |
NUM_RES_TYPES,
|
|
|
|
| 732 |
provide_soft_sequence_to_msa_and_profile: bool = True,
|
| 733 |
noise_scale: float | None = None,
|
| 734 |
step_scale: float | None = None,
|
| 735 |
+
max_inference_sigma: float | None = DEFAULT_MAX_INFERENCE_SIGMA,
|
| 736 |
output_attentions: bool | None = None,
|
| 737 |
output_hidden_states: bool | None = None,
|
| 738 |
return_dict: bool | None = None,
|
|
|
|
| 1095 |
seed: int | None = None,
|
| 1096 |
noise_scale: float | None = None,
|
| 1097 |
step_scale: float | None = None,
|
| 1098 |
+
max_inference_sigma: float | None = DEFAULT_MAX_INFERENCE_SIGMA,
|
| 1099 |
early_exit: bool = False,
|
| 1100 |
complex_id: str = "pred",
|
| 1101 |
verbose: bool = False,
|
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 |
-
|
| 992 |
-
|
|
|
|
|
|
|
|
|
|
| 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 != "
|
| 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 != "32ac853f73c3a103bc49f8896703b0b5a1f309b0efcaf9206f795e3090be745d":
|
| 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/ESMFold2-Fast.
|
| 2 |
# FastPLMs source is embedded in this model repository.
|
| 3 |
-
torch>=2.
|
| 4 |
-
transformers>=5.
|
| 5 |
-
huggingface-hub>=
|
| 6 |
-
tokenizers>=0.
|
| 7 |
-
safetensors>=0.
|
| 8 |
-
numpy>=
|
| 9 |
-
einops>=0.8
|
| 10 |
-
tqdm>=4.
|
| 11 |
-
accelerate>=1.
|
| 12 |
-
biopython>=1.
|
| 13 |
-
biotite>=1.
|
| 14 |
-
brotli>=1.
|
| 15 |
-
msgpack>=1.
|
| 16 |
-
msgpack-numpy>=0.4.8
|
| 17 |
-
omegaconf>=2.3
|
| 18 |
-
rdkit>=
|
| 19 |
-
scipy>=1.
|
| 20 |
-
zstandard>=0.
|
|
|
|
| 1 |
# Direct runtime dependencies for Synthyra/ESMFold2-Fast.
|
| 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
|