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,180 @@
|
|
|
1
|
+
from typing import Any, Iterable, Optional, Sequence, Type, Union
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
from fastembed.common import OnnxProvider
|
|
5
|
+
from fastembed.common.onnx_model import OnnxOutputContext
|
|
6
|
+
from fastembed.common.utils import define_cache_dir
|
|
7
|
+
from fastembed.sparse.sparse_embedding_base import (
|
|
8
|
+
SparseEmbedding,
|
|
9
|
+
SparseTextEmbeddingBase,
|
|
10
|
+
)
|
|
11
|
+
from fastembed.text.onnx_text_model import OnnxTextModel, TextEmbeddingWorker
|
|
12
|
+
|
|
13
|
+
supported_splade_models = [
|
|
14
|
+
{
|
|
15
|
+
"model": "prithivida/Splade_PP_en_v1",
|
|
16
|
+
"vocab_size": 30522,
|
|
17
|
+
"description": "Independent Implementation of SPLADE++ Model for English.",
|
|
18
|
+
"license": "apache-2.0",
|
|
19
|
+
"size_in_GB": 0.532,
|
|
20
|
+
"sources": {
|
|
21
|
+
"hf": "Qdrant/SPLADE_PP_en_v1",
|
|
22
|
+
},
|
|
23
|
+
"model_file": "model.onnx",
|
|
24
|
+
},
|
|
25
|
+
{
|
|
26
|
+
"model": "prithvida/Splade_PP_en_v1",
|
|
27
|
+
"vocab_size": 30522,
|
|
28
|
+
"description": "Independent Implementation of SPLADE++ Model for English.",
|
|
29
|
+
"license": "apache-2.0",
|
|
30
|
+
"size_in_GB": 0.532,
|
|
31
|
+
"sources": {
|
|
32
|
+
"hf": "Qdrant/SPLADE_PP_en_v1",
|
|
33
|
+
},
|
|
34
|
+
"model_file": "model.onnx",
|
|
35
|
+
},
|
|
36
|
+
]
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
|
40
|
+
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[SparseEmbedding]:
|
|
41
|
+
if output.attention_mask is None:
|
|
42
|
+
raise ValueError("attention_mask must be provided for document post-processing")
|
|
43
|
+
|
|
44
|
+
relu_log = np.log(1 + np.maximum(output.model_output, 0))
|
|
45
|
+
|
|
46
|
+
weighted_log = relu_log * np.expand_dims(output.attention_mask, axis=-1)
|
|
47
|
+
|
|
48
|
+
scores = np.max(weighted_log, axis=1)
|
|
49
|
+
|
|
50
|
+
# Score matrix of shape (batch_size, vocab_size)
|
|
51
|
+
# Most of the values are 0, only a few are non-zero
|
|
52
|
+
for row_scores in scores:
|
|
53
|
+
indices = row_scores.nonzero()[0]
|
|
54
|
+
scores = row_scores[indices]
|
|
55
|
+
yield SparseEmbedding(values=scores, indices=indices)
|
|
56
|
+
|
|
57
|
+
@classmethod
|
|
58
|
+
def list_supported_models(cls) -> list[dict[str, Any]]:
|
|
59
|
+
"""Lists the supported models.
|
|
60
|
+
|
|
61
|
+
Returns:
|
|
62
|
+
list[dict[str, Any]]: A list of dictionaries containing the model information.
|
|
63
|
+
"""
|
|
64
|
+
return supported_splade_models
|
|
65
|
+
|
|
66
|
+
def __init__(
|
|
67
|
+
self,
|
|
68
|
+
model_name: str,
|
|
69
|
+
cache_dir: Optional[str] = None,
|
|
70
|
+
threads: Optional[int] = None,
|
|
71
|
+
providers: Optional[Sequence[OnnxProvider]] = None,
|
|
72
|
+
cuda: bool = False,
|
|
73
|
+
device_ids: Optional[list[int]] = None,
|
|
74
|
+
lazy_load: bool = False,
|
|
75
|
+
device_id: Optional[int] = None,
|
|
76
|
+
**kwargs,
|
|
77
|
+
):
|
|
78
|
+
"""
|
|
79
|
+
Args:
|
|
80
|
+
model_name (str): The name of the model to use.
|
|
81
|
+
cache_dir (str, optional): The path to the cache directory.
|
|
82
|
+
Can be set using the `FASTEMBED_CACHE_PATH` env variable.
|
|
83
|
+
Defaults to `fastembed_cache` in the system's temp directory.
|
|
84
|
+
threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.
|
|
85
|
+
providers (Optional[Sequence[OnnxProvider]], optional): The list of onnxruntime providers to use.
|
|
86
|
+
Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
|
|
87
|
+
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
|
|
88
|
+
Defaults to False.
|
|
89
|
+
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
|
|
90
|
+
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
|
|
91
|
+
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
|
|
92
|
+
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
|
|
93
|
+
device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
|
|
94
|
+
|
|
95
|
+
Raises:
|
|
96
|
+
ValueError: If the model_name is not in the format <org>/<model> e.g. BAAI/bge-base-en.
|
|
97
|
+
"""
|
|
98
|
+
super().__init__(model_name, cache_dir, threads, **kwargs)
|
|
99
|
+
self.providers = providers
|
|
100
|
+
self.lazy_load = lazy_load
|
|
101
|
+
|
|
102
|
+
# List of device ids, that can be used for data parallel processing in workers
|
|
103
|
+
self.device_ids = device_ids
|
|
104
|
+
self.cuda = cuda
|
|
105
|
+
|
|
106
|
+
# This device_id will be used if we need to load model in current process
|
|
107
|
+
if device_id is not None:
|
|
108
|
+
self.device_id = device_id
|
|
109
|
+
elif self.device_ids is not None:
|
|
110
|
+
self.device_id = self.device_ids[0]
|
|
111
|
+
else:
|
|
112
|
+
self.device_id = None
|
|
113
|
+
|
|
114
|
+
self.model_description = self._get_model_description(model_name)
|
|
115
|
+
self.cache_dir = define_cache_dir(cache_dir)
|
|
116
|
+
|
|
117
|
+
self._model_dir = self.download_model(
|
|
118
|
+
self.model_description, self.cache_dir, local_files_only=self._local_files_only
|
|
119
|
+
)
|
|
120
|
+
|
|
121
|
+
if not self.lazy_load:
|
|
122
|
+
self.load_onnx_model()
|
|
123
|
+
|
|
124
|
+
def load_onnx_model(self) -> None:
|
|
125
|
+
self._load_onnx_model(
|
|
126
|
+
model_dir=self._model_dir,
|
|
127
|
+
model_file=self.model_description["model_file"],
|
|
128
|
+
threads=self.threads,
|
|
129
|
+
providers=self.providers,
|
|
130
|
+
cuda=self.cuda,
|
|
131
|
+
device_id=self.device_id,
|
|
132
|
+
)
|
|
133
|
+
|
|
134
|
+
def embed(
|
|
135
|
+
self,
|
|
136
|
+
documents: Union[str, Iterable[str]],
|
|
137
|
+
batch_size: int = 256,
|
|
138
|
+
parallel: Optional[int] = None,
|
|
139
|
+
**kwargs,
|
|
140
|
+
) -> Iterable[SparseEmbedding]:
|
|
141
|
+
"""
|
|
142
|
+
Encode a list of documents into list of embeddings.
|
|
143
|
+
We use mean pooling with attention so that the model can handle variable-length inputs.
|
|
144
|
+
|
|
145
|
+
Args:
|
|
146
|
+
documents: Iterator of documents or single document to embed
|
|
147
|
+
batch_size: Batch size for encoding -- higher values will use more memory, but be faster
|
|
148
|
+
parallel:
|
|
149
|
+
If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.
|
|
150
|
+
If 0, use all available cores.
|
|
151
|
+
If None, don't use data-parallel processing, use default onnxruntime threading instead.
|
|
152
|
+
|
|
153
|
+
Returns:
|
|
154
|
+
List of embeddings, one per document
|
|
155
|
+
"""
|
|
156
|
+
yield from self._embed_documents(
|
|
157
|
+
model_name=self.model_name,
|
|
158
|
+
cache_dir=str(self.cache_dir),
|
|
159
|
+
documents=documents,
|
|
160
|
+
batch_size=batch_size,
|
|
161
|
+
parallel=parallel,
|
|
162
|
+
providers=self.providers,
|
|
163
|
+
cuda=self.cuda,
|
|
164
|
+
device_ids=self.device_ids,
|
|
165
|
+
**kwargs,
|
|
166
|
+
)
|
|
167
|
+
|
|
168
|
+
@classmethod
|
|
169
|
+
def _get_worker_class(cls) -> Type[TextEmbeddingWorker]:
|
|
170
|
+
return SpladePPEmbeddingWorker
|
|
171
|
+
|
|
172
|
+
|
|
173
|
+
class SpladePPEmbeddingWorker(TextEmbeddingWorker):
|
|
174
|
+
def init_embedding(self, model_name: str, cache_dir: str, **kwargs) -> SpladePP:
|
|
175
|
+
return SpladePP(
|
|
176
|
+
model_name=model_name,
|
|
177
|
+
cache_dir=cache_dir,
|
|
178
|
+
threads=1,
|
|
179
|
+
**kwargs,
|
|
180
|
+
)
|
|
@@ -0,0 +1,120 @@
|
|
|
1
|
+
# This code is a modified copy of the `NLTKWordTokenizer` class from `NLTK` library.
|
|
2
|
+
|
|
3
|
+
import re
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class SimpleTokenizer:
|
|
7
|
+
@staticmethod
|
|
8
|
+
def tokenize(text: str) -> list[str]:
|
|
9
|
+
text = re.sub(r"[^\w]", " ", text.lower())
|
|
10
|
+
text = re.sub(r"\s+", " ", text)
|
|
11
|
+
|
|
12
|
+
return text.strip().split()
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
class WordTokenizer:
|
|
16
|
+
"""The tokenizer is "destructive" such that the regexes applied will munge the
|
|
17
|
+
input string to a state beyond re-construction.
|
|
18
|
+
"""
|
|
19
|
+
|
|
20
|
+
# Starting quotes.
|
|
21
|
+
STARTING_QUOTES = [
|
|
22
|
+
(re.compile("([«“‘„]|[`]+)", re.U), r" \1 "),
|
|
23
|
+
(re.compile(r"^\""), r"``"),
|
|
24
|
+
(re.compile(r"(``)"), r" \1 "),
|
|
25
|
+
(re.compile(r"([ \(\[{<])(\"|\'{2})"), r"\1 `` "),
|
|
26
|
+
(re.compile(r"(?i)(\')(?!re|ve|ll|m|t|s|d|n)(\w)\b", re.U), r"\1 \2"),
|
|
27
|
+
]
|
|
28
|
+
|
|
29
|
+
# Ending quotes.
|
|
30
|
+
ENDING_QUOTES = [
|
|
31
|
+
(re.compile("([»”’])", re.U), r" \1 "),
|
|
32
|
+
(re.compile(r"''"), " '' "),
|
|
33
|
+
(re.compile(r'"'), " '' "),
|
|
34
|
+
(re.compile(r"([^' ])('[sS]|'[mM]|'[dD]|') "), r"\1 \2 "),
|
|
35
|
+
(re.compile(r"([^' ])('ll|'LL|'re|'RE|'ve|'VE|n't|N'T) "), r"\1 \2 "),
|
|
36
|
+
]
|
|
37
|
+
|
|
38
|
+
# Punctuation.
|
|
39
|
+
PUNCTUATION = [
|
|
40
|
+
(re.compile(r'([^\.])(\.)([\]\)}>"\'' "»”’ " r"]*)\s*$", re.U), r"\1 \2 \3 "),
|
|
41
|
+
(re.compile(r"([:,])([^\d])"), r" \1 \2"),
|
|
42
|
+
(re.compile(r"([:,])$"), r" \1 "),
|
|
43
|
+
(
|
|
44
|
+
re.compile(r"\.{2,}", re.U),
|
|
45
|
+
r" \g<0> ",
|
|
46
|
+
),
|
|
47
|
+
(re.compile(r"[;@#$%&]"), r" \g<0> "),
|
|
48
|
+
(
|
|
49
|
+
re.compile(r'([^\.])(\.)([\]\)}>"\']*)\s*$'),
|
|
50
|
+
r"\1 \2\3 ",
|
|
51
|
+
), # Handles the final period.
|
|
52
|
+
(re.compile(r"[?!]"), r" \g<0> "),
|
|
53
|
+
(re.compile(r"([^'])' "), r"\1 ' "),
|
|
54
|
+
(
|
|
55
|
+
re.compile(r"[*]", re.U),
|
|
56
|
+
r" \g<0> ",
|
|
57
|
+
),
|
|
58
|
+
]
|
|
59
|
+
|
|
60
|
+
# Pads parentheses
|
|
61
|
+
PARENS_BRACKETS = (re.compile(r"[\]\[\(\)\{\}\<\>]"), r" \g<0> ")
|
|
62
|
+
DOUBLE_DASHES = (re.compile(r"--"), r" -- ")
|
|
63
|
+
|
|
64
|
+
# List of contractions adapted from Robert MacIntyre's tokenizer.
|
|
65
|
+
CONTRACTIONS2 = [
|
|
66
|
+
re.compile(pattern)
|
|
67
|
+
for pattern in (
|
|
68
|
+
r"(?i)\b(can)(?#X)(not)\b",
|
|
69
|
+
r"(?i)\b(d)(?#X)('ye)\b",
|
|
70
|
+
r"(?i)\b(gim)(?#X)(me)\b",
|
|
71
|
+
r"(?i)\b(gon)(?#X)(na)\b",
|
|
72
|
+
r"(?i)\b(got)(?#X)(ta)\b",
|
|
73
|
+
r"(?i)\b(lem)(?#X)(me)\b",
|
|
74
|
+
r"(?i)\b(more)(?#X)('n)\b",
|
|
75
|
+
r"(?i)\b(wan)(?#X)(na)(?=\s)",
|
|
76
|
+
)
|
|
77
|
+
]
|
|
78
|
+
CONTRACTIONS3 = [
|
|
79
|
+
re.compile(pattern) for pattern in (r"(?i) ('t)(?#X)(is)\b", r"(?i) ('t)(?#X)(was)\b")
|
|
80
|
+
]
|
|
81
|
+
|
|
82
|
+
@classmethod
|
|
83
|
+
def tokenize(cls, text: str) -> list[str]:
|
|
84
|
+
"""Return a tokenized copy of `text`.
|
|
85
|
+
|
|
86
|
+
>>> s = '''Good muffins cost $3.88 (roughly 3,36 euros)\nin New York.'''
|
|
87
|
+
>>> WordTokenizer().tokenize(s)
|
|
88
|
+
['Good', 'muffins', 'cost', '$', '3.88', '(', 'roughly', '3,36', 'euros', ')', 'in', 'New', 'York', '.']
|
|
89
|
+
|
|
90
|
+
Args:
|
|
91
|
+
text: The text to be tokenized.
|
|
92
|
+
|
|
93
|
+
Returns:
|
|
94
|
+
A list of tokens.
|
|
95
|
+
"""
|
|
96
|
+
for regexp, substitution in cls.STARTING_QUOTES:
|
|
97
|
+
text = regexp.sub(substitution, text)
|
|
98
|
+
|
|
99
|
+
for regexp, substitution in cls.PUNCTUATION:
|
|
100
|
+
text = regexp.sub(substitution, text)
|
|
101
|
+
|
|
102
|
+
# Handles parentheses.
|
|
103
|
+
regexp, substitution = cls.PARENS_BRACKETS
|
|
104
|
+
text = regexp.sub(substitution, text)
|
|
105
|
+
|
|
106
|
+
# Handles double dash.
|
|
107
|
+
regexp, substitution = cls.DOUBLE_DASHES
|
|
108
|
+
text = regexp.sub(substitution, text)
|
|
109
|
+
|
|
110
|
+
# add extra space to make things easier
|
|
111
|
+
text = " " + text + " "
|
|
112
|
+
|
|
113
|
+
for regexp, substitution in cls.ENDING_QUOTES:
|
|
114
|
+
text = regexp.sub(substitution, text)
|
|
115
|
+
|
|
116
|
+
for regexp in cls.CONTRACTIONS2:
|
|
117
|
+
text = regexp.sub(r" \1 \2 ", text)
|
|
118
|
+
for regexp in cls.CONTRACTIONS3:
|
|
119
|
+
text = regexp.sub(r" \1 \2 ", text)
|
|
120
|
+
return text.split()
|
|
@@ -0,0 +1,54 @@
|
|
|
1
|
+
from typing import Any, Iterable, Type
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
|
|
5
|
+
from fastembed.common.onnx_model import OnnxOutputContext
|
|
6
|
+
from fastembed.text.onnx_embedding import OnnxTextEmbedding, OnnxTextEmbeddingWorker
|
|
7
|
+
from fastembed.text.onnx_text_model import TextEmbeddingWorker
|
|
8
|
+
|
|
9
|
+
supported_clip_models = [
|
|
10
|
+
{
|
|
11
|
+
"model": "Qdrant/clip-ViT-B-32-text",
|
|
12
|
+
"dim": 512,
|
|
13
|
+
"description": "Text embeddings, Multimodal (text&image), English, 77 input tokens truncation, Prefixes for queries/documents: not necessary, 2021 year",
|
|
14
|
+
"license": "mit",
|
|
15
|
+
"size_in_GB": 0.25,
|
|
16
|
+
"sources": {
|
|
17
|
+
"hf": "Qdrant/clip-ViT-B-32-text",
|
|
18
|
+
},
|
|
19
|
+
"model_file": "model.onnx",
|
|
20
|
+
},
|
|
21
|
+
]
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
class CLIPOnnxEmbedding(OnnxTextEmbedding):
|
|
25
|
+
@classmethod
|
|
26
|
+
def _get_worker_class(cls) -> Type[TextEmbeddingWorker]:
|
|
27
|
+
return CLIPEmbeddingWorker
|
|
28
|
+
|
|
29
|
+
@classmethod
|
|
30
|
+
def list_supported_models(cls) -> list[dict[str, Any]]:
|
|
31
|
+
"""Lists the supported models.
|
|
32
|
+
|
|
33
|
+
Returns:
|
|
34
|
+
list[dict[str, Any]]: A list of dictionaries containing the model information.
|
|
35
|
+
"""
|
|
36
|
+
return supported_clip_models
|
|
37
|
+
|
|
38
|
+
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[np.ndarray]:
|
|
39
|
+
return output.model_output
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
class CLIPEmbeddingWorker(OnnxTextEmbeddingWorker):
|
|
43
|
+
def init_embedding(
|
|
44
|
+
self,
|
|
45
|
+
model_name: str,
|
|
46
|
+
cache_dir: str,
|
|
47
|
+
**kwargs,
|
|
48
|
+
) -> OnnxTextEmbedding:
|
|
49
|
+
return CLIPOnnxEmbedding(
|
|
50
|
+
model_name=model_name,
|
|
51
|
+
cache_dir=cache_dir,
|
|
52
|
+
threads=1,
|
|
53
|
+
**kwargs,
|
|
54
|
+
)
|
|
@@ -0,0 +1,72 @@
|
|
|
1
|
+
from typing import Any, Type
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
|
|
5
|
+
from fastembed.text.onnx_embedding import OnnxTextEmbedding, OnnxTextEmbeddingWorker
|
|
6
|
+
from fastembed.text.onnx_text_model import TextEmbeddingWorker
|
|
7
|
+
|
|
8
|
+
supported_multilingual_e5_models = [
|
|
9
|
+
{
|
|
10
|
+
"model": "intfloat/multilingual-e5-large",
|
|
11
|
+
"dim": 1024,
|
|
12
|
+
"description": "Text embeddings, Unimodal (text), Multilingual (~100 languages), 512 input tokens truncation, Prefixes for queries/documents: necessary, 2024 year.",
|
|
13
|
+
"license": "mit",
|
|
14
|
+
"size_in_GB": 2.24,
|
|
15
|
+
"sources": {
|
|
16
|
+
"url": "https://storage.googleapis.com/qdrant-fastembed/fast-multilingual-e5-large.tar.gz",
|
|
17
|
+
"hf": "qdrant/multilingual-e5-large-onnx",
|
|
18
|
+
},
|
|
19
|
+
"model_file": "model.onnx",
|
|
20
|
+
"additional_files": ["model.onnx_data"],
|
|
21
|
+
},
|
|
22
|
+
{
|
|
23
|
+
"model": "sentence-transformers/paraphrase-multilingual-mpnet-base-v2",
|
|
24
|
+
"dim": 768,
|
|
25
|
+
"description": "Text embeddings, Unimodal (text), Multilingual (~50 languages), 384 input tokens truncation, Prefixes for queries/documents: not necessary, 2021 year.",
|
|
26
|
+
"license": "apache-2.0",
|
|
27
|
+
"size_in_GB": 1.00,
|
|
28
|
+
"sources": {
|
|
29
|
+
"hf": "xenova/paraphrase-multilingual-mpnet-base-v2",
|
|
30
|
+
},
|
|
31
|
+
"model_file": "onnx/model.onnx",
|
|
32
|
+
},
|
|
33
|
+
]
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
class E5OnnxEmbedding(OnnxTextEmbedding):
|
|
37
|
+
@classmethod
|
|
38
|
+
def _get_worker_class(cls) -> Type["TextEmbeddingWorker"]:
|
|
39
|
+
return E5OnnxEmbeddingWorker
|
|
40
|
+
|
|
41
|
+
@classmethod
|
|
42
|
+
def list_supported_models(cls) -> list[dict[str, Any]]:
|
|
43
|
+
"""Lists the supported models.
|
|
44
|
+
|
|
45
|
+
Returns:
|
|
46
|
+
list[dict[str, Any]]: A list of dictionaries containing the model information.
|
|
47
|
+
"""
|
|
48
|
+
return supported_multilingual_e5_models
|
|
49
|
+
|
|
50
|
+
def _preprocess_onnx_input(
|
|
51
|
+
self, onnx_input: dict[str, np.ndarray], **kwargs
|
|
52
|
+
) -> dict[str, np.ndarray]:
|
|
53
|
+
"""
|
|
54
|
+
Preprocess the onnx input.
|
|
55
|
+
"""
|
|
56
|
+
onnx_input.pop("token_type_ids", None)
|
|
57
|
+
return onnx_input
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
class E5OnnxEmbeddingWorker(OnnxTextEmbeddingWorker):
|
|
61
|
+
def init_embedding(
|
|
62
|
+
self,
|
|
63
|
+
model_name: str,
|
|
64
|
+
cache_dir: str,
|
|
65
|
+
**kwargs,
|
|
66
|
+
) -> E5OnnxEmbedding:
|
|
67
|
+
return E5OnnxEmbedding(
|
|
68
|
+
model_name=model_name,
|
|
69
|
+
cache_dir=cache_dir,
|
|
70
|
+
threads=1,
|
|
71
|
+
**kwargs,
|
|
72
|
+
)
|