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.
Files changed (47) hide show
  1. fastembed/__init__.py +20 -0
  2. fastembed/common/__init__.py +3 -0
  3. fastembed/common/model_management.py +301 -0
  4. fastembed/common/onnx_model.py +132 -0
  5. fastembed/common/preprocessor_utils.py +82 -0
  6. fastembed/common/types.py +16 -0
  7. fastembed/common/utils.py +55 -0
  8. fastembed/embedding.py +24 -0
  9. fastembed/image/__init__.py +3 -0
  10. fastembed/image/image_embedding.py +97 -0
  11. fastembed/image/image_embedding_base.py +44 -0
  12. fastembed/image/onnx_embedding.py +211 -0
  13. fastembed/image/onnx_image_model.py +131 -0
  14. fastembed/image/transform/functional.py +150 -0
  15. fastembed/image/transform/operators.py +268 -0
  16. fastembed/late_interaction/__init__.py +5 -0
  17. fastembed/late_interaction/colbert.py +256 -0
  18. fastembed/late_interaction/jina_colbert.py +62 -0
  19. fastembed/late_interaction/late_interaction_embedding_base.py +62 -0
  20. fastembed/late_interaction/late_interaction_text_embedding.py +114 -0
  21. fastembed/parallel_processor.py +252 -0
  22. fastembed/rerank/cross_encoder/__init__.py +3 -0
  23. fastembed/rerank/cross_encoder/onnx_text_cross_encoder.py +224 -0
  24. fastembed/rerank/cross_encoder/onnx_text_model.py +150 -0
  25. fastembed/rerank/cross_encoder/text_cross_encoder.py +120 -0
  26. fastembed/rerank/cross_encoder/text_cross_encoder_base.py +58 -0
  27. fastembed/sparse/__init__.py +4 -0
  28. fastembed/sparse/bm25.py +347 -0
  29. fastembed/sparse/bm42.py +340 -0
  30. fastembed/sparse/sparse_embedding_base.py +83 -0
  31. fastembed/sparse/sparse_text_embedding.py +121 -0
  32. fastembed/sparse/splade_pp.py +180 -0
  33. fastembed/sparse/utils/tokenizer.py +120 -0
  34. fastembed/text/__init__.py +3 -0
  35. fastembed/text/clip_embedding.py +54 -0
  36. fastembed/text/e5_onnx_embedding.py +72 -0
  37. fastembed/text/onnx_embedding.py +333 -0
  38. fastembed/text/onnx_text_model.py +145 -0
  39. fastembed/text/pooled_embedding.py +92 -0
  40. fastembed/text/pooled_normalized_embedding.py +125 -0
  41. fastembed/text/text_embedding.py +107 -0
  42. fastembed/text/text_embedding_base.py +62 -0
  43. fastembed_gpu-0.5.1.dist-info/LICENSE +201 -0
  44. fastembed_gpu-0.5.1.dist-info/METADATA +262 -0
  45. fastembed_gpu-0.5.1.dist-info/NOTICE +14 -0
  46. fastembed_gpu-0.5.1.dist-info/RECORD +47 -0
  47. fastembed_gpu-0.5.1.dist-info/WHEEL +4 -0
@@ -0,0 +1,150 @@
1
+ import os
2
+ from multiprocessing import get_all_start_methods
3
+ from pathlib import Path
4
+ from typing import Any, Iterable, Optional, Sequence, Type
5
+
6
+ import numpy as np
7
+ from tokenizers import Encoding
8
+
9
+ from fastembed.common.onnx_model import (
10
+ EmbeddingWorker,
11
+ OnnxModel,
12
+ OnnxOutputContext,
13
+ OnnxProvider,
14
+ )
15
+ from fastembed.common.preprocessor_utils import load_tokenizer
16
+ from fastembed.common.utils import iter_batch
17
+ from fastembed.parallel_processor import ParallelWorkerPool
18
+
19
+
20
+ class OnnxCrossEncoderModel(OnnxModel[float]):
21
+ ONNX_OUTPUT_NAMES: Optional[list[str]] = None
22
+
23
+ @classmethod
24
+ def _get_worker_class(cls) -> Type["TextRerankerWorker"]:
25
+ raise NotImplementedError("Subclasses must implement this method")
26
+
27
+ def _load_onnx_model(
28
+ self,
29
+ model_dir: Path,
30
+ model_file: str,
31
+ threads: Optional[int],
32
+ providers: Optional[Sequence[OnnxProvider]] = None,
33
+ cuda: bool = False,
34
+ device_id: Optional[int] = None,
35
+ ) -> None:
36
+ super()._load_onnx_model(
37
+ model_dir=model_dir,
38
+ model_file=model_file,
39
+ threads=threads,
40
+ providers=providers,
41
+ cuda=cuda,
42
+ device_id=device_id,
43
+ )
44
+ self.tokenizer, _ = load_tokenizer(model_dir=model_dir)
45
+
46
+ def tokenize(self, pairs: list[tuple[str, str]], **_: Any) -> list[Encoding]:
47
+ return self.tokenizer.encode_batch(pairs)
48
+
49
+ def _build_onnx_input(self, tokenized_input):
50
+ input_names = {node.name for node in self.model.get_inputs()}
51
+ inputs = {
52
+ "input_ids": np.array([enc.ids for enc in tokenized_input], dtype=np.int64),
53
+ }
54
+ if "token_type_ids" in input_names:
55
+ inputs["token_type_ids"] = np.array(
56
+ [enc.type_ids for enc in tokenized_input], dtype=np.int64
57
+ )
58
+ if "attention_mask" in input_names:
59
+ inputs["attention_mask"] = np.array(
60
+ [enc.attention_mask for enc in tokenized_input], dtype=np.int64
61
+ )
62
+ return inputs
63
+
64
+ def onnx_embed(self, query: str, documents: list[str], **kwargs: Any) -> OnnxOutputContext:
65
+ pairs = [(query, doc) for doc in documents]
66
+ return self.onnx_embed_pairs(pairs, **kwargs)
67
+
68
+ def onnx_embed_pairs(self, pairs: list[tuple[str, str]], **kwargs: Any) -> OnnxOutputContext:
69
+ tokenized_input = self.tokenize(pairs, **kwargs)
70
+ inputs = self._build_onnx_input(tokenized_input)
71
+ onnx_input = self._preprocess_onnx_input(inputs, **kwargs)
72
+ outputs = self.model.run(self.ONNX_OUTPUT_NAMES, onnx_input)
73
+ relevant_output = outputs[0]
74
+ scores = relevant_output[:, 0]
75
+ return OnnxOutputContext(model_output=scores)
76
+
77
+ def _rerank_documents(
78
+ self, query: str, documents: Iterable[str], batch_size: int, **kwargs: Any
79
+ ) -> Iterable[float]:
80
+ if not hasattr(self, "model") or self.model is None:
81
+ self.load_onnx_model()
82
+ for batch in iter_batch(documents, batch_size):
83
+ yield from self._post_process_onnx_output(self.onnx_embed(query, batch, **kwargs))
84
+
85
+ def _rerank_pairs(
86
+ self,
87
+ model_name: str,
88
+ cache_dir: str,
89
+ pairs: Iterable[tuple[str, str]],
90
+ batch_size: int,
91
+ parallel: Optional[int] = None,
92
+ providers: Optional[Sequence[OnnxProvider]] = None,
93
+ cuda: bool = False,
94
+ device_ids: Optional[list[int]] = None,
95
+ **kwargs: Any,
96
+ ) -> Iterable[float]:
97
+ is_small = False
98
+
99
+ if isinstance(pairs, tuple):
100
+ pairs = [pairs]
101
+ is_small = True
102
+
103
+ if isinstance(pairs, list):
104
+ if len(pairs) < batch_size:
105
+ is_small = True
106
+
107
+ if parallel is None or is_small:
108
+ if not hasattr(self, "model") or self.model is None:
109
+ self.load_onnx_model()
110
+ for batch in iter_batch(pairs, batch_size):
111
+ yield from self._post_process_onnx_output(self.onnx_embed_pairs(batch, **kwargs))
112
+ else:
113
+ if parallel == 0:
114
+ parallel = os.cpu_count()
115
+
116
+ start_method = "forkserver" if "forkserver" in get_all_start_methods() else "spawn"
117
+ params = {
118
+ "model_name": model_name,
119
+ "cache_dir": cache_dir,
120
+ "providers": providers,
121
+ **kwargs,
122
+ }
123
+
124
+ pool = ParallelWorkerPool(
125
+ num_workers=parallel or 1,
126
+ worker=self._get_worker_class(),
127
+ cuda=cuda,
128
+ device_ids=device_ids,
129
+ start_method=start_method,
130
+ )
131
+ for batch in pool.ordered_map(iter_batch(pairs, batch_size), **params):
132
+ yield from self._post_process_onnx_output(batch)
133
+
134
+ def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[float]:
135
+ raise NotImplementedError("Subclasses must implement this method")
136
+
137
+ def _preprocess_onnx_input(
138
+ self, onnx_input: dict[str, np.ndarray], **kwargs: Any
139
+ ) -> dict[str, np.ndarray]:
140
+ """
141
+ Preprocess the onnx input.
142
+ """
143
+ return onnx_input
144
+
145
+
146
+ class TextRerankerWorker(EmbeddingWorker):
147
+ def process(self, items: Iterable[tuple[int, Any]]) -> Iterable[tuple[int, Any]]:
148
+ for idx, batch in items:
149
+ onnx_output = self.model.onnx_embed_pairs(batch)
150
+ yield idx, onnx_output
@@ -0,0 +1,120 @@
1
+ from typing import Any, Iterable, Optional, Sequence, Type
2
+
3
+ from fastembed.common import OnnxProvider
4
+ from fastembed.rerank.cross_encoder.onnx_text_cross_encoder import OnnxTextCrossEncoder
5
+ from fastembed.rerank.cross_encoder.text_cross_encoder_base import TextCrossEncoderBase
6
+
7
+
8
+ class TextCrossEncoder(TextCrossEncoderBase):
9
+ CROSS_ENCODER_REGISTRY: list[Type[TextCrossEncoderBase]] = [
10
+ OnnxTextCrossEncoder,
11
+ ]
12
+
13
+ @classmethod
14
+ def list_supported_models(cls) -> list[dict[str, Any]]:
15
+ """Lists the supported models.
16
+
17
+ Returns:
18
+ list[dict[str, Any]]: A list of dictionaries containing the model information.
19
+
20
+ Example:
21
+ ```
22
+ [
23
+ {
24
+ "model": "Xenova/ms-marco-MiniLM-L-6-v2",
25
+ "size_in_GB": 0.08,
26
+ "sources": {
27
+ "hf": "Xenova/ms-marco-MiniLM-L-6-v2",
28
+ },
29
+ "model_file": "onnx/model.onnx",
30
+ "description": "MiniLM-L-6-v2 model optimized for re-ranking tasks.",
31
+ "license": "apache-2.0",
32
+ }
33
+ ]
34
+ ```
35
+ """
36
+ result = []
37
+ for encoder in cls.CROSS_ENCODER_REGISTRY:
38
+ result.extend(encoder.list_supported_models())
39
+ return result
40
+
41
+ def __init__(
42
+ self,
43
+ model_name: str,
44
+ cache_dir: Optional[str] = None,
45
+ threads: Optional[int] = None,
46
+ providers: Optional[Sequence[OnnxProvider]] = None,
47
+ cuda: bool = False,
48
+ device_ids: Optional[list[int]] = None,
49
+ lazy_load: bool = False,
50
+ **kwargs: Any,
51
+ ):
52
+ super().__init__(model_name, cache_dir, threads, **kwargs)
53
+
54
+ for CROSS_ENCODER_TYPE in self.CROSS_ENCODER_REGISTRY:
55
+ supported_models = CROSS_ENCODER_TYPE.list_supported_models()
56
+ if any(model_name.lower() == model["model"].lower() for model in supported_models):
57
+ self.model = CROSS_ENCODER_TYPE(
58
+ model_name=model_name,
59
+ cache_dir=cache_dir,
60
+ threads=threads,
61
+ providers=providers,
62
+ cuda=cuda,
63
+ device_ids=device_ids,
64
+ lazy_load=lazy_load,
65
+ **kwargs,
66
+ )
67
+ return
68
+
69
+ raise ValueError(
70
+ f"Model {model_name} is not supported in TextCrossEncoder."
71
+ "Please check the supported models using `TextCrossEncoder.list_supported_models()`"
72
+ )
73
+
74
+ def rerank(
75
+ self, query: str, documents: Iterable[str], batch_size: int = 64, **kwargs: Any
76
+ ) -> Iterable[float]:
77
+ """Rerank a list of documents based on a query.
78
+
79
+ Args:
80
+ query: Query to rerank the documents against
81
+ documents: Iterator of documents to rerank
82
+ batch_size: Batch size for reranking
83
+
84
+ Returns:
85
+ Iterable of scores for each document
86
+ """
87
+ yield from self.model.rerank(query, documents, batch_size=batch_size, **kwargs)
88
+
89
+ def rerank_pairs(
90
+ self,
91
+ pairs: Iterable[tuple[str, str]],
92
+ batch_size: int = 64,
93
+ parallel: Optional[int] = None,
94
+ **kwargs: Any,
95
+ ) -> Iterable[float]:
96
+ """
97
+ Rerank a list of query-document pairs.
98
+
99
+ Args:
100
+ pairs (Iterable[tuple[str, str]]): An iterable of tuples, where each tuple contains a query and a document
101
+ to be scored together.
102
+ batch_size (int, optional): The number of query-document pairs to process in a single batch. Defaults to 64.
103
+ parallel (Optional[int], optional): The number of parallel processes to use for reranking.
104
+ If None, parallelization is disabled. Defaults to None.
105
+ **kwargs (Any): Additional arguments to pass to the underlying reranking model.
106
+
107
+ Returns:
108
+ Iterable[float]: An iterable of scores corresponding to each query-document pair in the input.
109
+ Higher scores indicate a stronger match between the query and the document.
110
+
111
+ Example:
112
+ >>> encoder = TextCrossEncoder("Xenova/ms-marco-MiniLM-L-6-v2")
113
+ >>> pairs = [("What is AI?", "Artificial intelligence is ..."), ("What is ML?", "Machine learning is ...")]
114
+ >>> scores = list(encoder.rerank_pairs(pairs))
115
+ >>> print(list(map(lambda x: round(x, 2), scores)))
116
+ [-1.24, -10.6]
117
+ """
118
+ yield from self.model.rerank_pairs(
119
+ pairs, batch_size=batch_size, parallel=parallel, **kwargs
120
+ )
@@ -0,0 +1,58 @@
1
+ from typing import Any, Iterable, Optional
2
+
3
+ from fastembed.common.model_management import ModelManagement
4
+
5
+
6
+ class TextCrossEncoderBase(ModelManagement):
7
+ def __init__(
8
+ self,
9
+ model_name: str,
10
+ cache_dir: Optional[str] = None,
11
+ threads: Optional[int] = None,
12
+ **kwargs,
13
+ ):
14
+ self.model_name = model_name
15
+ self.cache_dir = cache_dir
16
+ self.threads = threads
17
+ self._local_files_only = kwargs.pop("local_files_only", False)
18
+
19
+ def rerank(
20
+ self,
21
+ query: str,
22
+ documents: Iterable[str],
23
+ batch_size: int = 64,
24
+ **kwargs,
25
+ ) -> Iterable[float]:
26
+ """Rerank a list of documents given a query.
27
+
28
+ Args:
29
+ query (str): The query to rerank the documents.
30
+ documents (Iterable[str]): The list of texts to rerank.
31
+ batch_size (int): The batch size to use for reranking.
32
+ **kwargs: Additional keyword argument to pass to the rerank method.
33
+
34
+ Yields:
35
+ Iterable[float]: The scores of the reranked the documents.
36
+ """
37
+ raise NotImplementedError("This method should be overridden by subclasses")
38
+
39
+ def rerank_pairs(
40
+ self,
41
+ pairs: Iterable[tuple[str, str]],
42
+ batch_size: int = 64,
43
+ parallel: Optional[int] = None,
44
+ **kwargs: Any,
45
+ ) -> Iterable[float]:
46
+ """Rerank query-document pairs.
47
+ Args:
48
+ pairs (Iterable[tuple[str, str]]): Query-document pairs to rerank
49
+ batch_size (int): The batch size to use for reranking.
50
+ parallel: parallel:
51
+ If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.
52
+ If 0, use all available cores.
53
+ If None, don't use data-parallel processing, use default onnxruntime threading instead.
54
+ **kwargs: Additional keyword argument to pass to the rerank method.
55
+ Yields:
56
+ Iterable[float]: Scores for each individual pair
57
+ """
58
+ raise NotImplementedError("This method should be overridden by subclasses")
@@ -0,0 +1,4 @@
1
+ from fastembed.sparse.sparse_embedding_base import SparseEmbedding
2
+ from fastembed.sparse.sparse_text_embedding import SparseTextEmbedding
3
+
4
+ __all__ = ["SparseEmbedding", "SparseTextEmbedding"]
@@ -0,0 +1,347 @@
1
+ import os
2
+ from collections import defaultdict
3
+ from multiprocessing import get_all_start_methods
4
+ from pathlib import Path
5
+ from typing import Any, Iterable, Optional, Type, Union
6
+
7
+ import mmh3
8
+ import numpy as np
9
+ from py_rust_stemmers import SnowballStemmer
10
+ from fastembed.common.utils import (
11
+ define_cache_dir,
12
+ iter_batch,
13
+ get_all_punctuation,
14
+ remove_non_alphanumeric,
15
+ )
16
+ from fastembed.parallel_processor import ParallelWorkerPool, Worker
17
+ from fastembed.sparse.sparse_embedding_base import (
18
+ SparseEmbedding,
19
+ SparseTextEmbeddingBase,
20
+ )
21
+ from fastembed.sparse.utils.tokenizer import SimpleTokenizer
22
+
23
+ supported_languages = [
24
+ "arabic",
25
+ "azerbaijani",
26
+ "basque",
27
+ "bengali",
28
+ "catalan",
29
+ "chinese",
30
+ "danish",
31
+ "dutch",
32
+ "english",
33
+ "finnish",
34
+ "french",
35
+ "german",
36
+ "greek",
37
+ "hebrew",
38
+ "hinglish",
39
+ "hungarian",
40
+ "indonesian",
41
+ "italian",
42
+ "kazakh",
43
+ "nepali",
44
+ "norwegian",
45
+ "portuguese",
46
+ "romanian",
47
+ "russian",
48
+ "slovene",
49
+ "spanish",
50
+ "swedish",
51
+ "tajik",
52
+ "turkish",
53
+ ]
54
+
55
+ supported_bm25_models = [
56
+ {
57
+ "model": "Qdrant/bm25",
58
+ "description": "BM25 as sparse embeddings meant to be used with Qdrant",
59
+ "license": "apache-2.0",
60
+ "size_in_GB": 0.01,
61
+ "sources": {
62
+ "hf": "Qdrant/bm25",
63
+ },
64
+ "model_file": "mock.file", # bm25 does not require a model, so we just use a mock
65
+ "additional_files": [f"{lang}.txt" for lang in supported_languages],
66
+ "requires_idf": True,
67
+ },
68
+ ]
69
+
70
+
71
+ class Bm25(SparseTextEmbeddingBase):
72
+ """Implements traditional BM25 in a form of sparse embeddings.
73
+ Uses a count of tokens in the document to evaluate the importance of the token.
74
+
75
+ WARNING: This model is expected to be used with `modifier="idf"` in the sparse vector index of Qdrant.
76
+
77
+ BM25 formula:
78
+
79
+ score(q, d) = SUM[ IDF(q_i) * (f(q_i, d) * (k + 1)) / (f(q_i, d) + k * (1 - b + b * (|d| / avg_len))) ],
80
+
81
+ where IDF is the inverse document frequency, computed on Qdrant's side
82
+ f(q_i, d) is the term frequency of the token q_i in the document d
83
+ k, b, avg_len are hyperparameters, described below.
84
+
85
+ Args:
86
+ model_name (str): The name of the model to use.
87
+ cache_dir (str, optional): The path to the cache directory.
88
+ Can be set using the `FASTEMBED_CACHE_PATH` env variable.
89
+ Defaults to `fastembed_cache` in the system's temp directory.
90
+ k (float, optional): The k parameter in the BM25 formula. Defines the saturation of the term frequency.
91
+ I.e. defines how fast the moment when additional terms stop to increase the score. Defaults to 1.2.
92
+ b (float, optional): The b parameter in the BM25 formula. Defines the importance of the document length.
93
+ Defaults to 0.75.
94
+ avg_len (float, optional): The average length of the documents in the corpus. Defaults to 256.0.
95
+ language (str): Specifies the language for the stemmer.
96
+ disable_stemmer (bool): Disable the stemmer.
97
+ Raises:
98
+ ValueError: If the model_name is not in the format <org>/<model> e.g. BAAI/bge-base-en.
99
+ """
100
+
101
+ def __init__(
102
+ self,
103
+ model_name: str,
104
+ cache_dir: Optional[str] = None,
105
+ k: float = 1.2,
106
+ b: float = 0.75,
107
+ avg_len: float = 256.0,
108
+ language: str = "english",
109
+ token_max_length: int = 40,
110
+ disable_stemmer: bool = False,
111
+ **kwargs,
112
+ ):
113
+ super().__init__(model_name, cache_dir, **kwargs)
114
+
115
+ if language not in supported_languages:
116
+ raise ValueError(f"{language} language is not supported")
117
+ else:
118
+ self.language = language
119
+
120
+ self.k = k
121
+ self.b = b
122
+ self.avg_len = avg_len
123
+
124
+ model_description = self._get_model_description(model_name)
125
+ self.cache_dir = define_cache_dir(cache_dir)
126
+
127
+ self._model_dir = self.download_model(
128
+ model_description, self.cache_dir, local_files_only=self._local_files_only
129
+ )
130
+
131
+ self.token_max_length = token_max_length
132
+ self.punctuation = set(get_all_punctuation())
133
+ self.disable_stemmer = disable_stemmer
134
+
135
+ if disable_stemmer:
136
+ self.stopwords = set()
137
+ self.stemmer = None
138
+ else:
139
+ self.stopwords = set(self._load_stopwords(self._model_dir, self.language))
140
+ self.stemmer = SnowballStemmer(language)
141
+
142
+ self.tokenizer = SimpleTokenizer
143
+
144
+ @classmethod
145
+ def list_supported_models(cls) -> list[dict[str, Any]]:
146
+ """Lists the supported models.
147
+
148
+ Returns:
149
+ list[dict[str, Any]]: A list of dictionaries containing the model information.
150
+ """
151
+ return supported_bm25_models
152
+
153
+ @classmethod
154
+ def _load_stopwords(cls, model_dir: Path, language: str) -> list[str]:
155
+ stopwords_path = model_dir / f"{language}.txt"
156
+ if not stopwords_path.exists():
157
+ return []
158
+
159
+ with open(stopwords_path, "r") as f:
160
+ return f.read().splitlines()
161
+
162
+ def _embed_documents(
163
+ self,
164
+ model_name: str,
165
+ cache_dir: str,
166
+ documents: Union[str, Iterable[str]],
167
+ batch_size: int = 256,
168
+ parallel: Optional[int] = None,
169
+ ) -> Iterable[SparseEmbedding]:
170
+ is_small = False
171
+
172
+ if isinstance(documents, str):
173
+ documents = [documents]
174
+ is_small = True
175
+
176
+ if isinstance(documents, list):
177
+ if len(documents) < batch_size:
178
+ is_small = True
179
+
180
+ if parallel is None or is_small:
181
+ for batch in iter_batch(documents, batch_size):
182
+ yield from self.raw_embed(batch)
183
+ else:
184
+ if parallel == 0:
185
+ parallel = os.cpu_count()
186
+
187
+ start_method = "forkserver" if "forkserver" in get_all_start_methods() else "spawn"
188
+ params = {
189
+ "model_name": model_name,
190
+ "cache_dir": cache_dir,
191
+ "k": self.k,
192
+ "b": self.b,
193
+ "avg_len": self.avg_len,
194
+ "language": self.language,
195
+ "token_max_length": self.token_max_length,
196
+ "disable_stemmer": self.disable_stemmer,
197
+ }
198
+ pool = ParallelWorkerPool(
199
+ num_workers=parallel or 1,
200
+ worker=self._get_worker_class(),
201
+ start_method=start_method,
202
+ )
203
+ for batch in pool.ordered_map(iter_batch(documents, batch_size), **params):
204
+ for record in batch:
205
+ yield record
206
+
207
+ def embed(
208
+ self,
209
+ documents: Union[str, Iterable[str]],
210
+ batch_size: int = 256,
211
+ parallel: Optional[int] = None,
212
+ **kwargs,
213
+ ) -> Iterable[SparseEmbedding]:
214
+ """
215
+ Encode a list of documents into list of embeddings.
216
+ We use mean pooling with attention so that the model can handle variable-length inputs.
217
+
218
+ Args:
219
+ documents: Iterator of documents or single document to embed
220
+ batch_size: Batch size for encoding -- higher values will use more memory, but be faster
221
+ parallel:
222
+ If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.
223
+ If 0, use all available cores.
224
+ If None, don't use data-parallel processing, use default onnxruntime threading instead.
225
+
226
+ Returns:
227
+ List of embeddings, one per document
228
+ """
229
+ yield from self._embed_documents(
230
+ model_name=self.model_name,
231
+ cache_dir=str(self.cache_dir),
232
+ documents=documents,
233
+ batch_size=batch_size,
234
+ parallel=parallel,
235
+ )
236
+
237
+ def _stem(self, tokens: list[str]) -> list[str]:
238
+ stemmed_tokens = []
239
+ for token in tokens:
240
+ lower_token = token.lower()
241
+
242
+ if token in self.punctuation:
243
+ continue
244
+
245
+ if lower_token in self.stopwords:
246
+ continue
247
+
248
+ if len(token) > self.token_max_length:
249
+ continue
250
+
251
+ stemmed_token = self.stemmer.stem_word(lower_token) if self.stemmer else lower_token
252
+
253
+ if stemmed_token:
254
+ stemmed_tokens.append(stemmed_token)
255
+ return stemmed_tokens
256
+
257
+ def raw_embed(
258
+ self,
259
+ documents: list[str],
260
+ ) -> list[SparseEmbedding]:
261
+ embeddings = []
262
+ for document in documents:
263
+ document = remove_non_alphanumeric(document)
264
+ tokens = self.tokenizer.tokenize(document)
265
+ stemmed_tokens = self._stem(tokens)
266
+ token_id2value = self._term_frequency(stemmed_tokens)
267
+ embeddings.append(SparseEmbedding.from_dict(token_id2value))
268
+ return embeddings
269
+
270
+ def _term_frequency(self, tokens: list[str]) -> dict[int, float]:
271
+ """Calculate the term frequency part of the BM25 formula.
272
+
273
+ (
274
+ f(q_i, d) * (k + 1)
275
+ ) / (
276
+ f(q_i, d) + k * (1 - b + b * (|d| / avg_len))
277
+ )
278
+
279
+ Args:
280
+ tokens (list[str]): The list of tokens in the document.
281
+
282
+ Returns:
283
+ dict[int, float]: The token_id to term frequency mapping.
284
+ """
285
+ tf_map = {}
286
+ counter = defaultdict(int)
287
+ for stemmed_token in tokens:
288
+ counter[stemmed_token] += 1
289
+
290
+ doc_len = len(tokens)
291
+ for stemmed_token in counter:
292
+ token_id = self.compute_token_id(stemmed_token)
293
+ num_occurrences = counter[stemmed_token]
294
+ tf_map[token_id] = num_occurrences * (self.k + 1)
295
+ tf_map[token_id] /= num_occurrences + self.k * (
296
+ 1 - self.b + self.b * doc_len / self.avg_len
297
+ )
298
+ return tf_map
299
+
300
+ @classmethod
301
+ def compute_token_id(cls, token: str) -> int:
302
+ return abs(mmh3.hash(token))
303
+
304
+ def query_embed(self, query: Union[str, Iterable[str]], **kwargs) -> Iterable[SparseEmbedding]:
305
+ """To emulate BM25 behaviour, we don't need to use weights in the query, and
306
+ it's enough to just hash the tokens and assign a weight of 1.0 to them.
307
+ """
308
+ if isinstance(query, str):
309
+ query = [query]
310
+
311
+ for text in query:
312
+ text = remove_non_alphanumeric(text)
313
+ tokens = self.tokenizer.tokenize(text)
314
+ stemmed_tokens = self._stem(tokens)
315
+ token_ids = np.array(
316
+ list(set(self.compute_token_id(token) for token in stemmed_tokens)),
317
+ dtype=np.int32,
318
+ )
319
+ values = np.ones_like(token_ids)
320
+ yield SparseEmbedding(indices=token_ids, values=values)
321
+
322
+ @classmethod
323
+ def _get_worker_class(cls) -> Type["Bm25Worker"]:
324
+ return Bm25Worker
325
+
326
+
327
+ class Bm25Worker(Worker):
328
+ def __init__(
329
+ self,
330
+ model_name: str,
331
+ cache_dir: str,
332
+ **kwargs,
333
+ ):
334
+ self.model = self.init_embedding(model_name, cache_dir, **kwargs)
335
+
336
+ @classmethod
337
+ def start(cls, model_name: str, cache_dir: str, **kwargs: Any) -> "Bm25Worker":
338
+ return cls(model_name=model_name, cache_dir=cache_dir, **kwargs)
339
+
340
+ def process(self, items: Iterable[tuple[int, Any]]) -> Iterable[tuple[int, Any]]:
341
+ for idx, batch in items:
342
+ onnx_output = self.model.raw_embed(batch)
343
+ yield idx, onnx_output
344
+
345
+ @staticmethod
346
+ def init_embedding(model_name: str, cache_dir: str, **kwargs) -> Bm25:
347
+ return Bm25(model_name=model_name, cache_dir=cache_dir, **kwargs)