fastembed-gpu 0.7.4__tar.gz → 0.8.1__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.7.4 → fastembed_gpu-0.8.1}/LICENSE +1 -1
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/NOTICE +2 -0
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/PKG-INFO +12 -12
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/common/model_description.py +8 -7
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/common/model_management.py +108 -45
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/common/onnx_model.py +32 -15
- fastembed_gpu-0.8.1/fastembed/common/preprocessor_utils.py +151 -0
- fastembed_gpu-0.8.1/fastembed/common/types.py +27 -0
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/common/utils.py +12 -2
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/embedding.py +3 -3
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/image/image_embedding.py +17 -11
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/image/image_embedding_base.py +6 -6
- fastembed_gpu-0.8.1/fastembed/image/normalized_embedding.py +69 -0
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/image/onnx_embedding.py +16 -15
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/image/onnx_image_model.py +22 -18
- fastembed_gpu-0.8.1/fastembed/image/siglip_embedding.py +44 -0
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/image/transform/functional.py +99 -14
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/image/transform/operators.py +243 -13
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/late_interaction/colbert.py +25 -22
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/late_interaction/late_interaction_embedding_base.py +8 -8
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/late_interaction/late_interaction_text_embedding.py +12 -12
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/late_interaction/token_embeddings.py +3 -3
- fastembed_gpu-0.8.1/fastembed/late_interaction_multimodal/colmodernvbert.py +532 -0
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/late_interaction_multimodal/colpali.py +19 -18
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/late_interaction_multimodal/late_interaction_multimodal_embedding.py +18 -14
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/late_interaction_multimodal/late_interaction_multimodal_embedding_base.py +9 -9
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/late_interaction_multimodal/onnx_multimodal_model.py +28 -26
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/parallel_processor.py +10 -9
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/postprocess/muvera.py +1 -3
- fastembed_gpu-0.8.1/fastembed/rerank/cross_encoder/custom_text_cross_encoder.py +78 -0
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/rerank/cross_encoder/onnx_text_cross_encoder.py +15 -13
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/rerank/cross_encoder/onnx_text_model.py +15 -14
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/rerank/cross_encoder/text_cross_encoder.py +9 -8
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/rerank/cross_encoder/text_cross_encoder_base.py +4 -4
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/sparse/bm25.py +10 -12
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/sparse/bm42.py +18 -18
- fastembed_gpu-0.8.1/fastembed/sparse/if_splade.py +246 -0
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/sparse/minicoil.py +22 -22
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/sparse/sparse_embedding_base.py +8 -10
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/sparse/sparse_text_embedding.py +19 -13
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/sparse/splade_pp.py +17 -15
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/sparse/utils/sparse_vectors_converter.py +19 -22
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/sparse/utils/vocab_resolver.py +10 -3
- fastembed_gpu-0.8.1/fastembed/text/builtin_sentence_embedding.py +69 -0
- fastembed_gpu-0.8.1/fastembed/text/custom_text_embedding.py +144 -0
- fastembed_gpu-0.8.1/fastembed/text/last_token_normalized_embedding.py +94 -0
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/text/multitask_embedding.py +8 -10
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/text/onnx_embedding.py +56 -19
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/text/onnx_text_model.py +28 -22
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/text/pooled_normalized_embedding.py +2 -2
- fastembed_gpu-0.8.1/fastembed/text/siglip_embedding.py +71 -0
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/text/text_embedding.py +21 -29
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/text/text_embedding_base.py +8 -8
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/pyproject.toml +21 -20
- fastembed_gpu-0.7.4/fastembed/common/preprocessor_utils.py +0 -83
- fastembed_gpu-0.7.4/fastembed/common/types.py +0 -25
- fastembed_gpu-0.7.4/fastembed/rerank/cross_encoder/custom_text_cross_encoder.py +0 -46
- fastembed_gpu-0.7.4/fastembed/text/custom_text_embedding.py +0 -98
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/README.md +0 -0
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/__init__.py +0 -0
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/common/__init__.py +0 -0
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/image/__init__.py +0 -0
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/late_interaction/__init__.py +0 -0
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/late_interaction/jina_colbert.py +0 -0
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/late_interaction_multimodal/__init__.py +0 -0
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/postprocess/__init__.py +0 -0
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/py.typed +0 -0
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/rerank/cross_encoder/__init__.py +0 -0
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/sparse/__init__.py +0 -0
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/sparse/utils/minicoil_encoder.py +0 -0
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/sparse/utils/tokenizer.py +0 -0
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/text/__init__.py +0 -0
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/text/clip_embedding.py +0 -0
- {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/text/pooled_embedding.py +0 -0
|
@@ -186,7 +186,7 @@
|
|
|
186
186
|
same "printed page" as the copyright notice for easier
|
|
187
187
|
identification within third-party archives.
|
|
188
188
|
|
|
189
|
-
Copyright
|
|
189
|
+
Copyright 2026 Qdrant Solutions GmbH
|
|
190
190
|
|
|
191
191
|
Licensed under the Apache License, Version 2.0 (the "License");
|
|
192
192
|
you may not use this file except in compliance with the License.
|
|
@@ -15,6 +15,8 @@ These models are developed by Jina (https://jina.ai/) and are subject to Jina AI
|
|
|
15
15
|
This distribution includes the following Google models, each with its respective license:
|
|
16
16
|
- vidore/colpali-v1.3
|
|
17
17
|
- License: gemma
|
|
18
|
+
- google/embeddinggemma-300m
|
|
19
|
+
- License: gemma
|
|
18
20
|
|
|
19
21
|
Gemma is provided under and subject to the Gemma Terms of Use found at ai.google.dev/gemma/terms
|
|
20
22
|
|
|
@@ -1,37 +1,37 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: fastembed-gpu
|
|
3
|
-
Version: 0.
|
|
3
|
+
Version: 0.8.1
|
|
4
4
|
Summary: Fast, light, accurate library built for retrieval embedding generation
|
|
5
|
-
License: Apache
|
|
5
|
+
License: Apache-2.0
|
|
6
6
|
License-File: LICENSE
|
|
7
7
|
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.
|
|
12
|
-
Classifier: License ::
|
|
11
|
+
Requires-Python: >=3.10.0
|
|
12
|
+
Classifier: License :: OSI Approved :: Apache Software 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
|
|
18
17
|
Classifier: Programming Language :: Python :: 3.13
|
|
19
18
|
Classifier: Programming Language :: Python :: 3.14
|
|
19
|
+
Classifier: Programming Language :: Python :: 3.15
|
|
20
20
|
Requires-Dist: huggingface-hub (>=0.20,<2.0)
|
|
21
21
|
Requires-Dist: loguru (>=0.7.2,<0.8.0)
|
|
22
22
|
Requires-Dist: mmh3 (>=4.1.0,<6.0.0)
|
|
23
23
|
Requires-Dist: numpy (>=1.21) ; python_version == "3.11"
|
|
24
|
-
Requires-Dist: numpy (>=1.21,<2.1.0) ; python_version < "3.10"
|
|
25
24
|
Requires-Dist: numpy (>=1.21,<2.3.0) ; python_version == "3.10"
|
|
26
25
|
Requires-Dist: numpy (>=1.26) ; python_version == "3.12"
|
|
27
26
|
Requires-Dist: numpy (>=2.1.0) ; python_version == "3.13"
|
|
28
27
|
Requires-Dist: numpy (>=2.3.0) ; python_version >= "3.14"
|
|
29
|
-
Requires-Dist: onnxruntime-gpu (>1.
|
|
30
|
-
Requires-Dist: onnxruntime-gpu (>=1.17.0,!=1.20.0) ; python_version >= "3.
|
|
31
|
-
Requires-Dist: onnxruntime-gpu (>=1.17.0
|
|
32
|
-
Requires-Dist:
|
|
33
|
-
Requires-Dist: pillow (>=10.3.0,<
|
|
34
|
-
Requires-Dist: pillow (>=11.0.0,<
|
|
28
|
+
Requires-Dist: onnxruntime-gpu (>1.21.0,!=1.24.0,!=1.24.1) ; python_version == "3.13"
|
|
29
|
+
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"
|
|
30
|
+
Requires-Dist: onnxruntime-gpu (>=1.17.0,!=1.20.0,<1.24) ; python_version == "3.10"
|
|
31
|
+
Requires-Dist: onnxruntime-gpu (>=1.24.2) ; python_version >= "3.14"
|
|
32
|
+
Requires-Dist: pillow (>=10.3.0,<13.0) ; python_version >= "3.10" and python_version < "3.13"
|
|
33
|
+
Requires-Dist: pillow (>=11.0.0,<13.0) ; python_version == "3.13"
|
|
34
|
+
Requires-Dist: pillow (>=12.0.0,<13.0) ; python_version >= "3.14"
|
|
35
35
|
Requires-Dist: py-rust-stemmers (>=0.1.0,<0.2.0)
|
|
36
36
|
Requires-Dist: requests (>=2.31,<3.0)
|
|
37
37
|
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
|
|
3
|
+
from typing import Any
|
|
4
4
|
|
|
5
5
|
|
|
6
6
|
@dataclass(frozen=True)
|
|
7
7
|
class ModelSource:
|
|
8
|
-
hf:
|
|
9
|
-
url:
|
|
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:
|
|
37
|
-
tasks:
|
|
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,11 +42,12 @@ class DenseModelDescription(BaseModelDescription):
|
|
|
42
42
|
|
|
43
43
|
@dataclass(frozen=True)
|
|
44
44
|
class SparseModelDescription(BaseModelDescription):
|
|
45
|
-
requires_idf:
|
|
46
|
-
vocab_size:
|
|
45
|
+
requires_idf: bool | None = None
|
|
46
|
+
vocab_size: int | None = None
|
|
47
47
|
|
|
48
48
|
|
|
49
49
|
class PoolingType(str, Enum):
|
|
50
50
|
CLS = "CLS"
|
|
51
51
|
MEAN = "MEAN"
|
|
52
|
+
LAST_TOKEN = "LAST_TOKEN"
|
|
52
53
|
DISABLED = "DISABLED"
|
|
@@ -1,11 +1,14 @@
|
|
|
1
1
|
import os
|
|
2
2
|
import time
|
|
3
|
+
import gzip
|
|
3
4
|
import json
|
|
4
5
|
import shutil
|
|
5
6
|
import tarfile
|
|
7
|
+
import tempfile
|
|
8
|
+
import contextlib
|
|
6
9
|
from copy import deepcopy
|
|
7
|
-
from pathlib import Path
|
|
8
|
-
from typing import Any,
|
|
10
|
+
from pathlib import Path, PureWindowsPath
|
|
11
|
+
from typing import Any, TypeVar, Generic
|
|
9
12
|
|
|
10
13
|
import requests
|
|
11
14
|
from huggingface_hub import snapshot_download, model_info, list_repo_tree
|
|
@@ -21,6 +24,8 @@ from fastembed.common.model_description import BaseModelDescription
|
|
|
21
24
|
|
|
22
25
|
T = TypeVar("T", bound=BaseModelDescription)
|
|
23
26
|
|
|
27
|
+
_DOWNLOAD_CHUNK_SIZE = 256 * 1024
|
|
28
|
+
|
|
24
29
|
|
|
25
30
|
class ModelManagement(Generic[T]):
|
|
26
31
|
METADATA_FILE = "files_metadata.json"
|
|
@@ -97,9 +102,7 @@ class ModelManagement(Generic[T]):
|
|
|
97
102
|
str: The path to the downloaded file.
|
|
98
103
|
"""
|
|
99
104
|
|
|
100
|
-
|
|
101
|
-
return output_path
|
|
102
|
-
response = requests.get(url, stream=True)
|
|
105
|
+
response = requests.get(url, stream=True, timeout=(10, 120))
|
|
103
106
|
|
|
104
107
|
# Handle HTTP errors
|
|
105
108
|
if response.status_code == 403:
|
|
@@ -107,6 +110,8 @@ class ModelManagement(Generic[T]):
|
|
|
107
110
|
"Authentication Error: You do not have permission to access this resource. "
|
|
108
111
|
"Please check your credentials."
|
|
109
112
|
)
|
|
113
|
+
# Otherwise an error page gets written out as though it were the archive.
|
|
114
|
+
response.raise_for_status()
|
|
110
115
|
|
|
111
116
|
# Get the total size of the file
|
|
112
117
|
total_size_in_bytes = int(response.headers.get("content-length", 0))
|
|
@@ -124,7 +129,7 @@ class ModelManagement(Generic[T]):
|
|
|
124
129
|
disable=not show_progress,
|
|
125
130
|
) as progress_bar:
|
|
126
131
|
with open(output_path, "wb") as file:
|
|
127
|
-
for chunk in response.iter_content(chunk_size=
|
|
132
|
+
for chunk in response.iter_content(chunk_size=_DOWNLOAD_CHUNK_SIZE):
|
|
128
133
|
if chunk: # Filter out keep-alive new chunks
|
|
129
134
|
progress_bar.update(len(chunk))
|
|
130
135
|
file.write(chunk)
|
|
@@ -180,8 +185,8 @@ class ModelManagement(Generic[T]):
|
|
|
180
185
|
|
|
181
186
|
def _collect_file_metadata(
|
|
182
187
|
model_dir: Path, repo_files: list[RepoFile]
|
|
183
|
-
) -> dict[str, dict[str,
|
|
184
|
-
meta: dict[str, dict[str,
|
|
188
|
+
) -> dict[str, dict[str, int | str]]:
|
|
189
|
+
meta: dict[str, dict[str, int | str]] = {}
|
|
185
190
|
file_info_map = {f.path: f for f in repo_files}
|
|
186
191
|
for file_path in model_dir.rglob("*"):
|
|
187
192
|
if file_path.is_file() and file_path.name != cls.METADATA_FILE:
|
|
@@ -193,9 +198,7 @@ class ModelManagement(Generic[T]):
|
|
|
193
198
|
}
|
|
194
199
|
return meta
|
|
195
200
|
|
|
196
|
-
def _save_file_metadata(
|
|
197
|
-
model_dir: Path, meta: dict[str, dict[str, Union[int, str]]]
|
|
198
|
-
) -> None:
|
|
201
|
+
def _save_file_metadata(model_dir: Path, meta: dict[str, dict[str, int | str]]) -> None:
|
|
199
202
|
try:
|
|
200
203
|
if not model_dir.exists():
|
|
201
204
|
model_dir.mkdir(parents=True, exist_ok=True)
|
|
@@ -288,12 +291,18 @@ class ModelManagement(Generic[T]):
|
|
|
288
291
|
"""
|
|
289
292
|
Decompresses a .tar.gz file to a cache directory.
|
|
290
293
|
|
|
294
|
+
Nothing is deleted on failure, since `cache_dir` may hold more than this archive.
|
|
295
|
+
Cleaning up a partial extraction is the caller's job.
|
|
296
|
+
|
|
291
297
|
Args:
|
|
292
298
|
targz_path (str): Path to the .tar.gz file.
|
|
293
299
|
cache_dir (str): Path to the cache directory.
|
|
294
300
|
|
|
295
301
|
Returns:
|
|
296
302
|
cache_dir (str): Path to the cache directory.
|
|
303
|
+
|
|
304
|
+
Raises:
|
|
305
|
+
ValueError: If the archive is missing, corrupt, or holds an unsafe member.
|
|
297
306
|
"""
|
|
298
307
|
# Check if targz_path exists and is a file
|
|
299
308
|
if not os.path.isfile(targz_path):
|
|
@@ -306,20 +315,50 @@ class ModelManagement(Generic[T]):
|
|
|
306
315
|
try:
|
|
307
316
|
# Open the tar.gz file
|
|
308
317
|
with tarfile.open(targz_path, "r:gz") as tar:
|
|
309
|
-
|
|
310
|
-
|
|
311
|
-
|
|
312
|
-
|
|
313
|
-
|
|
314
|
-
|
|
315
|
-
|
|
316
|
-
|
|
317
|
-
|
|
318
|
-
|
|
319
|
-
|
|
318
|
+
if hasattr(tarfile, "data_filter"):
|
|
319
|
+
tar.extractall(path=cache_dir, filter="data")
|
|
320
|
+
else:
|
|
321
|
+
# No PEP 706 filter before 3.10.12, so vet the members by hand.
|
|
322
|
+
members = tar.getmembers()
|
|
323
|
+
for member in members:
|
|
324
|
+
cls._validate_tar_member(member)
|
|
325
|
+
tar.extractall(path=cache_dir, members=members)
|
|
326
|
+
# tarfile stops at the end-of-archive marker, short of the gzip trailer, so
|
|
327
|
+
# the CRC is only checked if the rest of the stream is read.
|
|
328
|
+
while tar.fileobj.read(1 << 20):
|
|
329
|
+
pass
|
|
330
|
+
except (tarfile.TarError, ValueError, EOFError, gzip.BadGzipFile) as e:
|
|
331
|
+
# gzip raises EOFError for a truncated stream and BadGzipFile for a corrupted one.
|
|
332
|
+
raise ValueError(f"An error occurred while decompressing {targz_path}: {e}") from e
|
|
320
333
|
|
|
321
334
|
return cache_dir
|
|
322
335
|
|
|
336
|
+
@staticmethod
|
|
337
|
+
def _is_unsafe_tar_path(path: str) -> bool:
|
|
338
|
+
"""Checks whether a tar member name or link target may escape the extraction dir.
|
|
339
|
+
|
|
340
|
+
Lexical on purpose: resolving against the extraction directory is unsound before
|
|
341
|
+
extraction, since `link/../escape` only escapes once an earlier member has been
|
|
342
|
+
written as a symlink. Any `..` component is therefore rejected outright.
|
|
343
|
+
"""
|
|
344
|
+
# PureWindowsPath splits on both separators, so `root` covers POSIX "/evil" as
|
|
345
|
+
# well as "\\evil", which escapes on Windows without being absolute.
|
|
346
|
+
windows_path = PureWindowsPath(path)
|
|
347
|
+
return bool(windows_path.drive or windows_path.root) or ".." in windows_path.parts
|
|
348
|
+
|
|
349
|
+
@classmethod
|
|
350
|
+
def _validate_tar_member(cls, member: tarfile.TarInfo) -> None:
|
|
351
|
+
"""Raises ValueError if a member could write outside the extraction directory."""
|
|
352
|
+
if cls._is_unsafe_tar_path(member.name):
|
|
353
|
+
raise ValueError(f"Unsafe tar member path: {member.name}")
|
|
354
|
+
|
|
355
|
+
if member.issym() or member.islnk():
|
|
356
|
+
if cls._is_unsafe_tar_path(member.linkname):
|
|
357
|
+
raise ValueError(f"Unsafe tar link target: {member.name} -> {member.linkname}")
|
|
358
|
+
elif not (member.isfile() or member.isdir()):
|
|
359
|
+
# Devices, fifos and the like have no place in a model archive.
|
|
360
|
+
raise ValueError(f"Unsupported tar member type: {member.name}")
|
|
361
|
+
|
|
323
362
|
@classmethod
|
|
324
363
|
def retrieve_model_gcs(
|
|
325
364
|
cls,
|
|
@@ -331,42 +370,58 @@ class ModelManagement(Generic[T]):
|
|
|
331
370
|
) -> Path:
|
|
332
371
|
fast_model_name = f"{'fast-' if deprecated_tar_struct else ''}{model_name.split('/')[-1]}"
|
|
333
372
|
cache_tmp_dir = Path(cache_dir) / "tmp"
|
|
334
|
-
model_tmp_dir = cache_tmp_dir / fast_model_name
|
|
335
373
|
model_dir = Path(cache_dir) / fast_model_name
|
|
336
374
|
|
|
337
375
|
# check if the model_dir and the model files are both present for macOS
|
|
338
376
|
if model_dir.exists() and len(list(model_dir.glob("*"))) > 0:
|
|
339
377
|
return model_dir
|
|
340
378
|
|
|
341
|
-
if
|
|
342
|
-
|
|
379
|
+
if local_files_only:
|
|
380
|
+
logger.error(
|
|
381
|
+
f"Could not find the model tar.gz file at {model_dir} and local_files_only=True."
|
|
382
|
+
)
|
|
383
|
+
raise ValueError(
|
|
384
|
+
f"Could not find the model tar.gz file at {model_dir} and local_files_only=True."
|
|
385
|
+
)
|
|
343
386
|
|
|
387
|
+
if cache_tmp_dir.is_symlink():
|
|
388
|
+
raise ValueError(
|
|
389
|
+
f"{cache_tmp_dir} is a symlink, refusing to stage downloads through it"
|
|
390
|
+
)
|
|
344
391
|
cache_tmp_dir.mkdir(parents=True, exist_ok=True)
|
|
345
392
|
|
|
346
|
-
|
|
347
|
-
|
|
348
|
-
|
|
349
|
-
|
|
350
|
-
|
|
351
|
-
if not local_files_only:
|
|
393
|
+
# The archive and everything extracted from it go in a directory of this attempt's own,
|
|
394
|
+
# so removing it undoes the attempt without touching any other download of the model.
|
|
395
|
+
staging_dir = Path(tempfile.mkdtemp(dir=cache_tmp_dir, prefix=f"{fast_model_name}-"))
|
|
396
|
+
try:
|
|
397
|
+
model_tar_gz = staging_dir / f"{fast_model_name}.tar.gz"
|
|
352
398
|
cls.download_file_from_gcs(
|
|
353
399
|
source_url,
|
|
354
400
|
output_path=str(model_tar_gz),
|
|
355
401
|
)
|
|
356
402
|
|
|
357
|
-
cls.decompress_to_cache(targz_path=str(model_tar_gz), cache_dir=str(
|
|
358
|
-
assert model_tmp_dir.exists(), f"Could not find {model_tmp_dir} in {cache_tmp_dir}"
|
|
403
|
+
cls.decompress_to_cache(targz_path=str(model_tar_gz), cache_dir=str(staging_dir))
|
|
359
404
|
|
|
360
|
-
|
|
361
|
-
|
|
362
|
-
|
|
363
|
-
|
|
364
|
-
|
|
365
|
-
|
|
366
|
-
|
|
367
|
-
|
|
368
|
-
|
|
369
|
-
|
|
405
|
+
model_tmp_dir = staging_dir / fast_model_name
|
|
406
|
+
if not model_tmp_dir.is_dir() or model_tmp_dir.is_symlink():
|
|
407
|
+
raise ValueError(
|
|
408
|
+
f"The archive from {source_url} has no {fast_model_name} directory"
|
|
409
|
+
)
|
|
410
|
+
|
|
411
|
+
# Replace a stale empty model_dir, which Windows will not rename onto. rmdir leaves
|
|
412
|
+
# anything else alone, including one another download has just filled.
|
|
413
|
+
with contextlib.suppress(OSError):
|
|
414
|
+
model_dir.rmdir()
|
|
415
|
+
|
|
416
|
+
try:
|
|
417
|
+
# Rename from the staging dir to the final name is atomic
|
|
418
|
+
model_tmp_dir.rename(model_dir)
|
|
419
|
+
except OSError:
|
|
420
|
+
# Another download of the same model finished first, so keep its copy.
|
|
421
|
+
if not (model_dir.is_dir() and any(model_dir.iterdir())):
|
|
422
|
+
raise
|
|
423
|
+
finally:
|
|
424
|
+
shutil.rmtree(staging_dir, ignore_errors=True)
|
|
370
425
|
|
|
371
426
|
return model_dir
|
|
372
427
|
|
|
@@ -397,7 +452,11 @@ class ModelManagement(Generic[T]):
|
|
|
397
452
|
Path: The path to the downloaded model directory.
|
|
398
453
|
"""
|
|
399
454
|
local_files_only = kwargs.get("local_files_only", False)
|
|
400
|
-
|
|
455
|
+
hf_offline = os.environ.get("HF_HUB_OFFLINE", "").strip().upper()
|
|
456
|
+
if not local_files_only and hf_offline in {"1", "TRUE", "YES", "ON"}:
|
|
457
|
+
local_files_only = True
|
|
458
|
+
kwargs["local_files_only"] = True
|
|
459
|
+
specific_model_path: str | None = kwargs.pop("specific_model_path", None)
|
|
401
460
|
if specific_model_path:
|
|
402
461
|
return Path(specific_model_path)
|
|
403
462
|
retries = 1 if local_files_only else retries
|
|
@@ -411,7 +470,7 @@ class ModelManagement(Generic[T]):
|
|
|
411
470
|
try:
|
|
412
471
|
cache_kwargs = deepcopy(kwargs)
|
|
413
472
|
cache_kwargs["local_files_only"] = True
|
|
414
|
-
|
|
473
|
+
resolved_path = Path(
|
|
415
474
|
cls.download_files_from_huggingface(
|
|
416
475
|
hf_source,
|
|
417
476
|
cache_dir=cache_dir,
|
|
@@ -419,6 +478,10 @@ class ModelManagement(Generic[T]):
|
|
|
419
478
|
**cache_kwargs,
|
|
420
479
|
)
|
|
421
480
|
)
|
|
481
|
+
if (resolved_path / model.model_file).exists() and all(
|
|
482
|
+
(resolved_path / file).exists() for file in extra_patterns
|
|
483
|
+
):
|
|
484
|
+
return resolved_path
|
|
422
485
|
except Exception:
|
|
423
486
|
pass
|
|
424
487
|
finally:
|
|
@@ -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,
|
|
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:
|
|
23
|
-
input_ids:
|
|
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]):
|
|
@@ -30,6 +31,18 @@ class OnnxModel(Generic[T]):
|
|
|
30
31
|
def _get_worker_class(cls) -> Type["EmbeddingWorker[T]"]:
|
|
31
32
|
raise NotImplementedError("Subclasses must implement this method")
|
|
32
33
|
|
|
34
|
+
def _get_worker_init_kwargs(self) -> dict[str, Any]:
|
|
35
|
+
"""Additional kwargs a worker process needs to reconstruct this model.
|
|
36
|
+
|
|
37
|
+
Workers are started with `spawn`/`forkserver`, hence they don't inherit class-level state
|
|
38
|
+
which has been set up in runtime, e.g. models registered via `add_custom_model`.
|
|
39
|
+
Such state has to be shipped to the workers explicitly.
|
|
40
|
+
|
|
41
|
+
Returns:
|
|
42
|
+
dict[str, Any]: kwargs to pass to `_get_worker_class().init_embedding`.
|
|
43
|
+
"""
|
|
44
|
+
return {}
|
|
45
|
+
|
|
33
46
|
def _post_process_onnx_output(self, output: OnnxOutputContext, **kwargs: Any) -> Iterable[T]:
|
|
34
47
|
"""Post-process the ONNX model output to convert it into a usable format.
|
|
35
48
|
|
|
@@ -43,8 +56,8 @@ class OnnxModel(Generic[T]):
|
|
|
43
56
|
raise NotImplementedError("Subclasses must implement this method")
|
|
44
57
|
|
|
45
58
|
def __init__(self) -> None:
|
|
46
|
-
self.model:
|
|
47
|
-
self.tokenizer:
|
|
59
|
+
self.model: ort.InferenceSession | None = None
|
|
60
|
+
self.tokenizer: Tokenizer | None = None
|
|
48
61
|
|
|
49
62
|
def _preprocess_onnx_input(
|
|
50
63
|
self, onnx_input: dict[str, NumpyArray], **kwargs: Any
|
|
@@ -58,25 +71,30 @@ class OnnxModel(Generic[T]):
|
|
|
58
71
|
self,
|
|
59
72
|
model_dir: Path,
|
|
60
73
|
model_file: str,
|
|
61
|
-
threads:
|
|
62
|
-
providers:
|
|
63
|
-
cuda: bool =
|
|
64
|
-
device_id:
|
|
65
|
-
extra_session_options:
|
|
74
|
+
threads: int | None,
|
|
75
|
+
providers: Sequence[OnnxProvider] | None = None,
|
|
76
|
+
cuda: bool | Device = Device.AUTO,
|
|
77
|
+
device_id: int | None = None,
|
|
78
|
+
extra_session_options: dict[str, Any] | None = None,
|
|
66
79
|
) -> None:
|
|
67
80
|
model_path = model_dir / model_file
|
|
68
81
|
# List of Execution Providers: https://onnxruntime.ai/docs/execution-providers
|
|
82
|
+
available_providers = ort.get_available_providers()
|
|
83
|
+
cuda_available = "CUDAExecutionProvider" in available_providers
|
|
84
|
+
explicit_cuda = cuda is True or cuda == Device.CUDA
|
|
69
85
|
|
|
70
|
-
if
|
|
86
|
+
if explicit_cuda and providers is not None:
|
|
71
87
|
warnings.warn(
|
|
72
|
-
f"`cuda` and `providers` are mutually exclusive parameters,
|
|
88
|
+
f"`cuda` and `providers` are mutually exclusive parameters, "
|
|
89
|
+
f"cuda: {cuda}, providers: {providers}. If you'd like to use providers, cuda should be one of "
|
|
90
|
+
f"[False, Device.CPU, Device.AUTO].",
|
|
73
91
|
category=UserWarning,
|
|
74
92
|
stacklevel=6,
|
|
75
93
|
)
|
|
76
94
|
|
|
77
95
|
if providers is not None:
|
|
78
96
|
onnx_providers = list(providers)
|
|
79
|
-
elif cuda:
|
|
97
|
+
elif explicit_cuda or (cuda == Device.AUTO and cuda_available):
|
|
80
98
|
if device_id is None:
|
|
81
99
|
onnx_providers = ["CUDAExecutionProvider"]
|
|
82
100
|
else:
|
|
@@ -84,7 +102,6 @@ class OnnxModel(Generic[T]):
|
|
|
84
102
|
else:
|
|
85
103
|
onnx_providers = ["CPUExecutionProvider"]
|
|
86
104
|
|
|
87
|
-
available_providers = ort.get_available_providers()
|
|
88
105
|
requested_provider_names: list[str] = []
|
|
89
106
|
for provider in onnx_providers:
|
|
90
107
|
# check providers available
|
|
@@ -0,0 +1,151 @@
|
|
|
1
|
+
import json
|
|
2
|
+
import sys
|
|
3
|
+
from typing import Any, Iterator
|
|
4
|
+
from pathlib import Path
|
|
5
|
+
|
|
6
|
+
from tokenizers import AddedToken, Tokenizer
|
|
7
|
+
|
|
8
|
+
from fastembed.image.transform.operators import Compose
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
def load_special_tokens(model_dir: Path) -> dict[str, Any]:
|
|
12
|
+
"""Read special_tokens_map.json, treating an absent file as an empty map."""
|
|
13
|
+
tokens_map_path = model_dir / "special_tokens_map.json"
|
|
14
|
+
if not tokens_map_path.exists():
|
|
15
|
+
return {}
|
|
16
|
+
|
|
17
|
+
with open(str(tokens_map_path)) as tokens_map_file:
|
|
18
|
+
tokens_map = json.load(tokens_map_file)
|
|
19
|
+
|
|
20
|
+
return tokens_map
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
def iter_special_tokens(tokens_map: dict[str, Any]) -> Iterator[str | dict[str, Any]]:
|
|
24
|
+
"""Yield the individual tokens declared in a special tokens map.
|
|
25
|
+
|
|
26
|
+
Most keys hold one token, but `additional_special_tokens` holds a list of them.
|
|
27
|
+
"""
|
|
28
|
+
for value in tokens_map.values():
|
|
29
|
+
if isinstance(value, list):
|
|
30
|
+
yield from value
|
|
31
|
+
else:
|
|
32
|
+
yield value
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
def _valid_context(value: Any) -> int | None:
|
|
36
|
+
"""Return `value` if it can be used as a truncation limit, `None` otherwise.
|
|
37
|
+
|
|
38
|
+
Config files do not always carry a real limit: transformers writes `model_max_length` as
|
|
39
|
+
1e30 when the value is unknown, and some repos ship a 0 or a null. `enable_truncation`
|
|
40
|
+
raises an `OverflowError` on the former and silently produces empty encodings on the
|
|
41
|
+
latter, so both are rejected here rather than passed through.
|
|
42
|
+
"""
|
|
43
|
+
if isinstance(value, bool) or not isinstance(value, int):
|
|
44
|
+
return None
|
|
45
|
+
if not 0 < value <= sys.maxsize:
|
|
46
|
+
return None
|
|
47
|
+
return value
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
def _resolve_max_context(tokenizer_config: dict[str, Any], model_dir: Path) -> int:
|
|
51
|
+
"""Pick the truncation limit, preferring the stricter of the two tokenizer config keys.
|
|
52
|
+
|
|
53
|
+
`config.json:max_position_embeddings` deliberately is not used as a fallback: it is the size
|
|
54
|
+
of the position table, not the usable context, and the two differ per architecture, e.g.
|
|
55
|
+
roberta reports 514 for a usable 512.
|
|
56
|
+
"""
|
|
57
|
+
candidates = [
|
|
58
|
+
context
|
|
59
|
+
for context in (
|
|
60
|
+
_valid_context(tokenizer_config.get("model_max_length")),
|
|
61
|
+
_valid_context(tokenizer_config.get("max_length")),
|
|
62
|
+
)
|
|
63
|
+
if context is not None
|
|
64
|
+
]
|
|
65
|
+
if not candidates:
|
|
66
|
+
raise ValueError(
|
|
67
|
+
f"Could not determine the maximum context length for {model_dir}. Set a positive "
|
|
68
|
+
"`model_max_length` or `max_length` in tokenizer_config.json."
|
|
69
|
+
)
|
|
70
|
+
|
|
71
|
+
return min(candidates)
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
def load_tokenizer(model_dir: Path) -> tuple[Tokenizer, dict[str, int]]:
|
|
75
|
+
tokenizer_path = model_dir / "tokenizer.json"
|
|
76
|
+
if not tokenizer_path.exists():
|
|
77
|
+
raise ValueError(f"Could not find tokenizer.json in {model_dir}")
|
|
78
|
+
|
|
79
|
+
tokenizer_config_path = model_dir / "tokenizer_config.json"
|
|
80
|
+
if not tokenizer_config_path.exists():
|
|
81
|
+
raise ValueError(f"Could not find tokenizer_config.json in {model_dir}")
|
|
82
|
+
|
|
83
|
+
# config.json is optional: transformers v5 no longer writes it for every model.
|
|
84
|
+
config_path = model_dir / "config.json"
|
|
85
|
+
config: dict[str, Any] = {}
|
|
86
|
+
if config_path.exists():
|
|
87
|
+
with open(str(config_path)) as config_file:
|
|
88
|
+
config = json.load(config_file)
|
|
89
|
+
|
|
90
|
+
with open(str(tokenizer_config_path)) as tokenizer_config_file:
|
|
91
|
+
tokenizer_config = json.load(tokenizer_config_file)
|
|
92
|
+
|
|
93
|
+
max_context = _resolve_max_context(tokenizer_config, model_dir)
|
|
94
|
+
|
|
95
|
+
tokens_map = load_special_tokens(model_dir)
|
|
96
|
+
|
|
97
|
+
tokenizer = Tokenizer.from_file(str(tokenizer_path))
|
|
98
|
+
tokenizer.enable_truncation(max_length=max_context)
|
|
99
|
+
|
|
100
|
+
# Registered before the padding is resolved: the map may name a pad token that
|
|
101
|
+
# tokenizer.json does not carry, and it only gets an id once it is added.
|
|
102
|
+
for token in iter_special_tokens(tokens_map):
|
|
103
|
+
if isinstance(token, str):
|
|
104
|
+
tokenizer.add_special_tokens([token])
|
|
105
|
+
elif isinstance(token, dict):
|
|
106
|
+
tokenizer.add_special_tokens([AddedToken(**token)])
|
|
107
|
+
|
|
108
|
+
# Padding is always normalized to batch-longest. A serialized fixed length shorter than the
|
|
109
|
+
# truncation limit leaves longer encodings untouched, which produces ragged batches, and a
|
|
110
|
+
# fixed length equal to it pads every batch to the maximum. Direction and pad token metadata
|
|
111
|
+
# are taken from the serialized settings, since some models pad on the left.
|
|
112
|
+
padding = tokenizer.padding or {}
|
|
113
|
+
pad_token = padding.get("pad_token") or tokenizer_config.get("pad_token")
|
|
114
|
+
if pad_token is None:
|
|
115
|
+
raise ValueError(f"Could not find a pad token for {model_dir}")
|
|
116
|
+
|
|
117
|
+
# The vocabulary is the last resort, not a hardcoded 0: that silently disagrees with
|
|
118
|
+
# `pad_token` for every model whose pad token is not the first entry.
|
|
119
|
+
pad_id = padding.get("pad_id", config.get("pad_token_id"))
|
|
120
|
+
if pad_id is None:
|
|
121
|
+
pad_id = tokenizer.token_to_id(pad_token)
|
|
122
|
+
if pad_id is None:
|
|
123
|
+
raise ValueError(f"Could not resolve an id for the pad token {pad_token!r} in {model_dir}")
|
|
124
|
+
|
|
125
|
+
tokenizer.enable_padding(
|
|
126
|
+
direction=padding.get("direction", "right"),
|
|
127
|
+
pad_id=pad_id,
|
|
128
|
+
pad_type_id=padding.get("pad_type_id", 0),
|
|
129
|
+
pad_token=pad_token,
|
|
130
|
+
pad_to_multiple_of=padding.get("pad_to_multiple_of"),
|
|
131
|
+
length=None,
|
|
132
|
+
)
|
|
133
|
+
|
|
134
|
+
special_token_to_id = {
|
|
135
|
+
token.content: token_id
|
|
136
|
+
for token_id, token in tokenizer.get_added_tokens_decoder().items()
|
|
137
|
+
if token.special
|
|
138
|
+
}
|
|
139
|
+
|
|
140
|
+
return tokenizer, special_token_to_id
|
|
141
|
+
|
|
142
|
+
|
|
143
|
+
def load_preprocessor(model_dir: Path) -> Compose:
|
|
144
|
+
preprocessor_config_path = model_dir / "preprocessor_config.json"
|
|
145
|
+
if not preprocessor_config_path.exists():
|
|
146
|
+
raise ValueError(f"Could not find preprocessor_config.json in {model_dir}")
|
|
147
|
+
|
|
148
|
+
with open(str(preprocessor_config_path)) as preprocessor_config_file:
|
|
149
|
+
preprocessor_config = json.load(preprocessor_config_file)
|
|
150
|
+
transforms = Compose.from_config(preprocessor_config)
|
|
151
|
+
return transforms
|
|
@@ -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
|
+
)
|