fastembed 0.2.5__tar.gz → 0.2.7__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 (26) hide show
  1. {fastembed-0.2.5 → fastembed-0.2.7}/PKG-INFO +34 -11
  2. {fastembed-0.2.5 → fastembed-0.2.7}/README.md +33 -9
  3. fastembed-0.2.7/fastembed/__init__.py +12 -0
  4. fastembed-0.2.7/fastembed/common/__init__.py +3 -0
  5. {fastembed-0.2.5 → fastembed-0.2.7}/fastembed/common/model_management.py +42 -24
  6. {fastembed-0.2.5 → fastembed-0.2.7}/fastembed/common/models.py +3 -1
  7. {fastembed-0.2.5 → fastembed-0.2.7}/fastembed/common/onnx_model.py +44 -9
  8. {fastembed-0.2.5 → fastembed-0.2.7}/fastembed/embedding.py +3 -2
  9. {fastembed-0.2.5 → fastembed-0.2.7}/fastembed/parallel_processor.py +3 -1
  10. fastembed-0.2.7/fastembed/sparse/__init__.py +4 -0
  11. {fastembed-0.2.5 → fastembed-0.2.7}/fastembed/sparse/sparse_embedding_base.py +8 -1
  12. {fastembed-0.2.5 → fastembed-0.2.7}/fastembed/sparse/sparse_text_embedding.py +6 -2
  13. {fastembed-0.2.5 → fastembed-0.2.7}/fastembed/sparse/splade_pp.py +33 -10
  14. {fastembed-0.2.5/fastembed → fastembed-0.2.7/fastembed/text}/__init__.py +0 -3
  15. {fastembed-0.2.5 → fastembed-0.2.7}/fastembed/text/e5_onnx_embedding.py +4 -1
  16. {fastembed-0.2.5 → fastembed-0.2.7}/fastembed/text/jina_onnx_embedding.py +7 -3
  17. {fastembed-0.2.5 → fastembed-0.2.7}/fastembed/text/onnx_embedding.py +110 -56
  18. {fastembed-0.2.5 → fastembed-0.2.7}/fastembed/text/text_embedding.py +6 -2
  19. {fastembed-0.2.5 → fastembed-0.2.7}/fastembed/text/text_embedding_base.py +8 -1
  20. {fastembed-0.2.5 → fastembed-0.2.7}/pyproject.toml +4 -4
  21. fastembed-0.2.5/fastembed/image/__init__.py +0 -0
  22. fastembed-0.2.5/fastembed/sparse/__init__.py +0 -0
  23. fastembed-0.2.5/fastembed/text/__init__.py +0 -0
  24. {fastembed-0.2.5 → fastembed-0.2.7}/LICENSE +0 -0
  25. {fastembed-0.2.5 → fastembed-0.2.7}/fastembed/common/utils.py +0 -0
  26. {fastembed-0.2.5/fastembed/common → fastembed-0.2.7/fastembed/image}/__init__.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.1
2
2
  Name: fastembed
3
- Version: 0.2.5
3
+ Version: 0.2.7
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
@@ -22,7 +22,7 @@ Requires-Dist: numpy (>=1.26) ; python_version >= "3.12"
22
22
  Requires-Dist: onnx (>=1.15.0,<2.0.0)
23
23
  Requires-Dist: onnxruntime (>=1.17.0,<2.0.0)
24
24
  Requires-Dist: requests (>=2.31,<3.0)
25
- Requires-Dist: tokenizers (>=0.15.1,<0.16.0)
25
+ Requires-Dist: tokenizers (>=0.15,<0.16)
26
26
  Requires-Dist: tqdm (>=4.66,<5.0)
27
27
  Project-URL: Repository, https://github.com/qdrant/fastembed
28
28
  Description-Content-Type: text/markdown
@@ -31,7 +31,7 @@ Description-Content-Type: text/markdown
31
31
 
32
32
  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
33
 
34
- The default text embedding (`TextEmbedding`) model is Flag Embedding, the top model 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/).
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/qdrant/Retrieval_with_FastEmbed/) and how to use [FastEmbed with Qdrant](https://qdrant.github.io/fastembed/qdrant/Usage_With_Qdrant/).
35
35
 
36
36
  ## 📈 Why FastEmbed?
37
37
 
@@ -43,16 +43,21 @@ The default text embedding (`TextEmbedding`) model is Flag Embedding, the top mo
43
43
 
44
44
  ## 🚀 Installation
45
45
 
46
- To install the FastEmbed library, pip works:
46
+ To install the FastEmbed library, pip works best. You can install it with or without GPU support:
47
47
 
48
48
  ```bash
49
49
  pip install fastembed
50
50
  ```
51
51
 
52
+ ### ⚡️ With GPU
53
+
54
+ ```bash
55
+ pip install fastembed-gpu
56
+ ```
57
+
52
58
  ## 📖 Quickstart
53
59
 
54
60
  ```python
55
- import numpy as np
56
61
  from fastembed import TextEmbedding
57
62
  from typing import List
58
63
 
@@ -72,6 +77,23 @@ embeddings_list = list(embedding_model.embed(documents))
72
77
  len(embeddings_list[0]) # Vector of 384 dimensions
73
78
  ```
74
79
 
80
+ ### ⚡️ FastEmbed on a GPU
81
+
82
+ FastEmbed supports running on GPU devices. It requires installation of the `fastembed-gpu` package.
83
+ Make sure not to have the `fastembed` package installed, as it might interfere with the `fastembed-gpu` package.
84
+
85
+ ```bash
86
+ pip install fastembed-gpu
87
+ ```
88
+
89
+ ```python
90
+ from fastembed import TextEmbedding
91
+
92
+ embedding_model = TextEmbedding(model_name="BAAI/bge-small-en-v1.5", providers=["CUDAExecutionProvider"])
93
+ print("The model BAAI/bge-small-en-v1.5 is ready to use on a GPU.")
94
+
95
+ ```
96
+
75
97
  ## Usage with Qdrant
76
98
 
77
99
  Installation with Qdrant Client in Python:
@@ -80,7 +102,13 @@ Installation with Qdrant Client in Python:
80
102
  pip install qdrant-client[fastembed]
81
103
  ```
82
104
 
83
- You might have to use ```pip install 'qdrant-client[fastembed]'``` on zsh.
105
+ or
106
+
107
+ ```bash
108
+ pip install qdrant-client[fastembed-gpu]
109
+ ```
110
+
111
+ You might have to use quotes ```pip install 'qdrant-client[fastembed]'``` on zsh.
84
112
 
85
113
  ```python
86
114
  from qdrant_client import QdrantClient
@@ -116,8 +144,3 @@ search_result = client.query(
116
144
  )
117
145
  print(search_result)
118
146
  ```
119
-
120
- #### Similar Work
121
-
122
- Ilyas M. wrote about using [FlagEmbeddings with Optimum](https://twitter.com/IlysMoutawwakil/status/1705215192425288017) over CUDA.
123
-
@@ -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, the top model 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,16 +14,21 @@ The default text embedding (`TextEmbedding`) model is Flag Embedding, the top mo
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
26
- import numpy as np
27
32
  from fastembed import TextEmbedding
28
33
  from typing import List
29
34
 
@@ -43,6 +48,23 @@ embeddings_list = list(embedding_model.embed(documents))
43
48
  len(embeddings_list[0]) # Vector of 384 dimensions
44
49
  ```
45
50
 
51
+ ### ⚡️ FastEmbed on a GPU
52
+
53
+ FastEmbed supports running on GPU devices. It requires installation of the `fastembed-gpu` package.
54
+ Make sure not to have the `fastembed` package installed, as it might interfere with the `fastembed-gpu` package.
55
+
56
+ ```bash
57
+ pip install fastembed-gpu
58
+ ```
59
+
60
+ ```python
61
+ from fastembed import TextEmbedding
62
+
63
+ embedding_model = TextEmbedding(model_name="BAAI/bge-small-en-v1.5", providers=["CUDAExecutionProvider"])
64
+ print("The model BAAI/bge-small-en-v1.5 is ready to use on a GPU.")
65
+
66
+ ```
67
+
46
68
  ## Usage with Qdrant
47
69
 
48
70
  Installation with Qdrant Client in Python:
@@ -51,7 +73,13 @@ Installation with Qdrant Client in Python:
51
73
  pip install qdrant-client[fastembed]
52
74
  ```
53
75
 
54
- You might have to use ```pip install 'qdrant-client[fastembed]'``` on zsh.
76
+ or
77
+
78
+ ```bash
79
+ pip install qdrant-client[fastembed-gpu]
80
+ ```
81
+
82
+ You might have to use quotes ```pip install 'qdrant-client[fastembed]'``` on zsh.
55
83
 
56
84
  ```python
57
85
  from qdrant_client import QdrantClient
@@ -86,8 +114,4 @@ search_result = client.query(
86
114
  query_text="This is a query document"
87
115
  )
88
116
  print(search_result)
89
- ```
90
-
91
- #### Similar Work
92
-
93
- Ilyas M. wrote about using [FlagEmbeddings with Optimum](https://twitter.com/IlysMoutawwakil/status/1705215192425288017) over CUDA.
117
+ ```
@@ -0,0 +1,12 @@
1
+ import importlib.metadata
2
+
3
+ from fastembed.text import TextEmbedding
4
+ from fastembed.sparse import SparseTextEmbedding, SparseEmbedding
5
+
6
+ try:
7
+ version = importlib.metadata.version("fastembed")
8
+ except importlib.metadata.PackageNotFoundError as _:
9
+ version = importlib.metadata.version("fastembed-gpu")
10
+
11
+ __version__ = version
12
+ __all__ = ["TextEmbedding", "SparseTextEmbedding", "SparseEmbedding"]
@@ -0,0 +1,3 @@
1
+ from fastembed.common.onnx_model import OnnxProvider
2
+
3
+ __all__ = ["OnnxProvider"]
@@ -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]]:
@@ -53,7 +36,7 @@ class ModelManagement:
53
36
  Dict[str, Any]: The model description.
54
37
  """
55
38
  for model in cls.list_supported_models():
56
- if model_name == model["model"]:
39
+ if model_name.lower() == model["model"].lower():
57
40
  return model
58
41
 
59
42
  raise ValueError(f"Model {model_name} is not supported in {cls.__name__}.")
@@ -92,7 +75,9 @@ class ModelManagement:
92
75
 
93
76
  show_progress = total_size_in_bytes and show_progress
94
77
 
95
- with tqdm(total=total_size_in_bytes, unit="iB", unit_scale=True, disable=not show_progress) as progress_bar:
78
+ with tqdm(
79
+ total=total_size_in_bytes, unit="iB", unit_scale=True, disable=not show_progress
80
+ ) as progress_bar:
96
81
  with open(output_path, "wb") as file:
97
82
  for chunk in response.iter_content(chunk_size=1024):
98
83
  if chunk: # Filter out keep-alive new chunks
@@ -101,20 +86,37 @@ class ModelManagement:
101
86
  return output_path
102
87
 
103
88
  @classmethod
104
- def download_files_from_huggingface(cls, hf_source_repo: str, cache_dir: Optional[str] = None) -> str:
89
+ def download_files_from_huggingface(
90
+ cls,
91
+ hf_source_repo: str,
92
+ cache_dir: Optional[str] = None,
93
+ extra_patterns: Optional[List[str]] = None,
94
+ **kwargs,
95
+ ) -> str:
105
96
  """
106
97
  Downloads a model from HuggingFace Hub.
107
98
  Args:
108
99
  hf_source_repo (str): Name of the model on HuggingFace Hub, e.g. "qdrant/all-MiniLM-L6-v2-onnx".
109
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.
110
103
  Returns:
111
104
  Path: The path to the model directory.
112
105
  """
106
+ allow_patterns = [
107
+ "config.json",
108
+ "tokenizer.json",
109
+ "tokenizer_config.json",
110
+ "special_tokens_map.json",
111
+ ]
112
+ if extra_patterns is not None:
113
+ allow_patterns.extend(extra_patterns)
113
114
 
114
115
  return snapshot_download(
115
116
  repo_id=hf_source_repo,
116
- ignore_patterns=["model.safetensors", "pytorch_model.bin"],
117
+ allow_patterns=allow_patterns,
117
118
  cache_dir=cache_dir,
119
+ local_files_only=kwargs.get("local_files_only", False),
118
120
  )
119
121
 
120
122
  @classmethod
@@ -171,6 +173,9 @@ class ModelManagement:
171
173
 
172
174
  model_tar_gz = Path(cache_dir) / f"{fast_model_name}.tar.gz"
173
175
 
176
+ if model_tar_gz.exists():
177
+ model_tar_gz.unlink()
178
+
174
179
  cls.download_file_from_gcs(
175
180
  source_url,
176
181
  output_path=str(model_tar_gz),
@@ -186,7 +191,7 @@ class ModelManagement:
186
191
  return model_dir
187
192
 
188
193
  @classmethod
189
- def download_model(cls, model: Dict[str, Any], cache_dir: Path) -> Path:
194
+ def download_model(cls, model: Dict[str, Any], cache_dir: Path, **kwargs) -> Path:
190
195
  """
191
196
  Downloads a model from HuggingFace Hub or Google Cloud Storage.
192
197
 
@@ -215,10 +220,23 @@ class ModelManagement:
215
220
  url_source = model.get("sources", {}).get("url")
216
221
 
217
222
  if hf_source:
223
+ extra_patterns = [model["model_file"]]
224
+ extra_patterns.extend(model.get("additional_files", []))
225
+
218
226
  try:
219
- return Path(cls.download_files_from_huggingface(hf_source, cache_dir=str(cache_dir)))
227
+ return Path(
228
+ cls.download_files_from_huggingface(
229
+ hf_source,
230
+ cache_dir=str(cache_dir),
231
+ extra_patterns=extra_patterns,
232
+ local_files_only=kwargs.get("local_files_only", False),
233
+ )
234
+ )
220
235
  except (EnvironmentError, RepositoryNotFoundError, ValueError) as e:
221
- logger.error(f"Could not download model from HuggingFace: {e}" "Falling back to other sources.")
236
+ logger.error(
237
+ f"Could not download model from HuggingFace: {e}"
238
+ "Falling back to other sources."
239
+ )
222
240
 
223
241
  if url_source:
224
242
  return cls.retrieve_model_gcs(model["model"], url_source, str(cache_dir))
@@ -33,7 +33,9 @@ def load_tokenizer(model_dir: Path, max_length: int = 512) -> Tokenizer:
33
33
 
34
34
  tokenizer = Tokenizer.from_file(str(tokenizer_path))
35
35
  tokenizer.enable_truncation(max_length=min(tokenizer_config["model_max_length"], max_length))
36
- tokenizer.enable_padding(pad_id=config.get("pad_token_id", 0), pad_token=tokenizer_config["pad_token"])
36
+ tokenizer.enable_padding(
37
+ pad_id=config.get("pad_token_id", 0), pad_token=tokenizer_config["pad_token"]
38
+ )
37
39
 
38
40
  for token in tokens_map.values():
39
41
  if isinstance(token, str):
@@ -1,19 +1,33 @@
1
1
  import os
2
2
  from multiprocessing import get_all_start_methods
3
3
  from pathlib import Path
4
- from typing import Any, Dict, Generic, Iterable, List, Optional, Tuple, Type, TypeVar, Union
4
+ from typing import (
5
+ Any,
6
+ Dict,
7
+ Generic,
8
+ Iterable,
9
+ List,
10
+ Optional,
11
+ Tuple,
12
+ Type,
13
+ TypeVar,
14
+ Union,
15
+ Sequence,
16
+ )
5
17
 
6
18
  import numpy as np
7
19
  import onnxruntime as ort
8
20
 
9
- from fastembed.common.model_management import locate_model_file
10
21
  from fastembed.common.models import load_tokenizer
11
22
  from fastembed.common.utils import iter_batch
12
23
  from fastembed.parallel_processor import ParallelWorkerPool, Worker
13
24
 
25
+
14
26
  # Holds type of the embedding result
15
27
  T = TypeVar("T")
16
28
 
29
+ OnnxProvider = Union[str, Tuple[str, Dict[Any, Any]]]
30
+
17
31
 
18
32
  class OnnxModel(Generic[T]):
19
33
  @classmethod
@@ -34,11 +48,26 @@ class OnnxModel(Generic[T]):
34
48
  """
35
49
  return onnx_input
36
50
 
37
- def load_onnx_model(self, model_dir: Path, threads: Optional[int], max_length: int) -> None:
38
- model_path = locate_model_file(model_dir, ["model.onnx", "model_optimized.onnx"])
51
+ def load_onnx_model(
52
+ self,
53
+ model_dir: Path,
54
+ model_file: str,
55
+ threads: Optional[int],
56
+ providers: Optional[Sequence[OnnxProvider]] = None,
57
+ ) -> None:
58
+ model_path = model_dir / model_file
39
59
 
40
60
  # List of Execution Providers: https://onnxruntime.ai/docs/execution-providers
41
- onnx_providers = ["CPUExecutionProvider"]
61
+
62
+ onnx_providers = ["CPUExecutionProvider"] if providers is None else list(providers)
63
+ available_providers = ort.get_available_providers()
64
+ for provider in onnx_providers:
65
+ # check providers available
66
+ provider_name = provider if isinstance(provider, str) else provider[0]
67
+ if provider_name not in available_providers:
68
+ raise ValueError(
69
+ f"Provider {provider_name} is not available. Available providers: {available_providers}"
70
+ )
42
71
 
43
72
  so = ort.SessionOptions()
44
73
  so.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
@@ -47,8 +76,10 @@ class OnnxModel(Generic[T]):
47
76
  so.intra_op_num_threads = threads
48
77
  so.inter_op_num_threads = threads
49
78
 
50
- self.tokenizer = load_tokenizer(model_dir=model_dir, max_length=max_length)
51
- self.model = ort.InferenceSession(str(model_path), providers=onnx_providers, sess_options=so)
79
+ self.tokenizer = load_tokenizer(model_dir=model_dir)
80
+ self.model = ort.InferenceSession(
81
+ str(model_path), providers=onnx_providers, sess_options=so
82
+ )
52
83
 
53
84
  def onnx_embed(self, documents: List[str]) -> Tuple[np.ndarray, np.ndarray]:
54
85
  encoded = self.tokenizer.encode_batch(documents)
@@ -58,7 +89,9 @@ class OnnxModel(Generic[T]):
58
89
  onnx_input = {
59
90
  "input_ids": np.array(input_ids, dtype=np.int64),
60
91
  "attention_mask": np.array(attention_mask, dtype=np.int64),
61
- "token_type_ids": np.array([np.zeros(len(e), dtype=np.int64) for e in input_ids], dtype=np.int64),
92
+ "token_type_ids": np.array(
93
+ [np.zeros(len(e), dtype=np.int64) for e in input_ids], dtype=np.int64
94
+ ),
62
95
  }
63
96
 
64
97
  onnx_input = self._preprocess_onnx_input(onnx_input)
@@ -97,7 +130,9 @@ class OnnxModel(Generic[T]):
97
130
  "model_name": model_name,
98
131
  "cache_dir": cache_dir,
99
132
  }
100
- pool = ParallelWorkerPool(parallel, self._get_worker_class(), start_method=start_method)
133
+ pool = ParallelWorkerPool(
134
+ parallel, self._get_worker_class(), start_method=start_method
135
+ )
101
136
  for batch in pool.ordered_map(iter_batch(documents, batch_size), **params):
102
137
  yield from self._post_process_onnx_output(batch)
103
138
 
@@ -2,10 +2,11 @@ from typing import Optional
2
2
 
3
3
  from loguru import logger
4
4
 
5
- from fastembed.text.text_embedding import TextEmbedding
5
+ from fastembed import TextEmbedding
6
6
 
7
7
  logger.warning(
8
- "DefaultEmbedding, FlagEmbedding, JinaEmbedding are deprecated." "Use from fastembed import TextEmbedding instead."
8
+ "DefaultEmbedding, FlagEmbedding, JinaEmbedding are deprecated."
9
+ "Use from fastembed import TextEmbedding instead."
9
10
  )
10
11
 
11
12
  DefaultEmbedding = TextEmbedding
@@ -128,7 +128,9 @@ class ParallelWorkerPool:
128
128
  yield buffer.pop(next_expected)
129
129
  next_expected += 1
130
130
 
131
- def semi_ordered_map(self, stream: Iterable[Any], *args: Any, **kwargs: Any) -> Iterable[Tuple[int, Any]]:
131
+ def semi_ordered_map(
132
+ self, stream: Iterable[Any], *args: Any, **kwargs: Any
133
+ ) -> Iterable[Tuple[int, Any]]:
132
134
  try:
133
135
  self.start(**kwargs)
134
136
 
@@ -0,0 +1,4 @@
1
+ from fastembed.sparse.sparse_embedding_base import SparseEmbedding
2
+ from fastembed.sparse.sparse_text_embedding import SparseTextEmbedding
3
+
4
+ __all__ = ["SparseEmbedding", "SparseTextEmbedding"]
@@ -22,10 +22,17 @@ class SparseEmbedding:
22
22
 
23
23
 
24
24
  class SparseTextEmbeddingBase(ModelManagement):
25
- def __init__(self, model_name: str, cache_dir: Optional[str] = None, threads: Optional[int] = None, **kwargs):
25
+ def __init__(
26
+ self,
27
+ model_name: str,
28
+ cache_dir: Optional[str] = None,
29
+ threads: Optional[int] = None,
30
+ **kwargs,
31
+ ):
26
32
  self.model_name = model_name
27
33
  self.cache_dir = cache_dir
28
34
  self.threads = threads
35
+ self._local_files_only = kwargs.pop("local_files_only", False)
29
36
 
30
37
  def embed(
31
38
  self,
@@ -1,5 +1,6 @@
1
- from typing import List, Type, Dict, Any, Union, Iterable, Optional
1
+ from typing import List, Type, Dict, Any, Union, Iterable, Optional, Sequence
2
2
 
3
+ from fastembed.common import OnnxProvider
3
4
  from fastembed.sparse.sparse_embedding_base import SparseTextEmbeddingBase, SparseEmbedding
4
5
  from fastembed.sparse.splade_pp import SpladePP
5
6
 
@@ -42,6 +43,7 @@ class SparseTextEmbedding(SparseTextEmbeddingBase):
42
43
  model_name: str,
43
44
  cache_dir: Optional[str] = None,
44
45
  threads: Optional[int] = None,
46
+ providers: Optional[Sequence[OnnxProvider]] = None,
45
47
  **kwargs,
46
48
  ):
47
49
  super().__init__(model_name, cache_dir, threads, **kwargs)
@@ -49,7 +51,9 @@ class SparseTextEmbedding(SparseTextEmbeddingBase):
49
51
  for EMBEDDING_MODEL_TYPE in self.EMBEDDINGS_REGISTRY:
50
52
  supported_models = EMBEDDING_MODEL_TYPE.list_supported_models()
51
53
  if any(model_name.lower() == model["model"].lower() for model in supported_models):
52
- self.model = EMBEDDING_MODEL_TYPE(model_name, cache_dir, threads, **kwargs)
54
+ self.model = EMBEDDING_MODEL_TYPE(
55
+ model_name, cache_dir, threads, providers=providers, **kwargs
56
+ )
53
57
  return
54
58
 
55
59
  raise ValueError(
@@ -1,8 +1,8 @@
1
- from typing import Any, Dict, Iterable, List, Optional, Tuple, Union
1
+ from typing import Any, Dict, Iterable, List, Optional, Tuple, Union, Type, Sequence
2
2
 
3
3
  import numpy as np
4
4
 
5
- from fastembed.common.onnx_model import EmbeddingWorker, OnnxModel
5
+ from fastembed.common.onnx_model import EmbeddingWorker, OnnxModel, OnnxProvider
6
6
  from fastembed.common.utils import define_cache_dir
7
7
  from fastembed.sparse.sparse_embedding_base import SparseEmbedding, SparseTextEmbeddingBase
8
8
 
@@ -10,18 +10,31 @@ supported_splade_models = [
10
10
  {
11
11
  "model": "prithvida/Splade_PP_en_v1",
12
12
  "vocab_size": 30522,
13
+ "description": "Misspelled version of the model. Retained for backward compatibility. Independent Implementation of SPLADE++ Model for English",
14
+ "size_in_GB": 0.532,
15
+ "sources": {
16
+ "hf": "Qdrant/SPLADE_PP_en_v1",
17
+ },
18
+ "model_file": "model.onnx",
19
+ },
20
+ {
21
+ "model": "prithivida/Splade_PP_en_v1",
22
+ "vocab_size": 30522,
13
23
  "description": "Independent Implementation of SPLADE++ Model for English",
14
24
  "size_in_GB": 0.532,
15
25
  "sources": {
16
26
  "hf": "Qdrant/SPLADE_PP_en_v1",
17
27
  },
28
+ "model_file": "model.onnx",
18
29
  },
19
30
  ]
20
31
 
21
32
 
22
33
  class SpladePP(SparseTextEmbeddingBase, OnnxModel[SparseEmbedding]):
23
34
  @classmethod
24
- def _post_process_onnx_output(cls, output: Tuple[np.ndarray, np.ndarray]) -> Iterable[SparseEmbedding]:
35
+ def _post_process_onnx_output(
36
+ cls, output: Tuple[np.ndarray, np.ndarray]
37
+ ) -> Iterable[SparseEmbedding]:
25
38
  logits, attention_mask = output
26
39
  relu_log = np.log(1 + np.maximum(logits, 0))
27
40
 
@@ -50,6 +63,7 @@ class SpladePP(SparseTextEmbeddingBase, OnnxModel[SparseEmbedding]):
50
63
  model_name: str,
51
64
  cache_dir: Optional[str] = None,
52
65
  threads: Optional[int] = None,
66
+ providers: Optional[Sequence[OnnxProvider]] = None,
53
67
  **kwargs,
54
68
  ):
55
69
  """
@@ -66,14 +80,19 @@ class SpladePP(SparseTextEmbeddingBase, OnnxModel[SparseEmbedding]):
66
80
 
67
81
  super().__init__(model_name, cache_dir, threads, **kwargs)
68
82
 
69
- self.model_name = model_name
70
- self._model_description = self._get_model_description(model_name)
83
+ model_description = self._get_model_description(model_name)
84
+ cache_dir = define_cache_dir(cache_dir)
71
85
 
72
- self._cache_dir = define_cache_dir(cache_dir)
73
- self._model_dir = self.download_model(self._model_description, self._cache_dir)
74
- self._max_length = 512
86
+ model_dir = self.download_model(
87
+ model_description, cache_dir, local_files_only=self._local_files_only
88
+ )
75
89
 
76
- self.load_onnx_model(self._model_dir, self.threads, self._max_length)
90
+ self.load_onnx_model(
91
+ model_dir=model_dir,
92
+ model_file=model_description["model_file"],
93
+ threads=threads,
94
+ providers=providers,
95
+ )
77
96
 
78
97
  def embed(
79
98
  self,
@@ -99,12 +118,16 @@ class SpladePP(SparseTextEmbeddingBase, OnnxModel[SparseEmbedding]):
99
118
  """
100
119
  yield from self._embed_documents(
101
120
  model_name=self.model_name,
102
- cache_dir=str(self._cache_dir),
121
+ cache_dir=str(self.cache_dir),
103
122
  documents=documents,
104
123
  batch_size=batch_size,
105
124
  parallel=parallel,
106
125
  )
107
126
 
127
+ @classmethod
128
+ def _get_worker_class(cls) -> Type[EmbeddingWorker]:
129
+ return SpladePPEmbeddingWorker
130
+
108
131
 
109
132
  class SpladePPEmbeddingWorker(EmbeddingWorker):
110
133
  def init_embedding(
@@ -1,6 +1,3 @@
1
- import importlib.metadata
2
-
3
1
  from fastembed.text.text_embedding import TextEmbedding
4
2
 
5
- __version__ = importlib.metadata.version("fastembed")
6
3
  __all__ = ["TextEmbedding"]
@@ -15,15 +15,18 @@ supported_multilingual_e5_models = [
15
15
  "url": "https://storage.googleapis.com/qdrant-fastembed/fast-multilingual-e5-large.tar.gz",
16
16
  "hf": "qdrant/multilingual-e5-large-onnx",
17
17
  },
18
+ "model_file": "model.onnx",
19
+ "additional_files": ["model.onnx_data"],
18
20
  },
19
21
  {
20
22
  "model": "sentence-transformers/paraphrase-multilingual-mpnet-base-v2",
21
23
  "dim": 768,
22
24
  "description": "Sentence-transformers model for tasks like clustering or semantic search",
23
- "size_in_GB": 1.11,
25
+ "size_in_GB": 1.00,
24
26
  "sources": {
25
27
  "hf": "xenova/paraphrase-multilingual-mpnet-base-v2",
26
28
  },
29
+ "model_file": "onnx/model.onnx",
27
30
  },
28
31
  ]
29
32
 
@@ -11,15 +11,17 @@ supported_jina_models = [
11
11
  "model": "jinaai/jina-embeddings-v2-base-en",
12
12
  "dim": 768,
13
13
  "description": "English embedding model supporting 8192 sequence length",
14
- "size_in_GB": 0.55,
14
+ "size_in_GB": 0.52,
15
15
  "sources": {"hf": "xenova/jina-embeddings-v2-base-en"},
16
+ "model_file": "onnx/model.onnx",
16
17
  },
17
18
  {
18
19
  "model": "jinaai/jina-embeddings-v2-small-en",
19
20
  "dim": 512,
20
21
  "description": "English embedding model supporting 8192 sequence length",
21
- "size_in_GB": 0.13,
22
+ "size_in_GB": 0.12,
22
23
  "sources": {"hf": "xenova/jina-embeddings-v2-small-en"},
24
+ "model_file": "onnx/model.onnx",
23
25
  },
24
26
  ]
25
27
 
@@ -49,7 +51,9 @@ class JinaOnnxEmbedding(OnnxTextEmbedding):
49
51
  return supported_jina_models
50
52
 
51
53
  @classmethod
52
- def _post_process_onnx_output(cls, output: Tuple[np.ndarray, np.ndarray]) -> Iterable[np.ndarray]:
54
+ def _post_process_onnx_output(
55
+ cls, output: Tuple[np.ndarray, np.ndarray]
56
+ ) -> Iterable[np.ndarray]:
53
57
  embeddings, attn_mask = output
54
58
  return normalize(cls.mean_pooling(embeddings, attn_mask)).astype(np.float32)
55
59
 
@@ -1,8 +1,8 @@
1
- from typing import Dict, Optional, Tuple, Union, Iterable, Type, List, Any
1
+ from typing import Dict, Optional, Tuple, Union, Iterable, Type, List, Any, Sequence
2
2
 
3
3
  import numpy as np
4
4
 
5
- from fastembed.common.onnx_model import OnnxModel, EmbeddingWorker
5
+ from fastembed.common.onnx_model import OnnxModel, EmbeddingWorker, OnnxProvider
6
6
  from fastembed.common.models import normalize
7
7
  from fastembed.common.utils import define_cache_dir
8
8
  from fastembed.text.text_embedding_base import TextEmbeddingBase
@@ -12,79 +12,64 @@ supported_onnx_models = [
12
12
  "model": "BAAI/bge-base-en",
13
13
  "dim": 768,
14
14
  "description": "Base English model",
15
- "size_in_GB": 0.5,
15
+ "size_in_GB": 0.42,
16
16
  "sources": {
17
17
  "url": "https://storage.googleapis.com/qdrant-fastembed/fast-bge-base-en.tar.gz",
18
18
  },
19
+ "model_file": "model_optimized.onnx",
19
20
  },
20
21
  {
21
22
  "model": "BAAI/bge-base-en-v1.5",
22
23
  "dim": 768,
23
24
  "description": "Base English model, v1.5",
24
- "size_in_GB": 0.44,
25
+ "size_in_GB": 0.21,
25
26
  "sources": {
26
27
  "url": "https://storage.googleapis.com/qdrant-fastembed/fast-bge-base-en-v1.5.tar.gz",
27
28
  "hf": "qdrant/bge-base-en-v1.5-onnx-q",
28
29
  },
29
- },
30
- {
31
- "model": "BAAI/bge-large-en-v1.5-quantized",
32
- "dim": 1024,
33
- "description": "Large English model, v1.5",
34
- "size_in_GB": 1.34,
35
- "sources": {
36
- "hf": "qdrant/bge-large-en-v1.5-onnx-q",
37
- },
30
+ "model_file": "model_optimized.onnx",
38
31
  },
39
32
  {
40
33
  "model": "BAAI/bge-large-en-v1.5",
41
34
  "dim": 1024,
42
35
  "description": "Large English model, v1.5",
43
- "size_in_GB": 1.34,
36
+ "size_in_GB": 1.20,
44
37
  "sources": {
45
38
  "hf": "qdrant/bge-large-en-v1.5-onnx",
46
39
  },
40
+ "model_file": "model.onnx",
47
41
  },
48
42
  {
49
43
  "model": "BAAI/bge-small-en",
50
44
  "dim": 384,
51
45
  "description": "Fast English model",
52
- "size_in_GB": 0.2,
46
+ "size_in_GB": 0.13,
53
47
  "sources": {
54
48
  "url": "https://storage.googleapis.com/qdrant-fastembed/BAAI-bge-small-en.tar.gz",
55
49
  },
50
+ "model_file": "model_optimized.onnx",
56
51
  },
57
- # {
58
- # "model": "BAAI/bge-small-en",
59
- # "dim": 384,
60
- # "description": "Fast English model",
61
- # "size_in_GB": 0.2,
62
- # "hf_sources": [],
63
- # "compressed_url_sources": [
64
- # "https://storage.googleapis.com/qdrant-fastembed/fast-bge-small-en.tar.gz",
65
- # "https://storage.googleapis.com/qdrant-fastembed/BAAI-bge-small-en.tar.gz"
66
- # ]
67
- # },
68
52
  {
69
53
  "model": "BAAI/bge-small-en-v1.5",
70
54
  "dim": 384,
71
55
  "description": "Fast and Default English model",
72
- "size_in_GB": 0.13,
56
+ "size_in_GB": 0.067,
73
57
  "sources": {
74
- "url": "https://storage.googleapis.com/qdrant-fastembed/fast-bge-small-en-v1.5.tar.gz",
75
58
  "hf": "qdrant/bge-small-en-v1.5-onnx-q",
76
59
  },
60
+ "model_file": "model_optimized.onnx",
77
61
  },
78
62
  {
79
63
  "model": "BAAI/bge-small-zh-v1.5",
80
64
  "dim": 512,
81
65
  "description": "Fast and recommended Chinese model",
82
- "size_in_GB": 0.1,
66
+ "size_in_GB": 0.09,
83
67
  "sources": {
84
68
  "url": "https://storage.googleapis.com/qdrant-fastembed/fast-bge-small-zh-v1.5.tar.gz",
85
69
  },
70
+ "model_file": "model_optimized.onnx",
86
71
  },
87
- { # todo: it is not a flag embedding
72
+ {
88
73
  "model": "sentence-transformers/all-MiniLM-L6-v2",
89
74
  "dim": 384,
90
75
  "description": "Sentence Transformer model, MiniLM-L6-v2",
@@ -93,56 +78,118 @@ supported_onnx_models = [
93
78
  "url": "https://storage.googleapis.com/qdrant-fastembed/sentence-transformers-all-MiniLM-L6-v2.tar.gz",
94
79
  "hf": "qdrant/all-MiniLM-L6-v2-onnx",
95
80
  },
81
+ "model_file": "model.onnx",
96
82
  },
97
83
  {
98
84
  "model": "sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2",
99
85
  "dim": 384,
100
86
  "description": "Sentence Transformer model, paraphrase-multilingual-MiniLM-L12-v2",
101
- "size_in_GB": 0.46,
87
+ "size_in_GB": 0.22,
102
88
  "sources": {
103
89
  "hf": "qdrant/paraphrase-multilingual-MiniLM-L12-v2-onnx-Q",
104
90
  },
91
+ "model_file": "model_optimized.onnx",
105
92
  },
106
93
  {
107
94
  "model": "nomic-ai/nomic-embed-text-v1",
108
95
  "dim": 768,
109
96
  "description": "8192 context length english model",
110
- "size_in_GB": 0.54,
97
+ "size_in_GB": 0.52,
111
98
  "sources": {
112
99
  "hf": "nomic-ai/nomic-embed-text-v1",
113
100
  },
101
+ "model_file": "onnx/model.onnx",
114
102
  },
115
103
  {
116
104
  "model": "nomic-ai/nomic-embed-text-v1.5",
117
105
  "dim": 768,
118
106
  "description": "8192 context length english model",
119
- "size_in_GB": 0.54,
107
+ "size_in_GB": 0.52,
108
+ "sources": {
109
+ "hf": "nomic-ai/nomic-embed-text-v1.5",
110
+ },
111
+ "model_file": "onnx/model.onnx",
112
+ },
113
+ {
114
+ "model": "nomic-ai/nomic-embed-text-v1.5-Q",
115
+ "dim": 768,
116
+ "description": "Quantized 8192 context length english model",
117
+ "size_in_GB": 0.13,
120
118
  "sources": {
121
119
  "hf": "nomic-ai/nomic-embed-text-v1.5",
122
120
  },
121
+ "model_file": "onnx/model_quantized.onnx",
123
122
  },
124
123
  {
125
124
  "model": "thenlper/gte-large",
126
125
  "dim": 1024,
127
126
  "description": "Large general text embeddings model",
128
- "size_in_GB": 1.34,
127
+ "size_in_GB": 1.20,
129
128
  "sources": {
130
129
  "hf": "qdrant/gte-large-onnx",
131
130
  },
131
+ "model_file": "model.onnx",
132
+ },
133
+ {
134
+ "model": "mixedbread-ai/mxbai-embed-large-v1",
135
+ "dim": 1024,
136
+ "description": "MixedBread Base sentence embedding model, does well on MTEB",
137
+ "size_in_GB": 0.64,
138
+ "sources": {
139
+ "hf": "mixedbread-ai/mxbai-embed-large-v1",
140
+ },
141
+ "model_file": "onnx/model.onnx",
142
+ },
143
+ {
144
+ "model": "snowflake/snowflake-arctic-embed-xs",
145
+ "dim": 384,
146
+ "description": "Based on all-MiniLM-L6-v2 model with only 22m parameters, ideal for latency/TCO budgets.",
147
+ "size_in_GB": 0.09,
148
+ "sources": {
149
+ "hf": "snowflake/snowflake-arctic-embed-xs",
150
+ },
151
+ "model_file": "onnx/model.onnx",
152
+ },
153
+ {
154
+ "model": "snowflake/snowflake-arctic-embed-s",
155
+ "dim": 384,
156
+ "description": "Based on infloat/e5-small-unsupervised, does not trade off retrieval accuracy for its small size.",
157
+ "size_in_GB": 0.13,
158
+ "sources": {
159
+ "hf": "snowflake/snowflake-arctic-embed-s",
160
+ },
161
+ "model_file": "onnx/model.onnx",
162
+ },
163
+ {
164
+ "model": "snowflake/snowflake-arctic-embed-m",
165
+ "dim": 768,
166
+ "description": "Based on intfloat/e5-base-unsupervised model, provides the best retrieval without slowing down inference.",
167
+ "size_in_GB": 0.43,
168
+ "sources": {
169
+ "hf": "Snowflake/snowflake-arctic-embed-m",
170
+ },
171
+ "model_file": "onnx/model.onnx",
172
+ },
173
+ {
174
+ "model": "snowflake/snowflake-arctic-embed-m-long",
175
+ "dim": 768,
176
+ "description": "Based on nomic-ai/nomic-embed-text-v1-unsupervised model, 8192 context-length model",
177
+ "size_in_GB": 0.54,
178
+ "sources": {
179
+ "hf": "snowflake/snowflake-arctic-embed-m-long",
180
+ },
181
+ "model_file": "onnx/model.onnx",
182
+ },
183
+ {
184
+ "model": "snowflake/snowflake-arctic-embed-l",
185
+ "dim": 1024,
186
+ "description": "Based on intfloat/e5-large-unsupervised, large model for most accurate retrieval.",
187
+ "size_in_GB": 1.02,
188
+ "sources": {
189
+ "hf": "snowflake/snowflake-arctic-embed-l",
190
+ },
191
+ "model_file": "onnx/model.onnx",
132
192
  },
133
- # {
134
- # "model": "sentence-transformers/all-MiniLM-L6-v2",
135
- # "dim": 384,
136
- # "description": "Sentence Transformer model, MiniLM-L6-v2",
137
- # "size_in_GB": 0.09,
138
- # "hf_sources": [
139
- # "qdrant/all-MiniLM-L6-v2-onnx"
140
- # ],
141
- # "compressed_url_sources": [
142
- # "https://storage.googleapis.com/qdrant-fastembed/fast-all-MiniLM-L6-v2.tar.gz",
143
- # "https://storage.googleapis.com/qdrant-fastembed/sentence-transformers-all-MiniLM-L6-v2.tar.gz"
144
- # ]
145
- # }
146
193
  ]
147
194
 
148
195
 
@@ -164,6 +211,7 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxModel[np.ndarray]):
164
211
  model_name: str = "BAAI/bge-small-en-v1.5",
165
212
  cache_dir: Optional[str] = None,
166
213
  threads: Optional[int] = None,
214
+ providers: Optional[Sequence[OnnxProvider]] = None,
167
215
  **kwargs,
168
216
  ):
169
217
  """
@@ -180,14 +228,18 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxModel[np.ndarray]):
180
228
 
181
229
  super().__init__(model_name, cache_dir, threads, **kwargs)
182
230
 
183
- self.model_name = model_name
184
- self._model_description = self._get_model_description(model_name)
185
-
186
- self._cache_dir = define_cache_dir(cache_dir)
187
- self._model_dir = self.download_model(self._model_description, self._cache_dir)
188
- self._max_length = 512
231
+ model_description = self._get_model_description(model_name)
232
+ cache_dir = define_cache_dir(cache_dir)
233
+ model_dir = self.download_model(
234
+ model_description, cache_dir, local_files_only=self._local_files_only
235
+ )
189
236
 
190
- self.load_onnx_model(self._model_dir, self.threads, self._max_length)
237
+ self.load_onnx_model(
238
+ model_dir=model_dir,
239
+ model_file=model_description["model_file"],
240
+ threads=threads,
241
+ providers=providers,
242
+ )
191
243
 
192
244
  def embed(
193
245
  self,
@@ -213,7 +265,7 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxModel[np.ndarray]):
213
265
  """
214
266
  yield from self._embed_documents(
215
267
  model_name=self.model_name,
216
- cache_dir=str(self._cache_dir),
268
+ cache_dir=str(self.cache_dir),
217
269
  documents=documents,
218
270
  batch_size=batch_size,
219
271
  parallel=parallel,
@@ -230,7 +282,9 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxModel[np.ndarray]):
230
282
  return onnx_input
231
283
 
232
284
  @classmethod
233
- def _post_process_onnx_output(cls, output: Tuple[np.ndarray, np.ndarray]) -> Iterable[np.ndarray]:
285
+ def _post_process_onnx_output(
286
+ cls, output: Tuple[np.ndarray, np.ndarray]
287
+ ) -> Iterable[np.ndarray]:
234
288
  embeddings, _ = output
235
289
  return normalize(embeddings[:, 0]).astype(np.float32)
236
290
 
@@ -1,7 +1,8 @@
1
- from typing import Any, Dict, Iterable, List, Optional, Type, Union
1
+ from typing import Any, Dict, Iterable, List, Optional, Type, Union, Sequence
2
2
 
3
3
  import numpy as np
4
4
 
5
+ from fastembed.common import OnnxProvider
5
6
  from fastembed.text.e5_onnx_embedding import E5OnnxEmbedding
6
7
  from fastembed.text.jina_onnx_embedding import JinaOnnxEmbedding
7
8
  from fastembed.text.onnx_embedding import OnnxTextEmbedding
@@ -49,6 +50,7 @@ class TextEmbedding(TextEmbeddingBase):
49
50
  model_name: str = "BAAI/bge-small-en-v1.5",
50
51
  cache_dir: Optional[str] = None,
51
52
  threads: Optional[int] = None,
53
+ providers: Optional[Sequence[OnnxProvider]] = None,
52
54
  **kwargs,
53
55
  ):
54
56
  super().__init__(model_name, cache_dir, threads, **kwargs)
@@ -56,7 +58,9 @@ class TextEmbedding(TextEmbeddingBase):
56
58
  for EMBEDDING_MODEL_TYPE in self.EMBEDDINGS_REGISTRY:
57
59
  supported_models = EMBEDDING_MODEL_TYPE.list_supported_models()
58
60
  if any(model_name.lower() == model["model"].lower() for model in supported_models):
59
- self.model = EMBEDDING_MODEL_TYPE(model_name, cache_dir, threads, **kwargs)
61
+ self.model = EMBEDDING_MODEL_TYPE(
62
+ model_name, cache_dir, threads, providers=providers, **kwargs
63
+ )
60
64
  return
61
65
 
62
66
  raise ValueError(
@@ -6,10 +6,17 @@ from fastembed.common.model_management import ModelManagement
6
6
 
7
7
 
8
8
  class TextEmbeddingBase(ModelManagement):
9
- def __init__(self, model_name: str, cache_dir: Optional[str] = None, threads: Optional[int] = None, **kwargs):
9
+ def __init__(
10
+ self,
11
+ model_name: str,
12
+ cache_dir: Optional[str] = None,
13
+ threads: Optional[int] = None,
14
+ **kwargs,
15
+ ):
10
16
  self.model_name = model_name
11
17
  self.cache_dir = cache_dir
12
18
  self.threads = threads
19
+ self._local_files_only = kwargs.pop("local_files_only", False)
13
20
 
14
21
  def embed(
15
22
  self,
@@ -1,6 +1,6 @@
1
1
  [tool.poetry]
2
2
  name = "fastembed"
3
- version = "0.2.5"
3
+ version = "0.2.7"
4
4
  description = "Fast, light, accurate library built for retrieval embedding generation"
5
5
  authors = ["NirantK <nirant.bits@gmail.com>"]
6
6
  license = "Apache License"
@@ -16,7 +16,7 @@ onnx = "^1.15.0"
16
16
  onnxruntime = "^1.17.0"
17
17
  tqdm = "^4.66"
18
18
  requests = "^2.31"
19
- tokenizers = "^0.15.1"
19
+ tokenizers = "^0.15"
20
20
  huggingface-hub = "^0.20"
21
21
  loguru = "^0.7.2"
22
22
  numpy = [
@@ -26,7 +26,7 @@ numpy = [
26
26
 
27
27
  [tool.poetry.group.dev.dependencies]
28
28
  pytest = "^7.4.2"
29
- ruff = "^0.2.2"
29
+ ruff = "^0.3.1"
30
30
  notebook = ">=7.0.2"
31
31
  pre-commit = {version = "^3.6.2", python = ">=3.9,<3.12" }
32
32
 
@@ -43,4 +43,4 @@ requires = ["poetry-core"]
43
43
  build-backend = "poetry.core.masonry.api"
44
44
 
45
45
  [tool.ruff]
46
- line-length = 120
46
+ line-length = 99
File without changes
File without changes
File without changes
File without changes