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 +3 -0
- semsift/chunk/__init__.py +5 -0
- semsift/chunk/chunkers.py +206 -0
- semsift/embed/__init__.py +12 -0
- semsift/embed/backends.py +433 -0
- semsift/embed/policy.py +89 -0
- semsift/embed/space.py +43 -0
- semsift/evals/__init__.py +7 -0
- semsift/evals/harness.py +233 -0
- semsift/fuse/__init__.py +7 -0
- semsift/fuse/lists.py +177 -0
- semsift/rerank/__init__.py +5 -0
- semsift/rerank/rerankers.py +241 -0
- semsift/search/__init__.py +5 -0
- semsift/search/composer.py +205 -0
- semsift/store/__init__.py +8 -0
- semsift/store/codec.py +48 -0
- semsift/store/filters.py +231 -0
- semsift/store/index.py +82 -0
- semsift/store/store.py +621 -0
- semsift-0.0.1.dist-info/METADATA +86 -0
- semsift-0.0.1.dist-info/RECORD +25 -0
- semsift-0.0.1.dist-info/WHEEL +5 -0
- semsift-0.0.1.dist-info/licenses/LICENSE +21 -0
- semsift-0.0.1.dist-info/top_level.txt +1 -0
semsift/__init__.py
ADDED
|
@@ -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"]
|