Compare commits

..
59 changed files with 559 additions and 1337 deletions
+18 -20
View File
@@ -1,6 +1,6 @@
name: Bug
description: File a bug report
title: "[Bug]: "
name: Bug/New Model Request
description: File a bug report/Request a new Model
title: "[Bug/Model Request]: "
body:
- type: markdown
attributes:
@@ -10,22 +10,11 @@ body:
id: what-happened
attributes:
label: What happened?
description: Describe the error you encountered.
placeholder: <Description>
description: Also tell us, what did you expect to happen?
placeholder: Tell us what you see!
value: "A bug happened!"
validations:
required: true
- type: textarea
id: expected
attributes:
label: What is the expected behaviour?
description: Describe the way you expected the code to behave.
placeholder: <Description>
- type: textarea
id: code-snippet
attributes:
label: A minimal reproducible example
description: It would really help us to fix the problem if you could provide a code snippet that reproduces the issue.
placeholder: <Code snippet>
- type: textarea
id: python-version
attributes:
@@ -34,12 +23,21 @@ body:
placeholder: Python3.10
validations:
required: true
- type: textarea
- type: dropdown
id: version
attributes:
label: FastEmbed version
label: Version
description: What version of FastEmbed are you running? python -c "import fastembed; print(fastembed.__version__)". If you're not on the latest, please upgrade and see if the problem persists.
placeholder: v0.4.2
options:
- 0.2.7 (Latest)
- 0.2.6
- 0.2.5
- 0.2.4
- 0.2.3
- 0.2.2
- 0.2.1
- 0.1.x
default: 0
validations:
required: true
- type: dropdown
+1 -1
View File
@@ -1,4 +1,4 @@
blank_issues_enabled: true
blank_issues_enabled: false
contact_links:
- name: GitHub Community Support
url: https://github.com/qdrant/fastembed/discussions
@@ -1,22 +0,0 @@
name: Feature
description: New functionality request
title: "[Feature]: "
body:
- type: markdown
attributes:
value: |
Thanks for taking the time to fill out this report!
- type: textarea
id: feature-description
attributes:
label: What feature would you like to request?
description: Please provide the description of the feature you would like to request.
placeholder: <Description>
validations:
required: true
- type: textarea
id: additional-info
attributes:
label: Is there any additional information you would like to provide?
description: Please provide any additional information that you think might be useful.
placeholder: <Info>
-22
View File
@@ -1,22 +0,0 @@
name: Model
description: Request a new model
title: "[Model]: "
body:
- type: markdown
attributes:
value: |
Thanks for taking the time to fill out this report!
- type: textarea
id: model-name
attributes:
label: Which model would you like to support?
description: Please provide the name of the model you would like to see supported.
placeholder: Link to the model (e.g. on HuggingFace)
validations:
required: true
- type: textarea
id: motivation
attributes:
label: What are the main advantages of this model?
description: Please describe the main advantages of this model comparing to the existing ones and provide links to benchmarks if there are any.
placeholder: <Description>
-19
View File
@@ -1,19 +0,0 @@
### All Submissions:
* [ ] Have you followed the guidelines in our Contributing document?
* [ ] Have you checked to ensure there aren't other open [Pull Requests](../../../pulls) for the same update/change?
<!-- You can erase any parts of this template not applicable to your Pull Request. -->
### New Feature Submissions:
* [ ] Does your submission pass the existing tests?
* [ ] Have you added tests for your feature?
* [ ] Have you installed `pre-commit` with `pip3 install pre-commit` and set up hooks with `pre-commit install`?
### New models submission:
* [ ] Have you added an explanation of why it's important to include this model?
* [ ] Have you added tests for the new model? Were canonical values for tests computed via the original model?
* [ ] Have you added the code snippet for how canonical values were computed?
* [ ] Have you successfully ran tests with your changes locally?
+4 -2
View File
@@ -15,11 +15,11 @@ jobs:
strategy:
matrix:
python-version:
- '3.8.x'
- '3.9.x'
- '3.10.x'
- '3.11.x'
- '3.12.x'
- '3.13.x'
os:
- ubuntu-latest
@@ -33,6 +33,8 @@ jobs:
uses: actions/setup-python@v5
with:
python-version: ${{ matrix.python-version }}
# - name: Setup tmate session
# uses: mxschmitt/action-tmate@v3
- name: Install dependencies
run: |
python -m pip install poetry
@@ -41,4 +43,4 @@ jobs:
- name: Run pytest
run: |
poetry run pytest
poetry run pytest
-2
View File
@@ -5,8 +5,6 @@ This product includes software developed by Qdrant
This distribution includes the following Jina AI models, each with its respective license:
- jinaai/jina-colbert-v2
- License: cc-by-nc-4.0
- jinaai/jina-reranker-v2-base-multilingual
- License: cc-by-nc-4.0
These models are developed by Jina (https://jina.ai/) and are subject to Jina AI's licensing terms.
+2 -16
View File
@@ -28,10 +28,10 @@ pip install fastembed-gpu
```python
from fastembed import TextEmbedding
from typing import List
# Example list of documents
documents: list[str] = [
documents: List[str] = [
"This is built to be faster and lighter than other embedding libraries e.g. Transformers, Sentence-Transformers, etc.",
"fastembed is supported by and maintained by Qdrant.",
]
@@ -137,20 +137,6 @@ embeddings = list(model.embed(images))
# ]
```
### 🔄 Rerankers
```python
from fastembed.rerank.cross_encoder import TextCrossEncoder
query = "Who is maintaining Qdrant?"
documents: list[str] = [
"This is built to be faster and lighter than other embedding libraries e.g. Transformers, Sentence-Transformers, etc.",
"fastembed is supported by and maintained by Qdrant.",
]
encoder = TextCrossEncoder(model_name="Xenova/ms-marco-MiniLM-L-6-v2")
scores = list(encoder.rerank(query, documents))
# [-11.48061752319336, 5.472434997558594]
```
## ⚡️ FastEmbed on a GPU
+3 -1
View File
@@ -65,13 +65,15 @@
}
],
"source": [
"from typing import List\n",
"\n",
"import numpy as np\n",
"\n",
"from fastembed import TextEmbedding\n",
"\n",
"\n",
"# Example list of documents\n",
"documents: list[str] = [\n",
"documents: List[str] = [\n",
" \"This is built to be faster and lighter than other embedding libraries e.g. Transformers, Sentence-Transformers, etc.\",\n",
" \"fastembed is supported by and maintained by Qdrant.\",\n",
"]\n",
+5 -1
View File
@@ -388,6 +388,8 @@
}
],
"source": [
"from typing import List\n",
"\n",
"import numpy as np\n",
"\n",
"from fastembed import TextEmbedding\n",
@@ -405,7 +407,9 @@
"id": "iPtoHf7GeV-i"
},
"outputs": [],
"source": "documents: list[str] = list(np.repeat(\"Demonstrating GPU acceleration in fastembed\", 500))"
"source": [
"documents: List[str] = list(np.repeat(\"Demonstrating GPU acceleration in fastembed\", 500))"
]
},
{
"cell_type": "code",
-88
View File
@@ -1,88 +0,0 @@
{
"cells": [
{
"cell_type": "markdown",
"metadata": {},
"source": [
"### Fastembed Multi-GPU Tutorial\n",
"This tutorial demonstrates how to leverage multi-GPU support in Fastembed. Fastembed supports embedding text and images utilizing modern GPUs for acceleration. Let's explore how to use Fastembed with multiple GPUs step by step."
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"#### Prerequisites\n",
"To get started, ensure you have the following installed:\n",
"- Python 3.9 or later\n",
"- Fastembed (`pip install fastembed-gpu`)\n",
"- Refer to [this](https://github.com/qdrant/fastembed/blob/main/docs/examples/FastEmbed_GPU.ipynb) tutorial if you have issues with GPU dependencies\n",
"- Access to a multi-GPU server"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"### Multi-GPU using cuda argument with TextEmbedding Model"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"from fastembed import TextEmbedding\n",
"\n",
"# define the documents to embed\n",
"docs = [\"hello world\", \"flag embedding\"] * 100\n",
"\n",
"# define gpu ids\n",
"device_ids = [0, 1]\n",
"\n",
"if __name__ == \"__main__\":\n",
" # initialize a TextEmbedding model using CUDA\n",
" text_model = TextEmbedding(\n",
" model_name=\"sentence-transformers/all-MiniLM-L6-v2\",\n",
" cuda=True,\n",
" device_ids=device_ids,\n",
" lazy_load=True,\n",
" )\n",
"\n",
" # generate embeddings\n",
" text_embeddings = list(text_model.embed(docs, batch_size=2, parallel=len(device_ids)))\n",
" print(text_embeddings)"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"In this snippet:\n",
"- `cuda=True` enables GPU acceleration.\n",
"- `device_ids=[0, 1]` specifies GPUs to use. Replace `[0, 1]` with available GPU IDs.\n",
"- `lazy_load=True`\n",
"\n",
"**NOTE**: When using multi-GPU settings, it is important to configure `parallel` and `lazy_load` properly to avoid inefficiencies:\n",
"\n",
"`parallel`: This parameter enables multi-GPU support by spawning child processes for each GPU specified in device_ids. To ensure proper utilization, the value of `parallel` must match the number of GPUs in device_ids. If using a single GPU, this parameter is not necessary.\n",
"\n",
"`lazy_load`: Enabling `lazy_load` prevents redundant memory usage. Without `lazy_load`, the model is initially loaded into the memory of the first GPU by the main process. When child processes are spawned for each GPU, the model is reloaded on the first GPU, causing redundant memory consumption and inefficiencies."
]
}
],
"metadata": {
"kernelspec": {
"display_name": ".venv",
"language": "python",
"name": "python3"
},
"language_info": {
"name": "python",
"version": "3.10.15"
}
},
"nbformat": 4,
"nbformat_minor": 2
}
@@ -36,7 +36,7 @@
"outputs": [],
"source": [
"import time\n",
"from typing import Callable\n",
"from typing import Callable, List, Tuple\n",
"\n",
"import matplotlib.pyplot as plt\n",
"import torch.nn.functional as F\n",
@@ -64,7 +64,6 @@
],
"source": [
"import fastembed\n",
"\n",
"fastembed.__version__"
]
},
@@ -99,7 +98,7 @@
}
],
"source": [
"documents: list[str] = [\n",
"documents: List[str] = [\n",
" \"Chandrayaan-3 is India's third lunar mission\",\n",
" \"It aimed to land a rover on the Moon's surface - joining the US, China and Russia\",\n",
" \"The mission is a follow-up to Chandrayaan-2, which had partial success\",\n",
@@ -156,7 +155,7 @@
" self.model = AutoModel.from_pretrained(model_id)\n",
" self.tokenizer = AutoTokenizer.from_pretrained(model_id)\n",
"\n",
" def embed(self, texts: list[str]):\n",
" def embed(self, texts: List[str]):\n",
" encoded_input = self.tokenizer(\n",
" texts, max_length=512, padding=True, truncation=True, return_tensors=\"pt\"\n",
" )\n",
@@ -255,7 +254,7 @@
"\n",
"def calculate_time_stats(\n",
" embed_func: Callable, documents: list, k: int\n",
") -> tuple[float, float, float]:\n",
") -> Tuple[float, float, float]:\n",
" times = []\n",
" for _ in range(k):\n",
" # Timing the embed_func call\n",
@@ -310,7 +309,7 @@
],
"source": [
"def plot_character_per_second_comparison(\n",
" hf_stats: tuple[float, float, float], fst_stats: tuple[float, float, float], documents: list\n",
" hf_stats: Tuple[float, float, float], fst_stats: Tuple[float, float, float], documents: list\n",
"):\n",
" # Calculating total characters in documents\n",
" total_characters = sum(len(doc) for doc in documents)\n",
@@ -44,7 +44,7 @@
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 21,
"metadata": {
"ExecuteTime": {
"end_time": "2024-03-30T00:45:24.814968Z",
@@ -58,6 +58,8 @@
},
"outputs": [],
"source": [
"from typing import List\n",
"\n",
"import numpy as np\n",
"from datasets import load_dataset\n",
"from peft import AutoPeftModelForCausalLM\n",
@@ -70,11 +72,11 @@
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 23,
"metadata": {},
"outputs": [],
"source": [
"hf_token = \"<YOUR_HF_TOKEN_HERE>\" # Get your token from https://huggingface.co/settings/token, needed for Gemma weights"
"hf_token = <YOUR_HF_TOKEN_HERE> # Get your token from https://huggingface.co/settings/token, needed for Gemma weights"
]
},
{
@@ -244,7 +246,7 @@
},
"outputs": [],
"source": [
"context_embeddings: list[np.ndarray] = list(\n",
"context_embeddings: List[np.ndarray] = list(\n",
" embedding_model.embed(contexts)\n",
") # Note the list() call - this is a generator"
]
+10 -9
View File
@@ -50,6 +50,7 @@
"outputs": [],
"source": [
"import json\n",
"from typing import List, Tuple\n",
"\n",
"import numpy as np\n",
"import pandas as pd\n",
@@ -488,11 +489,11 @@
}
],
"source": [
"def make_sparse_embedding(texts: list[str]):\n",
"def make_sparse_embedding(texts: List[str]):\n",
" return list(sparse_model.embed(texts, batch_size=32))\n",
"\n",
"\n",
"sparse_embedding: list[SparseEmbedding] = make_sparse_embedding(\n",
"sparse_embedding: List[SparseEmbedding] = make_sparse_embedding(\n",
" [\"Fastembed is a great library for text embeddings!\"]\n",
")\n",
"sparse_embedding"
@@ -661,7 +662,7 @@
},
"outputs": [],
"source": [
"def make_dense_embedding(texts: list[str]):\n",
"def make_dense_embedding(texts: List[str]):\n",
" return list(dense_model.embed(texts))\n",
"\n",
"\n",
@@ -871,7 +872,7 @@
},
"outputs": [],
"source": [
"def make_points(df: pd.DataFrame) -> list[PointStruct]:\n",
"def make_points(df: pd.DataFrame) -> List[PointStruct]:\n",
" sparse_vectors = df[\"sparse_embedding\"].tolist()\n",
" product_texts = df[\"combined_text\"].tolist()\n",
" dense_vectors = df[\"dense_embedding\"].tolist()\n",
@@ -898,7 +899,7 @@
" return points\n",
"\n",
"\n",
"points: list[PointStruct] = make_points(df)"
"points: List[PointStruct] = make_points(df)"
]
},
{
@@ -941,8 +942,8 @@
"source": [
"def search(query_text: str):\n",
" # # Compute sparse and dense vectors\n",
" query_sparse_vectors: list[SparseEmbedding] = make_sparse_embedding([query_text])\n",
" query_dense_vector: list[np.ndarray] = make_dense_embedding([query_text])\n",
" query_sparse_vectors: List[SparseEmbedding] = make_sparse_embedding([query_text])\n",
" query_dense_vector: List[np.ndarray] = make_dense_embedding([query_text])\n",
"\n",
" search_results = client.search_batch(\n",
" collection_name=collection_name,\n",
@@ -1074,7 +1075,7 @@
"metadata": {},
"outputs": [],
"source": [
"def rank_list(search_result: list[ScoredPoint]):\n",
"def rank_list(search_result: List[ScoredPoint]):\n",
" return [(point.id, rank + 1) for rank, point in enumerate(search_result)]\n",
"\n",
"\n",
@@ -1148,7 +1149,7 @@
],
"source": [
"def find_point_by_id(\n",
" client: QdrantClient, collection_name: str, rrf_rank_list: list[tuple[int, float]]\n",
" client: QdrantClient, collection_name: str, rrf_rank_list: List[Tuple[int, float]]\n",
"):\n",
" return client.retrieve(\n",
" collection_name=collection_name, ids=[item[0] for item in rrf_rank_list]\n",
+9 -14
View File
@@ -47,7 +47,7 @@
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 2,
"metadata": {
"ExecuteTime": {
"end_time": "2024-03-30T00:49:20.516644Z",
@@ -56,7 +56,8 @@
},
"outputs": [],
"source": [
"from fastembed import SparseTextEmbedding, SparseEmbedding"
"from fastembed import SparseTextEmbedding, SparseEmbedding\n",
"from typing import List"
]
},
{
@@ -133,7 +134,7 @@
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 5,
"metadata": {
"ExecuteTime": {
"end_time": "2024-03-30T00:49:28.624109Z",
@@ -142,7 +143,7 @@
},
"outputs": [],
"source": [
"documents: list[str] = [\n",
"documents: List[str] = [\n",
" \"Chandrayaan-3 is India's third lunar mission\",\n",
" \"It aimed to land a rover on the Moon's surface - joining the US, China and Russia\",\n",
" \"The mission is a follow-up to Chandrayaan-2, which had partial success\",\n",
@@ -156,7 +157,7 @@
" \"Chandrayaan-3 was launched from the Satish Dhawan Space Centre in Sriharikota\",\n",
" \"Chandrayaan-3 was launched earlier in the year 2023\",\n",
"]\n",
"sparse_embeddings_list: list[SparseEmbedding] = list(\n",
"sparse_embeddings_list: List[SparseEmbedding] = list(\n",
" model.embed(documents, batch_size=6)\n",
") # batch_size is optional, notice the generator"
]
@@ -234,9 +235,7 @@
"source": [
"# Let's print the first 5 features and their weights for better understanding.\n",
"for i in range(5):\n",
" print(\n",
" f\"Token at index {sparse_embeddings_list[0].indices[i]} has weight {sparse_embeddings_list[0].values[i]}\"\n",
" )"
" print(f\"Token at index {sparse_embeddings_list[0].indices[i]} has weight {sparse_embeddings_list[0].values[i]}\")"
]
},
{
@@ -262,9 +261,7 @@
"import json\n",
"from transformers import AutoTokenizer\n",
"\n",
"tokenizer = AutoTokenizer.from_pretrained(\n",
" SparseTextEmbedding.list_supported_models()[0][\"sources\"][\"hf\"]\n",
")"
"tokenizer = AutoTokenizer.from_pretrained(SparseTextEmbedding.list_supported_models()[0][\"sources\"][\"hf\"])"
]
},
{
@@ -329,9 +326,7 @@
" token_weight_dict[token] = weight\n",
"\n",
" # Sort the dictionary by weights\n",
" token_weight_dict = dict(\n",
" sorted(token_weight_dict.items(), key=lambda item: item[1], reverse=True)\n",
" )\n",
" token_weight_dict = dict(sorted(token_weight_dict.items(), key=lambda item: item[1], reverse=True))\n",
" return token_weight_dict\n",
"\n",
"\n",
+174 -310
View File
@@ -2,36 +2,38 @@
"cells": [
{
"cell_type": "code",
"execution_count": 1,
"metadata": {
"ExecuteTime": {
"end_time": "2024-11-13T09:01:03.324551Z",
"start_time": "2024-11-13T09:01:03.234711Z"
"end_time": "2024-05-31T18:13:23.806907Z",
"start_time": "2024-05-31T18:13:23.797078Z"
}
},
"outputs": [],
"source": [
"%load_ext autoreload\n",
"%autoreload 2"
],
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"The autoreload extension is already loaded. To reload it, use:\n",
" %reload_ext autoreload\n"
]
}
],
"execution_count": 10
]
},
{
"cell_type": "code",
"execution_count": 2,
"metadata": {
"ExecuteTime": {
"end_time": "2024-11-13T09:01:04.505772Z",
"start_time": "2024-11-13T09:01:04.493296Z"
"end_time": "2024-05-31T18:14:31.147674Z",
"start_time": "2024-05-31T18:14:31.134015Z"
}
},
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
"/home/hossam/.pyenv/versions/.venv/lib/python3.10/site-packages/tqdm/auto.py:21: TqdmWarning: IProgress not found. Please update jupyter and ipywidgets. See https://ipywidgets.readthedocs.io/en/stable/user_install.html\n",
" from .autonotebook import tqdm as notebook_tqdm\n"
]
}
],
"source": [
"import pandas as pd\n",
"\n",
@@ -40,11 +42,8 @@
" TextEmbedding,\n",
" LateInteractionTextEmbedding,\n",
" ImageEmbedding,\n",
")\n",
"from fastembed.rerank.cross_encoder import TextCrossEncoder"
],
"outputs": [],
"execution_count": 11
")"
]
},
{
"cell_type": "markdown",
@@ -55,79 +54,16 @@
},
{
"cell_type": "code",
"execution_count": 3,
"metadata": {
"ExecuteTime": {
"end_time": "2024-11-13T09:01:05.812271Z",
"start_time": "2024-11-13T09:01:05.795846Z"
"end_time": "2024-05-31T18:13:25.863008Z",
"start_time": "2024-05-31T18:13:25.837795Z"
}
},
"source": [
"supported_models = (\n",
" pd.DataFrame(TextEmbedding.list_supported_models())\n",
" .sort_values(\"size_in_GB\")\n",
" .drop(columns=[\"sources\", \"model_file\", \"additional_files\"])\n",
" .reset_index(drop=True)\n",
")\n",
"supported_models"
],
"outputs": [
{
"data": {
"text/plain": [
" model dim \\\n",
"0 BAAI/bge-small-en-v1.5 384 \n",
"1 BAAI/bge-small-zh-v1.5 512 \n",
"2 snowflake/snowflake-arctic-embed-xs 384 \n",
"3 sentence-transformers/all-MiniLM-L6-v2 384 \n",
"4 jinaai/jina-embeddings-v2-small-en 512 \n",
"5 BAAI/bge-small-en 384 \n",
"6 snowflake/snowflake-arctic-embed-s 384 \n",
"7 nomic-ai/nomic-embed-text-v1.5-Q 768 \n",
"8 BAAI/bge-base-en-v1.5 768 \n",
"9 sentence-transformers/paraphrase-multilingual-... 384 \n",
"10 Qdrant/clip-ViT-B-32-text 512 \n",
"11 jinaai/jina-embeddings-v2-base-de 768 \n",
"12 BAAI/bge-base-en 768 \n",
"13 snowflake/snowflake-arctic-embed-m 768 \n",
"14 nomic-ai/nomic-embed-text-v1.5 768 \n",
"15 jinaai/jina-embeddings-v2-base-en 768 \n",
"16 nomic-ai/nomic-embed-text-v1 768 \n",
"17 snowflake/snowflake-arctic-embed-m-long 768 \n",
"18 mixedbread-ai/mxbai-embed-large-v1 1024 \n",
"19 jinaai/jina-embeddings-v2-base-code 768 \n",
"20 sentence-transformers/paraphrase-multilingual-... 768 \n",
"21 snowflake/snowflake-arctic-embed-l 1024 \n",
"22 thenlper/gte-large 1024 \n",
"23 BAAI/bge-large-en-v1.5 1024 \n",
"24 intfloat/multilingual-e5-large 1024 \n",
"\n",
" description license size_in_GB \n",
"0 Text embeddings, Unimodal (text), English, 512... mit 0.067 \n",
"1 Text embeddings, Unimodal (text), Chinese, 512... mit 0.090 \n",
"2 Text embeddings, Unimodal (text), English, 512... apache-2.0 0.090 \n",
"3 Text embeddings, Unimodal (text), English, 256... apache-2.0 0.090 \n",
"4 Text embeddings, Unimodal (text), English, 819... apache-2.0 0.120 \n",
"5 Text embeddings, Unimodal (text), English, 512... mit 0.130 \n",
"6 Text embeddings, Unimodal (text), English, 512... apache-2.0 0.130 \n",
"7 Text embeddings, Multimodal (text, image), Eng... apache-2.0 0.130 \n",
"8 Text embeddings, Unimodal (text), English, 512... mit 0.210 \n",
"9 Text embeddings, Unimodal (text), Multilingual... apache-2.0 0.220 \n",
"10 Text embeddings, Multimodal (text&image), Engl... mit 0.250 \n",
"11 Text embeddings, Unimodal (text), Multilingual... apache-2.0 0.320 \n",
"12 Text embeddings, Unimodal (text), English, 512... mit 0.420 \n",
"13 Text embeddings, Unimodal (text), English, 512... apache-2.0 0.430 \n",
"14 Text embeddings, Multimodal (text, image), Eng... apache-2.0 0.520 \n",
"15 Text embeddings, Unimodal (text), English, 819... apache-2.0 0.520 \n",
"16 Text embeddings, Multimodal (text, image), Eng... apache-2.0 0.520 \n",
"17 Text embeddings, Unimodal (text), English, 204... apache-2.0 0.540 \n",
"18 Text embeddings, Unimodal (text), English, 512... apache-2.0 0.640 \n",
"19 Text embeddings, Unimodal (text), Multilingual... apache-2.0 0.640 \n",
"20 Text embeddings, Unimodal (text), Multilingual... apache-2.0 1.000 \n",
"21 Text embeddings, Unimodal (text), English, 512... apache-2.0 1.020 \n",
"22 Text embeddings, Unimodal (text), English, 512... mit 1.200 \n",
"23 Text embeddings, Unimodal (text), English, 512... mit 1.200 \n",
"24 Text embeddings, Unimodal (text), Multilingual... mit 2.240 "
],
"text/html": [
"<div>\n",
"<style scoped>\n",
@@ -358,14 +294,77 @@
" </tbody>\n",
"</table>\n",
"</div>"
],
"text/plain": [
" model dim \\\n",
"0 BAAI/bge-small-en-v1.5 384 \n",
"1 BAAI/bge-small-zh-v1.5 512 \n",
"2 snowflake/snowflake-arctic-embed-xs 384 \n",
"3 sentence-transformers/all-MiniLM-L6-v2 384 \n",
"4 jinaai/jina-embeddings-v2-small-en 512 \n",
"5 BAAI/bge-small-en 384 \n",
"6 snowflake/snowflake-arctic-embed-s 384 \n",
"7 nomic-ai/nomic-embed-text-v1.5-Q 768 \n",
"8 BAAI/bge-base-en-v1.5 768 \n",
"9 sentence-transformers/paraphrase-multilingual-... 384 \n",
"10 Qdrant/clip-ViT-B-32-text 512 \n",
"11 jinaai/jina-embeddings-v2-base-de 768 \n",
"12 BAAI/bge-base-en 768 \n",
"13 snowflake/snowflake-arctic-embed-m 768 \n",
"14 nomic-ai/nomic-embed-text-v1.5 768 \n",
"15 jinaai/jina-embeddings-v2-base-en 768 \n",
"16 nomic-ai/nomic-embed-text-v1 768 \n",
"17 snowflake/snowflake-arctic-embed-m-long 768 \n",
"18 mixedbread-ai/mxbai-embed-large-v1 1024 \n",
"19 jinaai/jina-embeddings-v2-base-code 768 \n",
"20 sentence-transformers/paraphrase-multilingual-... 768 \n",
"21 snowflake/snowflake-arctic-embed-l 1024 \n",
"22 thenlper/gte-large 1024 \n",
"23 BAAI/bge-large-en-v1.5 1024 \n",
"24 intfloat/multilingual-e5-large 1024 \n",
"\n",
" description license size_in_GB \n",
"0 Text embeddings, Unimodal (text), English, 512... mit 0.067 \n",
"1 Text embeddings, Unimodal (text), Chinese, 512... mit 0.090 \n",
"2 Text embeddings, Unimodal (text), English, 512... apache-2.0 0.090 \n",
"3 Text embeddings, Unimodal (text), English, 256... apache-2.0 0.090 \n",
"4 Text embeddings, Unimodal (text), English, 819... apache-2.0 0.120 \n",
"5 Text embeddings, Unimodal (text), English, 512... mit 0.130 \n",
"6 Text embeddings, Unimodal (text), English, 512... apache-2.0 0.130 \n",
"7 Text embeddings, Multimodal (text, image), Eng... apache-2.0 0.130 \n",
"8 Text embeddings, Unimodal (text), English, 512... mit 0.210 \n",
"9 Text embeddings, Unimodal (text), Multilingual... apache-2.0 0.220 \n",
"10 Text embeddings, Multimodal (text&image), Engl... mit 0.250 \n",
"11 Text embeddings, Unimodal (text), Multilingual... apache-2.0 0.320 \n",
"12 Text embeddings, Unimodal (text), English, 512... mit 0.420 \n",
"13 Text embeddings, Unimodal (text), English, 512... apache-2.0 0.430 \n",
"14 Text embeddings, Multimodal (text, image), Eng... apache-2.0 0.520 \n",
"15 Text embeddings, Unimodal (text), English, 819... apache-2.0 0.520 \n",
"16 Text embeddings, Multimodal (text, image), Eng... apache-2.0 0.520 \n",
"17 Text embeddings, Unimodal (text), English, 204... apache-2.0 0.540 \n",
"18 Text embeddings, Unimodal (text), English, 512... apache-2.0 0.640 \n",
"19 Text embeddings, Unimodal (text), Multilingual... apache-2.0 0.640 \n",
"20 Text embeddings, Unimodal (text), Multilingual... apache-2.0 1.000 \n",
"21 Text embeddings, Unimodal (text), English, 512... apache-2.0 1.020 \n",
"22 Text embeddings, Unimodal (text), English, 512... mit 1.200 \n",
"23 Text embeddings, Unimodal (text), English, 512... mit 1.200 \n",
"24 Text embeddings, Unimodal (text), Multilingual... mit 2.240 "
]
},
"execution_count": 12,
"execution_count": 3,
"metadata": {},
"output_type": "execute_result"
}
],
"execution_count": 12
"source": [
"supported_models = (\n",
" pd.DataFrame(TextEmbedding.list_supported_models())\n",
" .sort_values(\"size_in_GB\")\n",
" .drop(columns=[\"sources\", \"model_file\", \"additional_files\"])\n",
" .reset_index(drop=True)\n",
")\n",
"supported_models"
]
},
{
"cell_type": "markdown",
@@ -376,42 +375,16 @@
},
{
"cell_type": "code",
"execution_count": 4,
"metadata": {
"ExecuteTime": {
"end_time": "2024-11-13T09:01:07.038954Z",
"start_time": "2024-11-13T09:01:07.019656Z"
"end_time": "2024-05-31T18:13:27.124747Z",
"start_time": "2024-05-31T18:13:27.096212Z"
}
},
"source": [
"(\n",
" pd.DataFrame(SparseTextEmbedding.list_supported_models())\n",
" .sort_values(\"size_in_GB\")\n",
" .drop(columns=[\"sources\", \"model_file\", \"additional_files\"])\n",
" .reset_index(drop=True)\n",
")"
],
"outputs": [
{
"data": {
"text/plain": [
" model vocab_size \\\n",
"0 Qdrant/bm25 NaN \n",
"1 Qdrant/bm42-all-minilm-l6-v2-attentions 30522.0 \n",
"2 prithivida/Splade_PP_en_v1 30522.0 \n",
"3 prithvida/Splade_PP_en_v1 30522.0 \n",
"\n",
" description license size_in_GB \\\n",
"0 BM25 as sparse embeddings meant to be used wit... apache-2.0 0.010 \n",
"1 Light sparse embedding model, which assigns an... apache-2.0 0.090 \n",
"2 Independent Implementation of SPLADE++ Model f... apache-2.0 0.532 \n",
"3 Independent Implementation of SPLADE++ Model f... apache-2.0 0.532 \n",
"\n",
" requires_idf \n",
"0 True \n",
"1 True \n",
"2 NaN \n",
"3 NaN "
],
"text/html": [
"<div>\n",
"<style scoped>\n",
@@ -479,14 +452,40 @@
" </tbody>\n",
"</table>\n",
"</div>"
],
"text/plain": [
" model vocab_size \\\n",
"0 Qdrant/bm25 NaN \n",
"1 Qdrant/bm42-all-minilm-l6-v2-attentions 30522.0 \n",
"2 prithivida/Splade_PP_en_v1 30522.0 \n",
"3 prithvida/Splade_PP_en_v1 30522.0 \n",
"\n",
" description license size_in_GB \\\n",
"0 BM25 as sparse embeddings meant to be used wit... apache-2.0 0.010 \n",
"1 Light sparse embedding model, which assigns an... apache-2.0 0.090 \n",
"2 Independent Implementation of SPLADE++ Model f... apache-2.0 0.532 \n",
"3 Independent Implementation of SPLADE++ Model f... apache-2.0 0.532 \n",
"\n",
" requires_idf \n",
"0 True \n",
"1 True \n",
"2 NaN \n",
"3 NaN "
]
},
"execution_count": 13,
"execution_count": 4,
"metadata": {},
"output_type": "execute_result"
}
],
"execution_count": 13
"source": [
"(\n",
" pd.DataFrame(SparseTextEmbedding.list_supported_models())\n",
" .sort_values(\"size_in_GB\")\n",
" .drop(columns=[\"sources\", \"model_file\", \"additional_files\"])\n",
" .reset_index(drop=True)\n",
")"
]
},
{
"cell_type": "markdown",
@@ -499,40 +498,17 @@
},
{
"cell_type": "code",
"execution_count": 5,
"metadata": {
"collapsed": false,
"ExecuteTime": {
"end_time": "2024-11-13T09:01:08.074442Z",
"start_time": "2024-11-13T09:01:08.056138Z"
}
"end_time": "2024-05-31T18:14:34.370252Z",
"start_time": "2024-05-31T18:14:34.354270Z"
},
"collapsed": false
},
"source": [
"(\n",
" pd.DataFrame(LateInteractionTextEmbedding.list_supported_models())\n",
" .sort_values(\"size_in_GB\")\n",
" .drop(columns=[\"sources\", \"model_file\"])\n",
" .reset_index(drop=True)\n",
")"
],
"outputs": [
{
"data": {
"text/plain": [
" model dim \\\n",
"0 answerdotai/answerai-colbert-small-v1 96 \n",
"1 colbert-ir/colbertv2.0 128 \n",
"2 jinaai/jina-colbert-v2 128 \n",
"\n",
" description license \\\n",
"0 Text embeddings, Unimodal (text), Multilingual... apache-2.0 \n",
"1 Late interaction model mit \n",
"2 New model that expands capabilities of colbert... cc-by-nc-4.0 \n",
"\n",
" size_in_GB additional_files \n",
"0 0.13 NaN \n",
"1 0.44 NaN \n",
"2 2.24 [onnx/model.onnx_data] "
],
"text/html": [
"<div>\n",
"<style scoped>\n",
@@ -582,7 +558,7 @@
" <tr>\n",
" <th>2</th>\n",
" <td>jinaai/jina-colbert-v2</td>\n",
" <td>128</td>\n",
" <td>1024</td>\n",
" <td>New model that expands capabilities of colbert...</td>\n",
" <td>cc-by-nc-4.0</td>\n",
" <td>2.24</td>\n",
@@ -591,14 +567,37 @@
" </tbody>\n",
"</table>\n",
"</div>"
],
"text/plain": [
" model dim \\\n",
"0 answerdotai/answerai-colbert-small-v1 96 \n",
"1 colbert-ir/colbertv2.0 128 \n",
"2 jinaai/jina-colbert-v2 1024 \n",
"\n",
" description license \\\n",
"0 Text embeddings, Unimodal (text), Multilingual... apache-2.0 \n",
"1 Late interaction model mit \n",
"2 New model that expands capabilities of colbert... cc-by-nc-4.0 \n",
"\n",
" size_in_GB additional_files \n",
"0 0.13 NaN \n",
"1 0.44 NaN \n",
"2 2.24 [onnx/model.onnx_data] "
]
},
"execution_count": 14,
"execution_count": 5,
"metadata": {},
"output_type": "execute_result"
}
],
"execution_count": 14
"source": [
"(\n",
" pd.DataFrame(LateInteractionTextEmbedding.list_supported_models())\n",
" .sort_values(\"size_in_GB\")\n",
" .drop(columns=[\"sources\", \"model_file\"])\n",
" .reset_index(drop=True)\n",
")"
]
},
{
"cell_type": "markdown",
@@ -611,37 +610,17 @@
},
{
"cell_type": "code",
"execution_count": 6,
"metadata": {
"collapsed": false,
"ExecuteTime": {
"end_time": "2024-11-13T09:01:09.171647Z",
"start_time": "2024-11-13T09:01:09.150940Z"
}
"end_time": "2024-05-31T18:14:42.501881Z",
"start_time": "2024-05-31T18:14:42.484726Z"
},
"collapsed": false
},
"source": [
"(\n",
" pd.DataFrame(ImageEmbedding.list_supported_models())\n",
" .sort_values(\"size_in_GB\")\n",
" .drop(columns=[\"sources\", \"model_file\"])\n",
" .reset_index(drop=True)\n",
")"
],
"outputs": [
{
"data": {
"text/plain": [
" model dim \\\n",
"0 Qdrant/resnet50-onnx 2048 \n",
"1 Qdrant/clip-ViT-B-32-vision 512 \n",
"2 Qdrant/Unicom-ViT-B-32 512 \n",
"3 Qdrant/Unicom-ViT-B-16 768 \n",
"\n",
" description license size_in_GB \n",
"0 Image embeddings, Unimodal (image), 2016 year apache-2.0 0.10 \n",
"1 Image embeddings, Multimodal (text&image), 202... mit 0.34 \n",
"2 Image embeddings, Multimodal (text&image), 202... apache-2.0 0.48 \n",
"3 Image embeddings (more detailed than Unicom-Vi... apache-2.0 0.82 "
],
"text/html": [
"<div>\n",
"<style scoped>\n",
@@ -704,149 +683,39 @@
" </tbody>\n",
"</table>\n",
"</div>"
],
"text/plain": [
" model dim \\\n",
"0 Qdrant/resnet50-onnx 2048 \n",
"1 Qdrant/clip-ViT-B-32-vision 512 \n",
"2 Qdrant/Unicom-ViT-B-32 512 \n",
"3 Qdrant/Unicom-ViT-B-16 768 \n",
"\n",
" description license size_in_GB \n",
"0 Image embeddings, Unimodal (image), 2016 year apache-2.0 0.10 \n",
"1 Image embeddings, Multimodal (text&image), 202... mit 0.34 \n",
"2 Image embeddings, Multimodal (text&image), 202... apache-2.0 0.48 \n",
"3 Image embeddings (more detailed than Unicom-Vi... apache-2.0 0.82 "
]
},
"execution_count": 15,
"execution_count": 6,
"metadata": {},
"output_type": "execute_result"
}
],
"execution_count": 15
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Supported Rerank Cross Encoder Models"
]
},
{
"cell_type": "code",
"metadata": {
"ExecuteTime": {
"end_time": "2024-11-13T09:01:10.313943Z",
"start_time": "2024-11-13T09:01:10.298428Z"
}
},
"source": [
"(\n",
" pd.DataFrame(TextCrossEncoder.list_supported_models())\n",
" pd.DataFrame(ImageEmbedding.list_supported_models())\n",
" .sort_values(\"size_in_GB\")\n",
" .drop(columns=[\"sources\", \"model_file\"])\n",
" .reset_index(drop=True)\n",
")"
],
"outputs": [
{
"data": {
"text/plain": [
" model size_in_GB \\\n",
"0 Xenova/ms-marco-MiniLM-L-6-v2 0.08 \n",
"1 Xenova/ms-marco-MiniLM-L-12-v2 0.12 \n",
"2 jinaai/jina-reranker-v1-tiny-en 0.13 \n",
"3 jinaai/jina-reranker-v1-turbo-en 0.15 \n",
"4 BAAI/bge-reranker-base 1.04 \n",
"5 jinaai/jina-reranker-v2-base-multilingual 1.11 \n",
"\n",
" description license \n",
"0 MiniLM-L-6-v2 model optimized for re-ranking t... apache-2.0 \n",
"1 MiniLM-L-12-v2 model optimized for re-ranking ... apache-2.0 \n",
"2 Designed for blazing-fast re-ranking with 8K c... apache-2.0 \n",
"3 Designed for blazing-fast re-ranking with 8K c... apache-2.0 \n",
"4 BGE reranker base model for cross-encoder re-r... mit \n",
"5 A multi-lingual reranker model for cross-encod... cc-by-nc-4.0 "
],
"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>size_in_GB</th>\n",
" <th>description</th>\n",
" <th>license</th>\n",
" </tr>\n",
" </thead>\n",
" <tbody>\n",
" <tr>\n",
" <th>0</th>\n",
" <td>Xenova/ms-marco-MiniLM-L-6-v2</td>\n",
" <td>0.08</td>\n",
" <td>MiniLM-L-6-v2 model optimized for re-ranking t...</td>\n",
" <td>apache-2.0</td>\n",
" </tr>\n",
" <tr>\n",
" <th>1</th>\n",
" <td>Xenova/ms-marco-MiniLM-L-12-v2</td>\n",
" <td>0.12</td>\n",
" <td>MiniLM-L-12-v2 model optimized for re-ranking ...</td>\n",
" <td>apache-2.0</td>\n",
" </tr>\n",
" <tr>\n",
" <th>2</th>\n",
" <td>jinaai/jina-reranker-v1-tiny-en</td>\n",
" <td>0.13</td>\n",
" <td>Designed for blazing-fast re-ranking with 8K c...</td>\n",
" <td>apache-2.0</td>\n",
" </tr>\n",
" <tr>\n",
" <th>3</th>\n",
" <td>jinaai/jina-reranker-v1-turbo-en</td>\n",
" <td>0.15</td>\n",
" <td>Designed for blazing-fast re-ranking with 8K c...</td>\n",
" <td>apache-2.0</td>\n",
" </tr>\n",
" <tr>\n",
" <th>4</th>\n",
" <td>BAAI/bge-reranker-base</td>\n",
" <td>1.04</td>\n",
" <td>BGE reranker base model for cross-encoder re-r...</td>\n",
" <td>mit</td>\n",
" </tr>\n",
" <tr>\n",
" <th>5</th>\n",
" <td>jinaai/jina-reranker-v2-base-multilingual</td>\n",
" <td>1.11</td>\n",
" <td>A multi-lingual reranker model for cross-encod...</td>\n",
" <td>cc-by-nc-4.0</td>\n",
" </tr>\n",
" </tbody>\n",
"</table>\n",
"</div>"
]
},
"execution_count": 16,
"metadata": {},
"output_type": "execute_result"
}
],
"execution_count": 16
},
{
"metadata": {},
"cell_type": "code",
"outputs": [],
"execution_count": null,
"source": ""
]
}
],
"metadata": {
"kernelspec": {
"display_name": "Python 3.8.18 ('base')",
"display_name": ".venv",
"language": "python",
"name": "python3"
},
@@ -860,14 +729,9 @@
"name": "python",
"nbconvert_exporter": "python",
"pygments_lexer": "ipython3",
"version": "3.11.8"
"version": "3.10.15"
},
"orig_nbformat": 4,
"vscode": {
"interpreter": {
"hash": "c4a27af61e455bc18dcf16f5867a2ff0402fa12b01dd0f6ce3a79ae73ad15e91"
}
}
"orig_nbformat": 4
},
"nbformat": 4,
"nbformat_minor": 2
+2 -2
View File
@@ -26,14 +26,14 @@ pip install fastembed
```python
from fastembed import TextEmbedding
documents: list[str] = [
documents: List[str] = [
"passage: Hello, World!",
"query: Hello, World!",
"passage: This is an example passage.",
"fastembed is supported by and maintained by Qdrant."
]
embedding_model = TextEmbedding()
embeddings: list[np.ndarray] = embedding_model.embed(documents)
embeddings: List[np.ndarray] = embedding_model.embed(documents)
```
## Usage with Qdrant
+3 -2
View File
@@ -41,6 +41,7 @@
"metadata": {},
"outputs": [],
"source": [
"from typing import List\n",
"import numpy as np\n",
"from fastembed import TextEmbedding"
]
@@ -70,7 +71,7 @@
],
"source": [
"# Example list of documents\n",
"documents: list[str] = [\n",
"documents: List[str] = [\n",
" \"Maharana Pratap was a Rajput warrior king from Mewar\",\n",
" \"He fought against the Mughal Empire led by Akbar\",\n",
" \"The Battle of Haldighati in 1576 was his most famous battle\",\n",
@@ -86,7 +87,7 @@
"embedding_model = TextEmbedding(model_name=\"BAAI/bge-small-en\")\n",
"\n",
"# We'll use the passage_embed method to get the embeddings for the documents\n",
"embeddings: list[np.ndarray] = list(\n",
"embeddings: List[np.ndarray] = list(\n",
" embedding_model.passage_embed(documents)\n",
") # notice that we are casting the generator to a list\n",
"\n",
+3 -4
View File
@@ -46,6 +46,7 @@
"metadata": {},
"outputs": [],
"source": [
"from typing import List\n",
"from qdrant_client import QdrantClient"
]
},
@@ -66,7 +67,7 @@
"outputs": [],
"source": [
"# Example list of documents\n",
"documents: list[str] = [\n",
"documents: List[str] = [\n",
" \"Maharana Pratap was a Rajput warrior king from Mewar\",\n",
" \"He fought against the Mughal Empire led by Akbar\",\n",
" \"The Battle of Haldighati in 1576 was his most famous battle\",\n",
@@ -198,9 +199,7 @@
}
],
"source": [
"search_result = client.query(\n",
" collection_name=\"demo_collection\", query_text=\"This is a query document\"\n",
")\n",
"search_result = client.query(collection_name=\"demo_collection\", query_text=\"This is a query document\")\n",
"print(search_result)"
]
},
+5 -11
View File
@@ -19,7 +19,7 @@
"outputs": [],
"source": [
"from pathlib import Path\n",
"from typing import Any\n",
"from typing import List, Tuple, Any\n",
"\n",
"import numpy as np\n",
"import time\n",
@@ -91,11 +91,9 @@
" return last_hidden.sum(dim=1) / attention_mask.sum(dim=1)[..., None]\n",
"\n",
"\n",
"def hf_embed(model_id: str, inputs: list[str]):\n",
"def hf_embed(model_id: str, inputs: List[str]):\n",
" # Tokenize the input texts\n",
" batch_dict = hf_tokenizer(\n",
" inputs, max_length=512, padding=True, truncation=True, return_tensors=\"pt\"\n",
" )\n",
" batch_dict = hf_tokenizer(inputs, max_length=512, padding=True, truncation=True, return_tensors=\"pt\")\n",
"\n",
" outputs = hf_model(**batch_dict)\n",
" embeddings = average_pool(outputs.last_hidden_state, batch_dict[\"attention_mask\"])\n",
@@ -135,9 +133,7 @@
"optimization_config = AutoOptimizationConfig.O4()\n",
"optimizer = ORTOptimizer.from_pretrained(model)\n",
"\n",
"optimizer.optimize(\n",
" save_dir=save_dir, optimization_config=optimization_config, use_external_data_format=True\n",
")\n",
"optimizer.optimize(save_dir=save_dir, optimization_config=optimization_config, use_external_data_format=True)\n",
"model = ORTModelForFeatureExtraction.from_pretrained(save_dir)\n",
"\n",
"tokenizer.save_pretrained(save_dir)\n",
@@ -175,9 +171,7 @@
"metadata": {},
"outputs": [],
"source": [
"def measure_pipeline_time(\n",
" pipeline, input_texts: list[str], num_runs=10, **kwargs: Any\n",
") -> tuple[float, float]:\n",
"def measure_pipeline_time(pipeline, input_texts: List[str], num_runs=10, **kwargs: Any) -> Tuple[float, float]:\n",
" \"\"\"Measures the time it takes to run the pipeline on the input texts.\"\"\"\n",
" times = []\n",
" total_chars = sum(len(text) for text in input_texts)\n",
File diff suppressed because one or more lines are too long
+11 -23
View File
@@ -3,31 +3,27 @@ import time
import shutil
import tarfile
from pathlib import Path
from typing import Any, Optional
from typing import Any, Dict, List, Optional
import requests
from huggingface_hub import snapshot_download
from huggingface_hub.utils import (
RepositoryNotFoundError,
disable_progress_bars,
enable_progress_bars,
)
from huggingface_hub.utils import RepositoryNotFoundError
from loguru import logger
from tqdm import tqdm
class ModelManagement:
@classmethod
def list_supported_models(cls) -> list[dict[str, Any]]:
def list_supported_models(cls) -> List[Dict[str, Any]]:
"""Lists the supported models.
Returns:
list[dict[str, Any]]: A list of dictionaries containing the model information.
List[Dict[str, Any]]: A list of dictionaries containing the model information.
"""
raise NotImplementedError()
@classmethod
def _get_model_description(cls, model_name: str) -> dict[str, Any]:
def _get_model_description(cls, model_name: str) -> Dict[str, Any]:
"""
Gets the model description from the model_name.
@@ -38,7 +34,7 @@ class ModelManagement:
ValueError: If the model_name is not supported.
Returns:
dict[str, Any]: The model description.
Dict[str, Any]: The model description.
"""
for model in cls.list_supported_models():
if model_name.lower() == model["model"].lower():
@@ -97,8 +93,8 @@ class ModelManagement:
def download_files_from_huggingface(
cls,
hf_source_repo: str,
cache_dir: str,
extra_patterns: Optional[list[str]] = None,
cache_dir: Optional[str] = None,
extra_patterns: Optional[List[str]] = None,
local_files_only: bool = False,
**kwargs,
) -> str:
@@ -107,7 +103,7 @@ class ModelManagement:
Args:
hf_source_repo (str): Name of the model on HuggingFace Hub, e.g. "qdrant/all-MiniLM-L6-v2-onnx".
cache_dir (Optional[str]): The path to the cache directory.
extra_patterns (Optional[list[str]]): extra patterns to allow in the snapshot download, typically
extra_patterns (Optional[List[str]]): extra patterns to allow in the snapshot download, typically
includes the required model files.
local_files_only (bool, optional): Whether to only use local files. Defaults to False.
Returns:
@@ -123,12 +119,6 @@ class ModelManagement:
if extra_patterns is not None:
allow_patterns.extend(extra_patterns)
snapshot_dir = Path(cache_dir) / f"models--{hf_source_repo.replace('/', '--')}"
is_cached = snapshot_dir.exists()
if is_cached:
disable_progress_bars()
return snapshot_download(
repo_id=hf_source_repo,
allow_patterns=allow_patterns,
@@ -221,13 +211,13 @@ class ModelManagement:
@classmethod
def download_model(
cls, model: dict[str, Any], cache_dir: Path, retries: int = 3, **kwargs
cls, model: Dict[str, Any], cache_dir: Path, retries: int = 3, **kwargs
) -> Path:
"""
Downloads a model from HuggingFace Hub or Google Cloud Storage.
Args:
model (dict[str, Any]): The model description.
model (Dict[str, Any]): The model description.
Example:
```
{
@@ -275,8 +265,6 @@ class ModelManagement:
f"Could not download model from HuggingFace: {e} "
"Falling back to other sources."
)
finally:
enable_progress_bars()
if url_source or local_files_only:
try:
return cls.retrieve_model_gcs(
+14 -11
View File
@@ -1,7 +1,17 @@
import warnings
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Generic, Iterable, Optional, Sequence, Type, TypeVar
from typing import (
Any,
Dict,
Generic,
Iterable,
Optional,
Sequence,
Tuple,
Type,
TypeVar,
)
import numpy as np
import onnxruntime as ort
@@ -33,8 +43,8 @@ class OnnxModel(Generic[T]):
self.tokenizer = None
def _preprocess_onnx_input(
self, onnx_input: dict[str, np.ndarray], **kwargs
) -> dict[str, np.ndarray]:
self, onnx_input: Dict[str, np.ndarray], **kwargs
) -> Dict[str, np.ndarray]:
"""
Preprocess the onnx input.
"""
@@ -52,13 +62,6 @@ class OnnxModel(Generic[T]):
model_path = model_dir / model_file
# List of Execution Providers: https://onnxruntime.ai/docs/execution-providers
if cuda and providers is not None:
warnings.warn(
f"`cuda` and `providers` are mutually exclusive parameters, cuda: {cuda}, providers: {providers}",
category=UserWarning,
stacklevel=6,
)
if providers is not None:
onnx_providers = list(providers)
elif cuda:
@@ -128,5 +131,5 @@ class EmbeddingWorker(Worker):
def start(cls, model_name: str, cache_dir: str, **kwargs: Any) -> "EmbeddingWorker":
return cls(model_name=model_name, cache_dir=cache_dir, **kwargs)
def process(self, items: Iterable[tuple[int, Any]]) -> Iterable[tuple[int, Any]]:
def process(self, items: Iterable[Tuple[int, Any]]) -> Iterable[Tuple[int, Any]]:
raise NotImplementedError("Subclasses must implement this method")
+2 -2
View File
@@ -1,6 +1,6 @@
import json
from pathlib import Path
from typing import Tuple
from tokenizers import AddedToken, Tokenizer
from fastembed.image.transform.operators import Compose
@@ -17,7 +17,7 @@ def load_special_tokens(model_dir: Path) -> dict:
return tokens_map
def load_tokenizer(model_dir: Path) -> tuple[Tokenizer, dict]:
def load_tokenizer(model_dir: Path) -> Tuple[Tokenizer, dict]:
config_path = model_dir / "config.json"
if not config_path.exists():
raise ValueError(f"Could not find config.json in {model_dir}")
+2 -2
View File
@@ -1,7 +1,7 @@
import os
import sys
from PIL import Image
from typing import Any, Iterable, Union
from typing import Any, Dict, Iterable, Tuple, Union
if sys.version_info >= (3, 10):
from typing import TypeAlias
@@ -13,4 +13,4 @@ PathInput: TypeAlias = Union[str, os.PathLike]
PilInput: TypeAlias = Union[Image.Image, Iterable[Image.Image]]
ImageInput: TypeAlias = Union[PathInput, Iterable[PathInput], PilInput]
OnnxProvider: TypeAlias = Union[str, tuple[str, dict[Any, Any]]]
OnnxProvider: TypeAlias = Union[str, Tuple[str, Dict[Any, Any]]]
+6 -6
View File
@@ -1,13 +1,13 @@
import os
import sys
import re
import tempfile
import unicodedata
from pathlib import Path
from itertools import islice
from pathlib import Path
from typing import Generator, Iterable, Optional, Union
import unicodedata
import sys
import numpy as np
import re
from typing import Set
def normalize(input_array, p=2, dim=1, eps=1e-12) -> np.ndarray:
@@ -45,7 +45,7 @@ def define_cache_dir(cache_dir: Optional[str] = None) -> Path:
return cache_path
def get_all_punctuation() -> set[str]:
def get_all_punctuation() -> Set[str]:
return set(
chr(i) for i in range(sys.maxunicode) if unicodedata.category(chr(i)).startswith("P")
)
+5 -5
View File
@@ -1,4 +1,4 @@
from typing import Any, Iterable, Optional, Sequence, Type
from typing import Any, Dict, Iterable, List, Optional, Sequence, Type
import numpy as np
@@ -8,15 +8,15 @@ from fastembed.image.onnx_embedding import OnnxImageEmbedding
class ImageEmbedding(ImageEmbeddingBase):
EMBEDDINGS_REGISTRY: list[Type[ImageEmbeddingBase]] = [OnnxImageEmbedding]
EMBEDDINGS_REGISTRY: List[Type[ImageEmbeddingBase]] = [OnnxImageEmbedding]
@classmethod
def list_supported_models(cls) -> list[dict[str, Any]]:
def list_supported_models(cls) -> List[Dict[str, Any]]:
"""
Lists the supported models.
Returns:
list[dict[str, Any]]: A list of dictionaries containing the model information.
List[Dict[str, Any]]: A list of dictionaries containing the model information.
Example:
```
@@ -47,7 +47,7 @@ class ImageEmbedding(ImageEmbeddingBase):
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
device_ids: Optional[List[int]] = None,
lazy_load: bool = False,
**kwargs,
):
+7 -18
View File
@@ -1,4 +1,4 @@
from typing import Any, Iterable, Optional, Sequence, Type
from typing import Any, Dict, Iterable, List, Optional, Sequence, Type
import numpy as np
@@ -53,17 +53,6 @@ supported_onnx_models = [
},
"model_file": "model.onnx",
},
{
"model": "jinaai/jina-clip-v1",
"dim": 768,
"description": "Image embeddings, Multimodal (text&image), 2024 year",
"license": "apache-2.0",
"size_in_GB": 0.34,
"sources": {
"hf": "jinaai/jina-clip-v1",
},
"model_file": "onnx/vision_model.onnx",
},
]
@@ -75,7 +64,7 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[np.ndarray]):
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
device_ids: Optional[List[int]] = None,
lazy_load: bool = False,
device_id: Optional[int] = None,
**kwargs,
@@ -91,7 +80,7 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[np.ndarray]):
Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to False.
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
device_ids (Optional[List[int]], optional): The list of device ids to use for data parallel processing in
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
@@ -140,12 +129,12 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[np.ndarray]):
)
@classmethod
def list_supported_models(cls) -> list[dict[str, Any]]:
def list_supported_models(cls) -> List[Dict[str, Any]]:
"""
Lists the supported models.
Returns:
list[Dict[str, Any]]: A list of dictionaries containing the model information.
List[Dict[str, Any]]: A list of dictionaries containing the model information.
"""
return supported_onnx_models
@@ -189,8 +178,8 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[np.ndarray]):
return OnnxImageEmbeddingWorker
def _preprocess_onnx_input(
self, onnx_input: dict[str, np.ndarray], **kwargs
) -> dict[str, np.ndarray]:
self, onnx_input: Dict[str, np.ndarray], **kwargs
) -> Dict[str, np.ndarray]:
"""
Preprocess the onnx input.
"""
+7 -7
View File
@@ -2,7 +2,7 @@ import contextlib
import os
from multiprocessing import get_all_start_methods
from pathlib import Path
from typing import Any, Iterable, Optional, Sequence, Type
from typing import Any, Dict, Iterable, List, Optional, Sequence, Tuple, Type
import numpy as np
from PIL import Image
@@ -29,8 +29,8 @@ class OnnxImageModel(OnnxModel[T]):
self.processor = None
def _preprocess_onnx_input(
self, onnx_input: dict[str, np.ndarray], **kwargs
) -> dict[str, np.ndarray]:
self, onnx_input: Dict[str, np.ndarray], **kwargs
) -> Dict[str, np.ndarray]:
"""
Preprocess the onnx input.
"""
@@ -58,10 +58,10 @@ class OnnxImageModel(OnnxModel[T]):
def load_onnx_model(self) -> None:
raise NotImplementedError("Subclasses must implement this method")
def _build_onnx_input(self, encoded: np.ndarray) -> dict[str, np.ndarray]:
def _build_onnx_input(self, encoded: np.ndarray) -> Dict[str, np.ndarray]:
return {node.name: encoded for node in self.model.get_inputs()}
def onnx_embed(self, images: list[ImageInput], **kwargs) -> OnnxOutputContext:
def onnx_embed(self, images: List[ImageInput], **kwargs) -> OnnxOutputContext:
with contextlib.ExitStack():
image_files = [
Image.open(image) if not isinstance(image, Image.Image) else image
@@ -83,7 +83,7 @@ class OnnxImageModel(OnnxModel[T]):
parallel: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
device_ids: Optional[List[int]] = None,
**kwargs,
) -> Iterable[T]:
is_small = False
@@ -125,7 +125,7 @@ class OnnxImageModel(OnnxModel[T]):
class ImageEmbeddingWorker(EmbeddingWorker):
def process(self, items: Iterable[tuple[int, Any]]) -> Iterable[tuple[int, Any]]:
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
+8 -34
View File
@@ -1,4 +1,4 @@
from typing import Sized, Union
from typing import Sized, Tuple, Union
import numpy as np
from PIL import Image
@@ -14,7 +14,7 @@ def convert_to_rgb(image: Image.Image) -> Image.Image:
def center_crop(
image: Union[Image.Image, np.ndarray],
size: tuple[int, int],
size: Tuple[int, int],
) -> np.ndarray:
if isinstance(image, np.ndarray):
_, orig_height, orig_width = image.shape
@@ -62,8 +62,8 @@ def center_crop(
def normalize(
image: np.ndarray,
mean: Union[float, np.ndarray],
std: Union[float, np.ndarray],
mean=Union[float, np.ndarray],
std=Union[float, np.ndarray],
) -> np.ndarray:
if not isinstance(image, np.ndarray):
raise ValueError("image must be a numpy array")
@@ -96,10 +96,10 @@ def normalize(
def resize(
image: Image.Image,
size: Union[int, tuple[int, int]],
resample: Union[int, Image.Resampling] = Image.Resampling.BILINEAR,
) -> Image.Image:
image: Image,
size: Union[int, Tuple[int, int]],
resample: Image.Resampling = Image.Resampling.BILINEAR,
) -> Image:
if isinstance(size, tuple):
return image.resize(size, resample)
@@ -122,29 +122,3 @@ def pil2ndarray(image: Union[Image.Image, np.ndarray]):
if isinstance(image, Image.Image):
return np.asarray(image).transpose((2, 0, 1))
return image
def pad2square(
image: Image.Image,
size: int,
fill_color: Union[str, int, tuple[int, ...]] = 0,
) -> Image.Image:
height, width = image.height, image.width
left, right = 0, width
top, bottom = 0, height
crop_required = False
if width > size:
left = (width - size) // 2
right = left + size
crop_required = True
if height > size:
top = (height - size) // 2
bottom = top + size
crop_required = True
new_image = Image.new(mode="RGB", size=(size, size), color=fill_color)
new_image.paste(image.crop((left, top, right, bottom)) if crop_required else image)
return new_image
+29 -99
View File
@@ -1,4 +1,4 @@
from typing import Any, Union, Optional
from typing import Any, Dict, List, Tuple, Union
import numpy as np
from PIL import Image
@@ -10,111 +10,93 @@ from fastembed.image.transform.functional import (
pil2ndarray,
rescale,
resize,
pad2square,
)
class Transform:
def __call__(self, images: list) -> Union[list[Image.Image], list[np.ndarray]]:
def __call__(self, images: List) -> Union[List[Image.Image], List[np.ndarray]]:
raise NotImplementedError("Subclasses must implement this method")
class ConvertToRGB(Transform):
def __call__(self, images: list[Image.Image]) -> list[Image.Image]:
def __call__(self, images: List[Image.Image]) -> List[Image.Image]:
return [convert_to_rgb(image=image) for image in images]
class CenterCrop(Transform):
def __init__(self, size: tuple[int, int]):
def __init__(self, size: Tuple[int, int]):
self.size = size
def __call__(self, images: list[Image.Image]) -> list[np.ndarray]:
def __call__(self, images: List[Image.Image]) -> List[np.ndarray]:
return [center_crop(image=image, size=self.size) for image in images]
class Normalize(Transform):
def __init__(self, mean: Union[float, list[float]], std: Union[float, list[float]]):
def __init__(self, mean: Union[float, List[float]], std: Union[float, List[float]]):
self.mean = mean
self.std = std
def __call__(self, images: list[np.ndarray]) -> list[np.ndarray]:
def __call__(self, images: List[np.ndarray]) -> List[np.ndarray]:
return [normalize(image, mean=self.mean, std=self.std) for image in images]
class Resize(Transform):
def __init__(
self,
size: Union[int, tuple[int, int]],
size: Union[int, Tuple[int, int]],
resample: Image.Resampling = Image.Resampling.BICUBIC,
):
self.size = size
self.resample = resample
def __call__(self, images: list[Image.Image]) -> list[Image.Image]:
return [resize(image, size=self.size, resample=self.resample) for image in images]
def __call__(self, images: List[Image.Image]) -> List[Image.Image]:
return [
resize(image, size=self.size, resample=self.resample) for image in images
]
class Rescale(Transform):
def __init__(self, scale: float = 1 / 255):
self.scale = scale
def __call__(self, images: list[np.ndarray]) -> list[np.ndarray]:
def __call__(self, images: List[np.ndarray]) -> List[np.ndarray]:
return [rescale(image, scale=self.scale) for image in images]
class PILtoNDarray(Transform):
def __call__(self, images: list[Union[Image.Image, np.ndarray]]) -> list[np.ndarray]:
def __call__(
self, images: List[Union[Image.Image, np.ndarray]]
) -> List[np.ndarray]:
return [pil2ndarray(image) for image in images]
class PadtoSquare(Transform):
def __init__(
self,
size: int,
fill_color: Optional[Union[str, int, tuple[int, ...]]] = None,
):
self.size = size
self.fill_color = fill_color
def __call__(self, images: list[Image.Image]) -> list[Image.Image]:
return [
pad2square(image=image, size=self.size, fill_color=self.fill_color) for image in images
]
class Compose:
def __init__(self, transforms: list[Transform]):
def __init__(self, transforms: List[Transform]):
self.transforms = transforms
def __call__(
self, images: Union[list[Image.Image], list[np.ndarray]]
) -> Union[list[np.ndarray], list[Image.Image]]:
self, images: Union[List[Image.Image], List[np.ndarray]]
) -> Union[List[np.ndarray], List[Image.Image]]:
for transform in self.transforms:
images = transform(images)
return images
@classmethod
def from_config(cls, config: dict[str, Any]) -> "Compose":
def from_config(cls, config: Dict[str, Any]) -> "Compose":
"""Creates processor from a config dict.
Args:
config (dict[str, Any]): Configuration dictionary.
config (Dict[str, Any]): Configuration dictionary.
Valid keys:
- do_resize
- resize_mode
- size
- fill_color
- do_center_crop
- crop_size
- do_rescale
- rescale_factor
- do_normalize
- image_mean
- mean
- image_std
- std
- resample
- interpolation
Valid size keys (nested):
- {"height", "width"}
- {"shortest_edge"}
@@ -125,7 +107,6 @@ class Compose:
transforms = []
cls._get_convert_to_rgb(transforms, config)
cls._get_resize(transforms, config)
cls._get_pad2square(transforms, config)
cls._get_center_crop(transforms, config)
cls._get_pil2ndarray(transforms, config)
cls._get_rescale(transforms, config)
@@ -133,11 +114,11 @@ class Compose:
return cls(transforms=transforms)
@staticmethod
def _get_convert_to_rgb(transforms: list[Transform], config: dict[str, Any]):
def _get_convert_to_rgb(transforms: List[Transform], config: Dict[str, Any]):
transforms.append(ConvertToRGB())
@classmethod
def _get_resize(cls, transforms: list[Transform], config: dict[str, Any]):
@staticmethod
def _get_resize(transforms: List[Transform], config: Dict[str, Any]):
mode = config.get("image_processor_type", "CLIPImageProcessor")
if mode == "CLIPImageProcessor":
if config.get("do_resize", False):
@@ -180,27 +161,9 @@ class Compose:
resample=config.get("resample", Image.Resampling.BICUBIC),
)
)
elif mode == "JinaCLIPImageProcessor":
interpolation = config.get("interpolation")
if isinstance(interpolation, str):
resample = cls._interpolation_resolver(interpolation)
else:
resample = interpolation or Image.Resampling.BICUBIC
if "size" in config:
resize_mode = config.get("resize_mode", "shortest")
if resize_mode == "shortest":
transforms.append(
Resize(
size=config["size"],
resample=resample,
)
)
else:
raise ValueError(f"Preprocessor {mode} is not supported")
@staticmethod
def _get_center_crop(transforms: list[Transform], config: dict[str, Any]):
def _get_center_crop(transforms: List[Transform], config: Dict[str, Any]):
mode = config.get("image_processor_type", "CLIPImageProcessor")
if mode == "CLIPImageProcessor":
if config.get("do_center_crop", False):
@@ -214,55 +177,22 @@ class Compose:
transforms.append(CenterCrop(size=crop_size))
elif mode == "ConvNextFeatureExtractor":
pass
elif mode == "JinaCLIPImageProcessor":
pass
else:
raise ValueError(f"Preprocessor {mode} is not supported")
@staticmethod
def _get_pil2ndarray(transforms: list[Transform], config: dict[str, Any]):
def _get_pil2ndarray(transforms: List[Transform], config: Dict[str, Any]):
transforms.append(PILtoNDarray())
@staticmethod
def _get_rescale(transforms: list[Transform], config: dict[str, Any]):
def _get_rescale(transforms: List[Transform], config: Dict[str, Any]):
if config.get("do_rescale", True):
rescale_factor = config.get("rescale_factor", 1 / 255)
transforms.append(Rescale(scale=rescale_factor))
@staticmethod
def _get_normalize(transforms: list[Transform], config: dict[str, Any]):
def _get_normalize(transforms: List[Transform], config: Dict[str, Any]):
if config.get("do_normalize", False):
transforms.append(Normalize(mean=config["image_mean"], std=config["image_std"]))
elif "mean" in config and "std" in config:
transforms.append(Normalize(mean=config["mean"], std=config["std"]))
@staticmethod
def _get_pad2square(transforms: list[Transform], config: dict[str, Any]):
mode = config.get("image_processor_type", "CLIPImageProcessor")
if mode == "CLIPImageProcessor":
pass
elif mode == "ConvNextFeatureExtractor":
pass
elif mode == "JinaCLIPImageProcessor":
transforms.append(
PadtoSquare(
size=config["size"],
fill_color=config.get("fill_color", 0),
)
Normalize(mean=config["image_mean"], std=config["image_std"])
)
@staticmethod
def _interpolation_resolver(resample: Optional[str] = None) -> Image.Resampling:
interpolation_map = {
"nearest": Image.Resampling.NEAREST,
"lanczos": Image.Resampling.LANCZOS,
"bilinear": Image.Resampling.BILINEAR,
"bicubic": Image.Resampling.BICUBIC,
"box": Image.Resampling.BOX,
"hamming": Image.Resampling.HAMMING,
}
if resample and (method := interpolation_map.get(resample.lower())):
return method
raise ValueError(f"Unknown interpolation method: {resample}")
+22 -19
View File
@@ -1,5 +1,5 @@
import string
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Dict, Iterable, List, Optional, Sequence, Type, Union
import numpy as np
from tokenizers import Encoding
@@ -12,7 +12,6 @@ from fastembed.late_interaction.late_interaction_embedding_base import (
)
from fastembed.text.onnx_text_model import OnnxTextModel, TextEmbeddingWorker
supported_colbert_models = [
{
"model": "colbert-ir/colbertv2.0",
@@ -42,7 +41,7 @@ supported_colbert_models = [
class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[np.ndarray]):
QUERY_MARKER_TOKEN_ID = 1
DOCUMENT_MARKER_TOKEN_ID = 2
MIN_QUERY_LENGTH = 31 # it's 32, we add one additional special token in the beginning
MIN_QUERY_LENGTH = 32
MASK_TOKEN = "[MASK]"
def _post_process_onnx_output(
@@ -68,21 +67,25 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[np.ndarray]):
return output.model_output.astype(np.float32)
def _preprocess_onnx_input(
self, onnx_input: dict[str, np.ndarray], is_doc: bool = True, **kwargs: Any
) -> dict[str, np.ndarray]:
marker_token = self.DOCUMENT_MARKER_TOKEN_ID if is_doc else self.QUERY_MARKER_TOKEN_ID
onnx_input["input_ids"] = np.insert(onnx_input["input_ids"], 1, marker_token, axis=1)
onnx_input["attention_mask"] = np.insert(onnx_input["attention_mask"], 1, 1, axis=1)
self, onnx_input: Dict[str, np.ndarray], is_doc: bool = True
) -> Dict[str, np.ndarray]:
if is_doc:
onnx_input["input_ids"][:, 1] = self.DOCUMENT_MARKER_TOKEN_ID
else:
onnx_input["input_ids"][:, 1] = self.QUERY_MARKER_TOKEN_ID
return onnx_input
def tokenize(self, documents: list[str], is_doc: bool = True, **kwargs: Any) -> list[Encoding]:
def tokenize(self, documents: List[str], is_doc: bool = True) -> List[Encoding]:
return (
self._tokenize_documents(documents=documents)
if is_doc
else self._tokenize_query(query=next(iter(documents)))
)
def _tokenize_query(self, query: str) -> list[Encoding]:
def _tokenize_query(self, query: str) -> List[Encoding]:
# "@ " is added to a query to be replaced with a special query token
# make sure that "@ " is considered as a single token
query = f"@ {query}"
encoded = self.tokenizer.encode_batch([query])
# colbert authors recommend to pad queries with [MASK] tokens for query augmentation to improve performance
if len(encoded[0].ids) < self.MIN_QUERY_LENGTH:
@@ -101,16 +104,19 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[np.ndarray]):
self.tokenizer.enable_padding(**prev_padding)
return encoded
def _tokenize_documents(self, documents: list[str]) -> list[Encoding]:
def _tokenize_documents(self, documents: List[str]) -> List[Encoding]:
# "@ " is added to a document to be replaced with a special document token
# make sure that "@ " is considered as a single token
documents = ["@ " + doc for doc in documents]
encoded = self.tokenizer.encode_batch(documents)
return encoded
@classmethod
def list_supported_models(cls) -> list[dict[str, Any]]:
def list_supported_models(cls) -> List[Dict[str, Any]]:
"""Lists the supported models.
Returns:
list[dict[str, Any]]: A list of dictionaries containing the model information.
List[Dict[str, Any]]: A list of dictionaries containing the model information.
"""
return supported_colbert_models
@@ -121,7 +127,7 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[np.ndarray]):
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
device_ids: Optional[List[int]] = None,
lazy_load: bool = False,
device_id: Optional[int] = None,
**kwargs,
@@ -137,7 +143,7 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[np.ndarray]):
Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to False.
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
device_ids (Optional[List[int]], optional): The list of device ids to use for data parallel processing in
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
@@ -191,9 +197,6 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[np.ndarray]):
self.tokenizer.encode(symbol, add_special_tokens=False).ids[0]
for symbol in string.punctuation
}
current_max_length = self.tokenizer.truncation["max_length"]
# ensure not to overflow after adding document-marker
self.tokenizer.enable_truncation(max_length=current_max_length - 1)
def embed(
self,
@@ -229,7 +232,7 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[np.ndarray]):
**kwargs,
)
def query_embed(self, query: Union[str, Iterable[str]], **kwargs) -> Iterable[np.ndarray]:
def query_embed(self, query: Union[str, List[str]], **kwargs) -> Iterable[np.ndarray]:
if isinstance(query, str):
query = [query]
+11 -11
View File
@@ -1,11 +1,10 @@
from typing import Any, Type
from typing import Any, Dict, List, Type
import numpy as np
from fastembed.late_interaction.colbert import Colbert
from fastembed.text.onnx_text_model import TextEmbeddingWorker
supported_jina_colbert_models = [
{
"model": "jinaai/jina-colbert-v2",
@@ -25,7 +24,7 @@ supported_jina_colbert_models = [
class JinaColbert(Colbert):
QUERY_MARKER_TOKEN_ID = 250002
DOCUMENT_MARKER_TOKEN_ID = 250003
MIN_QUERY_LENGTH = 31 # it's 32, we add one additional special token in the beginning
MIN_QUERY_LENGTH = 32
MASK_TOKEN = "<mask>"
@classmethod
@@ -33,21 +32,22 @@ class JinaColbert(Colbert):
return JinaColbertEmbeddingWorker
@classmethod
def list_supported_models(cls) -> list[dict[str, Any]]:
def list_supported_models(cls) -> List[Dict[str, Any]]:
"""Lists the supported models.
Returns:
list[dict[str, Any]]: A list of dictionaries containing the model information.
List[Dict[str, Any]]: A list of dictionaries containing the model information.
"""
return supported_jina_colbert_models
def _preprocess_onnx_input(
self, onnx_input: dict[str, np.ndarray], is_doc: bool = True, **kwargs: Any
) -> dict[str, np.ndarray]:
onnx_input = super()._preprocess_onnx_input(onnx_input, is_doc)
# the attention mask for jina-colbert-v2 is always 1 in queries
if not is_doc:
self, onnx_input: Dict[str, np.ndarray], is_doc: bool = True
) -> Dict[str, np.ndarray]:
if is_doc:
onnx_input["input_ids"][:, 1] = self.DOCUMENT_MARKER_TOKEN_ID
else:
onnx_input["input_ids"][:, 1] = self.QUERY_MARKER_TOKEN_ID
# the attention mask for jina-colbert-v2 is always 1 in queries
onnx_input["attention_mask"][:] = 1
return onnx_input
@@ -1,4 +1,4 @@
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Dict, Iterable, List, Optional, Sequence, Type, Union
import numpy as np
@@ -11,15 +11,15 @@ from fastembed.late_interaction.late_interaction_embedding_base import (
class LateInteractionTextEmbedding(LateInteractionTextEmbeddingBase):
EMBEDDINGS_REGISTRY: list[Type[LateInteractionTextEmbeddingBase]] = [Colbert, JinaColbert]
EMBEDDINGS_REGISTRY: List[Type[LateInteractionTextEmbeddingBase]] = [Colbert, JinaColbert]
@classmethod
def list_supported_models(cls) -> list[dict[str, Any]]:
def list_supported_models(cls) -> List[Dict[str, Any]]:
"""
Lists the supported models.
Returns:
list[dict[str, Any]]: A list of dictionaries containing the model information.
List[Dict[str, Any]]: A list of dictionaries containing the model information.
Example:
```
@@ -50,7 +50,7 @@ class LateInteractionTextEmbedding(LateInteractionTextEmbeddingBase):
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
device_ids: Optional[List[int]] = None,
lazy_load: bool = False,
**kwargs,
):
+8 -9
View File
@@ -1,15 +1,14 @@
import logging
import os
from collections import defaultdict
from copy import deepcopy
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, Iterable, Optional, Type
from typing import Any, Dict, Iterable, List, Optional, Tuple, Type
from copy import deepcopy
# Single item should be processed in less than:
processing_timeout = 10 * 60 # seconds
@@ -25,10 +24,10 @@ class QueueSignals(str, Enum):
class Worker:
@classmethod
def start(cls, *args: Any, **kwargs: Any) -> "Worker":
def start(cls, **kwargs: Any) -> "Worker":
raise NotImplementedError()
def process(self, items: Iterable[tuple[int, Any]]) -> Iterable[tuple[int, Any]]:
def process(self, items: Iterable[Tuple[int, Any]]) -> Iterable[Tuple[int, Any]]:
raise NotImplementedError()
@@ -38,7 +37,7 @@ def _worker(
output_queue: Queue,
num_active_workers: BaseValue,
worker_id: int,
kwargs: Optional[dict[str, Any]] = None,
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.
@@ -94,7 +93,7 @@ class ParallelWorkerPool:
num_workers: int,
worker: Type[Worker],
start_method: Optional[str] = None,
device_ids: Optional[list[int]] = None,
device_ids: Optional[List[int]] = None,
cuda: bool = False,
):
self.worker_class = worker
@@ -102,7 +101,7 @@ class ParallelWorkerPool:
self.input_queue: Optional[Queue] = None
self.output_queue: Optional[Queue] = None
self.ctx: BaseContext = get_context(start_method)
self.processes: list[BaseProcess] = []
self.processes: List[BaseProcess] = []
self.queue_size = self.num_workers * max_internal_batch_size
self.emergency_shutdown = False
self.device_ids = device_ids
@@ -151,7 +150,7 @@ class ParallelWorkerPool:
def semi_ordered_map(
self, stream: Iterable[Any], *args: Any, **kwargs: Any
) -> Iterable[tuple[int, Any]]:
) -> Iterable[Tuple[int, Any]]:
try:
self.start(**kwargs)
@@ -1,15 +1,11 @@
from typing import Any, Iterable, Optional, Sequence, Type
from typing import List, Iterable, Dict, Any, Sequence, Optional
from loguru import logger
from fastembed.common import OnnxProvider
from fastembed.common.onnx_model import OnnxOutputContext
from fastembed.common.utils import define_cache_dir
from fastembed.rerank.cross_encoder.onnx_text_model import (
OnnxCrossEncoderModel,
TextRerankerWorker,
)
from fastembed.rerank.cross_encoder.onnx_text_model import OnnxCrossEncoderModel
from fastembed.rerank.cross_encoder.text_cross_encoder_base import TextCrossEncoderBase
from fastembed.common.utils import define_cache_dir
supported_onnx_models = [
{
@@ -42,46 +38,16 @@ supported_onnx_models = [
"description": "BGE reranker base model for cross-encoder re-ranking.",
"license": "mit",
},
{
"model": "jinaai/jina-reranker-v1-tiny-en",
"size_in_GB": 0.13,
"sources": {
"hf": "jinaai/jina-reranker-v1-tiny-en",
},
"model_file": "onnx/model.onnx",
"description": "Designed for blazing-fast re-ranking with 8K context length and fewer parameters than jina-reranker-v1-turbo-en.",
"license": "apache-2.0",
},
{
"model": "jinaai/jina-reranker-v1-turbo-en",
"size_in_GB": 0.15,
"sources": {
"hf": "jinaai/jina-reranker-v1-turbo-en",
},
"model_file": "onnx/model.onnx",
"description": "Designed for blazing-fast re-ranking with 8K context length.",
"license": "apache-2.0",
},
{
"model": "jinaai/jina-reranker-v2-base-multilingual",
"size_in_GB": 1.11,
"sources": {
"hf": "jinaai/jina-reranker-v2-base-multilingual",
},
"model_file": "onnx/model.onnx",
"description": "A multi-lingual reranker model for cross-encoder re-ranking with 1K context length and sliding window",
"license": "cc-by-nc-4.0",
},
]
class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
@classmethod
def list_supported_models(cls) -> list[dict[str, Any]]:
def list_supported_models(cls) -> List[Dict[str, Any]]:
"""Lists the supported models.
Returns:
list[dict[str, Any]]: A list of dictionaries containing the model information.
List[Dict[str, Any]]: A list of dictionaries containing the model information.
"""
return supported_onnx_models
@@ -92,10 +58,10 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
device_ids: Optional[List[int]] = None,
lazy_load: bool = False,
device_id: Optional[int] = None,
**kwargs: Any,
**kwargs,
):
"""
Args:
@@ -108,7 +74,7 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to False.
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
device_ids (Optional[List[int]], optional): The list of device ids to use for data parallel processing in
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
@@ -142,9 +108,7 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
self.model_description = self._get_model_description(model_name)
self.cache_dir = define_cache_dir(cache_dir)
self._model_dir = self.download_model(
self.model_description,
self.cache_dir,
local_files_only=self._local_files_only,
self.model_description, self.cache_dir, local_files_only=self._local_files_only
)
if not self.lazy_load:
@@ -165,7 +129,7 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
query: str,
documents: Iterable[str],
batch_size: int = 64,
**kwargs: Any,
**kwargs,
) -> Iterable[float]:
"""Reranks documents based on their relevance to a given query.
@@ -181,44 +145,3 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
yield from self._rerank_documents(
query=query, documents=documents, batch_size=batch_size, **kwargs
)
def rerank_pairs(
self,
pairs: Iterable[tuple[str, str]],
batch_size: int = 64,
parallel: Optional[int] = None,
**kwargs: Any,
) -> Iterable[float]:
yield from self._rerank_pairs(
model_name=self.model_name,
cache_dir=str(self.cache_dir),
pairs=pairs,
batch_size=batch_size,
parallel=parallel,
providers=self.providers,
cuda=self.cuda,
device_ids=self.device_ids,
**kwargs,
)
@classmethod
def _get_worker_class(cls) -> Type[TextRerankerWorker]:
return TextCrossEncoderWorker
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[float]:
return (float(elem) for elem in output.model_output)
class TextCrossEncoderWorker(TextRerankerWorker):
def init_embedding(
self,
model_name: str,
cache_dir: str,
**kwargs,
) -> OnnxTextCrossEncoder:
return OnnxTextCrossEncoder(
model_name=model_name,
cache_dir=cache_dir,
threads=1,
**kwargs,
)
@@ -1,28 +1,16 @@
import os
from multiprocessing import get_all_start_methods
from typing import Sequence, Optional, List, Dict, Iterable
from pathlib import Path
from typing import Any, Iterable, Optional, Sequence, Type
import numpy as np
from tokenizers import Encoding
from fastembed.common.onnx_model import (
EmbeddingWorker,
OnnxModel,
OnnxOutputContext,
OnnxProvider,
)
from fastembed.common.onnx_model import OnnxModel, OnnxProvider
from fastembed.common.preprocessor_utils import load_tokenizer
from fastembed.common.utils import iter_batch
from fastembed.parallel_processor import ParallelWorkerPool
class OnnxCrossEncoderModel(OnnxModel[float]):
ONNX_OUTPUT_NAMES: Optional[list[str]] = None
@classmethod
def _get_worker_class(cls) -> Type["TextRerankerWorker"]:
raise NotImplementedError("Subclasses must implement this method")
class OnnxCrossEncoderModel(OnnxModel):
ONNX_OUTPUT_NAMES: Optional[List[str]] = None
def _load_onnx_model(
self,
@@ -43,108 +31,40 @@ class OnnxCrossEncoderModel(OnnxModel[float]):
)
self.tokenizer, _ = load_tokenizer(model_dir=model_dir)
def tokenize(self, pairs: list[tuple[str, str]], **_: Any) -> list[Encoding]:
return self.tokenizer.encode_batch(pairs)
def tokenize(self, query: str, documents: List[str], **kwargs) -> List[Encoding]:
return self.tokenizer.encode_batch([(query, doc) for doc in documents])
def onnx_embed(self, query: str, documents: List[str], **kwargs) -> List[float]:
tokenized_input = self.tokenize(query, documents, **kwargs)
def _build_onnx_input(self, tokenized_input):
input_names = {node.name for node in self.model.get_inputs()}
inputs = {
"input_ids": np.array([enc.ids for enc in tokenized_input], dtype=np.int64),
"attention_mask": np.array(
[enc.attention_mask for enc in tokenized_input], dtype=np.int64
),
}
input_names = {node.name for node in self.model.get_inputs()}
if "token_type_ids" in input_names:
inputs["token_type_ids"] = np.array(
[enc.type_ids for enc in tokenized_input], dtype=np.int64
)
if "attention_mask" in input_names:
inputs["attention_mask"] = np.array(
[enc.attention_mask for enc in tokenized_input], dtype=np.int64
)
return inputs
def onnx_embed(self, query: str, documents: list[str], **kwargs: Any) -> OnnxOutputContext:
pairs = [(query, doc) for doc in documents]
return self.onnx_embed_pairs(pairs, **kwargs)
def onnx_embed_pairs(self, pairs: list[tuple[str, str]], **kwargs: Any) -> OnnxOutputContext:
tokenized_input = self.tokenize(pairs, **kwargs)
inputs = self._build_onnx_input(tokenized_input)
onnx_input = self._preprocess_onnx_input(inputs, **kwargs)
outputs = self.model.run(self.ONNX_OUTPUT_NAMES, onnx_input)
relevant_output = outputs[0]
scores = relevant_output[:, 0]
return OnnxOutputContext(model_output=scores)
return outputs[0][:, 0].tolist()
def _rerank_documents(
self, query: str, documents: Iterable[str], batch_size: int, **kwargs: Any
self, query: str, documents: Iterable[str], batch_size: int, **kwargs
) -> Iterable[float]:
if not hasattr(self, "model") or self.model is None:
self.load_onnx_model()
for batch in iter_batch(documents, batch_size):
yield from self._post_process_onnx_output(self.onnx_embed(query, batch, **kwargs))
def _rerank_pairs(
self,
model_name: str,
cache_dir: str,
pairs: Iterable[tuple[str, str]],
batch_size: int,
parallel: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
**kwargs: Any,
) -> Iterable[float]:
is_small = False
if isinstance(pairs, tuple):
pairs = [pairs]
is_small = True
if isinstance(pairs, list):
if len(pairs) < batch_size:
is_small = True
if parallel is None or is_small:
if not hasattr(self, "model") or self.model is None:
self.load_onnx_model()
for batch in iter_batch(pairs, batch_size):
yield from self._post_process_onnx_output(self.onnx_embed_pairs(batch, **kwargs))
else:
if parallel == 0:
parallel = os.cpu_count()
start_method = "forkserver" if "forkserver" in get_all_start_methods() else "spawn"
params = {
"model_name": model_name,
"cache_dir": cache_dir,
"providers": providers,
**kwargs,
}
pool = ParallelWorkerPool(
num_workers=parallel or 1,
worker=self._get_worker_class(),
cuda=cuda,
device_ids=device_ids,
start_method=start_method,
)
for batch in pool.ordered_map(iter_batch(pairs, batch_size), **params):
yield from self._post_process_onnx_output(batch)
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[float]:
raise NotImplementedError("Subclasses must implement this method")
yield from self.onnx_embed(query, batch, **kwargs)
def _preprocess_onnx_input(
self, onnx_input: dict[str, np.ndarray], **kwargs: Any
) -> dict[str, np.ndarray]:
self, onnx_input: Dict[str, np.ndarray], **kwargs
) -> Dict[str, np.ndarray]:
"""
Preprocess the onnx input.
"""
return onnx_input
class TextRerankerWorker(EmbeddingWorker):
def process(self, items: Iterable[tuple[int, Any]]) -> Iterable[tuple[int, Any]]:
for idx, batch in items:
onnx_output = self.model.onnx_embed_pairs(batch)
yield idx, onnx_output
@@ -1,21 +1,21 @@
from typing import Any, Iterable, Optional, Sequence, Type
from typing import Any, Dict, Iterable, List, Optional, Sequence, Type
from fastembed.common import OnnxProvider
from fastembed.rerank.cross_encoder.onnx_text_cross_encoder import OnnxTextCrossEncoder
from fastembed.rerank.cross_encoder.text_cross_encoder_base import TextCrossEncoderBase
from fastembed.rerank.cross_encoder.onnx_text_cross_encoder import OnnxTextCrossEncoder
from fastembed.common import OnnxProvider
class TextCrossEncoder(TextCrossEncoderBase):
CROSS_ENCODER_REGISTRY: list[Type[TextCrossEncoderBase]] = [
CROSS_ENCODER_REGISTRY: List[Type[TextCrossEncoderBase]] = [
OnnxTextCrossEncoder,
]
@classmethod
def list_supported_models(cls) -> list[dict[str, Any]]:
def list_supported_models(cls) -> List[Dict[str, Any]]:
"""Lists the supported models.
Returns:
list[dict[str, Any]]: A list of dictionaries containing the model information.
List[Dict[str, Any]]: A list of dictionaries containing the model information.
Example:
```
@@ -45,9 +45,9 @@ class TextCrossEncoder(TextCrossEncoderBase):
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
device_ids: Optional[List[int]] = None,
lazy_load: bool = False,
**kwargs: Any,
**kwargs,
):
super().__init__(model_name, cache_dir, threads, **kwargs)
@@ -72,7 +72,7 @@ class TextCrossEncoder(TextCrossEncoderBase):
)
def rerank(
self, query: str, documents: Iterable[str], batch_size: int = 64, **kwargs: Any
self, query: str, documents: Iterable[str], batch_size: int = 64, **kwargs
) -> Iterable[float]:
"""Rerank a list of documents based on a query.
@@ -85,36 +85,3 @@ class TextCrossEncoder(TextCrossEncoderBase):
Iterable of scores for each document
"""
yield from self.model.rerank(query, documents, batch_size=batch_size, **kwargs)
def rerank_pairs(
self,
pairs: Iterable[tuple[str, str]],
batch_size: int = 64,
parallel: Optional[int] = None,
**kwargs: Any,
) -> Iterable[float]:
"""
Rerank a list of query-document pairs.
Args:
pairs (Iterable[tuple[str, str]]): An iterable of tuples, where each tuple contains a query and a document
to be scored together.
batch_size (int, optional): The number of query-document pairs to process in a single batch. Defaults to 64.
parallel (Optional[int], optional): The number of parallel processes to use for reranking.
If None, parallelization is disabled. Defaults to None.
**kwargs (Any): Additional arguments to pass to the underlying reranking model.
Returns:
Iterable[float]: An iterable of scores corresponding to each query-document pair in the input.
Higher scores indicate a stronger match between the query and the document.
Example:
>>> encoder = TextCrossEncoder("Xenova/ms-marco-MiniLM-L-6-v2")
>>> pairs = [("What is AI?", "Artificial intelligence is ..."), ("What is ML?", "Machine learning is ...")]
>>> scores = list(encoder.rerank_pairs(pairs))
>>> print(list(map(lambda x: round(x, 2), scores)))
[-1.24, -10.6]
"""
yield from self.model.rerank_pairs(
pairs, batch_size=batch_size, parallel=parallel, **kwargs
)
@@ -1,4 +1,4 @@
from typing import Any, Iterable, Optional
from typing import Iterable, Optional
from fastembed.common.model_management import ModelManagement
@@ -23,7 +23,7 @@ class TextCrossEncoderBase(ModelManagement):
batch_size: int = 64,
**kwargs,
) -> Iterable[float]:
"""Rerank a list of documents given a query.
"""Reranks a list of documents given a query.
Args:
query (str): The query to rerank the documents.
@@ -32,27 +32,6 @@ class TextCrossEncoderBase(ModelManagement):
**kwargs: Additional keyword argument to pass to the rerank method.
Yields:
Iterable[float]: The scores of the reranked the documents.
"""
raise NotImplementedError("This method should be overridden by subclasses")
def rerank_pairs(
self,
pairs: Iterable[tuple[str, str]],
batch_size: int = 64,
parallel: Optional[int] = None,
**kwargs: Any,
) -> Iterable[float]:
"""Rerank query-document pairs.
Args:
pairs (Iterable[tuple[str, str]]): Query-document pairs to rerank
batch_size (int): The batch size to use for reranking.
parallel: 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.
**kwargs: Additional keyword argument to pass to the rerank method.
Yields:
Iterable[float]: Scores for each individual pair
Iterable[float]: The scores of reranked the documents.
"""
raise NotImplementedError("This method should be overridden by subclasses")
+15 -29
View File
@@ -2,7 +2,7 @@ import os
from collections import defaultdict
from multiprocessing import get_all_start_methods
from pathlib import Path
from typing import Any, Iterable, Optional, Type, Union
from typing import Any, Dict, Iterable, List, Optional, Tuple, Type, Union
import mmh3
import numpy as np
@@ -92,8 +92,6 @@ class Bm25(SparseTextEmbeddingBase):
b (float, optional): The b parameter in the BM25 formula. Defines the importance of the document length.
Defaults to 0.75.
avg_len (float, optional): The average length of the documents in the corpus. Defaults to 256.0.
language (str): Specifies the language for the stemmer.
disable_stemmer (bool): Disable the stemmer.
Raises:
ValueError: If the model_name is not in the format <org>/<model> e.g. BAAI/bge-base-en.
"""
@@ -107,7 +105,6 @@ class Bm25(SparseTextEmbeddingBase):
avg_len: float = 256.0,
language: str = "english",
token_max_length: int = 40,
disable_stemmer: bool = False,
**kwargs,
):
super().__init__(model_name, cache_dir, **kwargs)
@@ -130,28 +127,22 @@ class Bm25(SparseTextEmbeddingBase):
self.token_max_length = token_max_length
self.punctuation = set(get_all_punctuation())
self.disable_stemmer = disable_stemmer
if disable_stemmer:
self.stopwords = set()
self.stemmer = None
else:
self.stopwords = set(self._load_stopwords(self._model_dir, self.language))
self.stemmer = SnowballStemmer(language)
self.stopwords = set(self._load_stopwords(self._model_dir, self.language))
self.stemmer = SnowballStemmer(language)
self.tokenizer = SimpleTokenizer
@classmethod
def list_supported_models(cls) -> list[dict[str, Any]]:
def list_supported_models(cls) -> List[Dict[str, Any]]:
"""Lists the supported models.
Returns:
list[dict[str, Any]]: A list of dictionaries containing the model information.
List[Dict[str, Any]]: A list of dictionaries containing the model information.
"""
return supported_bm25_models
@classmethod
def _load_stopwords(cls, model_dir: Path, language: str) -> list[str]:
def _load_stopwords(cls, model_dir: Path, language: str) -> List[str]:
stopwords_path = model_dir / f"{language}.txt"
if not stopwords_path.exists():
return []
@@ -191,9 +182,6 @@ class Bm25(SparseTextEmbeddingBase):
"k": self.k,
"b": self.b,
"avg_len": self.avg_len,
"language": self.language,
"token_max_length": self.token_max_length,
"disable_stemmer": self.disable_stemmer,
}
pool = ParallelWorkerPool(
num_workers=parallel or 1,
@@ -234,21 +222,19 @@ class Bm25(SparseTextEmbeddingBase):
parallel=parallel,
)
def _stem(self, tokens: list[str]) -> list[str]:
def _stem(self, tokens: List[str]) -> List[str]:
stemmed_tokens = []
for token in tokens:
lower_token = token.lower()
if token in self.punctuation:
continue
if lower_token in self.stopwords:
if token.lower() in self.stopwords:
continue
if len(token) > self.token_max_length:
continue
stemmed_token = self.stemmer.stem_word(lower_token) if self.stemmer else lower_token
stemmed_token = self.stemmer.stem_word(token.lower())
if stemmed_token:
stemmed_tokens.append(stemmed_token)
@@ -256,8 +242,8 @@ class Bm25(SparseTextEmbeddingBase):
def raw_embed(
self,
documents: list[str],
) -> list[SparseEmbedding]:
documents: List[str],
) -> List[SparseEmbedding]:
embeddings = []
for document in documents:
document = remove_non_alphanumeric(document)
@@ -267,7 +253,7 @@ class Bm25(SparseTextEmbeddingBase):
embeddings.append(SparseEmbedding.from_dict(token_id2value))
return embeddings
def _term_frequency(self, tokens: list[str]) -> dict[int, float]:
def _term_frequency(self, tokens: List[str]) -> Dict[int, float]:
"""Calculate the term frequency part of the BM25 formula.
(
@@ -277,10 +263,10 @@ class Bm25(SparseTextEmbeddingBase):
)
Args:
tokens (list[str]): The list of tokens in the document.
tokens (List[str]): The list of tokens in the document.
Returns:
dict[int, float]: The token_id to term frequency mapping.
Dict[int, float]: The token_id to term frequency mapping.
"""
tf_map = {}
counter = defaultdict(int)
@@ -337,7 +323,7 @@ class Bm25Worker(Worker):
def start(cls, model_name: str, cache_dir: str, **kwargs: Any) -> "Bm25Worker":
return cls(model_name=model_name, cache_dir=cache_dir, **kwargs)
def process(self, items: Iterable[tuple[int, Any]]) -> Iterable[tuple[int, Any]]:
def process(self, items: Iterable[Tuple[int, Any]]) -> Iterable[Tuple[int, Any]]:
for idx, batch in items:
onnx_output = self.model.raw_embed(batch)
yield idx, onnx_output
+14 -14
View File
@@ -1,7 +1,7 @@
import math
import string
from pathlib import Path
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Dict, Iterable, List, Optional, Sequence, Tuple, Type, Union
import mmh3
import numpy as np
@@ -63,7 +63,7 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
providers: Optional[Sequence[OnnxProvider]] = None,
alpha: float = 0.5,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
device_ids: Optional[List[int]] = None,
lazy_load: bool = False,
device_id: Optional[int] = None,
**kwargs,
@@ -81,7 +81,7 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
It is recommended to only change this parameter based on training data for a specific dataset.
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to False.
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
device_ids (Optional[List[int]], optional): The list of device ids to use for data parallel processing in
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
@@ -141,7 +141,7 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
self.special_tokens_ids = set(self.special_token_to_id.values())
self.stopwords = set(self._load_stopwords(self._model_dir))
def _filter_pair_tokens(self, tokens: list[tuple[str, Any]]) -> list[tuple[str, Any]]:
def _filter_pair_tokens(self, tokens: List[Tuple[str, Any]]) -> List[Tuple[str, Any]]:
result = []
for token, value in tokens:
if token in self.stopwords or token in self.punctuation:
@@ -149,7 +149,7 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
result.append((token, value))
return result
def _stem_pair_tokens(self, tokens: list[tuple[str, Any]]) -> list[tuple[str, Any]]:
def _stem_pair_tokens(self, tokens: List[Tuple[str, Any]]) -> List[Tuple[str, Any]]:
result = []
for token, value in tokens:
processed_token = self.stemmer.stem_word(token)
@@ -158,8 +158,8 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
@classmethod
def _aggregate_weights(
cls, tokens: list[tuple[str, list[int]]], weights: list[float]
) -> list[tuple[str, float]]:
cls, tokens: List[Tuple[str, List[int]]], weights: List[float]
) -> List[Tuple[str, float]]:
result = []
for token, idxs in tokens:
sum_weight = sum(weights[idx] for idx in idxs)
@@ -167,8 +167,8 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
return result
def _reconstruct_bpe(
self, bpe_tokens: Iterable[tuple[int, str]]
) -> list[tuple[str, list[int]]]:
self, bpe_tokens: Iterable[Tuple[int, str]]
) -> List[Tuple[str, List[int]]]:
result = []
acc = ""
acc_idx = []
@@ -195,7 +195,7 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
return result
def _rescore_vector(self, vector: dict[str, float]) -> dict[int, float]:
def _rescore_vector(self, vector: Dict[str, float]) -> Dict[int, float]:
"""
Orders all tokens in the vector by their importance and generates a new score based on the importance order.
So that the scoring doesn't depend on absolute values assigned by the model, but on the relative importance.
@@ -246,16 +246,16 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
yield SparseEmbedding.from_dict(rescored)
@classmethod
def list_supported_models(cls) -> list[dict[str, Any]]:
def list_supported_models(cls) -> List[Dict[str, Any]]:
"""Lists the supported models.
Returns:
list[dict[str, Any]]: A list of dictionaries containing the model information.
List[Dict[str, Any]]: A list of dictionaries containing the model information.
"""
return supported_bm42_models
@classmethod
def _load_stopwords(cls, model_dir: Path) -> list[str]:
def _load_stopwords(cls, model_dir: Path) -> List[str]:
stopwords_path = model_dir / "stopwords.txt"
if not stopwords_path.exists():
return []
@@ -298,7 +298,7 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
)
@classmethod
def _query_rehash(cls, tokens: Iterable[str]) -> dict[int, float]:
def _query_rehash(cls, tokens: Iterable[str]) -> Dict[int, float]:
result = {}
for token in tokens:
token_id = abs(mmh3.hash(token))
+10 -6
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass
from typing import Iterable, Optional, Union
from typing import Dict, Iterable, Optional, Union
import numpy as np
@@ -11,17 +11,17 @@ class SparseEmbedding:
values: np.ndarray
indices: np.ndarray
def as_object(self) -> dict[str, np.ndarray]:
def as_object(self) -> Dict[str, np.ndarray]:
return {
"values": self.values,
"indices": self.indices,
}
def as_dict(self) -> dict[int, float]:
def as_dict(self) -> Dict[int, float]:
return {i: v for i, v in zip(self.indices, self.values)}
@classmethod
def from_dict(cls, data: dict[int, float]) -> "SparseEmbedding":
def from_dict(cls, data: Dict[int, float]) -> "SparseEmbedding":
if len(data) == 0:
return cls(values=np.array([]), indices=np.array([]))
indices, values = zip(*data.items())
@@ -50,7 +50,9 @@ class SparseTextEmbeddingBase(ModelManagement):
) -> Iterable[SparseEmbedding]:
raise NotImplementedError()
def passage_embed(self, texts: Iterable[str], **kwargs) -> Iterable[SparseEmbedding]:
def passage_embed(
self, texts: Iterable[str], **kwargs
) -> Iterable[SparseEmbedding]:
"""
Embeds a list of text passages into a list of embeddings.
@@ -65,7 +67,9 @@ class SparseTextEmbeddingBase(ModelManagement):
# This is model-specific, so that different models can have specialized implementations
yield from self.embed(texts, **kwargs)
def query_embed(self, query: Union[str, Iterable[str]], **kwargs) -> Iterable[SparseEmbedding]:
def query_embed(
self, query: Union[str, Iterable[str]], **kwargs
) -> Iterable[SparseEmbedding]:
"""
Embeds queries
+5 -5
View File
@@ -1,4 +1,4 @@
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Dict, Iterable, List, Optional, Sequence, Type, Union
from fastembed.common import OnnxProvider
from fastembed.sparse.bm25 import Bm25
@@ -12,15 +12,15 @@ import warnings
class SparseTextEmbedding(SparseTextEmbeddingBase):
EMBEDDINGS_REGISTRY: list[Type[SparseTextEmbeddingBase]] = [SpladePP, Bm42, Bm25]
EMBEDDINGS_REGISTRY: List[Type[SparseTextEmbeddingBase]] = [SpladePP, Bm42, Bm25]
@classmethod
def list_supported_models(cls) -> list[dict[str, Any]]:
def list_supported_models(cls) -> List[Dict[str, Any]]:
"""
Lists the supported models.
Returns:
list[dict[str, Any]]: A list of dictionaries containing the model information.
List[Dict[str, Any]]: A list of dictionaries containing the model information.
Example:
```
@@ -50,7 +50,7 @@ class SparseTextEmbedding(SparseTextEmbeddingBase):
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
device_ids: Optional[List[int]] = None,
lazy_load: bool = False,
**kwargs,
):
+5 -5
View File
@@ -1,4 +1,4 @@
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Dict, Iterable, List, Optional, Sequence, Type, Union
import numpy as np
from fastembed.common import OnnxProvider
@@ -55,11 +55,11 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
yield SparseEmbedding(values=scores, indices=indices)
@classmethod
def list_supported_models(cls) -> list[dict[str, Any]]:
def list_supported_models(cls) -> List[Dict[str, Any]]:
"""Lists the supported models.
Returns:
list[dict[str, Any]]: A list of dictionaries containing the model information.
List[Dict[str, Any]]: A list of dictionaries containing the model information.
"""
return supported_splade_models
@@ -70,7 +70,7 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
device_ids: Optional[List[int]] = None,
lazy_load: bool = False,
device_id: Optional[int] = None,
**kwargs,
@@ -86,7 +86,7 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to False.
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
device_ids (Optional[List[int]], optional): The list of device ids to use for data parallel processing in
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
+3 -3
View File
@@ -1,11 +1,11 @@
# This code is a modified copy of the `NLTKWordTokenizer` class from `NLTK` library.
import re
from typing import List
class SimpleTokenizer:
@staticmethod
def tokenize(text: str) -> list[str]:
def tokenize(text: str) -> List[str]:
text = re.sub(r"[^\w]", " ", text.lower())
text = re.sub(r"\s+", " ", text)
@@ -80,7 +80,7 @@ class WordTokenizer:
]
@classmethod
def tokenize(cls, text: str) -> list[str]:
def tokenize(cls, text: str) -> List[str]:
"""Return a tokenized copy of `text`.
>>> s = '''Good muffins cost $3.88 (roughly 3,36 euros)\nin New York.'''
+3 -3
View File
@@ -1,4 +1,4 @@
from typing import Any, Iterable, Type
from typing import Any, Dict, Iterable, List, Type
import numpy as np
@@ -27,11 +27,11 @@ class CLIPOnnxEmbedding(OnnxTextEmbedding):
return CLIPEmbeddingWorker
@classmethod
def list_supported_models(cls) -> list[dict[str, Any]]:
def list_supported_models(cls) -> List[Dict[str, Any]]:
"""Lists the supported models.
Returns:
list[dict[str, Any]]: A list of dictionaries containing the model information.
List[Dict[str, Any]]: A list of dictionaries containing the model information.
"""
return supported_clip_models
+5 -5
View File
@@ -1,4 +1,4 @@
from typing import Any, Type
from typing import Any, Dict, List, Type
import numpy as np
@@ -39,17 +39,17 @@ class E5OnnxEmbedding(OnnxTextEmbedding):
return E5OnnxEmbeddingWorker
@classmethod
def list_supported_models(cls) -> list[dict[str, Any]]:
def list_supported_models(cls) -> List[Dict[str, Any]]:
"""Lists the supported models.
Returns:
list[dict[str, Any]]: A list of dictionaries containing the model information.
List[Dict[str, Any]]: A list of dictionaries containing the model information.
"""
return supported_multilingual_e5_models
def _preprocess_onnx_input(
self, onnx_input: dict[str, np.ndarray], **kwargs
) -> dict[str, np.ndarray]:
self, onnx_input: Dict[str, np.ndarray], **kwargs
) -> Dict[str, np.ndarray]:
"""
Preprocess the onnx input.
"""
+8 -28
View File
@@ -1,4 +1,4 @@
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Dict, Iterable, List, Optional, Sequence, Type, Union
import numpy as np
@@ -16,7 +16,6 @@ supported_onnx_models = [
"license": "mit",
"size_in_GB": 0.42,
"sources": {
"hf": "Qdrant/fast-bge-base-en",
"url": "https://storage.googleapis.com/qdrant-fastembed/fast-bge-base-en.tar.gz",
},
"model_file": "model_optimized.onnx",
@@ -51,7 +50,6 @@ supported_onnx_models = [
"license": "mit",
"size_in_GB": 0.13,
"sources": {
"hf": "Qdrant/bge-small-en",
"url": "https://storage.googleapis.com/qdrant-fastembed/BAAI-bge-small-en.tar.gz",
},
"model_file": "model_optimized.onnx",
@@ -74,7 +72,6 @@ supported_onnx_models = [
"license": "mit",
"size_in_GB": 0.09,
"sources": {
"hf": "Qdrant/bge-small-zh-v1.5",
"url": "https://storage.googleapis.com/qdrant-fastembed/fast-bge-small-zh-v1.5.tar.gz",
},
"model_file": "model_optimized.onnx",
@@ -167,17 +164,6 @@ supported_onnx_models = [
},
"model_file": "onnx/model.onnx",
},
{
"model": "jinaai/jina-clip-v1",
"dim": 768,
"description": "Text embeddings, Multimodal (text&image), English, Prefixes for queries/documents: not necessary, 2024 year",
"license": "apache-2.0",
"size_in_GB": 0.55,
"sources": {
"hf": "jinaai/jina-clip-v1",
},
"model_file": "onnx/text_model.onnx",
},
]
@@ -185,12 +171,12 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[np.ndarray]):
"""Implementation of the Flag Embedding model."""
@classmethod
def list_supported_models(cls) -> list[dict[str, Any]]:
def list_supported_models(cls) -> List[Dict[str, Any]]:
"""
Lists the supported models.
Returns:
list[dict[str, Any]]: A list of dictionaries containing the model information.
List[Dict[str, Any]]: A list of dictionaries containing the model information.
"""
return supported_onnx_models
@@ -201,7 +187,7 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[np.ndarray]):
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
device_ids: Optional[List[int]] = None,
lazy_load: bool = False,
device_id: Optional[int] = None,
**kwargs,
@@ -217,7 +203,7 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[np.ndarray]):
Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to False.
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
device_ids (Optional[List[int]], optional): The list of device ids to use for data parallel processing in
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
@@ -290,8 +276,8 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[np.ndarray]):
return OnnxTextEmbeddingWorker
def _preprocess_onnx_input(
self, onnx_input: dict[str, np.ndarray], **kwargs
) -> dict[str, np.ndarray]:
self, onnx_input: Dict[str, np.ndarray], **kwargs
) -> Dict[str, np.ndarray]:
"""
Preprocess the onnx input.
"""
@@ -299,13 +285,7 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[np.ndarray]):
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[np.ndarray]:
embeddings = output.model_output
if embeddings.ndim == 3: # (batch_size, seq_len, embedding_dim)
processed_embeddings = embeddings[:, 0]
elif embeddings.ndim == 2: # (batch_size, embedding_dim)
processed_embeddings = embeddings
else:
raise ValueError(f"Unsupported embedding shape: {embeddings.shape}")
return normalize(processed_embeddings).astype(np.float32)
return normalize(embeddings[:, 0]).astype(np.float32)
def load_onnx_model(self) -> None:
self._load_onnx_model(
+9 -8
View File
@@ -1,7 +1,7 @@
import os
from multiprocessing import get_all_start_methods
from pathlib import Path
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Dict, Iterable, List, Optional, Sequence, Tuple, Type, Union
import numpy as np
from tokenizers import Encoding
@@ -14,7 +14,7 @@ from fastembed.parallel_processor import ParallelWorkerPool
class OnnxTextModel(OnnxModel[T]):
ONNX_OUTPUT_NAMES: Optional[list[str]] = None
ONNX_OUTPUT_NAMES: Optional[List[str]] = None
@classmethod
def _get_worker_class(cls) -> Type["TextEmbeddingWorker"]:
@@ -29,8 +29,8 @@ class OnnxTextModel(OnnxModel[T]):
self.special_token_to_id = {}
def _preprocess_onnx_input(
self, onnx_input: dict[str, np.ndarray], **kwargs
) -> dict[str, np.ndarray]:
self, onnx_input: Dict[str, np.ndarray], **kwargs
) -> Dict[str, np.ndarray]:
"""
Preprocess the onnx input.
"""
@@ -58,12 +58,12 @@ class OnnxTextModel(OnnxModel[T]):
def load_onnx_model(self) -> None:
raise NotImplementedError("Subclasses must implement this method")
def tokenize(self, documents: list[str], **kwargs) -> list[Encoding]:
def tokenize(self, documents: List[str], **kwargs) -> List[Encoding]:
return self.tokenizer.encode_batch(documents)
def onnx_embed(
self,
documents: list[str],
documents: List[str],
**kwargs,
) -> OnnxOutputContext:
encoded = self.tokenize(documents, **kwargs)
@@ -79,6 +79,7 @@ class OnnxTextModel(OnnxModel[T]):
onnx_input["token_type_ids"] = np.array(
[np.zeros(len(e), dtype=np.int64) for e in input_ids], dtype=np.int64
)
onnx_input = self._preprocess_onnx_input(onnx_input, **kwargs)
model_output = self.model.run(self.ONNX_OUTPUT_NAMES, onnx_input)
@@ -97,7 +98,7 @@ class OnnxTextModel(OnnxModel[T]):
parallel: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
device_ids: Optional[List[int]] = None,
**kwargs,
) -> Iterable[T]:
is_small = False
@@ -139,7 +140,7 @@ class OnnxTextModel(OnnxModel[T]):
class TextEmbeddingWorker(EmbeddingWorker):
def process(self, items: Iterable[tuple[int, Any]]) -> Iterable[tuple[int, Any]]:
def process(self, items: Iterable[Tuple[int, Any]]) -> Iterable[Tuple[int, Any]]:
for idx, batch in items:
onnx_output = self.model.onnx_embed(batch)
yield idx, onnx_output
+3 -3
View File
@@ -1,4 +1,4 @@
from typing import Any, Iterable, Type
from typing import Any, Dict, Iterable, List, Type
import numpy as np
@@ -60,11 +60,11 @@ class PooledEmbedding(OnnxTextEmbedding):
return pooled_embeddings
@classmethod
def list_supported_models(cls) -> list[dict[str, Any]]:
def list_supported_models(cls) -> List[Dict[str, Any]]:
"""Lists the supported models.
Returns:
list[dict[str, Any]]: A list of dictionaries containing the model information.
List[Dict[str, Any]]: A list of dictionaries containing the model information.
"""
return supported_pooled_models
+3 -30
View File
@@ -1,4 +1,4 @@
from typing import Any, Iterable, Type
from typing import Any, Dict, Iterable, List, Type
import numpy as np
@@ -57,33 +57,6 @@ supported_pooled_normalized_models = [
"sources": {"hf": "jinaai/jina-embeddings-v2-base-code"},
"model_file": "onnx/model.onnx",
},
{
"model": "jinaai/jina-embeddings-v2-base-zh",
"dim": 768,
"description": "Text embeddings, Unimodal (text), supports mixed Chinese-English input text, 8192 input tokens truncation, Prefixes for queries/documents: not necessary, 2024 year.",
"license": "apache-2.0",
"size_in_GB": 0.64,
"sources": {"hf": "jinaai/jina-embeddings-v2-base-zh"},
"model_file": "onnx/model.onnx",
},
{
"model": "jinaai/jina-embeddings-v2-base-es",
"dim": 768,
"description": "Text embeddings, Unimodal (text), supports mixed Spanish-English input text, 8192 input tokens truncation, Prefixes for queries/documents: not necessary, 2024 year.",
"license": "apache-2.0",
"size_in_GB": 0.64,
"sources": {"hf": "jinaai/jina-embeddings-v2-base-es"},
"model_file": "onnx/model.onnx",
},
{
"model": "thenlper/gte-base",
"dim": 768,
"description": "General text embeddings, Unimodal (text), supports English only input text, 512 input tokens truncation, Prefixes for queries/documents: not necessary, 2024 year.",
"license": "mit",
"size_in_GB": 0.44,
"sources": {"hf": "thenlper/gte-base"},
"model_file": "onnx/model.onnx",
},
]
@@ -93,11 +66,11 @@ class PooledNormalizedEmbedding(PooledEmbedding):
return PooledNormalizedEmbeddingWorker
@classmethod
def list_supported_models(cls) -> list[dict[str, Any]]:
def list_supported_models(cls) -> List[Dict[str, Any]]:
"""Lists the supported models.
Returns:
list[dict[str, Any]]: A list of dictionaries containing the model information.
List[Dict[str, Any]]: A list of dictionaries containing the model information.
"""
return supported_pooled_normalized_models
+6 -6
View File
@@ -1,4 +1,4 @@
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Dict, Iterable, List, Optional, Sequence, Type, Union
import numpy as np
@@ -12,7 +12,7 @@ from fastembed.text.text_embedding_base import TextEmbeddingBase
class TextEmbedding(TextEmbeddingBase):
EMBEDDINGS_REGISTRY: list[Type[TextEmbeddingBase]] = [
EMBEDDINGS_REGISTRY: List[Type[TextEmbeddingBase]] = [
OnnxTextEmbedding,
E5OnnxEmbedding,
CLIPOnnxEmbedding,
@@ -21,12 +21,12 @@ class TextEmbedding(TextEmbeddingBase):
]
@classmethod
def list_supported_models(cls) -> list[dict[str, Any]]:
def list_supported_models(cls) -> List[Dict[str, Any]]:
"""
Lists the supported models.
Returns:
list[dict[str, Any]]: A list of dictionaries containing the model information.
List[Dict[str, Any]]: A list of dictionaries containing the model information.
Example:
```
@@ -57,7 +57,7 @@ class TextEmbedding(TextEmbeddingBase):
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
device_ids: Optional[List[int]] = None,
lazy_load: bool = False,
**kwargs,
):
@@ -78,7 +78,7 @@ class TextEmbedding(TextEmbeddingBase):
return
raise ValueError(
f"Model {model_name} is not supported in TextEmbedding. "
f"Model {model_name} is not supported in TextEmbedding."
"Please check the supported models using `TextEmbedding.list_supported_models()`"
)
+9 -15
View File
@@ -1,6 +1,6 @@
[tool.poetry]
name = "fastembed-gpu"
version = "0.5.0"
version = "0.4.2"
description = "Fast, light, accurate library built for retrieval embedding generation"
authors = ["Qdrant Team <info@qdrant.tech>", "NirantK <nirant.bits@gmail.com>"]
license = "Apache License"
@@ -11,24 +11,18 @@ repository = "https://github.com/qdrant/fastembed"
keywords = ["vector", "embedding", "neural", "search", "qdrant", "sentence-transformers"]
[tool.poetry.dependencies]
python = ">=3.9.0"
onnx = ">=1.15.0"
numpy = [
{ version = ">=1.21,<2.1.0", python = "<3.10" },
{ version = ">=1.21", python = ">=3.10,<3.12" },
{ version = ">=1.26", python = ">=3.12,<3.13" },
{ version = ">=2.1.0", python = ">=3.13" }
]
onnxruntime-gpu = [
{ version = ">=1.17.0,<1.20.0", python = "<3.10" },
{ version = ">=1.17.0,!=1.20.0", python = ">=3.10,<3.13" },
{ version = ">1.20.0", python = ">=3.13" }
]
python = ">=3.8.0,<3.13"
onnx = "^1.15.0"
onnxruntime-gpu = ">=1.17.0,<1.20.0"
tqdm = "^4.66"
requests = "^2.31"
tokenizers = ">=0.15,<1.0"
huggingface-hub = ">=0.20,<1.0"
loguru = "^0.7.2"
numpy = [
{ version = ">=1.21, <2", python = "<3.12" },
{ version = ">=1.26, <2", python = ">=3.12" }
]
pillow = "^10.3.0"
mmh3 = "^4.1.0"
py-rust-stemmers = "^0.1.0"
@@ -37,7 +31,7 @@ py-rust-stemmers = "^0.1.0"
pytest = "^7.4.2"
ruff = ">=0.3.1,<1.0"
notebook = ">=7.0.2"
pre-commit = "^3.6.2"
pre-commit = {version = "^3.6.2", python = ">=3.9,<3.12" }
[tool.poetry.group.docs.dependencies]
mkdocs-material = "^9.5.10"
+6 -6
View File
@@ -9,7 +9,7 @@
# %%
import time
from typing import Callable
from typing import Callable, List, Tuple
import matplotlib.pyplot as plt
import torch.nn.functional as F
@@ -23,7 +23,7 @@ from fastembed.embedding import DefaultEmbedding
# data is a list of strings, each string is a document.
# %%
documents: list[str] = [
documents: List[str] = [
"Chandrayaan-3 is India's third lunar mission",
"It aimed to land a rover on the Moon's surface - joining the US, China and Russia",
"The mission is a follow-up to Chandrayaan-2, which had partial success",
@@ -56,7 +56,7 @@ class HF:
self.model = AutoModel.from_pretrained(model_id)
self.tokenizer = AutoTokenizer.from_pretrained(model_id)
def embed(self, texts: list[str]):
def embed(self, texts: List[str]):
encoded_input = self.tokenizer(
texts, max_length=512, padding=True, truncation=True, return_tensors="pt"
)
@@ -88,7 +88,7 @@ embedding_model = DefaultEmbedding()
# %%
def calculate_time_stats(
embed_func: Callable, documents: list, k: int
) -> tuple[float, float, float]:
) -> Tuple[float, float, float]:
times = []
for _ in range(k):
# Timing the embed_func call
@@ -111,8 +111,8 @@ print(f"FastEmbed (Average, Max, Min): {fst_stats}")
# %%
def plot_character_per_second_comparison(
hf_stats: tuple[float, float, float],
fst_stats: tuple[float, float, float],
hf_stats: Tuple[float, float, float],
fst_stats: Tuple[float, float, float],
documents: list,
):
# Calculating total characters in documents
-3
View File
@@ -21,9 +21,6 @@ CANONICAL_VECTOR_VALUES = {
"Qdrant/Unicom-ViT-B-32": np.array(
[0.0418, 0.0550, 0.0003, 0.0253, -0.0185, 0.0016, -0.0368, -0.0402, -0.0891, -0.0186]
),
"jinaai/jina-clip-v1": np.array(
[-0.029, 0.0216, 0.0396, 0.0283, -0.0023, 0.0151, 0.011, -0.0235, 0.0251, -0.0343]
),
}
-21
View File
@@ -151,27 +151,6 @@ def test_stem_case_insensitive_stopwords(bm25_instance):
assert result == expected, f"Expected {expected}, but got {result}"
@pytest.mark.parametrize("disable_stemmer", [True, False])
def test_disable_stemmer_behavior(disable_stemmer):
# Setup
model = Bm25("Qdrant/bm25", language="english", disable_stemmer=disable_stemmer)
model.stopwords = {"the", "is", "a"}
model.punctuation = {".", ",", "!"}
# Test data
tokens = ["The", "quick", "brown", "fox", "is", "a", "test", "sentence", ".", "!"]
# Execute
result = model._stem(tokens)
# Assert
if disable_stemmer:
expected = ["quick", "brown", "fox", "test", "sentence"] # no stemming, lower case only
else:
expected = ["quick", "brown", "fox", "test", "sentenc"]
assert result == expected, f"Expected {expected}, but got {result}"
@pytest.mark.parametrize(
"model_name",
["prithivida/Splade_PP_en_v1"],
+16 -60
View File
@@ -10,48 +10,34 @@ CANONICAL_SCORE_VALUES = {
"Xenova/ms-marco-MiniLM-L-6-v2": np.array([8.500708, -2.541011]),
"Xenova/ms-marco-MiniLM-L-12-v2": np.array([9.330912, -2.0380247]),
"BAAI/bge-reranker-base": np.array([6.15733337, -3.65939403]),
"jinaai/jina-reranker-v1-tiny-en": np.array([2.5911, 0.1122]),
"jinaai/jina-reranker-v1-turbo-en": np.array([1.8295, -2.8908]),
"jinaai/jina-reranker-v2-base-multilingual": np.array([1.6533, -1.6455]),
}
SELECTED_MODELS = {
"Xenova": "Xenova/ms-marco-MiniLM-L-6-v2",
"BAAI": "BAAI/bge-reranker-base",
"jinaai": "jinaai/jina-reranker-v1-tiny-en",
}
@pytest.mark.parametrize(
"model_name",
[model_name for model_name in CANONICAL_SCORE_VALUES],
)
def test_rerank(model_name):
def test_rerank():
is_ci = os.getenv("CI")
model = TextCrossEncoder(model_name=model_name)
for model_desc in TextCrossEncoder.list_supported_models():
if not is_ci and model_desc["size_in_GB"] > 1:
continue
query = "What is the capital of France?"
documents = ["Paris is the capital of France.", "Berlin is the capital of Germany."]
scores = np.array(list(model.rerank(query, documents)))
model_name = model_desc["model"]
model = TextCrossEncoder(model_name=model_name)
pairs = [(query, doc) for doc in documents]
scores2 = np.array(list(model.rerank_pairs(pairs)))
assert np.allclose(
scores, scores2, atol=1e-5
), f"Model: {model_name}, Scores: {scores}, Scores2: {scores2}"
query = "What is the capital of France?"
documents = ["Paris is the capital of France.", "Berlin is the capital of Germany."]
scores = np.array(list(model.rerank(query, documents)))
canonical_scores = CANONICAL_SCORE_VALUES[model_name]
assert np.allclose(
scores, canonical_scores, atol=1e-3
), f"Model: {model_name}, Scores: {scores}, Expected: {canonical_scores}"
if is_ci:
delete_model_cache(model.model._model_dir)
canonical_scores = CANONICAL_SCORE_VALUES[model_name]
assert np.allclose(
scores, canonical_scores, atol=1e-3
), f"Model: {model_name}, Scores: {scores}, Expected: {canonical_scores}"
if is_ci:
delete_model_cache(model.model._model_dir)
@pytest.mark.parametrize(
"model_name",
[model_name for model_name in SELECTED_MODELS.values()],
["Xenova/ms-marco-MiniLM-L-6-v2", "Xenova/ms-marco-MiniLM-L-12-v2", "BAAI/bge-reranker-base"],
)
def test_batch_rerank(model_name):
is_ci = os.getenv("CI")
@@ -62,12 +48,6 @@ def test_batch_rerank(model_name):
documents = ["Paris is the capital of France.", "Berlin is the capital of Germany."] * 50
scores = np.array(list(model.rerank(query, documents, batch_size=10)))
pairs = [(query, doc) for doc in documents]
scores2 = np.array(list(model.rerank_pairs(pairs)))
assert np.allclose(
scores, scores2, atol=1e-5
), f"Model: {model_name}, Scores: {scores}, Scores2: {scores2}"
canonical_scores = np.tile(CANONICAL_SCORE_VALUES[model_name], 50)
assert scores.shape == canonical_scores.shape, f"Unexpected shape for model {model_name}"
@@ -93,27 +73,3 @@ def test_lazy_load(model_name):
if is_ci:
delete_model_cache(model.model._model_dir)
@pytest.mark.parametrize(
"model_name",
[model_name for model_name in SELECTED_MODELS.values()],
)
def test_rerank_pairs_parallel(model_name):
is_ci = os.getenv("CI")
model = TextCrossEncoder(model_name=model_name)
query = "What is the capital of France?"
documents = ["Paris is the capital of France.", "Berlin is the capital of Germany."] * 10
pairs = [(query, doc) for doc in documents]
scores_parallel = np.array(list(model.rerank_pairs(pairs, parallel=2, batch_size=10)))
scores_sequential = np.array(list(model.rerank_pairs(pairs, batch_size=10)))
assert np.allclose(
scores_parallel, scores_sequential, atol=1e-5
), f"Model: {model_name}, Scores (Parallel): {scores_parallel}, Scores (Sequential): {scores_sequential}"
canonical_scores = CANONICAL_SCORE_VALUES[model_name]
assert np.allclose(
scores_parallel[: len(canonical_scores)], canonical_scores, atol=1e-3
), f"Model: {model_name}, Scores (Parallel): {scores_parallel}, Expected: {canonical_scores}"
if is_ci:
delete_model_cache(model.model._model_dir)
-4
View File
@@ -41,8 +41,6 @@ CANONICAL_VECTOR_VALUES = {
"jinaai/jina-embeddings-v2-base-en": np.array([-0.0332, -0.0509, 0.0287, -0.0043, -0.0077]),
"jinaai/jina-embeddings-v2-base-de": np.array([-0.0085, 0.0417, 0.0342, 0.0309, -0.0149]),
"jinaai/jina-embeddings-v2-base-code": np.array([0.0145, -0.0164, 0.0136, -0.0170, 0.0734]),
"jinaai/jina-embeddings-v2-base-zh": np.array([0.0381, 0.0286, -0.0231, 0.0052, -0.0151]),
"jinaai/jina-embeddings-v2-base-es": np.array([-0.0108, -0.0092, -0.0373, 0.0171, -0.0301]),
"nomic-ai/nomic-embed-text-v1": np.array([0.3708, 0.2031, -0.3406, -0.2114, -0.3230]),
"nomic-ai/nomic-embed-text-v1.5": np.array(
[-0.15407836, -0.03053198, -3.9138033, 0.1910364, 0.13224715]
@@ -64,8 +62,6 @@ CANONICAL_VECTOR_VALUES = {
),
"snowflake/snowflake-arctic-embed-l": np.array([0.0189, -0.0673, 0.0183, 0.0124, 0.0146]),
"Qdrant/clip-ViT-B-32-text": np.array([0.0083, 0.0103, -0.0138, 0.0199, -0.0069]),
"thenlper/gte-base": np.array([0.0038, 0.0355, 0.0181, 0.0092, 0.0654]),
"jinaai/jina-clip-v1": np.array([-0.0862, -0.0101, -0.0056, 0.0375, -0.0472]),
}
+1 -9
View File
@@ -1,5 +1,4 @@
import shutil
import traceback
from pathlib import Path
from typing import Union
@@ -15,12 +14,6 @@ def delete_model_cache(model_dir: Union[str, Path]) -> None:
Args:
model_dir (Union[str, Path]): The path to the model cache directory.
"""
def on_error(func, path, exc_info):
print("Failed to remove: ", path)
print("Exception: ", exc_info)
traceback.print_exception(*exc_info)
if isinstance(model_dir, str):
model_dir = Path(model_dir)
@@ -28,5 +21,4 @@ def delete_model_cache(model_dir: Union[str, Path]) -> None:
model_dir = model_dir.parent.parent
if model_dir.exists():
# todo: PermissionDenied is raised on blobs removal in Windows, with blobs > 2GB
shutil.rmtree(model_dir, onerror=on_error)
shutil.rmtree(model_dir)