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.
Files changed (74) hide show
  1. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/LICENSE +1 -1
  2. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/NOTICE +2 -0
  3. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/PKG-INFO +12 -12
  4. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/common/model_description.py +8 -7
  5. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/common/model_management.py +108 -45
  6. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/common/onnx_model.py +32 -15
  7. fastembed_gpu-0.8.1/fastembed/common/preprocessor_utils.py +151 -0
  8. fastembed_gpu-0.8.1/fastembed/common/types.py +27 -0
  9. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/common/utils.py +12 -2
  10. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/embedding.py +3 -3
  11. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/image/image_embedding.py +17 -11
  12. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/image/image_embedding_base.py +6 -6
  13. fastembed_gpu-0.8.1/fastembed/image/normalized_embedding.py +69 -0
  14. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/image/onnx_embedding.py +16 -15
  15. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/image/onnx_image_model.py +22 -18
  16. fastembed_gpu-0.8.1/fastembed/image/siglip_embedding.py +44 -0
  17. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/image/transform/functional.py +99 -14
  18. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/image/transform/operators.py +243 -13
  19. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/late_interaction/colbert.py +25 -22
  20. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/late_interaction/late_interaction_embedding_base.py +8 -8
  21. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/late_interaction/late_interaction_text_embedding.py +12 -12
  22. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/late_interaction/token_embeddings.py +3 -3
  23. fastembed_gpu-0.8.1/fastembed/late_interaction_multimodal/colmodernvbert.py +532 -0
  24. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/late_interaction_multimodal/colpali.py +19 -18
  25. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/late_interaction_multimodal/late_interaction_multimodal_embedding.py +18 -14
  26. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/late_interaction_multimodal/late_interaction_multimodal_embedding_base.py +9 -9
  27. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/late_interaction_multimodal/onnx_multimodal_model.py +28 -26
  28. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/parallel_processor.py +10 -9
  29. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/postprocess/muvera.py +1 -3
  30. fastembed_gpu-0.8.1/fastembed/rerank/cross_encoder/custom_text_cross_encoder.py +78 -0
  31. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/rerank/cross_encoder/onnx_text_cross_encoder.py +15 -13
  32. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/rerank/cross_encoder/onnx_text_model.py +15 -14
  33. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/rerank/cross_encoder/text_cross_encoder.py +9 -8
  34. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/rerank/cross_encoder/text_cross_encoder_base.py +4 -4
  35. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/sparse/bm25.py +10 -12
  36. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/sparse/bm42.py +18 -18
  37. fastembed_gpu-0.8.1/fastembed/sparse/if_splade.py +246 -0
  38. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/sparse/minicoil.py +22 -22
  39. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/sparse/sparse_embedding_base.py +8 -10
  40. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/sparse/sparse_text_embedding.py +19 -13
  41. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/sparse/splade_pp.py +17 -15
  42. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/sparse/utils/sparse_vectors_converter.py +19 -22
  43. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/sparse/utils/vocab_resolver.py +10 -3
  44. fastembed_gpu-0.8.1/fastembed/text/builtin_sentence_embedding.py +69 -0
  45. fastembed_gpu-0.8.1/fastembed/text/custom_text_embedding.py +144 -0
  46. fastembed_gpu-0.8.1/fastembed/text/last_token_normalized_embedding.py +94 -0
  47. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/text/multitask_embedding.py +8 -10
  48. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/text/onnx_embedding.py +56 -19
  49. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/text/onnx_text_model.py +28 -22
  50. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/text/pooled_normalized_embedding.py +2 -2
  51. fastembed_gpu-0.8.1/fastembed/text/siglip_embedding.py +71 -0
  52. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/text/text_embedding.py +21 -29
  53. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/text/text_embedding_base.py +8 -8
  54. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/pyproject.toml +21 -20
  55. fastembed_gpu-0.7.4/fastembed/common/preprocessor_utils.py +0 -83
  56. fastembed_gpu-0.7.4/fastembed/common/types.py +0 -25
  57. fastembed_gpu-0.7.4/fastembed/rerank/cross_encoder/custom_text_cross_encoder.py +0 -46
  58. fastembed_gpu-0.7.4/fastembed/text/custom_text_embedding.py +0 -98
  59. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/README.md +0 -0
  60. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/__init__.py +0 -0
  61. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/common/__init__.py +0 -0
  62. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/image/__init__.py +0 -0
  63. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/late_interaction/__init__.py +0 -0
  64. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/late_interaction/jina_colbert.py +0 -0
  65. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/late_interaction_multimodal/__init__.py +0 -0
  66. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/postprocess/__init__.py +0 -0
  67. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/py.typed +0 -0
  68. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/rerank/cross_encoder/__init__.py +0 -0
  69. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/sparse/__init__.py +0 -0
  70. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/sparse/utils/minicoil_encoder.py +0 -0
  71. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/sparse/utils/tokenizer.py +0 -0
  72. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/text/__init__.py +0 -0
  73. {fastembed_gpu-0.7.4 → fastembed_gpu-0.8.1}/fastembed/text/clip_embedding.py +0 -0
  74. {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 [yyyy] [name of copyright owner]
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.7.4
3
+ Version: 0.8.1
4
4
  Summary: Fast, light, accurate library built for retrieval embedding generation
5
- License: Apache License
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.9.0
12
- Classifier: License :: Other/Proprietary 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.20.0) ; python_version >= "3.13"
30
- Requires-Dist: onnxruntime-gpu (>=1.17.0,!=1.20.0) ; python_version >= "3.10" and python_version < "3.13"
31
- Requires-Dist: onnxruntime-gpu (>=1.17.0,<1.20.0) ; python_version < "3.10"
32
- Requires-Dist: pillow (>=10.3.0,<11.0) ; python_version < "3.10"
33
- Requires-Dist: pillow (>=10.3.0,<12.0) ; python_version >= "3.10" and python_version < "3.13"
34
- Requires-Dist: pillow (>=11.0.0,<12.0) ; python_version >= "3.13"
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 Optional, Any
3
+ from typing import Any
4
4
 
5
5
 
6
6
  @dataclass(frozen=True)
7
7
  class ModelSource:
8
- hf: Optional[str] = None
9
- url: Optional[str] = None
8
+ hf: str | None = None
9
+ url: str | None = None
10
10
  _deprecated_tar_struct: bool = False
11
11
 
12
12
  @property
@@ -33,8 +33,8 @@ class BaseModelDescription:
33
33
 
34
34
  @dataclass(frozen=True)
35
35
  class DenseModelDescription(BaseModelDescription):
36
- dim: Optional[int] = None
37
- tasks: Optional[dict[str, Any]] = field(default_factory=dict)
36
+ dim: int | None = None
37
+ tasks: dict[str, Any] | None = field(default_factory=dict)
38
38
 
39
39
  def __post_init__(self) -> None:
40
40
  assert self.dim is not None, "dim is required for dense model description"
@@ -42,11 +42,12 @@ class DenseModelDescription(BaseModelDescription):
42
42
 
43
43
  @dataclass(frozen=True)
44
44
  class SparseModelDescription(BaseModelDescription):
45
- requires_idf: Optional[bool] = None
46
- vocab_size: Optional[int] = None
45
+ requires_idf: bool | None = None
46
+ vocab_size: int | None = None
47
47
 
48
48
 
49
49
  class PoolingType(str, Enum):
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, Optional, Union, TypeVar, Generic
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
- if os.path.exists(output_path):
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=1024):
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, Union[int, str]]]:
184
- meta: dict[str, dict[str, Union[int, 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
- # Extract all files into the cache directory
310
- tar.extractall(
311
- path=cache_dir,
312
- )
313
- except tarfile.TarError as e:
314
- # If any error occurs while opening or extracting the tar.gz file,
315
- # delete the cache directory (if it was created in this function)
316
- # and raise the error again
317
- if "tmp" in cache_dir:
318
- shutil.rmtree(cache_dir)
319
- raise ValueError(f"An error occurred while decompressing {targz_path}: {e}")
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 model_tmp_dir.exists():
342
- shutil.rmtree(model_tmp_dir)
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
- model_tar_gz = Path(cache_dir) / f"{fast_model_name}.tar.gz"
347
-
348
- if model_tar_gz.exists():
349
- model_tar_gz.unlink()
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(cache_tmp_dir))
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
- model_tar_gz.unlink()
361
- # Rename from tmp to final name is atomic
362
- model_tmp_dir.rename(model_dir)
363
- else:
364
- logger.error(
365
- f"Could not find the model tar.gz file at {model_dir} and local_files_only=True."
366
- )
367
- raise ValueError(
368
- f"Could not find the model tar.gz file at {model_dir} and local_files_only=True."
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
- specific_model_path: Optional[str] = kwargs.pop("specific_model_path", None)
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
- return Path(
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, Optional, Sequence, Type, TypeVar
4
+ from typing import Any, Generic, Iterable, Sequence, Type, TypeVar
5
5
 
6
6
  import numpy as np
7
7
  import onnxruntime as ort
@@ -9,7 +9,7 @@ import onnxruntime as ort
9
9
  from numpy.typing import NDArray
10
10
  from tokenizers import Tokenizer
11
11
 
12
- from fastembed.common.types import OnnxProvider, NumpyArray
12
+ from fastembed.common.types import OnnxProvider, NumpyArray, Device
13
13
  from fastembed.parallel_processor import Worker
14
14
 
15
15
  # Holds type of the embedding result
@@ -19,8 +19,9 @@ T = TypeVar("T")
19
19
  @dataclass
20
20
  class OnnxOutputContext:
21
21
  model_output: NumpyArray
22
- attention_mask: Optional[NDArray[np.int64]] = None
23
- input_ids: Optional[NDArray[np.int64]] = None
22
+ attention_mask: NDArray[np.int64] | None = None
23
+ input_ids: NDArray[np.int64] | None = None
24
+ metadata: dict[str, Any] | None = None
24
25
 
25
26
 
26
27
  class OnnxModel(Generic[T]):
@@ -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: Optional[ort.InferenceSession] = None
47
- self.tokenizer: Optional[Tokenizer] = None
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: Optional[int],
62
- providers: Optional[Sequence[OnnxProvider]] = None,
63
- cuda: bool = False,
64
- device_id: Optional[int] = None,
65
- extra_session_options: Optional[dict[str, Any]] = None,
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 cuda and providers is not None:
86
+ if explicit_cuda and providers is not None:
71
87
  warnings.warn(
72
- f"`cuda` and `providers` are mutually exclusive parameters, cuda: {cuda}, providers: {providers}",
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
+ )