fastembed-gpu 0.6.0__tar.gz → 0.7.0__tar.gz

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 (61) hide show
  1. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/PKG-INFO +18 -1
  2. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/README.md +17 -0
  3. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/common/model_description.py +5 -0
  4. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/common/model_management.py +3 -1
  5. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/common/onnx_model.py +10 -1
  6. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/common/preprocessor_utils.py +3 -3
  7. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/common/types.py +1 -0
  8. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/image/onnx_embedding.py +4 -3
  9. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/image/onnx_image_model.py +12 -3
  10. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/image/transform/functional.py +10 -10
  11. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/late_interaction/colbert.py +4 -4
  12. fastembed_gpu-0.7.0/fastembed/late_interaction/token_embeddings.py +83 -0
  13. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/late_interaction_multimodal/colpali.py +2 -2
  14. fastembed_gpu-0.7.0/fastembed/rerank/cross_encoder/custom_text_cross_encoder.py +46 -0
  15. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/rerank/cross_encoder/onnx_text_cross_encoder.py +3 -1
  16. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/rerank/cross_encoder/onnx_text_model.py +12 -1
  17. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/rerank/cross_encoder/text_cross_encoder.py +38 -1
  18. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/sparse/bm25.py +2 -12
  19. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/sparse/bm42.py +3 -1
  20. fastembed_gpu-0.7.0/fastembed/sparse/minicoil.py +349 -0
  21. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/sparse/sparse_text_embedding.py +2 -1
  22. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/sparse/splade_pp.py +5 -3
  23. fastembed_gpu-0.7.0/fastembed/sparse/utils/minicoil_encoder.py +146 -0
  24. fastembed_gpu-0.7.0/fastembed/sparse/utils/sparse_vectors_converter.py +247 -0
  25. fastembed_gpu-0.7.0/fastembed/sparse/utils/vocab_resolver.py +202 -0
  26. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/text/clip_embedding.py +3 -1
  27. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/text/custom_text_embedding.py +3 -1
  28. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/text/multitask_embedding.py +26 -15
  29. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/text/onnx_embedding.py +8 -3
  30. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/text/onnx_text_model.py +14 -3
  31. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/text/pooled_embedding.py +5 -2
  32. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/text/pooled_normalized_embedding.py +5 -3
  33. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/pyproject.toml +2 -2
  34. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/LICENSE +0 -0
  35. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/NOTICE +0 -0
  36. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/__init__.py +0 -0
  37. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/common/__init__.py +0 -0
  38. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/common/utils.py +0 -0
  39. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/embedding.py +0 -0
  40. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/image/__init__.py +0 -0
  41. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/image/image_embedding.py +0 -0
  42. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/image/image_embedding_base.py +0 -0
  43. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/image/transform/operators.py +0 -0
  44. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/late_interaction/__init__.py +0 -0
  45. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/late_interaction/jina_colbert.py +0 -0
  46. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/late_interaction/late_interaction_embedding_base.py +0 -0
  47. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/late_interaction/late_interaction_text_embedding.py +0 -0
  48. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/late_interaction_multimodal/__init__.py +0 -0
  49. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/late_interaction_multimodal/late_interaction_multimodal_embedding.py +0 -0
  50. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/late_interaction_multimodal/late_interaction_multimodal_embedding_base.py +0 -0
  51. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/late_interaction_multimodal/onnx_multimodal_model.py +0 -0
  52. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/parallel_processor.py +0 -0
  53. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/py.typed +0 -0
  54. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/rerank/cross_encoder/__init__.py +0 -0
  55. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/rerank/cross_encoder/text_cross_encoder_base.py +0 -0
  56. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/sparse/__init__.py +0 -0
  57. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/sparse/sparse_embedding_base.py +0 -0
  58. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/sparse/utils/tokenizer.py +0 -0
  59. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/text/__init__.py +0 -0
  60. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/text/text_embedding.py +0 -0
  61. {fastembed_gpu-0.6.0 → fastembed_gpu-0.7.0}/fastembed/text/text_embedding_base.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.3
2
2
  Name: fastembed-gpu
3
- Version: 0.6.0
3
+ Version: 0.7.0
4
4
  Summary: Fast, light, accurate library built for retrieval embedding generation
5
5
  License: Apache License
6
6
  Keywords: vector,embedding,neural,search,qdrant,sentence-transformers
@@ -225,6 +225,23 @@ scores = list(encoder.rerank(query, documents))
225
225
  # [-11.48061752319336, 5.472434997558594]
226
226
  ```
227
227
 
228
+ Text cross encoders can also be extended with models which are not in the list of supported models.
229
+
230
+ ```python
231
+ from fastembed.rerank.cross_encoder import TextCrossEncoder
232
+ from fastembed.common.model_description import ModelSource
233
+
234
+ TextCrossEncoder.add_custom_model(
235
+ model="Xenova/ms-marco-MiniLM-L-4-v2",
236
+ model_file="onnx/model.onnx",
237
+ sources=ModelSource(hf="Xenova/ms-marco-MiniLM-L-4-v2"),
238
+ )
239
+ model = TextCrossEncoder(model_name="Xenova/ms-marco-MiniLM-L-4-v2")
240
+ scores = list(model.rerank_pairs(
241
+ [("What is AI?", "Artificial intelligence is ..."), ("What is ML?", "Machine learning is ..."),]
242
+ ))
243
+ ```
244
+
228
245
  ## ⚡️ FastEmbed on a GPU
229
246
 
230
247
  FastEmbed supports running on GPU devices.
@@ -190,6 +190,23 @@ scores = list(encoder.rerank(query, documents))
190
190
  # [-11.48061752319336, 5.472434997558594]
191
191
  ```
192
192
 
193
+ Text cross encoders can also be extended with models which are not in the list of supported models.
194
+
195
+ ```python
196
+ from fastembed.rerank.cross_encoder import TextCrossEncoder
197
+ from fastembed.common.model_description import ModelSource
198
+
199
+ TextCrossEncoder.add_custom_model(
200
+ model="Xenova/ms-marco-MiniLM-L-4-v2",
201
+ model_file="onnx/model.onnx",
202
+ sources=ModelSource(hf="Xenova/ms-marco-MiniLM-L-4-v2"),
203
+ )
204
+ model = TextCrossEncoder(model_name="Xenova/ms-marco-MiniLM-L-4-v2")
205
+ scores = list(model.rerank_pairs(
206
+ [("What is AI?", "Artificial intelligence is ..."), ("What is ML?", "Machine learning is ..."),]
207
+ ))
208
+ ```
209
+
193
210
  ## ⚡️ FastEmbed on a GPU
194
211
 
195
212
  FastEmbed supports running on GPU devices.
@@ -7,6 +7,11 @@ from typing import Optional, Any
7
7
  class ModelSource:
8
8
  hf: Optional[str] = None
9
9
  url: Optional[str] = None
10
+ _deprecated_tar_struct: bool = False
11
+
12
+ @property
13
+ def deprecated_tar_struct(self) -> bool:
14
+ return self._deprecated_tar_struct
10
15
 
11
16
  def __post_init__(self) -> None:
12
17
  if self.hf is None and self.url is None:
@@ -330,9 +330,10 @@ class ModelManagement(Generic[T]):
330
330
  model_name: str,
331
331
  source_url: str,
332
332
  cache_dir: str,
333
+ deprecated_tar_struct: bool = False,
333
334
  local_files_only: bool = False,
334
335
  ) -> Path:
335
- fast_model_name = f"fast-{model_name.split('/')[-1]}"
336
+ fast_model_name = f"{'fast-' if deprecated_tar_struct else ''}{model_name.split('/')[-1]}"
336
337
  cache_tmp_dir = Path(cache_dir) / "tmp"
337
338
  model_tmp_dir = cache_tmp_dir / fast_model_name
338
339
  model_dir = Path(cache_dir) / fast_model_name
@@ -438,6 +439,7 @@ class ModelManagement(Generic[T]):
438
439
  model.model,
439
440
  str(url_source),
440
441
  str(cache_dir),
442
+ deprecated_tar_struct=model.sources.deprecated_tar_struct,
441
443
  local_files_only=local_files_only,
442
444
  )
443
445
  except Exception:
@@ -28,7 +28,16 @@ class OnnxModel(Generic[T]):
28
28
  def _get_worker_class(cls) -> Type["EmbeddingWorker[T]"]:
29
29
  raise NotImplementedError("Subclasses must implement this method")
30
30
 
31
- def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[T]:
31
+ def _post_process_onnx_output(self, output: OnnxOutputContext, **kwargs: Any) -> Iterable[T]:
32
+ """Post-process the ONNX model output to convert it into a usable format.
33
+
34
+ Args:
35
+ output (OnnxOutputContext): The raw output from the ONNX model.
36
+ **kwargs: Additional keyword arguments that may be needed by specific implementations.
37
+
38
+ Returns:
39
+ Iterable[T]: Post-processed output as an iterable of type T.
40
+ """
32
41
  raise NotImplementedError("Subclasses must implement this method")
33
42
 
34
43
  def __init__(self) -> None:
@@ -36,9 +36,9 @@ def load_tokenizer(model_dir: Path) -> tuple[Tokenizer, dict[str, int]]:
36
36
 
37
37
  with open(str(tokenizer_config_path)) as tokenizer_config_file:
38
38
  tokenizer_config = json.load(tokenizer_config_file)
39
- assert (
40
- "model_max_length" in tokenizer_config or "max_length" in tokenizer_config
41
- ), "Models without model_max_length or max_length are not supported."
39
+ assert "model_max_length" in tokenizer_config or "max_length" in tokenizer_config, (
40
+ "Models without model_max_length or max_length are not supported."
41
+ )
42
42
  if "model_max_length" not in tokenizer_config:
43
43
  max_context = tokenizer_config["max_length"]
44
44
  elif "max_length" not in tokenizer_config:
@@ -16,6 +16,7 @@ ImageInput: TypeAlias = Union[PathInput, Image.Image]
16
16
 
17
17
  OnnxProvider: TypeAlias = Union[str, tuple[str, dict[Any, Any]]]
18
18
  NumpyArray = Union[
19
+ NDArray[np.float64],
19
20
  NDArray[np.float32],
20
21
  NDArray[np.float16],
21
22
  NDArray[np.int8],
@@ -1,6 +1,5 @@
1
1
  from typing import Any, Iterable, Optional, Sequence, Type, Union
2
2
 
3
- import numpy as np
4
3
 
5
4
  from fastembed.common.types import NumpyArray
6
5
  from fastembed.common import ImageInput, OnnxProvider
@@ -194,8 +193,10 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
194
193
 
195
194
  return onnx_input
196
195
 
197
- def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[NumpyArray]:
198
- return normalize(output.model_output).astype(np.float32)
196
+ def _post_process_onnx_output(
197
+ self, output: OnnxOutputContext, **kwargs: Any
198
+ ) -> Iterable[NumpyArray]:
199
+ return normalize(output.model_output)
199
200
 
200
201
 
201
202
  class OnnxImageEmbeddingWorker(ImageEmbeddingWorker[NumpyArray]):
@@ -23,7 +23,16 @@ class OnnxImageModel(OnnxModel[T]):
23
23
  def _get_worker_class(cls) -> Type["ImageEmbeddingWorker[T]"]:
24
24
  raise NotImplementedError("Subclasses must implement this method")
25
25
 
26
- def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[T]:
26
+ def _post_process_onnx_output(self, output: OnnxOutputContext, **kwargs: Any) -> Iterable[T]:
27
+ """Post-process the ONNX model output to convert it into a usable format.
28
+
29
+ Args:
30
+ output (OnnxOutputContext): The raw output from the ONNX model.
31
+ **kwargs: Additional keyword arguments that may be needed by specific implementations.
32
+
33
+ Returns:
34
+ Iterable[T]: Post-processed output as an iterable of type T.
35
+ """
27
36
  raise NotImplementedError("Subclasses must implement this method")
28
37
 
29
38
  def __init__(self) -> None:
@@ -104,7 +113,7 @@ class OnnxImageModel(OnnxModel[T]):
104
113
  self.load_onnx_model()
105
114
 
106
115
  for batch in iter_batch(images, batch_size):
107
- yield from self._post_process_onnx_output(self.onnx_embed(batch))
116
+ yield from self._post_process_onnx_output(self.onnx_embed(batch), **kwargs)
108
117
  else:
109
118
  if parallel == 0:
110
119
  parallel = os.cpu_count()
@@ -125,7 +134,7 @@ class OnnxImageModel(OnnxModel[T]):
125
134
  start_method=start_method,
126
135
  )
127
136
  for batch in pool.ordered_map(iter_batch(images, batch_size), **params):
128
- yield from self._post_process_onnx_output(batch) # type: ignore
137
+ yield from self._post_process_onnx_output(batch, **kwargs) # type: ignore
129
138
 
130
139
 
131
140
  class ImageEmbeddingWorker(EmbeddingWorker[T]):
@@ -72,26 +72,26 @@ def normalize(
72
72
  if not np.issubdtype(image.dtype, np.floating):
73
73
  image = image.astype(np.float32)
74
74
 
75
- mean = mean if isinstance(mean, list) else [mean] * num_channels
75
+ mean_list = mean if isinstance(mean, list) else [mean] * num_channels
76
76
 
77
- if len(mean) != num_channels:
77
+ if len(mean_list) != num_channels:
78
78
  raise ValueError(
79
79
  f"mean must have the same number of channels as the image, image has {num_channels} channels, got "
80
- f"{len(mean)}"
80
+ f"{len(mean_list)}"
81
81
  )
82
82
 
83
- mean_arr = np.array(mean, dtype=np.float32)
83
+ mean_arr = np.array(mean_list, dtype=np.float32)
84
84
 
85
- std = std if isinstance(std, list) else [std] * num_channels
86
- if len(std) != num_channels:
85
+ std_list = std if isinstance(std, list) else [std] * num_channels
86
+ if len(std_list) != num_channels:
87
87
  raise ValueError(
88
- f"std must have the same number of channels as the image, image has {num_channels} channels, got {len(std)}"
88
+ f"std must have the same number of channels as the image, image has {num_channels} channels, got {len(std_list)}"
89
89
  )
90
90
 
91
- std_arr = np.array(std, dtype=np.float32)
91
+ std_arr = np.array(std_list, dtype=np.float32)
92
92
 
93
- image = ((image.T - mean_arr) / std_arr).T
94
- return image
93
+ image_upd = ((image.T - mean_arr) / std_arr).T
94
+ return image_upd
95
95
 
96
96
 
97
97
  def resize(
@@ -43,10 +43,10 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
43
43
  MASK_TOKEN = "[MASK]"
44
44
 
45
45
  def _post_process_onnx_output(
46
- self, output: OnnxOutputContext, is_doc: bool = True
46
+ self, output: OnnxOutputContext, is_doc: bool = True, **kwargs: Any
47
47
  ) -> Iterable[NumpyArray]:
48
48
  if not is_doc:
49
- return output.model_output.astype(np.float32)
49
+ return output.model_output
50
50
 
51
51
  if output.input_ids is None or output.attention_mask is None:
52
52
  raise ValueError(
@@ -58,11 +58,11 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
58
58
  if token_id in self.skip_list or token_id == self.pad_token_id:
59
59
  output.attention_mask[i, j] = 0
60
60
 
61
- output.model_output *= np.expand_dims(output.attention_mask, 2).astype(np.float32)
61
+ output.model_output *= np.expand_dims(output.attention_mask, 2)
62
62
  norm = np.linalg.norm(output.model_output, ord=2, axis=2, keepdims=True)
63
63
  norm_clamped = np.maximum(norm, 1e-12)
64
64
  output.model_output /= norm_clamped
65
- return output.model_output.astype(np.float32)
65
+ return output.model_output
66
66
 
67
67
  def _preprocess_onnx_input(
68
68
  self, onnx_input: dict[str, NumpyArray], is_doc: bool = True, **kwargs: Any
@@ -0,0 +1,83 @@
1
+ from dataclasses import asdict
2
+ from typing import Union, Iterable, Optional, Any, Type
3
+
4
+ from fastembed.common.model_description import DenseModelDescription, ModelSource
5
+ from fastembed.common.onnx_model import OnnxOutputContext
6
+ from fastembed.common.types import NumpyArray
7
+ from fastembed.late_interaction.late_interaction_embedding_base import (
8
+ LateInteractionTextEmbeddingBase,
9
+ )
10
+ from fastembed.text.onnx_embedding import OnnxTextEmbedding
11
+ from fastembed.text.onnx_text_model import TextEmbeddingWorker
12
+ import numpy as np
13
+
14
+ supported_token_embeddings_models = [
15
+ DenseModelDescription(
16
+ model="jinaai/jina-embeddings-v2-small-en-tokens",
17
+ dim=512,
18
+ description="Text embeddings, Unimodal (text), English, 8192 input tokens truncation,"
19
+ " Prefixes for queries/documents: not necessary, 2023 year.",
20
+ license="apache-2.0",
21
+ size_in_GB=0.12,
22
+ sources=ModelSource(hf="xenova/jina-embeddings-v2-small-en"),
23
+ model_file="onnx/model.onnx",
24
+ ),
25
+ ]
26
+
27
+
28
+ class TokenEmbeddingsModel(OnnxTextEmbedding, LateInteractionTextEmbeddingBase):
29
+ @classmethod
30
+ def _list_supported_models(cls) -> list[DenseModelDescription]:
31
+ """Lists the supported models.
32
+
33
+ Returns:
34
+ list[DenseModelDescription]: A list of DenseModelDescription objects containing the model information.
35
+ """
36
+ return supported_token_embeddings_models
37
+
38
+ @classmethod
39
+ def list_supported_models(cls) -> list[dict[str, Any]]:
40
+ """Lists the supported models.
41
+
42
+ Returns:
43
+ list[dict[str, Any]]: A list of dictionaries containing the model information.
44
+ """
45
+ return [asdict(model) for model in cls._list_supported_models()]
46
+
47
+ @classmethod
48
+ def _get_worker_class(cls) -> Type[TextEmbeddingWorker[NumpyArray]]:
49
+ return TokensEmbeddingWorker
50
+
51
+ def _post_process_onnx_output(
52
+ self, output: OnnxOutputContext, **kwargs: Any
53
+ ) -> Iterable[NumpyArray]:
54
+ # Size: (batch_size, sequence_length, hidden_size)
55
+ embeddings = output.model_output
56
+ # Size: (batch_size, sequence_length)
57
+ assert output.attention_mask is not None
58
+ masks = output.attention_mask
59
+
60
+ # For each document we only select those embeddings that are not masked out
61
+ for i in range(embeddings.shape[0]):
62
+ yield embeddings[i, masks[i] == 1]
63
+
64
+ def embed(
65
+ self,
66
+ documents: Union[str, Iterable[str]],
67
+ batch_size: int = 256,
68
+ parallel: Optional[int] = None,
69
+ **kwargs: Any,
70
+ ) -> Iterable[NumpyArray]:
71
+ yield from super().embed(documents, batch_size=batch_size, parallel=parallel, **kwargs)
72
+
73
+
74
+ class TokensEmbeddingWorker(TextEmbeddingWorker[NumpyArray]):
75
+ def init_embedding(
76
+ self, model_name: str, cache_dir: str, **kwargs: Any
77
+ ) -> TokenEmbeddingsModel:
78
+ return TokenEmbeddingsModel(
79
+ model_name=model_name,
80
+ cache_dir=cache_dir,
81
+ threads=1,
82
+ **kwargs,
83
+ )
@@ -142,7 +142,7 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
142
142
  assert self.model_description.dim is not None, "Model dim is not defined"
143
143
  return output.model_output.reshape(
144
144
  output.model_output.shape[0], -1, self.model_description.dim
145
- ).astype(np.float32)
145
+ )
146
146
 
147
147
  def _post_process_onnx_text_output(
148
148
  self,
@@ -157,7 +157,7 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
157
157
  Returns:
158
158
  Iterable[NumpyArray]: Post-processed output as NumPy arrays.
159
159
  """
160
- return output.model_output.astype(np.float32)
160
+ return output.model_output
161
161
 
162
162
  def tokenize(self, documents: list[str], **kwargs: Any) -> list[Encoding]:
163
163
  texts_query: list[str] = []
@@ -0,0 +1,46 @@
1
+ from typing import Optional, Sequence, Any
2
+
3
+ from fastembed.common import OnnxProvider
4
+ from fastembed.common.model_description import BaseModelDescription
5
+ from fastembed.rerank.cross_encoder.onnx_text_cross_encoder import OnnxTextCrossEncoder
6
+
7
+
8
+ class CustomTextCrossEncoder(OnnxTextCrossEncoder):
9
+ SUPPORTED_MODELS: list[BaseModelDescription] = []
10
+
11
+ def __init__(
12
+ self,
13
+ model_name: str,
14
+ cache_dir: Optional[str] = None,
15
+ threads: Optional[int] = None,
16
+ providers: Optional[Sequence[OnnxProvider]] = None,
17
+ cuda: bool = False,
18
+ device_ids: Optional[list[int]] = None,
19
+ lazy_load: bool = False,
20
+ device_id: Optional[int] = None,
21
+ specific_model_path: Optional[str] = None,
22
+ **kwargs: Any,
23
+ ):
24
+ super().__init__(
25
+ model_name=model_name,
26
+ cache_dir=cache_dir,
27
+ threads=threads,
28
+ providers=providers,
29
+ cuda=cuda,
30
+ device_ids=device_ids,
31
+ lazy_load=lazy_load,
32
+ device_id=device_id,
33
+ specific_model_path=specific_model_path,
34
+ **kwargs,
35
+ )
36
+
37
+ @classmethod
38
+ def _list_supported_models(cls) -> list[BaseModelDescription]:
39
+ return cls.SUPPORTED_MODELS
40
+
41
+ @classmethod
42
+ def add_model(
43
+ cls,
44
+ model_description: BaseModelDescription,
45
+ ) -> None:
46
+ cls.SUPPORTED_MODELS.append(model_description)
@@ -196,7 +196,9 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
196
196
  def _get_worker_class(cls) -> Type[TextRerankerWorker]:
197
197
  return TextCrossEncoderWorker
198
198
 
199
- def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[float]:
199
+ def _post_process_onnx_output(
200
+ self, output: OnnxOutputContext, **kwargs: Any
201
+ ) -> Iterable[float]:
200
202
  return (float(elem) for elem in output.model_output)
201
203
 
202
204
 
@@ -133,7 +133,18 @@ class OnnxCrossEncoderModel(OnnxModel[float]):
133
133
  for batch in pool.ordered_map(iter_batch(pairs, batch_size), **params):
134
134
  yield from self._post_process_onnx_output(batch) # type: ignore
135
135
 
136
- def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[float]:
136
+ def _post_process_onnx_output(
137
+ self, output: OnnxOutputContext, **kwargs: Any
138
+ ) -> Iterable[float]:
139
+ """Post-process the ONNX model output to convert it into a usable format.
140
+
141
+ Args:
142
+ output (OnnxOutputContext): The raw output from the ONNX model.
143
+ **kwargs: Additional keyword arguments that may be needed by specific implementations.
144
+
145
+ Returns:
146
+ Iterable[float]: Post-processed output as an iterable of float values.
147
+ """
137
148
  raise NotImplementedError("Subclasses must implement this method")
138
149
 
139
150
  def _preprocess_onnx_input(
@@ -3,13 +3,19 @@ from dataclasses import asdict
3
3
 
4
4
  from fastembed.common import OnnxProvider
5
5
  from fastembed.rerank.cross_encoder.onnx_text_cross_encoder import OnnxTextCrossEncoder
6
+ from fastembed.rerank.cross_encoder.custom_text_cross_encoder import CustomTextCrossEncoder
7
+
6
8
  from fastembed.rerank.cross_encoder.text_cross_encoder_base import TextCrossEncoderBase
7
- from fastembed.common.model_description import BaseModelDescription
9
+ from fastembed.common.model_description import (
10
+ ModelSource,
11
+ BaseModelDescription,
12
+ )
8
13
 
9
14
 
10
15
  class TextCrossEncoder(TextCrossEncoderBase):
11
16
  CROSS_ENCODER_REGISTRY: list[Type[TextCrossEncoderBase]] = [
12
17
  OnnxTextCrossEncoder,
18
+ CustomTextCrossEncoder,
13
19
  ]
14
20
 
15
21
  @classmethod
@@ -124,3 +130,34 @@ class TextCrossEncoder(TextCrossEncoderBase):
124
130
  yield from self.model.rerank_pairs(
125
131
  pairs, batch_size=batch_size, parallel=parallel, **kwargs
126
132
  )
133
+
134
+ @classmethod
135
+ def add_custom_model(
136
+ cls,
137
+ model: str,
138
+ sources: ModelSource,
139
+ model_file: str = "onnx/model.onnx",
140
+ description: str = "",
141
+ license: str = "",
142
+ size_in_gb: float = 0.0,
143
+ additional_files: Optional[list[str]] = None,
144
+ ) -> None:
145
+ registered_models = cls._list_supported_models()
146
+ for registered_model in registered_models:
147
+ if model == registered_model.model:
148
+ raise ValueError(
149
+ f"Model {model} is already registered in CrossEncoderModel, if you still want to add this model, "
150
+ f"please use another model name"
151
+ )
152
+
153
+ CustomTextCrossEncoder.add_model(
154
+ BaseModelDescription(
155
+ model=model,
156
+ sources=sources,
157
+ model_file=model_file,
158
+ description=description,
159
+ license=license,
160
+ size_in_GB=size_in_gb,
161
+ additional_files=additional_files or [],
162
+ )
163
+ )
@@ -21,13 +21,9 @@ from fastembed.sparse.sparse_embedding_base import (
21
21
  from fastembed.sparse.utils.tokenizer import SimpleTokenizer
22
22
  from fastembed.common.model_description import SparseModelDescription, ModelSource
23
23
 
24
+
24
25
  supported_languages = [
25
26
  "arabic",
26
- "azerbaijani",
27
- "basque",
28
- "bengali",
29
- "catalan",
30
- "chinese",
31
27
  "danish",
32
28
  "dutch",
33
29
  "english",
@@ -35,21 +31,15 @@ supported_languages = [
35
31
  "french",
36
32
  "german",
37
33
  "greek",
38
- "hebrew",
39
- "hinglish",
40
34
  "hungarian",
41
- "indonesian",
42
35
  "italian",
43
- "kazakh",
44
- "nepali",
45
36
  "norwegian",
46
37
  "portuguese",
47
38
  "romanian",
48
39
  "russian",
49
- "slovene",
50
40
  "spanish",
51
41
  "swedish",
52
- "tajik",
42
+ "tamil",
53
43
  "turkish",
54
44
  ]
55
45
 
@@ -217,7 +217,9 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
217
217
 
218
218
  return new_vector
219
219
 
220
- def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[SparseEmbedding]:
220
+ def _post_process_onnx_output(
221
+ self, output: OnnxOutputContext, **kwargs: Any
222
+ ) -> Iterable[SparseEmbedding]:
221
223
  if output.input_ids is None:
222
224
  raise ValueError("input_ids must be provided for document post-processing")
223
225