fastembed-gpu 0.5.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.
- fastembed/__init__.py +20 -0
- fastembed/common/__init__.py +3 -0
- fastembed/common/model_management.py +301 -0
- fastembed/common/onnx_model.py +132 -0
- fastembed/common/preprocessor_utils.py +82 -0
- fastembed/common/types.py +16 -0
- fastembed/common/utils.py +55 -0
- fastembed/embedding.py +24 -0
- fastembed/image/__init__.py +3 -0
- fastembed/image/image_embedding.py +97 -0
- fastembed/image/image_embedding_base.py +44 -0
- fastembed/image/onnx_embedding.py +211 -0
- fastembed/image/onnx_image_model.py +131 -0
- fastembed/image/transform/functional.py +150 -0
- fastembed/image/transform/operators.py +268 -0
- fastembed/late_interaction/__init__.py +5 -0
- fastembed/late_interaction/colbert.py +256 -0
- fastembed/late_interaction/jina_colbert.py +62 -0
- fastembed/late_interaction/late_interaction_embedding_base.py +62 -0
- fastembed/late_interaction/late_interaction_text_embedding.py +114 -0
- fastembed/parallel_processor.py +252 -0
- fastembed/rerank/cross_encoder/__init__.py +3 -0
- fastembed/rerank/cross_encoder/onnx_text_cross_encoder.py +224 -0
- fastembed/rerank/cross_encoder/onnx_text_model.py +150 -0
- fastembed/rerank/cross_encoder/text_cross_encoder.py +120 -0
- fastembed/rerank/cross_encoder/text_cross_encoder_base.py +58 -0
- fastembed/sparse/__init__.py +4 -0
- fastembed/sparse/bm25.py +347 -0
- fastembed/sparse/bm42.py +340 -0
- fastembed/sparse/sparse_embedding_base.py +83 -0
- fastembed/sparse/sparse_text_embedding.py +121 -0
- fastembed/sparse/splade_pp.py +180 -0
- fastembed/sparse/utils/tokenizer.py +120 -0
- fastembed/text/__init__.py +3 -0
- fastembed/text/clip_embedding.py +54 -0
- fastembed/text/e5_onnx_embedding.py +72 -0
- fastembed/text/onnx_embedding.py +333 -0
- fastembed/text/onnx_text_model.py +145 -0
- fastembed/text/pooled_embedding.py +92 -0
- fastembed/text/pooled_normalized_embedding.py +125 -0
- fastembed/text/text_embedding.py +107 -0
- fastembed/text/text_embedding_base.py +62 -0
- fastembed_gpu-0.5.1.dist-info/LICENSE +201 -0
- fastembed_gpu-0.5.1.dist-info/METADATA +262 -0
- fastembed_gpu-0.5.1.dist-info/NOTICE +14 -0
- fastembed_gpu-0.5.1.dist-info/RECORD +47 -0
- fastembed_gpu-0.5.1.dist-info/WHEEL +4 -0
|
@@ -0,0 +1,114 @@
|
|
|
1
|
+
from typing import Any, Iterable, Optional, Sequence, Type, Union
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
|
|
5
|
+
from fastembed.common import OnnxProvider
|
|
6
|
+
from fastembed.late_interaction.colbert import Colbert
|
|
7
|
+
from fastembed.late_interaction.jina_colbert import JinaColbert
|
|
8
|
+
from fastembed.late_interaction.late_interaction_embedding_base import (
|
|
9
|
+
LateInteractionTextEmbeddingBase,
|
|
10
|
+
)
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class LateInteractionTextEmbedding(LateInteractionTextEmbeddingBase):
|
|
14
|
+
EMBEDDINGS_REGISTRY: list[Type[LateInteractionTextEmbeddingBase]] = [Colbert, JinaColbert]
|
|
15
|
+
|
|
16
|
+
@classmethod
|
|
17
|
+
def list_supported_models(cls) -> list[dict[str, Any]]:
|
|
18
|
+
"""
|
|
19
|
+
Lists the supported models.
|
|
20
|
+
|
|
21
|
+
Returns:
|
|
22
|
+
list[dict[str, Any]]: A list of dictionaries containing the model information.
|
|
23
|
+
|
|
24
|
+
Example:
|
|
25
|
+
```
|
|
26
|
+
[
|
|
27
|
+
{
|
|
28
|
+
"model": "colbert-ir/colbertv2.0",
|
|
29
|
+
"dim": 128,
|
|
30
|
+
"description": "Late interaction model",
|
|
31
|
+
"license": "mit",
|
|
32
|
+
"size_in_GB": 0.44,
|
|
33
|
+
"sources": {
|
|
34
|
+
"hf": "colbert-ir/colbertv2.0",
|
|
35
|
+
},
|
|
36
|
+
"model_file": "model.onnx",
|
|
37
|
+
},
|
|
38
|
+
]
|
|
39
|
+
```
|
|
40
|
+
"""
|
|
41
|
+
result = []
|
|
42
|
+
for embedding in cls.EMBEDDINGS_REGISTRY:
|
|
43
|
+
result.extend(embedding.list_supported_models())
|
|
44
|
+
return result
|
|
45
|
+
|
|
46
|
+
def __init__(
|
|
47
|
+
self,
|
|
48
|
+
model_name: str,
|
|
49
|
+
cache_dir: Optional[str] = None,
|
|
50
|
+
threads: Optional[int] = None,
|
|
51
|
+
providers: Optional[Sequence[OnnxProvider]] = None,
|
|
52
|
+
cuda: bool = False,
|
|
53
|
+
device_ids: Optional[list[int]] = None,
|
|
54
|
+
lazy_load: bool = False,
|
|
55
|
+
**kwargs,
|
|
56
|
+
):
|
|
57
|
+
super().__init__(model_name, cache_dir, threads, **kwargs)
|
|
58
|
+
for EMBEDDING_MODEL_TYPE in self.EMBEDDINGS_REGISTRY:
|
|
59
|
+
supported_models = EMBEDDING_MODEL_TYPE.list_supported_models()
|
|
60
|
+
if any(model_name.lower() == model["model"].lower() for model in supported_models):
|
|
61
|
+
self.model = EMBEDDING_MODEL_TYPE(
|
|
62
|
+
model_name,
|
|
63
|
+
cache_dir,
|
|
64
|
+
threads=threads,
|
|
65
|
+
providers=providers,
|
|
66
|
+
cuda=cuda,
|
|
67
|
+
device_ids=device_ids,
|
|
68
|
+
lazy_load=lazy_load,
|
|
69
|
+
**kwargs,
|
|
70
|
+
)
|
|
71
|
+
return
|
|
72
|
+
|
|
73
|
+
raise ValueError(
|
|
74
|
+
f"Model {model_name} is not supported in LateInteractionTextEmbedding."
|
|
75
|
+
"Please check the supported models using `LateInteractionTextEmbedding.list_supported_models()`"
|
|
76
|
+
)
|
|
77
|
+
|
|
78
|
+
def embed(
|
|
79
|
+
self,
|
|
80
|
+
documents: Union[str, Iterable[str]],
|
|
81
|
+
batch_size: int = 256,
|
|
82
|
+
parallel: Optional[int] = None,
|
|
83
|
+
**kwargs,
|
|
84
|
+
) -> Iterable[np.ndarray]:
|
|
85
|
+
"""
|
|
86
|
+
Encode a list of documents into list of embeddings.
|
|
87
|
+
We use mean pooling with attention so that the model can handle variable-length inputs.
|
|
88
|
+
|
|
89
|
+
Args:
|
|
90
|
+
documents: Iterator of documents or single document to embed
|
|
91
|
+
batch_size: Batch size for encoding -- higher values will use more memory, but be faster
|
|
92
|
+
parallel:
|
|
93
|
+
If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.
|
|
94
|
+
If 0, use all available cores.
|
|
95
|
+
If None, don't use data-parallel processing, use default onnxruntime threading instead.
|
|
96
|
+
|
|
97
|
+
Returns:
|
|
98
|
+
List of embeddings, one per document
|
|
99
|
+
"""
|
|
100
|
+
yield from self.model.embed(documents, batch_size, parallel, **kwargs)
|
|
101
|
+
|
|
102
|
+
def query_embed(self, query: Union[str, Iterable[str]], **kwargs) -> Iterable[np.ndarray]:
|
|
103
|
+
"""
|
|
104
|
+
Embeds queries
|
|
105
|
+
|
|
106
|
+
Args:
|
|
107
|
+
query (Union[str, Iterable[str]]): The query to embed, or an iterable e.g. list of queries.
|
|
108
|
+
|
|
109
|
+
Returns:
|
|
110
|
+
Iterable[np.ndarray]: The embeddings.
|
|
111
|
+
"""
|
|
112
|
+
|
|
113
|
+
# This is model-specific, so that different models can have specialized implementations
|
|
114
|
+
yield from self.model.query_embed(query, **kwargs)
|
|
@@ -0,0 +1,252 @@
|
|
|
1
|
+
import logging
|
|
2
|
+
import os
|
|
3
|
+
from collections import defaultdict
|
|
4
|
+
from copy import deepcopy
|
|
5
|
+
from enum import Enum
|
|
6
|
+
from multiprocessing import Queue, get_context
|
|
7
|
+
from multiprocessing.context import BaseContext
|
|
8
|
+
from multiprocessing.process import BaseProcess
|
|
9
|
+
from multiprocessing.sharedctypes import Synchronized as BaseValue
|
|
10
|
+
from queue import Empty
|
|
11
|
+
from typing import Any, Iterable, Optional, Type
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
# Single item should be processed in less than:
|
|
15
|
+
processing_timeout = 10 * 60 # seconds
|
|
16
|
+
|
|
17
|
+
max_internal_batch_size = 200
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
class QueueSignals(str, Enum):
|
|
21
|
+
stop = "stop"
|
|
22
|
+
confirm = "confirm"
|
|
23
|
+
error = "error"
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
class Worker:
|
|
27
|
+
@classmethod
|
|
28
|
+
def start(cls, *args: Any, **kwargs: Any) -> "Worker":
|
|
29
|
+
raise NotImplementedError()
|
|
30
|
+
|
|
31
|
+
def process(self, items: Iterable[tuple[int, Any]]) -> Iterable[tuple[int, Any]]:
|
|
32
|
+
raise NotImplementedError()
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
def _worker(
|
|
36
|
+
worker_class: Type[Worker],
|
|
37
|
+
input_queue: Queue,
|
|
38
|
+
output_queue: Queue,
|
|
39
|
+
num_active_workers: BaseValue,
|
|
40
|
+
worker_id: int,
|
|
41
|
+
kwargs: Optional[dict[str, Any]] = None,
|
|
42
|
+
) -> None:
|
|
43
|
+
"""
|
|
44
|
+
A worker that pulls data pints off the input queue, and places the execution result on the output queue.
|
|
45
|
+
When there are no data pints left on the input queue, it decrements
|
|
46
|
+
num_active_workers to signal completion.
|
|
47
|
+
"""
|
|
48
|
+
|
|
49
|
+
if kwargs is None:
|
|
50
|
+
kwargs = {}
|
|
51
|
+
|
|
52
|
+
logging.info(
|
|
53
|
+
f"Reader worker: {worker_id} PID: {os.getpid()} Device: {kwargs.get('device_id', 'CPU')}"
|
|
54
|
+
)
|
|
55
|
+
try:
|
|
56
|
+
worker = worker_class.start(**kwargs)
|
|
57
|
+
|
|
58
|
+
# Keep going until you get an item that's None.
|
|
59
|
+
def input_queue_iterable() -> Iterable[Any]:
|
|
60
|
+
while True:
|
|
61
|
+
item = input_queue.get()
|
|
62
|
+
if item == QueueSignals.stop:
|
|
63
|
+
break
|
|
64
|
+
yield item
|
|
65
|
+
|
|
66
|
+
for processed_item in worker.process(input_queue_iterable()):
|
|
67
|
+
output_queue.put(processed_item)
|
|
68
|
+
except Exception as e: # pylint: disable=broad-except
|
|
69
|
+
logging.exception(e)
|
|
70
|
+
output_queue.put(QueueSignals.error)
|
|
71
|
+
finally:
|
|
72
|
+
# It's important that we close and join the queue here before
|
|
73
|
+
# decrementing num_active_workers. Otherwise our parent may join us
|
|
74
|
+
# before the queue's feeder thread has passed all buffered items to
|
|
75
|
+
# the underlying pipe resulting in a deadlock.
|
|
76
|
+
#
|
|
77
|
+
# See:
|
|
78
|
+
# https://docs.python.org/3.6/library/multiprocessing.html?highlight=process#pipes-and-queues
|
|
79
|
+
# https://docs.python.org/3.6/library/multiprocessing.html?highlight=process#programming-guidelines
|
|
80
|
+
input_queue.close()
|
|
81
|
+
output_queue.close()
|
|
82
|
+
input_queue.join_thread()
|
|
83
|
+
output_queue.join_thread()
|
|
84
|
+
|
|
85
|
+
with num_active_workers.get_lock():
|
|
86
|
+
num_active_workers.value -= 1
|
|
87
|
+
|
|
88
|
+
logging.info(f"Reader worker {worker_id} finished")
|
|
89
|
+
|
|
90
|
+
|
|
91
|
+
class ParallelWorkerPool:
|
|
92
|
+
def __init__(
|
|
93
|
+
self,
|
|
94
|
+
num_workers: int,
|
|
95
|
+
worker: Type[Worker],
|
|
96
|
+
start_method: Optional[str] = None,
|
|
97
|
+
device_ids: Optional[list[int]] = None,
|
|
98
|
+
cuda: bool = False,
|
|
99
|
+
):
|
|
100
|
+
self.worker_class = worker
|
|
101
|
+
self.num_workers = num_workers
|
|
102
|
+
self.input_queue: Optional[Queue] = None
|
|
103
|
+
self.output_queue: Optional[Queue] = None
|
|
104
|
+
self.ctx: BaseContext = get_context(start_method)
|
|
105
|
+
self.processes: list[BaseProcess] = []
|
|
106
|
+
self.queue_size = self.num_workers * max_internal_batch_size
|
|
107
|
+
self.emergency_shutdown = False
|
|
108
|
+
self.device_ids = device_ids
|
|
109
|
+
self.cuda = cuda
|
|
110
|
+
self.num_active_workers: Optional[BaseValue] = None
|
|
111
|
+
|
|
112
|
+
def start(self, **kwargs: Any) -> None:
|
|
113
|
+
self.input_queue = self.ctx.Queue(self.queue_size)
|
|
114
|
+
self.output_queue = self.ctx.Queue(self.queue_size)
|
|
115
|
+
|
|
116
|
+
ctx_value = self.ctx.Value("i", self.num_workers)
|
|
117
|
+
assert isinstance(ctx_value, BaseValue)
|
|
118
|
+
self.num_active_workers = ctx_value
|
|
119
|
+
|
|
120
|
+
for worker_id in range(0, self.num_workers):
|
|
121
|
+
worker_kwargs = deepcopy(kwargs)
|
|
122
|
+
if self.device_ids:
|
|
123
|
+
device_id = self.device_ids[worker_id % len(self.device_ids)]
|
|
124
|
+
worker_kwargs["device_id"] = device_id
|
|
125
|
+
worker_kwargs["cuda"] = self.cuda
|
|
126
|
+
|
|
127
|
+
assert hasattr(self.ctx, "Process")
|
|
128
|
+
process = self.ctx.Process(
|
|
129
|
+
target=_worker,
|
|
130
|
+
args=(
|
|
131
|
+
self.worker_class,
|
|
132
|
+
self.input_queue,
|
|
133
|
+
self.output_queue,
|
|
134
|
+
self.num_active_workers,
|
|
135
|
+
worker_id,
|
|
136
|
+
worker_kwargs,
|
|
137
|
+
),
|
|
138
|
+
)
|
|
139
|
+
process.start()
|
|
140
|
+
self.processes.append(process)
|
|
141
|
+
|
|
142
|
+
def ordered_map(self, stream: Iterable[Any], *args: Any, **kwargs: Any) -> Iterable[Any]:
|
|
143
|
+
buffer = defaultdict(Any)
|
|
144
|
+
next_expected = 0
|
|
145
|
+
|
|
146
|
+
for idx, item in self.semi_ordered_map(stream, *args, **kwargs):
|
|
147
|
+
buffer[idx] = item
|
|
148
|
+
while next_expected in buffer:
|
|
149
|
+
yield buffer.pop(next_expected)
|
|
150
|
+
next_expected += 1
|
|
151
|
+
|
|
152
|
+
def semi_ordered_map(
|
|
153
|
+
self, stream: Iterable[Any], *args: Any, **kwargs: Any
|
|
154
|
+
) -> Iterable[tuple[int, Any]]:
|
|
155
|
+
try:
|
|
156
|
+
self.start(**kwargs)
|
|
157
|
+
|
|
158
|
+
assert self.input_queue is not None, "Input queue was not initialized"
|
|
159
|
+
assert self.output_queue is not None, "Output queue was not initialized"
|
|
160
|
+
|
|
161
|
+
pushed = 0
|
|
162
|
+
read = 0
|
|
163
|
+
for idx, item in enumerate(stream):
|
|
164
|
+
self.check_worker_health()
|
|
165
|
+
if pushed - read < self.queue_size:
|
|
166
|
+
try:
|
|
167
|
+
out_item = self.output_queue.get_nowait()
|
|
168
|
+
except Empty:
|
|
169
|
+
out_item = None
|
|
170
|
+
else:
|
|
171
|
+
try:
|
|
172
|
+
out_item = self.output_queue.get(timeout=processing_timeout)
|
|
173
|
+
except Empty as e:
|
|
174
|
+
self.join_or_terminate()
|
|
175
|
+
raise e
|
|
176
|
+
|
|
177
|
+
if out_item is not None:
|
|
178
|
+
if out_item == QueueSignals.error:
|
|
179
|
+
self.join_or_terminate()
|
|
180
|
+
raise RuntimeError("Thread unexpectedly terminated")
|
|
181
|
+
yield out_item
|
|
182
|
+
read += 1
|
|
183
|
+
|
|
184
|
+
self.input_queue.put((idx, item))
|
|
185
|
+
pushed += 1
|
|
186
|
+
|
|
187
|
+
for _ in range(self.num_workers):
|
|
188
|
+
self.input_queue.put(QueueSignals.stop)
|
|
189
|
+
|
|
190
|
+
while read < pushed:
|
|
191
|
+
self.check_worker_health()
|
|
192
|
+
out_item = self.output_queue.get(timeout=processing_timeout)
|
|
193
|
+
if out_item == QueueSignals.error:
|
|
194
|
+
self.join_or_terminate()
|
|
195
|
+
raise RuntimeError("Thread unexpectedly terminated")
|
|
196
|
+
yield out_item
|
|
197
|
+
read += 1
|
|
198
|
+
finally:
|
|
199
|
+
assert self.input_queue is not None, "Input queue is None"
|
|
200
|
+
assert self.output_queue is not None, "Output queue is None"
|
|
201
|
+
self.join()
|
|
202
|
+
self.input_queue.close()
|
|
203
|
+
self.output_queue.close()
|
|
204
|
+
if self.emergency_shutdown:
|
|
205
|
+
self.input_queue.cancel_join_thread()
|
|
206
|
+
self.output_queue.cancel_join_thread()
|
|
207
|
+
else:
|
|
208
|
+
self.input_queue.join_thread()
|
|
209
|
+
self.output_queue.join_thread()
|
|
210
|
+
|
|
211
|
+
def check_worker_health(self) -> None:
|
|
212
|
+
"""
|
|
213
|
+
Checks if any worker process has terminated unexpectedly
|
|
214
|
+
"""
|
|
215
|
+
for process in self.processes:
|
|
216
|
+
if not process.is_alive() and process.exitcode != 0:
|
|
217
|
+
self.emergency_shutdown = True
|
|
218
|
+
self.join_or_terminate()
|
|
219
|
+
raise RuntimeError(
|
|
220
|
+
f"Worker PID: {process.pid} terminated unexpectedly with code {process.exitcode}"
|
|
221
|
+
)
|
|
222
|
+
|
|
223
|
+
def join_or_terminate(self, timeout: Optional[int] = 1) -> None:
|
|
224
|
+
"""
|
|
225
|
+
Emergency shutdown
|
|
226
|
+
@param timeout:
|
|
227
|
+
@return:
|
|
228
|
+
"""
|
|
229
|
+
for process in self.processes:
|
|
230
|
+
process.join(timeout=timeout)
|
|
231
|
+
if process.is_alive():
|
|
232
|
+
process.terminate()
|
|
233
|
+
self.processes.clear()
|
|
234
|
+
|
|
235
|
+
def join(self) -> None:
|
|
236
|
+
for process in self.processes:
|
|
237
|
+
process.join()
|
|
238
|
+
self.processes.clear()
|
|
239
|
+
|
|
240
|
+
def __del__(self) -> None:
|
|
241
|
+
"""
|
|
242
|
+
Terminate processes if the user hasn't joined. This is necessary as
|
|
243
|
+
leaving stray processes running can corrupt shared state. In brief,
|
|
244
|
+
we've observed shared memory counters being reused (when the memory was
|
|
245
|
+
free from the perspective of the parent process) while the stray
|
|
246
|
+
workers still held a reference to them.
|
|
247
|
+
For a discussion of using destructors in Python in this manner, see
|
|
248
|
+
https://eli.thegreenplace.net/2009/06/12/safely-using-destructors-in-python/.
|
|
249
|
+
"""
|
|
250
|
+
for process in self.processes:
|
|
251
|
+
if process.is_alive():
|
|
252
|
+
process.terminate()
|
|
@@ -0,0 +1,224 @@
|
|
|
1
|
+
from typing import Any, Iterable, Optional, Sequence, Type
|
|
2
|
+
|
|
3
|
+
from loguru import logger
|
|
4
|
+
|
|
5
|
+
from fastembed.common import OnnxProvider
|
|
6
|
+
from fastembed.common.onnx_model import OnnxOutputContext
|
|
7
|
+
from fastembed.common.utils import define_cache_dir
|
|
8
|
+
from fastembed.rerank.cross_encoder.onnx_text_model import (
|
|
9
|
+
OnnxCrossEncoderModel,
|
|
10
|
+
TextRerankerWorker,
|
|
11
|
+
)
|
|
12
|
+
from fastembed.rerank.cross_encoder.text_cross_encoder_base import TextCrossEncoderBase
|
|
13
|
+
|
|
14
|
+
supported_onnx_models = [
|
|
15
|
+
{
|
|
16
|
+
"model": "Xenova/ms-marco-MiniLM-L-6-v2",
|
|
17
|
+
"size_in_GB": 0.08,
|
|
18
|
+
"sources": {
|
|
19
|
+
"hf": "Xenova/ms-marco-MiniLM-L-6-v2",
|
|
20
|
+
},
|
|
21
|
+
"model_file": "onnx/model.onnx",
|
|
22
|
+
"description": "MiniLM-L-6-v2 model optimized for re-ranking tasks.",
|
|
23
|
+
"license": "apache-2.0",
|
|
24
|
+
},
|
|
25
|
+
{
|
|
26
|
+
"model": "Xenova/ms-marco-MiniLM-L-12-v2",
|
|
27
|
+
"size_in_GB": 0.12,
|
|
28
|
+
"sources": {
|
|
29
|
+
"hf": "Xenova/ms-marco-MiniLM-L-12-v2",
|
|
30
|
+
},
|
|
31
|
+
"model_file": "onnx/model.onnx",
|
|
32
|
+
"description": "MiniLM-L-12-v2 model optimized for re-ranking tasks.",
|
|
33
|
+
"license": "apache-2.0",
|
|
34
|
+
},
|
|
35
|
+
{
|
|
36
|
+
"model": "BAAI/bge-reranker-base",
|
|
37
|
+
"size_in_GB": 1.04,
|
|
38
|
+
"sources": {
|
|
39
|
+
"hf": "BAAI/bge-reranker-base",
|
|
40
|
+
},
|
|
41
|
+
"model_file": "onnx/model.onnx",
|
|
42
|
+
"description": "BGE reranker base model for cross-encoder re-ranking.",
|
|
43
|
+
"license": "mit",
|
|
44
|
+
},
|
|
45
|
+
{
|
|
46
|
+
"model": "jinaai/jina-reranker-v1-tiny-en",
|
|
47
|
+
"size_in_GB": 0.13,
|
|
48
|
+
"sources": {
|
|
49
|
+
"hf": "jinaai/jina-reranker-v1-tiny-en",
|
|
50
|
+
},
|
|
51
|
+
"model_file": "onnx/model.onnx",
|
|
52
|
+
"description": "Designed for blazing-fast re-ranking with 8K context length and fewer parameters than jina-reranker-v1-turbo-en.",
|
|
53
|
+
"license": "apache-2.0",
|
|
54
|
+
},
|
|
55
|
+
{
|
|
56
|
+
"model": "jinaai/jina-reranker-v1-turbo-en",
|
|
57
|
+
"size_in_GB": 0.15,
|
|
58
|
+
"sources": {
|
|
59
|
+
"hf": "jinaai/jina-reranker-v1-turbo-en",
|
|
60
|
+
},
|
|
61
|
+
"model_file": "onnx/model.onnx",
|
|
62
|
+
"description": "Designed for blazing-fast re-ranking with 8K context length.",
|
|
63
|
+
"license": "apache-2.0",
|
|
64
|
+
},
|
|
65
|
+
{
|
|
66
|
+
"model": "jinaai/jina-reranker-v2-base-multilingual",
|
|
67
|
+
"size_in_GB": 1.11,
|
|
68
|
+
"sources": {
|
|
69
|
+
"hf": "jinaai/jina-reranker-v2-base-multilingual",
|
|
70
|
+
},
|
|
71
|
+
"model_file": "onnx/model.onnx",
|
|
72
|
+
"description": "A multi-lingual reranker model for cross-encoder re-ranking with 1K context length and sliding window",
|
|
73
|
+
"license": "cc-by-nc-4.0",
|
|
74
|
+
},
|
|
75
|
+
]
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
|
|
79
|
+
@classmethod
|
|
80
|
+
def list_supported_models(cls) -> list[dict[str, Any]]:
|
|
81
|
+
"""Lists the supported models.
|
|
82
|
+
|
|
83
|
+
Returns:
|
|
84
|
+
list[dict[str, Any]]: A list of dictionaries containing the model information.
|
|
85
|
+
"""
|
|
86
|
+
return supported_onnx_models
|
|
87
|
+
|
|
88
|
+
def __init__(
|
|
89
|
+
self,
|
|
90
|
+
model_name: str,
|
|
91
|
+
cache_dir: Optional[str] = None,
|
|
92
|
+
threads: Optional[int] = None,
|
|
93
|
+
providers: Optional[Sequence[OnnxProvider]] = None,
|
|
94
|
+
cuda: bool = False,
|
|
95
|
+
device_ids: Optional[list[int]] = None,
|
|
96
|
+
lazy_load: bool = False,
|
|
97
|
+
device_id: Optional[int] = None,
|
|
98
|
+
**kwargs: Any,
|
|
99
|
+
):
|
|
100
|
+
"""
|
|
101
|
+
Args:
|
|
102
|
+
model_name (str): The name of the model to use.
|
|
103
|
+
cache_dir (str, optional): The path to the cache directory.
|
|
104
|
+
Can be set using the `FASTEMBED_CACHE_PATH` env variable.
|
|
105
|
+
Defaults to `fastembed_cache` in the system's temp directory.
|
|
106
|
+
threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.
|
|
107
|
+
providers (Optional[Sequence[OnnxProvider]], optional): The list of onnxruntime providers to use.
|
|
108
|
+
Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
|
|
109
|
+
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
|
|
110
|
+
Defaults to False.
|
|
111
|
+
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
|
|
112
|
+
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
|
|
113
|
+
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
|
|
114
|
+
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
|
|
115
|
+
device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
|
|
116
|
+
|
|
117
|
+
Raises:
|
|
118
|
+
ValueError: If the model_name is not in the format <org>/<model> e.g. Xenova/ms-marco-MiniLM-L-6-v2.
|
|
119
|
+
"""
|
|
120
|
+
super().__init__(model_name, cache_dir, threads, **kwargs)
|
|
121
|
+
self.providers = providers
|
|
122
|
+
self.lazy_load = lazy_load
|
|
123
|
+
|
|
124
|
+
# List of device ids, that can be used for data parallel processing in workers
|
|
125
|
+
self.device_ids = device_ids
|
|
126
|
+
self.cuda = cuda
|
|
127
|
+
|
|
128
|
+
if self.device_ids is not None and len(self.device_ids) > 1:
|
|
129
|
+
logger.warning(
|
|
130
|
+
"Parallel execution is currently not supported for cross encoders, "
|
|
131
|
+
f"only the first device will be used for inference: {self.device_ids[0]}."
|
|
132
|
+
)
|
|
133
|
+
|
|
134
|
+
# This device_id will be used if we need to load model in current process
|
|
135
|
+
if device_id is not None:
|
|
136
|
+
self.device_id = device_id
|
|
137
|
+
elif self.device_ids is not None:
|
|
138
|
+
self.device_id = self.device_ids[0]
|
|
139
|
+
else:
|
|
140
|
+
self.device_id = None
|
|
141
|
+
|
|
142
|
+
self.model_description = self._get_model_description(model_name)
|
|
143
|
+
self.cache_dir = define_cache_dir(cache_dir)
|
|
144
|
+
self._model_dir = self.download_model(
|
|
145
|
+
self.model_description,
|
|
146
|
+
self.cache_dir,
|
|
147
|
+
local_files_only=self._local_files_only,
|
|
148
|
+
)
|
|
149
|
+
|
|
150
|
+
if not self.lazy_load:
|
|
151
|
+
self.load_onnx_model()
|
|
152
|
+
|
|
153
|
+
def load_onnx_model(self) -> None:
|
|
154
|
+
self._load_onnx_model(
|
|
155
|
+
model_dir=self._model_dir,
|
|
156
|
+
model_file=self.model_description["model_file"],
|
|
157
|
+
threads=self.threads,
|
|
158
|
+
providers=self.providers,
|
|
159
|
+
cuda=self.cuda,
|
|
160
|
+
device_id=self.device_id,
|
|
161
|
+
)
|
|
162
|
+
|
|
163
|
+
def rerank(
|
|
164
|
+
self,
|
|
165
|
+
query: str,
|
|
166
|
+
documents: Iterable[str],
|
|
167
|
+
batch_size: int = 64,
|
|
168
|
+
**kwargs: Any,
|
|
169
|
+
) -> Iterable[float]:
|
|
170
|
+
"""Reranks documents based on their relevance to a given query.
|
|
171
|
+
|
|
172
|
+
Args:
|
|
173
|
+
query (str): The query string to which document relevance is calculated.
|
|
174
|
+
documents (Iterable[str]): Iterable of documents to be reranked.
|
|
175
|
+
batch_size (int, optional): The number of documents processed in each batch. Higher batch sizes improve speed
|
|
176
|
+
but require more memory. Default is 64.
|
|
177
|
+
Returns:
|
|
178
|
+
Iterable[float]: An iterable of relevance scores for each document.
|
|
179
|
+
"""
|
|
180
|
+
|
|
181
|
+
yield from self._rerank_documents(
|
|
182
|
+
query=query, documents=documents, batch_size=batch_size, **kwargs
|
|
183
|
+
)
|
|
184
|
+
|
|
185
|
+
def rerank_pairs(
|
|
186
|
+
self,
|
|
187
|
+
pairs: Iterable[tuple[str, str]],
|
|
188
|
+
batch_size: int = 64,
|
|
189
|
+
parallel: Optional[int] = None,
|
|
190
|
+
**kwargs: Any,
|
|
191
|
+
) -> Iterable[float]:
|
|
192
|
+
yield from self._rerank_pairs(
|
|
193
|
+
model_name=self.model_name,
|
|
194
|
+
cache_dir=str(self.cache_dir),
|
|
195
|
+
pairs=pairs,
|
|
196
|
+
batch_size=batch_size,
|
|
197
|
+
parallel=parallel,
|
|
198
|
+
providers=self.providers,
|
|
199
|
+
cuda=self.cuda,
|
|
200
|
+
device_ids=self.device_ids,
|
|
201
|
+
**kwargs,
|
|
202
|
+
)
|
|
203
|
+
|
|
204
|
+
@classmethod
|
|
205
|
+
def _get_worker_class(cls) -> Type[TextRerankerWorker]:
|
|
206
|
+
return TextCrossEncoderWorker
|
|
207
|
+
|
|
208
|
+
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[float]:
|
|
209
|
+
return (float(elem) for elem in output.model_output)
|
|
210
|
+
|
|
211
|
+
|
|
212
|
+
class TextCrossEncoderWorker(TextRerankerWorker):
|
|
213
|
+
def init_embedding(
|
|
214
|
+
self,
|
|
215
|
+
model_name: str,
|
|
216
|
+
cache_dir: str,
|
|
217
|
+
**kwargs,
|
|
218
|
+
) -> OnnxTextCrossEncoder:
|
|
219
|
+
return OnnxTextCrossEncoder(
|
|
220
|
+
model_name=model_name,
|
|
221
|
+
cache_dir=cache_dir,
|
|
222
|
+
threads=1,
|
|
223
|
+
**kwargs,
|
|
224
|
+
)
|