semsift 0.0.1__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
semsift/__init__.py ADDED
@@ -0,0 +1,3 @@
1
+ """Building blocks for hybrid keyword and vector retrieval."""
2
+
3
+ __version__ = "0.0.1"
@@ -0,0 +1,5 @@
1
+ """Splitting text into chunks for the store."""
2
+
3
+ from .chunkers import Chunk, TextChunker, TreeSitterChunker
4
+
5
+ __all__ = ["Chunk", "TextChunker", "TreeSitterChunker"]
@@ -0,0 +1,206 @@
1
+ """Splitting text into chunks whose spans double as citation spans."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from bisect import bisect_right
6
+ from dataclasses import dataclass
7
+
8
+
9
+ @dataclass(frozen=True)
10
+ class Chunk:
11
+ """`text` is exactly `source[start:end]`; lines are 1-based and inclusive.
12
+
13
+ `context` is the enclosing definitions (e.g. a class name) and
14
+ `symbols` the names defined inside, when the chunker knows them.
15
+ """
16
+
17
+ text: str
18
+ start: int
19
+ end: int
20
+ start_line: int
21
+ end_line: int
22
+ context: tuple[str, ...] = ()
23
+ symbols: tuple[str, ...] = ()
24
+
25
+
26
+ def _build(text: str, spans: list[tuple[int, int]], min_chars: int,
27
+ extras: dict[int, tuple[tuple[str, ...], tuple[str, ...]]] | None = None) -> list[Chunk]:
28
+ newlines = [i for i, ch in enumerate(text)
29
+ if ch == "\n" or (ch == "\r" and (i + 1 == len(text) or text[i + 1] != "\n"))]
30
+ out = []
31
+ for start, end in spans:
32
+ body = text[start:end]
33
+ if len(body.strip()) < min_chars:
34
+ continue
35
+ context, symbols = (extras or {}).get(start, ((), ()))
36
+ out.append(Chunk(body, start, end, bisect_right(newlines, start - 1) + 1,
37
+ bisect_right(newlines, max(start, end - 1) - 1) + 1,
38
+ context, symbols))
39
+ return out
40
+
41
+
42
+ def _check(max_chars: int, min_chars: int) -> None:
43
+ if max_chars <= 0:
44
+ raise ValueError("max_chars must be positive")
45
+ if not 0 <= min_chars <= max_chars:
46
+ raise ValueError("min_chars must be between 0 and max_chars")
47
+
48
+
49
+ class TextChunker:
50
+ """Line-based windows of at most `max_chars`, split at natural breaks.
51
+
52
+ With `markdown`, an ATX heading outside a fenced block always starts a
53
+ chunk; turn it off for text where `#` begins a comment. After a blank
54
+ line, a chunk at least half full ends before the next paragraph. A
55
+ line longer than `max_chars` is cut into pieces. Chunks whose stripped
56
+ text is shorter than `min_chars` are dropped, so the spans tile the
57
+ text only when `min_chars` is 0.
58
+ """
59
+
60
+ def __init__(self, max_chars: int = 750, min_chars: int = 1, *,
61
+ markdown: bool = True) -> None:
62
+ _check(max_chars, min_chars)
63
+ self.max_chars, self.min_chars, self.markdown = max_chars, min_chars, markdown
64
+
65
+ def spans(self, text: str) -> list[tuple[int, int]]:
66
+ spans: list[tuple[int, int]] = []
67
+ start = pos = 0
68
+ after_blank = False
69
+ fence: tuple[str, int] | None = None
70
+ for line in text.splitlines(keepends=True):
71
+ size = pos - start
72
+ blank = not line.strip()
73
+ marker = _markdown_fence(line)
74
+ heading = self.markdown and fence is None and _markdown_heading(line)
75
+ if size and (heading
76
+ or (after_blank and not blank and size >= self.max_chars // 2)
77
+ or size + len(line) > self.max_chars):
78
+ spans.append((start, pos))
79
+ start = pos
80
+ if len(line) > self.max_chars:
81
+ for cut in range(pos, pos + len(line), self.max_chars):
82
+ spans.append((cut, min(cut + self.max_chars, pos + len(line))))
83
+ start = pos + len(line)
84
+ pos += len(line)
85
+ after_blank = blank
86
+ if marker is not None:
87
+ char, width, rest = marker
88
+ if fence is None:
89
+ fence = (char, width)
90
+ elif char == fence[0] and width >= fence[1] and not rest.strip():
91
+ fence = None
92
+ if pos > start:
93
+ spans.append((start, pos))
94
+ return spans
95
+
96
+ def chunk(self, text: str) -> list[Chunk]:
97
+ return _build(text, self.spans(text), self.min_chars)
98
+
99
+
100
+ class TreeSitterChunker:
101
+ """Syntax-aligned chunks from tree-sitter-language-pack's chunker.
102
+
103
+ Needs the `tree-sitter` extra. The language pack downloads a
104
+ grammar the first time a language is parsed. A language it does not
105
+ know, a grammar it cannot download, or a failure to parse falls back
106
+ to `fallback`: a TextChunker of the same bounds that does not read
107
+ markdown, by default. So does a source over `max_source_bytes`, or a
108
+ parse that runs past `parse_timeout_ms`, so one generated or minified
109
+ file cannot stall indexing.
110
+ """
111
+
112
+ def __init__(self, max_chars: int = 750, min_chars: int = 1,
113
+ fallback: TextChunker | None = None, *,
114
+ max_source_bytes: int = 5_000_000, parse_timeout_ms: int = 5_000) -> None:
115
+ _check(max_chars, min_chars)
116
+ if (isinstance(max_source_bytes, bool)
117
+ or not isinstance(max_source_bytes, int) or max_source_bytes <= 0
118
+ or isinstance(parse_timeout_ms, bool)
119
+ or not isinstance(parse_timeout_ms, int) or parse_timeout_ms <= 0):
120
+ raise ValueError("max_source_bytes and parse_timeout_ms must be positive integers")
121
+ self.max_chars, self.min_chars = max_chars, min_chars
122
+ self.max_source_bytes, self.parse_timeout_ms = max_source_bytes, parse_timeout_ms
123
+ self.fallback = fallback or TextChunker(max_chars, min_chars, markdown=False)
124
+
125
+ def supports(self, language: str) -> bool:
126
+ """Whether the pack knows `language`; its grammar may still need a download."""
127
+ import tree_sitter_language_pack as pack
128
+
129
+ return language in pack.manifest_languages()
130
+
131
+ def chunk(self, text: str, language: str) -> list[Chunk]:
132
+ import tree_sitter_language_pack as pack
133
+
134
+ if not text.strip():
135
+ return self.fallback.chunk(text)
136
+ if not self.supports(language) or len(text.encode()) > self.max_source_bytes:
137
+ return self.fallback.chunk(text)
138
+ try:
139
+ result = pack.process(text, pack.ProcessConfig(
140
+ language=language, structure=False, imports=False, exports=False,
141
+ chunk_max_size=self.max_chars, max_source_bytes=self.max_source_bytes,
142
+ parse_timeout_ms=self.parse_timeout_ms))
143
+ except pack.Error:
144
+ return self.fallback.chunk(text)
145
+ pieces = sorted(result.chunks, key=lambda c: c.start_byte)
146
+ if not pieces:
147
+ return self.fallback.chunk(text)
148
+ data = text.encode()
149
+ # Chunks are cut at their start bytes, so they tile the text even
150
+ # if the pack leaves a gap; each start is moved back to a
151
+ # character boundary before converting to a character offset.
152
+ starts = sorted({0} | {_char_start(data, c.start_byte) for c in pieces})
153
+ chars = _char_offsets(data, starts)
154
+ extras = {}
155
+ for c in pieces:
156
+ meta = c.metadata
157
+ if meta is not None:
158
+ key = chars[_char_start(data, c.start_byte)]
159
+ extras.setdefault(key, (tuple(meta.context_path), tuple(meta.symbols_defined)))
160
+ bounds = [chars[b] for b in starts] + [len(text)]
161
+ spans = [(a, b) for a, b in zip(bounds, bounds[1:]) if b > a]
162
+ if any(end - start > self.max_chars for start, end in spans):
163
+ return self.fallback.chunk(text)
164
+ return _build(text, spans, self.min_chars, extras)
165
+
166
+
167
+ def _char_start(data: bytes, offset: int) -> int:
168
+ """`offset` moved back to the first byte of the UTF-8 character it is in."""
169
+ offset = min(max(offset, 0), len(data))
170
+ while 0 < offset < len(data) and data[offset] & 0xC0 == 0x80:
171
+ offset -= 1
172
+ return offset
173
+
174
+
175
+ def _char_offsets(data: bytes, byte_offsets: list[int]) -> dict[int, int]:
176
+ """Character offset of each sorted byte offset, which must be on a boundary."""
177
+ out, chars, prev = {}, 0, 0
178
+ for b in byte_offsets:
179
+ chars += len(data[prev:b].decode())
180
+ out[b] = chars
181
+ prev = b
182
+ return out
183
+
184
+
185
+ def _markdown_heading(line: str) -> bool:
186
+ body = line.rstrip("\r\n")
187
+ stripped = body.lstrip(" ")
188
+ if len(body) - len(stripped) > 3:
189
+ return False
190
+ width = len(stripped) - len(stripped.lstrip("#"))
191
+ return 1 <= width <= 6 and (width == len(stripped) or stripped[width] in " \t")
192
+
193
+
194
+ def _markdown_fence(line: str) -> tuple[str, int, str] | None:
195
+ body = line.rstrip("\r\n")
196
+ stripped = body.lstrip(" ")
197
+ if len(body) - len(stripped) > 3 or not stripped or stripped[0] not in "`~":
198
+ return None
199
+ char = stripped[0]
200
+ width = len(stripped) - len(stripped.lstrip(char))
201
+ if width < 3:
202
+ return None
203
+ rest = stripped[width:]
204
+ if char == "`" and "`" in rest:
205
+ return None
206
+ return char, width, rest
@@ -0,0 +1,12 @@
1
+ """Text to vectors.
2
+
3
+ Imports nothing else from semsift, so it can become its own package.
4
+ """
5
+
6
+ from .backends import (Encoder, FakeEncoder, HttpEncoder, OnnxEncoder,
7
+ StaticEncoder)
8
+ from .policy import resolve_prefix
9
+ from .space import VectorSpace
10
+
11
+ __all__ = ["Encoder", "FakeEncoder", "HttpEncoder", "OnnxEncoder",
12
+ "StaticEncoder", "VectorSpace", "resolve_prefix"]
@@ -0,0 +1,433 @@
1
+ """Text to vectors: the encoder protocol and its backends.
2
+
3
+ The same encoder must encode documents and queries; `space` names what
4
+ has to match for stored vectors to stay comparable.
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ import hashlib
10
+ import json
11
+ import math
12
+ import urllib.request
13
+ from pathlib import Path
14
+ from typing import Protocol, Sequence
15
+
16
+ from .policy import default_pooling, resolve_prefix
17
+ from .space import VectorSpace
18
+
19
+
20
+ class Encoder(Protocol):
21
+ """Backends implement `_encode`; callers use `encode`/`encode_query`.
22
+
23
+ Only the public pair applies the side-specific prefixes, so a backend
24
+ that encoded through `_encode` directly would skip the model's
25
+ query/document markers.
26
+ """
27
+
28
+ name: str
29
+ backend: str
30
+ dims: int
31
+ query_prefix: str
32
+ doc_prefix: str
33
+
34
+ @property
35
+ def space(self) -> VectorSpace: ...
36
+
37
+ def _encode(self, texts: Sequence[str]) -> list[Sequence[float]]: ...
38
+
39
+ def encode(self, texts: Sequence[str]) -> list[Sequence[float]]: ...
40
+
41
+ def encode_query(self, texts: Sequence[str]) -> list[Sequence[float]]: ...
42
+
43
+
44
+ class _PrefixMixin:
45
+ """Side-specific prefixes and the vector space, in one place.
46
+
47
+ Retrieval models are trained with markers that differ by side. bge
48
+ marks only the query; e5 wants `query: ` and `passage: `; nomic wants
49
+ `search_query: ` and `search_document: `. Applying the wrong one, or
50
+ only half the pair, puts queries and documents in different regions
51
+ of the space and nothing errors -- recall just drops. Backends define
52
+ `_encode` and inherit both public methods so no backend can get this
53
+ subtly different.
54
+ """
55
+
56
+ name: str
57
+ backend: str
58
+ dims: int
59
+ query_prefix: str = ""
60
+ doc_prefix: str = ""
61
+ pooling: str = ""
62
+ variant: str = ""
63
+
64
+ def _prefixes(self, query_prefix: str | None,
65
+ doc_prefix: str | None) -> None:
66
+ """`None` takes the model family's default; `''` means none."""
67
+ self.query_prefix = resolve_prefix(query_prefix, self.name, "query")
68
+ self.doc_prefix = resolve_prefix(doc_prefix, self.name, "doc")
69
+
70
+ @property
71
+ def space(self) -> VectorSpace:
72
+ return VectorSpace(model=self.name, backend=self.backend,
73
+ variant=self.variant, dims=self.dims,
74
+ doc_prefix=self.doc_prefix, pooling=self.pooling)
75
+
76
+ def encode(self, texts: Sequence[str]) -> list[Sequence[float]]:
77
+ """Documents, with the model's document-side marker if it has one."""
78
+ if not self.doc_prefix:
79
+ return self._encode(texts) # type: ignore[attr-defined]
80
+ return self._encode( # type: ignore[attr-defined]
81
+ [self.doc_prefix + t for t in texts]
82
+ )
83
+
84
+ def encode_query(self, texts: Sequence[str]) -> list[Sequence[float]]:
85
+ return self._encode( # type: ignore[attr-defined]
86
+ [self.query_prefix + t for t in texts] if self.query_prefix
87
+ else list(texts)
88
+ )
89
+
90
+
91
+ def resolve_model_source(model: str) -> str:
92
+ """A hub id -> a local snapshot path, without the network if cached.
93
+
94
+ `StaticModel.from_pretrained` defaults `force_download=True`, so a
95
+ bare model id re-resolves through huggingface_hub on every
96
+ construction, dominating startup. One process per query is exactly
97
+ the CLI's shape, so that cost is paid on every invocation.
98
+
99
+ Falls back to a networked resolve when nothing is cached, and to
100
+ the id itself when that fails too -- the model loader's own error
101
+ is more useful than one invented here.
102
+ """
103
+ if Path(model).exists():
104
+ return model
105
+ from huggingface_hub import snapshot_download
106
+
107
+ for local_only in (True, False):
108
+ try:
109
+ return snapshot_download(model, local_files_only=local_only)
110
+ except Exception:
111
+ continue
112
+ return model
113
+
114
+
115
+ class StaticEncoder(_PrefixMixin):
116
+ """model2vec. Runs in process, no server."""
117
+
118
+ backend = "static"
119
+
120
+ def __init__(self, model: str, *, query_prefix: str | None = None,
121
+ doc_prefix: str | None = None) -> None:
122
+ from model2vec import StaticModel
123
+
124
+ self.name = model
125
+ self._prefixes(query_prefix, doc_prefix)
126
+ self._model = StaticModel.from_pretrained(resolve_model_source(model))
127
+ self.dims = int(self._model.dim)
128
+
129
+ def _encode(self, texts: Sequence[str]) -> list[Sequence[float]]:
130
+ return [list(map(float, v)) for v in self._model.encode(list(texts))]
131
+
132
+
133
+ class HttpEncoder(_PrefixMixin):
134
+ """OpenAI-compatible /v1/embeddings, e.g. a local llama-server.
135
+
136
+ urllib only; no client library.
137
+ """
138
+
139
+ backend = "http"
140
+
141
+ def __init__(self, endpoint: str, model: str, *, api_key: str = "",
142
+ probe: bool = True, query_prefix: str | None = None,
143
+ doc_prefix: str | None = None) -> None:
144
+ self.name = model
145
+ self._prefixes(query_prefix, doc_prefix)
146
+ self.endpoint = endpoint.rstrip("/")
147
+ self.variant = self.endpoint
148
+ self._api_key = api_key
149
+ self.dims = 0
150
+ if probe:
151
+ # Learned eagerly, with one request, because `dims` is part of
152
+ # the vector space. Left at 0 until the first real encode, a
153
+ # consumer comparing a fresh encoder's space against a stored
154
+ # one would see a change on every run.
155
+ self.dims = len(self._encode(["probe"])[0])
156
+
157
+ #: Requests carry at most this many texts. A large single POST can
158
+ #: exceed the server's batch limit or time out, so the transport
159
+ #: splits it without requiring callers to do so.
160
+ BATCH = 64
161
+
162
+ def _encode(self, texts: Sequence[str]) -> list[Sequence[float]]:
163
+ out: list[Sequence[float]] = []
164
+ for i in range(0, len(texts), self.BATCH):
165
+ out.extend(self._post(list(texts[i : i + self.BATCH])))
166
+ return out
167
+
168
+ def _post(self, texts: list[str]) -> list[Sequence[float]]:
169
+ headers = {"Content-Type": "application/json"}
170
+ if self._api_key:
171
+ headers["Authorization"] = f"Bearer {self._api_key}"
172
+ body = json.dumps({"input": texts, "model": self.name}).encode()
173
+ req = urllib.request.Request(
174
+ f"{self.endpoint}/v1/embeddings", data=body, headers=headers, method="POST"
175
+ )
176
+ with urllib.request.urlopen(req, timeout=300) as resp:
177
+ payload = json.loads(resp.read().decode())
178
+ # Order is not guaranteed by the spec. Every input index must be
179
+ # present exactly once; sorting an incomplete response would attach
180
+ # later vectors to the wrong texts.
181
+ items = payload["data"]
182
+ indices = [item.get("index") for item in items]
183
+ expected = list(range(len(texts)))
184
+ if (any(type(index) is not int for index in indices)
185
+ or sorted(indices) != expected):
186
+ raise ValueError(
187
+ f"embedding response indices {indices!r}; expected {expected!r}")
188
+ items = sorted(items, key=lambda d: d["index"])
189
+ vectors = [item["embedding"] for item in items]
190
+ widths = {len(vector) for vector in vectors}
191
+ if len(widths) > 1:
192
+ raise ValueError(
193
+ f"embedding response has mixed vector widths {sorted(widths)}")
194
+ if widths == {0}:
195
+ raise ValueError("embedding response contains empty vectors")
196
+ if any(type(value) not in (int, float) or not math.isfinite(value)
197
+ for vector in vectors for value in vector):
198
+ raise ValueError("embedding response contains a non-finite value")
199
+ if widths and self.dims and widths != {self.dims}:
200
+ raise ValueError(
201
+ f"embedding response width {next(iter(widths))}; expected {self.dims}")
202
+ if vectors and not self.dims:
203
+ self.dims = len(vectors[0])
204
+ return vectors
205
+
206
+
207
+ class FakeEncoder(_PrefixMixin):
208
+ """Deterministic vectors derived from the text hash. Tests only."""
209
+
210
+ backend = "fake"
211
+
212
+ def __init__(self, dims: int = 8, *, query_prefix: str | None = None,
213
+ doc_prefix: str | None = None) -> None:
214
+ if dims <= 0:
215
+ raise ValueError("dims must be positive")
216
+ self.name = "fake"
217
+ self.dims = dims
218
+ self._prefixes(query_prefix, doc_prefix)
219
+
220
+ def _encode(self, texts: Sequence[str]) -> list[Sequence[float]]:
221
+ out: list[Sequence[float]] = []
222
+ for text in texts:
223
+ encoded = text.encode()
224
+ digest = (hashlib.blake2b(encoded, digest_size=self.dims).digest()
225
+ if self.dims <= 64
226
+ else hashlib.shake_256(encoded).digest(self.dims))
227
+ out.append([(b / 255.0) - 0.5 for b in digest])
228
+ return out
229
+
230
+
231
+ #: Sequences per inference call. Narrow rather than wide: peak memory
232
+ #: scales with the batch, while throughput does not once the batch is
233
+ #: grouped by length.
234
+ _BATCH = 16
235
+
236
+
237
+ class OnnxEncoder(_PrefixMixin):
238
+ """A sentence transformer through onnxruntime. No torch, no server.
239
+
240
+ It costs an optional dependency and is slower than static embedding.
241
+
242
+ `dims` is learned by encoding once at construction rather than read
243
+ from config.json, because the pooling choice -- not the hidden size
244
+ alone -- decides the output width.
245
+ """
246
+
247
+ backend = "onnx"
248
+
249
+ def __init__(self, repo_id: str, *, local_only: bool = False,
250
+ pooling: str | None = None, filename: str = "onnx/model.onnx",
251
+ providers: str = "auto", query_prefix: str | None = None,
252
+ doc_prefix: str | None = None) -> None:
253
+ try:
254
+ import onnxruntime as ort
255
+ except ImportError as exc: # pragma: no cover - environment
256
+ raise ValueError(
257
+ "OnnxEncoder requires onnxruntime: pip install onnxruntime"
258
+ ) from exc
259
+ from huggingface_hub import hf_hub_download
260
+ from tokenizers import Tokenizer
261
+
262
+ self.name = repo_id
263
+ self._prefixes(query_prefix, doc_prefix)
264
+ self.variant = filename
265
+ self.pooling = pooling or default_pooling(repo_id)
266
+ if self.pooling not in ("cls", "mean", "last"):
267
+ raise ValueError(f"unknown pooling {self.pooling!r}")
268
+ get = lambda f: hf_hub_download( # noqa: E731
269
+ repo_id, f, local_files_only=local_only
270
+ )
271
+ self._tok = Tokenizer.from_file(get("tokenizer.json"))
272
+ self._tok.enable_truncation(512)
273
+ self._tok.enable_padding()
274
+ path = get(filename)
275
+ # Models over 2 GB store weights beside the graph; fetch the
276
+ # sidecar or the session loads a graph with no tensors.
277
+ if filename.endswith(".onnx"):
278
+ try:
279
+ get(filename + "_data")
280
+ except Exception: # most models have no sidecar
281
+ pass
282
+ self._sess = _session(ort, path, providers)
283
+ self.providers = self._sess.get_providers()
284
+ self._inputs = {i.name for i in self._sess.get_inputs()}
285
+ # Decoder-style exports (Qwen3-Embedding and friends) declare
286
+ # position_ids and a full empty KV cache as required inputs, two
287
+ # entries per layer. Encoder exports (bge, MiniLM) have none of
288
+ # this, so it stays empty for them.
289
+ self._kv = [
290
+ (i.name, i.shape, i.type) for i in self._sess.get_inputs()
291
+ if i.name.startswith("past_key_values.")
292
+ ]
293
+ self._wants_positions = "position_ids" in self._inputs
294
+ self.dims = len(self._encode(["probe"])[0])
295
+
296
+ def _decoder_inputs(self, batch: int, length: int) -> dict:
297
+ """position_ids and a zero-length KV cache, sized from the graph."""
298
+ import numpy as np
299
+
300
+ extra: dict = {}
301
+ if self._wants_positions:
302
+ extra["position_ids"] = np.tile(
303
+ np.arange(length, dtype=np.int64), (batch, 1)
304
+ )
305
+ for name, shape, dtype in self._kv:
306
+ # shape is [batch, heads, past_sequence_length, head_dim];
307
+ # past length is 0 because we never reuse a cache.
308
+ dims = [batch, int(shape[1]), 0, int(shape[3])]
309
+ extra[name] = np.zeros(
310
+ dims, dtype=np.float16 if "float16" in dtype else np.float32
311
+ )
312
+ return extra
313
+
314
+ def _encode(self, texts: Sequence[str]) -> list[Sequence[float]]:
315
+ import numpy as np
316
+
317
+ # A batch is padded to its longest member. Grouping by length first
318
+ # limits padding, then restoring the caller's order preserves the
319
+ # input-to-output mapping.
320
+ order = sorted(range(len(texts)), key=lambda i: len(texts[i]))
321
+ out: list[Sequence[float]] = [()] * len(texts)
322
+ for i in range(0, len(order), _BATCH):
323
+ window = order[i : i + _BATCH]
324
+ enc = self._tok.encode_batch([texts[j] for j in window])
325
+ feed = {
326
+ "input_ids": np.array([e.ids for e in enc], dtype=np.int64),
327
+ "attention_mask": np.array(
328
+ [e.attention_mask for e in enc], dtype=np.int64
329
+ ),
330
+ "token_type_ids": np.array(
331
+ [e.type_ids for e in enc], dtype=np.int64
332
+ ),
333
+ }
334
+ run = {k: v for k, v in feed.items() if k in self._inputs}
335
+ if self._kv or self._wants_positions:
336
+ run.update(
337
+ self._decoder_inputs(*feed["input_ids"].shape)
338
+ )
339
+ raw = self._sess.run(None, run)[0]
340
+ vec = (raw if raw.ndim == 2
341
+ else _pool(raw, feed["attention_mask"], self.pooling))
342
+ vec = vec / (np.linalg.norm(vec, axis=1, keepdims=True) + 1e-9)
343
+ for j, v in zip(window, vec):
344
+ out[j] = v.tolist()
345
+ return out
346
+
347
+
348
+ def _pool(hidden, mask, how: str):
349
+ """Token vectors -> one sentence vector.
350
+
351
+ Padding must be excluded for `mean` and `last`, which is why the
352
+ attention mask is threaded through rather than taking `hidden[:, -1]`:
353
+ with right padding the final row is a pad token for every sequence
354
+ shorter than the batch maximum, so the naive version would embed
355
+ padding and the error would be invisible.
356
+ """
357
+ import numpy as np
358
+
359
+ if how == "cls":
360
+ return hidden[:, 0]
361
+ m = mask.astype("float32")
362
+ if how == "mean":
363
+ return (hidden * m[:, :, None]).sum(1) / np.maximum(m.sum(1, keepdims=True), 1e-9)
364
+ idx = np.maximum(m.sum(1).astype("int64") - 1, 0)
365
+ return hidden[np.arange(hidden.shape[0]), idx]
366
+
367
+
368
+ #: Set once `register_execution_provider_library` has run for webgpu.
369
+ #: Registration is process-wide and a second call raises.
370
+ _WEBGPU_NAME: str | None = None
371
+
372
+
373
+ def _webgpu(ort) -> str:
374
+ """Register the webgpu plugin library and return its provider name."""
375
+ global _WEBGPU_NAME
376
+
377
+ if _WEBGPU_NAME is None:
378
+ try:
379
+ import onnxruntime_ep_webgpu as plugin
380
+ except ImportError as exc:
381
+ raise ValueError(
382
+ "providers='webgpu' requires the plugin: pip install"
383
+ " onnxruntime-ep-webgpu, or use providers='cpu'"
384
+ ) from exc
385
+ name = plugin.get_ep_name()
386
+ ort.register_execution_provider_library(name, plugin.get_library_path())
387
+ _WEBGPU_NAME = name
388
+ return _WEBGPU_NAME
389
+
390
+
391
+ def _session(ort, path: str, choice: str):
392
+ """An inference session bound to the chosen execution provider.
393
+
394
+ webgpu ships as a plugin provider, which `providers=` cannot reach:
395
+ a name onnxruntime does not recognise there is dropped and the
396
+ session runs on CPU reporting success. Plugins attach by device
397
+ instead, so they take a separate path rather than a longer list.
398
+ """
399
+ if choice != "webgpu":
400
+ return ort.InferenceSession(path, providers=_providers(ort, choice))
401
+ name = _webgpu(ort)
402
+ devices = [d for d in ort.get_ep_devices() if d.ep_name == name]
403
+ if not devices:
404
+ raise ValueError(
405
+ f"{name} registered but exposes no device; "
406
+ "use providers='cpu'"
407
+ )
408
+ opts = ort.SessionOptions()
409
+ opts.add_provider_for_devices(devices, {})
410
+ return ort.InferenceSession(path, opts)
411
+
412
+
413
+ def _providers(ort, choice: str) -> list[str]:
414
+ """Execution providers for onnxruntime, most preferred first.
415
+
416
+ 'auto' takes CoreML or CUDA when the build offers it. CPU is always
417
+ appended as the fallback, because CoreML silently declines operators
418
+ it cannot compile and a session with no usable provider would fail to
419
+ construct.
420
+ """
421
+ available = ort.get_available_providers()
422
+ if choice == "cpu":
423
+ return ["CPUExecutionProvider"]
424
+ if choice == "auto":
425
+ wanted = [p for p in ("CoreMLExecutionProvider", "CUDAExecutionProvider")
426
+ if p in available]
427
+ return wanted + ["CPUExecutionProvider"]
428
+ if choice not in available:
429
+ raise ValueError(
430
+ f"execution provider {choice!r} not available; "
431
+ f"onnxruntime offers {available}"
432
+ )
433
+ return [choice, "CPUExecutionProvider"]