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.
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/NOTICE +8 -0
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/PKG-INFO +44 -7
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/README.md +38 -0
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/__init__.py +2 -0
- fastembed_gpu-0.6.0/fastembed/common/__init__.py +3 -0
- fastembed_gpu-0.6.0/fastembed/common/model_description.py +47 -0
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/common/model_management.py +185 -30
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/common/onnx_model.py +20 -16
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/common/preprocessor_utils.py +4 -3
- fastembed_gpu-0.6.0/fastembed/common/types.py +24 -0
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/common/utils.py +17 -3
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/embedding.py +2 -2
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/image/image_embedding.py +17 -13
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/image/image_embedding_base.py +9 -9
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/image/onnx_embedding.py +72 -75
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/image/onnx_image_model.py +18 -14
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/image/transform/functional.py +30 -31
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/image/transform/operators.py +26 -25
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/late_interaction/colbert.py +57 -50
- fastembed_gpu-0.6.0/fastembed/late_interaction/jina_colbert.py +58 -0
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/late_interaction/late_interaction_embedding_base.py +12 -14
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/late_interaction/late_interaction_text_embedding.py +16 -11
- fastembed_gpu-0.6.0/fastembed/late_interaction_multimodal/__init__.py +5 -0
- fastembed_gpu-0.6.0/fastembed/late_interaction_multimodal/colpali.py +300 -0
- fastembed_gpu-0.6.0/fastembed/late_interaction_multimodal/late_interaction_multimodal_embedding.py +130 -0
- fastembed_gpu-0.6.0/fastembed/late_interaction_multimodal/late_interaction_multimodal_embedding_base.py +67 -0
- fastembed_gpu-0.6.0/fastembed/late_interaction_multimodal/onnx_multimodal_model.py +271 -0
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/parallel_processor.py +1 -1
- fastembed_gpu-0.6.0/fastembed/py.typed +1 -0
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/rerank/cross_encoder/onnx_text_cross_encoder.py +60 -69
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/rerank/cross_encoder/onnx_text_model.py +29 -10
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/rerank/cross_encoder/text_cross_encoder.py +11 -5
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/rerank/cross_encoder/text_cross_encoder_base.py +4 -3
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/sparse/bm25.py +37 -29
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/sparse/bm42.py +50 -44
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/sparse/sparse_embedding_base.py +16 -11
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/sparse/sparse_text_embedding.py +15 -7
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/sparse/splade_pp.py +37 -36
- fastembed_gpu-0.6.0/fastembed/text/clip_embedding.py +54 -0
- fastembed_gpu-0.6.0/fastembed/text/custom_text_embedding.py +91 -0
- fastembed_gpu-0.6.0/fastembed/text/multitask_embedding.py +100 -0
- fastembed_gpu-0.6.0/fastembed/text/onnx_embedding.py +337 -0
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/text/onnx_text_model.py +18 -17
- fastembed_gpu-0.6.0/fastembed/text/pooled_embedding.py +133 -0
- fastembed_gpu-0.6.0/fastembed/text/pooled_normalized_embedding.py +162 -0
- fastembed_gpu-0.6.0/fastembed/text/text_embedding.py +180 -0
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/text/text_embedding_base.py +12 -14
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/pyproject.toml +17 -9
- fastembed_gpu-0.5.0/fastembed/common/__init__.py +0 -3
- fastembed_gpu-0.5.0/fastembed/common/types.py +0 -16
- fastembed_gpu-0.5.0/fastembed/late_interaction/jina_colbert.py +0 -62
- fastembed_gpu-0.5.0/fastembed/text/clip_embedding.py +0 -54
- fastembed_gpu-0.5.0/fastembed/text/e5_onnx_embedding.py +0 -72
- fastembed_gpu-0.5.0/fastembed/text/onnx_embedding.py +0 -333
- fastembed_gpu-0.5.0/fastembed/text/pooled_embedding.py +0 -92
- fastembed_gpu-0.5.0/fastembed/text/pooled_normalized_embedding.py +0 -125
- fastembed_gpu-0.5.0/fastembed/text/text_embedding.py +0 -107
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/LICENSE +0 -0
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/image/__init__.py +0 -0
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/late_interaction/__init__.py +0 -0
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/rerank/cross_encoder/__init__.py +0 -0
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/sparse/__init__.py +0 -0
- {fastembed_gpu-0.5.0 → fastembed_gpu-0.6.0}/fastembed/sparse/utils/tokenizer.py +0 -0
- {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
|
+
Metadata-Version: 2.3
|
|
2
2
|
Name: fastembed-gpu
|
|
3
|
-
Version: 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,<
|
|
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
|
|
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,<
|
|
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,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[
|
|
32
|
+
list[T]: A list of dictionaries containing the model information.
|
|
26
33
|
"""
|
|
27
34
|
raise NotImplementedError()
|
|
28
35
|
|
|
29
36
|
@classmethod
|
|
30
|
-
def
|
|
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
|
-
|
|
77
|
+
T: The model description.
|
|
42
78
|
"""
|
|
43
|
-
for model in cls.
|
|
44
|
-
if model_name.lower() == model
|
|
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:
|
|
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 (
|
|
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
|
-
|
|
124
|
-
|
|
212
|
+
|
|
213
|
+
allow_patterns.extend(extra_patterns)
|
|
125
214
|
|
|
126
215
|
snapshot_dir = Path(cache_dir) / f"models--{hf_source_repo.replace('/', '--')}"
|
|
127
|
-
|
|
216
|
+
metadata_file = snapshot_dir / cls.METADATA_FILE
|
|
128
217
|
|
|
129
|
-
if
|
|
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
|
-
|
|
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,
|
|
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 (
|
|
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.
|
|
253
|
-
url_source = model.
|
|
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
|
|
261
|
-
extra_patterns.extend(model.
|
|
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=
|
|
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
|
|
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
|
|
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
|
|
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:
|
|
19
|
-
attention_mask: Optional[np.
|
|
20
|
-
input_ids: Optional[np.
|
|
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,
|
|
37
|
-
) -> dict[str,
|
|
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]]:
|