Compare commits

..
37 Commits
Author SHA1 Message Date
Andrey Vasnetsov f02d713e93 Merge pull request #24 from qdrant/streaming-inference
implement data-parallel inference and up version
2023-10-16 14:31:40 +02:00
Andrey Vasnetsov c901da0820 Merge pull request #26 from qdrant/add_model_size
* feat(embedding.py): add size_in_GB information for each model
2023-10-16 14:26:13 +02:00
Nirant f79de09ff7 Merge branch 'streaming-inference' into add_model_size 2023-10-16 17:54:33 +05:30
NirantK 7799180b18 * fix(embedding.py): update return type of list_supported_models method to include Union[int, float] for values in the dictionary 2023-10-16 17:52:43 +05:30
NirantK e24ea64e21 * test(test_onnx_embeddings.py): skip specific model if size_in_GB is greater than 1 2023-10-16 17:52:35 +05:30
NirantK 299c76c099 * feat(embedding.py): add size_in_GB information for each model 2023-10-16 17:50:28 +05:30
generall 35ec40b3a3 review fixes 2023-10-16 14:19:58 +02:00
NirantK b35cb28eeb * refactor(embedding.py): reorder import statements in alphabetical order
* feat(embedding.py): add optional 'threads' parameter to DefaultEmbedding constructor
2023-10-16 17:39:21 +05:30
Nirant c953a083cc Merge branch 'main' into streaming-inference 2023-10-16 17:33:36 +05:30
Nirant 66c5cf76c6 Merge pull request #25 from qdrant/support-v1.5-models
* feat(embedding.py): add support for v1.5 models
2023-10-16 17:29:33 +05:30
NirantK 9e5d37846c * feat(embedding.py): add support for BAAI/bge-small-en-v1.5 and BAAI/bge-base-en-v1.5 models
* feat(embedding.py): change default model to v1.5
2023-10-16 17:08:58 +05:30
generall c719fc696d disable large models on non-ubuntu CI 2023-10-16 13:27:28 +02:00
generall 24dc24b02d implement data-parallel inference and up version 2023-10-16 13:07:07 +02:00
NirantK c408b7e13e * chore(main.html): add utm parameters to Qdrant Cloud link 2023-10-10 19:01:49 +05:30
NirantK d2bdfee4e0 * docs(examples): update comparison notebook with more accurate description of embeddings similarity 2023-10-10 17:56:25 +05:30
NirantK ac4375516f Rename nbs; add Cosine similarity check 2023-10-10 17:55:37 +05:30
NirantK ca6f9d629a Add skeleton 2023-10-05 19:20:01 +05:30
NirantK e0e7e5721e Add generator note to the comment in Python block 2023-10-05 19:17:32 +05:30
NirantK bc402694bd Update code to handle Generator 2023-10-05 19:17:18 +05:30
Nirant 3139fb7275 Merge pull request #15 from qdrant/v0.0.5
bump version 0.0.5
2023-10-05 16:46:26 +05:30
generall 3fbc878ddd bump version 0.0.5 2023-10-04 10:20:58 +02:00
Andrey Vasnetsov fa8684f7ef Merge pull request #13 from qdrant/0.5-suggestions
remove bulky dependencies
2023-10-04 10:04:06 +02:00
generall 288cee1d16 upd dependencies 2023-10-03 21:35:46 +02:00
generall 8c6d4d2b52 remove optimum + more test + fix batching embed + ci on other machines 2023-10-03 21:19:58 +02:00
NirantK 050d80ab51 * chore(docs): update index.md with FastEmbed library information and usage examples
* feat(docs): add installation instructions for FastEmbed with Qdrant Client
2023-09-28 10:52:38 +05:30
NirantK 49c2b3e7a9 * chore(docs): update installation command for fastembed in Getting Started.ipynb
* fix(docs): remove unnecessary casting of generator to list in code cell
2023-09-28 10:43:14 +05:30
NirantK 5e2ced87f5 * chore(README.md): update links and fix formatting in README.md 2023-09-27 18:00:19 +05:30
NirantK ba1a12a0f0 * docs(README.md): update description of FastEmbed library
* fix(README.md): fix typo in the description of FastEmbed library
2023-09-27 17:59:21 +05:30
NirantK 983964c432 * docs(README.md): update links and descriptions in the README file
* feat(README.md): add usage example with Qdrant client
2023-09-27 17:57:51 +05:30
NirantK 4949158eff * chore(Usage_With_Qdrant.ipynb): update notebook title and remove experimental note
* docs(Usage_With_Qdrant.ipynb): add support for qdrant-client[fastembed] installation
2023-09-27 17:53:05 +05:30
NirantK c68c1029a4 * docs(README.md): add link to supported models dosc 2023-09-27 17:44:01 +05:30
NirantK dd2a9e9bcc * docs(Supported_Models.ipynb): update model descriptions 2023-09-27 17:41:39 +05:30
NirantK 71289926c0 * feat(Supported_Models.ipynb): add example notebook for supported models 2023-09-27 17:40:40 +05:30
NirantK 85a9cc08ec * chore(embedding.py): add list_supported_models method to Embedding class 2023-09-27 17:39:42 +05:30
NirantK 589105c84c * fix(embedding.py): handle single string input in embed_documents method 2023-09-27 17:27:56 +05:30
NirantK ac6b8c9402 * chore(pyproject.toml): add onnx dependency to the project 2023-09-27 13:26:27 +05:30
NirantK b28ff3f8d6 * chore(pyproject.toml): update version from 0.0.4 to 0.0.5a1
* chore(pyproject.toml): remove onnxruntime-silicon dependency for macOS
2023-09-26 21:25:59 +05:30
15 changed files with 2002 additions and 2367 deletions
+5 -1
View File
@@ -20,6 +20,8 @@ jobs:
- '3.11.x'
os:
- ubuntu-latest
- macos-latest
- windows-latest
runs-on: ${{ matrix.os }}
@@ -37,5 +39,7 @@ jobs:
poetry config virtualenvs.create false
poetry install --no-interaction --no-ansi
- name: Run tests
run: pytest
run: |
export IS_UBUNTU_CI=$(test "${{ matrix.os }}" = "ubuntu-latest" && echo "true" || echo "false")
pytest
shell: bash
+45 -24
View File
@@ -1,22 +1,19 @@
# ⚡️ What is FastEmbed?
FastEmbed is an easy to use -- lightweight, fast, Python library built for retrieval embedding generation.
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.
The default embedding supports "query" and "passage" prefixes for the input text. The default model is Flag Embedding, which is top of the [MTEB](https://huggingface.co/spaces/mteb/leaderboard) leaderboard. Here is an example for [Retrieval Embedding Generation](https://qdrant.github.io/fastembed/examples/Retrieval%20with%20FastEmbed/)
The default embedding supports "query" and "passage" prefixes for the input text. The default model is Flag Embedding, which is top of the [MTEB](https://huggingface.co/spaces/mteb/leaderboard) leaderboard. Here is an example for [Retrieval Embedding Generation](https://qdrant.github.io/fastembed/examples/Retrieval%20with%20FastEmbed/) and how to use [FastEmbed with Qdrant](https://qdrant.github.io/fastembed/examples/Usage_With_Qdrant/).
1. Light
1. Light & Fast
- Quantized model weights
- ONNX Runtime for inference
- No hidden dependencies on PyTorch or TensorFlow via Huggingface Transformers
- ONNX Runtime, no PyTorch dependency
- CPU-first design
- Data-parallelism for encoding of large datasets
2. Accuracy/Recall
- Better than OpenAI Ada-002
- Default is Flag Embedding, which is top of the [MTEB](https://huggingface.co/spaces/mteb/leaderboard) leaderboard
3. Fast
- About 2x faster than Huggingface (PyTorch) transformers on single queries
- Lot faster for batches!
- ONNX Runtime allows you to use dedicated runtimes for even higher throughput and lower latency
- List of [supported models](https://qdrant.github.io/fastembed/examples/Supported_Models/) - including multilingual models
## 🚀 Installation
@@ -35,27 +32,51 @@ documents: List[str] = [
"passage: Hello, World!",
"query: Hello, World!", # these are two different embedding
"passage: This is an example passage.",
# You can leave out the prefix but it's recommended
"fastembed is supported by and maintained by Qdrant."
"fastembed is supported by and maintained by Qdrant." # You can leave out the prefix but it's recommended
]
embedding_model = Embedding(model_name="BAAI/bge-base-en", max_length=512)
embeddings: List[np.ndarray] = list(embedding_model.embed(documents))
embeddings: List[np.ndarray] = list(embedding_model.embed(documents)) # Note the list() call - this is a generator
```
### Why fast?
## Usage with Qdrant
It's important we justify the "fast" in FastEmbed. FastEmbed is fast because:
Installation with Qdrant Client in Python:
1. Quantized model weights
2. ONNX Runtime which allows for fast inference on CPU and other dedicated runtimes
```bash
pip install qdrant-client[fastembed]
```
### Why light?
1. No hidden dependencies on PyTorch or TensorFlow via Huggingface Transformers
2. We do use the tokenizer from Huggingface Transformers, but it's a light dependency
Might have to use ```pip install 'qdrant-client[fastembed]'``` on zsh.
### Why accurate?
1. Better than OpenAI Ada-002
2. Top of the Embedding leaderboards e.g. [MTEB](https://huggingface.co/spaces/mteb/leaderboard)
```python
from qdrant_client import QdrantClient
# Initialize the client
client = QdrantClient(":memory:") # or QdrantClient(path="path/to/db")
# Prepare your documents, metadata, and IDs
docs = ["Qdrant has Langchain integrations", "Qdrant also has Llama Index integrations"]
metadata = [
{"source": "Langchain-docs"},
{"source": "Linkedin-docs"},
]
ids = [42, 2]
# Use the new add method
client.add(
collection_name="demo_collection",
documents=docs,
metadata=metadata,
ids=ids
)
search_result = client.query(
collection_name="demo_collection",
query_text="This is a query document"
)
print(search_result)
```
#### Similar Work
Ilyas M. wrote about using [FlagEmbeddings with Optimum](https://twitter.com/IlysMoutawwakil/status/1705215192425288017) over CUDA.
Ilyas M. wrote about using [FlagEmbeddings with Optimum](https://twitter.com/IlysMoutawwakil/status/1705215192425288017) over CUDA.
+44 -18
View File
@@ -16,12 +16,12 @@
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 1,
"id": "ada95c6a",
"metadata": {},
"outputs": [],
"source": [
"!pip install fastembed --upgrade # Install fastembed"
"!pip install fastembed --upgrade --quiet # Install fastembed "
]
},
{
@@ -34,10 +34,25 @@
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 2,
"id": "b61c6552",
"metadata": {},
"outputs": [],
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
"Asking to truncate to max_length but no maximum length is provided and the model has no predefined maximum length. Default to no truncation.\n"
]
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"torch.Size([384])\n"
]
}
],
"source": [
"from typing import List\n",
"import numpy as np\n",
@@ -51,10 +66,7 @@
"]\n",
"# Initialize the DefaultEmbedding class with the desired parameters\n",
"embedding_model = DefaultEmbedding(model_name=\"BAAI/bge-small-en\", max_length=512)\n",
"embeddings: List[np.ndarray] = list(\n",
" embedding_model.embed(documents)\n",
") # notice that we are casting the generator to a list\n",
"\n",
"embeddings: List[np.ndarray] = embedding_model.embed(documents)\n",
"print(embeddings[0].shape)"
]
},
@@ -78,7 +90,7 @@
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 3,
"id": "c0a6f634",
"metadata": {},
"outputs": [],
@@ -107,7 +119,7 @@
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 4,
"id": "145a56ce",
"metadata": {},
"outputs": [],
@@ -139,7 +151,7 @@
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 5,
"id": "272c8915",
"metadata": {},
"outputs": [],
@@ -159,14 +171,20 @@
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 6,
"id": "8013eee9",
"metadata": {},
"outputs": [],
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
"Asking to truncate to max_length but no maximum length is provided and the model has no predefined maximum length. Default to no truncation.\n"
]
}
],
"source": [
"embeddings: List[np.ndarray] = list(\n",
" embedding_model.embed(documents)\n",
") # notice that we are casting the generator to a list"
"embeddings: List[np.ndarray] = embedding_model.embed(documents)"
]
},
{
@@ -179,10 +197,18 @@
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 7,
"id": "0d8c8e08",
"metadata": {},
"outputs": [],
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"torch.Size([384])\n"
]
}
],
"source": [
"print(embeddings[0].shape) # (384,) or similar output"
]
File diff suppressed because one or more lines are too long
+115
View File
@@ -0,0 +1,115 @@
{
"cells": [
{
"cell_type": "code",
"execution_count": 1,
"metadata": {},
"outputs": [
{
"data": {
"text/html": [
"<div>\n",
"<style scoped>\n",
" .dataframe tbody tr th:only-of-type {\n",
" vertical-align: middle;\n",
" }\n",
"\n",
" .dataframe tbody tr th {\n",
" vertical-align: top;\n",
" }\n",
"\n",
" .dataframe thead th {\n",
" text-align: right;\n",
" }\n",
"</style>\n",
"<table border=\"1\" class=\"dataframe\">\n",
" <thead>\n",
" <tr style=\"text-align: right;\">\n",
" <th></th>\n",
" <th>model</th>\n",
" <th>dim</th>\n",
" <th>description</th>\n",
" </tr>\n",
" </thead>\n",
" <tbody>\n",
" <tr>\n",
" <th>0</th>\n",
" <td>BAAI/bge-small-en</td>\n",
" <td>384</td>\n",
" <td>Fast and Default English model</td>\n",
" </tr>\n",
" <tr>\n",
" <th>1</th>\n",
" <td>BAAI/bge-base-en</td>\n",
" <td>768</td>\n",
" <td>Base English model</td>\n",
" </tr>\n",
" <tr>\n",
" <th>2</th>\n",
" <td>sentence-transformers/all-MiniLM-L6-v2</td>\n",
" <td>384</td>\n",
" <td>Sentence Transformer model, MiniLM-L6-v2</td>\n",
" </tr>\n",
" <tr>\n",
" <th>3</th>\n",
" <td>intfloat/multilingual-e5-large</td>\n",
" <td>1024</td>\n",
" <td>Multilingual model, e5-large. Recommend using this model for non-English languages. Recommend using this via Torch implementation of FastEmbed</td>\n",
" </tr>\n",
" </tbody>\n",
"</table>\n",
"</div>"
],
"text/plain": [
" model dim \\\n",
"0 BAAI/bge-small-en 384 \n",
"1 BAAI/bge-base-en 768 \n",
"2 sentence-transformers/all-MiniLM-L6-v2 384 \n",
"3 intfloat/multilingual-e5-large 1024 \n",
"\n",
" description \n",
"0 Fast and Default English model \n",
"1 Base English model \n",
"2 Sentence Transformer model, MiniLM-L6-v2 \n",
"3 Multilingual model, e5-large. Recommend using this model for non-English languages. Recommend using this via Torch implementation of FastEmbed "
]
},
"execution_count": 1,
"metadata": {},
"output_type": "execute_result"
}
],
"source": [
"%load_ext autoreload\n",
"%autoreload 2\n",
"\n",
"from fastembed.embedding import Embedding\n",
"import pandas as pd\n",
"pd.set_option('display.max_colwidth', None)\n",
"pd.DataFrame(Embedding.list_supported_models())"
]
}
],
"metadata": {
"kernelspec": {
"display_name": "fst",
"language": "python",
"name": "python3"
},
"language_info": {
"codemirror_mode": {
"name": "ipython",
"version": 3
},
"file_extension": ".py",
"mimetype": "text/x-python",
"name": "python",
"nbconvert_exporter": "python",
"pygments_lexer": "ipython3",
"version": "3.9.17"
},
"orig_nbformat": 4
},
"nbformat": 4,
"nbformat_minor": 2
}
File diff suppressed because one or more lines are too long
+75 -44
View File
@@ -5,9 +5,7 @@
"cell_type": "markdown",
"metadata": {},
"source": [
"# [Experimental] Usage With Qdrant\n",
"\n",
"> **Note:** This notebook is experimental and is subject to change. For working with this, use the dev branch of QdrantClient.\n",
"# Usage With Qdrant\n",
"\n",
"This notebook demonstrates how to use FastEmbed and Qdrant to perform vector search and retrieval. Qdrant is an open-source vector similarity search engine that is used to store, organize, and query collections of high-dimensional vectors. \n",
"\n",
@@ -28,13 +26,20 @@
},
{
"cell_type": "code",
"execution_count": 14,
"execution_count": 1,
"metadata": {},
"outputs": [],
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"\u001b[33mDEPRECATION: pytorch-lightning 1.6.5 has a non-standard dependency specifier torch>=1.8.*. pip 23.3 will enforce this behaviour change. A possible replacement is to upgrade to a newer version of pytorch-lightning or contact the author to suggest that they release a version with a conforming dependency specifiers. Discussion can be found at https://github.com/pypa/pip/issues/12063\u001b[0m\u001b[33m\n",
"\u001b[0m"
]
}
],
"source": [
"# !pip install fastembed --quiet --upgrade\n",
"\n",
"# !pip install git+https://github.com/qdrant/qdrant_client.git@dev"
"!pip install 'qdrant-client[fastembed]' --quiet --upgrade"
]
},
{
@@ -46,7 +51,7 @@
},
{
"cell_type": "code",
"execution_count": 15,
"execution_count": 2,
"metadata": {},
"outputs": [],
"source": [
@@ -68,7 +73,7 @@
},
{
"cell_type": "code",
"execution_count": 16,
"execution_count": 3,
"metadata": {},
"outputs": [],
"source": [
@@ -105,25 +110,32 @@
},
{
"cell_type": "code",
"execution_count": 17,
"execution_count": 4,
"metadata": {},
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
"Asking to truncate to max_length but no maximum length is provided and the model has no predefined maximum length. Default to no truncation.\n"
]
},
{
"data": {
"text/plain": [
"['901ff9d7a90a4e56afe655d1de8e7c06',\n",
" '7bae5aa894164398a9be68d47d72ff7a',\n",
" 'd940bba124e24469ae3166c3110c62f8',\n",
" 'c8dcf956a4f444c6bfe69a00cbeb85ce',\n",
" 'd2aeb24c51f549c5b21851be05cb4048',\n",
" '6d4c672bafef4db68ae72c8986ee8a72',\n",
" 'ff198613768a4361a6d8230b40ee7f78',\n",
" 'a133ad72876b48d48e716533c2d78cf0',\n",
" 'bdf2b016136e4f57b7ac4b65534031b8',\n",
" '5f4f279f304c47839e1a74b0f28de136']"
"['77e1e4724dd243b08608f57d5692f6aa',\n",
" '74841e5dc3594646bda2c6a6d2795dbd',\n",
" '6ef39a9445604d0da84d04f760cd7cf7',\n",
" 'e659503d3b3748ef90f23c778274835b',\n",
" 'b999675068cd413f93faa0cc890c3819',\n",
" '8e452f2935cf4e4b80d8eea68c2aad58',\n",
" '28ed4fd4592c48c9a0519618d51bb86e',\n",
" '59378c784c5f49109bef65fdc4061334',\n",
" 'a78c9b598f7942749156334283a6f24f',\n",
" 'f72bb24701c64fabb0182c9e757b581b']"
]
},
"execution_count": 17,
"execution_count": 4,
"metadata": {},
"output_type": "execute_result"
}
@@ -150,38 +162,57 @@
},
{
"cell_type": "code",
"execution_count": 18,
"execution_count": 5,
"metadata": {},
"outputs": [
{
"data": {
"text/plain": [
"[42, 2]"
]
},
"execution_count": 5,
"metadata": {},
"output_type": "execute_result"
}
],
"source": [
"# Prepare your documents, metadata, and IDs\n",
"docs = [\"Qdrant has Langchain integrations\", \"Qdrant also has Llama Index integrations\"]\n",
"metadata = [\n",
" {\"source\": \"Langchain-docs\"},\n",
" {\"source\": \"Linkedin-docs\"},\n",
"]\n",
"ids = [42, 2]\n",
"\n",
"# Use the new add method\n",
"client.add(\n",
" collection_name=\"demo_collection\",\n",
" documents=docs,\n",
" metadata=metadata,\n",
" ids=ids\n",
")"
]
},
{
"cell_type": "code",
"execution_count": 6,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"Top 5 results:\n",
"Rank 1: Maharana Pratap was a Rajput warrior king from Mewar. Score: 0.77\n",
"Rank 2: Maharana Pratap is considered a symbol of Rajput resistance against foreign rule. Score: 0.77\n",
"Rank 3: His legacy is celebrated in Rajasthan through festivals and monuments. Score: 0.69\n",
"Rank 4: He had 11 wives and 17 sons, including Amar Singh I who succeeded him as ruler of Mewar. Score: 0.68\n",
"Rank 5: He fought against the Mughal Empire led by Akbar. Score: 0.67\n"
"[QueryResponse(id='42', embedding=None, metadata={'document': 'Qdrant has Langchain integrations', 'source': 'Langchain-docs'}, document='Qdrant has Langchain integrations', score=0.8496814051311954), QueryResponse(id='2', embedding=None, metadata={'document': 'Qdrant also has Llama Index integrations', 'source': 'Linkedin-docs'}, document='Qdrant also has Llama Index integrations', score=0.8478494193031256)]\n"
]
}
],
"source": [
"from qdrant_client.qdrant_fastembed import QueryResponse\n",
"\n",
"\n",
"def print_top_k_results(results: List[QueryResponse], k: int = 5):\n",
" print(f\"Top {k} results:\")\n",
" for i, result in enumerate(results[:k]):\n",
" print(f\"Rank {i + 1}: {result.document}. Score: {result.score:.2f}\")\n",
"\n",
"\n",
"query_text = \"Who is Maharana Pratap?\"\n",
"results = client.query(\n",
" collection_name=\"test_collection\", query_text=query_text, limit=7\n",
") # Returns limit most relevant documents\n",
"\n",
"print_top_k_results(results)"
"search_result = client.query(\n",
" collection_name=\"demo_collection\",\n",
" query_text=[\"This is a query document\"]\n",
")\n",
"print(search_result)"
]
},
{
File diff suppressed because one or more lines are too long
+49 -18
View File
@@ -1,10 +1,19 @@
# ⚡️ What is FastEmbed?
FastEmbed is an easy to use -- lightweight, fast, Python library built for retrieval embedding generation.
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.
The default embedding supports "query" and "passage" prefixes for the input text. The default model is [Flag Embedding](https://github.com/FlagOpen/FlagEmbedding), which is top of the [MTEB](https://huggingface.co/spaces/mteb/leaderboard) leaderboard.
The default embedding supports "query" and "passage" prefixes for the input text. The default model is Flag Embedding, which is top of the [MTEB](https://huggingface.co/spaces/mteb/leaderboard) leaderboard. Here is an example for [Retrieval Embedding Generation](https://qdrant.github.io/fastembed/examples/Retrieval%20with%20FastEmbed/) and how to use [FastEmbed with Qdrant](https://qdrant.github.io/fastembed/examples/Usage_With_Qdrant/).
Advanced user? Skip ahead to [Retrieval with FastEmbed](https://qdrant.github.io/fastembed/examples/Retrieval_with_FastEmbed/)
1. Light & Fast
- Quantized model weights
- ONNX Runtime for inference via [Optimum](github.com/huggingface/optimum)
2. Accuracy/Recall
- Better than OpenAI Ada-002
- Default is Flag Embedding, which is top of the [MTEB](https://huggingface.co/spaces/mteb/leaderboard) leaderboard
- List of [supported models](https://qdrant.github.io/fastembed/examples/Supported_Models/) - including multilingual models
## 🚀 Installation
To install the FastEmbed library, pip works:
@@ -15,31 +24,53 @@ pip install fastembed
## 📖 Usage
```python
from fastembed.embedding import DefaultEmbedding
from fastembed.embedding import FlagEmbedding as Embedding
documents: List[str] = [
"passage: Hello, World!",
"query: Hello, World!", # these are two different embedding
"passage: This is an example passage.",
# You can leave out the prefix but it's recommended
"fastembed is supported by and maintained by Qdrant."
"fastembed is supported by and maintained by Qdrant." # You can leave out the prefix but it's recommended
]
embedding_model = DefaultEmbedding()
embeddings: List[np.ndarray] = list(embedding_model.embed(documents))
embedding_model = Embedding(model_name="BAAI/bge-base-en", max_length=512)
embeddings: List[np.ndarray] = embedding_model.embed(documents) # If you use
```
## 🚒 Under the hood
## Usage with Qdrant
### Why fast?
Installation with Qdrant Client in Python:
It's important we justify the "fast" in FastEmbed. FastEmbed is fast because:
```bash
pip install qdrant-client[fastembed]
```
1. Quantized model weights
2. ONNX Runtime which allows for inference on CPU, GPU, and other dedicated runtimes
Might have to use ```pip install 'qdrant-client[fastembed]'``` on zsh.
### Why light?
1. No hidden dependencies on PyTorch or TensorFlow via Huggingface Transformers
```python
from qdrant_client import QdrantClient
### Why accurate?
1. Better than OpenAI Ada-002
2. Top of the Embedding leaderboards e.g. [MTEB](https://huggingface.co/spaces/mteb/leaderboard)
# Initialize the client
client = QdrantClient(":memory:") # or QdrantClient(path="path/to/db")
# Prepare your documents, metadata, and IDs
docs = ["Qdrant has Langchain integrations", "Qdrant also has Llama Index integrations"]
metadata = [
{"source": "Langchain-docs"},
{"source": "Linkedin-docs"},
]
ids = [42, 2]
# Use the new add method
client.add(
collection_name="demo_collection",
documents=docs,
metadata=metadata,
ids=ids
)
search_result = client.query(
collection_name="demo_collection",
query_text="This is a query document"
)
print(search_result)
```
+2 -1
View File
@@ -20,7 +20,8 @@
</span>
<strong>Qdrant Discord server</strong>
</a>
to get help and share your work! Or check out <a rel="me" href="https://login.cloud.qdrant.io/">Qdrant Cloud</a> to
to get help and share your work! Or check out <a rel="me"
href="https://cloud.qdrant.io?utm_source=twitter&utm_medium=website&utm_campaign=fastembed">Qdrant Cloud</a> to
get started with vector search!
</div>
{% endblock %}
+237 -42
View File
@@ -1,24 +1,159 @@
import json
import os
import shutil
import tarfile
from abc import ABC, abstractmethod
from itertools import islice
from multiprocessing import get_all_start_methods
from pathlib import Path
from typing import Iterable, List
from typing import Any, Dict, Generator, Iterable, List, Optional, Tuple, Union
import numpy as np
import onnxruntime as ort
import requests
from optimum.onnxruntime import ORTModelForFeatureExtraction
from tokenizers import AddedToken, Tokenizer
from tqdm import tqdm
from transformers import AutoTokenizer
from fastembed.parallel_processor import ParallelWorkerPool, Worker
def normalize(input_array, p=2.0, dim=1, eps=1e-12):
def iter_batch(iterable: Union[Iterable, Generator], size: int) -> Iterable:
"""
>>> list(iter_batch([1,2,3,4,5], 3))
[[1, 2, 3], [4, 5]]
"""
source_iter = iter(iterable)
while source_iter:
b = list(islice(source_iter, size))
if len(b) == 0:
break
yield b
def normalize(input_array, p=2, dim=1, eps=1e-12):
# Calculate the Lp norm along the specified dimension
norm = np.linalg.norm(input_array, ord=p, axis=dim, keepdims=True)
norm = np.maximum(norm, eps) # Avoid division by zero
norm = np.maximum(norm, eps) # Avoid division by zero
normalized_array = input_array / norm
return normalized_array
class EmbeddingModel(ABC):
@classmethod
def load_tokenizer(cls, model_dir: Path, max_length: int = 512) -> Tokenizer:
config_path = model_dir / "config.json"
if not config_path.exists():
raise ValueError(f"Could not find config.json in {model_dir}")
tokenizer_path = model_dir / "tokenizer.json"
if not tokenizer_path.exists():
raise ValueError(f"Could not find tokenizer.json in {model_dir}")
tokenizer_config_path = model_dir / "tokenizer_config.json"
if not tokenizer_config_path.exists():
raise ValueError(f"Could not find tokenizer_config.json in {model_dir}")
tokens_map_path = model_dir / "special_tokens_map.json"
if not tokens_map_path.exists():
raise ValueError(f"Could not find special_tokens_map.json in {model_dir}")
config = json.load(open(str(config_path)))
tokenizer_config = json.load(open(str(tokenizer_config_path)))
tokens_map = json.load(open(str(tokens_map_path)))
tokenizer = Tokenizer.from_file(str(tokenizer_path))
tokenizer.enable_truncation(max_length=min(tokenizer_config["model_max_length"], max_length))
tokenizer.enable_padding(pad_id=config["pad_token_id"], pad_token=tokenizer_config["pad_token"])
for token in tokens_map.values():
if isinstance(token, str):
tokenizer.add_special_tokens([token])
elif isinstance(token, dict):
tokenizer.add_special_tokens([AddedToken(**token)])
return tokenizer
def __init__(
self,
path: Path,
model_name: str,
max_length: int = 512,
max_threads: int = None,
):
self.path = path
self.model_name = model_name
model_path = self.path / "model.onnx"
optimized_model_path = self.path / "model_optimized.onnx"
# List of Execution Providers: https://onnxruntime.ai/docs/execution-providers
onnx_providers = ["CPUExecutionProvider"]
if not model_path.exists():
# Rename file model_optimized.onnx to model.onnx if it exists
if optimized_model_path.exists():
optimized_model_path.rename(model_path)
else:
raise ValueError(f"Could not find model.onnx in {self.path}")
# Hacky support for multilingual model
self.exclude_token_type_ids = False
if model_name == "intfloat/multilingual-e5-large":
self.exclude_token_type_ids = True
so = ort.SessionOptions()
so.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
if max_threads is not None:
so.intra_op_num_threads = max_threads
so.inter_op_num_threads = max_threads
self.tokenizer = self.load_tokenizer(self.path, max_length=max_length)
self.model = ort.InferenceSession(str(model_path), providers=onnx_providers, sess_options=so)
def onnx_embed(self, documents: List[str]) -> np.ndarray:
encoded = self.tokenizer.encode_batch(documents)
input_ids = np.array([e.ids for e in encoded])
attention_mask = np.array([e.attention_mask for e in encoded])
onnx_input = {
"input_ids": np.array(input_ids, dtype=np.int64),
"attention_mask": np.array(attention_mask, dtype=np.int64),
}
if not self.exclude_token_type_ids:
onnx_input["token_type_ids"] = np.array(
[np.zeros(len(e), dtype=np.int64) for e in input_ids], dtype=np.int64
)
model_output = self.model.run(None, onnx_input)
last_hidden_state = model_output[0][:, 0]
embeddings = normalize(last_hidden_state).astype(np.float32)
return embeddings
class EmbeddingWorker(Worker):
def __init__(
self,
path: Path,
model_name: str,
max_length: int = 512,
):
self.model = EmbeddingModel(path=path, model_name=model_name, max_length=max_length, max_threads=1)
@classmethod
def start(cls, path: Path, model_name: str, max_length: int = 512, **kwargs: Any) -> "EmbeddingWorker":
return cls(
path=path,
model_name=model_name,
max_length=max_length,
)
def process(self, items: Iterable[Tuple[int, Any]]) -> Iterable[Tuple[int, Any]]:
for idx, batch in items:
embeddings = self.model.onnx_embed(batch)
yield idx, embeddings
class Embedding(ABC):
"""
Abstract class for embeddings.
@@ -43,6 +178,50 @@ class Embedding(ABC):
def embed(self, texts: List[str]) -> List[np.ndarray]:
raise NotImplementedError
@classmethod
def list_supported_models(cls) -> List[Dict[str, Union[str, Union[int, float]]]]:
"""
Lists the supported models.
"""
return [
{
"model": "BAAI/bge-small-en",
"dim": 384,
"description": "Fast English model",
"size_in_GB": 0.2
},
{
"model": "BAAI/bge-small-en-v1.5",
"dim": 384,
"description": "Fast and Default English model",
"size_in_GB": 0.13
},
{
"model": "BAAI/bge-base-en",
"dim": 768,
"description": "Base English model",
"size_in_GB": 0.5
},
{
"model": "BAAI/bge-base-en-v1.5",
"dim": 768,
"description": "Base English model, v1.5",
"size_in_GB": 0.44
},
{
"model": "sentence-transformers/all-MiniLM-L6-v2",
"dim": 384,
"description": "Sentence Transformer model, MiniLM-L6-v2",
"size_in_GB": 0.09
},
{
"model": "intfloat/multilingual-e5-large",
"dim": 1024,
"description": "Multilingual model, e5-large. Recommend using this model for non-English languages",
"size_in_GB": 2.24
},
]
@classmethod
def download_file_from_gcs(cls, url: str, output_path: str, show_progress: bool = True) -> str:
"""
@@ -159,7 +338,7 @@ class Embedding(ABC):
try:
self.download_file_from_gcs(
f"https://storage.googleapis.com/qdrant-fastembed/{fast_model_name}.tar.gz",
output_path=str(model_tar_gz),
output_path=str(model_tar_gz),
)
except PermissionError:
simple_model_name = model_name.replace("/", "-")
@@ -190,7 +369,7 @@ class Embedding(ABC):
for i in range(0, len(texts), batch_size):
# Prepend "passage: " to each text
yield from self.embed([f"passage: {t}" for t in texts[i : i + batch_size]])
yield from self.embed([f"passage: {t}" for t in texts[i: i + batch_size]])
def query_embed(self, query: str) -> Iterable[np.ndarray]:
"""
@@ -220,60 +399,79 @@ class FlagEmbedding(Embedding):
"""
def __init__(
self,
model_name: str = "BAAI/bge-small-en",
max_length: int = 512,
cache_dir: str = None,
self,
model_name: str = "BAAI/bge-small-en",
max_length: int = 512,
cache_dir: str = None,
threads: int = None,
):
"""
Args:
model_name (str): The name of the model to use.
max_length (int, optional): The maximum number of tokens. Defaults to 512. Unknown behavior for values > 512.
cache_dir (str, optional): The path to the cache directory. Defaults to `local_cache` in the current directory.
threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.
Raises:
ValueError: If the model_name is not in the format <org>/<model> e.g. BAAI/bge-base-en.
"""
self.model_name = model_name
if cache_dir is None:
cache_dir = Path(".").resolve() / "local_cache"
cache_dir.mkdir(parents=True, exist_ok=True)
model_dir = self.retrieve_model(model_name, cache_dir)
if not (model_dir / "tokenizer.json").exists():
raise ValueError(f"Could not find tokenizer.json in {model_dir}")
if not (model_dir / "model.onnx").exists():
# Rename file model_optimized.onnx to model.onnx if it exists
if (model_dir / "model_optimized.onnx").exists():
(model_dir / "model_optimized.onnx").rename(model_dir / "model.onnx")
else:
raise ValueError(f"Could not find model.onnx in {model_dir}")
self.tokenizer = AutoTokenizer.from_pretrained(str(model_dir))
self.model = ORTModelForFeatureExtraction.from_pretrained(str(model_dir))
self._cache_dir = cache_dir
self._model_dir = self.retrieve_model(model_name, cache_dir)
self._max_length = max_length
def onnx_embed(self, documents: List[str]) -> Iterable[np.ndarray]:
encoded_input = self.tokenizer(documents, padding=True, truncation=True, return_tensors='pt')
model_output = self.model(**encoded_input)
embeddings = model_output[0][:, 0]
return normalize(embeddings, p=2, dim=1)
self.model = EmbeddingModel(self._model_dir, self.model_name, max_length=max_length,
max_threads=threads)
def embed(self, documents: List[str], batch_size: int = 256) -> Iterable[np.ndarray]:
def embed(
self, documents: Union[str, Iterable[str]], batch_size: int = 256, parallel: int = None
) -> Iterable[np.ndarray]:
"""
Encode a list of documents into list of embeddings.
We use mean pooling with attention so that the model can handle variable-length inputs.
Args:
documents: List of documents to embed
documents: Iterator of documents or single document to embed
batch_size: Batch size for encoding -- higher values will use more memory, but be faster
parallel:
If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.
If 0, use all available cores.
If None, don't use data-parallel processing, use default onnxruntime threading instead.
Returns:
List of embeddings, one per document
"""
# TODO: Replace loop with parallelized batching
if len(documents) >= batch_size:
for i in range(0, len(documents), batch_size):
batch = documents[i : i + batch_size]
return self.onnx_embed(batch)
is_small = False
if isinstance(documents, str):
documents = [documents]
is_small = True
if isinstance(documents, list):
if len(documents) < batch_size:
is_small = True
if parallel == 0:
parallel = os.cpu_count()
if parallel is None or is_small:
for batch in iter_batch(documents, batch_size):
yield from self.model.onnx_embed(batch)
else:
return self.onnx_embed(documents)
start_method = "forkserver" if "forkserver" in get_all_start_methods() else "spawn"
params = {
"path": self._model_dir,
"model_name": self.model_name,
"max_length": self._max_length,
}
pool = ParallelWorkerPool(parallel, EmbeddingWorker, start_method=start_method)
for batch in pool.ordered_map(iter_batch(documents, batch_size), **params):
yield from batch
class DefaultEmbedding(FlagEmbedding):
@@ -286,14 +484,12 @@ class DefaultEmbedding(FlagEmbedding):
def __init__(
self,
model_name: str = "BAAI/bge-small-en",
onnx_providers: List[str] = None,
model_name: str = "BAAI/bge-small-en-v1.5",
max_length: int = 512,
cache_dir: str = None,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
):
# if onnx_providers is None:
# onnx_providers = [ONNXProviders.CPU]
super().__init__(model_name, max_length=max_length, cache_dir=cache_dir)
super().__init__(model_name, max_length=max_length, cache_dir=cache_dir, threads=threads)
class OpenAIEmbedding(Embedding):
@@ -306,4 +502,3 @@ class OpenAIEmbedding(Embedding):
# Use your OpenAI model to embed the texts
# return self.model.embed(texts)
raise NotImplementedError
raise NotImplementedError
+207
View File
@@ -0,0 +1,207 @@
import logging
import os
from collections import defaultdict
from enum import Enum
from multiprocessing import Queue, get_context
from multiprocessing.context import BaseContext
from multiprocessing.process import BaseProcess
from multiprocessing.sharedctypes import Synchronized as BaseValue
from queue import Empty
from typing import Any, Dict, Iterable, List, Optional, Type, Tuple
# Single item should be processed in less than:
processing_timeout = 10 * 60 # seconds
max_internal_batch_size = 200
class QueueSignals(str, Enum):
stop = "stop"
confirm = "confirm"
error = "error"
class Worker:
@classmethod
def start(cls, **kwargs: Any) -> "Worker":
raise NotImplementedError()
def process(self, items: Iterable[Tuple[int, Any]]) -> Iterable[Tuple[int, Any]]:
raise NotImplementedError()
def _worker(
worker_class: Type[Worker],
input_queue: Queue,
output_queue: Queue,
num_active_workers: BaseValue,
worker_id: int,
kwargs: Optional[Dict[str, Any]] = None,
) -> None:
"""
A worker that pulls data pints off the input queue, and places the execution result on the output queue.
When there are no data pints left on the input queue, it decrements
num_active_workers to signal completion.
"""
if kwargs is None:
kwargs = {}
logging.info(f"Reader worker: {worker_id} PID: {os.getpid()}")
try:
worker = worker_class.start(**kwargs)
# Keep going until you get an item that's None.
def input_queue_iterable() -> Iterable[Any]:
while True:
item = input_queue.get()
if item == QueueSignals.stop:
break
yield item
for processed_item in worker.process(input_queue_iterable()):
output_queue.put(processed_item)
except Exception as e: # pylint: disable=broad-except
logging.exception(e)
output_queue.put(QueueSignals.error)
finally:
# It's important that we close and join the queue here before
# decrementing num_active_workers. Otherwise our parent may join us
# before the queue's feeder thread has passed all buffered items to
# the underlying pipe resulting in a deadlock.
#
# See:
# https://docs.python.org/3.6/library/multiprocessing.html?highlight=process#pipes-and-queues
# https://docs.python.org/3.6/library/multiprocessing.html?highlight=process#programming-guidelines
output_queue.close()
output_queue.join_thread()
with num_active_workers.get_lock():
num_active_workers.value -= 1
logging.info(f"Reader worker {worker_id} finished")
class ParallelWorkerPool:
def __init__(self, num_workers: int, worker: Type[Worker], start_method: Optional[str] = None):
self.worker_class = worker
self.num_workers = num_workers
self.input_queue: Optional[Queue] = None
self.output_queue: Optional[Queue] = None
self.ctx: BaseContext = get_context(start_method)
self.processes: List[BaseProcess] = []
self.queue_size = self.num_workers * max_internal_batch_size
self.num_active_workers: Optional[BaseValue] = None
def start(self, **kwargs: Any) -> None:
self.input_queue = self.ctx.Queue(self.queue_size)
self.output_queue = self.ctx.Queue(self.queue_size)
ctx_value = self.ctx.Value("i", self.num_workers)
assert isinstance(ctx_value, BaseValue)
self.num_active_workers = ctx_value
for worker_id in range(0, self.num_workers):
assert hasattr(self.ctx, "Process")
process = self.ctx.Process(
target=_worker,
args=(
self.worker_class,
self.input_queue,
self.output_queue,
self.num_active_workers,
worker_id,
kwargs.copy(),
),
)
process.start()
self.processes.append(process)
def ordered_map(self, stream: Iterable[Any], *args: Any, **kwargs: Any) -> Iterable[Any]:
buffer = defaultdict(Any)
next_expected = 0
for idx, item in self.semi_ordered_map(stream, *args, **kwargs):
buffer[idx] = item
while next_expected in buffer:
yield buffer.pop(next_expected)
next_expected += 1
def semi_ordered_map(self, stream: Iterable[Any], *args: Any, **kwargs: Any) -> Iterable[Tuple[int, Any]]:
try:
self.start(**kwargs)
assert self.input_queue is not None, "Input queue was not initialized"
assert self.output_queue is not None, "Output queue was not initialized"
pushed = 0
read = 0
for idx, item in enumerate(stream):
if pushed - read < self.queue_size:
try:
out_item = self.output_queue.get_nowait()
except Empty:
out_item = None
else:
try:
out_item = self.output_queue.get(timeout=processing_timeout)
except Empty as e:
self.join_or_terminate()
raise e
if out_item is not None:
if out_item == QueueSignals.error:
self.join_or_terminate()
raise RuntimeError("Thread unexpectedly terminated")
yield out_item
read += 1
self.input_queue.put((idx, item))
pushed += 1
for _ in range(self.num_workers):
self.input_queue.put(QueueSignals.stop)
while read < pushed:
out_item = self.output_queue.get(timeout=processing_timeout)
if out_item == QueueSignals.error:
self.join_or_terminate()
raise RuntimeError("Thread unexpectedly terminated")
yield out_item
read += 1
finally:
assert self.input_queue is not None, "Input queue is None"
assert self.output_queue is not None, "Output queue is None"
self.input_queue.close()
self.output_queue.close()
def join_or_terminate(self, timeout: Optional[int] = 1) -> None:
"""
Emergency shutdown
@param timeout:
@return:
"""
for process in self.processes:
process.join(timeout=timeout)
if process.is_alive():
process.terminate()
self.processes.clear()
def join(self) -> None:
for process in self.processes:
process.join()
self.processes.clear()
def __del__(self) -> None:
"""
Terminate processes if the user hasn't joined. This is necessary as
leaving stray processes running can corrupt shared state. In brief,
we've observed shared memory counters being reused (when the memory was
free from the perspective of the parent process) while the stray
workers still held a reference to them.
For a discussion of using destructors in Python in this manner, see
https://eli.thegreenplace.net/2009/06/12/safely-using-destructors-in-python/.
"""
for process in self.processes:
process.terminate()
Generated
+549 -1884
View File
File diff suppressed because it is too large Load Diff
+8 -13
View File
@@ -1,6 +1,6 @@
[tool.poetry]
name = "fastembed"
version = "0.0.4"
version = "0.1.0"
description = "Fast, light, accurate library built for retrieval embedding generation"
authors = ["NirantK <nirant.bits@gmail.com>"]
license = "Apache License"
@@ -12,29 +12,24 @@ keywords = ["vector", "embedding", "neural", "search", "qdrant", "sentence-trans
[tool.poetry.dependencies]
python = ">=3.8.0,<3.12"
onnxruntime = "^1.15.1"
torch = ">=2.0.0, !=2.0.1"
optimum = ">1.12.0"
tqdm = "^4.65.0"
requests = "^2.31.0"
tokenizers = "^0.13.3"
onnx = "^1.11"
onnxruntime = "^1.15"
tqdm = "^4.65"
requests = "^2.31"
tokenizers = "^0.13"
[tool.poetry.dev-dependencies]
[tool.poetry.group.dev.dependencies]
pytest = "^7.4.2"
ruff = "^0.0.277"
isort = "^5.12.0"
black = "^23.7.0"
onnx = "^1.11.0"
notebook = ">=7.0.2"
mkdocs-material = "^9.1.21"
mkdocstrings = "^0.22.0"
pillow = "^10.0.0"
cairosvg = "^2.7.1"
mknotebooks = "^0.8.0"
pytest = "^7.4.0"
[tool.poetry.dependencies.onnxruntime-silicon]
version = "^1.15.0"
markers = "sys_platform == 'darwin'" # This makes it macOS specific
[build-system]
requires = ["poetry-core"]
+55 -6
View File
@@ -1,13 +1,62 @@
import numpy as np
import os
from fastembed.embedding import DefaultEmbedding
import numpy as np
from tqdm import tqdm
from fastembed.embedding import DefaultEmbedding, Embedding
CANONICAL_VECTOR_VALUES = {
"BAAI/bge-small-en": np.array([-0.0232, -0.0255, 0.0174, -0.0639, -0.0006]),
"BAAI/bge-small-en-v1.5": np.array([0.01522374, -0.02271799, 0.00860278, -0.07424029, 0.00386434]),
"BAAI/bge-base-en": np.array([0.0115, 0.0372, 0.0295, 0.0121, 0.0346]),
"BAAI/bge-base-en-v1.5": np.array([0.01129394, 0.05493144, 0.02615099, 0.00328772, 0.02996045]),
"sentence-transformers/all-MiniLM-L6-v2": np.array([0.0259, 0.0058, 0.0114, 0.0380, -0.0233]),
"intfloat/multilingual-e5-large": np.array([0.0098, 0.0045, 0.0066, -0.0354, 0.0070]),
}
def test_default_embedding():
is_ubuntu_ci = os.getenv("IS_UBUNTU_CI")
for model_desc in Embedding.list_supported_models():
if is_ubuntu_ci == "false" and model_desc["size_in_GB"] > 1:
continue
dim = model_desc["dim"]
model = DefaultEmbedding(model_name=model_desc["model"])
docs = ["hello world", "flag embedding"]
embeddings = list(model.embed(docs))
embeddings = np.stack(embeddings, axis=0)
assert embeddings.shape == (2, dim)
canonical_vector = CANONICAL_VECTOR_VALUES[model_desc["model"]]
assert np.allclose(embeddings[0, :canonical_vector.shape[0]], canonical_vector, atol=1e-3), model_desc["model"]
def test_batch_embedding():
model = DefaultEmbedding()
docs = ["hello world", "flag embedding"]
embeddings = np.array(model.embed(docs))
assert embeddings.shape == (2, 384)
docs = ["hello world", "flag embedding"] * 100
embeddings = list(model.embed(docs, batch_size=10))
embeddings = np.stack(embeddings, axis=0)
test_default_embedding()
assert embeddings.shape == (200, 384)
def test_parallel_processing():
model = DefaultEmbedding()
docs = ["hello world", "flag embedding"] * 100
embeddings = list(model.embed(docs, batch_size=10, parallel=2))
embeddings = np.stack(embeddings, axis=0)
embeddings_2 = list(model.embed(docs, batch_size=10, parallel=None))
embeddings_2 = np.stack(embeddings_2, axis=0)
embeddings_3 = list(model.embed(docs, batch_size=10, parallel=0))
embeddings_3 = np.stack(embeddings_3, axis=0)
assert embeddings.shape == (200, 384)
assert np.allclose(embeddings, embeddings_2, atol=1e-3)
assert np.allclose(embeddings, embeddings_3, atol=1e-3)