fastembed 0.2.4__tar.gz → 0.2.6__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 (25) hide show
  1. {fastembed-0.2.4 → fastembed-0.2.6}/PKG-INFO +2 -3
  2. {fastembed-0.2.4 → fastembed-0.2.6}/README.md +1 -2
  3. fastembed-0.2.6/fastembed/__init__.py +7 -0
  4. {fastembed-0.2.4 → fastembed-0.2.6}/fastembed/common/model_management.py +14 -5
  5. {fastembed-0.2.4 → fastembed-0.2.6}/fastembed/common/models.py +3 -1
  6. {fastembed-0.2.4 → fastembed-0.2.6}/fastembed/common/onnx_model.py +9 -3
  7. {fastembed-0.2.4 → fastembed-0.2.6}/fastembed/embedding.py +3 -2
  8. {fastembed-0.2.4 → fastembed-0.2.6}/fastembed/parallel_processor.py +3 -1
  9. fastembed-0.2.6/fastembed/sparse/__init__.py +4 -0
  10. {fastembed-0.2.4 → fastembed-0.2.6}/fastembed/sparse/sparse_embedding_base.py +7 -1
  11. {fastembed-0.2.4 → fastembed-0.2.6}/fastembed/sparse/sparse_text_embedding.py +1 -1
  12. {fastembed-0.2.4 → fastembed-0.2.6}/fastembed/sparse/splade_pp.py +18 -3
  13. {fastembed-0.2.4 → fastembed-0.2.6}/fastembed/text/e5_onnx_embedding.py +1 -1
  14. {fastembed-0.2.4 → fastembed-0.2.6}/fastembed/text/jina_onnx_embedding.py +5 -3
  15. {fastembed-0.2.4 → fastembed-0.2.6}/fastembed/text/onnx_embedding.py +23 -22
  16. {fastembed-0.2.4 → fastembed-0.2.6}/fastembed/text/text_embedding.py +1 -1
  17. {fastembed-0.2.4 → fastembed-0.2.6}/fastembed/text/text_embedding_base.py +7 -1
  18. {fastembed-0.2.4 → fastembed-0.2.6}/pyproject.toml +3 -3
  19. fastembed-0.2.4/fastembed/sparse/__init__.py +0 -0
  20. fastembed-0.2.4/fastembed/text/__init__.py +0 -0
  21. {fastembed-0.2.4 → fastembed-0.2.6}/LICENSE +0 -0
  22. {fastembed-0.2.4 → fastembed-0.2.6}/fastembed/common/__init__.py +0 -0
  23. {fastembed-0.2.4 → fastembed-0.2.6}/fastembed/common/utils.py +0 -0
  24. {fastembed-0.2.4 → fastembed-0.2.6}/fastembed/image/__init__.py +0 -0
  25. {fastembed-0.2.4/fastembed → fastembed-0.2.6/fastembed/text}/__init__.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.1
2
2
  Name: fastembed
3
- Version: 0.2.4
3
+ Version: 0.2.6
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
@@ -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/examples/Retrieval_with_FastEmbed/) and how to use [FastEmbed with Qdrant](https://qdrant.github.io/fastembed/examples/Usage_With_Qdrant/).
35
35
 
36
36
  ## 📈 Why FastEmbed?
37
37
 
@@ -52,7 +52,6 @@ pip install fastembed
52
52
  ## 📖 Quickstart
53
53
 
54
54
  ```python
55
- import numpy as np
56
55
  from fastembed import TextEmbedding
57
56
  from typing import List
58
57
 
@@ -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/examples/Retrieval_with_FastEmbed/) and how to use [FastEmbed with Qdrant](https://qdrant.github.io/fastembed/examples/Usage_With_Qdrant/).
6
6
 
7
7
  ## 📈 Why FastEmbed?
8
8
 
@@ -23,7 +23,6 @@ pip install fastembed
23
23
  ## 📖 Quickstart
24
24
 
25
25
  ```python
26
- import numpy as np
27
26
  from fastembed import TextEmbedding
28
27
  from typing import List
29
28
 
@@ -0,0 +1,7 @@
1
+ import importlib.metadata
2
+
3
+ from fastembed.text import TextEmbedding
4
+ from fastembed.sparse import SparseTextEmbedding, SparseEmbedding
5
+
6
+ __version__ = importlib.metadata.version("fastembed")
7
+ __all__ = ["TextEmbedding", "SparseTextEmbedding", "SparseEmbedding"]
@@ -53,7 +53,7 @@ class ModelManagement:
53
53
  Dict[str, Any]: The model description.
54
54
  """
55
55
  for model in cls.list_supported_models():
56
- if model_name == model["model"]:
56
+ if model_name.lower() == model["model"].lower():
57
57
  return model
58
58
 
59
59
  raise ValueError(f"Model {model_name} is not supported in {cls.__name__}.")
@@ -92,7 +92,9 @@ class ModelManagement:
92
92
 
93
93
  show_progress = total_size_in_bytes and show_progress
94
94
 
95
- with tqdm(total=total_size_in_bytes, unit="iB", unit_scale=True, disable=not show_progress) as progress_bar:
95
+ with tqdm(
96
+ total=total_size_in_bytes, unit="iB", unit_scale=True, disable=not show_progress
97
+ ) as progress_bar:
96
98
  with open(output_path, "wb") as file:
97
99
  for chunk in response.iter_content(chunk_size=1024):
98
100
  if chunk: # Filter out keep-alive new chunks
@@ -101,7 +103,9 @@ class ModelManagement:
101
103
  return output_path
102
104
 
103
105
  @classmethod
104
- def download_files_from_huggingface(cls, hf_source_repo: str, cache_dir: Optional[str] = None) -> str:
106
+ def download_files_from_huggingface(
107
+ cls, hf_source_repo: str, cache_dir: Optional[str] = None
108
+ ) -> str:
105
109
  """
106
110
  Downloads a model from HuggingFace Hub.
107
111
  Args:
@@ -216,9 +220,14 @@ class ModelManagement:
216
220
 
217
221
  if hf_source:
218
222
  try:
219
- return Path(cls.download_files_from_huggingface(hf_source, cache_dir=str(cache_dir)))
223
+ return Path(
224
+ cls.download_files_from_huggingface(hf_source, cache_dir=str(cache_dir))
225
+ )
220
226
  except (EnvironmentError, RepositoryNotFoundError, ValueError) as e:
221
- logger.error(f"Could not download model from HuggingFace: {e}" "Falling back to other sources.")
227
+ logger.error(
228
+ f"Could not download model from HuggingFace: {e}"
229
+ "Falling back to other sources."
230
+ )
222
231
 
223
232
  if url_source:
224
233
  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):
@@ -48,7 +48,9 @@ class OnnxModel(Generic[T]):
48
48
  so.inter_op_num_threads = threads
49
49
 
50
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)
51
+ self.model = ort.InferenceSession(
52
+ str(model_path), providers=onnx_providers, sess_options=so
53
+ )
52
54
 
53
55
  def onnx_embed(self, documents: List[str]) -> Tuple[np.ndarray, np.ndarray]:
54
56
  encoded = self.tokenizer.encode_batch(documents)
@@ -58,7 +60,9 @@ class OnnxModel(Generic[T]):
58
60
  onnx_input = {
59
61
  "input_ids": np.array(input_ids, dtype=np.int64),
60
62
  "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),
63
+ "token_type_ids": np.array(
64
+ [np.zeros(len(e), dtype=np.int64) for e in input_ids], dtype=np.int64
65
+ ),
62
66
  }
63
67
 
64
68
  onnx_input = self._preprocess_onnx_input(onnx_input)
@@ -97,7 +101,9 @@ class OnnxModel(Generic[T]):
97
101
  "model_name": model_name,
98
102
  "cache_dir": cache_dir,
99
103
  }
100
- pool = ParallelWorkerPool(parallel, self._get_worker_class(), start_method=start_method)
104
+ pool = ParallelWorkerPool(
105
+ parallel, self._get_worker_class(), start_method=start_method
106
+ )
101
107
  for batch in pool.ordered_map(iter_batch(documents, batch_size), **params):
102
108
  yield from self._post_process_onnx_output(batch)
103
109
 
@@ -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,7 +22,13 @@ 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
@@ -48,7 +48,7 @@ class SparseTextEmbedding(SparseTextEmbeddingBase):
48
48
 
49
49
  for EMBEDDING_MODEL_TYPE in self.EMBEDDINGS_REGISTRY:
50
50
  supported_models = EMBEDDING_MODEL_TYPE.list_supported_models()
51
- if any(model_name == model["model"] for model in supported_models):
51
+ if any(model_name.lower() == model["model"].lower() for model in supported_models):
52
52
  self.model = EMBEDDING_MODEL_TYPE(model_name, cache_dir, threads, **kwargs)
53
53
  return
54
54
 
@@ -1,4 +1,4 @@
1
- from typing import Any, Dict, Iterable, List, Optional, Tuple, Union
1
+ from typing import Any, Dict, Iterable, List, Optional, Tuple, Union, Type
2
2
 
3
3
  import numpy as np
4
4
 
@@ -8,7 +8,16 @@ from fastembed.sparse.sparse_embedding_base import SparseEmbedding, SparseTextEm
8
8
 
9
9
  supported_splade_models = [
10
10
  {
11
- "model": "prithvida/SPLADE_PP_en_v1",
11
+ "model": "prithvida/Splade_PP_en_v1",
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
+ },
19
+ {
20
+ "model": "prithivida/Splade_PP_en_v1",
12
21
  "vocab_size": 30522,
13
22
  "description": "Independent Implementation of SPLADE++ Model for English",
14
23
  "size_in_GB": 0.532,
@@ -21,7 +30,9 @@ supported_splade_models = [
21
30
 
22
31
  class SpladePP(SparseTextEmbeddingBase, OnnxModel[SparseEmbedding]):
23
32
  @classmethod
24
- def _post_process_onnx_output(cls, output: Tuple[np.ndarray, np.ndarray]) -> Iterable[SparseEmbedding]:
33
+ def _post_process_onnx_output(
34
+ cls, output: Tuple[np.ndarray, np.ndarray]
35
+ ) -> Iterable[SparseEmbedding]:
25
36
  logits, attention_mask = output
26
37
  relu_log = np.log(1 + np.maximum(logits, 0))
27
38
 
@@ -105,6 +116,10 @@ class SpladePP(SparseTextEmbeddingBase, OnnxModel[SparseEmbedding]):
105
116
  parallel=parallel,
106
117
  )
107
118
 
119
+ @classmethod
120
+ def _get_worker_class(cls) -> Type[EmbeddingWorker]:
121
+ return SpladePPEmbeddingWorker
122
+
108
123
 
109
124
  class SpladePPEmbeddingWorker(EmbeddingWorker):
110
125
  def init_embedding(
@@ -20,7 +20,7 @@ supported_multilingual_e5_models = [
20
20
  "model": "sentence-transformers/paraphrase-multilingual-mpnet-base-v2",
21
21
  "dim": 768,
22
22
  "description": "Sentence-transformers model for tasks like clustering or semantic search",
23
- "size_in_GB": 1.11,
23
+ "size_in_GB": 1.00,
24
24
  "sources": {
25
25
  "hf": "xenova/paraphrase-multilingual-mpnet-base-v2",
26
26
  },
@@ -11,14 +11,14 @@ 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
16
  },
17
17
  {
18
18
  "model": "jinaai/jina-embeddings-v2-small-en",
19
19
  "dim": 512,
20
20
  "description": "English embedding model supporting 8192 sequence length",
21
- "size_in_GB": 0.13,
21
+ "size_in_GB": 0.12,
22
22
  "sources": {"hf": "xenova/jina-embeddings-v2-small-en"},
23
23
  },
24
24
  ]
@@ -49,7 +49,9 @@ class JinaOnnxEmbedding(OnnxTextEmbedding):
49
49
  return supported_jina_models
50
50
 
51
51
  @classmethod
52
- def _post_process_onnx_output(cls, output: Tuple[np.ndarray, np.ndarray]) -> Iterable[np.ndarray]:
52
+ def _post_process_onnx_output(
53
+ cls, output: Tuple[np.ndarray, np.ndarray]
54
+ ) -> Iterable[np.ndarray]:
53
55
  embeddings, attn_mask = output
54
56
  return normalize(cls.mean_pooling(embeddings, attn_mask)).astype(np.float32)
55
57
 
@@ -12,7 +12,7 @@ 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
  },
@@ -21,26 +21,17 @@ supported_onnx_models = [
21
21
  "model": "BAAI/bge-base-en-v1.5",
22
22
  "dim": 768,
23
23
  "description": "Base English model, v1.5",
24
- "size_in_GB": 0.44,
24
+ "size_in_GB": 0.21,
25
25
  "sources": {
26
26
  "url": "https://storage.googleapis.com/qdrant-fastembed/fast-bge-base-en-v1.5.tar.gz",
27
27
  "hf": "qdrant/bge-base-en-v1.5-onnx-q",
28
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
- },
38
- },
39
30
  {
40
31
  "model": "BAAI/bge-large-en-v1.5",
41
32
  "dim": 1024,
42
33
  "description": "Large English model, v1.5",
43
- "size_in_GB": 1.34,
34
+ "size_in_GB": 1.20,
44
35
  "sources": {
45
36
  "hf": "qdrant/bge-large-en-v1.5-onnx",
46
37
  },
@@ -49,7 +40,7 @@ supported_onnx_models = [
49
40
  "model": "BAAI/bge-small-en",
50
41
  "dim": 384,
51
42
  "description": "Fast English model",
52
- "size_in_GB": 0.2,
43
+ "size_in_GB": 0.13,
53
44
  "sources": {
54
45
  "url": "https://storage.googleapis.com/qdrant-fastembed/BAAI-bge-small-en.tar.gz",
55
46
  },
@@ -69,9 +60,8 @@ supported_onnx_models = [
69
60
  "model": "BAAI/bge-small-en-v1.5",
70
61
  "dim": 384,
71
62
  "description": "Fast and Default English model",
72
- "size_in_GB": 0.13,
63
+ "size_in_GB": 0.067,
73
64
  "sources": {
74
- "url": "https://storage.googleapis.com/qdrant-fastembed/fast-bge-small-en-v1.5.tar.gz",
75
65
  "hf": "qdrant/bge-small-en-v1.5-onnx-q",
76
66
  },
77
67
  },
@@ -79,12 +69,12 @@ supported_onnx_models = [
79
69
  "model": "BAAI/bge-small-zh-v1.5",
80
70
  "dim": 512,
81
71
  "description": "Fast and recommended Chinese model",
82
- "size_in_GB": 0.1,
72
+ "size_in_GB": 0.09,
83
73
  "sources": {
84
74
  "url": "https://storage.googleapis.com/qdrant-fastembed/fast-bge-small-zh-v1.5.tar.gz",
85
75
  },
86
76
  },
87
- { # todo: it is not a flag embedding
77
+ {
88
78
  "model": "sentence-transformers/all-MiniLM-L6-v2",
89
79
  "dim": 384,
90
80
  "description": "Sentence Transformer model, MiniLM-L6-v2",
@@ -98,7 +88,7 @@ supported_onnx_models = [
98
88
  "model": "sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2",
99
89
  "dim": 384,
100
90
  "description": "Sentence Transformer model, paraphrase-multilingual-MiniLM-L12-v2",
101
- "size_in_GB": 0.46,
91
+ "size_in_GB": 0.22,
102
92
  "sources": {
103
93
  "hf": "qdrant/paraphrase-multilingual-MiniLM-L12-v2-onnx-Q",
104
94
  },
@@ -107,7 +97,7 @@ supported_onnx_models = [
107
97
  "model": "nomic-ai/nomic-embed-text-v1",
108
98
  "dim": 768,
109
99
  "description": "8192 context length english model",
110
- "size_in_GB": 0.54,
100
+ "size_in_GB": 0.52,
111
101
  "sources": {
112
102
  "hf": "nomic-ai/nomic-embed-text-v1",
113
103
  },
@@ -116,7 +106,7 @@ supported_onnx_models = [
116
106
  "model": "nomic-ai/nomic-embed-text-v1.5",
117
107
  "dim": 768,
118
108
  "description": "8192 context length english model",
119
- "size_in_GB": 0.54,
109
+ "size_in_GB": 0.52,
120
110
  "sources": {
121
111
  "hf": "nomic-ai/nomic-embed-text-v1.5",
122
112
  },
@@ -125,7 +115,7 @@ supported_onnx_models = [
125
115
  "model": "thenlper/gte-large",
126
116
  "dim": 1024,
127
117
  "description": "Large general text embeddings model",
128
- "size_in_GB": 1.34,
118
+ "size_in_GB": 1.20,
129
119
  "sources": {
130
120
  "hf": "qdrant/gte-large-onnx",
131
121
  },
@@ -143,6 +133,15 @@ supported_onnx_models = [
143
133
  # "https://storage.googleapis.com/qdrant-fastembed/sentence-transformers-all-MiniLM-L6-v2.tar.gz"
144
134
  # ]
145
135
  # }
136
+ {
137
+ "model": "mixedbread-ai/mxbai-embed-large-v1",
138
+ "dim": 1024,
139
+ "description": "MixedBread Base sentence embedding model, does well on MTEB",
140
+ "size_in_GB": 0.64,
141
+ "sources": {
142
+ "hf": "mixedbread-ai/mxbai-embed-large-v1",
143
+ },
144
+ },
146
145
  ]
147
146
 
148
147
 
@@ -230,7 +229,9 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxModel[np.ndarray]):
230
229
  return onnx_input
231
230
 
232
231
  @classmethod
233
- def _post_process_onnx_output(cls, output: Tuple[np.ndarray, np.ndarray]) -> Iterable[np.ndarray]:
232
+ def _post_process_onnx_output(
233
+ cls, output: Tuple[np.ndarray, np.ndarray]
234
+ ) -> Iterable[np.ndarray]:
234
235
  embeddings, _ = output
235
236
  return normalize(embeddings[:, 0]).astype(np.float32)
236
237
 
@@ -55,7 +55,7 @@ class TextEmbedding(TextEmbeddingBase):
55
55
 
56
56
  for EMBEDDING_MODEL_TYPE in self.EMBEDDINGS_REGISTRY:
57
57
  supported_models = EMBEDDING_MODEL_TYPE.list_supported_models()
58
- if any(model_name == model["model"] for model in supported_models):
58
+ if any(model_name.lower() == model["model"].lower() for model in supported_models):
59
59
  self.model = EMBEDDING_MODEL_TYPE(model_name, cache_dir, threads, **kwargs)
60
60
  return
61
61
 
@@ -6,7 +6,13 @@ 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
@@ -1,6 +1,6 @@
1
1
  [tool.poetry]
2
2
  name = "fastembed"
3
- version = "0.2.4"
3
+ version = "0.2.6"
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"
@@ -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