fastembed 0.2.6__tar.gz → 0.3.0__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (42) hide show
  1. {fastembed-0.2.6 → fastembed-0.3.0}/PKG-INFO +46 -13
  2. {fastembed-0.2.6 → fastembed-0.3.0}/README.md +38 -8
  3. fastembed-0.3.0/fastembed/__init__.py +20 -0
  4. fastembed-0.3.0/fastembed/common/__init__.py +3 -0
  5. {fastembed-0.2.6 → fastembed-0.3.0}/fastembed/common/model_management.py +31 -21
  6. fastembed-0.3.0/fastembed/common/onnx_model.py +113 -0
  7. fastembed-0.2.6/fastembed/common/models.py → fastembed-0.3.0/fastembed/common/preprocessor_utils.py +35 -15
  8. fastembed-0.3.0/fastembed/common/types.py +14 -0
  9. {fastembed-0.2.6 → fastembed-0.3.0}/fastembed/common/utils.py +10 -0
  10. fastembed-0.3.0/fastembed/image/__init__.py +4 -0
  11. fastembed-0.3.0/fastembed/image/image_embedding.py +87 -0
  12. fastembed-0.3.0/fastembed/image/image_embedding_base.py +39 -0
  13. fastembed-0.3.0/fastembed/image/onnx_embedding.py +134 -0
  14. fastembed-0.3.0/fastembed/image/onnx_image_model.py +107 -0
  15. fastembed-0.3.0/fastembed/image/transform/functional.py +125 -0
  16. fastembed-0.3.0/fastembed/image/transform/operators.py +168 -0
  17. fastembed-0.3.0/fastembed/late_interaction/__init__.py +4 -0
  18. fastembed-0.3.0/fastembed/late_interaction/colbert.py +196 -0
  19. fastembed-0.3.0/fastembed/late_interaction/late_interaction_embedding_base.py +60 -0
  20. fastembed-0.3.0/fastembed/late_interaction/late_interaction_text_embedding.py +104 -0
  21. fastembed-0.3.0/fastembed/sparse/bm42.py +281 -0
  22. fastembed-0.3.0/fastembed/sparse/sparse_embedding_base.py +81 -0
  23. {fastembed-0.2.6 → fastembed-0.3.0}/fastembed/sparse/sparse_text_embedding.py +20 -2
  24. {fastembed-0.2.6 → fastembed-0.3.0}/fastembed/sparse/splade_pp.py +26 -17
  25. fastembed-0.3.0/fastembed/text/clip_embedding.py +49 -0
  26. {fastembed-0.2.6 → fastembed-0.3.0}/fastembed/text/e5_onnx_embedding.py +5 -2
  27. {fastembed-0.2.6 → fastembed-0.3.0}/fastembed/text/jina_onnx_embedding.py +10 -7
  28. {fastembed-0.2.6 → fastembed-0.3.0}/fastembed/text/onnx_embedding.py +99 -45
  29. fastembed-0.3.0/fastembed/text/onnx_text_model.py +123 -0
  30. {fastembed-0.2.6 → fastembed-0.3.0}/fastembed/text/text_embedding.py +8 -2
  31. {fastembed-0.2.6 → fastembed-0.3.0}/fastembed/text/text_embedding_base.py +1 -0
  32. {fastembed-0.2.6 → fastembed-0.3.0}/pyproject.toml +9 -5
  33. fastembed-0.2.6/fastembed/__init__.py +0 -7
  34. fastembed-0.2.6/fastembed/common/__init__.py +0 -0
  35. fastembed-0.2.6/fastembed/common/onnx_model.py +0 -136
  36. fastembed-0.2.6/fastembed/image/__init__.py +0 -0
  37. fastembed-0.2.6/fastembed/sparse/sparse_embedding_base.py +0 -43
  38. {fastembed-0.2.6 → fastembed-0.3.0}/LICENSE +0 -0
  39. {fastembed-0.2.6 → fastembed-0.3.0}/fastembed/embedding.py +0 -0
  40. {fastembed-0.2.6 → fastembed-0.3.0}/fastembed/parallel_processor.py +0 -0
  41. {fastembed-0.2.6 → fastembed-0.3.0}/fastembed/sparse/__init__.py +0 -0
  42. {fastembed-0.2.6 → fastembed-0.3.0}/fastembed/text/__init__.py +0 -0
@@ -1,12 +1,12 @@
1
1
  Metadata-Version: 2.1
2
2
  Name: fastembed
3
- Version: 0.2.6
3
+ Version: 0.3.0
4
4
  Summary: Fast, light, accurate library built for retrieval embedding generation
5
5
  Home-page: https://github.com/qdrant/fastembed
6
6
  License: Apache License
7
7
  Keywords: vector,embedding,neural,search,qdrant,sentence-transformers
8
- Author: NirantK
9
- Author-email: nirant.bits@gmail.com
8
+ Author: Qdrant Team
9
+ Author-email: info@qdrant.tech
10
10
  Requires-Python: >=3.8.0,<3.13
11
11
  Classifier: License :: Other/Proprietary License
12
12
  Classifier: Programming Language :: Python :: 3
@@ -15,14 +15,18 @@ Classifier: Programming Language :: Python :: 3.9
15
15
  Classifier: Programming Language :: Python :: 3.10
16
16
  Classifier: Programming Language :: Python :: 3.11
17
17
  Classifier: Programming Language :: Python :: 3.12
18
- Requires-Dist: huggingface-hub (>=0.20,<0.21)
18
+ Requires-Dist: PyStemmer (>=2.2.0,<3.0.0)
19
+ Requires-Dist: huggingface-hub (>=0.20,<1.0)
19
20
  Requires-Dist: loguru (>=0.7.2,<0.8.0)
21
+ Requires-Dist: mmh3 (>=4.0,<5.0)
20
22
  Requires-Dist: numpy (>=1.21) ; python_version < "3.12"
21
23
  Requires-Dist: numpy (>=1.26) ; python_version >= "3.12"
22
24
  Requires-Dist: onnx (>=1.15.0,<2.0.0)
23
25
  Requires-Dist: onnxruntime (>=1.17.0,<2.0.0)
26
+ Requires-Dist: pillow (>=10.3.0,<11.0.0)
24
27
  Requires-Dist: requests (>=2.31,<3.0)
25
- Requires-Dist: tokenizers (>=0.15.1,<0.16.0)
28
+ Requires-Dist: snowballstemmer (>=2.2.0,<3.0.0)
29
+ Requires-Dist: tokenizers (>=0.15,<1.0)
26
30
  Requires-Dist: tqdm (>=4.66,<5.0)
27
31
  Project-URL: Repository, https://github.com/qdrant/fastembed
28
32
  Description-Content-Type: text/markdown
@@ -31,7 +35,7 @@ Description-Content-Type: text/markdown
31
35
 
32
36
  FastEmbed is a lightweight, fast, Python library built for embedding generation. We [support popular text models](https://qdrant.github.io/fastembed/examples/Supported_Models/). Please [open a GitHub issue](https://github.com/qdrant/fastembed/issues/new) if you want us to add a new model.
33
37
 
34
- The default text embedding (`TextEmbedding`) model is Flag Embedding, presented in the [MTEB](https://huggingface.co/spaces/mteb/leaderboard) leaderboard. It supports "query" and "passage" prefixes for the input text. Here is an example for [Retrieval Embedding Generation](https://qdrant.github.io/fastembed/examples/Retrieval_with_FastEmbed/) and how to use [FastEmbed with Qdrant](https://qdrant.github.io/fastembed/examples/Usage_With_Qdrant/).
38
+ The default text embedding (`TextEmbedding`) model is Flag Embedding, presented in the [MTEB](https://huggingface.co/spaces/mteb/leaderboard) leaderboard. It supports "query" and "passage" prefixes for the input text. Here is an example for [Retrieval Embedding Generation](https://qdrant.github.io/fastembed/qdrant/Retrieval_with_FastEmbed/) and how to use [FastEmbed with Qdrant](https://qdrant.github.io/fastembed/qdrant/Usage_With_Qdrant/).
35
39
 
36
40
  ## 📈 Why FastEmbed?
37
41
 
@@ -43,12 +47,18 @@ The default text embedding (`TextEmbedding`) model is Flag Embedding, presented
43
47
 
44
48
  ## 🚀 Installation
45
49
 
46
- To install the FastEmbed library, pip works:
50
+ To install the FastEmbed library, pip works best. You can install it with or without GPU support:
47
51
 
48
52
  ```bash
49
53
  pip install fastembed
50
54
  ```
51
55
 
56
+ ### ⚡️ With GPU
57
+
58
+ ```bash
59
+ pip install fastembed-gpu
60
+ ```
61
+
52
62
  ## 📖 Quickstart
53
63
 
54
64
  ```python
@@ -71,6 +81,28 @@ embeddings_list = list(embedding_model.embed(documents))
71
81
  len(embeddings_list[0]) # Vector of 384 dimensions
72
82
  ```
73
83
 
84
+ ### ⚡️ FastEmbed on a GPU
85
+
86
+ FastEmbed supports running on GPU devices.
87
+ It requires installation of the `fastembed-gpu` package.
88
+
89
+ ```bash
90
+ pip install fastembed-gpu
91
+ ```
92
+
93
+ Check our [example](https://qdrant.github.io/fastembed/examples/FastEmbed_GPU/) for the detailed instructions and CUDA 12.x support.
94
+
95
+ ```python
96
+ from fastembed import TextEmbedding
97
+
98
+ embedding_model = TextEmbedding(
99
+ model_name="BAAI/bge-small-en-v1.5",
100
+ providers=["CUDAExecutionProvider"]
101
+ )
102
+ print("The model BAAI/bge-small-en-v1.5 is ready to use on a GPU.")
103
+
104
+ ```
105
+
74
106
  ## Usage with Qdrant
75
107
 
76
108
  Installation with Qdrant Client in Python:
@@ -79,7 +111,13 @@ Installation with Qdrant Client in Python:
79
111
  pip install qdrant-client[fastembed]
80
112
  ```
81
113
 
82
- You might have to use ```pip install 'qdrant-client[fastembed]'``` on zsh.
114
+ or
115
+
116
+ ```bash
117
+ pip install qdrant-client[fastembed-gpu]
118
+ ```
119
+
120
+ You might have to use quotes ```pip install 'qdrant-client[fastembed]'``` on zsh.
83
121
 
84
122
  ```python
85
123
  from qdrant_client import QdrantClient
@@ -115,8 +153,3 @@ search_result = client.query(
115
153
  )
116
154
  print(search_result)
117
155
  ```
118
-
119
- #### Similar Work
120
-
121
- Ilyas M. wrote about using [FlagEmbeddings with Optimum](https://twitter.com/IlysMoutawwakil/status/1705215192425288017) over CUDA.
122
-
@@ -2,7 +2,7 @@
2
2
 
3
3
  FastEmbed is a lightweight, fast, Python library built for embedding generation. We [support popular text models](https://qdrant.github.io/fastembed/examples/Supported_Models/). Please [open a GitHub issue](https://github.com/qdrant/fastembed/issues/new) if you want us to add a new model.
4
4
 
5
- The default text embedding (`TextEmbedding`) model is Flag Embedding, presented in the [MTEB](https://huggingface.co/spaces/mteb/leaderboard) leaderboard. It supports "query" and "passage" prefixes for the input text. Here is an example for [Retrieval Embedding Generation](https://qdrant.github.io/fastembed/examples/Retrieval_with_FastEmbed/) and how to use [FastEmbed with Qdrant](https://qdrant.github.io/fastembed/examples/Usage_With_Qdrant/).
5
+ The default text embedding (`TextEmbedding`) model is Flag Embedding, presented in the [MTEB](https://huggingface.co/spaces/mteb/leaderboard) leaderboard. It supports "query" and "passage" prefixes for the input text. Here is an example for [Retrieval Embedding Generation](https://qdrant.github.io/fastembed/qdrant/Retrieval_with_FastEmbed/) and how to use [FastEmbed with Qdrant](https://qdrant.github.io/fastembed/qdrant/Usage_With_Qdrant/).
6
6
 
7
7
  ## 📈 Why FastEmbed?
8
8
 
@@ -14,12 +14,18 @@ The default text embedding (`TextEmbedding`) model is Flag Embedding, presented
14
14
 
15
15
  ## 🚀 Installation
16
16
 
17
- To install the FastEmbed library, pip works:
17
+ To install the FastEmbed library, pip works best. You can install it with or without GPU support:
18
18
 
19
19
  ```bash
20
20
  pip install fastembed
21
21
  ```
22
22
 
23
+ ### ⚡️ With GPU
24
+
25
+ ```bash
26
+ pip install fastembed-gpu
27
+ ```
28
+
23
29
  ## 📖 Quickstart
24
30
 
25
31
  ```python
@@ -42,6 +48,28 @@ embeddings_list = list(embedding_model.embed(documents))
42
48
  len(embeddings_list[0]) # Vector of 384 dimensions
43
49
  ```
44
50
 
51
+ ### ⚡️ FastEmbed on a GPU
52
+
53
+ FastEmbed supports running on GPU devices.
54
+ It requires installation of the `fastembed-gpu` package.
55
+
56
+ ```bash
57
+ pip install fastembed-gpu
58
+ ```
59
+
60
+ Check our [example](https://qdrant.github.io/fastembed/examples/FastEmbed_GPU/) for the detailed instructions and CUDA 12.x support.
61
+
62
+ ```python
63
+ from fastembed import TextEmbedding
64
+
65
+ embedding_model = TextEmbedding(
66
+ model_name="BAAI/bge-small-en-v1.5",
67
+ providers=["CUDAExecutionProvider"]
68
+ )
69
+ print("The model BAAI/bge-small-en-v1.5 is ready to use on a GPU.")
70
+
71
+ ```
72
+
45
73
  ## Usage with Qdrant
46
74
 
47
75
  Installation with Qdrant Client in Python:
@@ -50,7 +78,13 @@ Installation with Qdrant Client in Python:
50
78
  pip install qdrant-client[fastembed]
51
79
  ```
52
80
 
53
- You might have to use ```pip install 'qdrant-client[fastembed]'``` on zsh.
81
+ or
82
+
83
+ ```bash
84
+ pip install qdrant-client[fastembed-gpu]
85
+ ```
86
+
87
+ You might have to use quotes ```pip install 'qdrant-client[fastembed]'``` on zsh.
54
88
 
55
89
  ```python
56
90
  from qdrant_client import QdrantClient
@@ -85,8 +119,4 @@ search_result = client.query(
85
119
  query_text="This is a query document"
86
120
  )
87
121
  print(search_result)
88
- ```
89
-
90
- #### Similar Work
91
-
92
- Ilyas M. wrote about using [FlagEmbeddings with Optimum](https://twitter.com/IlysMoutawwakil/status/1705215192425288017) over CUDA.
122
+ ```
@@ -0,0 +1,20 @@
1
+ import importlib.metadata
2
+
3
+ from fastembed.image import ImageEmbedding
4
+ from fastembed.text import TextEmbedding
5
+ from fastembed.sparse import SparseTextEmbedding, SparseEmbedding
6
+ from fastembed.late_interaction import LateInteractionTextEmbedding
7
+
8
+ try:
9
+ version = importlib.metadata.version("fastembed")
10
+ except importlib.metadata.PackageNotFoundError as _:
11
+ version = importlib.metadata.version("fastembed-gpu")
12
+
13
+ __version__ = version
14
+ __all__ = [
15
+ "TextEmbedding",
16
+ "SparseTextEmbedding",
17
+ "SparseEmbedding",
18
+ "ImageEmbedding",
19
+ "LateInteractionTextEmbedding",
20
+ ]
@@ -0,0 +1,3 @@
1
+ from fastembed.common.types import OnnxProvider, ImageInput, PathInput
2
+
3
+ __all__ = ["OnnxProvider", "ImageInput", "PathInput"]
@@ -11,23 +11,6 @@ from tqdm import tqdm
11
11
  from loguru import logger
12
12
 
13
13
 
14
- def locate_model_file(model_dir: Path, file_names: List[str]) -> Path:
15
- """
16
- Find model path for both TransformerJS style `onnx` subdirectory structure and direct model weights structure used
17
- by Optimum and Qdrant
18
- """
19
- if not model_dir.is_dir():
20
- raise ValueError(f"Provided model path '{model_dir}' is not a directory.")
21
-
22
- for file_name in file_names:
23
- file_paths = [path for path in model_dir.rglob(file_name) if path.is_file()]
24
-
25
- if file_paths:
26
- return file_paths[0]
27
-
28
- raise ValueError(f"Could not find either of {', '.join(file_names)} in {model_dir}")
29
-
30
-
31
14
  class ModelManagement:
32
15
  @classmethod
33
16
  def list_supported_models(cls) -> List[Dict[str, Any]]:
@@ -104,21 +87,37 @@ class ModelManagement:
104
87
 
105
88
  @classmethod
106
89
  def download_files_from_huggingface(
107
- cls, hf_source_repo: str, cache_dir: Optional[str] = None
90
+ cls,
91
+ hf_source_repo: str,
92
+ cache_dir: Optional[str] = None,
93
+ extra_patterns: Optional[List[str]] = None,
94
+ **kwargs,
108
95
  ) -> str:
109
96
  """
110
97
  Downloads a model from HuggingFace Hub.
111
98
  Args:
112
99
  hf_source_repo (str): Name of the model on HuggingFace Hub, e.g. "qdrant/all-MiniLM-L6-v2-onnx".
113
100
  cache_dir (Optional[str]): The path to the cache directory.
101
+ extra_patterns (Optional[List[str]]): extra patterns to allow in the snapshot download, typically
102
+ includes the required model files.
114
103
  Returns:
115
104
  Path: The path to the model directory.
116
105
  """
106
+ allow_patterns = [
107
+ "config.json",
108
+ "tokenizer.json",
109
+ "tokenizer_config.json",
110
+ "special_tokens_map.json",
111
+ "preprocessor_config.json",
112
+ ]
113
+ if extra_patterns is not None:
114
+ allow_patterns.extend(extra_patterns)
117
115
 
118
116
  return snapshot_download(
119
117
  repo_id=hf_source_repo,
120
- ignore_patterns=["model.safetensors", "pytorch_model.bin"],
118
+ allow_patterns=allow_patterns,
121
119
  cache_dir=cache_dir,
120
+ local_files_only=kwargs.get("local_files_only", False),
122
121
  )
123
122
 
124
123
  @classmethod
@@ -175,6 +174,9 @@ class ModelManagement:
175
174
 
176
175
  model_tar_gz = Path(cache_dir) / f"{fast_model_name}.tar.gz"
177
176
 
177
+ if model_tar_gz.exists():
178
+ model_tar_gz.unlink()
179
+
178
180
  cls.download_file_from_gcs(
179
181
  source_url,
180
182
  output_path=str(model_tar_gz),
@@ -190,7 +192,7 @@ class ModelManagement:
190
192
  return model_dir
191
193
 
192
194
  @classmethod
193
- def download_model(cls, model: Dict[str, Any], cache_dir: Path) -> Path:
195
+ def download_model(cls, model: Dict[str, Any], cache_dir: Path, **kwargs) -> Path:
194
196
  """
195
197
  Downloads a model from HuggingFace Hub or Google Cloud Storage.
196
198
 
@@ -219,9 +221,17 @@ class ModelManagement:
219
221
  url_source = model.get("sources", {}).get("url")
220
222
 
221
223
  if hf_source:
224
+ extra_patterns = [model["model_file"]]
225
+ extra_patterns.extend(model.get("additional_files", []))
226
+
222
227
  try:
223
228
  return Path(
224
- cls.download_files_from_huggingface(hf_source, cache_dir=str(cache_dir))
229
+ cls.download_files_from_huggingface(
230
+ hf_source,
231
+ cache_dir=str(cache_dir),
232
+ extra_patterns=extra_patterns,
233
+ local_files_only=kwargs.get("local_files_only", False),
234
+ )
225
235
  )
226
236
  except (EnvironmentError, RepositoryNotFoundError, ValueError) as e:
227
237
  logger.error(
@@ -0,0 +1,113 @@
1
+ from dataclasses import dataclass
2
+ from pathlib import Path
3
+ from typing import Any, Dict, Generic, Iterable, Optional, Tuple, Type, TypeVar, Sequence
4
+ import warnings
5
+
6
+ import numpy as np
7
+ import onnxruntime as ort
8
+
9
+ from fastembed.common.types import OnnxProvider
10
+ from fastembed.parallel_processor import Worker
11
+
12
+
13
+ # Holds type of the embedding result
14
+ T = TypeVar("T")
15
+
16
+
17
+ @dataclass
18
+ class OnnxOutputContext:
19
+ model_output: np.ndarray
20
+ attention_mask: Optional[np.ndarray] = None
21
+ input_ids: Optional[np.ndarray] = None
22
+
23
+
24
+ class OnnxModel(Generic[T]):
25
+ @classmethod
26
+ def _get_worker_class(cls) -> Type["EmbeddingWorker"]:
27
+ raise NotImplementedError("Subclasses must implement this method")
28
+
29
+ def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[T]:
30
+ raise NotImplementedError("Subclasses must implement this method")
31
+
32
+ def __init__(self) -> None:
33
+ self.model = None
34
+ self.tokenizer = None
35
+
36
+ def _preprocess_onnx_input(
37
+ self, onnx_input: Dict[str, np.ndarray], **kwargs
38
+ ) -> Dict[str, np.ndarray]:
39
+ """
40
+ Preprocess the onnx input.
41
+ """
42
+ return onnx_input
43
+
44
+ def load_onnx_model(
45
+ self,
46
+ model_dir: Path,
47
+ model_file: str,
48
+ threads: Optional[int],
49
+ providers: Optional[Sequence[OnnxProvider]] = None,
50
+ ) -> None:
51
+ model_path = model_dir / model_file
52
+ # List of Execution Providers: https://onnxruntime.ai/docs/execution-providers
53
+
54
+ onnx_providers = ["CPUExecutionProvider"] if providers is None else list(providers)
55
+ available_providers = ort.get_available_providers()
56
+ requested_provider_names = []
57
+ for provider in onnx_providers:
58
+ # check providers available
59
+ provider_name = provider if isinstance(provider, str) else provider[0]
60
+ requested_provider_names.append(provider_name)
61
+ if provider_name not in available_providers:
62
+ raise ValueError(
63
+ f"Provider {provider_name} is not available. Available providers: {available_providers}"
64
+ )
65
+
66
+ so = ort.SessionOptions()
67
+ so.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
68
+
69
+ if threads is not None:
70
+ so.intra_op_num_threads = threads
71
+ so.inter_op_num_threads = threads
72
+
73
+ self.model = ort.InferenceSession(
74
+ str(model_path), providers=onnx_providers, sess_options=so
75
+ )
76
+ if "CUDAExecutionProvider" in requested_provider_names:
77
+ current_providers = self.model.get_providers()
78
+ if "CUDAExecutionProvider" not in current_providers:
79
+ warnings.warn(
80
+ f"Attempt to set CUDAExecutionProvider failed. Current providers: {current_providers}."
81
+ "If you are using CUDA 12.x, install onnxruntime-gpu via "
82
+ "`pip install onnxruntime-gpu --extra-index-url https://aiinfra.pkgs.visualstudio.com/PublicPackages/_packaging/onnxruntime-cuda-12/pypi/simple/`",
83
+ RuntimeWarning,
84
+ )
85
+
86
+ def onnx_embed(self, *args, **kwargs) -> OnnxOutputContext:
87
+ raise NotImplementedError("Subclasses must implement this method")
88
+
89
+
90
+ class EmbeddingWorker(Worker):
91
+ def init_embedding(
92
+ self,
93
+ model_name: str,
94
+ cache_dir: str,
95
+ ) -> OnnxModel:
96
+ raise NotImplementedError()
97
+
98
+ def __init__(
99
+ self,
100
+ model_name: str,
101
+ cache_dir: str,
102
+ ):
103
+ self.model = self.init_embedding(model_name, cache_dir)
104
+
105
+ @classmethod
106
+ def start(cls, model_name: str, cache_dir: str, **kwargs: Any) -> "EmbeddingWorker":
107
+ return cls(
108
+ model_name=model_name,
109
+ cache_dir=cache_dir,
110
+ )
111
+
112
+ def process(self, items: Iterable[Tuple[int, Any]]) -> Iterable[Tuple[int, Any]]:
113
+ raise NotImplementedError("Subclasses must implement this method")
@@ -1,11 +1,24 @@
1
1
  import json
2
2
  from pathlib import Path
3
+ from typing import Tuple
3
4
 
4
- import numpy as np
5
5
  from tokenizers import Tokenizer, AddedToken
6
6
 
7
+ from fastembed.image.transform.operators import Compose
7
8
 
8
- def load_tokenizer(model_dir: Path, max_length: int = 512) -> Tokenizer:
9
+
10
+ def load_special_tokens(model_dir: Path) -> dict:
11
+ tokens_map_path = model_dir / "special_tokens_map.json"
12
+ if not tokens_map_path.exists():
13
+ raise ValueError(f"Could not find special_tokens_map.json in {model_dir}")
14
+
15
+ with open(str(tokens_map_path)) as tokens_map_file:
16
+ tokens_map = json.load(tokens_map_file)
17
+
18
+ return tokens_map
19
+
20
+
21
+ def load_tokenizer(model_dir: Path, max_length: int = 512) -> Tuple[Tokenizer, dict]:
9
22
  config_path = model_dir / "config.json"
10
23
  if not config_path.exists():
11
24
  raise ValueError(f"Could not find config.json in {model_dir}")
@@ -18,18 +31,13 @@ def load_tokenizer(model_dir: Path, max_length: int = 512) -> Tokenizer:
18
31
  if not tokenizer_config_path.exists():
19
32
  raise ValueError(f"Could not find tokenizer_config.json in {model_dir}")
20
33
 
21
- tokens_map_path = model_dir / "special_tokens_map.json"
22
- if not tokens_map_path.exists():
23
- raise ValueError(f"Could not find special_tokens_map.json in {model_dir}")
24
-
25
34
  with open(str(config_path)) as config_file:
26
35
  config = json.load(config_file)
27
36
 
28
37
  with open(str(tokenizer_config_path)) as tokenizer_config_file:
29
38
  tokenizer_config = json.load(tokenizer_config_file)
30
39
 
31
- with open(str(tokens_map_path)) as tokens_map_file:
32
- tokens_map = json.load(tokens_map_file)
40
+ tokens_map = load_special_tokens(model_dir)
33
41
 
34
42
  tokenizer = Tokenizer.from_file(str(tokenizer_path))
35
43
  tokenizer.enable_truncation(max_length=min(tokenizer_config["model_max_length"], max_length))
@@ -43,12 +51,24 @@ def load_tokenizer(model_dir: Path, max_length: int = 512) -> Tokenizer:
43
51
  elif isinstance(token, dict):
44
52
  tokenizer.add_special_tokens([AddedToken(**token)])
45
53
 
46
- return tokenizer
54
+ special_token_to_id = {}
55
+
56
+ for token in tokens_map.values():
57
+ if isinstance(token, str):
58
+ special_token_to_id[token] = tokenizer.token_to_id(token)
59
+ elif isinstance(token, dict):
60
+ token_str = token.get("content", "")
61
+ special_token_to_id[token_str] = tokenizer.token_to_id(token_str)
62
+
63
+ return tokenizer, special_token_to_id
64
+
47
65
 
66
+ def load_preprocessor(model_dir: Path) -> Compose:
67
+ preprocessor_config_path = model_dir / "preprocessor_config.json"
68
+ if not preprocessor_config_path.exists():
69
+ raise ValueError(f"Could not find preprocessor_config.json in {model_dir}")
48
70
 
49
- def normalize(input_array, p=2, dim=1, eps=1e-12) -> np.ndarray:
50
- # Calculate the Lp norm along the specified dimension
51
- norm = np.linalg.norm(input_array, ord=p, axis=dim, keepdims=True)
52
- norm = np.maximum(norm, eps) # Avoid division by zero
53
- normalized_array = input_array / norm
54
- return normalized_array
71
+ with open(str(preprocessor_config_path)) as preprocessor_config_file:
72
+ preprocessor_config = json.load(preprocessor_config_file)
73
+ transforms = Compose.from_config(preprocessor_config)
74
+ return transforms
@@ -0,0 +1,14 @@
1
+ import os
2
+ import sys
3
+ from typing import Union, Iterable, Tuple, Dict, Any
4
+
5
+ if sys.version_info >= (3, 10):
6
+ from typing import TypeAlias
7
+ else:
8
+ from typing_extensions import TypeAlias
9
+
10
+
11
+ PathInput: TypeAlias = Union[str, os.PathLike]
12
+ ImageInput: TypeAlias = Union[PathInput, Iterable[PathInput]]
13
+
14
+ OnnxProvider: TypeAlias = Union[str, Tuple[str, Dict[Any, Any]]]
@@ -4,6 +4,16 @@ from itertools import islice
4
4
  from pathlib import Path
5
5
  from typing import Union, Iterable, Generator, Optional
6
6
 
7
+ import numpy as np
8
+
9
+
10
+ def normalize(input_array, p=2, dim=1, eps=1e-12) -> np.ndarray:
11
+ # Calculate the Lp norm along the specified dimension
12
+ norm = np.linalg.norm(input_array, ord=p, axis=dim, keepdims=True)
13
+ norm = np.maximum(norm, eps) # Avoid division by zero
14
+ normalized_array = input_array / norm
15
+ return normalized_array
16
+
7
17
 
8
18
  def iter_batch(iterable: Union[Iterable, Generator], size: int) -> Iterable:
9
19
  """
@@ -0,0 +1,4 @@
1
+ from fastembed.image.image_embedding import ImageEmbedding
2
+
3
+
4
+ __all__ = ["ImageEmbedding"]
@@ -0,0 +1,87 @@
1
+ from typing import Any, Dict, Iterable, List, Optional, Type, Sequence
2
+
3
+ import numpy as np
4
+
5
+ from fastembed.common import ImageInput, OnnxProvider
6
+ from fastembed.image.image_embedding_base import ImageEmbeddingBase
7
+ from fastembed.image.onnx_embedding import OnnxImageEmbedding
8
+
9
+
10
+ class ImageEmbedding(ImageEmbeddingBase):
11
+ EMBEDDINGS_REGISTRY: List[Type[ImageEmbeddingBase]] = [OnnxImageEmbedding]
12
+
13
+ @classmethod
14
+ def list_supported_models(cls) -> List[Dict[str, Any]]:
15
+ """
16
+ Lists the supported models.
17
+
18
+ Returns:
19
+ List[Dict[str, Any]]: A list of dictionaries containing the model information.
20
+
21
+ Example:
22
+ ```
23
+ [
24
+ {
25
+ "model": "Qdrant/clip-ViT-B-32-vision",
26
+ "dim": 512,
27
+ "description": "CLIP vision encoder based on ViT-B/32",
28
+ "size_in_GB": 0.33,
29
+ "sources": {
30
+ "hf": "Qdrant/clip-ViT-B-32-vision",
31
+ },
32
+ "model_file": "model.onnx",
33
+ }
34
+ ]
35
+ ```
36
+ """
37
+ result = []
38
+ for embedding in cls.EMBEDDINGS_REGISTRY:
39
+ result.extend(embedding.list_supported_models())
40
+ return result
41
+
42
+ def __init__(
43
+ self,
44
+ model_name: str,
45
+ cache_dir: Optional[str] = None,
46
+ threads: Optional[int] = None,
47
+ providers: Optional[Sequence[OnnxProvider]] = None,
48
+ **kwargs,
49
+ ):
50
+ super().__init__(model_name, cache_dir, threads, **kwargs)
51
+
52
+ for EMBEDDING_MODEL_TYPE in self.EMBEDDINGS_REGISTRY:
53
+ supported_models = EMBEDDING_MODEL_TYPE.list_supported_models()
54
+ if any(model_name.lower() == model["model"].lower() for model in supported_models):
55
+ self.model = EMBEDDING_MODEL_TYPE(
56
+ model_name, cache_dir, threads, providers=providers, **kwargs
57
+ )
58
+ return
59
+
60
+ raise ValueError(
61
+ f"Model {model_name} is not supported in TextEmbedding."
62
+ "Please check the supported models using `TextEmbedding.list_supported_models()`"
63
+ )
64
+
65
+ def embed(
66
+ self,
67
+ images: ImageInput,
68
+ batch_size: int = 16,
69
+ parallel: Optional[int] = None,
70
+ **kwargs,
71
+ ) -> Iterable[np.ndarray]:
72
+ """
73
+ Encode a list of documents into list of embeddings.
74
+ We use mean pooling with attention so that the model can handle variable-length inputs.
75
+
76
+ Args:
77
+ images: Iterator of image paths or single image path to embed
78
+ batch_size: Batch size for encoding -- higher values will use more memory, but be faster
79
+ parallel:
80
+ If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.
81
+ If 0, use all available cores.
82
+ If None, don't use data-parallel processing, use default onnxruntime threading instead.
83
+
84
+ Returns:
85
+ List of embeddings, one per document
86
+ """
87
+ yield from self.model.embed(images, batch_size, parallel, **kwargs)
@@ -0,0 +1,39 @@
1
+ from typing import Iterable, Optional
2
+
3
+ import numpy as np
4
+
5
+ from fastembed.common.model_management import ModelManagement
6
+ from fastembed.common.types import ImageInput
7
+
8
+
9
+ class ImageEmbeddingBase(ModelManagement):
10
+ def __init__(
11
+ self,
12
+ model_name: str,
13
+ cache_dir: Optional[str] = None,
14
+ threads: Optional[int] = None,
15
+ **kwargs,
16
+ ):
17
+ self.model_name = model_name
18
+ self.cache_dir = cache_dir
19
+ self.threads = threads
20
+ self._local_files_only = kwargs.pop("local_files_only", False)
21
+
22
+ def embed(
23
+ self,
24
+ images: ImageInput,
25
+ batch_size: int = 16,
26
+ parallel: Optional[int] = None,
27
+ **kwargs,
28
+ ) -> Iterable[np.ndarray]:
29
+ """
30
+ Embeds a list of images into a list of embeddings.
31
+
32
+ Args:
33
+ images - The list of image paths to preprocess and embed.
34
+ **kwargs: Additional keyword argument to pass to the embed method.
35
+
36
+ Yields:
37
+ Iterable[np.ndarray]: The embeddings.
38
+ """
39
+ raise NotImplementedError()