File size: 17,848 Bytes
9641d1d
feb1b1c
 
9641d1d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
feb1b1c
9641d1d
 
 
 
 
 
 
 
 
4bb0bdf
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9641d1d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
bcdd459
 
 
 
9641d1d
 
 
 
 
 
 
 
 
 
bcdd459
 
9641d1d
 
4bb0bdf
 
 
 
9641d1d
 
4bb0bdf
9641d1d
 
 
 
 
 
 
 
 
 
 
 
 
 
830d137
9641d1d
 
 
4bb0bdf
 
 
 
 
 
 
 
 
9641d1d
 
 
 
4bb0bdf
 
 
 
9641d1d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4bb0bdf
 
 
 
 
 
 
 
 
9641d1d
 
 
 
 
 
4bb0bdf
9641d1d
 
 
 
 
 
 
 
 
 
4bb0bdf
 
 
 
 
 
 
 
9641d1d
 
 
 
 
 
4bb0bdf
 
9641d1d
 
 
 
 
 
 
4bb0bdf
 
 
9641d1d
 
 
4bb0bdf
 
 
9641d1d
 
 
 
 
 
 
 
 
 
 
4bb0bdf
9641d1d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4bb0bdf
 
 
 
 
 
 
 
 
830d137
 
 
 
 
 
 
 
 
 
 
bcdd459
 
 
 
 
 
 
 
 
 
9641d1d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4bb0bdf
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9641d1d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4bb0bdf
9641d1d
 
 
4bb0bdf
9641d1d
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471

from __future__ import annotations

import importlib
import os
from collections.abc import Mapping, Sequence
from pathlib import Path
from typing import TYPE_CHECKING, Final, Protocol, cast

import numpy as np
import numpy.typing as npt

from redstack.ports._types import FloatMatrix
from redstack.ports.embedding import EmbeddingError

#: Environment variables whose value ``"online"`` forbids importing this module.
_ONLINE_PROFILE_ENVS: Final[tuple[str, ...]] = (
    "REDSTACK_EXECUTION_PROFILE",
    "REDSTACK_RUNTIME_PROFILE",
)
#: A boolean online flag whose truthy value also forbids import.
_ONLINE_FLAG_ENV: Final[str] = "REDSTACK_ONLINE"
_TRUTHY: Final[frozenset[str]] = frozenset({"1", "true", "yes", "on"})
#: Default per-call batch size (throughput hint; never affects results).
_DEFAULT_BATCH_SIZE: Final[int] = 32
#: Default ONNX export opset and the accepted parity floor.
_DEFAULT_OPSET: Final[int] = 17
_PARITY_FLOOR: Final[float] = 0.999

#: Explicit override for auto device selection (build-time only; never read by
#: anything under pipelines.online). Unset means "auto-detect".
_DEVICE_ENV_VAR: Final[str] = "REDSTACK_OFFLINE_DEVICE"
_VALID_DEVICES: Final[frozenset[str]] = frozenset({"cpu", "mps", "cuda"})


def _select_device(torch_module: _Torch) -> str:
    """Resolve the offline encode device: override, else best available accelerator.

    Priority: an explicit ``REDSTACK_OFFLINE_DEVICE`` env override, then CUDA, then
    Apple MPS, else CPU. This picks the device for ``encode`` only —
    ``export_onnx`` always traces on CPU regardless of this choice (Adapters §4).

    Raises:
        EmbeddingError: the env override names a device outside ``_VALID_DEVICES``.
    """
    override = os.environ.get(_DEVICE_ENV_VAR, "").strip().lower()
    if override:
        if override not in _VALID_DEVICES:
            raise EmbeddingError(
                f"{_DEVICE_ENV_VAR}={override!r} must be one of "
                f"{sorted(_VALID_DEVICES)}"
            )
        return override
    if torch_module.cuda.is_available():
        return "cuda"
    if torch_module.backends.mps.is_available():
        return "mps"
    return "cpu"


def _guard_offline_only() -> None:
    """Raise if an online execution marker is set (import + construction guard).

    Raises:
        RuntimeError: an online profile/flag is present in the environment.
    """
    for key in _ONLINE_PROFILE_ENVS:
        if os.environ.get(key, "").strip().lower() == "online":
            raise RuntimeError(
                f"adapters.st_embedder is offline-only but {key}=online is set"
            )
    if os.environ.get(_ONLINE_FLAG_ENV, "").strip().lower() in _TRUTHY:
        raise RuntimeError(
            f"adapters.st_embedder is offline-only but {_ONLINE_FLAG_ENV} is truthy"
        )


# Import-time guard (defence in depth alongside the import-linter contract).
_guard_offline_only()


# --------------------------------------------------------------------------- #
# Minimal structural views over the untyped runtimes (loaded via importlib so no
# untyped ``import torch`` / ``import sentence_transformers`` statement enters
# the typed surface; concrete objects are narrowed by ``cast``).
# --------------------------------------------------------------------------- #
class _StModel(Protocol):
    def encode(
        self,
        sentences: Sequence[str],
        *,
        batch_size: int,
        convert_to_numpy: bool,
        normalize_embeddings: bool,
    ) -> FloatMatrix: ...
    def get_sentence_embedding_dimension(self) -> int: ...
    def __getitem__(self, index: int) -> object: ...
    @property
    def tokenizer(self) -> _HfTokenizer: ...


class _FastTokenizerHandle(Protocol):
    def to_str(self) -> str: ...


class _HfTokenizer(Protocol):
    def __call__(
        self,
        text: Sequence[str],
        *,
        padding: bool,
        truncation: bool,
        max_length: int,
        return_tensors: str,
    ) -> dict[str, object]: ...
    @property
    def backend_tokenizer(self) -> _FastTokenizerHandle: ...


class _TorchModule(Protocol):
    def to(self, device: str) -> _TorchModule: ...


class _Pooling(Protocol):
    @property
    def auto_model(self) -> _TorchModule: ...


class _TorchOnnx(Protocol):
    def export(
        self,
        model: object,
        args: object,
        f: str,
        *,
        input_names: list[str],
        output_names: list[str],
        dynamic_axes: Mapping[str, Mapping[int, str]],
        opset_version: int,
        do_constant_folding: bool,
        dynamo: bool,
    ) -> None: ...


class _TorchAccelerator(Protocol):
    def is_available(self) -> bool: ...


class _TorchBackends(Protocol):
    @property
    def mps(self) -> _TorchAccelerator: ...


class _Torch(Protocol):
    def set_num_threads(self, n: int) -> None: ...
    @property
    def onnx(self) -> _TorchOnnx: ...
    @property
    def cuda(self) -> _TorchAccelerator: ...
    @property
    def backends(self) -> _TorchBackends: ...


class _OrtSession(Protocol):
    def run(
        self, output_names: list[str], input_feed: Mapping[str, npt.NDArray[np.int64]]
    ) -> list[npt.NDArray[np.float32]]: ...


class SentenceTransformerEmbeddingAdapter:
    """Offline sentence-transformers encoder + ONNX-twin exporter.

    Constructed only inside ``pipelines/offline``. Loads the pinned-revision
    model in eval mode under the HuggingFace offline environment; serves
    ``encode`` and exports the onnx twin with parity verification.
    """

    __slots__ = (
        "_batch_size",
        "_device",
        "_dim",
        "_model",
        "_model_id",
        "_normalize",
        "_torch",
    )

    def __init__(
        self,
        model_id: str,
        *,
        revision: str | None = None,
        device: str = "auto",
        torch_num_threads: int = 1,
        normalize: bool = True,
        batch_size: int = _DEFAULT_BATCH_SIZE,
        dim: int | None = None,
    ) -> None:
        """Load the pinned model under the offline environment.

        Args:
            model_id: The pinned sentence-transformers model id (provenance).
            revision: Pinned model revision for reproducibility.
            device: Compute device for :meth:`encode` — ``"auto"`` (default)
                detects CUDA, then Apple MPS, then falls back to CPU; an explicit
                ``"cpu"``/``"mps"``/``"cuda"`` pins that device. ``"auto"`` is
                itself overridable via the ``REDSTACK_OFFLINE_DEVICE`` env var.
                :meth:`export_onnx` always traces on CPU regardless of this value
                (Adapters §4) — the device choice only affects encode throughput,
                never the exported artifact's correctness.
            torch_num_threads: Pinned torch thread count (CPU path only).
            normalize: Apply L2 normalization to outputs (fixed contract).
            batch_size: Default encode batch size.
            dim: Optional dimensionality override; otherwise read from the model.

        Raises:
            RuntimeError: an online marker is set.
            EmbeddingError: the model could not be loaded, or an explicit
                ``REDSTACK_OFFLINE_DEVICE`` override names an unknown device.
        """
        _guard_offline_only()
        os.environ.setdefault("HF_HUB_OFFLINE", "1")
        os.environ.setdefault("TRANSFORMERS_OFFLINE", "1")

        try:
            torch_module = cast("_Torch", importlib.import_module("torch"))
            resolved_device = (
                _select_device(torch_module) if device == "auto" else device
            )
            torch_module.set_num_threads(torch_num_threads)
            st_module = importlib.import_module("sentence_transformers")
            transformer_cls = st_module.SentenceTransformer
            loaded = transformer_cls(
                model_id, revision=revision, device=resolved_device
            )
        except (ImportError, OSError, ValueError, RuntimeError) as exc:
            raise EmbeddingError(
                f"cannot load sentence-transformers model {model_id!r}: {exc}"
            ) from exc

        model = cast("_StModel", loaded)
        self._torch: Final[_Torch] = torch_module
        self._model: Final[_StModel] = model
        self._model_id: Final[str] = model_id
        self._normalize: Final[bool] = normalize
        self._batch_size: Final[int] = batch_size
        self._device: Final[str] = resolved_device
        self._dim: Final[int] = (
            dim if dim is not None else int(model.get_sentence_embedding_dimension())
        )

    # ------------------------------------------------------------------ #
    # Port surface.
    # ------------------------------------------------------------------ #
    @property
    def dim(self) -> int:
        """The fixed output dimensionality."""
        return self._dim

    @property
    def model_id(self) -> str:
        """The stable model identifier for provenance."""
        return self._model_id

    @property
    def device(self) -> str:
        """The resolved compute device used by ``encode`` (``cpu``/``mps``/``cuda``).

        Build provenance only (Adapters §4 device policy) — ``export_onnx``
        always traces on CPU regardless of this value.
        """
        return self._device

    @property
    def opset(self) -> int:
        """The pinned ONNX opset :meth:`export_onnx` exports at by default.

        Required by the ``OnnxExportCapable`` Protocol (a ``runtime_checkable``
        Protocol checks attribute *presence*, not signature — without this
        property ``isinstance(adapter, OnnxExportCapable)`` is ``False`` even
        though :meth:`export_onnx` itself is fully implemented).
        """
        return _DEFAULT_OPSET

    @property
    def tokenizer_json(self) -> str:
        """The fast tokenizer's ``tokenizers.Tokenizer.from_str`` JSON payload.

        The online ``OnnxEmbeddingModelAdapter`` fallback encoder tokenizes
        through this exact serialization, so it must travel as its own
        artifact alongside ``model/encoder.onnx`` (Adapters §4).
        """
        return self._model.tokenizer.backend_tokenizer.to_str()

    def encode(
        self, texts: Sequence[str], *, batch_size: int | None = None
    ) -> FloatMatrix:
        """Encode pre-composed documents into a read-only ``(len(texts), dim)`` matrix.

        Output is ``float32``, each row L2-normalized within epsilon, row order
        equal to input order regardless of batching.

        Raises:
            EmbeddingError: the encode operation failed.
        """
        n = len(texts)
        if n == 0:
            empty = np.empty((0, self._dim), dtype=np.float32)
            empty.flags.writeable = False
            return empty

        step = batch_size if batch_size is not None and batch_size > 0 else self._batch_size
        try:
            raw = self._model.encode(
                list(texts),
                batch_size=step,
                convert_to_numpy=True,
                normalize_embeddings=self._normalize,
            )
        except Exception as exc:  # the library raises bare exceptions
            raise EmbeddingError(f"sentence-transformers encode failed: {exc}") from exc

        matrix = np.asarray(raw, dtype=np.float32)
        if matrix.ndim != 2 or matrix.shape != (n, self._dim):
            raise EmbeddingError(
                f"encoded shape {matrix.shape} != expected {(n, self._dim)}"
            )
        matrix.flags.writeable = False
        return matrix

    # ------------------------------------------------------------------ #
    # ONNX export + parity (offline build responsibility, Adapters §4).
    # ------------------------------------------------------------------ #
    def export_onnx(
        self,
        output_path: Path,
        *,
        opset: int = _DEFAULT_OPSET,
        sample_texts: Sequence[str] | None = None,
        max_seq_length: int = 256,
    ) -> float:
        """Export the transformer twin to ``output_path`` and verify st↔onnx parity.

        Exports the underlying HuggingFace transformer (token-embedding output);
        the online onnx adapter applies the matching mean pooling + L2 norm. A
        sample is encoded through both paths and the mean cosine is asserted
        ``>= 0.999``.

        Args:
            output_path: Destination ``.onnx`` path.
            opset: Pinned ONNX opset version.
            sample_texts: Texts for the parity check (a small default if omitted).
            max_seq_length: Tokenization truncation length for the export sample.

        Returns:
            The mean cosine similarity between st and onnx sentence embeddings.

        Raises:
            EmbeddingError: export failed or parity fell below the floor.
        """
        samples = list(sample_texts) if sample_texts else ["the quick brown fox", "hello world"]
        try:
            transformer = cast("_Pooling", self._model[0]).auto_model
            tokenizer = self._model.tokenizer
            tokens = tokenizer(
                samples,
                padding=True,
                truncation=True,
                max_length=max_seq_length,
                return_tensors="pt",
            )
            input_ids = tokens["input_ids"]
            attention_mask = tokens["attention_mask"]
            output_path.parent.mkdir(parents=True, exist_ok=True)
            # torch.onnx.export traces most reliably off a CPU-resident model
            # (the tokenizer's "pt" tensors are already CPU); this is a fixed
            # small-sample trace, not the encode throughput path, so it is
            # always pinned to CPU regardless of self._device and restored
            # afterward so a later encode() still runs on the configured device.
            transformer.to("cpu")
            try:
                self._torch.onnx.export(
                    transformer,
                    (input_ids, attention_mask),
                    str(output_path),
                    input_names=["input_ids", "attention_mask"],
                    output_names=["last_hidden_state"],
                    dynamic_axes={
                        "input_ids": {0: "batch", 1: "sequence"},
                        "attention_mask": {0: "batch", 1: "sequence"},
                        "last_hidden_state": {0: "batch", 1: "sequence"},
                    },
                    opset_version=opset,
                    do_constant_folding=True,
                    # Force the legacy TorchScript-based exporter: the dynamic_axes
                    # kwarg above is that exporter's API, and the newer dynamo=True
                    # default (torch>=2.5) requires an additional `onnxscript`
                    # dependency this offline build does not declare.
                    dynamo=False,
                )
            finally:
                transformer.to(self._device)
        except Exception as exc:  # torch/transformers raise bare exceptions
            raise EmbeddingError(f"onnx export failed: {exc}") from exc

        parity = self._parity_cosine(output_path, samples, max_seq_length)
        if parity < _PARITY_FLOOR:
            raise EmbeddingError(
                f"st<->onnx parity {parity:.6f} below floor {_PARITY_FLOOR}"
            )
        return parity

    def _parity_cosine(
        self, onnx_path: Path, samples: Sequence[str], max_seq_length: int
    ) -> float:
        """Mean cosine between st embeddings and onnx-pooled embeddings."""
        try:
            ort_module = importlib.import_module("onnxruntime")
            session = cast(
                "_OrtSession",
                ort_module.InferenceSession(
                    str(onnx_path), providers=["CPUExecutionProvider"]
                ),
            )
            tokenizer = self._model.tokenizer
            tokens = tokenizer(
                list(samples),
                padding=True,
                truncation=True,
                max_length=max_seq_length,
                return_tensors="np",
            )
            input_ids = np.asarray(
                cast("npt.NDArray[np.int64]", tokens["input_ids"]), dtype=np.int64
            )
            attention_mask = np.asarray(
                cast("npt.NDArray[np.int64]", tokens["attention_mask"]), dtype=np.int64
            )
            raw = session.run(
                ["last_hidden_state"],
                {"input_ids": input_ids, "attention_mask": attention_mask},
            )
        except Exception as exc:
            raise EmbeddingError(f"onnx parity inference failed: {exc}") from exc

        token_output = np.asarray(raw[0], dtype=np.float32)
        mask = attention_mask.astype(np.float32)[:, :, None]
        summed = np.sum(token_output * mask, axis=1)
        counts = np.clip(mask.sum(axis=1), a_min=1e-9, a_max=None)
        onnx_vecs = self._unit_rows((summed / counts).astype(np.float32, copy=False))

        st_vecs = self._unit_rows(np.asarray(self.encode(list(samples)), dtype=np.float32))
        cosines = np.sum(onnx_vecs * st_vecs, axis=1)
        return float(np.mean(cosines))

    @staticmethod
    def _unit_rows(matrix: FloatMatrix) -> FloatMatrix:
        norms = np.clip(np.linalg.norm(matrix, axis=1, keepdims=True), a_min=1e-12, a_max=None)
        unit: FloatMatrix = (matrix / norms).astype(np.float32, copy=False)
        return unit


if TYPE_CHECKING:
    from redstack.ports.embedding import DeviceReporting, EmbeddingModelPort

    # Compile-time structural conformance to the frozen port surface.
    _PORT_CONFORMANCE: type[EmbeddingModelPort] = SentenceTransformerEmbeddingAdapter
    _DEVICE_CONFORMANCE: type[DeviceReporting] = SentenceTransformerEmbeddingAdapter


__all__: tuple[str, ...] = ("SentenceTransformerEmbeddingAdapter",)