fastembed-gpu 0.7.4__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.4 → fastembed_gpu-0.8.0}/PKG-INFO +9 -10
  2. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/common/model_description.py +7 -7
  3. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/common/model_management.py +9 -7
  4. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/common/onnx_model.py +20 -15
  5. {fastembed_gpu-0.7.4 → 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.4 → fastembed_gpu-0.8.0}/fastembed/common/utils.py +2 -2
  8. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/embedding.py +3 -3
  9. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/image/image_embedding.py +10 -10
  10. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/image/image_embedding_base.py +6 -6
  11. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/image/onnx_embedding.py +17 -15
  12. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/image/onnx_image_model.py +19 -17
  13. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/image/transform/functional.py +81 -9
  14. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/image/transform/operators.py +243 -13
  15. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/late_interaction/colbert.py +23 -22
  16. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/late_interaction/late_interaction_embedding_base.py +8 -8
  17. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/late_interaction/late_interaction_text_embedding.py +12 -12
  18. {fastembed_gpu-0.7.4 → 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.4 → fastembed_gpu-0.8.0}/fastembed/late_interaction_multimodal/colpali.py +19 -18
  21. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/late_interaction_multimodal/late_interaction_multimodal_embedding.py +18 -14
  22. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/late_interaction_multimodal/late_interaction_multimodal_embedding_base.py +9 -9
  23. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/late_interaction_multimodal/onnx_multimodal_model.py +28 -26
  24. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/parallel_processor.py +10 -9
  25. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/postprocess/muvera.py +1 -3
  26. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/rerank/cross_encoder/custom_text_cross_encoder.py +9 -8
  27. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/rerank/cross_encoder/onnx_text_cross_encoder.py +15 -13
  28. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/rerank/cross_encoder/onnx_text_model.py +14 -14
  29. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/rerank/cross_encoder/text_cross_encoder.py +9 -8
  30. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/rerank/cross_encoder/text_cross_encoder_base.py +4 -4
  31. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/sparse/bm25.py +10 -12
  32. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/sparse/bm42.py +18 -18
  33. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/sparse/minicoil.py +22 -22
  34. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/sparse/sparse_embedding_base.py +8 -10
  35. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/sparse/sparse_text_embedding.py +11 -12
  36. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/sparse/splade_pp.py +17 -15
  37. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/sparse/utils/sparse_vectors_converter.py +19 -22
  38. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/text/custom_text_embedding.py +10 -11
  39. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/text/multitask_embedding.py +8 -10
  40. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/text/onnx_embedding.py +17 -16
  41. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/text/onnx_text_model.py +18 -20
  42. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/text/text_embedding.py +13 -13
  43. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/text/text_embedding_base.py +8 -8
  44. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/pyproject.toml +16 -15
  45. fastembed_gpu-0.7.4/fastembed/common/types.py +0 -25
  46. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/LICENSE +0 -0
  47. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/NOTICE +0 -0
  48. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/README.md +0 -0
  49. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/__init__.py +0 -0
  50. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/common/__init__.py +0 -0
  51. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/image/__init__.py +0 -0
  52. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/late_interaction/__init__.py +0 -0
  53. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/late_interaction/jina_colbert.py +0 -0
  54. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/late_interaction_multimodal/__init__.py +0 -0
  55. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/postprocess/__init__.py +0 -0
  56. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/py.typed +0 -0
  57. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/rerank/cross_encoder/__init__.py +0 -0
  58. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/sparse/__init__.py +0 -0
  59. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/sparse/utils/minicoil_encoder.py +0 -0
  60. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/sparse/utils/tokenizer.py +0 -0
  61. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/sparse/utils/vocab_resolver.py +0 -0
  62. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/text/__init__.py +0 -0
  63. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/text/clip_embedding.py +0 -0
  64. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/text/pooled_embedding.py +0 -0
  65. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.0}/fastembed/text/pooled_normalized_embedding.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: fastembed-gpu
3
- Version: 0.7.4
3
+ Version: 0.8.0
4
4
  Summary: Fast, light, accurate library built for retrieval embedding generation
5
5
  License: Apache License
6
6
  License-File: LICENSE
@@ -8,10 +8,9 @@ License-File: NOTICE
8
8
  Keywords: vector,embedding,neural,search,qdrant,sentence-transformers
9
9
  Author: Qdrant Team
10
10
  Author-email: info@qdrant.tech
11
- Requires-Python: >=3.9.0
11
+ Requires-Python: >=3.10.0
12
12
  Classifier: License :: Other/Proprietary License
13
13
  Classifier: Programming Language :: Python :: 3
14
- Classifier: Programming Language :: Python :: 3.9
15
14
  Classifier: Programming Language :: Python :: 3.10
16
15
  Classifier: Programming Language :: Python :: 3.11
17
16
  Classifier: Programming Language :: Python :: 3.12
@@ -21,17 +20,17 @@ Requires-Dist: huggingface-hub (>=0.20,<2.0)
21
20
  Requires-Dist: loguru (>=0.7.2,<0.8.0)
22
21
  Requires-Dist: mmh3 (>=4.1.0,<6.0.0)
23
22
  Requires-Dist: numpy (>=1.21) ; python_version == "3.11"
24
- Requires-Dist: numpy (>=1.21,<2.1.0) ; python_version < "3.10"
25
23
  Requires-Dist: numpy (>=1.21,<2.3.0) ; python_version == "3.10"
26
24
  Requires-Dist: numpy (>=1.26) ; python_version == "3.12"
27
25
  Requires-Dist: numpy (>=2.1.0) ; python_version == "3.13"
28
26
  Requires-Dist: numpy (>=2.3.0) ; python_version >= "3.14"
29
- Requires-Dist: onnxruntime-gpu (>1.20.0) ; python_version >= "3.13"
30
- Requires-Dist: onnxruntime-gpu (>=1.17.0,!=1.20.0) ; python_version >= "3.10" and python_version < "3.13"
31
- Requires-Dist: onnxruntime-gpu (>=1.17.0,<1.20.0) ; python_version < "3.10"
32
- Requires-Dist: pillow (>=10.3.0,<11.0) ; python_version < "3.10"
33
- Requires-Dist: pillow (>=10.3.0,<12.0) ; python_version >= "3.10" and python_version < "3.13"
34
- Requires-Dist: pillow (>=11.0.0,<12.0) ; python_version >= "3.13"
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"
35
34
  Requires-Dist: py-rust-stemmers (>=0.1.0,<0.2.0)
36
35
  Requires-Dist: requests (>=2.31,<3.0)
37
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):
@@ -5,7 +5,7 @@ import shutil
5
5
  import tarfile
6
6
  from copy import deepcopy
7
7
  from pathlib import Path
8
- from typing import Any, Optional, Union, TypeVar, Generic
8
+ from typing import Any, TypeVar, Generic
9
9
 
10
10
  import requests
11
11
  from huggingface_hub import snapshot_download, model_info, list_repo_tree
@@ -180,8 +180,8 @@ class ModelManagement(Generic[T]):
180
180
 
181
181
  def _collect_file_metadata(
182
182
  model_dir: Path, repo_files: list[RepoFile]
183
- ) -> dict[str, dict[str, Union[int, str]]]:
184
- meta: dict[str, dict[str, Union[int, str]]] = {}
183
+ ) -> dict[str, dict[str, int | str]]:
184
+ meta: dict[str, dict[str, int | str]] = {}
185
185
  file_info_map = {f.path: f for f in repo_files}
186
186
  for file_path in model_dir.rglob("*"):
187
187
  if file_path.is_file() and file_path.name != cls.METADATA_FILE:
@@ -193,9 +193,7 @@ class ModelManagement(Generic[T]):
193
193
  }
194
194
  return meta
195
195
 
196
- def _save_file_metadata(
197
- model_dir: Path, meta: dict[str, dict[str, Union[int, str]]]
198
- ) -> None:
196
+ def _save_file_metadata(model_dir: Path, meta: dict[str, dict[str, int | str]]) -> None:
199
197
  try:
200
198
  if not model_dir.exists():
201
199
  model_dir.mkdir(parents=True, exist_ok=True)
@@ -397,7 +395,11 @@ class ModelManagement(Generic[T]):
397
395
  Path: The path to the downloaded model directory.
398
396
  """
399
397
  local_files_only = kwargs.get("local_files_only", False)
400
- 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)
401
403
  if specific_model_path:
402
404
  return Path(specific_model_path)
403
405
  retries = 1 if local_files_only else retries
@@ -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,8 +19,9 @@ 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]):
@@ -43,8 +44,8 @@ class OnnxModel(Generic[T]):
43
44
  raise NotImplementedError("Subclasses must implement this method")
44
45
 
45
46
  def __init__(self) -> None:
46
- self.model: Optional[ort.InferenceSession] = None
47
- self.tokenizer: Optional[Tokenizer] = None
47
+ self.model: ort.InferenceSession | None = None
48
+ self.tokenizer: Tokenizer | None = None
48
49
 
49
50
  def _preprocess_onnx_input(
50
51
  self, onnx_input: dict[str, NumpyArray], **kwargs: Any
@@ -58,25 +59,30 @@ class OnnxModel(Generic[T]):
58
59
  self,
59
60
  model_dir: Path,
60
61
  model_file: str,
61
- threads: Optional[int],
62
- providers: Optional[Sequence[OnnxProvider]] = None,
63
- cuda: bool = False,
64
- device_id: Optional[int] = None,
65
- extra_session_options: Optional[dict[str, Any]] = 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,
66
67
  ) -> None:
67
68
  model_path = model_dir / model_file
68
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
69
73
 
70
- if cuda and providers is not None:
74
+ if explicit_cuda and providers is not None:
71
75
  warnings.warn(
72
- 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].",
73
79
  category=UserWarning,
74
80
  stacklevel=6,
75
81
  )
76
82
 
77
83
  if providers is not None:
78
84
  onnx_providers = list(providers)
79
- elif cuda:
85
+ elif explicit_cuda or (cuda == Device.AUTO and cuda_available):
80
86
  if device_id is None:
81
87
  onnx_providers = ["CUDAExecutionProvider"]
82
88
  else:
@@ -84,7 +90,6 @@ class OnnxModel(Generic[T]):
84
90
  else:
85
91
  onnx_providers = ["CPUExecutionProvider"]
86
92
 
87
- available_providers = ort.get_available_providers()
88
93
  requested_provider_names: list[str] = []
89
94
  for provider in onnx_providers:
90
95
  # check providers available
@@ -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.
@@ -105,7 +107,7 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
105
107
  self.cuda = cuda
106
108
 
107
109
  # This device_id will be used if we need to load model in current process
108
- self.device_id: Optional[int] = None
110
+ self.device_id: int | None = None
109
111
  if device_id is not None:
110
112
  self.device_id = device_id
111
113
  elif self.device_ids is not None:
@@ -150,9 +152,9 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
150
152
 
151
153
  def embed(
152
154
  self,
153
- images: Union[ImageInput, Iterable[ImageInput]],
155
+ images: ImageInput | Iterable[ImageInput],
154
156
  batch_size: int = 16,
155
- parallel: Optional[int] = None,
157
+ parallel: int | None = None,
156
158
  **kwargs: Any,
157
159
  ) -> Iterable[NumpyArray]:
158
160
  """
@@ -2,13 +2,13 @@ import contextlib
2
2
  import os
3
3
  from multiprocessing import get_all_start_methods
4
4
  from pathlib import Path
5
- from typing import Any, Iterable, Optional, Sequence, Type, Union
5
+ from typing import Any, Iterable, Sequence, Type
6
6
 
7
7
  import numpy as np
8
8
  from PIL import Image
9
9
 
10
10
  from fastembed.image.transform.operators import Compose
11
- from fastembed.common.types import NumpyArray
11
+ from fastembed.common.types import NumpyArray, Device
12
12
  from fastembed.common import ImageInput, OnnxProvider
13
13
  from fastembed.common.onnx_model import EmbeddingWorker, OnnxModel, OnnxOutputContext, T
14
14
  from fastembed.common.preprocessor_utils import load_preprocessor
@@ -37,7 +37,7 @@ class OnnxImageModel(OnnxModel[T]):
37
37
 
38
38
  def __init__(self) -> None:
39
39
  super().__init__()
40
- self.processor: Optional[Compose] = None
40
+ self.processor: Compose | None = None
41
41
 
42
42
  def _preprocess_onnx_input(
43
43
  self, onnx_input: dict[str, NumpyArray], **kwargs: Any
@@ -51,11 +51,11 @@ class OnnxImageModel(OnnxModel[T]):
51
51
  self,
52
52
  model_dir: Path,
53
53
  model_file: str,
54
- threads: Optional[int],
55
- providers: Optional[Sequence[OnnxProvider]] = None,
56
- cuda: bool = False,
57
- device_id: Optional[int] = None,
58
- extra_session_options: Optional[dict[str, Any]] = None,
54
+ threads: int | None,
55
+ providers: Sequence[OnnxProvider] | None = None,
56
+ cuda: bool | Device = Device.AUTO,
57
+ device_id: int | None = None,
58
+ extra_session_options: dict[str, Any] | None = None,
59
59
  ) -> None:
60
60
  super()._load_onnx_model(
61
61
  model_dir=model_dir,
@@ -76,9 +76,11 @@ class OnnxImageModel(OnnxModel[T]):
76
76
  return {input_name: encoded}
77
77
 
78
78
  def onnx_embed(self, images: list[ImageInput], **kwargs: Any) -> OnnxOutputContext:
79
- with contextlib.ExitStack():
79
+ with contextlib.ExitStack() as stack:
80
80
  image_files = [
81
- Image.open(image) if not isinstance(image, Image.Image) else image
81
+ stack.enter_context(Image.open(image))
82
+ if not isinstance(image, Image.Image)
83
+ else image
82
84
  for image in images
83
85
  ]
84
86
  assert self.processor is not None, "Processor is not initialized"
@@ -93,15 +95,15 @@ class OnnxImageModel(OnnxModel[T]):
93
95
  self,
94
96
  model_name: str,
95
97
  cache_dir: str,
96
- images: Union[ImageInput, Iterable[ImageInput]],
98
+ images: ImageInput | Iterable[ImageInput],
97
99
  batch_size: int = 256,
98
- parallel: Optional[int] = None,
99
- providers: Optional[Sequence[OnnxProvider]] = None,
100
- cuda: bool = False,
101
- device_ids: Optional[list[int]] = None,
100
+ parallel: int | None = None,
101
+ providers: Sequence[OnnxProvider] | None = None,
102
+ cuda: bool | Device = Device.AUTO,
103
+ device_ids: list[int] | None = None,
102
104
  local_files_only: bool = False,
103
- specific_model_path: Optional[str] = None,
104
- extra_session_options: Optional[dict[str, Any]] = None,
105
+ specific_model_path: str | None = None,
106
+ extra_session_options: dict[str, Any] | None = None,
105
107
  **kwargs: Any,
106
108
  ) -> Iterable[T]:
107
109
  is_small = False
@@ -1,5 +1,3 @@
1
- from typing import Union
2
-
3
1
  import numpy as np
4
2
  from PIL import Image
5
3
 
@@ -15,7 +13,7 @@ def convert_to_rgb(image: Image.Image) -> Image.Image:
15
13
 
16
14
 
17
15
  def center_crop(
18
- image: Union[Image.Image, NumpyArray],
16
+ image: Image.Image | NumpyArray,
19
17
  size: tuple[int, int],
20
18
  ) -> NumpyArray:
21
19
  if isinstance(image, np.ndarray):
@@ -64,8 +62,8 @@ def center_crop(
64
62
 
65
63
  def normalize(
66
64
  image: NumpyArray,
67
- mean: Union[float, list[float]],
68
- std: Union[float, list[float]],
65
+ mean: float | list[float],
66
+ std: float | list[float],
69
67
  ) -> NumpyArray:
70
68
  num_channels = image.shape[1] if len(image.shape) == 4 else image.shape[0]
71
69
 
@@ -96,8 +94,8 @@ def normalize(
96
94
 
97
95
  def resize(
98
96
  image: Image.Image,
99
- size: Union[int, tuple[int, int]],
100
- resample: Union[int, Image.Resampling] = Image.Resampling.BILINEAR,
97
+ size: int | tuple[int, int],
98
+ resample: int | Image.Resampling = Image.Resampling.BILINEAR,
101
99
  ) -> Image.Image:
102
100
  if isinstance(size, tuple):
103
101
  return image.resize(size, resample)
@@ -117,7 +115,7 @@ def rescale(image: NumpyArray, scale: float, dtype: type = np.float32) -> NumpyA
117
115
  return (image * scale).astype(dtype)
118
116
 
119
117
 
120
- def pil2ndarray(image: Union[Image.Image, NumpyArray]) -> NumpyArray:
118
+ def pil2ndarray(image: Image.Image | NumpyArray) -> NumpyArray:
121
119
  if isinstance(image, Image.Image):
122
120
  return np.asarray(image).transpose((2, 0, 1))
123
121
  return image
@@ -126,7 +124,7 @@ def pil2ndarray(image: Union[Image.Image, NumpyArray]) -> NumpyArray:
126
124
  def pad2square(
127
125
  image: Image.Image,
128
126
  size: int,
129
- fill_color: Union[str, int, tuple[int, ...]] = 0,
127
+ fill_color: str | int | tuple[int, ...] = 0,
130
128
  ) -> Image.Image:
131
129
  height, width = image.height, image.width
132
130
 
@@ -147,3 +145,77 @@ def pad2square(
147
145
  new_image = Image.new(mode="RGB", size=(size, size), color=fill_color)
148
146
  new_image.paste(image.crop((left, top, right, bottom)) if crop_required else image)
149
147
  return new_image
148
+
149
+
150
+ def resize_longest_edge(
151
+ image: Image.Image,
152
+ max_size: int,
153
+ resample: int | Image.Resampling = Image.Resampling.LANCZOS,
154
+ ) -> Image.Image:
155
+ height, width = image.height, image.width
156
+ aspect_ratio = width / height
157
+
158
+ if width >= height:
159
+ # Width is longer
160
+ new_width = max_size
161
+ new_height = int(new_width / aspect_ratio)
162
+ else:
163
+ # Height is longer
164
+ new_height = max_size
165
+ new_width = int(new_height * aspect_ratio)
166
+
167
+ # Ensure even dimensions
168
+ if new_height % 2 != 0:
169
+ new_height += 1
170
+ if new_width % 2 != 0:
171
+ new_width += 1
172
+
173
+ return image.resize((new_width, new_height), resample)
174
+
175
+
176
+ def crop_ndarray(
177
+ image: NumpyArray,
178
+ x1: int,
179
+ y1: int,
180
+ x2: int,
181
+ y2: int,
182
+ channel_first: bool = True,
183
+ ) -> NumpyArray:
184
+ if channel_first:
185
+ # (C, H, W) format
186
+ return image[:, y1:y2, x1:x2]
187
+ else:
188
+ # (H, W, C) format
189
+ return image[y1:y2, x1:x2, :]
190
+
191
+
192
+ def resize_ndarray(
193
+ image: NumpyArray,
194
+ size: tuple[int, int],
195
+ resample: int | Image.Resampling = Image.Resampling.LANCZOS,
196
+ channel_first: bool = True,
197
+ ) -> NumpyArray:
198
+ # Convert to PIL-friendly format (H, W, C)
199
+ if channel_first:
200
+ img_hwc = image.transpose((1, 2, 0))
201
+ else:
202
+ img_hwc = image
203
+
204
+ # Handle different dtypes
205
+ if img_hwc.dtype == np.float32 or img_hwc.dtype == np.float64:
206
+ # Assume normalized, scale to 0-255 for PIL
207
+ img_hwc_scaled = (img_hwc * 255).astype(np.uint8)
208
+ pil_img = Image.fromarray(img_hwc_scaled, mode="RGB")
209
+ resized = pil_img.resize(size, resample)
210
+ result = np.array(resized).astype(np.float32) / 255.0
211
+ else:
212
+ # uint8 or similar
213
+ pil_img = Image.fromarray(img_hwc.astype(np.uint8), mode="RGB")
214
+ resized = pil_img.resize(size, resample)
215
+ result = np.array(resized)
216
+
217
+ # Convert back to original format
218
+ if channel_first:
219
+ result = result.transpose((2, 0, 1))
220
+
221
+ return result