fastembed-gpu 0.5.0__tar.gz → 0.6.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 (64) hide show
  1. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/NOTICE +8 -0
  2. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/PKG-INFO +44 -7
  3. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/README.md +38 -0
  4. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/__init__.py +2 -0
  5. fastembed_gpu-0.6.0/fastembed/common/__init__.py +3 -0
  6. fastembed_gpu-0.6.0/fastembed/common/model_description.py +47 -0
  7. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/common/model_management.py +185 -30
  8. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/common/onnx_model.py +20 -16
  9. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/common/preprocessor_utils.py +4 -3
  10. fastembed_gpu-0.6.0/fastembed/common/types.py +24 -0
  11. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/common/utils.py +17 -3
  12. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/embedding.py +2 -2
  13. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/image/image_embedding.py +17 -13
  14. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/image/image_embedding_base.py +9 -9
  15. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/image/onnx_embedding.py +72 -75
  16. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/image/onnx_image_model.py +18 -14
  17. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/image/transform/functional.py +30 -31
  18. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/image/transform/operators.py +26 -25
  19. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/late_interaction/colbert.py +57 -50
  20. fastembed_gpu-0.6.0/fastembed/late_interaction/jina_colbert.py +58 -0
  21. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/late_interaction/late_interaction_embedding_base.py +12 -14
  22. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/late_interaction/late_interaction_text_embedding.py +16 -11
  23. fastembed_gpu-0.6.0/fastembed/late_interaction_multimodal/__init__.py +5 -0
  24. fastembed_gpu-0.6.0/fastembed/late_interaction_multimodal/colpali.py +300 -0
  25. fastembed_gpu-0.6.0/fastembed/late_interaction_multimodal/late_interaction_multimodal_embedding.py +130 -0
  26. fastembed_gpu-0.6.0/fastembed/late_interaction_multimodal/late_interaction_multimodal_embedding_base.py +67 -0
  27. fastembed_gpu-0.6.0/fastembed/late_interaction_multimodal/onnx_multimodal_model.py +271 -0
  28. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/parallel_processor.py +1 -1
  29. fastembed_gpu-0.6.0/fastembed/py.typed +1 -0
  30. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/rerank/cross_encoder/onnx_text_cross_encoder.py +60 -69
  31. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/rerank/cross_encoder/onnx_text_model.py +29 -10
  32. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/rerank/cross_encoder/text_cross_encoder.py +11 -5
  33. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/rerank/cross_encoder/text_cross_encoder_base.py +4 -3
  34. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/sparse/bm25.py +37 -29
  35. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/sparse/bm42.py +50 -44
  36. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/sparse/sparse_embedding_base.py +16 -11
  37. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/sparse/sparse_text_embedding.py +15 -7
  38. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/sparse/splade_pp.py +37 -36
  39. fastembed_gpu-0.6.0/fastembed/text/clip_embedding.py +54 -0
  40. fastembed_gpu-0.6.0/fastembed/text/custom_text_embedding.py +91 -0
  41. fastembed_gpu-0.6.0/fastembed/text/multitask_embedding.py +100 -0
  42. fastembed_gpu-0.6.0/fastembed/text/onnx_embedding.py +337 -0
  43. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/text/onnx_text_model.py +18 -17
  44. fastembed_gpu-0.6.0/fastembed/text/pooled_embedding.py +133 -0
  45. fastembed_gpu-0.6.0/fastembed/text/pooled_normalized_embedding.py +162 -0
  46. fastembed_gpu-0.6.0/fastembed/text/text_embedding.py +180 -0
  47. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/text/text_embedding_base.py +12 -14
  48. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/pyproject.toml +17 -9
  49. fastembed_gpu-0.5.0/fastembed/common/__init__.py +0 -3
  50. fastembed_gpu-0.5.0/fastembed/common/types.py +0 -16
  51. fastembed_gpu-0.5.0/fastembed/late_interaction/jina_colbert.py +0 -62
  52. fastembed_gpu-0.5.0/fastembed/text/clip_embedding.py +0 -54
  53. fastembed_gpu-0.5.0/fastembed/text/e5_onnx_embedding.py +0 -72
  54. fastembed_gpu-0.5.0/fastembed/text/onnx_embedding.py +0 -333
  55. fastembed_gpu-0.5.0/fastembed/text/pooled_embedding.py +0 -92
  56. fastembed_gpu-0.5.0/fastembed/text/pooled_normalized_embedding.py +0 -125
  57. fastembed_gpu-0.5.0/fastembed/text/text_embedding.py +0 -107
  58. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/LICENSE +0 -0
  59. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/image/__init__.py +0 -0
  60. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/late_interaction/__init__.py +0 -0
  61. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/rerank/cross_encoder/__init__.py +0 -0
  62. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/sparse/__init__.py +0 -0
  63. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/sparse/utils/tokenizer.py +0 -0
  64. {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/text/__init__.py +0 -0
@@ -7,8 +7,16 @@ This distribution includes the following Jina AI models, each with its respectiv
7
7
  - License: cc-by-nc-4.0
8
8
  - jinaai/jina-reranker-v2-base-multilingual
9
9
  - License: cc-by-nc-4.0
10
+ - jinaai/jina-embeddings-v3
11
+ - License: cc-by-nc-4.0
10
12
 
11
13
  These models are developed by Jina (https://jina.ai/) and are subject to Jina AI's licensing terms.
12
14
 
15
+ This distribution includes the following Google models, each with its respective license:
16
+ - vidore/colpali-v1.3
17
+ - License: gemma
18
+
19
+ Gemma is provided under and subject to the Gemma Terms of Use found at ai.google.dev/gemma/terms
20
+
13
21
  Additional Notes:
14
22
  This project also includes third-party libraries with their respective licenses. Please refer to the documentation of each library for details regarding its usage and licensing terms.
@@ -1,8 +1,7 @@
1
- Metadata-Version: 2.1
1
+ Metadata-Version: 2.3
2
2
  Name: fastembed-gpu
3
- Version: 0.5.0
3
+ Version: 0.6.0
4
4
  Summary: Fast, light, accurate library built for retrieval embedding generation
5
- Home-page: https://github.com/qdrant/fastembed
6
5
  License: Apache License
7
6
  Keywords: vector,embedding,neural,search,qdrant,sentence-transformers
8
7
  Author: Qdrant Team
@@ -17,20 +16,20 @@ Classifier: Programming Language :: Python :: 3.12
17
16
  Classifier: Programming Language :: Python :: 3.13
18
17
  Requires-Dist: huggingface-hub (>=0.20,<1.0)
19
18
  Requires-Dist: loguru (>=0.7.2,<0.8.0)
20
- Requires-Dist: mmh3 (>=4.1.0,<5.0.0)
19
+ Requires-Dist: mmh3 (>=4.1.0,<6.0.0)
21
20
  Requires-Dist: numpy (>=1.21) ; python_version >= "3.10" and python_version < "3.12"
22
21
  Requires-Dist: numpy (>=1.21,<2.1.0) ; python_version < "3.10"
23
- Requires-Dist: numpy (>=1.26) ; python_version >= "3.12" and python_version < "3.13"
22
+ Requires-Dist: numpy (>=1.26) ; python_version == "3.12"
24
23
  Requires-Dist: numpy (>=2.1.0) ; python_version >= "3.13"
25
- Requires-Dist: onnx (>=1.15.0)
26
24
  Requires-Dist: onnxruntime-gpu (>1.20.0) ; python_version >= "3.13"
27
25
  Requires-Dist: onnxruntime-gpu (>=1.17.0,!=1.20.0) ; python_version >= "3.10" and python_version < "3.13"
28
26
  Requires-Dist: onnxruntime-gpu (>=1.17.0,<1.20.0) ; python_version < "3.10"
29
- Requires-Dist: pillow (>=10.3.0,<11.0.0)
27
+ Requires-Dist: pillow (>=10.3.0,<12.0.0)
30
28
  Requires-Dist: py-rust-stemmers (>=0.1.0,<0.2.0)
31
29
  Requires-Dist: requests (>=2.31,<3.0)
32
30
  Requires-Dist: tokenizers (>=0.15,<1.0)
33
31
  Requires-Dist: tqdm (>=4.66,<5.0)
32
+ Project-URL: Homepage, https://github.com/qdrant/fastembed
34
33
  Project-URL: Repository, https://github.com/qdrant/fastembed
35
34
  Description-Content-Type: text/markdown
36
35
 
@@ -99,6 +98,23 @@ embeddings = list(model.embed(documents))
99
98
 
100
99
  ```
101
100
 
101
+ Dense text embedding can also be extended with models which are not in the list of supported models.
102
+
103
+ ```python
104
+ from fastembed import TextEmbedding
105
+ from fastembed.common.model_description import PoolingType, ModelSource
106
+
107
+ TextEmbedding.add_custom_model(
108
+ model="intfloat/multilingual-e5-small",
109
+ pooling=PoolingType.MEAN,
110
+ normalization=True,
111
+ sources=ModelSource(hf="intfloat/multilingual-e5-small"), # can be used with an `url` to load files from a private storage
112
+ dim=384,
113
+ model_file="onnx/model.onnx", # can be used to load an already supported model with another optimization or quantization, e.g. onnx/model_O4.onnx
114
+ )
115
+ model = TextEmbedding(model_name="intfloat/multilingual-e5-small")
116
+ embeddings = list(model.embed(documents))
117
+ ```
102
118
 
103
119
 
104
120
  ### 🔱 Sparse text embeddings
@@ -173,6 +189,27 @@ embeddings = list(model.embed(images))
173
189
  # ]
174
190
  ```
175
191
 
192
+ ### Late interaction multimodal models (ColPali)
193
+
194
+ ```python
195
+ from fastembed import LateInteractionMultimodalEmbedding
196
+
197
+ doc_images = [
198
+ "./path/to/qdrant_pdf_doc_1_screenshot.jpg",
199
+ "./path/to/colpali_pdf_doc_2_screenshot.jpg",
200
+ ]
201
+
202
+ query = "What is Qdrant?"
203
+
204
+ model = LateInteractionMultimodalEmbedding(model_name="Qdrant/colpali-v1.3-fp16")
205
+ doc_images_embeddings = list(model.embed_image(doc_images))
206
+ # shape (2, 1030, 128)
207
+ # [array([[-0.03353882, -0.02090454, ..., -0.15576172, -0.07678223]], dtype=float32)]
208
+ query_embedding = model.embed_text(query)
209
+ # shape (1, 20, 128)
210
+ # [array([[-0.00218201, 0.14758301, ..., -0.02207947, 0.16833496]], dtype=float32)]
211
+ ```
212
+
176
213
  ### 🔄 Rerankers
177
214
  ```python
178
215
  from fastembed.rerank.cross_encoder import TextCrossEncoder
@@ -63,6 +63,23 @@ embeddings = list(model.embed(documents))
63
63
 
64
64
  ```
65
65
 
66
+ Dense text embedding can also be extended with models which are not in the list of supported models.
67
+
68
+ ```python
69
+ from fastembed import TextEmbedding
70
+ from fastembed.common.model_description import PoolingType, ModelSource
71
+
72
+ TextEmbedding.add_custom_model(
73
+ model="intfloat/multilingual-e5-small",
74
+ pooling=PoolingType.MEAN,
75
+ normalization=True,
76
+ sources=ModelSource(hf="intfloat/multilingual-e5-small"), # can be used with an `url` to load files from a private storage
77
+ dim=384,
78
+ model_file="onnx/model.onnx", # can be used to load an already supported model with another optimization or quantization, e.g. onnx/model_O4.onnx
79
+ )
80
+ model = TextEmbedding(model_name="intfloat/multilingual-e5-small")
81
+ embeddings = list(model.embed(documents))
82
+ ```
66
83
 
67
84
 
68
85
  ### 🔱 Sparse text embeddings
@@ -137,6 +154,27 @@ embeddings = list(model.embed(images))
137
154
  # ]
138
155
  ```
139
156
 
157
+ ### Late interaction multimodal models (ColPali)
158
+
159
+ ```python
160
+ from fastembed import LateInteractionMultimodalEmbedding
161
+
162
+ doc_images = [
163
+ "./path/to/qdrant_pdf_doc_1_screenshot.jpg",
164
+ "./path/to/colpali_pdf_doc_2_screenshot.jpg",
165
+ ]
166
+
167
+ query = "What is Qdrant?"
168
+
169
+ model = LateInteractionMultimodalEmbedding(model_name="Qdrant/colpali-v1.3-fp16")
170
+ doc_images_embeddings = list(model.embed_image(doc_images))
171
+ # shape (2, 1030, 128)
172
+ # [array([[-0.03353882, -0.02090454, ..., -0.15576172, -0.07678223]], dtype=float32)]
173
+ query_embedding = model.embed_text(query)
174
+ # shape (1, 20, 128)
175
+ # [array([[-0.00218201, 0.14758301, ..., -0.02207947, 0.16833496]], dtype=float32)]
176
+ ```
177
+
140
178
  ### 🔄 Rerankers
141
179
  ```python
142
180
  from fastembed.rerank.cross_encoder import TextCrossEncoder
@@ -2,6 +2,7 @@ import importlib.metadata
2
2
 
3
3
  from fastembed.image import ImageEmbedding
4
4
  from fastembed.late_interaction import LateInteractionTextEmbedding
5
+ from fastembed.late_interaction_multimodal import LateInteractionMultimodalEmbedding
5
6
  from fastembed.sparse import SparseEmbedding, SparseTextEmbedding
6
7
  from fastembed.text import TextEmbedding
7
8
 
@@ -17,4 +18,5 @@ __all__ = [
17
18
  "SparseEmbedding",
18
19
  "ImageEmbedding",
19
20
  "LateInteractionTextEmbedding",
21
+ "LateInteractionMultimodalEmbedding",
20
22
  ]
@@ -0,0 +1,3 @@
1
+ from fastembed.common.types import ImageInput, OnnxProvider, PathInput
2
+
3
+ __all__ = ["OnnxProvider", "ImageInput", "PathInput"]
@@ -0,0 +1,47 @@
1
+ from dataclasses import dataclass, field
2
+ from enum import Enum
3
+ from typing import Optional, Any
4
+
5
+
6
+ @dataclass(frozen=True)
7
+ class ModelSource:
8
+ hf: Optional[str] = None
9
+ url: Optional[str] = None
10
+
11
+ def __post_init__(self) -> None:
12
+ if self.hf is None and self.url is None:
13
+ raise ValueError(
14
+ f"At least one source should be set, current sources: hf={self.hf}, url={self.url}"
15
+ )
16
+
17
+
18
+ @dataclass(frozen=True)
19
+ class BaseModelDescription:
20
+ model: str
21
+ sources: ModelSource
22
+ model_file: str
23
+ description: str
24
+ license: str
25
+ size_in_GB: float
26
+ additional_files: list[str] = field(default_factory=list)
27
+
28
+
29
+ @dataclass(frozen=True)
30
+ class DenseModelDescription(BaseModelDescription):
31
+ dim: Optional[int] = None
32
+ tasks: Optional[dict[str, Any]] = field(default_factory=dict)
33
+
34
+ def __post_init__(self) -> None:
35
+ assert self.dim is not None, "dim is required for dense model description"
36
+
37
+
38
+ @dataclass(frozen=True)
39
+ class SparseModelDescription(BaseModelDescription):
40
+ requires_idf: Optional[bool] = None
41
+ vocab_size: Optional[int] = None
42
+
43
+
44
+ class PoolingType(str, Enum):
45
+ CLS = "CLS"
46
+ MEAN = "MEAN"
47
+ DISABLED = "DISABLED"
@@ -1,12 +1,14 @@
1
1
  import os
2
2
  import time
3
+ import json
3
4
  import shutil
4
5
  import tarfile
5
6
  from pathlib import Path
6
- from typing import Any, Optional
7
+ from typing import Any, Optional, Union, TypeVar, Generic
7
8
 
8
9
  import requests
9
- from huggingface_hub import snapshot_download
10
+ from huggingface_hub import snapshot_download, model_info, list_repo_tree
11
+ from huggingface_hub.hf_api import RepoFile
10
12
  from huggingface_hub.utils import (
11
13
  RepositoryNotFoundError,
12
14
  disable_progress_bars,
@@ -14,20 +16,54 @@ from huggingface_hub.utils import (
14
16
  )
15
17
  from loguru import logger
16
18
  from tqdm import tqdm
19
+ from fastembed.common.model_description import BaseModelDescription
17
20
 
21
+ T = TypeVar("T", bound=BaseModelDescription)
22
+
23
+
24
+ class ModelManagement(Generic[T]):
25
+ METADATA_FILE = "files_metadata.json"
18
26
 
19
- class ModelManagement:
20
27
  @classmethod
21
28
  def list_supported_models(cls) -> list[dict[str, Any]]:
22
29
  """Lists the supported models.
23
30
 
24
31
  Returns:
25
- list[dict[str, Any]]: A list of dictionaries containing the model information.
32
+ list[T]: A list of dictionaries containing the model information.
26
33
  """
27
34
  raise NotImplementedError()
28
35
 
29
36
  @classmethod
30
- def _get_model_description(cls, model_name: str) -> dict[str, Any]:
37
+ def add_custom_model(
38
+ cls,
39
+ *args: Any,
40
+ **kwargs: Any,
41
+ ) -> None:
42
+ """Add a custom model to the existing embedding classes based on the passed model descriptions
43
+
44
+ Model description dict should contain the fields same as in one of the model descriptions presented
45
+ in fastembed.common.model_description
46
+
47
+ E.g. for BaseModelDescription:
48
+ model: str
49
+ sources: ModelSource
50
+ model_file: str
51
+ description: str
52
+ license: str
53
+ size_in_GB: float
54
+ additional_files: list[str]
55
+
56
+ Returns:
57
+ None
58
+ """
59
+ raise NotImplementedError()
60
+
61
+ @classmethod
62
+ def _list_supported_models(cls) -> list[T]:
63
+ raise NotImplementedError()
64
+
65
+ @classmethod
66
+ def _get_model_description(cls, model_name: str) -> T:
31
67
  """
32
68
  Gets the model description from the model_name.
33
69
 
@@ -38,10 +74,10 @@ class ModelManagement:
38
74
  ValueError: If the model_name is not supported.
39
75
 
40
76
  Returns:
41
- dict[str, Any]: The model description.
77
+ T: The model description.
42
78
  """
43
- for model in cls.list_supported_models():
44
- if model_name.lower() == model["model"].lower():
79
+ for model in cls._list_supported_models():
80
+ if model_name.lower() == model.model.lower():
45
81
  return model
46
82
 
47
83
  raise ValueError(f"Model {model_name} is not supported in {cls.__name__}.")
@@ -98,21 +134,74 @@ class ModelManagement:
98
134
  cls,
99
135
  hf_source_repo: str,
100
136
  cache_dir: str,
101
- extra_patterns: Optional[list[str]] = None,
137
+ extra_patterns: list[str],
102
138
  local_files_only: bool = False,
103
- **kwargs,
139
+ **kwargs: Any,
104
140
  ) -> str:
105
141
  """
106
142
  Downloads a model from HuggingFace Hub.
107
143
  Args:
108
144
  hf_source_repo (str): Name of the model on HuggingFace Hub, e.g. "qdrant/all-MiniLM-L6-v2-onnx".
109
145
  cache_dir (Optional[str]): The path to the cache directory.
110
- extra_patterns (Optional[list[str]]): extra patterns to allow in the snapshot download, typically
146
+ extra_patterns (list[str]): extra patterns to allow in the snapshot download, typically
111
147
  includes the required model files.
112
148
  local_files_only (bool, optional): Whether to only use local files. Defaults to False.
113
149
  Returns:
114
150
  Path: The path to the model directory.
115
151
  """
152
+
153
+ def _verify_files_from_metadata(
154
+ model_dir: Path, stored_metadata: dict[str, Any], repo_files: list[RepoFile]
155
+ ) -> bool:
156
+ try:
157
+ for rel_path, meta in stored_metadata.items():
158
+ file_path = model_dir / rel_path
159
+
160
+ if not file_path.exists():
161
+ return False
162
+
163
+ if repo_files: # online verification
164
+ file_info = next((f for f in repo_files if f.path == file_path.name), None)
165
+ if (
166
+ not file_info
167
+ or file_info.size != meta["size"]
168
+ or file_info.blob_id != meta["blob_id"]
169
+ ):
170
+ return False
171
+
172
+ else: # offline verification
173
+ if file_path.stat().st_size != meta["size"]:
174
+ return False
175
+ return True
176
+ except (OSError, KeyError) as e:
177
+ logger.error(f"Error verifying files: {str(e)}")
178
+ return False
179
+
180
+ def _collect_file_metadata(
181
+ model_dir: Path, repo_files: list[RepoFile]
182
+ ) -> dict[str, dict[str, Union[int, str]]]:
183
+ meta: dict[str, dict[str, Union[int, str]]] = {}
184
+ file_info_map = {f.path: f for f in repo_files}
185
+ for file_path in model_dir.rglob("*"):
186
+ if file_path.is_file() and file_path.name != cls.METADATA_FILE:
187
+ repo_file = file_info_map.get(file_path.name)
188
+ if repo_file:
189
+ meta[str(file_path.relative_to(model_dir))] = {
190
+ "size": repo_file.size,
191
+ "blob_id": repo_file.blob_id,
192
+ }
193
+ return meta
194
+
195
+ def _save_file_metadata(
196
+ model_dir: Path, meta: dict[str, dict[str, Union[int, str]]]
197
+ ) -> None:
198
+ try:
199
+ if not model_dir.exists():
200
+ model_dir.mkdir(parents=True, exist_ok=True)
201
+ (model_dir / cls.METADATA_FILE).write_text(json.dumps(meta))
202
+ except (OSError, ValueError) as e:
203
+ logger.warning(f"Error saving metadata: {str(e)}")
204
+
116
205
  allow_patterns = [
117
206
  "config.json",
118
207
  "tokenizer.json",
@@ -120,16 +209,59 @@ class ModelManagement:
120
209
  "special_tokens_map.json",
121
210
  "preprocessor_config.json",
122
211
  ]
123
- if extra_patterns is not None:
124
- allow_patterns.extend(extra_patterns)
212
+
213
+ allow_patterns.extend(extra_patterns)
125
214
 
126
215
  snapshot_dir = Path(cache_dir) / f"models--{hf_source_repo.replace('/', '--')}"
127
- is_cached = snapshot_dir.exists()
216
+ metadata_file = snapshot_dir / cls.METADATA_FILE
128
217
 
129
- if is_cached:
218
+ if local_files_only:
219
+ disable_progress_bars()
220
+ if metadata_file.exists():
221
+ metadata = json.loads(metadata_file.read_text())
222
+ verified = _verify_files_from_metadata(snapshot_dir, metadata, repo_files=[])
223
+ if not verified:
224
+ logger.warning(
225
+ "Local file sizes do not match the metadata."
226
+ ) # 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
+ result = snapshot_download(
233
+ repo_id=hf_source_repo,
234
+ allow_patterns=allow_patterns,
235
+ cache_dir=cache_dir,
236
+ local_files_only=local_files_only,
237
+ **kwargs,
238
+ )
239
+ return result
240
+
241
+ repo_revision = model_info(hf_source_repo).sha
242
+ repo_tree = list(list_repo_tree(hf_source_repo, revision=repo_revision, repo_type="model"))
243
+
244
+ allowed_extensions = {".json", ".onnx", ".txt"}
245
+ repo_files = (
246
+ [
247
+ f
248
+ for f in repo_tree
249
+ if isinstance(f, RepoFile) and Path(f.path).suffix in allowed_extensions
250
+ ]
251
+ if repo_tree
252
+ else []
253
+ )
254
+
255
+ verified_metadata = False
256
+
257
+ if snapshot_dir.exists() and metadata_file.exists():
258
+ metadata = json.loads(metadata_file.read_text())
259
+ verified_metadata = _verify_files_from_metadata(snapshot_dir, metadata, repo_files)
260
+
261
+ if verified_metadata:
130
262
  disable_progress_bars()
131
263
 
132
- return snapshot_download(
264
+ result = snapshot_download(
133
265
  repo_id=hf_source_repo,
134
266
  allow_patterns=allow_patterns,
135
267
  cache_dir=cache_dir,
@@ -137,8 +269,26 @@ class ModelManagement:
137
269
  **kwargs,
138
270
  )
139
271
 
272
+ if (
273
+ not verified_metadata
274
+ ): # metadata is not up-to-date, update it and check whether the files have been
275
+ # downloaded correctly
276
+ metadata = _collect_file_metadata(snapshot_dir, repo_files)
277
+
278
+ download_successful = _verify_files_from_metadata(
279
+ snapshot_dir, metadata, repo_files=[]
280
+ ) # offline verification
281
+ if not download_successful:
282
+ raise ValueError(
283
+ "Files have been corrupted during downloading process. "
284
+ "Please check your internet connection and try again."
285
+ )
286
+ _save_file_metadata(snapshot_dir, metadata)
287
+
288
+ return result
289
+
140
290
  @classmethod
141
- def decompress_to_cache(cls, targz_path: str, cache_dir: str):
291
+ def decompress_to_cache(cls, targz_path: str, cache_dir: str) -> str:
142
292
  """
143
293
  Decompresses a .tar.gz file to a cache directory.
144
294
 
@@ -176,7 +326,11 @@ class ModelManagement:
176
326
 
177
327
  @classmethod
178
328
  def retrieve_model_gcs(
179
- cls, model_name: str, source_url: str, cache_dir: str, local_files_only: bool = False
329
+ cls,
330
+ model_name: str,
331
+ source_url: str,
332
+ cache_dir: str,
333
+ local_files_only: bool = False,
180
334
  ) -> Path:
181
335
  fast_model_name = f"fast-{model_name.split('/')[-1]}"
182
336
  cache_tmp_dir = Path(cache_dir) / "tmp"
@@ -220,14 +374,12 @@ class ModelManagement:
220
374
  return model_dir
221
375
 
222
376
  @classmethod
223
- def download_model(
224
- cls, model: dict[str, Any], cache_dir: Path, retries: int = 3, **kwargs
225
- ) -> Path:
377
+ def download_model(cls, model: T, cache_dir: str, retries: int = 3, **kwargs: Any) -> Path:
226
378
  """
227
379
  Downloads a model from HuggingFace Hub or Google Cloud Storage.
228
380
 
229
381
  Args:
230
- model (dict[str, Any]): The model description.
382
+ model (T): The model description.
231
383
  Example:
232
384
  ```
233
385
  {
@@ -248,23 +400,26 @@ class ModelManagement:
248
400
  Path: The path to the downloaded model directory.
249
401
  """
250
402
  local_files_only = kwargs.get("local_files_only", False)
403
+ specific_model_path: Optional[str] = kwargs.pop("specific_model_path", None)
404
+ if specific_model_path:
405
+ return Path(specific_model_path)
251
406
  retries = 1 if local_files_only else retries
252
- hf_source = model.get("sources", {}).get("hf")
253
- url_source = model.get("sources", {}).get("url")
407
+ hf_source = model.sources.hf
408
+ url_source = model.sources.url
254
409
 
255
410
  sleep = 3.0
256
411
  while retries > 0:
257
412
  retries -= 1
258
413
 
259
414
  if hf_source:
260
- extra_patterns = [model["model_file"]]
261
- extra_patterns.extend(model.get("additional_files", []))
415
+ extra_patterns = [model.model_file]
416
+ extra_patterns.extend(model.additional_files)
262
417
 
263
418
  try:
264
419
  return Path(
265
420
  cls.download_files_from_huggingface(
266
421
  hf_source,
267
- cache_dir=str(cache_dir),
422
+ cache_dir=cache_dir,
268
423
  extra_patterns=extra_patterns,
269
424
  **kwargs,
270
425
  )
@@ -280,8 +435,8 @@ class ModelManagement:
280
435
  if url_source or local_files_only:
281
436
  try:
282
437
  return cls.retrieve_model_gcs(
283
- model["model"],
284
- url_source,
438
+ model.model,
439
+ str(url_source),
285
440
  str(cache_dir),
286
441
  local_files_only=local_files_only,
287
442
  )
@@ -298,4 +453,4 @@ class ModelManagement:
298
453
  time.sleep(sleep)
299
454
  sleep *= 3
300
455
 
301
- raise ValueError(f"Could not load model {model['model']} from any source.")
456
+ raise ValueError(f"Could not load model {model.model} from any source.")
@@ -6,7 +6,10 @@ from typing import Any, Generic, Iterable, Optional, Sequence, Type, TypeVar
6
6
  import numpy as np
7
7
  import onnxruntime as ort
8
8
 
9
- from fastembed.common.types import OnnxProvider
9
+ from numpy.typing import NDArray
10
+ from tokenizers import Tokenizer
11
+
12
+ from fastembed.common.types import OnnxProvider, NumpyArray
10
13
  from fastembed.parallel_processor import Worker
11
14
 
12
15
  # Holds type of the embedding result
@@ -15,26 +18,26 @@ T = TypeVar("T")
15
18
 
16
19
  @dataclass
17
20
  class OnnxOutputContext:
18
- model_output: np.ndarray
19
- attention_mask: Optional[np.ndarray] = None
20
- input_ids: Optional[np.ndarray] = None
21
+ model_output: NumpyArray
22
+ attention_mask: Optional[NDArray[np.int64]] = None
23
+ input_ids: Optional[NDArray[np.int64]] = None
21
24
 
22
25
 
23
26
  class OnnxModel(Generic[T]):
24
27
  @classmethod
25
- def _get_worker_class(cls) -> Type["EmbeddingWorker"]:
28
+ def _get_worker_class(cls) -> Type["EmbeddingWorker[T]"]:
26
29
  raise NotImplementedError("Subclasses must implement this method")
27
30
 
28
31
  def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[T]:
29
32
  raise NotImplementedError("Subclasses must implement this method")
30
33
 
31
34
  def __init__(self) -> None:
32
- self.model = None
33
- self.tokenizer = None
35
+ self.model: Optional[ort.InferenceSession] = None
36
+ self.tokenizer: Optional[Tokenizer] = None
34
37
 
35
38
  def _preprocess_onnx_input(
36
- self, onnx_input: dict[str, np.ndarray], **kwargs
37
- ) -> dict[str, np.ndarray]:
39
+ self, onnx_input: dict[str, NumpyArray], **kwargs: Any
40
+ ) -> dict[str, NumpyArray]:
38
41
  """
39
42
  Preprocess the onnx input.
40
43
  """
@@ -70,7 +73,7 @@ class OnnxModel(Generic[T]):
70
73
  onnx_providers = ["CPUExecutionProvider"]
71
74
 
72
75
  available_providers = ort.get_available_providers()
73
- requested_provider_names = []
76
+ requested_provider_names: list[str] = []
74
77
  for provider in onnx_providers:
75
78
  # check providers available
76
79
  provider_name = provider if isinstance(provider, str) else provider[0]
@@ -91,6 +94,7 @@ class OnnxModel(Generic[T]):
91
94
  str(model_path), providers=onnx_providers, sess_options=so
92
95
  )
93
96
  if "CUDAExecutionProvider" in requested_provider_names:
97
+ assert self.model is not None
94
98
  current_providers = self.model.get_providers()
95
99
  if "CUDAExecutionProvider" not in current_providers:
96
100
  warnings.warn(
@@ -103,29 +107,29 @@ class OnnxModel(Generic[T]):
103
107
  def load_onnx_model(self) -> None:
104
108
  raise NotImplementedError("Subclasses must implement this method")
105
109
 
106
- def onnx_embed(self, *args, **kwargs) -> OnnxOutputContext:
110
+ def onnx_embed(self, *args: Any, **kwargs: Any) -> OnnxOutputContext:
107
111
  raise NotImplementedError("Subclasses must implement this method")
108
112
 
109
113
 
110
- class EmbeddingWorker(Worker):
114
+ class EmbeddingWorker(Worker, Generic[T]):
111
115
  def init_embedding(
112
116
  self,
113
117
  model_name: str,
114
118
  cache_dir: str,
115
- **kwargs,
116
- ) -> OnnxModel:
119
+ **kwargs: Any,
120
+ ) -> OnnxModel[T]:
117
121
  raise NotImplementedError()
118
122
 
119
123
  def __init__(
120
124
  self,
121
125
  model_name: str,
122
126
  cache_dir: str,
123
- **kwargs,
127
+ **kwargs: Any,
124
128
  ):
125
129
  self.model = self.init_embedding(model_name, cache_dir, **kwargs)
126
130
 
127
131
  @classmethod
128
- def start(cls, model_name: str, cache_dir: str, **kwargs: Any) -> "EmbeddingWorker":
132
+ def start(cls, model_name: str, cache_dir: str, **kwargs: Any) -> "EmbeddingWorker[T]":
129
133
  return cls(model_name=model_name, cache_dir=cache_dir, **kwargs)
130
134
 
131
135
  def process(self, items: Iterable[tuple[int, Any]]) -> Iterable[tuple[int, Any]]: