fastembed-gpu 0.7.3__tar.gz → 0.8.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 (65) hide show
  1. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/PKG-INFO +18 -12
  2. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/common/model_description.py +7 -7
  3. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/common/model_management.py +35 -18
  4. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/common/onnx_model.py +57 -14
  5. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/common/preprocessor_utils.py +4 -3
  6. fastembed_gpu-0.8.0/fastembed/common/types.py +27 -0
  7. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/common/utils.py +2 -2
  8. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/embedding.py +3 -3
  9. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/image/image_embedding.py +10 -10
  10. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/image/image_embedding_base.py +6 -6
  11. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/image/onnx_embedding.py +20 -15
  12. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/image/onnx_image_model.py +23 -15
  13. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/image/transform/functional.py +81 -9
  14. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/image/transform/operators.py +243 -13
  15. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/late_interaction/colbert.py +58 -22
  16. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/late_interaction/late_interaction_embedding_base.py +16 -7
  17. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/late_interaction/late_interaction_text_embedding.py +38 -11
  18. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/late_interaction/token_embeddings.py +3 -3
  19. fastembed_gpu-0.8.0/fastembed/late_interaction_multimodal/colmodernvbert.py +532 -0
  20. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/late_interaction_multimodal/colpali.py +40 -18
  21. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/late_interaction_multimodal/late_interaction_multimodal_embedding.py +38 -13
  22. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/late_interaction_multimodal/late_interaction_multimodal_embedding_base.py +16 -8
  23. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/late_interaction_multimodal/onnx_multimodal_model.py +35 -23
  24. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/parallel_processor.py +10 -9
  25. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/postprocess/muvera.py +1 -3
  26. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/rerank/cross_encoder/custom_text_cross_encoder.py +9 -8
  27. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/rerank/cross_encoder/onnx_text_cross_encoder.py +32 -13
  28. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/rerank/cross_encoder/onnx_text_model.py +32 -12
  29. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/rerank/cross_encoder/text_cross_encoder.py +23 -8
  30. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/rerank/cross_encoder/text_cross_encoder_base.py +8 -4
  31. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/sparse/bm25.py +18 -11
  32. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/sparse/bm42.py +37 -19
  33. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/sparse/minicoil.py +39 -23
  34. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/sparse/sparse_embedding_base.py +11 -9
  35. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/sparse/sparse_text_embedding.py +24 -11
  36. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/sparse/splade_pp.py +24 -14
  37. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/sparse/utils/sparse_vectors_converter.py +19 -22
  38. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/text/custom_text_embedding.py +10 -11
  39. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/text/multitask_embedding.py +8 -10
  40. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/text/onnx_embedding.py +24 -16
  41. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/text/onnx_text_model.py +34 -15
  42. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/text/text_embedding.py +26 -12
  43. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/text/text_embedding_base.py +11 -7
  44. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/pyproject.toml +23 -13
  45. fastembed_gpu-0.7.3/fastembed/common/types.py +0 -25
  46. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/LICENSE +0 -0
  47. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/NOTICE +0 -0
  48. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/README.md +0 -0
  49. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/__init__.py +0 -0
  50. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/common/__init__.py +0 -0
  51. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/image/__init__.py +0 -0
  52. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/late_interaction/__init__.py +0 -0
  53. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/late_interaction/jina_colbert.py +0 -0
  54. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/late_interaction_multimodal/__init__.py +0 -0
  55. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/postprocess/__init__.py +0 -0
  56. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/py.typed +0 -0
  57. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/rerank/cross_encoder/__init__.py +0 -0
  58. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/sparse/__init__.py +0 -0
  59. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/sparse/utils/minicoil_encoder.py +0 -0
  60. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/sparse/utils/tokenizer.py +0 -0
  61. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/sparse/utils/vocab_resolver.py +0 -0
  62. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/text/__init__.py +0 -0
  63. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/text/clip_embedding.py +0 -0
  64. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/text/pooled_embedding.py +0 -0
  65. {fastembed_gpu-0.7.3 → fastembed_gpu-0.8.0}/fastembed/text/pooled_normalized_embedding.py +0 -0
@@ -1,30 +1,36 @@
1
- Metadata-Version: 2.3
1
+ Metadata-Version: 2.4
2
2
  Name: fastembed-gpu
3
- Version: 0.7.3
3
+ Version: 0.8.0
4
4
  Summary: Fast, light, accurate library built for retrieval embedding generation
5
5
  License: Apache License
6
+ License-File: LICENSE
7
+ License-File: NOTICE
6
8
  Keywords: vector,embedding,neural,search,qdrant,sentence-transformers
7
9
  Author: Qdrant Team
8
10
  Author-email: info@qdrant.tech
9
- Requires-Python: >=3.9.0
11
+ Requires-Python: >=3.10.0
10
12
  Classifier: License :: Other/Proprietary License
11
13
  Classifier: Programming Language :: Python :: 3
12
- Classifier: Programming Language :: Python :: 3.9
13
14
  Classifier: Programming Language :: Python :: 3.10
14
15
  Classifier: Programming Language :: Python :: 3.11
15
16
  Classifier: Programming Language :: Python :: 3.12
16
17
  Classifier: Programming Language :: Python :: 3.13
17
- Requires-Dist: huggingface-hub (>=0.20,<1.0)
18
+ Classifier: Programming Language :: Python :: 3.14
19
+ Requires-Dist: huggingface-hub (>=0.20,<2.0)
18
20
  Requires-Dist: loguru (>=0.7.2,<0.8.0)
19
21
  Requires-Dist: mmh3 (>=4.1.0,<6.0.0)
20
- Requires-Dist: numpy (>=1.21) ; python_version >= "3.10" and python_version < "3.12"
21
- Requires-Dist: numpy (>=1.21,<2.1.0) ; python_version < "3.10"
22
+ Requires-Dist: numpy (>=1.21) ; python_version == "3.11"
23
+ Requires-Dist: numpy (>=1.21,<2.3.0) ; python_version == "3.10"
22
24
  Requires-Dist: numpy (>=1.26) ; python_version == "3.12"
23
- Requires-Dist: numpy (>=2.1.0) ; python_version >= "3.13"
24
- Requires-Dist: onnxruntime-gpu (>1.20.0) ; python_version >= "3.13"
25
- Requires-Dist: onnxruntime-gpu (>=1.17.0,!=1.20.0) ; python_version >= "3.10" and python_version < "3.13"
26
- Requires-Dist: onnxruntime-gpu (>=1.17.0,<1.20.0) ; python_version < "3.10"
27
- Requires-Dist: pillow (>=10.3.0,<12.0.0)
25
+ Requires-Dist: numpy (>=2.1.0) ; python_version == "3.13"
26
+ Requires-Dist: numpy (>=2.3.0) ; python_version >= "3.14"
27
+ Requires-Dist: onnxruntime-gpu (>1.21.0,!=1.24.0,!=1.24.1) ; python_version == "3.13"
28
+ Requires-Dist: onnxruntime-gpu (>=1.17.0,!=1.20.0,!=1.24.0,!=1.24.1) ; python_version >= "3.11" and python_version < "3.13"
29
+ Requires-Dist: onnxruntime-gpu (>=1.17.0,!=1.20.0,<1.24) ; python_version == "3.10"
30
+ Requires-Dist: onnxruntime-gpu (>=1.24.2) ; python_version >= "3.14"
31
+ Requires-Dist: pillow (>=10.3.0,<13.0) ; python_version >= "3.10" and python_version < "3.13"
32
+ Requires-Dist: pillow (>=11.0.0,<13.0) ; python_version == "3.13"
33
+ Requires-Dist: pillow (>=12.0.0,<13.0) ; python_version >= "3.14"
28
34
  Requires-Dist: py-rust-stemmers (>=0.1.0,<0.2.0)
29
35
  Requires-Dist: requests (>=2.31,<3.0)
30
36
  Requires-Dist: tokenizers (>=0.15,<1.0)
@@ -1,12 +1,12 @@
1
1
  from dataclasses import dataclass, field
2
2
  from enum import Enum
3
- from typing import Optional, Any
3
+ from typing import Any
4
4
 
5
5
 
6
6
  @dataclass(frozen=True)
7
7
  class ModelSource:
8
- hf: Optional[str] = None
9
- url: Optional[str] = None
8
+ hf: str | None = None
9
+ url: str | None = None
10
10
  _deprecated_tar_struct: bool = False
11
11
 
12
12
  @property
@@ -33,8 +33,8 @@ class BaseModelDescription:
33
33
 
34
34
  @dataclass(frozen=True)
35
35
  class DenseModelDescription(BaseModelDescription):
36
- dim: Optional[int] = None
37
- tasks: Optional[dict[str, Any]] = field(default_factory=dict)
36
+ dim: int | None = None
37
+ tasks: dict[str, Any] | None = field(default_factory=dict)
38
38
 
39
39
  def __post_init__(self) -> None:
40
40
  assert self.dim is not None, "dim is required for dense model description"
@@ -42,8 +42,8 @@ class DenseModelDescription(BaseModelDescription):
42
42
 
43
43
  @dataclass(frozen=True)
44
44
  class SparseModelDescription(BaseModelDescription):
45
- requires_idf: Optional[bool] = None
46
- vocab_size: Optional[int] = None
45
+ requires_idf: bool | None = None
46
+ vocab_size: int | None = None
47
47
 
48
48
 
49
49
  class PoolingType(str, Enum):
@@ -3,8 +3,9 @@ import time
3
3
  import json
4
4
  import shutil
5
5
  import tarfile
6
+ from copy import deepcopy
6
7
  from pathlib import Path
7
- from typing import Any, Optional, Union, TypeVar, Generic
8
+ from typing import Any, TypeVar, Generic
8
9
 
9
10
  import requests
10
11
  from huggingface_hub import snapshot_download, model_info, list_repo_tree
@@ -179,8 +180,8 @@ class ModelManagement(Generic[T]):
179
180
 
180
181
  def _collect_file_metadata(
181
182
  model_dir: Path, repo_files: list[RepoFile]
182
- ) -> dict[str, dict[str, Union[int, str]]]:
183
- meta: dict[str, dict[str, Union[int, str]]] = {}
183
+ ) -> dict[str, dict[str, int | str]]:
184
+ meta: dict[str, dict[str, int | str]] = {}
184
185
  file_info_map = {f.path: f for f in repo_files}
185
186
  for file_path in model_dir.rglob("*"):
186
187
  if file_path.is_file() and file_path.name != cls.METADATA_FILE:
@@ -192,9 +193,7 @@ class ModelManagement(Generic[T]):
192
193
  }
193
194
  return meta
194
195
 
195
- def _save_file_metadata(
196
- model_dir: Path, meta: dict[str, dict[str, Union[int, str]]]
197
- ) -> None:
196
+ def _save_file_metadata(model_dir: Path, meta: dict[str, dict[str, int | str]]) -> None:
198
197
  try:
199
198
  if not model_dir.exists():
200
199
  model_dir.mkdir(parents=True, exist_ok=True)
@@ -224,11 +223,6 @@ class ModelManagement(Generic[T]):
224
223
  logger.warning(
225
224
  "Local file sizes do not match the metadata."
226
225
  ) # do not raise, still make an attempt to load the model
227
- else:
228
- logger.warning(
229
- "Metadata file not found. Proceeding without checking local files."
230
- ) # if users have downloaded models from hf manually, or they're updating from previous versions of
231
- # fastembed
232
226
  result = snapshot_download(
233
227
  repo_id=hf_source_repo,
234
228
  allow_patterns=allow_patterns,
@@ -401,21 +395,43 @@ class ModelManagement(Generic[T]):
401
395
  Path: The path to the downloaded model directory.
402
396
  """
403
397
  local_files_only = kwargs.get("local_files_only", False)
404
- specific_model_path: Optional[str] = kwargs.pop("specific_model_path", None)
398
+ hf_offline = os.environ.get("HF_HUB_OFFLINE", "").strip().upper()
399
+ if not local_files_only and hf_offline in {"1", "TRUE", "YES", "ON"}:
400
+ local_files_only = True
401
+ kwargs["local_files_only"] = True
402
+ specific_model_path: str | None = kwargs.pop("specific_model_path", None)
405
403
  if specific_model_path:
406
404
  return Path(specific_model_path)
407
405
  retries = 1 if local_files_only else retries
408
406
  hf_source = model.sources.hf
409
407
  url_source = model.sources.url
410
408
 
409
+ extra_patterns = [model.model_file]
410
+ extra_patterns.extend(model.additional_files)
411
+
412
+ if hf_source:
413
+ try:
414
+ cache_kwargs = deepcopy(kwargs)
415
+ cache_kwargs["local_files_only"] = True
416
+ return Path(
417
+ cls.download_files_from_huggingface(
418
+ hf_source,
419
+ cache_dir=cache_dir,
420
+ extra_patterns=extra_patterns,
421
+ **cache_kwargs,
422
+ )
423
+ )
424
+ except Exception:
425
+ pass
426
+ finally:
427
+ enable_progress_bars()
428
+
411
429
  sleep = 3.0
412
430
  while retries > 0:
413
431
  retries -= 1
414
432
 
415
- if hf_source:
416
- extra_patterns = [model.model_file]
417
- extra_patterns.extend(model.additional_files)
418
-
433
+ if hf_source and not local_files_only:
434
+ # we have already tried loading with `local_files_only=True` via hf and we failed
419
435
  try:
420
436
  return Path(
421
437
  cls.download_files_from_huggingface(
@@ -448,11 +464,12 @@ class ModelManagement(Generic[T]):
448
464
 
449
465
  if local_files_only:
450
466
  logger.error("Could not find model in cache_dir")
467
+ break
451
468
  else:
452
469
  logger.error(
453
470
  f"Could not download model from either source, sleeping for {sleep} seconds, {retries} retries left."
454
471
  )
455
- time.sleep(sleep)
456
- sleep *= 3
472
+ time.sleep(sleep)
473
+ sleep *= 3
457
474
 
458
475
  raise ValueError(f"Could not load model {model.model} from any source.")
@@ -1,7 +1,7 @@
1
1
  import warnings
2
2
  from dataclasses import dataclass
3
3
  from pathlib import Path
4
- from typing import Any, Generic, Iterable, Optional, Sequence, Type, TypeVar
4
+ from typing import Any, Generic, Iterable, Sequence, Type, TypeVar
5
5
 
6
6
  import numpy as np
7
7
  import onnxruntime as ort
@@ -9,7 +9,7 @@ import onnxruntime as ort
9
9
  from numpy.typing import NDArray
10
10
  from tokenizers import Tokenizer
11
11
 
12
- from fastembed.common.types import OnnxProvider, NumpyArray
12
+ from fastembed.common.types import OnnxProvider, NumpyArray, Device
13
13
  from fastembed.parallel_processor import Worker
14
14
 
15
15
  # Holds type of the embedding result
@@ -19,11 +19,14 @@ T = TypeVar("T")
19
19
  @dataclass
20
20
  class OnnxOutputContext:
21
21
  model_output: NumpyArray
22
- attention_mask: Optional[NDArray[np.int64]] = None
23
- input_ids: Optional[NDArray[np.int64]] = None
22
+ attention_mask: NDArray[np.int64] | None = None
23
+ input_ids: NDArray[np.int64] | None = None
24
+ metadata: dict[str, Any] | None = None
24
25
 
25
26
 
26
27
  class OnnxModel(Generic[T]):
28
+ EXPOSED_SESSION_OPTIONS = ("enable_cpu_mem_arena",)
29
+
27
30
  @classmethod
28
31
  def _get_worker_class(cls) -> Type["EmbeddingWorker[T]"]:
29
32
  raise NotImplementedError("Subclasses must implement this method")
@@ -41,8 +44,8 @@ class OnnxModel(Generic[T]):
41
44
  raise NotImplementedError("Subclasses must implement this method")
42
45
 
43
46
  def __init__(self) -> None:
44
- self.model: Optional[ort.InferenceSession] = None
45
- self.tokenizer: Optional[Tokenizer] = None
47
+ self.model: ort.InferenceSession | None = None
48
+ self.tokenizer: Tokenizer | None = None
46
49
 
47
50
  def _preprocess_onnx_input(
48
51
  self, onnx_input: dict[str, NumpyArray], **kwargs: Any
@@ -56,24 +59,30 @@ class OnnxModel(Generic[T]):
56
59
  self,
57
60
  model_dir: Path,
58
61
  model_file: str,
59
- threads: Optional[int],
60
- providers: Optional[Sequence[OnnxProvider]] = None,
61
- cuda: bool = False,
62
- device_id: Optional[int] = None,
62
+ threads: int | None,
63
+ providers: Sequence[OnnxProvider] | None = None,
64
+ cuda: bool | Device = Device.AUTO,
65
+ device_id: int | None = None,
66
+ extra_session_options: dict[str, Any] | None = None,
63
67
  ) -> None:
64
68
  model_path = model_dir / model_file
65
69
  # List of Execution Providers: https://onnxruntime.ai/docs/execution-providers
70
+ available_providers = ort.get_available_providers()
71
+ cuda_available = "CUDAExecutionProvider" in available_providers
72
+ explicit_cuda = cuda is True or cuda == Device.CUDA
66
73
 
67
- if cuda and providers is not None:
74
+ if explicit_cuda and providers is not None:
68
75
  warnings.warn(
69
- f"`cuda` and `providers` are mutually exclusive parameters, cuda: {cuda}, providers: {providers}",
76
+ f"`cuda` and `providers` are mutually exclusive parameters, "
77
+ f"cuda: {cuda}, providers: {providers}. If you'd like to use providers, cuda should be one of "
78
+ f"[False, Device.CPU, Device.AUTO].",
70
79
  category=UserWarning,
71
80
  stacklevel=6,
72
81
  )
73
82
 
74
83
  if providers is not None:
75
84
  onnx_providers = list(providers)
76
- elif cuda:
85
+ elif explicit_cuda or (cuda == Device.AUTO and cuda_available):
77
86
  if device_id is None:
78
87
  onnx_providers = ["CUDAExecutionProvider"]
79
88
  else:
@@ -81,7 +90,6 @@ class OnnxModel(Generic[T]):
81
90
  else:
82
91
  onnx_providers = ["CPUExecutionProvider"]
83
92
 
84
- available_providers = ort.get_available_providers()
85
93
  requested_provider_names: list[str] = []
86
94
  for provider in onnx_providers:
87
95
  # check providers available
@@ -99,6 +107,9 @@ class OnnxModel(Generic[T]):
99
107
  so.intra_op_num_threads = threads
100
108
  so.inter_op_num_threads = threads
101
109
 
110
+ if extra_session_options is not None:
111
+ self.add_extra_session_options(so, extra_session_options)
112
+
102
113
  self.model = ort.InferenceSession(
103
114
  str(model_path), providers=onnx_providers, sess_options=so
104
115
  )
@@ -113,6 +124,38 @@ class OnnxModel(Generic[T]):
113
124
  RuntimeWarning,
114
125
  )
115
126
 
127
+ @classmethod
128
+ def _select_exposed_session_options(cls, model_kwargs: dict[str, Any]) -> dict[str, Any]:
129
+ """A convenience method to select the exposed session options in models
130
+
131
+ Args:
132
+ model_kwargs (dict[str, Any]): The model kwargs.
133
+
134
+ Returns:
135
+ dict[str, Any]: a dict with filtered exposed session options.
136
+ """
137
+ return {k: v for k, v in model_kwargs.items() if k in cls.EXPOSED_SESSION_OPTIONS}
138
+
139
+ @classmethod
140
+ def add_extra_session_options(
141
+ cls, session_options: ort.SessionOptions, extra_options: dict[str, Any]
142
+ ) -> None:
143
+ """Add extra session options to the existing options object in-place
144
+
145
+ Args:
146
+ session_options (ort.SessionOptions): The existing session options object.
147
+ extra_options (dict[str, Any]): The extra session options available in cls.EXPOSED_SESSION_OPTIONS.
148
+
149
+ Returns:
150
+ None
151
+ """
152
+ for option in extra_options:
153
+ assert (
154
+ option in cls.EXPOSED_SESSION_OPTIONS
155
+ ), f"{option} is unknown or not exposed (exposed options: {cls.EXPOSED_SESSION_OPTIONS})"
156
+ if "enable_cpu_mem_arena" in extra_options:
157
+ session_options.enable_cpu_mem_arena = extra_options["enable_cpu_mem_arena"]
158
+
116
159
  def load_onnx_model(self) -> None:
117
160
  raise NotImplementedError("Subclasses must implement this method")
118
161
 
@@ -50,9 +50,10 @@ def load_tokenizer(model_dir: Path) -> tuple[Tokenizer, dict[str, int]]:
50
50
 
51
51
  tokenizer = Tokenizer.from_file(str(tokenizer_path))
52
52
  tokenizer.enable_truncation(max_length=max_context)
53
- tokenizer.enable_padding(
54
- pad_id=config.get("pad_token_id", 0), pad_token=tokenizer_config["pad_token"]
55
- )
53
+ if not tokenizer.padding:
54
+ tokenizer.enable_padding(
55
+ pad_id=config.get("pad_token_id", 0), pad_token=tokenizer_config["pad_token"]
56
+ )
56
57
 
57
58
  for token in tokens_map.values():
58
59
  if isinstance(token, str):
@@ -0,0 +1,27 @@
1
+ from enum import Enum
2
+ from pathlib import Path
3
+ from typing import Any, TypeAlias
4
+
5
+ import numpy as np
6
+ from numpy.typing import NDArray
7
+ from PIL import Image
8
+
9
+
10
+ class Device(str, Enum):
11
+ CPU = "cpu"
12
+ CUDA = "cuda"
13
+ AUTO = "auto"
14
+
15
+
16
+ PathInput: TypeAlias = str | Path
17
+ ImageInput: TypeAlias = PathInput | Image.Image
18
+
19
+ OnnxProvider: TypeAlias = str | tuple[str, dict[Any, Any]]
20
+ NumpyArray: TypeAlias = (
21
+ NDArray[np.float64]
22
+ | NDArray[np.float32]
23
+ | NDArray[np.float16]
24
+ | NDArray[np.int8]
25
+ | NDArray[np.int64]
26
+ | NDArray[np.int32]
27
+ )
@@ -5,7 +5,7 @@ import tempfile
5
5
  import unicodedata
6
6
  from pathlib import Path
7
7
  from itertools import islice
8
- from typing import Iterable, Optional, TypeVar
8
+ from typing import Iterable, TypeVar
9
9
 
10
10
  import numpy as np
11
11
  from numpy.typing import NDArray
@@ -45,7 +45,7 @@ def iter_batch(iterable: Iterable[T], size: int) -> Iterable[list[T]]:
45
45
  yield b
46
46
 
47
47
 
48
- def define_cache_dir(cache_dir: Optional[str] = None) -> Path:
48
+ def define_cache_dir(cache_dir: str | None = None) -> Path:
49
49
  """
50
50
  Define the cache directory for fastembed
51
51
  """
@@ -1,4 +1,4 @@
1
- from typing import Optional, Any
1
+ from typing import Any
2
2
 
3
3
  from loguru import logger
4
4
 
@@ -17,8 +17,8 @@ class JinaEmbedding(TextEmbedding):
17
17
  def __init__(
18
18
  self,
19
19
  model_name: str = "jinaai/jina-embeddings-v2-base-en",
20
- cache_dir: Optional[str] = None,
21
- threads: Optional[int] = None,
20
+ cache_dir: str | None = None,
21
+ threads: int | None = None,
22
22
  **kwargs: Any,
23
23
  ):
24
24
  super().__init__(model_name, cache_dir, threads, **kwargs)
@@ -1,7 +1,7 @@
1
- from typing import Any, Iterable, Optional, Sequence, Type, Union
1
+ from typing import Any, Iterable, Sequence, Type
2
2
  from dataclasses import asdict
3
3
 
4
- from fastembed.common.types import NumpyArray
4
+ from fastembed.common.types import NumpyArray, Device
5
5
  from fastembed.common import ImageInput, OnnxProvider
6
6
  from fastembed.image.image_embedding_base import ImageEmbeddingBase
7
7
  from fastembed.image.onnx_embedding import OnnxImageEmbedding
@@ -48,11 +48,11 @@ class ImageEmbedding(ImageEmbeddingBase):
48
48
  def __init__(
49
49
  self,
50
50
  model_name: str,
51
- cache_dir: Optional[str] = None,
52
- threads: Optional[int] = None,
53
- providers: Optional[Sequence[OnnxProvider]] = None,
54
- cuda: bool = False,
55
- device_ids: Optional[list[int]] = None,
51
+ cache_dir: str | None = None,
52
+ threads: int | None = None,
53
+ providers: Sequence[OnnxProvider] | None = None,
54
+ cuda: bool | Device = Device.AUTO,
55
+ device_ids: list[int] | None = None,
56
56
  lazy_load: bool = False,
57
57
  **kwargs: Any,
58
58
  ):
@@ -98,7 +98,7 @@ class ImageEmbedding(ImageEmbeddingBase):
98
98
  ValueError: If the model name is not found in the supported models.
99
99
  """
100
100
  descriptions = cls._list_supported_models()
101
- embedding_size: Optional[int] = None
101
+ embedding_size: int | None = None
102
102
  for description in descriptions:
103
103
  if description.model.lower() == model_name.lower():
104
104
  embedding_size = description.dim
@@ -113,9 +113,9 @@ class ImageEmbedding(ImageEmbeddingBase):
113
113
 
114
114
  def embed(
115
115
  self,
116
- images: Union[ImageInput, Iterable[ImageInput]],
116
+ images: ImageInput | Iterable[ImageInput],
117
117
  batch_size: int = 16,
118
- parallel: Optional[int] = None,
118
+ parallel: int | None = None,
119
119
  **kwargs: Any,
120
120
  ) -> Iterable[NumpyArray]:
121
121
  """
@@ -1,4 +1,4 @@
1
- from typing import Iterable, Optional, Any, Union
1
+ from typing import Iterable, Any
2
2
 
3
3
  from fastembed.common.model_description import DenseModelDescription
4
4
  from fastembed.common.types import NumpyArray
@@ -10,21 +10,21 @@ class ImageEmbeddingBase(ModelManagement[DenseModelDescription]):
10
10
  def __init__(
11
11
  self,
12
12
  model_name: str,
13
- cache_dir: Optional[str] = None,
14
- threads: Optional[int] = None,
13
+ cache_dir: str | None = None,
14
+ threads: int | None = None,
15
15
  **kwargs: Any,
16
16
  ):
17
17
  self.model_name = model_name
18
18
  self.cache_dir = cache_dir
19
19
  self.threads = threads
20
20
  self._local_files_only = kwargs.pop("local_files_only", False)
21
- self._embedding_size: Optional[int] = None
21
+ self._embedding_size: int | None = None
22
22
 
23
23
  def embed(
24
24
  self,
25
- images: Union[ImageInput, Iterable[ImageInput]],
25
+ images: ImageInput | Iterable[ImageInput],
26
26
  batch_size: int = 16,
27
- parallel: Optional[int] = None,
27
+ parallel: int | None = None,
28
28
  **kwargs: Any,
29
29
  ) -> Iterable[NumpyArray]:
30
30
  """
@@ -1,7 +1,7 @@
1
- from typing import Any, Iterable, Optional, Sequence, Type, Union
1
+ from typing import Any, Iterable, Sequence, Type
2
2
 
3
3
 
4
- from fastembed.common.types import NumpyArray
4
+ from fastembed.common.types import NumpyArray, Device
5
5
  from fastembed.common import ImageInput, OnnxProvider
6
6
  from fastembed.common.onnx_model import OnnxOutputContext
7
7
  from fastembed.common.utils import define_cache_dir, normalize
@@ -63,14 +63,15 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
63
63
  def __init__(
64
64
  self,
65
65
  model_name: str,
66
- cache_dir: Optional[str] = None,
67
- threads: Optional[int] = None,
68
- providers: Optional[Sequence[OnnxProvider]] = None,
69
- cuda: bool = False,
70
- device_ids: Optional[list[int]] = None,
66
+
67
+ cache_dir: str | None = None,
68
+ threads: int | None = None,
69
+ providers: Sequence[OnnxProvider] | None = None,
70
+ cuda: bool | Device = Device.AUTO,
71
+ device_ids: list[int] | None = None,
71
72
  lazy_load: bool = False,
72
- device_id: Optional[int] = None,
73
- specific_model_path: Optional[str] = None,
73
+ device_id: int | None = None,
74
+ specific_model_path: str | None = None,
74
75
  **kwargs: Any,
75
76
  ):
76
77
  """
@@ -82,10 +83,11 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
82
83
  threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.
83
84
  providers (Optional[Sequence[OnnxProvider]], optional): The list of onnxruntime providers to use.
84
85
  Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
85
- cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
86
- Defaults to False.
86
+ cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`
87
+ Defaults to Device.AUTO.
87
88
  device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
88
- workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
89
+ workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, mutually exclusive
90
+ with `providers`. Defaults to None.
89
91
  lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
90
92
  Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
91
93
  device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
@@ -98,13 +100,14 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
98
100
  super().__init__(model_name, cache_dir, threads, **kwargs)
99
101
  self.providers = providers
100
102
  self.lazy_load = lazy_load
103
+ self._extra_session_options = self._select_exposed_session_options(kwargs)
101
104
 
102
105
  # List of device ids, that can be used for data parallel processing in workers
103
106
  self.device_ids = device_ids
104
107
  self.cuda = cuda
105
108
 
106
109
  # This device_id will be used if we need to load model in current process
107
- self.device_id: Optional[int] = None
110
+ self.device_id: int | None = None
108
111
  if device_id is not None:
109
112
  self.device_id = device_id
110
113
  elif self.device_ids is not None:
@@ -134,6 +137,7 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
134
137
  providers=self.providers,
135
138
  cuda=self.cuda,
136
139
  device_id=self.device_id,
140
+ extra_session_options=self._extra_session_options,
137
141
  )
138
142
 
139
143
  @classmethod
@@ -148,9 +152,9 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
148
152
 
149
153
  def embed(
150
154
  self,
151
- images: Union[ImageInput, Iterable[ImageInput]],
155
+ images: ImageInput | Iterable[ImageInput],
152
156
  batch_size: int = 16,
153
- parallel: Optional[int] = None,
157
+ parallel: int | None = None,
154
158
  **kwargs: Any,
155
159
  ) -> Iterable[NumpyArray]:
156
160
  """
@@ -180,6 +184,7 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
180
184
  device_ids=self.device_ids,
181
185
  local_files_only=self._local_files_only,
182
186
  specific_model_path=self._specific_model_path,
187
+ extra_session_options=self._extra_session_options,
183
188
  **kwargs,
184
189
  )
185
190