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,268 @@
1
+ from typing import Any, Union, Optional
2
+
3
+ import numpy as np
4
+ from PIL import Image
5
+
6
+ from fastembed.image.transform.functional import (
7
+ center_crop,
8
+ convert_to_rgb,
9
+ normalize,
10
+ pil2ndarray,
11
+ rescale,
12
+ resize,
13
+ pad2square,
14
+ )
15
+
16
+
17
+ class Transform:
18
+ def __call__(self, images: list) -> Union[list[Image.Image], list[np.ndarray]]:
19
+ raise NotImplementedError("Subclasses must implement this method")
20
+
21
+
22
+ class ConvertToRGB(Transform):
23
+ def __call__(self, images: list[Image.Image]) -> list[Image.Image]:
24
+ return [convert_to_rgb(image=image) for image in images]
25
+
26
+
27
+ class CenterCrop(Transform):
28
+ def __init__(self, size: tuple[int, int]):
29
+ self.size = size
30
+
31
+ def __call__(self, images: list[Image.Image]) -> list[np.ndarray]:
32
+ return [center_crop(image=image, size=self.size) for image in images]
33
+
34
+
35
+ class Normalize(Transform):
36
+ def __init__(self, mean: Union[float, list[float]], std: Union[float, list[float]]):
37
+ self.mean = mean
38
+ self.std = std
39
+
40
+ def __call__(self, images: list[np.ndarray]) -> list[np.ndarray]:
41
+ return [normalize(image, mean=self.mean, std=self.std) for image in images]
42
+
43
+
44
+ class Resize(Transform):
45
+ def __init__(
46
+ self,
47
+ size: Union[int, tuple[int, int]],
48
+ resample: Image.Resampling = Image.Resampling.BICUBIC,
49
+ ):
50
+ self.size = size
51
+ self.resample = resample
52
+
53
+ def __call__(self, images: list[Image.Image]) -> list[Image.Image]:
54
+ return [resize(image, size=self.size, resample=self.resample) for image in images]
55
+
56
+
57
+ class Rescale(Transform):
58
+ def __init__(self, scale: float = 1 / 255):
59
+ self.scale = scale
60
+
61
+ def __call__(self, images: list[np.ndarray]) -> list[np.ndarray]:
62
+ return [rescale(image, scale=self.scale) for image in images]
63
+
64
+
65
+ class PILtoNDarray(Transform):
66
+ def __call__(self, images: list[Union[Image.Image, np.ndarray]]) -> list[np.ndarray]:
67
+ return [pil2ndarray(image) for image in images]
68
+
69
+
70
+ class PadtoSquare(Transform):
71
+ def __init__(
72
+ self,
73
+ size: int,
74
+ fill_color: Optional[Union[str, int, tuple[int, ...]]] = None,
75
+ ):
76
+ self.size = size
77
+ self.fill_color = fill_color
78
+
79
+ def __call__(self, images: list[Image.Image]) -> list[Image.Image]:
80
+ return [
81
+ pad2square(image=image, size=self.size, fill_color=self.fill_color) for image in images
82
+ ]
83
+
84
+
85
+ class Compose:
86
+ def __init__(self, transforms: list[Transform]):
87
+ self.transforms = transforms
88
+
89
+ def __call__(
90
+ self, images: Union[list[Image.Image], list[np.ndarray]]
91
+ ) -> Union[list[np.ndarray], list[Image.Image]]:
92
+ for transform in self.transforms:
93
+ images = transform(images)
94
+ return images
95
+
96
+ @classmethod
97
+ def from_config(cls, config: dict[str, Any]) -> "Compose":
98
+ """Creates processor from a config dict.
99
+ Args:
100
+ config (dict[str, Any]): Configuration dictionary.
101
+
102
+ Valid keys:
103
+ - do_resize
104
+ - resize_mode
105
+ - size
106
+ - fill_color
107
+ - do_center_crop
108
+ - crop_size
109
+ - do_rescale
110
+ - rescale_factor
111
+ - do_normalize
112
+ - image_mean
113
+ - mean
114
+ - image_std
115
+ - std
116
+ - resample
117
+ - interpolation
118
+ Valid size keys (nested):
119
+ - {"height", "width"}
120
+ - {"shortest_edge"}
121
+
122
+ Returns:
123
+ Compose: Image processor.
124
+ """
125
+ transforms = []
126
+ cls._get_convert_to_rgb(transforms, config)
127
+ cls._get_resize(transforms, config)
128
+ cls._get_pad2square(transforms, config)
129
+ cls._get_center_crop(transforms, config)
130
+ cls._get_pil2ndarray(transforms, config)
131
+ cls._get_rescale(transforms, config)
132
+ cls._get_normalize(transforms, config)
133
+ return cls(transforms=transforms)
134
+
135
+ @staticmethod
136
+ def _get_convert_to_rgb(transforms: list[Transform], config: dict[str, Any]):
137
+ transforms.append(ConvertToRGB())
138
+
139
+ @classmethod
140
+ def _get_resize(cls, transforms: list[Transform], config: dict[str, Any]):
141
+ mode = config.get("image_processor_type", "CLIPImageProcessor")
142
+ if mode == "CLIPImageProcessor":
143
+ if config.get("do_resize", False):
144
+ size = config["size"]
145
+ if "shortest_edge" in size:
146
+ size = size["shortest_edge"]
147
+ elif "height" in size and "width" in size:
148
+ size = (size["height"], size["width"])
149
+ else:
150
+ raise ValueError(
151
+ "Size must contain either 'shortest_edge' or 'height' and 'width'."
152
+ )
153
+ transforms.append(
154
+ Resize(
155
+ size=size,
156
+ resample=config.get("resample", Image.Resampling.BICUBIC),
157
+ )
158
+ )
159
+ elif mode == "ConvNextFeatureExtractor":
160
+ if "size" in config and "shortest_edge" not in config["size"]:
161
+ raise ValueError(
162
+ f"Size dictionary must contain 'shortest_edge' key. Got {config['size'].keys()}"
163
+ )
164
+ shortest_edge = config["size"]["shortest_edge"]
165
+ crop_pct = config.get("crop_pct", 0.875)
166
+ if shortest_edge < 384:
167
+ # maintain same ratio, resizing shortest edge to shortest_edge/crop_pct
168
+ resize_shortest_edge = int(shortest_edge / crop_pct)
169
+ transforms.append(
170
+ Resize(
171
+ size=resize_shortest_edge,
172
+ resample=config.get("resample", Image.Resampling.BICUBIC),
173
+ )
174
+ )
175
+ transforms.append(CenterCrop(size=(shortest_edge, shortest_edge)))
176
+ else:
177
+ transforms.append(
178
+ Resize(
179
+ size=(shortest_edge, shortest_edge),
180
+ resample=config.get("resample", Image.Resampling.BICUBIC),
181
+ )
182
+ )
183
+ elif mode == "JinaCLIPImageProcessor":
184
+ interpolation = config.get("interpolation")
185
+ if isinstance(interpolation, str):
186
+ resample = cls._interpolation_resolver(interpolation)
187
+ else:
188
+ resample = interpolation or Image.Resampling.BICUBIC
189
+
190
+ if "size" in config:
191
+ resize_mode = config.get("resize_mode", "shortest")
192
+ if resize_mode == "shortest":
193
+ transforms.append(
194
+ Resize(
195
+ size=config["size"],
196
+ resample=resample,
197
+ )
198
+ )
199
+ else:
200
+ raise ValueError(f"Preprocessor {mode} is not supported")
201
+
202
+ @staticmethod
203
+ def _get_center_crop(transforms: list[Transform], config: dict[str, Any]):
204
+ mode = config.get("image_processor_type", "CLIPImageProcessor")
205
+ if mode == "CLIPImageProcessor":
206
+ if config.get("do_center_crop", False):
207
+ crop_size = config["crop_size"]
208
+ if isinstance(crop_size, int):
209
+ crop_size = (crop_size, crop_size)
210
+ elif isinstance(crop_size, dict):
211
+ crop_size = (crop_size["height"], crop_size["width"])
212
+ else:
213
+ raise ValueError(f"Invalid crop size: {crop_size}")
214
+ transforms.append(CenterCrop(size=crop_size))
215
+ elif mode == "ConvNextFeatureExtractor":
216
+ pass
217
+ elif mode == "JinaCLIPImageProcessor":
218
+ pass
219
+ else:
220
+ raise ValueError(f"Preprocessor {mode} is not supported")
221
+
222
+ @staticmethod
223
+ def _get_pil2ndarray(transforms: list[Transform], config: dict[str, Any]):
224
+ transforms.append(PILtoNDarray())
225
+
226
+ @staticmethod
227
+ def _get_rescale(transforms: list[Transform], config: dict[str, Any]):
228
+ if config.get("do_rescale", True):
229
+ rescale_factor = config.get("rescale_factor", 1 / 255)
230
+ transforms.append(Rescale(scale=rescale_factor))
231
+
232
+ @staticmethod
233
+ def _get_normalize(transforms: list[Transform], config: dict[str, Any]):
234
+ if config.get("do_normalize", False):
235
+ transforms.append(Normalize(mean=config["image_mean"], std=config["image_std"]))
236
+ elif "mean" in config and "std" in config:
237
+ transforms.append(Normalize(mean=config["mean"], std=config["std"]))
238
+
239
+ @staticmethod
240
+ def _get_pad2square(transforms: list[Transform], config: dict[str, Any]):
241
+ mode = config.get("image_processor_type", "CLIPImageProcessor")
242
+ if mode == "CLIPImageProcessor":
243
+ pass
244
+ elif mode == "ConvNextFeatureExtractor":
245
+ pass
246
+ elif mode == "JinaCLIPImageProcessor":
247
+ transforms.append(
248
+ PadtoSquare(
249
+ size=config["size"],
250
+ fill_color=config.get("fill_color", 0),
251
+ )
252
+ )
253
+
254
+ @staticmethod
255
+ def _interpolation_resolver(resample: Optional[str] = None) -> Image.Resampling:
256
+ interpolation_map = {
257
+ "nearest": Image.Resampling.NEAREST,
258
+ "lanczos": Image.Resampling.LANCZOS,
259
+ "bilinear": Image.Resampling.BILINEAR,
260
+ "bicubic": Image.Resampling.BICUBIC,
261
+ "box": Image.Resampling.BOX,
262
+ "hamming": Image.Resampling.HAMMING,
263
+ }
264
+
265
+ if resample and (method := interpolation_map.get(resample.lower())):
266
+ return method
267
+
268
+ raise ValueError(f"Unknown interpolation method: {resample}")
@@ -0,0 +1,5 @@
1
+ from fastembed.late_interaction.late_interaction_text_embedding import (
2
+ LateInteractionTextEmbedding,
3
+ )
4
+
5
+ __all__ = ["LateInteractionTextEmbedding"]
@@ -0,0 +1,256 @@
1
+ import string
2
+ from typing import Any, Iterable, Optional, Sequence, Type, Union
3
+
4
+ import numpy as np
5
+ from tokenizers import Encoding
6
+
7
+ from fastembed.common import OnnxProvider
8
+ from fastembed.common.onnx_model import OnnxOutputContext
9
+ from fastembed.common.utils import define_cache_dir
10
+ from fastembed.late_interaction.late_interaction_embedding_base import (
11
+ LateInteractionTextEmbeddingBase,
12
+ )
13
+ from fastembed.text.onnx_text_model import OnnxTextModel, TextEmbeddingWorker
14
+
15
+
16
+ supported_colbert_models = [
17
+ {
18
+ "model": "colbert-ir/colbertv2.0",
19
+ "dim": 128,
20
+ "description": "Late interaction model",
21
+ "license": "mit",
22
+ "size_in_GB": 0.44,
23
+ "sources": {
24
+ "hf": "colbert-ir/colbertv2.0",
25
+ },
26
+ "model_file": "model.onnx",
27
+ },
28
+ {
29
+ "model": "answerdotai/answerai-colbert-small-v1",
30
+ "dim": 96,
31
+ "description": "Text embeddings, Unimodal (text), Multilingual (~100 languages), 512 input tokens truncation, 2024 year",
32
+ "license": "apache-2.0",
33
+ "size_in_GB": 0.13,
34
+ "sources": {
35
+ "hf": "answerdotai/answerai-colbert-small-v1",
36
+ },
37
+ "model_file": "vespa_colbert.onnx",
38
+ },
39
+ ]
40
+
41
+
42
+ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[np.ndarray]):
43
+ QUERY_MARKER_TOKEN_ID = 1
44
+ DOCUMENT_MARKER_TOKEN_ID = 2
45
+ MIN_QUERY_LENGTH = 31 # it's 32, we add one additional special token in the beginning
46
+ MASK_TOKEN = "[MASK]"
47
+
48
+ def _post_process_onnx_output(
49
+ self, output: OnnxOutputContext, is_doc: bool = True
50
+ ) -> Iterable[np.ndarray]:
51
+ if not is_doc:
52
+ return output.model_output.astype(np.float32)
53
+
54
+ if output.input_ids is None or output.attention_mask is None:
55
+ raise ValueError(
56
+ "input_ids and attention_mask must be provided for document post-processing"
57
+ )
58
+
59
+ for i, token_sequence in enumerate(output.input_ids):
60
+ for j, token_id in enumerate(token_sequence):
61
+ if token_id in self.skip_list or token_id == self.pad_token_id:
62
+ output.attention_mask[i, j] = 0
63
+
64
+ output.model_output *= np.expand_dims(output.attention_mask, 2).astype(np.float32)
65
+ norm = np.linalg.norm(output.model_output, ord=2, axis=2, keepdims=True)
66
+ norm_clamped = np.maximum(norm, 1e-12)
67
+ output.model_output /= norm_clamped
68
+ return output.model_output.astype(np.float32)
69
+
70
+ def _preprocess_onnx_input(
71
+ self, onnx_input: dict[str, np.ndarray], is_doc: bool = True, **kwargs: Any
72
+ ) -> dict[str, np.ndarray]:
73
+ marker_token = self.DOCUMENT_MARKER_TOKEN_ID if is_doc else self.QUERY_MARKER_TOKEN_ID
74
+ onnx_input["input_ids"] = np.insert(onnx_input["input_ids"], 1, marker_token, axis=1)
75
+ onnx_input["attention_mask"] = np.insert(onnx_input["attention_mask"], 1, 1, axis=1)
76
+ return onnx_input
77
+
78
+ def tokenize(self, documents: list[str], is_doc: bool = True, **kwargs: Any) -> list[Encoding]:
79
+ return (
80
+ self._tokenize_documents(documents=documents)
81
+ if is_doc
82
+ else self._tokenize_query(query=next(iter(documents)))
83
+ )
84
+
85
+ def _tokenize_query(self, query: str) -> list[Encoding]:
86
+ encoded = self.tokenizer.encode_batch([query])
87
+ # colbert authors recommend to pad queries with [MASK] tokens for query augmentation to improve performance
88
+ if len(encoded[0].ids) < self.MIN_QUERY_LENGTH:
89
+ prev_padding = None
90
+ if self.tokenizer.padding:
91
+ prev_padding = self.tokenizer.padding
92
+ self.tokenizer.enable_padding(
93
+ pad_token=self.MASK_TOKEN,
94
+ pad_id=self.mask_token_id,
95
+ length=self.MIN_QUERY_LENGTH,
96
+ )
97
+ encoded = self.tokenizer.encode_batch([query])
98
+ if prev_padding is None:
99
+ self.tokenizer.no_padding()
100
+ else:
101
+ self.tokenizer.enable_padding(**prev_padding)
102
+ return encoded
103
+
104
+ def _tokenize_documents(self, documents: list[str]) -> list[Encoding]:
105
+ encoded = self.tokenizer.encode_batch(documents)
106
+ return encoded
107
+
108
+ @classmethod
109
+ def list_supported_models(cls) -> list[dict[str, Any]]:
110
+ """Lists the supported models.
111
+
112
+ Returns:
113
+ list[dict[str, Any]]: A list of dictionaries containing the model information.
114
+ """
115
+ return supported_colbert_models
116
+
117
+ def __init__(
118
+ self,
119
+ model_name: str,
120
+ cache_dir: Optional[str] = None,
121
+ threads: Optional[int] = None,
122
+ providers: Optional[Sequence[OnnxProvider]] = None,
123
+ cuda: bool = False,
124
+ device_ids: Optional[list[int]] = None,
125
+ lazy_load: bool = False,
126
+ device_id: Optional[int] = None,
127
+ **kwargs,
128
+ ):
129
+ """
130
+ Args:
131
+ model_name (str): The name of the model to use.
132
+ cache_dir (str, optional): The path to the cache directory.
133
+ Can be set using the `FASTEMBED_CACHE_PATH` env variable.
134
+ Defaults to `fastembed_cache` in the system's temp directory.
135
+ threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.
136
+ providers (Optional[Sequence[OnnxProvider]], optional): The list of onnxruntime providers to use.
137
+ Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
138
+ cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
139
+ Defaults to False.
140
+ device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
141
+ workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
142
+ lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
143
+ Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
144
+ device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
145
+
146
+ Raises:
147
+ ValueError: If the model_name is not in the format <org>/<model> e.g. BAAI/bge-base-en.
148
+ """
149
+
150
+ super().__init__(model_name, cache_dir, threads, **kwargs)
151
+ self.providers = providers
152
+ self.lazy_load = lazy_load
153
+
154
+ # List of device ids, that can be used for data parallel processing in workers
155
+ self.device_ids = device_ids
156
+ self.cuda = cuda
157
+
158
+ # This device_id will be used if we need to load model in current process
159
+ if device_id is not None:
160
+ self.device_id = device_id
161
+ elif self.device_ids is not None:
162
+ self.device_id = self.device_ids[0]
163
+ else:
164
+ self.device_id = None
165
+
166
+ self.model_description = self._get_model_description(model_name)
167
+ self.cache_dir = define_cache_dir(cache_dir)
168
+
169
+ self._model_dir = self.download_model(
170
+ self.model_description, self.cache_dir, local_files_only=self._local_files_only
171
+ )
172
+ self.mask_token_id = None
173
+ self.pad_token_id = None
174
+ self.skip_list = set()
175
+
176
+ if not self.lazy_load:
177
+ self.load_onnx_model()
178
+
179
+ def load_onnx_model(self) -> None:
180
+ self._load_onnx_model(
181
+ model_dir=self._model_dir,
182
+ model_file=self.model_description["model_file"],
183
+ threads=self.threads,
184
+ providers=self.providers,
185
+ cuda=self.cuda,
186
+ device_id=self.device_id,
187
+ )
188
+ self.mask_token_id = self.special_token_to_id[self.MASK_TOKEN]
189
+ self.pad_token_id = self.tokenizer.padding["pad_id"]
190
+ self.skip_list = {
191
+ self.tokenizer.encode(symbol, add_special_tokens=False).ids[0]
192
+ for symbol in string.punctuation
193
+ }
194
+ current_max_length = self.tokenizer.truncation["max_length"]
195
+ # ensure not to overflow after adding document-marker
196
+ self.tokenizer.enable_truncation(max_length=current_max_length - 1)
197
+
198
+ def embed(
199
+ self,
200
+ documents: Union[str, Iterable[str]],
201
+ batch_size: int = 256,
202
+ parallel: Optional[int] = None,
203
+ **kwargs,
204
+ ) -> Iterable[np.ndarray]:
205
+ """
206
+ Encode a list of documents into list of embeddings.
207
+ We use mean pooling with attention so that the model can handle variable-length inputs.
208
+
209
+ Args:
210
+ documents: Iterator of documents or single document to embed
211
+ batch_size: Batch size for encoding -- higher values will use more memory, but be faster
212
+ parallel:
213
+ If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.
214
+ If 0, use all available cores.
215
+ If None, don't use data-parallel processing, use default onnxruntime threading instead.
216
+
217
+ Returns:
218
+ List of embeddings, one per document
219
+ """
220
+ yield from self._embed_documents(
221
+ model_name=self.model_name,
222
+ cache_dir=str(self.cache_dir),
223
+ documents=documents,
224
+ batch_size=batch_size,
225
+ parallel=parallel,
226
+ providers=self.providers,
227
+ cuda=self.cuda,
228
+ device_ids=self.device_ids,
229
+ **kwargs,
230
+ )
231
+
232
+ def query_embed(self, query: Union[str, Iterable[str]], **kwargs) -> Iterable[np.ndarray]:
233
+ if isinstance(query, str):
234
+ query = [query]
235
+
236
+ if not hasattr(self, "model") or self.model is None:
237
+ self.load_onnx_model()
238
+
239
+ for text in query:
240
+ yield from self._post_process_onnx_output(
241
+ self.onnx_embed([text], is_doc=False), is_doc=False
242
+ )
243
+
244
+ @classmethod
245
+ def _get_worker_class(cls) -> Type[TextEmbeddingWorker]:
246
+ return ColbertEmbeddingWorker
247
+
248
+
249
+ class ColbertEmbeddingWorker(TextEmbeddingWorker):
250
+ def init_embedding(self, model_name: str, cache_dir: str, **kwargs) -> Colbert:
251
+ return Colbert(
252
+ model_name=model_name,
253
+ cache_dir=cache_dir,
254
+ threads=1,
255
+ **kwargs,
256
+ )
@@ -0,0 +1,62 @@
1
+ from typing import Any, Type
2
+
3
+ import numpy as np
4
+
5
+ from fastembed.late_interaction.colbert import Colbert
6
+ from fastembed.text.onnx_text_model import TextEmbeddingWorker
7
+
8
+
9
+ supported_jina_colbert_models = [
10
+ {
11
+ "model": "jinaai/jina-colbert-v2",
12
+ "dim": 128,
13
+ "description": "New model that expands capabilities of colbert-v1 with multilingual and context length of 8192, 2024 year",
14
+ "license": "cc-by-nc-4.0",
15
+ "size_in_GB": 2.24,
16
+ "sources": {
17
+ "hf": "jinaai/jina-colbert-v2",
18
+ },
19
+ "model_file": "onnx/model.onnx",
20
+ "additional_files": ["onnx/model.onnx_data"],
21
+ },
22
+ ]
23
+
24
+
25
+ class JinaColbert(Colbert):
26
+ QUERY_MARKER_TOKEN_ID = 250002
27
+ DOCUMENT_MARKER_TOKEN_ID = 250003
28
+ MIN_QUERY_LENGTH = 31 # it's 32, we add one additional special token in the beginning
29
+ MASK_TOKEN = "<mask>"
30
+
31
+ @classmethod
32
+ def _get_worker_class(cls) -> Type[TextEmbeddingWorker]:
33
+ return JinaColbertEmbeddingWorker
34
+
35
+ @classmethod
36
+ def list_supported_models(cls) -> list[dict[str, Any]]:
37
+ """Lists the supported models.
38
+
39
+ Returns:
40
+ list[dict[str, Any]]: A list of dictionaries containing the model information.
41
+ """
42
+ return supported_jina_colbert_models
43
+
44
+ def _preprocess_onnx_input(
45
+ self, onnx_input: dict[str, np.ndarray], is_doc: bool = True, **kwargs: Any
46
+ ) -> dict[str, np.ndarray]:
47
+ onnx_input = super()._preprocess_onnx_input(onnx_input, is_doc)
48
+
49
+ # the attention mask for jina-colbert-v2 is always 1 in queries
50
+ if not is_doc:
51
+ onnx_input["attention_mask"][:] = 1
52
+ return onnx_input
53
+
54
+
55
+ class JinaColbertEmbeddingWorker(TextEmbeddingWorker):
56
+ def init_embedding(self, model_name: str, cache_dir: str, **kwargs) -> JinaColbert:
57
+ return JinaColbert(
58
+ model_name=model_name,
59
+ cache_dir=cache_dir,
60
+ threads=1,
61
+ **kwargs,
62
+ )
@@ -0,0 +1,62 @@
1
+ from typing import Iterable, Optional, Union
2
+
3
+ import numpy as np
4
+
5
+ from fastembed.common.model_management import ModelManagement
6
+
7
+
8
+ class LateInteractionTextEmbeddingBase(ModelManagement):
9
+ def __init__(
10
+ self,
11
+ model_name: str,
12
+ cache_dir: Optional[str] = None,
13
+ threads: Optional[int] = None,
14
+ **kwargs,
15
+ ):
16
+ self.model_name = model_name
17
+ self.cache_dir = cache_dir
18
+ self.threads = threads
19
+ self._local_files_only = kwargs.pop("local_files_only", False)
20
+
21
+ def embed(
22
+ self,
23
+ documents: Union[str, Iterable[str]],
24
+ batch_size: int = 256,
25
+ parallel: Optional[int] = None,
26
+ **kwargs,
27
+ ) -> Iterable[np.ndarray]:
28
+ raise NotImplementedError()
29
+
30
+ def passage_embed(self, texts: Iterable[str], **kwargs) -> Iterable[np.ndarray]:
31
+ """
32
+ Embeds a list of text passages into a list of embeddings.
33
+
34
+ Args:
35
+ texts (Iterable[str]): The list of texts to embed.
36
+ **kwargs: Additional keyword argument to pass to the embed method.
37
+
38
+ Yields:
39
+ Iterable[np.ndarray]: The embeddings.
40
+ """
41
+
42
+ # This is model-specific, so that different models can have specialized implementations
43
+ yield from self.embed(texts, **kwargs)
44
+
45
+ def query_embed(
46
+ self, query: Union[str, Iterable[str]], **kwargs
47
+ ) -> Iterable[np.ndarray]:
48
+ """
49
+ Embeds queries
50
+
51
+ Args:
52
+ query (Union[str, Iterable[str]]): The query to embed, or an iterable e.g. list of queries.
53
+
54
+ Returns:
55
+ Iterable[np.ndarray]: The embeddings.
56
+ """
57
+
58
+ # This is model-specific, so that different models can have specialized implementations
59
+ if isinstance(query, str):
60
+ yield from self.embed([query], **kwargs)
61
+ if isinstance(query, Iterable):
62
+ yield from self.embed(query, **kwargs)