mirror of
https://github.com/qdrant/fastembed.git
synced 2026-09-22 05:57:51 -05:00
Compare commits
43
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f6d70544e8 | ||
|
|
60a6e8c090 | ||
|
|
3a906a3af6 | ||
|
|
1376c631c2 | ||
|
|
44e4867896 | ||
|
|
df7846596d | ||
|
|
23ec0994ec | ||
|
|
2701880877 | ||
|
|
f58310ceca | ||
|
|
73e121fa93 | ||
|
|
fcdabf3230 | ||
|
|
64561fdeb4 | ||
|
|
03f7111a32 | ||
|
|
0cf7203595 | ||
|
|
51c6bb14e7 | ||
|
|
3f3d90bf2f | ||
|
|
bf4ef9d513 | ||
|
|
48c59dc7b3 | ||
|
|
7099dec962 | ||
|
|
a2660f8a3a | ||
|
|
e7d9abaee1 | ||
|
|
01097708fe | ||
|
|
d725974fc4 | ||
|
|
14c067e15e | ||
|
|
c8fff66b18 | ||
|
|
85aaae4c08 | ||
|
|
dfd25d41c9 | ||
|
|
316c33634b | ||
|
|
490c340a76 | ||
|
|
ec8f06978f | ||
|
|
cbe00107ec | ||
|
|
99164c9050 | ||
|
|
6ecab6d40d | ||
|
|
d8c592032b | ||
|
|
da603b8b7d | ||
|
|
562d604375 | ||
|
|
8b1a98a6a3 | ||
|
|
3c9b147e0a | ||
|
|
6abd415f4a | ||
|
|
8184acbb39 | ||
|
|
47cf7f9f92 | ||
|
|
f7896c81f3 | ||
|
|
4a59d09248 |
@@ -16,7 +16,7 @@ body:
|
||||
validations:
|
||||
required: true
|
||||
- type: textarea
|
||||
id: Python version
|
||||
id: python-version
|
||||
attributes:
|
||||
label: What Python version are you on? e.g. python --version
|
||||
description: Also tell us, what package manager are you using e.g. conda, pip, poetry?
|
||||
@@ -29,7 +29,8 @@ body:
|
||||
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.
|
||||
options:
|
||||
- 0.2.6 (Latest)
|
||||
- 0.2.7 (Latest)
|
||||
- 0.2.6
|
||||
- 0.2.5
|
||||
- 0.2.4
|
||||
- 0.2.3
|
||||
|
||||
@@ -15,7 +15,6 @@ on:
|
||||
tags:
|
||||
- 'v*' # Push events to every version tag
|
||||
|
||||
|
||||
jobs:
|
||||
deploy:
|
||||
|
||||
|
||||
@@ -1,8 +1,9 @@
|
||||
name: Tests
|
||||
run-name: Tests (gpu)
|
||||
|
||||
on:
|
||||
push:
|
||||
branches: [ master, main ]
|
||||
branches: [ master, main, gpu ]
|
||||
schedule:
|
||||
- cron: 0 0 * * *
|
||||
pull_request:
|
||||
@@ -23,8 +24,6 @@ jobs:
|
||||
- '3.12.x'
|
||||
os:
|
||||
- ubuntu-latest
|
||||
- macos-latest
|
||||
- windows-latest
|
||||
|
||||
runs-on: ${{ matrix.os }}
|
||||
|
||||
|
||||
@@ -14,10 +14,14 @@ The default text embedding (`TextEmbedding`) model is Flag Embedding, presented
|
||||
|
||||
## 🚀 Installation
|
||||
|
||||
To install the FastEmbed library, pip works:
|
||||
To install the FastEmbed library, pip works best. You can install it with or without GPU support:
|
||||
|
||||
```bash
|
||||
pip install fastembed
|
||||
|
||||
# or with GPU support
|
||||
|
||||
pip install fastembed-gpu
|
||||
```
|
||||
|
||||
## 📖 Quickstart
|
||||
@@ -42,6 +46,120 @@ embeddings_list = list(embedding_model.embed(documents))
|
||||
len(embeddings_list[0]) # Vector of 384 dimensions
|
||||
```
|
||||
|
||||
Fastembed supports a variety of models for different tasks and modalities.
|
||||
The list of all the available models can be found [here](https://qdrant.github.io/fastembed/examples/Supported_Models/)
|
||||
### 🎒 Dense text embeddings
|
||||
|
||||
```python
|
||||
from fastembed import TextEmbedding
|
||||
|
||||
model = TextEmbedding(model_name="BAAI/bge-small-en-v1.5")
|
||||
embeddings = list(embedding_model.embed(documents))
|
||||
|
||||
# [
|
||||
# array([-0.1115, 0.0097, 0.0052, 0.0195, ...], dtype=float32),
|
||||
# array([-0.1019, 0.0635, -0.0332, 0.0522, ...], dtype=float32)
|
||||
# ]
|
||||
|
||||
```
|
||||
|
||||
|
||||
|
||||
### 🔱 Sparse text embeddings
|
||||
|
||||
* SPLADE++
|
||||
|
||||
```python
|
||||
from fastembed import SparseTextEmbedding
|
||||
|
||||
model = SparseTextEmbedding(model_name="prithivida/Splade_PP_en_v1")
|
||||
embeddings = list(embedding_model.embed(documents))
|
||||
|
||||
# [
|
||||
# SparseEmbedding(indices=[ 17, 123, 919, ... ], values=[0.71, 0.22, 0.39, ...]),
|
||||
# SparseEmbedding(indices=[ 38, 12, 91, ... ], values=[0.11, 0.22, 0.39, ...])
|
||||
# ]
|
||||
```
|
||||
|
||||
<!--
|
||||
* BM42 - ([link](ToDo))
|
||||
|
||||
```
|
||||
from fastembed import SparseTextEmbedding
|
||||
|
||||
model = SparseTextEmbedding(model_name="Qdrant/bm42-all-minilm-l6-v2-attentions")
|
||||
embeddings = list(embedding_model.embed(documents))
|
||||
|
||||
# [
|
||||
# SparseEmbedding(indices=[ 17, 123, 919, ... ], values=[0.71, 0.22, 0.39, ...]),
|
||||
# SparseEmbedding(indices=[ 38, 12, 91, ... ], values=[0.11, 0.22, 0.39, ...])
|
||||
# ]
|
||||
```
|
||||
-->
|
||||
|
||||
### 🦥 Late interaction models (aka ColBERT)
|
||||
|
||||
|
||||
```python
|
||||
from fastembed import LateInteractionTextEmbedding
|
||||
|
||||
model = LateInteractionTextEmbedding(model_name="colbert-ir/colbertv2.0")
|
||||
embeddings = list(embedding_model.embed(documents))
|
||||
|
||||
# [
|
||||
# array([
|
||||
# [-0.1115, 0.0097, 0.0052, 0.0195, ...],
|
||||
# [-0.1019, 0.0635, -0.0332, 0.0522, ...],
|
||||
# ]),
|
||||
# array([
|
||||
# [-0.9019, 0.0335, -0.0032, 0.0991, ...],
|
||||
# [-0.2115, 0.8097, 0.1052, 0.0195, ...],
|
||||
# ]),
|
||||
# ]
|
||||
```
|
||||
|
||||
### 🖼️ Image embeddings
|
||||
|
||||
```python
|
||||
from fastembed import ImageEmbedding
|
||||
|
||||
images = [
|
||||
"./path/to/image1.jpg",
|
||||
"./path/to/image2.jpg",
|
||||
]
|
||||
|
||||
model = ImageEmbedding(model_name="Qdrant/clip-ViT-B-32-vision")
|
||||
embeddings = list(embedding_model.embed(images))
|
||||
|
||||
# [
|
||||
# array([-0.1115, 0.0097, 0.0052, 0.0195, ...], dtype=float32),
|
||||
# array([-0.1019, 0.0635, -0.0332, 0.0522, ...], dtype=float32)
|
||||
# ]
|
||||
```
|
||||
|
||||
|
||||
## ⚡️ FastEmbed on a GPU
|
||||
|
||||
FastEmbed supports running on GPU devices.
|
||||
It requires installation of the `fastembed-gpu` package.
|
||||
|
||||
```bash
|
||||
pip install fastembed-gpu
|
||||
```
|
||||
|
||||
Check our [example](https://qdrant.github.io/fastembed/examples/FastEmbed_GPU/) for the detailed instructions and CUDA 12.x support.
|
||||
|
||||
```python
|
||||
from fastembed import TextEmbedding
|
||||
|
||||
embedding_model = TextEmbedding(
|
||||
model_name="BAAI/bge-small-en-v1.5",
|
||||
providers=["CUDAExecutionProvider"]
|
||||
)
|
||||
print("The model BAAI/bge-small-en-v1.5 is ready to use on a GPU.")
|
||||
|
||||
```
|
||||
|
||||
## Usage with Qdrant
|
||||
|
||||
Installation with Qdrant Client in Python:
|
||||
@@ -50,7 +168,13 @@ Installation with Qdrant Client in Python:
|
||||
pip install qdrant-client[fastembed]
|
||||
```
|
||||
|
||||
You might have to use ```pip install 'qdrant-client[fastembed]'``` on zsh.
|
||||
or
|
||||
|
||||
```bash
|
||||
pip install qdrant-client[fastembed-gpu]
|
||||
```
|
||||
|
||||
You might have to use quotes ```pip install 'qdrant-client[fastembed]'``` on zsh.
|
||||
|
||||
```python
|
||||
from qdrant_client import QdrantClient
|
||||
@@ -85,8 +209,4 @@ search_result = client.query(
|
||||
query_text="This is a query document"
|
||||
)
|
||||
print(search_result)
|
||||
```
|
||||
|
||||
#### Similar Work
|
||||
|
||||
Ilyas M. wrote about using [FlagEmbeddings with Optimum](https://twitter.com/IlysMoutawwakil/status/1705215192425288017) over CUDA.
|
||||
```
|
||||
+41
@@ -0,0 +1,41 @@
|
||||
# Releasing FastEmbed
|
||||
|
||||
This is a guide how to release `fastembed` and `fastembed-gpu` packages.
|
||||
|
||||
## How to
|
||||
|
||||
1. Accumulate changes in the `main` branch.
|
||||
2. Bump the version in `pyproject.toml`
|
||||
|
||||
3. Rebase the `gpu` branch on `main` and resolve conflicts if occurred:
|
||||
|
||||
```bash
|
||||
git checkout gpu
|
||||
git rebase main
|
||||
git push origin gpu
|
||||
```
|
||||
|
||||
4. Draft release notes
|
||||
5. Checkout to `main` and create a tag, e.g.:
|
||||
|
||||
```bash
|
||||
git checkout main
|
||||
git tag -a v0.1.0 -m "Release v0.1.0"
|
||||
```
|
||||
|
||||
6. Checkout `gpu` and create a tag, e.g.:
|
||||
|
||||
```bash
|
||||
git checkout gpu
|
||||
git tag -a v0.1.0-gpu -m "Release v0.1.0"
|
||||
```
|
||||
|
||||
7. Push tags:
|
||||
|
||||
```bash
|
||||
git push --tags
|
||||
```
|
||||
|
||||
8. Verify that both packages have been published successfully on PyPI. Try installing them and verify imports.
|
||||
9. Create a release on GitHub with the written release notes.
|
||||
|
||||
+11
-11
@@ -11,7 +11,7 @@
|
||||
"\n",
|
||||
"## Quick Start\n",
|
||||
"\n",
|
||||
"The fastembed package is designed to be easy to use. We'll be using `TextEmbedding` class. It takes a list of strings as input and returns an generator of vectors. If you're seeing generators for the first time, don't worry, you can convert it to a list using `list()`.\n",
|
||||
"The fastembed package is designed to be easy to use. We'll be using `TextEmbedding` class. It takes a list of strings as input and returns a generator of vectors.\n",
|
||||
"\n",
|
||||
"> 💡 You can learn more about generators from [Python Wiki](https://wiki.python.org/moin/Generators)"
|
||||
]
|
||||
@@ -23,7 +23,7 @@
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"!pip install -Uqq fastembed # Install fastembed"
|
||||
"!pip install -Uqq fastembed"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -65,10 +65,13 @@
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"from fastembed import TextEmbedding\n",
|
||||
"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",
|
||||
" \"This is built to be faster and lighter than other embedding libraries e.g. Transformers, Sentence-Transformers, etc.\",\n",
|
||||
@@ -79,9 +82,8 @@
|
||||
"embedding_model = TextEmbedding()\n",
|
||||
"print(\"The model BAAI/bge-small-en-v1.5 is ready to use.\")\n",
|
||||
"\n",
|
||||
"embeddings_generator = embedding_model.embed(documents) # reminder this is a generator\n",
|
||||
"embeddings_generator = embedding_model.embed(documents)\n",
|
||||
"embeddings_list = list(embeddings_generator)\n",
|
||||
"# you can also convert the generator to a list, and that to a numpy array\n",
|
||||
"len(embeddings_list[0]) # Vector of 384 dimensions"
|
||||
]
|
||||
},
|
||||
@@ -113,7 +115,7 @@
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"embeddings_generator = embedding_model.embed(documents) # reminder this is a generator\n",
|
||||
"embeddings_generator = embedding_model.embed(documents)\n",
|
||||
"\n",
|
||||
"for doc, vector in zip(documents, embeddings_generator):\n",
|
||||
" print(\"Document:\", doc)\n",
|
||||
@@ -138,9 +140,7 @@
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"embeddings_list = np.array(\n",
|
||||
" list(embedding_model.embed(documents))\n",
|
||||
") # you can also convert the generator to a list, and that to a numpy array\n",
|
||||
"embeddings_list = np.array(list(embedding_model.embed(documents)))\n",
|
||||
"embeddings_list.shape"
|
||||
]
|
||||
},
|
||||
@@ -185,7 +185,7 @@
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"multilingual_large_model = TextEmbedding(\"intfloat/multilingual-e5-large\") # This can take a few minutes to download"
|
||||
"multilingual_large_model = TextEmbedding(\"intfloat/multilingual-e5-large\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
|
||||
@@ -0,0 +1,410 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "d14d29ebd3592ecb",
|
||||
"metadata": {
|
||||
"collapsed": false
|
||||
},
|
||||
"source": [
|
||||
"# Late Interaction Text Embedding Models\n",
|
||||
"\n",
|
||||
"As of version 0.3.0 FastEmbed supports Late Interaction Text Embedding Models and currently available with one of the most popular embedding model of the family - ColBERT.\n",
|
||||
"\n",
|
||||
"## What is a Late Interaction Text Embedding Model?\n",
|
||||
"\n",
|
||||
"Late Interaction Text Embedding Model is a kind of information retrieval model which performs query and documents interactions at the scoring stage.\n",
|
||||
"In order to better understand it, we can compare it to the models without interaction. \n",
|
||||
"For instance, if you take a sentence-transformer model, compute embeddings for your documents, compute embeddings for your queries, and just compare them by cosine similarity, then you're retrieving points without interaction.\n",
|
||||
"\n",
|
||||
"It is a pretty much easy and straightforward approach, however we might be sacrificing some precision due to its simplicity. It is caused by several facts: \n",
|
||||
"- there is no interaction between queries and documents at the early stage (embedding generation) nor at the late stage (during scoring). \n",
|
||||
"- we are trying to encapsulate all the document information in only one pooled embedding, and obviously, some information might be lost.\n",
|
||||
"\n",
|
||||
"Late Interaction Text Embedding models are trying to address it by computing embeddings for each token in queries and documents, and then finding the most similar ones via model specific operation, e.g. ColBERT (Contextual Late Interaction over BERT) uses MaxSim operation.\n",
|
||||
"With this approach we can have not only a better representation of the documents, but also make queries and documents more aware one of another.\n",
|
||||
"\n",
|
||||
"For more information on ColBERT and MaxSim operation, you can check out [this blogpost](https://jina.ai/news/what-is-colbert-and-late-interaction-and-why-they-matter-in-search/) by Jina AI.\n",
|
||||
"\n",
|
||||
"## ColBERT in FastEmbed\n",
|
||||
"\n",
|
||||
"FastEmbed provides a simple way to use ColBERT model, similar to the ones it has with `TextEmbedding`.\n",
|
||||
" "
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"id": "7f1053b17c810be5",
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-06-03T17:20:26.927643Z",
|
||||
"start_time": "2024-06-03T17:20:25.128994Z"
|
||||
},
|
||||
"collapsed": false
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"/Users/joein/work/qdrant/fastembed/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"
|
||||
]
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"text/plain": "[{'model': 'colbert-ir/colbertv2.0',\n 'dim': 128,\n 'description': 'Late interaction model',\n 'size_in_GB': 0.44,\n 'sources': {'hf': 'colbert-ir/colbertv2.0'},\n 'model_file': 'model.onnx'}]"
|
||||
},
|
||||
"execution_count": 1,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"from fastembed import LateInteractionTextEmbedding\n",
|
||||
"\n",
|
||||
"LateInteractionTextEmbedding.list_supported_models()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"id": "c2c15893df422631",
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-06-03T17:23:35.764183Z",
|
||||
"start_time": "2024-06-03T17:23:21.630277Z"
|
||||
},
|
||||
"collapsed": false
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Fetching 5 files: 0%| | 0/5 [00:00<?, ?it/s]\n",
|
||||
"config.json: 100%|██████████| 743/743 [00:00<00:00, 4.56MB/s]\n",
|
||||
"\n",
|
||||
"tokenizer_config.json: 100%|██████████| 405/405 [00:00<00:00, 3.34MB/s]\n",
|
||||
"Fetching 5 files: 20%|██ | 1/5 [00:00<00:01, 3.64it/s]\n",
|
||||
"tokenizer.json: 0%| | 0.00/466k [00:00<?, ?B/s]\u001b[A\n",
|
||||
"\n",
|
||||
"special_tokens_map.json: 100%|██████████| 112/112 [00:00<00:00, 727kB/s]\n",
|
||||
"\n",
|
||||
"tokenizer.json: 100%|██████████| 466k/466k [00:00<00:00, 1.48MB/s]\u001b[A\n",
|
||||
"\n",
|
||||
"model.onnx: 0%| | 0.00/436M [00:00<?, ?B/s]\u001b[A\n",
|
||||
"model.onnx: 2%|▏ | 10.5M/436M [00:00<00:34, 12.2MB/s]\u001b[A\n",
|
||||
"model.onnx: 5%|▍ | 21.0M/436M [00:01<00:20, 20.3MB/s]\u001b[A\n",
|
||||
"model.onnx: 7%|▋ | 31.5M/436M [00:01<00:15, 25.7MB/s]\u001b[A\n",
|
||||
"model.onnx: 10%|▉ | 41.9M/436M [00:01<00:13, 29.4MB/s]\u001b[A\n",
|
||||
"model.onnx: 12%|█▏ | 52.4M/436M [00:01<00:12, 31.9MB/s]\u001b[A\n",
|
||||
"model.onnx: 14%|█▍ | 62.9M/436M [00:02<00:11, 33.7MB/s]\u001b[A\n",
|
||||
"model.onnx: 17%|█▋ | 73.4M/436M [00:02<00:10, 34.9MB/s]\u001b[A\n",
|
||||
"model.onnx: 19%|█▉ | 83.9M/436M [00:02<00:09, 35.5MB/s]\u001b[A\n",
|
||||
"model.onnx: 22%|██▏ | 94.4M/436M [00:03<00:09, 36.1MB/s]\u001b[A\n",
|
||||
"model.onnx: 24%|██▍ | 105M/436M [00:03<00:09, 36.6MB/s] \u001b[A\n",
|
||||
"model.onnx: 26%|██▋ | 115M/436M [00:03<00:08, 36.9MB/s]\u001b[A\n",
|
||||
"model.onnx: 29%|██▉ | 126M/436M [00:03<00:08, 37.1MB/s]\u001b[A\n",
|
||||
"model.onnx: 31%|███▏ | 136M/436M [00:04<00:08, 37.3MB/s]\u001b[A\n",
|
||||
"model.onnx: 34%|███▎ | 147M/436M [00:04<00:07, 37.4MB/s]\u001b[A\n",
|
||||
"model.onnx: 36%|███▌ | 157M/436M [00:04<00:07, 37.4MB/s]\u001b[A\n",
|
||||
"model.onnx: 38%|███▊ | 168M/436M [00:05<00:07, 37.5MB/s]\u001b[A\n",
|
||||
"model.onnx: 41%|████ | 178M/436M [00:05<00:06, 37.6MB/s]\u001b[A\n",
|
||||
"model.onnx: 43%|████▎ | 189M/436M [00:05<00:06, 37.6MB/s]\u001b[A\n",
|
||||
"model.onnx: 46%|████▌ | 199M/436M [00:05<00:06, 37.6MB/s]\u001b[A\n",
|
||||
"model.onnx: 48%|████▊ | 210M/436M [00:06<00:06, 37.5MB/s]\u001b[A\n",
|
||||
"model.onnx: 50%|█████ | 220M/436M [00:06<00:05, 37.5MB/s]\u001b[A\n",
|
||||
"model.onnx: 53%|█████▎ | 231M/436M [00:06<00:05, 37.6MB/s]\u001b[A\n",
|
||||
"model.onnx: 55%|█████▌ | 241M/436M [00:06<00:05, 37.6MB/s]\u001b[A\n",
|
||||
"model.onnx: 58%|█████▊ | 252M/436M [00:07<00:04, 37.6MB/s]\u001b[A\n",
|
||||
"model.onnx: 60%|██████ | 262M/436M [00:07<00:04, 37.7MB/s]\u001b[A\n",
|
||||
"model.onnx: 63%|██████▎ | 273M/436M [00:07<00:04, 37.7MB/s]\u001b[A\n",
|
||||
"model.onnx: 65%|██████▍ | 283M/436M [00:08<00:04, 36.0MB/s]\u001b[A\n",
|
||||
"model.onnx: 67%|██████▋ | 294M/436M [00:08<00:03, 36.4MB/s]\u001b[A\n",
|
||||
"model.onnx: 70%|██████▉ | 304M/436M [00:08<00:03, 36.8MB/s]\u001b[A\n",
|
||||
"model.onnx: 72%|███████▏ | 315M/436M [00:08<00:03, 37.0MB/s]\u001b[A\n",
|
||||
"model.onnx: 75%|███████▍ | 325M/436M [00:09<00:02, 37.3MB/s]\u001b[A\n",
|
||||
"model.onnx: 77%|███████▋ | 336M/436M [00:09<00:03, 30.8MB/s]\u001b[A\n",
|
||||
"model.onnx: 79%|███████▉ | 346M/436M [00:10<00:02, 32.6MB/s]\u001b[A\n",
|
||||
"model.onnx: 82%|████████▏ | 357M/436M [00:10<00:02, 33.9MB/s]\u001b[A\n",
|
||||
"model.onnx: 84%|████████▍ | 367M/436M [00:10<00:01, 34.8MB/s]\u001b[A\n",
|
||||
"model.onnx: 87%|████████▋ | 377M/436M [00:10<00:01, 35.7MB/s]\u001b[A\n",
|
||||
"model.onnx: 89%|████████▉ | 388M/436M [00:11<00:01, 36.2MB/s]\u001b[A\n",
|
||||
"model.onnx: 91%|█████████▏| 398M/436M [00:11<00:01, 36.6MB/s]\u001b[A\n",
|
||||
"model.onnx: 94%|█████████▍| 409M/436M [00:11<00:00, 36.9MB/s]\u001b[A\n",
|
||||
"model.onnx: 96%|█████████▌| 419M/436M [00:11<00:00, 37.1MB/s]\u001b[A\n",
|
||||
"model.onnx: 99%|█████████▊| 430M/436M [00:12<00:00, 37.3MB/s]\u001b[A\n",
|
||||
"model.onnx: 100%|██████████| 436M/436M [00:12<00:00, 35.1MB/s]\u001b[A\n",
|
||||
"Fetching 5 files: 100%|██████████| 5/5 [00:13<00:00, 2.68s/it]\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"embedding_model = LateInteractionTextEmbedding(\"colbert-ir/colbertv2.0\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 16,
|
||||
"id": "e560b5fa7d63bea3",
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-06-03T17:39:33.400876Z",
|
||||
"start_time": "2024-06-03T17:39:33.397431Z"
|
||||
},
|
||||
"collapsed": false
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"documents = [\n",
|
||||
" \"ColBERT is a late interaction text embedding model, however, there are also other models such as TwinBERT.\",\n",
|
||||
" \"On the contrary to the late interaction models, the early interaction models contains interaction steps at embedding generation process\",\n",
|
||||
"]\n",
|
||||
"queries = [\n",
|
||||
" \"Are there any other late interaction text embedding models except ColBERT?\",\n",
|
||||
" \"What is the difference between late interaction and early interaction text embedding models?\",\n",
|
||||
"]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "347ad924a3449743",
|
||||
"metadata": {
|
||||
"collapsed": false
|
||||
},
|
||||
"source": [
|
||||
"*NOTE*: ColBERT computes query and documents embeddings differently, make sure to use the corresponding methods."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 17,
|
||||
"id": "496fbf51e4eaaae",
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-06-03T17:39:34.379885Z",
|
||||
"start_time": "2024-06-03T17:39:34.316257Z"
|
||||
},
|
||||
"collapsed": false
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"document_embeddings = list(\n",
|
||||
" embedding_model.embed(documents)\n",
|
||||
") # embed and qury_embed return generators,\n",
|
||||
"# which we need to evaluate by writing them to a list\n",
|
||||
"query_embeddings = list(embedding_model.query_embed(queries))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 18,
|
||||
"id": "50595bb0498f0c7c",
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-06-03T17:39:34.793528Z",
|
||||
"start_time": "2024-06-03T17:39:34.788545Z"
|
||||
},
|
||||
"collapsed": false
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": "((26, 128), (32, 128))"
|
||||
},
|
||||
"execution_count": 18,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"document_embeddings[0].shape, query_embeddings[0].shape"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "13e43f2c24a7d5fc",
|
||||
"metadata": {
|
||||
"collapsed": false
|
||||
},
|
||||
"source": [
|
||||
"Don't worry about query embeddings having the bigger shape in this case. \n",
|
||||
"ColBERT authors recommend to pad queries with [MASK] tokens to 32 tokens.\n",
|
||||
"They also recommends to truncate queries to 32 tokens, however we don't do that in FastEmbed, so you can put some straight into the queries."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "bb1a4011effd3699",
|
||||
"metadata": {
|
||||
"collapsed": false
|
||||
},
|
||||
"source": [
|
||||
"## MaxSim operator"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "e9ea4cf82521f2de",
|
||||
"metadata": {
|
||||
"collapsed": false
|
||||
},
|
||||
"source": [
|
||||
"Qdrant will support ColBERT as of the next version (v1.10), however, at the moment, you can compute embedding similarities manually. "
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 19,
|
||||
"id": "f84392f63d2c6076",
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-06-03T17:39:36.431622Z",
|
||||
"start_time": "2024-06-03T17:39:36.427363Z"
|
||||
},
|
||||
"collapsed": false
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"def compute_relevance_scores(query_embedding: np.array, document_embeddings: np.array, k: int):\n",
|
||||
" \"\"\"\n",
|
||||
" Compute relevance scores for top-k documents given a query.\n",
|
||||
"\n",
|
||||
" :param query_embedding: Numpy array representing the query embedding, shape: [num_query_terms, embedding_dim]\n",
|
||||
" :param document_embeddings: Numpy array representing embeddings for documents, shape: [num_documents, max_doc_length, embedding_dim]\n",
|
||||
" :param k: Number of top documents to return\n",
|
||||
" :return: Indices of the top-k documents based on their relevance scores\n",
|
||||
" \"\"\"\n",
|
||||
" # Compute batch dot-product of query_embedding and document_embeddings\n",
|
||||
" # Resulting shape: [num_documents, num_query_terms, max_doc_length]\n",
|
||||
" scores = np.matmul(query_embedding, document_embeddings.transpose(0, 2, 1))\n",
|
||||
"\n",
|
||||
" # Apply max-pooling across document terms (axis=2) to find the max similarity per query term\n",
|
||||
" # Shape after max-pool: [num_documents, num_query_terms]\n",
|
||||
" max_scores_per_query_term = np.max(scores, axis=2)\n",
|
||||
"\n",
|
||||
" # Sum the scores across query terms to get the total score for each document\n",
|
||||
" # Shape after sum: [num_documents]\n",
|
||||
" total_scores = np.sum(max_scores_per_query_term, axis=1)\n",
|
||||
"\n",
|
||||
" # Sort the documents based on their total scores and get the indices of the top-k documents\n",
|
||||
" sorted_indices = np.argsort(total_scores)[::-1][:k]\n",
|
||||
"\n",
|
||||
" return sorted_indices"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 20,
|
||||
"id": "c61d07bed7b60e35",
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-06-03T17:39:37.053383Z",
|
||||
"start_time": "2024-06-03T17:39:37.050926Z"
|
||||
},
|
||||
"collapsed": false
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Sorted document indices: [0 1]\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"sorted_indices = compute_relevance_scores(\n",
|
||||
" np.array(query_embeddings[0]), np.array(document_embeddings), k=3\n",
|
||||
")\n",
|
||||
"print(\"Sorted document indices:\", sorted_indices)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 22,
|
||||
"id": "b24df2569970d9e8",
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-06-03T17:40:52.276846Z",
|
||||
"start_time": "2024-06-03T17:40:52.273789Z"
|
||||
},
|
||||
"collapsed": false
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Query: Are there any other late interaction text embedding models except ColBERT?\n",
|
||||
"Document: ColBERT is a late interaction text embedding model, however, there are also other models such as TwinBERT.\n",
|
||||
"Document: On the contrary to the late interaction models, the early interaction models contains interaction steps at embedding generation process\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"print(f\"Query: {queries[0]}\")\n",
|
||||
"for index in sorted_indices:\n",
|
||||
" print(f\"Document: {documents[index]}\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "6de537c37aff3927",
|
||||
"metadata": {
|
||||
"collapsed": false
|
||||
},
|
||||
"source": [
|
||||
"## Use-case recommendation"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "37e3525d3259cd2b",
|
||||
"metadata": {
|
||||
"collapsed": false
|
||||
},
|
||||
"source": [
|
||||
"Despite ColBERT allows to compute embeddings independently and spare some workload offline, it still computes more resources than no interaction models. Due to this, it might be more reasonable to use ColBERT not as a first-stage retriever, but as a re-ranker.\n",
|
||||
"\n",
|
||||
"The first-stage retriever would then be a no-interaction model, which e.g. retrieves first 100 or 500 examples, and leave the final ranking to the ColBERT model."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "cfa922793454b4ad",
|
||||
"metadata": {
|
||||
"collapsed": false
|
||||
},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 2
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython2",
|
||||
"version": "2.7.6"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,436 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "ntGNDuSCeAR2"
|
||||
},
|
||||
"source": [
|
||||
"# FastEmbed on GPU\n",
|
||||
"\n",
|
||||
"As of version 0.2.7 FastEmbed supports GPU acceleration.\n",
|
||||
"\n",
|
||||
"This notebook covers the installation process and usage of fastembed on GPU.\n",
|
||||
"\n",
|
||||
"## Installation\n",
|
||||
"\n",
|
||||
"Fastembed depends on `onnxruntime` and inherits its scheme of GPU support.\n",
|
||||
"\n",
|
||||
"In order to use GPU with onnx models, you would need to have `onnxruntime-gpu` package, which substitutes all the `onnxruntime` functionality.\n",
|
||||
"Fastembed mimics this behavior and requires `fastembed-gpu` package to be installed."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "GK2XADwUeEK7"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"!pip install fastembed-gpu"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "3aiGPqjCeGzo"
|
||||
},
|
||||
"source": [
|
||||
"**NOTE**: `onnxruntime-gpu` and `onnxruntime` can't be installed in the same environment. If you have `onnxruntime` installed, you would need to uninstall it before installing `onnxruntime-gpu`. Same is true for `fastembed` and `fastembed-gpu`.\n",
|
||||
"\n",
|
||||
"### CUDA 12.x support\n",
|
||||
"\n",
|
||||
"By default `onnxruntime-gpu` is shipped with CUDA 11.8 support.\n",
|
||||
"CUDA 12.x support requires installation of `onnxruntime-gpu` with providing of a direct url:"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"colab": {
|
||||
"base_uri": "https://localhost:8080/"
|
||||
},
|
||||
"id": "OoSfWFFZeJ5t",
|
||||
"outputId": "417b9332-6a7b-4000-c74b-4ed2b5b76590"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"!pip install onnxruntime-gpu -i https://aiinfra.pkgs.visualstudio.com/PublicPackages/_packaging/onnxruntime-cuda-12/pypi/simple/ -qq\n",
|
||||
"!pip install fastembed-gpu -qqq"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "3xx3r-9jgAMi"
|
||||
},
|
||||
"source": [
|
||||
"You can check your CUDA version using such commands as `nvidia-smi` or `nvcc --version`\n",
|
||||
"\n",
|
||||
"Google Colab notebooks have CUDA 12.x."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "Igv5RXhSeO68"
|
||||
},
|
||||
"source": [
|
||||
"### CUDA drivers\n",
|
||||
"\n",
|
||||
"FastEmbed does not include CUDA drivers and CuDNN libraries.\n",
|
||||
"You would need to take care of the environment setup on your own.\n",
|
||||
"Dependencies required for the chosen onnxruntime version can be found [here](https://onnxruntime.ai/docs/execution-providers/CUDA-ExecutionProvider.html#requirements)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Usage"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"metadata": {
|
||||
"colab": {
|
||||
"base_uri": "https://localhost:8080/",
|
||||
"height": 334,
|
||||
"referenced_widgets": [
|
||||
"aacf08a7aa444b64a2efad1967d28a53",
|
||||
"5606aa785de74d65a9928b31c0be8a53",
|
||||
"d4ec9d3b74ec4412894da2161ed2bddf",
|
||||
"8edd544c3e074ec1813e5b9d1aef43d9",
|
||||
"9898890f8a75468ea20e3ce319d0b6e2",
|
||||
"da3b18abb16241a0a7191ee9afcb0510",
|
||||
"258a619168824253a6a329efdc51ebe6",
|
||||
"53c7cdc967d24faba0b5c659c94c50b8",
|
||||
"e9348d8be28d408e8e760c71b21ab294",
|
||||
"0ba06e0816714f2fbdec8260f160abc0",
|
||||
"0b96563334964d449dd34f35b6b3e715",
|
||||
"11c2eec490e8479b944eec7f30cb1ca2",
|
||||
"91463da0d1c5466795e06ab586002259",
|
||||
"30f4f7833406474f89ef0700b00a33aa",
|
||||
"4302c304ec6a4b5985797e300bd7e353",
|
||||
"2605640c7b824ed7aa137d404e14b774",
|
||||
"b02efe3a33d04f06aa8938719ab35671",
|
||||
"50408e5d052343b1a1b44a0fae0f801d",
|
||||
"1be01c95d9e84f8ea88367c987a72fdc",
|
||||
"a109c13bc93a449186424542dc330be8",
|
||||
"4adce304ce1947b5a01dde10bbb3bb8c",
|
||||
"a761366a37e44837a25e0f25b18efed2",
|
||||
"94512b9055e546389471197b76ad5449",
|
||||
"072dca00bd7b4918a178f90ccabf698a",
|
||||
"48a856c59ef74cc3834521b1bf616541",
|
||||
"c020c503aeaa464cad643ade5ee3ae24",
|
||||
"a4e7e40c0bbd4f878c20a9f65fe3a048",
|
||||
"e8c0a1c339fd47668d944a9defad79d4",
|
||||
"cd782d35c6bd40c0a60d57b1828a7251",
|
||||
"04f638ab08da4d20928644c4ba03f8ef",
|
||||
"17f20477fc79475f97adf1c1f64a4192",
|
||||
"96f7b5a2e224462e9fcffd03f906a593",
|
||||
"755cd32d9fc9407c80a160f45c802d1e",
|
||||
"a886258e7cd14c048b58391d7b772901",
|
||||
"bc3e48f826a74840867a6209e622b75e",
|
||||
"125b2ac0f78043bba7eca53474ca44c4",
|
||||
"82f186d1ffb4435d94a6c7e9025242ef",
|
||||
"77000333e5ca4094be291ad82d4a627a",
|
||||
"7fe64fb53055431488d002c76c8e331e",
|
||||
"7de59ae9919f4a5bb2b6e601a3c02412",
|
||||
"97a69423a6644eab87fc636e182f23a4",
|
||||
"4df936d1065b41f4bf02ed394fdf7b7e",
|
||||
"3918bd1affa3454e8e9044a418a056ea",
|
||||
"163b27ae0bce41e5b48efcb4b3fd780d",
|
||||
"94631fd6e0744085bc79c3121de4a9f7",
|
||||
"31cd98d66bc54418b35e70fbbc0fa3c0",
|
||||
"6d21627a638b4ddca6fe7bfb80a621b5",
|
||||
"b37bed9dc4fe45c08b8397288fe5b1a9",
|
||||
"164fef95d1414177a40d563f5682f6a3",
|
||||
"1a9a0ea53448413a8e4b360b7bb69e26",
|
||||
"dd1a4483b4b045c6929e3d2cf1338f63",
|
||||
"496ddd8e05f949cd8cbba8e677f476ac",
|
||||
"2813be951d7f48b2aad1dd4a444ce3eb",
|
||||
"8e9a2c2dd21942edbdfecb3b7dffc70b",
|
||||
"08a10fe247f1425db044cfc13f2fb384",
|
||||
"b8786aded92d421592bc7623c5c7899e",
|
||||
"c91a20a9433d4016ba2db69fa50e0b4d",
|
||||
"e997820738594c6dadb061908d7afdc1",
|
||||
"a5fc751f81ae498f9aa55ece0e6853b2",
|
||||
"2aee4fc8cda64c5eb8722be81e48e0ca",
|
||||
"3a53e8624dff48b3959875ef58ee99ce",
|
||||
"50a70044f77542108fe188598e70797e",
|
||||
"13cf998b35ae4507a63e797f6fa3eada",
|
||||
"6209eb6a68cf4a378767ef34d0d9216d",
|
||||
"7395db766b944af9b41d6b56c9ada0b1",
|
||||
"42122c317ec648688f0164a1adb5df28"
|
||||
]
|
||||
},
|
||||
"id": "Ttf4YggPeQQK",
|
||||
"outputId": "aa75129d-9e2d-4c88-cf03-251dd43a11b1"
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"/usr/local/lib/python3.10/dist-packages/huggingface_hub/utils/_token.py:88: UserWarning: \n",
|
||||
"The secret `HF_TOKEN` does not exist in your Colab secrets.\n",
|
||||
"To authenticate with the Hugging Face Hub, create a token in your settings tab (https://huggingface.co/settings/tokens), set it as secret in your Google Colab and restart your session.\n",
|
||||
"You will be able to reuse this secret in all of your notebooks.\n",
|
||||
"Please note that authentication is recommended but still optional to access public models or datasets.\n",
|
||||
" warnings.warn(\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"application/vnd.jupyter.widget-view+json": {
|
||||
"model_id": "aacf08a7aa444b64a2efad1967d28a53",
|
||||
"version_major": 2,
|
||||
"version_minor": 0
|
||||
},
|
||||
"text/plain": [
|
||||
"Fetching 5 files: 0%| | 0/5 [00:00<?, ?it/s]"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"application/vnd.jupyter.widget-view+json": {
|
||||
"model_id": "11c2eec490e8479b944eec7f30cb1ca2",
|
||||
"version_major": 2,
|
||||
"version_minor": 0
|
||||
},
|
||||
"text/plain": [
|
||||
"tokenizer_config.json: 0%| | 0.00/1.24k [00:00<?, ?B/s]"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"application/vnd.jupyter.widget-view+json": {
|
||||
"model_id": "94512b9055e546389471197b76ad5449",
|
||||
"version_major": 2,
|
||||
"version_minor": 0
|
||||
},
|
||||
"text/plain": [
|
||||
"config.json: 0%| | 0.00/706 [00:00<?, ?B/s]"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"application/vnd.jupyter.widget-view+json": {
|
||||
"model_id": "a886258e7cd14c048b58391d7b772901",
|
||||
"version_major": 2,
|
||||
"version_minor": 0
|
||||
},
|
||||
"text/plain": [
|
||||
"special_tokens_map.json: 0%| | 0.00/695 [00:00<?, ?B/s]"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"application/vnd.jupyter.widget-view+json": {
|
||||
"model_id": "94631fd6e0744085bc79c3121de4a9f7",
|
||||
"version_major": 2,
|
||||
"version_minor": 0
|
||||
},
|
||||
"text/plain": [
|
||||
"tokenizer.json: 0%| | 0.00/711k [00:00<?, ?B/s]"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"application/vnd.jupyter.widget-view+json": {
|
||||
"model_id": "b8786aded92d421592bc7623c5c7899e",
|
||||
"version_major": 2,
|
||||
"version_minor": 0
|
||||
},
|
||||
"text/plain": [
|
||||
"model_optimized.onnx: 0%| | 0.00/66.5M [00:00<?, ?B/s]"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"['CUDAExecutionProvider', 'CPUExecutionProvider']"
|
||||
]
|
||||
},
|
||||
"execution_count": 2,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"from typing import List\n",
|
||||
"\n",
|
||||
"import numpy as np\n",
|
||||
"\n",
|
||||
"from fastembed import TextEmbedding\n",
|
||||
"\n",
|
||||
"embedding_model_gpu = TextEmbedding(\n",
|
||||
" model_name=\"BAAI/bge-small-en-v1.5\", providers=[\"CUDAExecutionProvider\"]\n",
|
||||
")\n",
|
||||
"embedding_model_gpu.model.model.get_providers()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 3,
|
||||
"metadata": {
|
||||
"id": "iPtoHf7GeV-i"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"documents: List[str] = list(np.repeat(\"Demonstrating GPU acceleration in fastembed\", 500))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 4,
|
||||
"metadata": {
|
||||
"colab": {
|
||||
"base_uri": "https://localhost:8080/"
|
||||
},
|
||||
"id": "islhyLf4ed-H",
|
||||
"outputId": "8c8ed09b-9eac-438f-97bc-578751975148"
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"43.4 ms ± 2.06 ms per loop (mean ± std. dev. of 7 runs, 1 loop each)\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"%%timeit\n",
|
||||
"list(embedding_model_gpu.embed(documents))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 5,
|
||||
"metadata": {
|
||||
"colab": {
|
||||
"base_uri": "https://localhost:8080/",
|
||||
"height": 67,
|
||||
"referenced_widgets": [
|
||||
"9c306ce5188c45feb8dfb9089592591c",
|
||||
"296ff54c6e61441f978084df59626598",
|
||||
"d6d42b4f245a49b7ba7769e23a3202fc",
|
||||
"39ce7754480147759c16a3089d8105af",
|
||||
"8253960a069d4106863a75faae54b90d",
|
||||
"7ccf959452af4c0b873c7567747f0816",
|
||||
"ac9d0b5a5b1f401e90a1cc9ffe6d4b4c",
|
||||
"0aada067dec3472f9aba1772d6b775a5",
|
||||
"07597b1287e04653b80c47a771549376",
|
||||
"054be1dd9f084cae911745b692ccd929",
|
||||
"ab19e8e831694e308a4b79f05aff728e"
|
||||
]
|
||||
},
|
||||
"id": "bOKVUvWJegYJ",
|
||||
"outputId": "dde74917-08b0-4ce2-9a2b-cc31e02cafb2"
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"application/vnd.jupyter.widget-view+json": {
|
||||
"model_id": "9c306ce5188c45feb8dfb9089592591c",
|
||||
"version_major": 2,
|
||||
"version_minor": 0
|
||||
},
|
||||
"text/plain": [
|
||||
"Fetching 5 files: 0%| | 0/5 [00:00<?, ?it/s]"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"['CPUExecutionProvider']"
|
||||
]
|
||||
},
|
||||
"execution_count": 5,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"embedding_model_cpu = TextEmbedding(model_name=\"BAAI/bge-small-en-v1.5\")\n",
|
||||
"embedding_model_cpu.model.model.get_providers()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 6,
|
||||
"metadata": {
|
||||
"colab": {
|
||||
"base_uri": "https://localhost:8080/"
|
||||
},
|
||||
"id": "0NJj9RvSfASP",
|
||||
"outputId": "526f5280-99bd-454e-8af8-6a860ad96e54"
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"4.33 s ± 591 ms per loop (mean ± std. dev. of 7 runs, 1 loop each)\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"%%timeit\n",
|
||||
"list(embedding_model_cpu.embed(documents))"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"accelerator": "GPU",
|
||||
"colab": {
|
||||
"gpuType": "T4",
|
||||
"provenance": []
|
||||
},
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.10.12"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 1
|
||||
}
|
||||
@@ -328,7 +328,9 @@
|
||||
],
|
||||
"source": [
|
||||
"source_df = dataset.to_pandas()\n",
|
||||
"df = source_df.drop_duplicates(subset=[\"product_text\", \"product_title\", \"product_bullet_point\", \"product_brand\"])\n",
|
||||
"df = source_df.drop_duplicates(\n",
|
||||
" subset=[\"product_text\", \"product_title\", \"product_bullet_point\", \"product_brand\"]\n",
|
||||
")\n",
|
||||
"df = df.dropna(subset=[\"product_text\", \"product_title\", \"product_bullet_point\", \"product_brand\"])\n",
|
||||
"df.head()"
|
||||
]
|
||||
@@ -367,7 +369,9 @@
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df[\"combined_text\"] = df[\"product_title\"] + \"\\n\" + df[\"product_text\"] + \"\\n\" + df[\"product_bullet_point\"]"
|
||||
"df[\"combined_text\"] = (\n",
|
||||
" df[\"product_title\"] + \"\\n\" + df[\"product_text\"] + \"\\n\" + df[\"product_bullet_point\"]\n",
|
||||
")"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -489,7 +493,9 @@
|
||||
" return list(sparse_model.embed(texts, batch_size=32))\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"sparse_embedding: List[SparseEmbedding] = make_sparse_embedding([\"Fastembed is a great library for text embeddings!\"])\n",
|
||||
"sparse_embedding: List[SparseEmbedding] = make_sparse_embedding(\n",
|
||||
" [\"Fastembed is a great library for text embeddings!\"]\n",
|
||||
")\n",
|
||||
"sparse_embedding"
|
||||
]
|
||||
},
|
||||
@@ -628,7 +634,9 @@
|
||||
" token_weight_dict[token] = weight\n",
|
||||
"\n",
|
||||
" # Sort the dictionary by weights\n",
|
||||
" token_weight_dict = dict(sorted(token_weight_dict.items(), key=lambda item: item[1], reverse=True))\n",
|
||||
" token_weight_dict = dict(\n",
|
||||
" sorted(token_weight_dict.items(), key=lambda item: item[1], reverse=True)\n",
|
||||
" )\n",
|
||||
" return token_weight_dict\n",
|
||||
"\n",
|
||||
"\n",
|
||||
@@ -870,14 +878,21 @@
|
||||
" dense_vectors = df[\"dense_embedding\"].tolist()\n",
|
||||
" rows = df.to_dict(orient=\"records\")\n",
|
||||
" points = []\n",
|
||||
" for idx, (text, sparse_vector, dense_vector) in enumerate(zip(product_texts, sparse_vectors, dense_vectors)):\n",
|
||||
" sparse_vector = SparseVector(indices=sparse_vector.indices.tolist(), values=sparse_vector.values.tolist())\n",
|
||||
" for idx, (text, sparse_vector, dense_vector) in enumerate(\n",
|
||||
" zip(product_texts, sparse_vectors, dense_vectors)\n",
|
||||
" ):\n",
|
||||
" sparse_vector = SparseVector(\n",
|
||||
" indices=sparse_vector.indices.tolist(), values=sparse_vector.values.tolist()\n",
|
||||
" )\n",
|
||||
" point = PointStruct(\n",
|
||||
" id=idx,\n",
|
||||
" payload={\"text\": text, \"product_id\": rows[idx][\"product_id\"]}, # Add any additional payload if necessary\n",
|
||||
" payload={\n",
|
||||
" \"text\": text,\n",
|
||||
" \"product_id\": rows[idx][\"product_id\"],\n",
|
||||
" }, # Add any additional payload if necessary\n",
|
||||
" vector={\n",
|
||||
" \"text-sparse\": sparse_vector,\n",
|
||||
" \"text-dense\": dense_vector,\n",
|
||||
" \"text-dense\": dense_vector.tolist(),\n",
|
||||
" },\n",
|
||||
" )\n",
|
||||
" points.append(point)\n",
|
||||
@@ -936,7 +951,7 @@
|
||||
" SearchRequest(\n",
|
||||
" vector=NamedVector(\n",
|
||||
" name=\"text-dense\",\n",
|
||||
" vector=query_dense_vector[0],\n",
|
||||
" vector=query_dense_vector[0].tolist(),\n",
|
||||
" ),\n",
|
||||
" limit=10,\n",
|
||||
" with_payload=True,\n",
|
||||
@@ -1133,8 +1148,12 @@
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"def find_point_by_id(client: QdrantClient, collection_name: str, rrf_rank_list: List[Tuple[int, float]]):\n",
|
||||
" return client.retrieve(collection_name=collection_name, ids=[item[0] for item in rrf_rank_list])\n",
|
||||
"def find_point_by_id(\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",
|
||||
" )\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"find_point_by_id(client, collection_name, rrf_rank_list)"
|
||||
|
||||
@@ -0,0 +1,128 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "aa0a86859809102",
|
||||
"metadata": {
|
||||
"collapsed": false
|
||||
},
|
||||
"source": [
|
||||
"# Image Embedding\n",
|
||||
"As of version 0.3.0 fastembed supports computation of image embeddings.\n",
|
||||
"\n",
|
||||
"The process is as easy and straightforward as with text embeddings. Let's see how it works."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 5,
|
||||
"id": "cea8fd5c019571fe",
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-06-02T11:35:40.126023Z",
|
||||
"start_time": "2024-06-02T11:35:39.864701Z"
|
||||
},
|
||||
"collapsed": false
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Fetching 3 files: 100%|██████████| 3/3 [00:00<00:00, 47482.69it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"text/plain": "[array([0. , 0. , 0. , ..., 0. , 0.01139933,\n 0. ], dtype=float32),\n array([0.02169187, 0. , 0. , ..., 0. , 0.00848291,\n 0. ], dtype=float32)]"
|
||||
},
|
||||
"execution_count": 5,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"from fastembed import ImageEmbedding\n",
|
||||
"\n",
|
||||
"model = ImageEmbedding(\"Qdrant/resnet50-onnx\")\n",
|
||||
"\n",
|
||||
"embeddings_generator = model.embed(\n",
|
||||
" [\"../../tests/misc/image.jpeg\", \"../../tests/misc/small_image.jpeg\"]\n",
|
||||
")\n",
|
||||
"embeddings_list = list(embeddings_generator)\n",
|
||||
"embeddings_list"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "3f838f18523ad1e0",
|
||||
"metadata": {
|
||||
"collapsed": false
|
||||
},
|
||||
"source": [
|
||||
"## Preprocessing\n",
|
||||
"\n",
|
||||
"Preprocessing is encapsulated in the ImageEmbedding class, applied operations are identical to the ones provided by [Hugging Face Transformers](https://huggingface.co/docs/transformers/en/index).\n",
|
||||
"You don't need to think about batching, opening/closing files, resizing images, etc., Fastembed will take care of it."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "894b33ff9b385d72",
|
||||
"metadata": {
|
||||
"collapsed": false
|
||||
},
|
||||
"source": [
|
||||
"## Supported models\n",
|
||||
"\n",
|
||||
"List of supported image embedding models can either be found [here](https://qdrant.github.io/fastembed/examples/Supported_Models/#supported-image-embedding-models) or by calling the `ImageEmbedding.list_supported_models()` method."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 6,
|
||||
"id": "6d6a4cbbd2200d14",
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-06-02T11:40:19.313226Z",
|
||||
"start_time": "2024-06-02T11:40:19.309845Z"
|
||||
},
|
||||
"collapsed": false
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": "[{'model': 'Qdrant/clip-ViT-B-32-vision',\n 'dim': 512,\n 'description': 'CLIP vision encoder based on ViT-B/32',\n 'size_in_GB': 0.34,\n 'sources': {'hf': 'Qdrant/clip-ViT-B-32-vision'},\n 'model_file': 'model.onnx'},\n {'model': 'Qdrant/resnet50-onnx',\n 'dim': 2048,\n 'description': 'ResNet-50 from `Deep Residual Learning for Image Recognition <https://arxiv.org/abs/1512.03385>`__.',\n 'size_in_GB': 0.1,\n 'sources': {'hf': 'Qdrant/resnet50-onnx'},\n 'model_file': 'model.onnx'}]"
|
||||
},
|
||||
"execution_count": 6,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"ImageEmbedding.list_supported_models()"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 2
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython2",
|
||||
"version": "2.7.6"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -2,14 +2,23 @@
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"execution_count": 4,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-03-30T11:18:52.052764Z",
|
||||
"start_time": "2024-03-30T11:18:52.039616Z"
|
||||
"end_time": "2024-05-31T18:13:23.806907Z",
|
||||
"start_time": "2024-05-31T18:13:23.797078Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"The autoreload extension is already loaded. To reload it, use:\n",
|
||||
" %reload_ext autoreload\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"%load_ext autoreload\n",
|
||||
"%autoreload 2"
|
||||
@@ -18,12 +27,22 @@
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-05-31T18:14:31.147674Z",
|
||||
"start_time": "2024-05-31T18:14:31.134015Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import pandas as pd\n",
|
||||
"\n",
|
||||
"from fastembed import SparseTextEmbedding, TextEmbedding"
|
||||
"from fastembed import (\n",
|
||||
" SparseTextEmbedding,\n",
|
||||
" TextEmbedding,\n",
|
||||
" LateInteractionTextEmbedding,\n",
|
||||
" ImageEmbedding,\n",
|
||||
")"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -35,8 +54,13 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"metadata": {},
|
||||
"execution_count": 6,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-05-31T18:13:25.863008Z",
|
||||
"start_time": "2024-05-31T18:13:25.837795Z"
|
||||
}
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
@@ -117,97 +141,111 @@
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>7</th>\n",
|
||||
" <td>nomic-ai/nomic-embed-text-v1.5-Q</td>\n",
|
||||
" <td>768</td>\n",
|
||||
" <td>Quantized 8192 context length english model</td>\n",
|
||||
" <td>0.130</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>8</th>\n",
|
||||
" <td>BAAI/bge-base-en-v1.5</td>\n",
|
||||
" <td>768</td>\n",
|
||||
" <td>Base English model, v1.5</td>\n",
|
||||
" <td>0.210</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>8</th>\n",
|
||||
" <th>9</th>\n",
|
||||
" <td>sentence-transformers/paraphrase-multilingual-...</td>\n",
|
||||
" <td>384</td>\n",
|
||||
" <td>Sentence Transformer model, paraphrase-multili...</td>\n",
|
||||
" <td>0.220</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>9</th>\n",
|
||||
" <th>10</th>\n",
|
||||
" <td>Qdrant/clip-ViT-B-32-text</td>\n",
|
||||
" <td>512</td>\n",
|
||||
" <td>CLIP text encoder</td>\n",
|
||||
" <td>0.250</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>11</th>\n",
|
||||
" <td>BAAI/bge-base-en</td>\n",
|
||||
" <td>768</td>\n",
|
||||
" <td>Base English model</td>\n",
|
||||
" <td>0.420</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>10</th>\n",
|
||||
" <th>12</th>\n",
|
||||
" <td>snowflake/snowflake-arctic-embed-m</td>\n",
|
||||
" <td>768</td>\n",
|
||||
" <td>Based on intfloat/e5-base-unsupervised model, ...</td>\n",
|
||||
" <td>0.430</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>11</th>\n",
|
||||
" <td>jinaai/jina-embeddings-v2-base-en</td>\n",
|
||||
" <td>768</td>\n",
|
||||
" <td>English embedding model supporting 8192 sequen...</td>\n",
|
||||
" <td>0.520</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>12</th>\n",
|
||||
" <th>13</th>\n",
|
||||
" <td>nomic-ai/nomic-embed-text-v1</td>\n",
|
||||
" <td>768</td>\n",
|
||||
" <td>8192 context length english model</td>\n",
|
||||
" <td>0.520</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>13</th>\n",
|
||||
" <th>14</th>\n",
|
||||
" <td>jinaai/jina-embeddings-v2-base-en</td>\n",
|
||||
" <td>768</td>\n",
|
||||
" <td>English embedding model supporting 8192 sequen...</td>\n",
|
||||
" <td>0.520</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>15</th>\n",
|
||||
" <td>nomic-ai/nomic-embed-text-v1.5</td>\n",
|
||||
" <td>768</td>\n",
|
||||
" <td>8192 context length english model</td>\n",
|
||||
" <td>0.520</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>14</th>\n",
|
||||
" <th>16</th>\n",
|
||||
" <td>snowflake/snowflake-arctic-embed-m-long</td>\n",
|
||||
" <td>768</td>\n",
|
||||
" <td>Based on nomic-ai/nomic-embed-text-v1-unsuperv...</td>\n",
|
||||
" <td>0.540</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>15</th>\n",
|
||||
" <th>17</th>\n",
|
||||
" <td>mixedbread-ai/mxbai-embed-large-v1</td>\n",
|
||||
" <td>1024</td>\n",
|
||||
" <td>MixedBread Base sentence embedding model, does...</td>\n",
|
||||
" <td>0.640</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>16</th>\n",
|
||||
" <th>18</th>\n",
|
||||
" <td>sentence-transformers/paraphrase-multilingual-...</td>\n",
|
||||
" <td>768</td>\n",
|
||||
" <td>Sentence-transformers model for tasks like clu...</td>\n",
|
||||
" <td>1.000</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>17</th>\n",
|
||||
" <th>19</th>\n",
|
||||
" <td>snowflake/snowflake-arctic-embed-l</td>\n",
|
||||
" <td>1024</td>\n",
|
||||
" <td>Based on intfloat/e5-large-unsupervised, large...</td>\n",
|
||||
" <td>1.020</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>18</th>\n",
|
||||
" <th>20</th>\n",
|
||||
" <td>BAAI/bge-large-en-v1.5</td>\n",
|
||||
" <td>1024</td>\n",
|
||||
" <td>Large English model, v1.5</td>\n",
|
||||
" <td>1.200</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>19</th>\n",
|
||||
" <th>21</th>\n",
|
||||
" <td>thenlper/gte-large</td>\n",
|
||||
" <td>1024</td>\n",
|
||||
" <td>Large general text embeddings model</td>\n",
|
||||
" <td>1.200</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>20</th>\n",
|
||||
" <th>22</th>\n",
|
||||
" <td>intfloat/multilingual-e5-large</td>\n",
|
||||
" <td>1024</td>\n",
|
||||
" <td>Multilingual model, e5-large. Recommend using ...</td>\n",
|
||||
@@ -226,20 +264,22 @@
|
||||
"4 jinaai/jina-embeddings-v2-small-en 512 \n",
|
||||
"5 snowflake/snowflake-arctic-embed-s 384 \n",
|
||||
"6 BAAI/bge-small-en 384 \n",
|
||||
"7 BAAI/bge-base-en-v1.5 768 \n",
|
||||
"8 sentence-transformers/paraphrase-multilingual-... 384 \n",
|
||||
"9 BAAI/bge-base-en 768 \n",
|
||||
"10 snowflake/snowflake-arctic-embed-m 768 \n",
|
||||
"11 jinaai/jina-embeddings-v2-base-en 768 \n",
|
||||
"12 nomic-ai/nomic-embed-text-v1 768 \n",
|
||||
"13 nomic-ai/nomic-embed-text-v1.5 768 \n",
|
||||
"14 snowflake/snowflake-arctic-embed-m-long 768 \n",
|
||||
"15 mixedbread-ai/mxbai-embed-large-v1 1024 \n",
|
||||
"16 sentence-transformers/paraphrase-multilingual-... 768 \n",
|
||||
"17 snowflake/snowflake-arctic-embed-l 1024 \n",
|
||||
"18 BAAI/bge-large-en-v1.5 1024 \n",
|
||||
"19 thenlper/gte-large 1024 \n",
|
||||
"20 intfloat/multilingual-e5-large 1024 \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 BAAI/bge-base-en 768 \n",
|
||||
"12 snowflake/snowflake-arctic-embed-m 768 \n",
|
||||
"13 nomic-ai/nomic-embed-text-v1 768 \n",
|
||||
"14 jinaai/jina-embeddings-v2-base-en 768 \n",
|
||||
"15 nomic-ai/nomic-embed-text-v1.5 768 \n",
|
||||
"16 snowflake/snowflake-arctic-embed-m-long 768 \n",
|
||||
"17 mixedbread-ai/mxbai-embed-large-v1 1024 \n",
|
||||
"18 sentence-transformers/paraphrase-multilingual-... 768 \n",
|
||||
"19 snowflake/snowflake-arctic-embed-l 1024 \n",
|
||||
"20 BAAI/bge-large-en-v1.5 1024 \n",
|
||||
"21 thenlper/gte-large 1024 \n",
|
||||
"22 intfloat/multilingual-e5-large 1024 \n",
|
||||
"\n",
|
||||
" description size_in_GB \n",
|
||||
"0 Fast and Default English model 0.067 \n",
|
||||
@@ -249,23 +289,25 @@
|
||||
"4 English embedding model supporting 8192 sequen... 0.120 \n",
|
||||
"5 Based on infloat/e5-small-unsupervised, does n... 0.130 \n",
|
||||
"6 Fast English model 0.130 \n",
|
||||
"7 Base English model, v1.5 0.210 \n",
|
||||
"8 Sentence Transformer model, paraphrase-multili... 0.220 \n",
|
||||
"9 Base English model 0.420 \n",
|
||||
"10 Based on intfloat/e5-base-unsupervised model, ... 0.430 \n",
|
||||
"11 English embedding model supporting 8192 sequen... 0.520 \n",
|
||||
"12 8192 context length english model 0.520 \n",
|
||||
"7 Quantized 8192 context length english model 0.130 \n",
|
||||
"8 Base English model, v1.5 0.210 \n",
|
||||
"9 Sentence Transformer model, paraphrase-multili... 0.220 \n",
|
||||
"10 CLIP text encoder 0.250 \n",
|
||||
"11 Base English model 0.420 \n",
|
||||
"12 Based on intfloat/e5-base-unsupervised model, ... 0.430 \n",
|
||||
"13 8192 context length english model 0.520 \n",
|
||||
"14 Based on nomic-ai/nomic-embed-text-v1-unsuperv... 0.540 \n",
|
||||
"15 MixedBread Base sentence embedding model, does... 0.640 \n",
|
||||
"16 Sentence-transformers model for tasks like clu... 1.000 \n",
|
||||
"17 Based on intfloat/e5-large-unsupervised, large... 1.020 \n",
|
||||
"18 Large English model, v1.5 1.200 \n",
|
||||
"19 Large general text embeddings model 1.200 \n",
|
||||
"20 Multilingual model, e5-large. Recommend using ... 2.240 "
|
||||
"14 English embedding model supporting 8192 sequen... 0.520 \n",
|
||||
"15 8192 context length english model 0.520 \n",
|
||||
"16 Based on nomic-ai/nomic-embed-text-v1-unsuperv... 0.540 \n",
|
||||
"17 MixedBread Base sentence embedding model, does... 0.640 \n",
|
||||
"18 Sentence-transformers model for tasks like clu... 1.000 \n",
|
||||
"19 Based on intfloat/e5-large-unsupervised, large... 1.020 \n",
|
||||
"20 Large English model, v1.5 1.200 \n",
|
||||
"21 Large general text embeddings model 1.200 \n",
|
||||
"22 Multilingual model, e5-large. Recommend using ... 2.240 "
|
||||
]
|
||||
},
|
||||
"execution_count": 2,
|
||||
"execution_count": 6,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
@@ -274,7 +316,7 @@
|
||||
"supported_models = (\n",
|
||||
" pd.DataFrame(TextEmbedding.list_supported_models())\n",
|
||||
" .sort_values(\"size_in_GB\")\n",
|
||||
" .drop(columns=\"sources\")\n",
|
||||
" .drop(columns=[\"sources\", \"model_file\", \"additional_files\"])\n",
|
||||
" .reset_index(drop=True)\n",
|
||||
")\n",
|
||||
"supported_models"
|
||||
@@ -289,11 +331,11 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 7,
|
||||
"execution_count": 8,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-03-30T11:19:01.564291Z",
|
||||
"start_time": "2024-03-30T11:19:01.538768Z"
|
||||
"end_time": "2024-05-31T18:13:27.124747Z",
|
||||
"start_time": "2024-05-31T18:13:27.096212Z"
|
||||
}
|
||||
},
|
||||
"outputs": [
|
||||
@@ -322,51 +364,225 @@
|
||||
" <th>vocab_size</th>\n",
|
||||
" <th>description</th>\n",
|
||||
" <th>size_in_GB</th>\n",
|
||||
" <th>sources</th>\n",
|
||||
" </tr>\n",
|
||||
" </thead>\n",
|
||||
" <tbody>\n",
|
||||
" <tr>\n",
|
||||
" <th>0</th>\n",
|
||||
" <td>Qdrant/bm42-all-minilm-l6-v2-attentions</td>\n",
|
||||
" <td>30522</td>\n",
|
||||
" <td>Light sparse embedding model, which assigns an...</td>\n",
|
||||
" <td>0.090</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>1</th>\n",
|
||||
" <td>prithvida/Splade_PP_en_v1</td>\n",
|
||||
" <td>30522</td>\n",
|
||||
" <td>Misspelled version of the model. Retained for ...</td>\n",
|
||||
" <td>0.532</td>\n",
|
||||
" <td>{'hf': 'Qdrant/SPLADE_PP_en_v1'}</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>1</th>\n",
|
||||
" <th>2</th>\n",
|
||||
" <td>prithivida/Splade_PP_en_v1</td>\n",
|
||||
" <td>30522</td>\n",
|
||||
" <td>Independent Implementation of SPLADE++ Model f...</td>\n",
|
||||
" <td>0.532</td>\n",
|
||||
" <td>{'hf': 'Qdrant/SPLADE_PP_en_v1'}</td>\n",
|
||||
" </tr>\n",
|
||||
" </tbody>\n",
|
||||
"</table>\n",
|
||||
"</div>"
|
||||
],
|
||||
"text/plain": [
|
||||
" model vocab_size \\\n",
|
||||
"0 prithvida/Splade_PP_en_v1 30522 \n",
|
||||
"1 prithivida/Splade_PP_en_v1 30522 \n",
|
||||
" model vocab_size \\\n",
|
||||
"0 Qdrant/bm42-all-minilm-l6-v2-attentions 30522 \n",
|
||||
"1 prithvida/Splade_PP_en_v1 30522 \n",
|
||||
"2 prithivida/Splade_PP_en_v1 30522 \n",
|
||||
"\n",
|
||||
" description size_in_GB \\\n",
|
||||
"0 Misspelled version of the model. Retained for ... 0.532 \n",
|
||||
"1 Independent Implementation of SPLADE++ Model f... 0.532 \n",
|
||||
"\n",
|
||||
" sources \n",
|
||||
"0 {'hf': 'Qdrant/SPLADE_PP_en_v1'} \n",
|
||||
"1 {'hf': 'Qdrant/SPLADE_PP_en_v1'} "
|
||||
" description size_in_GB \n",
|
||||
"0 Light sparse embedding model, which assigns an... 0.090 \n",
|
||||
"1 Misspelled version of the model. Retained for ... 0.532 \n",
|
||||
"2 Independent Implementation of SPLADE++ Model f... 0.532 "
|
||||
]
|
||||
},
|
||||
"execution_count": 7,
|
||||
"execution_count": 8,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"pd.DataFrame(SparseTextEmbedding.list_supported_models())"
|
||||
"(\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",
|
||||
"metadata": {
|
||||
"collapsed": false
|
||||
},
|
||||
"source": [
|
||||
"## Supported Late Interaction Text Embedding Models"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 10,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-05-31T18:14:34.370252Z",
|
||||
"start_time": "2024-05-31T18:14:34.354270Z"
|
||||
},
|
||||
"collapsed": false
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/html": [
|
||||
"<div>\n",
|
||||
"<style scoped>\n",
|
||||
" .dataframe tbody tr th:only-of-type {\n",
|
||||
" vertical-align: middle;\n",
|
||||
" }\n",
|
||||
"\n",
|
||||
" .dataframe tbody tr th {\n",
|
||||
" vertical-align: top;\n",
|
||||
" }\n",
|
||||
"\n",
|
||||
" .dataframe thead th {\n",
|
||||
" text-align: right;\n",
|
||||
" }\n",
|
||||
"</style>\n",
|
||||
"<table border=\"1\" class=\"dataframe\">\n",
|
||||
" <thead>\n",
|
||||
" <tr style=\"text-align: right;\">\n",
|
||||
" <th></th>\n",
|
||||
" <th>model</th>\n",
|
||||
" <th>dim</th>\n",
|
||||
" <th>description</th>\n",
|
||||
" <th>size_in_GB</th>\n",
|
||||
" </tr>\n",
|
||||
" </thead>\n",
|
||||
" <tbody>\n",
|
||||
" <tr>\n",
|
||||
" <th>0</th>\n",
|
||||
" <td>colbert-ir/colbertv2.0</td>\n",
|
||||
" <td>128</td>\n",
|
||||
" <td>Late interaction model</td>\n",
|
||||
" <td>0.44</td>\n",
|
||||
" </tr>\n",
|
||||
" </tbody>\n",
|
||||
"</table>\n",
|
||||
"</div>"
|
||||
],
|
||||
"text/plain": [
|
||||
" model dim description size_in_GB\n",
|
||||
"0 colbert-ir/colbertv2.0 128 Late interaction model 0.44"
|
||||
]
|
||||
},
|
||||
"execution_count": 10,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"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",
|
||||
"metadata": {
|
||||
"collapsed": false
|
||||
},
|
||||
"source": [
|
||||
"## Supported Image Embedding Models"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 12,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-05-31T18:14:42.501881Z",
|
||||
"start_time": "2024-05-31T18:14:42.484726Z"
|
||||
},
|
||||
"collapsed": false
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/html": [
|
||||
"<div>\n",
|
||||
"<style scoped>\n",
|
||||
" .dataframe tbody tr th:only-of-type {\n",
|
||||
" vertical-align: middle;\n",
|
||||
" }\n",
|
||||
"\n",
|
||||
" .dataframe tbody tr th {\n",
|
||||
" vertical-align: top;\n",
|
||||
" }\n",
|
||||
"\n",
|
||||
" .dataframe thead th {\n",
|
||||
" text-align: right;\n",
|
||||
" }\n",
|
||||
"</style>\n",
|
||||
"<table border=\"1\" class=\"dataframe\">\n",
|
||||
" <thead>\n",
|
||||
" <tr style=\"text-align: right;\">\n",
|
||||
" <th></th>\n",
|
||||
" <th>model</th>\n",
|
||||
" <th>dim</th>\n",
|
||||
" <th>description</th>\n",
|
||||
" <th>size_in_GB</th>\n",
|
||||
" </tr>\n",
|
||||
" </thead>\n",
|
||||
" <tbody>\n",
|
||||
" <tr>\n",
|
||||
" <th>0</th>\n",
|
||||
" <td>Qdrant/resnet50-onnx</td>\n",
|
||||
" <td>2048</td>\n",
|
||||
" <td>ResNet-50 from `Deep Residual Learning for Ima...</td>\n",
|
||||
" <td>0.10</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>1</th>\n",
|
||||
" <td>Qdrant/clip-ViT-B-32-vision</td>\n",
|
||||
" <td>512</td>\n",
|
||||
" <td>CLIP vision encoder based on ViT-B/32</td>\n",
|
||||
" <td>0.34</td>\n",
|
||||
" </tr>\n",
|
||||
" </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",
|
||||
"\n",
|
||||
" description size_in_GB \n",
|
||||
"0 ResNet-50 from `Deep Residual Learning for Ima... 0.10 \n",
|
||||
"1 CLIP vision encoder based on ViT-B/32 0.34 "
|
||||
]
|
||||
},
|
||||
"execution_count": 12,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"(\n",
|
||||
" pd.DataFrame(ImageEmbedding.list_supported_models()).sort_values(\"size_in_GB\")\n",
|
||||
" .drop(columns=[\"sources\", \"model_file\"])\n",
|
||||
" .reset_index(drop=True)\n",
|
||||
")"
|
||||
]
|
||||
}
|
||||
],
|
||||
@@ -386,7 +602,7 @@
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.10.13"
|
||||
"version": "3.11.4"
|
||||
},
|
||||
"orig_nbformat": 4,
|
||||
"vscode": {
|
||||
|
||||
@@ -14,23 +14,33 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 11,
|
||||
"metadata": {},
|
||||
"execution_count": 1,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-06-06T17:00:06.460001Z",
|
||||
"start_time": "2024-06-06T17:00:04.214098Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"!pip install matplotlib tqdm pandas numpy --quiet"
|
||||
"!pip install matplotlib tqdm pandas numpy datasets --quiet --upgrade"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 12,
|
||||
"execution_count": 2,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-06-06T17:00:07.041784Z",
|
||||
"start_time": "2024-06-06T17:00:06.461658Z"
|
||||
},
|
||||
"id": "WBVTItUX4yyr"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import pandas as pd\n",
|
||||
"from datasets import load_dataset\n",
|
||||
"from tqdm import tqdm"
|
||||
]
|
||||
},
|
||||
@@ -52,8 +62,12 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 13,
|
||||
"execution_count": 3,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-06-06T17:01:09.343230Z",
|
||||
"start_time": "2024-06-06T17:00:07.042526Z"
|
||||
},
|
||||
"colab": {
|
||||
"base_uri": "https://localhost:8080/",
|
||||
"height": 250
|
||||
@@ -61,58 +75,24 @@
|
||||
"id": "REJpFqkG7EG2",
|
||||
"outputId": "7a43c0ae-fbcc-45fe-fd58-bfe691297b22"
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"100%|██████████| 26/26 [00:10<00:00, 2.45it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"(1000000, 1536)"
|
||||
]
|
||||
},
|
||||
"execution_count": 13,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def get_openai_vectors(force_download: bool = False):\n",
|
||||
" res = []\n",
|
||||
" for i in tqdm(range(26)):\n",
|
||||
" if force_download:\n",
|
||||
" !wget https://huggingface.co/api/datasets/KShivendu/dbpedia-entities-openai-1M/parquet/KShivendu--dbpedia-entities-openai-1M/train/{i}.parquet\n",
|
||||
" df = pd.read_parquet(f\"{i}.parquet\", engine=\"pyarrow\")\n",
|
||||
" res.append(np.stack(df.openai))\n",
|
||||
" del df\n",
|
||||
"\n",
|
||||
" openai_vectors = np.concatenate(res)\n",
|
||||
" del res\n",
|
||||
" return openai_vectors\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"openai_vectors = get_openai_vectors(force_download=False)\n",
|
||||
"openai_vectors.shape"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## ㆓ Binary Conversion\n",
|
||||
"\n",
|
||||
"Here, we will use 0 as the threshold for the binary conversion. All values greater than 0 will be set to 1, and others will remain 0. This is a simple and effective way to convert continuous values into binary values for OpenAI embeddings."
|
||||
"# Download from Huggingface Hub\n",
|
||||
"ds = load_dataset(\n",
|
||||
" \"Qdrant/dbpedia-entities-openai3-text-embedding-3-large-3072-100K\", split=\"train\"\n",
|
||||
")\n",
|
||||
"openai_vectors = np.array(ds[\"text-embedding-3-large-3072-embedding\"])\n",
|
||||
"del ds"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 14,
|
||||
"execution_count": 4,
|
||||
"metadata": {
|
||||
"id": "0JM2-Bj2Jkab"
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-06-06T17:01:10.900963Z",
|
||||
"start_time": "2024-06-06T17:01:09.344842Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
@@ -120,6 +100,30 @@
|
||||
"openai_bin[openai_vectors > 0] = 1"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 5,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-06-06T17:01:10.906827Z",
|
||||
"start_time": "2024-06-06T17:01:10.901820Z"
|
||||
}
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": "3072"
|
||||
},
|
||||
"execution_count": 5,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"n_dim = openai_vectors.shape[1]\n",
|
||||
"n_dim"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
@@ -131,8 +135,12 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 15,
|
||||
"execution_count": 6,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-06-06T17:01:10.909730Z",
|
||||
"start_time": "2024-06-06T17:01:10.908166Z"
|
||||
},
|
||||
"id": "FqshI-GlIERd"
|
||||
},
|
||||
"outputs": [],
|
||||
@@ -141,7 +149,7 @@
|
||||
" scores = np.dot(openai_vectors, openai_vectors[idx])\n",
|
||||
" dot_results = np.argsort(scores)[-limit:][::-1]\n",
|
||||
"\n",
|
||||
" bin_scores = 1536 - np.logical_xor(openai_bin, openai_bin[idx]).sum(axis=1)\n",
|
||||
" bin_scores = n_dim - np.logical_xor(openai_bin, openai_bin[idx]).sum(axis=1)\n",
|
||||
" bin_results = np.argsort(bin_scores)[-(limit * oversampling) :][::-1]\n",
|
||||
"\n",
|
||||
" return len(set(dot_results).intersection(set(bin_results))) / limit"
|
||||
@@ -156,8 +164,12 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 18,
|
||||
"execution_count": 7,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-06-06T17:01:25.206592Z",
|
||||
"start_time": "2024-06-06T17:01:10.911971Z"
|
||||
},
|
||||
"colab": {
|
||||
"base_uri": "https://localhost:8080/"
|
||||
},
|
||||
@@ -169,110 +181,128 @@
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
" 0%| | 0/4 [00:00<?, ?it/s]"
|
||||
" 0%| | 0/4 [00:00<?, ?it/s]\n",
|
||||
" 0%| | 0/2 [00:00<?, ?it/s]\u001b[A\n",
|
||||
" 50%|█████ | 1/2 [00:02<00:02, 2.05s/it]\u001b[A"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"{'sampling_rate': 1, 'limit': 10, 'recall': 0.8}\n"
|
||||
"{'sampling_rate': 1, 'limit': 3, 'mean_acc': 0.9}\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"100%|██████████| 2/2 [00:33<00:00, 16.98s/it]\n",
|
||||
" 25%|██▌ | 1/4 [00:33<01:41, 33.96s/it]"
|
||||
"\n",
|
||||
"100%|██████████| 2/2 [00:04<00:00, 2.02s/it]\u001b[A\n",
|
||||
" 25%|██▌ | 1/4 [00:04<00:12, 4.05s/it]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"{'sampling_rate': 1, 'limit': 100, 'recall': 0.708}\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": []
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"{'sampling_rate': 2, 'limit': 10, 'recall': 0.95}\n"
|
||||
"{'sampling_rate': 1, 'limit': 10, 'mean_acc': 0.8300000000000001}\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"100%|██████████| 2/2 [00:32<00:00, 16.38s/it]\n",
|
||||
" 50%|█████ | 2/4 [01:06<01:06, 33.26s/it]"
|
||||
"\n",
|
||||
" 0%| | 0/2 [00:00<?, ?it/s]\u001b[A\n",
|
||||
" 50%|█████ | 1/2 [00:01<00:01, 1.72s/it]\u001b[A"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"{'sampling_rate': 2, 'limit': 100, 'recall': 0.877}\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": []
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"{'sampling_rate': 3, 'limit': 10, 'recall': 0.96}\n"
|
||||
"{'sampling_rate': 2, 'limit': 3, 'mean_acc': 1.0}\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"100%|██████████| 2/2 [00:32<00:00, 16.49s/it]\n",
|
||||
" 75%|███████▌ | 3/4 [01:39<00:33, 33.13s/it]"
|
||||
"\n",
|
||||
"100%|██████████| 2/2 [00:03<00:00, 1.76s/it]\u001b[A\n",
|
||||
" 50%|█████ | 2/4 [00:07<00:07, 3.75s/it]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"{'sampling_rate': 3, 'limit': 100, 'recall': 0.937}\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": []
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"{'sampling_rate': 5, 'limit': 10, 'recall': 0.9800000000000001}\n"
|
||||
"{'sampling_rate': 2, 'limit': 10, 'mean_acc': 0.9700000000000001}\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"100%|██████████| 2/2 [00:32<00:00, 16.47s/it]\n",
|
||||
"100%|██████████| 4/4 [02:12<00:00, 33.17s/it]"
|
||||
"\n",
|
||||
" 0%| | 0/2 [00:00<?, ?it/s]\u001b[A\n",
|
||||
" 50%|█████ | 1/2 [00:01<00:01, 1.72s/it]\u001b[A"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"{'sampling_rate': 5, 'limit': 100, 'recall': 0.977}\n"
|
||||
"{'sampling_rate': 3, 'limit': 3, 'mean_acc': 1.0}\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"100%|██████████| 2/2 [00:03<00:00, 1.69s/it]\u001b[A\n",
|
||||
" 75%|███████▌ | 3/4 [00:10<00:03, 3.58s/it]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"{'sampling_rate': 3, 'limit': 10, 'mean_acc': 0.9800000000000001}\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
" 0%| | 0/2 [00:00<?, ?it/s]\u001b[A\n",
|
||||
" 50%|█████ | 1/2 [00:01<00:01, 1.68s/it]\u001b[A"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"{'sampling_rate': 5, 'limit': 3, 'mean_acc': 1.0}\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"100%|██████████| 2/2 [00:03<00:00, 1.65s/it]\u001b[A\n",
|
||||
"100%|██████████| 4/4 [00:14<00:00, 3.57s/it]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"{'sampling_rate': 5, 'limit': 10, 'mean_acc': 0.99}\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -285,119 +315,53 @@
|
||||
],
|
||||
"source": [
|
||||
"number_of_samples = 10\n",
|
||||
"limits = [10, 100]\n",
|
||||
"limits = [3, 10]\n",
|
||||
"sampling_rate = [1, 2, 3, 5]\n",
|
||||
"results = []\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"def mean_accuracy(number_of_samples, limit, sampling_rate):\n",
|
||||
" return np.mean([accuracy(i, limit=limit, oversampling=sampling_rate) for i in range(number_of_samples)])\n",
|
||||
" return np.mean(\n",
|
||||
" [accuracy(i, limit=limit, oversampling=sampling_rate) for i in range(number_of_samples)]\n",
|
||||
" )\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"for i in tqdm(sampling_rate):\n",
|
||||
" for j in tqdm(limits):\n",
|
||||
" result = {\"sampling_rate\": i, \"limit\": j, \"recall\": mean_accuracy(number_of_samples, j, i)}\n",
|
||||
" result = {\n",
|
||||
" \"sampling_rate\": i,\n",
|
||||
" \"limit\": j,\n",
|
||||
" \"mean_acc\": mean_accuracy(number_of_samples, j, i),\n",
|
||||
" }\n",
|
||||
" print(result)\n",
|
||||
" results.append(result)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 19,
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## ㆓ Binary Conversion\n",
|
||||
"\n",
|
||||
"Here, we will use 0 as the threshold for the binary conversion. All values greater than 0 will be set to 1, and others will remain 0. This is a simple and effective way to convert continuous values into binary values for OpenAI embeddings."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 8,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-06-06T17:01:25.247495Z",
|
||||
"start_time": "2024-06-06T17:01:25.213508Z"
|
||||
}
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/html": [
|
||||
"<div>\n",
|
||||
"<style scoped>\n",
|
||||
" .dataframe tbody tr th:only-of-type {\n",
|
||||
" vertical-align: middle;\n",
|
||||
" }\n",
|
||||
"\n",
|
||||
" .dataframe tbody tr th {\n",
|
||||
" vertical-align: top;\n",
|
||||
" }\n",
|
||||
"\n",
|
||||
" .dataframe thead th {\n",
|
||||
" text-align: right;\n",
|
||||
" }\n",
|
||||
"</style>\n",
|
||||
"<table border=\"1\" class=\"dataframe\">\n",
|
||||
" <thead>\n",
|
||||
" <tr style=\"text-align: right;\">\n",
|
||||
" <th></th>\n",
|
||||
" <th>sampling_rate</th>\n",
|
||||
" <th>limit</th>\n",
|
||||
" <th>recall</th>\n",
|
||||
" </tr>\n",
|
||||
" </thead>\n",
|
||||
" <tbody>\n",
|
||||
" <tr>\n",
|
||||
" <th>0</th>\n",
|
||||
" <td>1</td>\n",
|
||||
" <td>10</td>\n",
|
||||
" <td>0.800</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>1</th>\n",
|
||||
" <td>1</td>\n",
|
||||
" <td>100</td>\n",
|
||||
" <td>0.708</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>2</th>\n",
|
||||
" <td>2</td>\n",
|
||||
" <td>10</td>\n",
|
||||
" <td>0.950</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>3</th>\n",
|
||||
" <td>2</td>\n",
|
||||
" <td>100</td>\n",
|
||||
" <td>0.877</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>4</th>\n",
|
||||
" <td>3</td>\n",
|
||||
" <td>10</td>\n",
|
||||
" <td>0.960</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>5</th>\n",
|
||||
" <td>3</td>\n",
|
||||
" <td>100</td>\n",
|
||||
" <td>0.937</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>6</th>\n",
|
||||
" <td>5</td>\n",
|
||||
" <td>10</td>\n",
|
||||
" <td>0.980</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>7</th>\n",
|
||||
" <td>5</td>\n",
|
||||
" <td>100</td>\n",
|
||||
" <td>0.977</td>\n",
|
||||
" </tr>\n",
|
||||
" </tbody>\n",
|
||||
"</table>\n",
|
||||
"</div>"
|
||||
],
|
||||
"text/plain": [
|
||||
" sampling_rate limit recall\n",
|
||||
"0 1 10 0.800\n",
|
||||
"1 1 100 0.708\n",
|
||||
"2 2 10 0.950\n",
|
||||
"3 2 100 0.877\n",
|
||||
"4 3 10 0.960\n",
|
||||
"5 3 100 0.937\n",
|
||||
"6 5 10 0.980\n",
|
||||
"7 5 100 0.977"
|
||||
]
|
||||
"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>sampling_rate</th>\n <th>limit</th>\n <th>mean_acc</th>\n </tr>\n </thead>\n <tbody>\n <tr>\n <th>0</th>\n <td>1</td>\n <td>3</td>\n <td>0.90</td>\n </tr>\n <tr>\n <th>1</th>\n <td>1</td>\n <td>10</td>\n <td>0.83</td>\n </tr>\n <tr>\n <th>2</th>\n <td>2</td>\n <td>3</td>\n <td>1.00</td>\n </tr>\n <tr>\n <th>3</th>\n <td>2</td>\n <td>10</td>\n <td>0.97</td>\n </tr>\n <tr>\n <th>4</th>\n <td>3</td>\n <td>3</td>\n <td>1.00</td>\n </tr>\n <tr>\n <th>5</th>\n <td>3</td>\n <td>10</td>\n <td>0.98</td>\n </tr>\n <tr>\n <th>6</th>\n <td>5</td>\n <td>3</td>\n <td>1.00</td>\n </tr>\n <tr>\n <th>7</th>\n <td>5</td>\n <td>10</td>\n <td>0.99</td>\n </tr>\n </tbody>\n</table>\n</div>",
|
||||
"text/plain": " sampling_rate limit mean_acc\n0 1 3 0.90\n1 1 10 0.83\n2 2 3 1.00\n3 2 10 0.97\n4 3 3 1.00\n5 3 10 0.98\n6 5 3 1.00\n7 5 10 0.99"
|
||||
},
|
||||
"execution_count": 19,
|
||||
"execution_count": 8,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
@@ -408,22 +372,13 @@
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"| sampling_rate | limit | accuracy |\n",
|
||||
"|---------------|-------|----------|\n",
|
||||
"| 1 | 10 | 0.800 |\n",
|
||||
"| 1 | 100 | 0.708 |\n",
|
||||
"| 2 | 10 | 0.950 |\n",
|
||||
"| 2 | 100 | 0.877 |\n",
|
||||
"| 4 | 10 | 0.970 |\n",
|
||||
"| 4 | 100 | 0.956 |\n",
|
||||
"| 8 | 10 | 0.990 |\n",
|
||||
"| 8 | 100 | 0.990 |\n",
|
||||
"| 16 | 10 | 1.000 |\n",
|
||||
"| 16 | 100 | 0.998 |"
|
||||
]
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"collapsed": false
|
||||
},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
@@ -432,7 +387,8 @@
|
||||
"provenance": []
|
||||
},
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
@@ -445,7 +401,7 @@
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.9.17"
|
||||
"version": "3.10.13"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
|
||||
+11
-12
@@ -2,17 +2,17 @@
|
||||
|
||||
FastEmbed is a lightweight, fast, Python library built for embedding generation. We [support popular text models](https://qdrant.github.io/fastembed/examples/Supported_Models/). Please [open a Github issue](https://github.com/qdrant/fastembed/issues/new) if you want us to add a new model.
|
||||
|
||||
The default embedding supports "query" and "passage" prefixes for the input text. The default model is Flag Embedding, which is top of the [MTEB](https://huggingface.co/spaces/mteb/leaderboard) leaderboard. Here is an example for [Retrieval Embedding Generation](https://qdrant.github.io/fastembed/examples/Retrieval%20with%20FastEmbed/) and how to use [FastEmbed with Qdrant](https://qdrant.github.io/fastembed/examples/Usage_With_Qdrant/).
|
||||
|
||||
1. Light & Fast
|
||||
- Quantized model weights
|
||||
- ONNX Runtime for inference via [Optimum](https://github.com/huggingface/optimum)
|
||||
- ONNX Runtime for inference
|
||||
|
||||
2. Accuracy/Recall
|
||||
- Better than OpenAI Ada-002
|
||||
- Default is Flag Embedding, which is top of the [MTEB](https://huggingface.co/spaces/mteb/leaderboard) leaderboard
|
||||
- Default is Flag Embedding, which has shown good results on the [MTEB](https://huggingface.co/spaces/mteb/leaderboard) leaderboard
|
||||
- List of [supported models](https://qdrant.github.io/fastembed/examples/Supported_Models/) - including multilingual models
|
||||
|
||||
Here is an example for [Retrieval Embedding Generation](https://qdrant.github.io/fastembed/examples/Retrieval%20with%20FastEmbed/) and how to use [FastEmbed with Qdrant](https://qdrant.github.io/fastembed/examples/Usage_With_Qdrant/).
|
||||
|
||||
## 🚀 Installation
|
||||
|
||||
To install the FastEmbed library, pip works:
|
||||
@@ -24,16 +24,16 @@ pip install fastembed
|
||||
## 📖 Usage
|
||||
|
||||
```python
|
||||
from fastembed.embedding import FlagEmbedding as Embedding
|
||||
from fastembed import TextEmbedding
|
||||
|
||||
documents: List[str] = [
|
||||
"passage: Hello, World!",
|
||||
"query: Hello, World!", # these are two different embedding
|
||||
"query: Hello, World!",
|
||||
"passage: This is an example passage.",
|
||||
"fastembed is supported by and maintained by Qdrant." # You can leave out the prefix but it's recommended
|
||||
"fastembed is supported by and maintained by Qdrant."
|
||||
]
|
||||
embedding_model = Embedding(model_name="BAAI/bge-base-en", max_length=512)
|
||||
embeddings: List[np.ndarray] = embedding_model.embed(documents) # If you use
|
||||
embedding_model = TextEmbedding()
|
||||
embeddings: List[np.ndarray] = embedding_model.embed(documents)
|
||||
```
|
||||
|
||||
## Usage with Qdrant
|
||||
@@ -50,17 +50,16 @@ Might have to use ```pip install 'qdrant-client[fastembed]'``` on zsh.
|
||||
from qdrant_client import QdrantClient
|
||||
|
||||
# Initialize the client
|
||||
client = QdrantClient(":memory:") # or QdrantClient(path="path/to/db")
|
||||
client = QdrantClient(":memory:") # Using an in-process Qdrant
|
||||
|
||||
# Prepare your documents, metadata, and IDs
|
||||
docs = ["Qdrant has Langchain integrations", "Qdrant also has Llama Index integrations"]
|
||||
metadata = [
|
||||
{"source": "Langchain-docs"},
|
||||
{"source": "Linkedin-docs"},
|
||||
{"source": "Llama-index-docs"},
|
||||
]
|
||||
ids = [42, 2]
|
||||
|
||||
# Use the new add method
|
||||
client.add(
|
||||
collection_name="demo_collection",
|
||||
documents=docs,
|
||||
|
||||
File diff suppressed because one or more lines are too long
@@ -0,0 +1,122 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"id": "4bdb2a91-fa2a-4cee-ad5a-176cc957394d",
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-05-23T12:15:28.171586Z",
|
||||
"start_time": "2024-05-23T12:15:28.076314Z"
|
||||
}
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"ename": "ModuleNotFoundError",
|
||||
"evalue": "No module named 'torch'",
|
||||
"output_type": "error",
|
||||
"traceback": [
|
||||
"\u001B[0;31m---------------------------------------------------------------------------\u001B[0m",
|
||||
"\u001B[0;31mModuleNotFoundError\u001B[0m Traceback (most recent call last)",
|
||||
"Cell \u001B[0;32mIn[1], line 1\u001B[0m\n\u001B[0;32m----> 1\u001B[0m \u001B[38;5;28;01mimport\u001B[39;00m \u001B[38;5;21;01mtorch\u001B[39;00m\n\u001B[1;32m 2\u001B[0m \u001B[38;5;28;01mimport\u001B[39;00m \u001B[38;5;21;01mtorch\u001B[39;00m\u001B[38;5;21;01m.\u001B[39;00m\u001B[38;5;21;01monnx\u001B[39;00m\n\u001B[1;32m 3\u001B[0m \u001B[38;5;28;01mimport\u001B[39;00m \u001B[38;5;21;01mtorchvision\u001B[39;00m\u001B[38;5;21;01m.\u001B[39;00m\u001B[38;5;21;01mmodels\u001B[39;00m \u001B[38;5;28;01mas\u001B[39;00m \u001B[38;5;21;01mmodels\u001B[39;00m\n",
|
||||
"\u001B[0;31mModuleNotFoundError\u001B[0m: No module named 'torch'"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"import torch\n",
|
||||
"import torch.onnx\n",
|
||||
"import torchvision.models as models\n",
|
||||
"import torchvision.transforms as transforms\n",
|
||||
"from PIL import Image\n",
|
||||
"import numpy as np\n",
|
||||
"from tests.config import TEST_MISC_DIR\n",
|
||||
"\n",
|
||||
"# Load pre-trained ResNet-50 model\n",
|
||||
"resnet = models.resnet50(pretrained=True)\n",
|
||||
"resnet = torch.nn.Sequential(*(list(resnet.children())[:-1])) # Remove the last fully connected layer\n",
|
||||
"resnet.eval()\n",
|
||||
"\n",
|
||||
"# Define preprocessing transform\n",
|
||||
"preprocess = transforms.Compose([\n",
|
||||
" transforms.Resize(256),\n",
|
||||
" transforms.CenterCrop(224),\n",
|
||||
" transforms.ToTensor(),\n",
|
||||
" transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n",
|
||||
"])\n",
|
||||
"\n",
|
||||
"# Load and preprocess the image\n",
|
||||
"def preprocess_image(image_path):\n",
|
||||
" input_image = Image.open(image_path)\n",
|
||||
" input_tensor = preprocess(input_image)\n",
|
||||
" input_batch = input_tensor.unsqueeze(0) # Add batch dimension\n",
|
||||
" return input_batch\n",
|
||||
"\n",
|
||||
"# Example input for exporting\n",
|
||||
"input_image = preprocess_image('example.jpg')\n",
|
||||
"\n",
|
||||
"# Export the model to ONNX with dynamic axes\n",
|
||||
"torch.onnx.export(\n",
|
||||
" resnet, \n",
|
||||
" input_image, \n",
|
||||
" \"model.onnx\", \n",
|
||||
" export_params=True, \n",
|
||||
" opset_version=9, \n",
|
||||
" input_names=['input'], \n",
|
||||
" output_names=['output'],\n",
|
||||
" dynamic_axes={'input': {0: 'batch_size'}, 'output': {0: 'batch_size'}}\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"# Load ONNX model\n",
|
||||
"import onnx\n",
|
||||
"import onnxruntime as ort\n",
|
||||
"\n",
|
||||
"onnx_model = onnx.load(\"model.onnx\")\n",
|
||||
"ort_session = ort.InferenceSession(\"model.onnx\")\n",
|
||||
"\n",
|
||||
"# Run inference and extract feature vectors\n",
|
||||
"def extract_feature_vectors(image_paths):\n",
|
||||
" input_images = [preprocess_image(image_path) for image_path in image_paths]\n",
|
||||
" input_batch = torch.cat(input_images, dim=0) # Combine images into a single batch\n",
|
||||
" ort_inputs = {ort_session.get_inputs()[0].name: input_batch.numpy()}\n",
|
||||
" ort_outs = ort_session.run(None, ort_inputs)\n",
|
||||
" return ort_outs[0]\n",
|
||||
"\n",
|
||||
"# Example usage\n",
|
||||
"images = [TEST_MISC_DIR / \"image.jpeg\", str(TEST_MISC_DIR / \"small_image.jpeg\")] # Replace with your image paths\n",
|
||||
"feature_vectors = extract_feature_vectors(images)\n",
|
||||
"print(\"Feature vector shape:\", feature_vectors.shape)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"outputs": [],
|
||||
"source": [],
|
||||
"metadata": {
|
||||
"collapsed": false
|
||||
},
|
||||
"id": "baa650c4cb3e0e6d"
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.12.2"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -12,4 +12,6 @@ tokenizer = AutoTokenizer.from_pretrained(model_id)
|
||||
# print("Model already exported")
|
||||
# except FileNotFoundError:
|
||||
print(f"Exporting model to {output_dir}")
|
||||
main_export(model_id, output=output_dir, no_post_process=True, model_kwargs=model_kwargs)
|
||||
main_export(
|
||||
model_id, output=output_dir, no_post_process=True, model_kwargs=model_kwargs
|
||||
)
|
||||
|
||||
@@ -17,10 +17,14 @@ input_ids = tokenizer_output["input_ids"]
|
||||
attention_mask = tokenizer_output["attention_mask"]
|
||||
print(attention_mask)
|
||||
# Prepare the input
|
||||
input_ids = np.array(input_ids).astype(np.int64) # Replace your_input_ids with actual input data
|
||||
input_ids = np.array(input_ids).astype(
|
||||
np.int64
|
||||
) # Replace your_input_ids with actual input data
|
||||
|
||||
# Run the ONNX model
|
||||
outputs = ort_session.run(None, {"input_ids": input_ids, "attention_mask": attention_mask})
|
||||
outputs = ort_session.run(
|
||||
None, {"input_ids": input_ids, "attention_mask": attention_mask}
|
||||
)
|
||||
|
||||
# Get the attention weights
|
||||
attentions = outputs[-1]
|
||||
|
||||
+16
-3
@@ -1,7 +1,20 @@
|
||||
import importlib.metadata
|
||||
|
||||
from fastembed.image import ImageEmbedding
|
||||
from fastembed.late_interaction import LateInteractionTextEmbedding
|
||||
from fastembed.sparse import SparseEmbedding, SparseTextEmbedding
|
||||
from fastembed.text import TextEmbedding
|
||||
from fastembed.sparse import SparseTextEmbedding, SparseEmbedding
|
||||
|
||||
__version__ = importlib.metadata.version("fastembed")
|
||||
__all__ = ["TextEmbedding", "SparseTextEmbedding", "SparseEmbedding"]
|
||||
try:
|
||||
version = importlib.metadata.version("fastembed")
|
||||
except importlib.metadata.PackageNotFoundError as _:
|
||||
version = importlib.metadata.version("fastembed-gpu")
|
||||
|
||||
__version__ = version
|
||||
__all__ = [
|
||||
"TextEmbedding",
|
||||
"SparseTextEmbedding",
|
||||
"SparseEmbedding",
|
||||
"ImageEmbedding",
|
||||
"LateInteractionTextEmbedding",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
from fastembed.common.types import ImageInput, OnnxProvider, PathInput
|
||||
|
||||
__all__ = ["OnnxProvider", "ImageInput", "PathInput"]
|
||||
|
||||
@@ -2,13 +2,13 @@ import os
|
||||
import shutil
|
||||
import tarfile
|
||||
from pathlib import Path
|
||||
from typing import List, Optional, Dict, Any
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import requests
|
||||
from huggingface_hub import snapshot_download
|
||||
from huggingface_hub.utils import RepositoryNotFoundError
|
||||
from tqdm import tqdm
|
||||
from loguru import logger
|
||||
from tqdm import tqdm
|
||||
|
||||
|
||||
class ModelManagement:
|
||||
@@ -42,7 +42,9 @@ class ModelManagement:
|
||||
raise ValueError(f"Model {model_name} is not supported in {cls.__name__}.")
|
||||
|
||||
@classmethod
|
||||
def download_file_from_gcs(cls, url: str, output_path: str, show_progress: bool = True) -> str:
|
||||
def download_file_from_gcs(
|
||||
cls, url: str, output_path: str, show_progress: bool = True
|
||||
) -> str:
|
||||
"""
|
||||
Downloads a file from Google Cloud Storage.
|
||||
|
||||
@@ -71,12 +73,17 @@ class ModelManagement:
|
||||
|
||||
# Warn if the total size is zero
|
||||
if total_size_in_bytes == 0:
|
||||
print(f"Warning: Content-length header is missing or zero in the response from {url}.")
|
||||
print(
|
||||
f"Warning: Content-length header is missing or zero in the response from {url}."
|
||||
)
|
||||
|
||||
show_progress = total_size_in_bytes and show_progress
|
||||
|
||||
with tqdm(
|
||||
total=total_size_in_bytes, unit="iB", unit_scale=True, disable=not show_progress
|
||||
total=total_size_in_bytes,
|
||||
unit="iB",
|
||||
unit_scale=True,
|
||||
disable=not show_progress,
|
||||
) as progress_bar:
|
||||
with open(output_path, "wb") as file:
|
||||
for chunk in response.iter_content(chunk_size=1024):
|
||||
@@ -91,6 +98,7 @@ class ModelManagement:
|
||||
hf_source_repo: str,
|
||||
cache_dir: Optional[str] = None,
|
||||
extra_patterns: Optional[List[str]] = None,
|
||||
**kwargs,
|
||||
) -> str:
|
||||
"""
|
||||
Downloads a model from HuggingFace Hub.
|
||||
@@ -107,6 +115,7 @@ class ModelManagement:
|
||||
"tokenizer.json",
|
||||
"tokenizer_config.json",
|
||||
"special_tokens_map.json",
|
||||
"preprocessor_config.json",
|
||||
]
|
||||
if extra_patterns is not None:
|
||||
allow_patterns.extend(extra_patterns)
|
||||
@@ -115,6 +124,7 @@ class ModelManagement:
|
||||
repo_id=hf_source_repo,
|
||||
allow_patterns=allow_patterns,
|
||||
cache_dir=cache_dir,
|
||||
local_files_only=kwargs.get("local_files_only", False),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@@ -153,7 +163,9 @@ class ModelManagement:
|
||||
return cache_dir
|
||||
|
||||
@classmethod
|
||||
def retrieve_model_gcs(cls, model_name: str, source_url: str, cache_dir: str) -> Path:
|
||||
def retrieve_model_gcs(
|
||||
cls, model_name: str, source_url: str, cache_dir: str
|
||||
) -> Path:
|
||||
fast_model_name = f"fast-{model_name.split('/')[-1]}"
|
||||
|
||||
cache_tmp_dir = Path(cache_dir) / "tmp"
|
||||
@@ -179,8 +191,12 @@ class ModelManagement:
|
||||
output_path=str(model_tar_gz),
|
||||
)
|
||||
|
||||
cls.decompress_to_cache(targz_path=str(model_tar_gz), cache_dir=str(cache_tmp_dir))
|
||||
assert model_tmp_dir.exists(), f"Could not find {model_tmp_dir} in {cache_tmp_dir}"
|
||||
cls.decompress_to_cache(
|
||||
targz_path=str(model_tar_gz), cache_dir=str(cache_tmp_dir)
|
||||
)
|
||||
assert (
|
||||
model_tmp_dir.exists()
|
||||
), f"Could not find {model_tmp_dir} in {cache_tmp_dir}"
|
||||
|
||||
model_tar_gz.unlink()
|
||||
# Rename from tmp to final name is atomic
|
||||
@@ -189,7 +205,7 @@ class ModelManagement:
|
||||
return model_dir
|
||||
|
||||
@classmethod
|
||||
def download_model(cls, model: Dict[str, Any], cache_dir: Path) -> Path:
|
||||
def download_model(cls, model: Dict[str, Any], cache_dir: Path, **kwargs) -> Path:
|
||||
"""
|
||||
Downloads a model from HuggingFace Hub or Google Cloud Storage.
|
||||
|
||||
@@ -224,7 +240,10 @@ class ModelManagement:
|
||||
try:
|
||||
return Path(
|
||||
cls.download_files_from_huggingface(
|
||||
hf_source, cache_dir=str(cache_dir), extra_patterns=extra_patterns
|
||||
hf_source,
|
||||
cache_dir=str(cache_dir),
|
||||
extra_patterns=extra_patterns,
|
||||
local_files_only=kwargs.get("local_files_only", False),
|
||||
)
|
||||
)
|
||||
except (EnvironmentError, RepositoryNotFoundError, ValueError) as e:
|
||||
|
||||
@@ -1,34 +1,50 @@
|
||||
import os
|
||||
from multiprocessing import get_all_start_methods
|
||||
import warnings
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Generic, Iterable, List, Optional, Tuple, Type, TypeVar, Union
|
||||
from typing import (
|
||||
Any,
|
||||
Dict,
|
||||
Generic,
|
||||
Iterable,
|
||||
Optional,
|
||||
Sequence,
|
||||
Tuple,
|
||||
Type,
|
||||
TypeVar,
|
||||
)
|
||||
|
||||
import numpy as np
|
||||
import onnxruntime as ort
|
||||
|
||||
from fastembed.common.models import load_tokenizer
|
||||
from fastembed.common.utils import iter_batch
|
||||
from fastembed.parallel_processor import ParallelWorkerPool, Worker
|
||||
|
||||
from fastembed.common.types import OnnxProvider
|
||||
from fastembed.parallel_processor import Worker
|
||||
|
||||
# Holds type of the embedding result
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
@dataclass
|
||||
class OnnxOutputContext:
|
||||
model_output: np.ndarray
|
||||
attention_mask: Optional[np.ndarray] = None
|
||||
input_ids: Optional[np.ndarray] = None
|
||||
|
||||
|
||||
class OnnxModel(Generic[T]):
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type["EmbeddingWorker"]:
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
@classmethod
|
||||
def _post_process_onnx_output(cls, output: Tuple[np.ndarray, np.ndarray]) -> Iterable[T]:
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[T]:
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.model = None
|
||||
self.tokenizer = None
|
||||
|
||||
def _preprocess_onnx_input(self, onnx_input: Dict[str, np.ndarray]) -> Dict[str, np.ndarray]:
|
||||
def _preprocess_onnx_input(
|
||||
self, onnx_input: Dict[str, np.ndarray], **kwargs
|
||||
) -> Dict[str, np.ndarray]:
|
||||
"""
|
||||
Preprocess the onnx input.
|
||||
"""
|
||||
@@ -39,11 +55,24 @@ class OnnxModel(Generic[T]):
|
||||
model_dir: Path,
|
||||
model_file: str,
|
||||
threads: Optional[int],
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
) -> None:
|
||||
model_path = model_dir / model_file
|
||||
|
||||
# List of Execution Providers: https://onnxruntime.ai/docs/execution-providers
|
||||
onnx_providers = ["CPUExecutionProvider"]
|
||||
|
||||
onnx_providers = (
|
||||
["CPUExecutionProvider"] if providers is None else list(providers)
|
||||
)
|
||||
available_providers = ort.get_available_providers()
|
||||
requested_provider_names = []
|
||||
for provider in onnx_providers:
|
||||
# check providers available
|
||||
provider_name = provider if isinstance(provider, str) else provider[0]
|
||||
requested_provider_names.append(provider_name)
|
||||
if provider_name not in available_providers:
|
||||
raise ValueError(
|
||||
f"Provider {provider_name} is not available. Available providers: {available_providers}"
|
||||
)
|
||||
|
||||
so = ort.SessionOptions()
|
||||
so.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
|
||||
@@ -52,65 +81,21 @@ class OnnxModel(Generic[T]):
|
||||
so.intra_op_num_threads = threads
|
||||
so.inter_op_num_threads = threads
|
||||
|
||||
self.tokenizer = load_tokenizer(model_dir=model_dir)
|
||||
self.model = ort.InferenceSession(
|
||||
str(model_path), providers=onnx_providers, sess_options=so
|
||||
)
|
||||
if "CUDAExecutionProvider" in requested_provider_names:
|
||||
current_providers = self.model.get_providers()
|
||||
if "CUDAExecutionProvider" not in current_providers:
|
||||
warnings.warn(
|
||||
f"Attempt to set CUDAExecutionProvider failed. Current providers: {current_providers}."
|
||||
"If you are using CUDA 12.x, install onnxruntime-gpu via "
|
||||
"`pip install onnxruntime-gpu --extra-index-url https://aiinfra.pkgs.visualstudio.com/PublicPackages/_packaging/onnxruntime-cuda-12/pypi/simple/`",
|
||||
RuntimeWarning,
|
||||
)
|
||||
|
||||
def onnx_embed(self, documents: List[str]) -> Tuple[np.ndarray, np.ndarray]:
|
||||
encoded = self.tokenizer.encode_batch(documents)
|
||||
input_ids = np.array([e.ids for e in encoded])
|
||||
attention_mask = np.array([e.attention_mask for e in encoded])
|
||||
|
||||
onnx_input = {
|
||||
"input_ids": np.array(input_ids, dtype=np.int64),
|
||||
"attention_mask": np.array(attention_mask, dtype=np.int64),
|
||||
"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)
|
||||
|
||||
model_output = self.model.run(None, onnx_input)
|
||||
embeddings = model_output[0]
|
||||
return embeddings, attention_mask
|
||||
|
||||
def _embed_documents(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
documents: Union[str, Iterable[str]],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
) -> Iterable[T]:
|
||||
is_small = False
|
||||
|
||||
if isinstance(documents, str):
|
||||
documents = [documents]
|
||||
is_small = True
|
||||
|
||||
if isinstance(documents, list):
|
||||
if len(documents) < batch_size:
|
||||
is_small = True
|
||||
|
||||
if parallel == 0:
|
||||
parallel = os.cpu_count()
|
||||
|
||||
if parallel is None or is_small:
|
||||
for batch in iter_batch(documents, batch_size):
|
||||
yield from self._post_process_onnx_output(self.onnx_embed(batch))
|
||||
else:
|
||||
start_method = "forkserver" if "forkserver" in get_all_start_methods() else "spawn"
|
||||
params = {
|
||||
"model_name": model_name,
|
||||
"cache_dir": cache_dir,
|
||||
}
|
||||
pool = ParallelWorkerPool(
|
||||
parallel, self._get_worker_class(), start_method=start_method
|
||||
)
|
||||
for batch in pool.ordered_map(iter_batch(documents, batch_size), **params):
|
||||
yield from self._post_process_onnx_output(batch)
|
||||
def onnx_embed(self, *args, **kwargs) -> OnnxOutputContext:
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
|
||||
class EmbeddingWorker(Worker):
|
||||
@@ -118,6 +103,7 @@ class EmbeddingWorker(Worker):
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
**kwargs,
|
||||
) -> OnnxModel:
|
||||
raise NotImplementedError()
|
||||
|
||||
@@ -125,17 +111,13 @@ class EmbeddingWorker(Worker):
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
**kwargs,
|
||||
):
|
||||
self.model = self.init_embedding(model_name, cache_dir)
|
||||
self.model = self.init_embedding(model_name, cache_dir, **kwargs)
|
||||
|
||||
@classmethod
|
||||
def start(cls, model_name: str, cache_dir: str, **kwargs: Any) -> "EmbeddingWorker":
|
||||
return cls(
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
)
|
||||
return cls(model_name=model_name, cache_dir=cache_dir, **kwargs)
|
||||
|
||||
def process(self, items: Iterable[Tuple[int, Any]]) -> Iterable[Tuple[int, Any]]:
|
||||
for idx, batch in items:
|
||||
embeddings, attn_mask = self.model.onnx_embed(batch)
|
||||
yield idx, (embeddings, attn_mask)
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
@@ -1,11 +1,24 @@
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Tuple
|
||||
|
||||
import numpy as np
|
||||
from tokenizers import Tokenizer, AddedToken
|
||||
from tokenizers import AddedToken, Tokenizer
|
||||
|
||||
from fastembed.image.transform.operators import Compose
|
||||
|
||||
|
||||
def load_tokenizer(model_dir: Path, max_length: int = 512) -> Tokenizer:
|
||||
def load_special_tokens(model_dir: Path) -> dict:
|
||||
tokens_map_path = model_dir / "special_tokens_map.json"
|
||||
if not tokens_map_path.exists():
|
||||
raise ValueError(f"Could not find special_tokens_map.json in {model_dir}")
|
||||
|
||||
with open(str(tokens_map_path)) as tokens_map_file:
|
||||
tokens_map = json.load(tokens_map_file)
|
||||
|
||||
return tokens_map
|
||||
|
||||
|
||||
def load_tokenizer(model_dir: Path, max_length: int = 512) -> 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}")
|
||||
@@ -18,21 +31,18 @@ def load_tokenizer(model_dir: Path, max_length: int = 512) -> Tokenizer:
|
||||
if not tokenizer_config_path.exists():
|
||||
raise ValueError(f"Could not find tokenizer_config.json in {model_dir}")
|
||||
|
||||
tokens_map_path = model_dir / "special_tokens_map.json"
|
||||
if not tokens_map_path.exists():
|
||||
raise ValueError(f"Could not find special_tokens_map.json in {model_dir}")
|
||||
|
||||
with open(str(config_path)) as config_file:
|
||||
config = json.load(config_file)
|
||||
|
||||
with open(str(tokenizer_config_path)) as tokenizer_config_file:
|
||||
tokenizer_config = json.load(tokenizer_config_file)
|
||||
|
||||
with open(str(tokens_map_path)) as tokens_map_file:
|
||||
tokens_map = json.load(tokens_map_file)
|
||||
tokens_map = load_special_tokens(model_dir)
|
||||
|
||||
tokenizer = Tokenizer.from_file(str(tokenizer_path))
|
||||
tokenizer.enable_truncation(max_length=min(tokenizer_config["model_max_length"], max_length))
|
||||
tokenizer.enable_truncation(
|
||||
max_length=min(tokenizer_config["model_max_length"], max_length)
|
||||
)
|
||||
tokenizer.enable_padding(
|
||||
pad_id=config.get("pad_token_id", 0), pad_token=tokenizer_config["pad_token"]
|
||||
)
|
||||
@@ -43,12 +53,24 @@ def load_tokenizer(model_dir: Path, max_length: int = 512) -> Tokenizer:
|
||||
elif isinstance(token, dict):
|
||||
tokenizer.add_special_tokens([AddedToken(**token)])
|
||||
|
||||
return tokenizer
|
||||
special_token_to_id = {}
|
||||
|
||||
for token in tokens_map.values():
|
||||
if isinstance(token, str):
|
||||
special_token_to_id[token] = tokenizer.token_to_id(token)
|
||||
elif isinstance(token, dict):
|
||||
token_str = token.get("content", "")
|
||||
special_token_to_id[token_str] = tokenizer.token_to_id(token_str)
|
||||
|
||||
return tokenizer, special_token_to_id
|
||||
|
||||
|
||||
def normalize(input_array, p=2, dim=1, eps=1e-12) -> np.ndarray:
|
||||
# Calculate the Lp norm along the specified dimension
|
||||
norm = np.linalg.norm(input_array, ord=p, axis=dim, keepdims=True)
|
||||
norm = np.maximum(norm, eps) # Avoid division by zero
|
||||
normalized_array = input_array / norm
|
||||
return normalized_array
|
||||
def load_preprocessor(model_dir: Path) -> Compose:
|
||||
preprocessor_config_path = model_dir / "preprocessor_config.json"
|
||||
if not preprocessor_config_path.exists():
|
||||
raise ValueError(f"Could not find preprocessor_config.json in {model_dir}")
|
||||
|
||||
with open(str(preprocessor_config_path)) as preprocessor_config_file:
|
||||
preprocessor_config = json.load(preprocessor_config_file)
|
||||
transforms = Compose.from_config(preprocessor_config)
|
||||
return transforms
|
||||
@@ -0,0 +1,14 @@
|
||||
import os
|
||||
import sys
|
||||
from typing import Any, Dict, Iterable, Tuple, Union
|
||||
|
||||
if sys.version_info >= (3, 10):
|
||||
from typing import TypeAlias
|
||||
else:
|
||||
from typing_extensions import TypeAlias
|
||||
|
||||
|
||||
PathInput: TypeAlias = Union[str, os.PathLike]
|
||||
ImageInput: TypeAlias = Union[PathInput, Iterable[PathInput]]
|
||||
|
||||
OnnxProvider: TypeAlias = Union[str, Tuple[str, Dict[Any, Any]]]
|
||||
@@ -2,7 +2,17 @@ import os
|
||||
import tempfile
|
||||
from itertools import islice
|
||||
from pathlib import Path
|
||||
from typing import Union, Iterable, Generator, Optional
|
||||
from typing import Generator, Iterable, Optional, Union
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
def normalize(input_array, p=2, dim=1, eps=1e-12) -> np.ndarray:
|
||||
# Calculate the Lp norm along the specified dimension
|
||||
norm = np.linalg.norm(input_array, ord=p, axis=dim, keepdims=True)
|
||||
norm = np.maximum(norm, eps) # Avoid division by zero
|
||||
normalized_array = input_array / norm
|
||||
return normalized_array
|
||||
|
||||
|
||||
def iter_batch(iterable: Union[Iterable, Generator], size: int) -> Iterable:
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
from fastembed.image.image_embedding import ImageEmbedding
|
||||
|
||||
__all__ = ["ImageEmbedding"]
|
||||
|
||||
@@ -0,0 +1,94 @@
|
||||
from typing import Any, Dict, Iterable, List, Optional, Sequence, Type
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastembed.common import ImageInput, OnnxProvider
|
||||
from fastembed.image.image_embedding_base import ImageEmbeddingBase
|
||||
from fastembed.image.onnx_embedding import OnnxImageEmbedding
|
||||
|
||||
|
||||
class ImageEmbedding(ImageEmbeddingBase):
|
||||
EMBEDDINGS_REGISTRY: List[Type[ImageEmbeddingBase]] = [OnnxImageEmbedding]
|
||||
|
||||
@classmethod
|
||||
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.
|
||||
|
||||
Example:
|
||||
```
|
||||
[
|
||||
{
|
||||
"model": "Qdrant/clip-ViT-B-32-vision",
|
||||
"dim": 512,
|
||||
"description": "CLIP vision encoder based on ViT-B/32",
|
||||
"size_in_GB": 0.33,
|
||||
"sources": {
|
||||
"hf": "Qdrant/clip-ViT-B-32-vision",
|
||||
},
|
||||
"model_file": "model.onnx",
|
||||
}
|
||||
]
|
||||
```
|
||||
"""
|
||||
result = []
|
||||
for embedding in cls.EMBEDDINGS_REGISTRY:
|
||||
result.extend(embedding.list_supported_models())
|
||||
return result
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(model_name, cache_dir, threads, **kwargs)
|
||||
|
||||
for EMBEDDING_MODEL_TYPE in self.EMBEDDINGS_REGISTRY:
|
||||
supported_models = EMBEDDING_MODEL_TYPE.list_supported_models()
|
||||
if any(
|
||||
model_name.lower() == model["model"].lower()
|
||||
for model in supported_models
|
||||
):
|
||||
self.model = EMBEDDING_MODEL_TYPE(
|
||||
model_name,
|
||||
cache_dir,
|
||||
threads=threads,
|
||||
providers=providers,
|
||||
**kwargs,
|
||||
)
|
||||
return
|
||||
|
||||
raise ValueError(
|
||||
f"Model {model_name} is not supported in TextEmbedding."
|
||||
"Please check the supported models using `TextEmbedding.list_supported_models()`"
|
||||
)
|
||||
|
||||
def embed(
|
||||
self,
|
||||
images: ImageInput,
|
||||
batch_size: int = 16,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs,
|
||||
) -> Iterable[np.ndarray]:
|
||||
"""
|
||||
Encode a list of documents into list of embeddings.
|
||||
We use mean pooling with attention so that the model can handle variable-length inputs.
|
||||
|
||||
Args:
|
||||
images: Iterator of image paths or single image path to embed
|
||||
batch_size: Batch size for encoding -- higher values will use more memory, but be faster
|
||||
parallel:
|
||||
If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.
|
||||
If 0, use all available cores.
|
||||
If None, don't use data-parallel processing, use default onnxruntime threading instead.
|
||||
|
||||
Returns:
|
||||
List of embeddings, one per document
|
||||
"""
|
||||
yield from self.model.embed(images, batch_size, parallel, **kwargs)
|
||||
@@ -0,0 +1,44 @@
|
||||
from typing import Iterable, Optional
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastembed.common.model_management import ModelManagement
|
||||
from fastembed.common.types import ImageInput
|
||||
|
||||
|
||||
class ImageEmbeddingBase(ModelManagement):
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
**kwargs,
|
||||
):
|
||||
self.model_name = model_name
|
||||
self.cache_dir = cache_dir
|
||||
self.threads = threads
|
||||
self._local_files_only = kwargs.pop("local_files_only", False)
|
||||
|
||||
def embed(
|
||||
self,
|
||||
images: ImageInput,
|
||||
batch_size: int = 16,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs,
|
||||
) -> Iterable[np.ndarray]:
|
||||
"""
|
||||
Embeds a list of images into a list of embeddings.
|
||||
|
||||
Args:
|
||||
images - The list of image paths to preprocess and embed.
|
||||
batch_size: Batch size for encoding
|
||||
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 embed method.
|
||||
|
||||
Yields:
|
||||
Iterable[np.ndarray]: The embeddings.
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
@@ -0,0 +1,131 @@
|
||||
from typing import Any, Dict, Iterable, List, Optional, Sequence, Type
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastembed.common import ImageInput, OnnxProvider
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.common.utils import define_cache_dir, normalize
|
||||
from fastembed.image.image_embedding_base import ImageEmbeddingBase
|
||||
from fastembed.image.onnx_image_model import ImageEmbeddingWorker, OnnxImageModel
|
||||
|
||||
supported_onnx_models = [
|
||||
{
|
||||
"model": "Qdrant/clip-ViT-B-32-vision",
|
||||
"dim": 512,
|
||||
"description": "CLIP vision encoder based on ViT-B/32",
|
||||
"size_in_GB": 0.34,
|
||||
"sources": {
|
||||
"hf": "Qdrant/clip-ViT-B-32-vision",
|
||||
},
|
||||
"model_file": "model.onnx",
|
||||
},
|
||||
{
|
||||
"model": "Qdrant/resnet50-onnx",
|
||||
"dim": 2048,
|
||||
"description": "ResNet-50 from `Deep Residual Learning for Image Recognition <https://arxiv.org/abs/1512.03385>`__.",
|
||||
"size_in_GB": 0.1,
|
||||
"sources": {
|
||||
"hf": "Qdrant/resnet50-onnx",
|
||||
},
|
||||
"model_file": "model.onnx",
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[np.ndarray]):
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
model_name (str): The name of the model to use.
|
||||
cache_dir (str, optional): The path to the cache directory.
|
||||
Can be set using the `FASTEMBED_CACHE_PATH` env variable.
|
||||
Defaults to `fastembed_cache` in the system's temp directory.
|
||||
threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.
|
||||
|
||||
Raises:
|
||||
ValueError: If the model_name is not in the format <org>/<model> e.g. BAAI/bge-base-en.
|
||||
"""
|
||||
|
||||
super().__init__(model_name, cache_dir, threads, **kwargs)
|
||||
|
||||
model_description = self._get_model_description(model_name)
|
||||
self.cache_dir = define_cache_dir(cache_dir)
|
||||
model_dir = self.download_model(
|
||||
model_description, self.cache_dir, local_files_only=self._local_files_only
|
||||
)
|
||||
|
||||
self.load_onnx_model(
|
||||
model_dir=model_dir,
|
||||
model_file=model_description["model_file"],
|
||||
threads=threads,
|
||||
providers=providers,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
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.
|
||||
"""
|
||||
return supported_onnx_models
|
||||
|
||||
def embed(
|
||||
self,
|
||||
images: ImageInput,
|
||||
batch_size: int = 16,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs,
|
||||
) -> Iterable[np.ndarray]:
|
||||
"""
|
||||
Encode a list of images into list of embeddings.
|
||||
We use mean pooling with attention so that the model can handle variable-length inputs.
|
||||
|
||||
Args:
|
||||
images: Iterator of image paths or single image path to embed
|
||||
batch_size: Batch size for encoding -- higher values will use more memory, but be faster
|
||||
parallel:
|
||||
If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.
|
||||
If 0, use all available cores.
|
||||
If None, don't use data-parallel processing, use default onnxruntime threading instead.
|
||||
|
||||
Returns:
|
||||
List of embeddings, one per document
|
||||
"""
|
||||
yield from self._embed_images(
|
||||
model_name=self.model_name,
|
||||
cache_dir=str(self.cache_dir),
|
||||
images=images,
|
||||
batch_size=batch_size,
|
||||
parallel=parallel,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type["ImageEmbeddingWorker"]:
|
||||
return OnnxImageEmbeddingWorker
|
||||
|
||||
def _preprocess_onnx_input(
|
||||
self, onnx_input: Dict[str, np.ndarray], **kwargs
|
||||
) -> Dict[str, np.ndarray]:
|
||||
"""
|
||||
Preprocess the onnx input.
|
||||
"""
|
||||
|
||||
return onnx_input
|
||||
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[np.ndarray]:
|
||||
return normalize(output.model_output).astype(np.float32)
|
||||
|
||||
|
||||
class OnnxImageEmbeddingWorker(ImageEmbeddingWorker):
|
||||
def init_embedding(self, model_name: str, cache_dir: str, **kwargs) -> OnnxImageEmbedding:
|
||||
return OnnxImageEmbedding(model_name=model_name, cache_dir=cache_dir, threads=1, **kwargs)
|
||||
@@ -0,0 +1,108 @@
|
||||
import contextlib
|
||||
import os
|
||||
from multiprocessing import get_all_start_methods
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Iterable, List, Optional, Sequence, Tuple, Type
|
||||
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
|
||||
from fastembed.common import ImageInput, OnnxProvider, PathInput
|
||||
from fastembed.common.onnx_model import EmbeddingWorker, OnnxModel, OnnxOutputContext, T
|
||||
from fastembed.common.preprocessor_utils import load_preprocessor
|
||||
from fastembed.common.utils import iter_batch
|
||||
from fastembed.parallel_processor import ParallelWorkerPool
|
||||
|
||||
# Holds type of the embedding result
|
||||
|
||||
|
||||
class OnnxImageModel(OnnxModel[T]):
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type["ImageEmbeddingWorker"]:
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[T]:
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.processor = None
|
||||
|
||||
def _preprocess_onnx_input(
|
||||
self, onnx_input: Dict[str, np.ndarray], **kwargs
|
||||
) -> Dict[str, np.ndarray]:
|
||||
"""
|
||||
Preprocess the onnx input.
|
||||
"""
|
||||
return onnx_input
|
||||
|
||||
def load_onnx_model(
|
||||
self,
|
||||
model_dir: Path,
|
||||
model_file: str,
|
||||
threads: Optional[int],
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
) -> None:
|
||||
super().load_onnx_model(
|
||||
model_dir=model_dir,
|
||||
model_file=model_file,
|
||||
threads=threads,
|
||||
providers=providers,
|
||||
)
|
||||
self.processor = load_preprocessor(model_dir=model_dir)
|
||||
|
||||
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[PathInput], **kwargs) -> OnnxOutputContext:
|
||||
with contextlib.ExitStack():
|
||||
image_files = [Image.open(image) for image in images]
|
||||
encoded = self.processor(image_files)
|
||||
onnx_input = self._build_onnx_input(encoded)
|
||||
onnx_input = self._preprocess_onnx_input(onnx_input)
|
||||
model_output = self.model.run(None, onnx_input)
|
||||
embeddings = model_output[0].reshape(len(images), -1)
|
||||
return OnnxOutputContext(model_output=embeddings)
|
||||
|
||||
def _embed_images(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
images: ImageInput,
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs,
|
||||
) -> Iterable[T]:
|
||||
is_small = False
|
||||
|
||||
if isinstance(images, str) or isinstance(images, Path):
|
||||
images = [images]
|
||||
is_small = True
|
||||
|
||||
if isinstance(images, list):
|
||||
if len(images) < batch_size:
|
||||
is_small = True
|
||||
|
||||
if parallel == 0:
|
||||
parallel = os.cpu_count()
|
||||
|
||||
if parallel is None or is_small:
|
||||
for batch in iter_batch(images, batch_size):
|
||||
yield from self._post_process_onnx_output(self.onnx_embed(batch))
|
||||
else:
|
||||
start_method = (
|
||||
"forkserver" if "forkserver" in get_all_start_methods() else "spawn"
|
||||
)
|
||||
params = {"model_name": model_name, "cache_dir": cache_dir, **kwargs}
|
||||
pool = ParallelWorkerPool(
|
||||
parallel, self._get_worker_class(), start_method=start_method
|
||||
)
|
||||
for batch in pool.ordered_map(iter_batch(images, batch_size), **params):
|
||||
yield from self._post_process_onnx_output(batch)
|
||||
|
||||
|
||||
class ImageEmbeddingWorker(EmbeddingWorker):
|
||||
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
|
||||
@@ -0,0 +1,124 @@
|
||||
from typing import Sized, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
|
||||
|
||||
def convert_to_rgb(image: Image.Image) -> Image.Image:
|
||||
if image.mode == "RGB":
|
||||
return image
|
||||
|
||||
image = image.convert("RGB")
|
||||
return image
|
||||
|
||||
|
||||
def center_crop(
|
||||
image: Union[Image.Image, np.ndarray],
|
||||
size: Tuple[int, int],
|
||||
) -> np.ndarray:
|
||||
if isinstance(image, np.ndarray):
|
||||
_, orig_height, orig_width = image.shape
|
||||
else:
|
||||
orig_height, orig_width = image.height, image.width
|
||||
# (H, W, C) -> (C, H, W)
|
||||
image = np.array(image).transpose((2, 0, 1))
|
||||
|
||||
crop_height, crop_width = size
|
||||
|
||||
# left upper corner (0, 0)
|
||||
top = (orig_height - crop_height) // 2
|
||||
bottom = top + crop_height
|
||||
left = (orig_width - crop_width) // 2
|
||||
right = left + crop_width
|
||||
|
||||
# Check if cropped area is within image boundaries
|
||||
if top >= 0 and bottom <= orig_height and left >= 0 and right <= orig_width:
|
||||
image = image[..., top:bottom, left:right]
|
||||
return image
|
||||
|
||||
# Padding with zeros
|
||||
new_height = max(crop_height, orig_height)
|
||||
new_width = max(crop_width, orig_width)
|
||||
new_shape = image.shape[:-2] + (new_height, new_width)
|
||||
new_image = np.zeros_like(image, shape=new_shape)
|
||||
|
||||
top_pad = (new_height - orig_height) // 2
|
||||
bottom_pad = top_pad + orig_height
|
||||
left_pad = (new_width - orig_width) // 2
|
||||
right_pad = left_pad + orig_width
|
||||
new_image[..., top_pad:bottom_pad, left_pad:right_pad] = image
|
||||
|
||||
top += top_pad
|
||||
bottom += top_pad
|
||||
left += left_pad
|
||||
right += left_pad
|
||||
|
||||
new_image = new_image[
|
||||
..., max(0, top) : min(new_height, bottom), max(0, left) : min(new_width, right)
|
||||
]
|
||||
|
||||
return new_image
|
||||
|
||||
|
||||
def normalize(
|
||||
image: 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")
|
||||
|
||||
num_channels = image.shape[1] if len(image.shape) == 4 else image.shape[0]
|
||||
|
||||
if not np.issubdtype(image.dtype, np.floating):
|
||||
image = image.astype(np.float32)
|
||||
|
||||
if isinstance(mean, Sized):
|
||||
if len(mean) != num_channels:
|
||||
raise ValueError(
|
||||
f"mean must have {num_channels} elements if it is an iterable, got {len(mean)}"
|
||||
)
|
||||
else:
|
||||
mean = [mean] * num_channels
|
||||
mean = np.array(mean, dtype=image.dtype)
|
||||
|
||||
if isinstance(std, Sized):
|
||||
if len(std) != num_channels:
|
||||
raise ValueError(
|
||||
f"std must have {num_channels} elements if it is an iterable, got {len(std)}"
|
||||
)
|
||||
else:
|
||||
std = [std] * num_channels
|
||||
std = np.array(std, dtype=image.dtype)
|
||||
|
||||
image = ((image.T - mean) / std).T
|
||||
return image
|
||||
|
||||
|
||||
def resize(
|
||||
image: Image,
|
||||
size: Union[int, Tuple[int, int]],
|
||||
resample: Image.Resampling = Image.Resampling.BILINEAR,
|
||||
) -> Image:
|
||||
if isinstance(size, tuple):
|
||||
return image.resize(size, resample)
|
||||
|
||||
height, width = image.height, image.width
|
||||
short, long = (width, height) if width <= height else (height, width)
|
||||
|
||||
new_short, new_long = size, int(size * long / short)
|
||||
if width <= height:
|
||||
new_size = (new_short, new_long)
|
||||
else:
|
||||
new_size = (new_long, new_short)
|
||||
return image.resize(new_size, resample)
|
||||
|
||||
|
||||
def rescale(image: np.ndarray, scale: float, dtype=np.float32) -> np.ndarray:
|
||||
return (image * scale).astype(dtype)
|
||||
|
||||
|
||||
def pil2ndarray(image: Union[Image.Image, np.ndarray]):
|
||||
if isinstance(image, Image.Image):
|
||||
return np.asarray(image).transpose((2, 0, 1))
|
||||
return image
|
||||
@@ -0,0 +1,198 @@
|
||||
from typing import Any, Dict, List, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
|
||||
from fastembed.image.transform.functional import (
|
||||
center_crop,
|
||||
convert_to_rgb,
|
||||
normalize,
|
||||
pil2ndarray,
|
||||
rescale,
|
||||
resize,
|
||||
)
|
||||
|
||||
|
||||
class Transform:
|
||||
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]:
|
||||
return [convert_to_rgb(image=image) for image in images]
|
||||
|
||||
|
||||
class CenterCrop(Transform):
|
||||
def __init__(self, size: Tuple[int, int]):
|
||||
self.size = size
|
||||
|
||||
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]]):
|
||||
self.mean = mean
|
||||
self.std = std
|
||||
|
||||
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]],
|
||||
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
|
||||
]
|
||||
|
||||
|
||||
class Rescale(Transform):
|
||||
def __init__(self, scale: float = 1 / 255):
|
||||
self.scale = scale
|
||||
|
||||
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]:
|
||||
return [pil2ndarray(image) for image in images]
|
||||
|
||||
|
||||
class Compose:
|
||||
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]]:
|
||||
for transform in self.transforms:
|
||||
images = transform(images)
|
||||
return images
|
||||
|
||||
@classmethod
|
||||
def from_config(cls, config: Dict[str, Any]) -> "Compose":
|
||||
"""Creates processor from a config dict.
|
||||
Args:
|
||||
config (Dict[str, Any]): Configuration dictionary.
|
||||
|
||||
Valid keys:
|
||||
- do_resize
|
||||
- size
|
||||
- do_center_crop
|
||||
- crop_size
|
||||
- do_rescale
|
||||
- rescale_factor
|
||||
- do_normalize
|
||||
- image_mean
|
||||
- image_std
|
||||
Valid size keys (nested):
|
||||
- {"height", "width"}
|
||||
- {"shortest_edge"}
|
||||
|
||||
Returns:
|
||||
Compose: Image processor.
|
||||
"""
|
||||
transforms = []
|
||||
cls._get_convert_to_rgb(transforms, config)
|
||||
cls._get_resize(transforms, config)
|
||||
cls._get_center_crop(transforms, config)
|
||||
cls._get_pil2ndarray(transforms, config)
|
||||
cls._get_rescale(transforms, config)
|
||||
cls._get_normalize(transforms, config)
|
||||
return cls(transforms=transforms)
|
||||
|
||||
@staticmethod
|
||||
def _get_convert_to_rgb(transforms: List[Transform], config: Dict[str, Any]):
|
||||
transforms.append(ConvertToRGB())
|
||||
|
||||
@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):
|
||||
size = config["size"]
|
||||
if "shortest_edge" in size:
|
||||
size = size["shortest_edge"]
|
||||
elif "height" in size and "width" in size:
|
||||
size = (size["height"], size["width"])
|
||||
else:
|
||||
raise ValueError(
|
||||
"Size must contain either 'shortest_edge' or 'height' and 'width'."
|
||||
)
|
||||
transforms.append(
|
||||
Resize(
|
||||
size=size,
|
||||
resample=config.get("resample", Image.Resampling.BICUBIC),
|
||||
)
|
||||
)
|
||||
elif mode == "ConvNextFeatureExtractor":
|
||||
if "size" in config and "shortest_edge" not in config["size"]:
|
||||
raise ValueError(
|
||||
f"Size dictionary must contain 'shortest_edge' key. Got {config['size'].keys()}"
|
||||
)
|
||||
shortest_edge = config["size"]["shortest_edge"]
|
||||
crop_pct = config.get("crop_pct", 0.875)
|
||||
if shortest_edge < 384:
|
||||
# maintain same ratio, resizing shortest edge to shortest_edge/crop_pct
|
||||
resize_shortest_edge = int(shortest_edge / crop_pct)
|
||||
transforms.append(
|
||||
Resize(
|
||||
size=resize_shortest_edge,
|
||||
resample=config.get("resample", Image.Resampling.BICUBIC),
|
||||
)
|
||||
)
|
||||
transforms.append(CenterCrop(size=(shortest_edge, shortest_edge)))
|
||||
else:
|
||||
transforms.append(
|
||||
Resize(
|
||||
size=(shortest_edge, shortest_edge),
|
||||
resample=config.get("resample", Image.Resampling.BICUBIC),
|
||||
)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
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):
|
||||
crop_size = config["crop_size"]
|
||||
if isinstance(crop_size, int):
|
||||
crop_size = (crop_size, crop_size)
|
||||
elif isinstance(crop_size, dict):
|
||||
crop_size = (crop_size["height"], crop_size["width"])
|
||||
else:
|
||||
raise ValueError(f"Invalid crop size: {crop_size}")
|
||||
transforms.append(CenterCrop(size=crop_size))
|
||||
elif mode == "ConvNextFeatureExtractor":
|
||||
pass
|
||||
else:
|
||||
raise ValueError(f"Preprocessor {mode} is not supported")
|
||||
|
||||
@staticmethod
|
||||
def _get_pil2ndarray(transforms: List[Transform], config: Dict[str, Any]):
|
||||
transforms.append(PILtoNDarray())
|
||||
|
||||
@staticmethod
|
||||
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]):
|
||||
if config.get("do_normalize", False):
|
||||
transforms.append(
|
||||
Normalize(mean=config["image_mean"], std=config["image_std"])
|
||||
)
|
||||
@@ -0,0 +1,5 @@
|
||||
from fastembed.late_interaction.late_interaction_text_embedding import (
|
||||
LateInteractionTextEmbedding,
|
||||
)
|
||||
|
||||
__all__ = ["LateInteractionTextEmbedding"]
|
||||
@@ -0,0 +1,194 @@
|
||||
import string
|
||||
from typing import Any, Dict, Iterable, List, Optional, Sequence, Type, Union
|
||||
|
||||
import numpy as np
|
||||
from tokenizers import Encoding
|
||||
|
||||
from fastembed.common import OnnxProvider
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.common.utils import define_cache_dir
|
||||
from fastembed.late_interaction.late_interaction_embedding_base import (
|
||||
LateInteractionTextEmbeddingBase,
|
||||
)
|
||||
from fastembed.text.onnx_text_model import OnnxTextModel, TextEmbeddingWorker
|
||||
|
||||
supported_colbert_models = [
|
||||
{
|
||||
"model": "colbert-ir/colbertv2.0",
|
||||
"dim": 128,
|
||||
"description": "Late interaction model",
|
||||
"size_in_GB": 0.44,
|
||||
"sources": {
|
||||
"hf": "colbert-ir/colbertv2.0",
|
||||
},
|
||||
"model_file": "model.onnx",
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[np.ndarray]):
|
||||
QUERY_MARKER_TOKEN_ID = 1
|
||||
DOCUMENT_MARKER_TOKEN_ID = 2
|
||||
MIN_QUERY_LENGTH = 32
|
||||
MASK_TOKEN = "[MASK]"
|
||||
|
||||
def _post_process_onnx_output(
|
||||
self, output: OnnxOutputContext, is_doc: bool = True
|
||||
) -> Iterable[np.ndarray]:
|
||||
if not is_doc:
|
||||
return output.model_output.astype(np.float32)
|
||||
|
||||
for i, token_sequence in enumerate(output.input_ids):
|
||||
for j, token_id in enumerate(token_sequence):
|
||||
if token_id in self.skip_list or token_id == self.pad_token_id:
|
||||
output.attention_mask[i, j] = 0
|
||||
|
||||
output.model_output *= np.expand_dims(output.attention_mask, 2).astype(np.float32)
|
||||
norm = np.linalg.norm(output.model_output, ord=2, axis=2, keepdims=True)
|
||||
norm_clamped = np.maximum(norm, 1e-12)
|
||||
output.model_output /= norm_clamped
|
||||
return output.model_output.astype(np.float32)
|
||||
|
||||
def _preprocess_onnx_input(
|
||||
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) -> 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]:
|
||||
# ". " is added to a query to be replaced with a special query 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:
|
||||
prev_padding = None
|
||||
if self.tokenizer.padding:
|
||||
prev_padding = self.tokenizer.padding
|
||||
self.tokenizer.enable_padding(
|
||||
pad_token=self.MASK_TOKEN,
|
||||
pad_id=self.mask_token_id,
|
||||
length=self.MIN_QUERY_LENGTH,
|
||||
)
|
||||
encoded = self.tokenizer.encode_batch(query)
|
||||
if prev_padding is None:
|
||||
self.tokenizer.no_padding()
|
||||
else:
|
||||
self.tokenizer.enable_padding(**prev_padding)
|
||||
return encoded
|
||||
|
||||
def _tokenize_documents(self, documents: List[str]) -> List[Encoding]:
|
||||
# ". " is added to a document to be replaced with a special document 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]]:
|
||||
"""Lists the supported models.
|
||||
|
||||
Returns:
|
||||
List[Dict[str, Any]]: A list of dictionaries containing the model information.
|
||||
"""
|
||||
return supported_colbert_models
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
model_name (str): The name of the model to use.
|
||||
cache_dir (str, optional): The path to the cache directory.
|
||||
Can be set using the `FASTEMBED_CACHE_PATH` env variable.
|
||||
Defaults to `fastembed_cache` in the system's temp directory.
|
||||
threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.
|
||||
|
||||
Raises:
|
||||
ValueError: If the model_name is not in the format <org>/<model> e.g. BAAI/bge-base-en.
|
||||
"""
|
||||
|
||||
super().__init__(model_name, cache_dir, threads, **kwargs)
|
||||
|
||||
model_description = self._get_model_description(model_name)
|
||||
self.cache_dir = define_cache_dir(cache_dir)
|
||||
|
||||
model_dir = self.download_model(
|
||||
model_description, self.cache_dir, local_files_only=self._local_files_only
|
||||
)
|
||||
|
||||
self.load_onnx_model(
|
||||
model_dir=model_dir,
|
||||
model_file=model_description["model_file"],
|
||||
threads=threads,
|
||||
providers=providers,
|
||||
)
|
||||
self.mask_token_id = self.special_token_to_id["[MASK]"]
|
||||
self.pad_token_id = self.tokenizer.padding["pad_id"]
|
||||
|
||||
self.skip_list = {
|
||||
self.tokenizer.encode(symbol, add_special_tokens=False).ids[0]
|
||||
for symbol in string.punctuation
|
||||
}
|
||||
|
||||
def embed(
|
||||
self,
|
||||
documents: Union[str, Iterable[str]],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs,
|
||||
) -> Iterable[np.ndarray]:
|
||||
"""
|
||||
Encode a list of documents into list of embeddings.
|
||||
We use mean pooling with attention so that the model can handle variable-length inputs.
|
||||
|
||||
Args:
|
||||
documents: Iterator of documents or single document to embed
|
||||
batch_size: Batch size for encoding -- higher values will use more memory, but be faster
|
||||
parallel:
|
||||
If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.
|
||||
If 0, use all available cores.
|
||||
If None, don't use data-parallel processing, use default onnxruntime threading instead.
|
||||
|
||||
Returns:
|
||||
List of embeddings, one per document
|
||||
"""
|
||||
yield from self._embed_documents(
|
||||
model_name=self.model_name,
|
||||
cache_dir=str(self.cache_dir),
|
||||
documents=documents,
|
||||
batch_size=batch_size,
|
||||
parallel=parallel,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def query_embed(self, query: Union[str, List[str]], **kwargs) -> np.ndarray:
|
||||
if isinstance(query, str):
|
||||
query = [query]
|
||||
|
||||
for text in query:
|
||||
yield from self._post_process_onnx_output(
|
||||
self.onnx_embed([text], is_doc=False), is_doc=False
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type[TextEmbeddingWorker]:
|
||||
return ColbertEmbeddingWorker
|
||||
|
||||
|
||||
class ColbertEmbeddingWorker(TextEmbeddingWorker):
|
||||
def init_embedding(self, model_name: str, cache_dir: str, **kwargs) -> Colbert:
|
||||
return Colbert(model_name=model_name, cache_dir=cache_dir, threads=1, **kwargs)
|
||||
@@ -0,0 +1,62 @@
|
||||
from typing import Iterable, Optional, Union
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastembed.common.model_management import ModelManagement
|
||||
|
||||
|
||||
class LateInteractionTextEmbeddingBase(ModelManagement):
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
**kwargs,
|
||||
):
|
||||
self.model_name = model_name
|
||||
self.cache_dir = cache_dir
|
||||
self.threads = threads
|
||||
self._local_files_only = kwargs.pop("local_files_only", False)
|
||||
|
||||
def embed(
|
||||
self,
|
||||
documents: Union[str, Iterable[str]],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs,
|
||||
) -> Iterable[np.ndarray]:
|
||||
raise NotImplementedError()
|
||||
|
||||
def passage_embed(self, texts: Iterable[str], **kwargs) -> Iterable[np.ndarray]:
|
||||
"""
|
||||
Embeds a list of text passages into a list of embeddings.
|
||||
|
||||
Args:
|
||||
texts (Iterable[str]): The list of texts to embed.
|
||||
**kwargs: Additional keyword argument to pass to the embed method.
|
||||
|
||||
Yields:
|
||||
Iterable[np.ndarray]: The embeddings.
|
||||
"""
|
||||
|
||||
# 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[np.ndarray]:
|
||||
"""
|
||||
Embeds queries
|
||||
|
||||
Args:
|
||||
query (Union[str, Iterable[str]]): The query to embed, or an iterable e.g. list of queries.
|
||||
|
||||
Returns:
|
||||
Iterable[np.ndarray]: The embeddings.
|
||||
"""
|
||||
|
||||
# This is model-specific, so that different models can have specialized implementations
|
||||
if isinstance(query, str):
|
||||
yield from self.embed([query], **kwargs)
|
||||
if isinstance(query, Iterable):
|
||||
yield from self.embed(query, **kwargs)
|
||||
@@ -0,0 +1,109 @@
|
||||
from typing import Any, Dict, Iterable, List, Optional, Sequence, Type, Union
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastembed.common import OnnxProvider
|
||||
from fastembed.late_interaction.colbert import Colbert
|
||||
from fastembed.late_interaction.late_interaction_embedding_base import (
|
||||
LateInteractionTextEmbeddingBase,
|
||||
)
|
||||
|
||||
|
||||
class LateInteractionTextEmbedding(LateInteractionTextEmbeddingBase):
|
||||
EMBEDDINGS_REGISTRY: List[Type[LateInteractionTextEmbeddingBase]] = [
|
||||
Colbert,
|
||||
]
|
||||
|
||||
@classmethod
|
||||
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.
|
||||
|
||||
Example:
|
||||
```
|
||||
[
|
||||
{
|
||||
"model": "prithvida/SPLADE_PP_en_v1",
|
||||
"vocab_size": 30522,
|
||||
"description": "Independent Implementation of SPLADE++ Model for English",
|
||||
"size_in_GB": 0.532,
|
||||
"sources": {
|
||||
"hf": "qdrant/SPLADE_PP_en_v1",
|
||||
},
|
||||
}
|
||||
]
|
||||
```
|
||||
"""
|
||||
result = []
|
||||
for embedding in cls.EMBEDDINGS_REGISTRY:
|
||||
result.extend(embedding.list_supported_models())
|
||||
return result
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(model_name, cache_dir, threads, **kwargs)
|
||||
|
||||
for EMBEDDING_MODEL_TYPE in self.EMBEDDINGS_REGISTRY:
|
||||
supported_models = EMBEDDING_MODEL_TYPE.list_supported_models()
|
||||
if any(
|
||||
model_name.lower() == model["model"].lower()
|
||||
for model in supported_models
|
||||
):
|
||||
self.model = EMBEDDING_MODEL_TYPE(
|
||||
model_name, cache_dir, threads, providers=providers, **kwargs
|
||||
)
|
||||
return
|
||||
|
||||
raise ValueError(
|
||||
f"Model {model_name} is not supported in SparseTextEmbedding."
|
||||
"Please check the supported models using `SparseTextEmbedding.list_supported_models()`"
|
||||
)
|
||||
|
||||
def embed(
|
||||
self,
|
||||
documents: Union[str, Iterable[str]],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs,
|
||||
) -> Iterable[np.ndarray]:
|
||||
"""
|
||||
Encode a list of documents into list of embeddings.
|
||||
We use mean pooling with attention so that the model can handle variable-length inputs.
|
||||
|
||||
Args:
|
||||
documents: Iterator of documents or single document to embed
|
||||
batch_size: Batch size for encoding -- higher values will use more memory, but be faster
|
||||
parallel:
|
||||
If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.
|
||||
If 0, use all available cores.
|
||||
If None, don't use data-parallel processing, use default onnxruntime threading instead.
|
||||
|
||||
Returns:
|
||||
List of embeddings, one per document
|
||||
"""
|
||||
yield from self.model.embed(documents, batch_size, parallel, **kwargs)
|
||||
|
||||
def query_embed(
|
||||
self, query: Union[str, Iterable[str]], **kwargs
|
||||
) -> Iterable[np.ndarray]:
|
||||
"""
|
||||
Embeds queries
|
||||
|
||||
Args:
|
||||
query (Union[str, Iterable[str]]): The query to embed, or an iterable e.g. list of queries.
|
||||
|
||||
Returns:
|
||||
Iterable[np.ndarray]: The embeddings.
|
||||
"""
|
||||
|
||||
# This is model-specific, so that different models can have specialized implementations
|
||||
yield from self.model.query_embed(query, **kwargs)
|
||||
@@ -83,7 +83,9 @@ def _worker(
|
||||
|
||||
|
||||
class ParallelWorkerPool:
|
||||
def __init__(self, num_workers: int, worker: Type[Worker], start_method: Optional[str] = None):
|
||||
def __init__(
|
||||
self, num_workers: int, worker: Type[Worker], start_method: Optional[str] = None
|
||||
):
|
||||
self.worker_class = worker
|
||||
self.num_workers = num_workers
|
||||
self.input_queue: Optional[Queue] = None
|
||||
@@ -118,7 +120,9 @@ class ParallelWorkerPool:
|
||||
process.start()
|
||||
self.processes.append(process)
|
||||
|
||||
def ordered_map(self, stream: Iterable[Any], *args: Any, **kwargs: Any) -> Iterable[Any]:
|
||||
def ordered_map(
|
||||
self, stream: Iterable[Any], *args: Any, **kwargs: Any
|
||||
) -> Iterable[Any]:
|
||||
buffer = defaultdict(Any)
|
||||
next_expected = 0
|
||||
|
||||
|
||||
@@ -0,0 +1,284 @@
|
||||
import os
|
||||
import string
|
||||
from collections import defaultdict
|
||||
from multiprocessing import get_all_start_methods
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Iterable, List, Optional, Tuple, Type, Union
|
||||
|
||||
import mmh3
|
||||
import numpy as np
|
||||
from snowballstemmer import stemmer as get_stemmer
|
||||
|
||||
from fastembed.common.utils import define_cache_dir, iter_batch
|
||||
from fastembed.parallel_processor import ParallelWorkerPool, Worker
|
||||
from fastembed.sparse.sparse_embedding_base import (
|
||||
SparseEmbedding,
|
||||
SparseTextEmbeddingBase,
|
||||
)
|
||||
from fastembed.sparse.utils.tokenizer import WordTokenizer
|
||||
|
||||
supported_bm25_models = [
|
||||
{
|
||||
"model": "Qdrant/bm25",
|
||||
"description": "BM25 as sparse embeddings meant to be used with Qdrant",
|
||||
"size_in_GB": 0.01,
|
||||
"sources": {
|
||||
"hf": "Qdrant/bm25",
|
||||
},
|
||||
"model_file": "mock.file", # bm25 does not require a model, so we just use a mock
|
||||
"additional_files": ["stopwords.txt"],
|
||||
},
|
||||
]
|
||||
|
||||
MODEL_TO_LANGUAGE = {
|
||||
"Qdrant/bm25": "english",
|
||||
}
|
||||
|
||||
|
||||
class Bm25(SparseTextEmbeddingBase):
|
||||
"""Implements traditional BM25 in a form of sparse embeddings.
|
||||
Uses a count of tokens in the document to evaluate the importance of the token.
|
||||
|
||||
WARNING: This model is expected to be used with `modifier="idf"` in the sparse vector index of Qdrant.
|
||||
|
||||
BM25 formula:
|
||||
|
||||
score(q, d) = SUM[ IDF(q_i) * (f(q_i, d) * (k + 1)) / (f(q_i, d) + k * (1 - b + b * (|d| / avg_len))) ],
|
||||
|
||||
where IDF is the inverse document frequency, computed on Qdrant's side
|
||||
f(q_i, d) is the term frequency of the token q_i in the document d
|
||||
k, b, avg_len are hyperparameters, described below.
|
||||
|
||||
Args:
|
||||
model_name (str): The name of the model to use.
|
||||
cache_dir (str, optional): The path to the cache directory.
|
||||
Can be set using the `FASTEMBED_CACHE_PATH` env variable.
|
||||
Defaults to `fastembed_cache` in the system's temp directory.
|
||||
k (float, optional): The k parameter in the BM25 formula. Defines the saturation of the term frequency.
|
||||
I.e. defines how fast the moment when additional terms stop to increase the score. Defaults to 1.2.
|
||||
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.
|
||||
Raises:
|
||||
ValueError: If the model_name is not in the format <org>/<model> e.g. BAAI/bge-base-en.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: Optional[str] = None,
|
||||
k: float = 1.2,
|
||||
b: float = 0.75,
|
||||
avg_len: float = 256.0,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(model_name, cache_dir, **kwargs)
|
||||
|
||||
self.k = k
|
||||
self.b = b
|
||||
self.avg_len = avg_len
|
||||
|
||||
model_description = self._get_model_description(model_name)
|
||||
self.cache_dir = define_cache_dir(cache_dir)
|
||||
|
||||
model_dir = self.download_model(
|
||||
model_description, self.cache_dir, local_files_only=self._local_files_only
|
||||
)
|
||||
|
||||
self.punctuation = set(string.punctuation)
|
||||
self.stopwords = set(self._load_stopwords(model_dir))
|
||||
self.stemmer = get_stemmer(MODEL_TO_LANGUAGE[model_name])
|
||||
self.tokenizer = WordTokenizer
|
||||
|
||||
@classmethod
|
||||
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.
|
||||
"""
|
||||
return supported_bm25_models
|
||||
|
||||
@classmethod
|
||||
def _load_stopwords(cls, model_dir: Path) -> List[str]:
|
||||
stopwords_path = model_dir / "stopwords.txt"
|
||||
if not stopwords_path.exists():
|
||||
return []
|
||||
|
||||
with open(stopwords_path, "r") as f:
|
||||
return f.read().splitlines()
|
||||
|
||||
def _embed_documents(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
documents: Union[str, Iterable[str]],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
is_small = False
|
||||
|
||||
if isinstance(documents, str):
|
||||
documents = [documents]
|
||||
is_small = True
|
||||
|
||||
if isinstance(documents, list):
|
||||
if len(documents) < batch_size:
|
||||
is_small = True
|
||||
|
||||
if parallel == 0:
|
||||
parallel = os.cpu_count()
|
||||
|
||||
if parallel is None or is_small:
|
||||
for batch in iter_batch(documents, batch_size):
|
||||
yield from self.raw_embed(batch)
|
||||
else:
|
||||
start_method = "forkserver" if "forkserver" in get_all_start_methods() else "spawn"
|
||||
params = {
|
||||
"model_name": model_name,
|
||||
"cache_dir": cache_dir,
|
||||
"k": self.k,
|
||||
"b": self.b,
|
||||
"avg_len": self.avg_len,
|
||||
}
|
||||
pool = ParallelWorkerPool(
|
||||
parallel, self._get_worker_class(), start_method=start_method
|
||||
)
|
||||
for batch in pool.ordered_map(iter_batch(documents, batch_size), **params):
|
||||
for record in batch:
|
||||
yield record
|
||||
|
||||
def embed(
|
||||
self,
|
||||
documents: Union[str, Iterable[str]],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs,
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
"""
|
||||
Encode a list of documents into list of embeddings.
|
||||
We use mean pooling with attention so that the model can handle variable-length inputs.
|
||||
|
||||
Args:
|
||||
documents: Iterator of documents or single document to embed
|
||||
batch_size: Batch size for encoding -- higher values will use more memory, but be faster
|
||||
parallel:
|
||||
If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.
|
||||
If 0, use all available cores.
|
||||
If None, don't use data-parallel processing, use default onnxruntime threading instead.
|
||||
|
||||
Returns:
|
||||
List of embeddings, one per document
|
||||
"""
|
||||
yield from self._embed_documents(
|
||||
model_name=self.model_name,
|
||||
cache_dir=str(self.cache_dir),
|
||||
documents=documents,
|
||||
batch_size=batch_size,
|
||||
parallel=parallel,
|
||||
)
|
||||
|
||||
def _stem(self, tokens: List[str]) -> List[str]:
|
||||
stemmed_tokens = []
|
||||
for token in tokens:
|
||||
if token in self.punctuation:
|
||||
continue
|
||||
|
||||
if token in self.stopwords:
|
||||
continue
|
||||
|
||||
stemmed_token = self.stemmer.stemWord(token)
|
||||
|
||||
if stemmed_token:
|
||||
stemmed_tokens.append(stemmed_token)
|
||||
return stemmed_tokens
|
||||
|
||||
def raw_embed(
|
||||
self,
|
||||
documents: List[str],
|
||||
) -> List[SparseEmbedding]:
|
||||
embeddings = []
|
||||
for document in documents:
|
||||
tokens = self.tokenizer.tokenize(document)
|
||||
stemmed_tokens = self._stem(tokens)
|
||||
token_id2value = self._term_frequency(stemmed_tokens)
|
||||
embeddings.append(SparseEmbedding.from_dict(token_id2value))
|
||||
return embeddings
|
||||
|
||||
def _term_frequency(self, tokens: List[str]) -> Dict[int, float]:
|
||||
"""Calculate the term frequency part of the BM25 formula.
|
||||
|
||||
(
|
||||
f(q_i, d) * (k + 1)
|
||||
) / (
|
||||
f(q_i, d) + k * (1 - b + b * (|d| / avg_len))
|
||||
)
|
||||
|
||||
Args:
|
||||
tokens (List[str]): The list of tokens in the document.
|
||||
|
||||
Returns:
|
||||
Dict[int, float]: The token_id to term frequency mapping.
|
||||
"""
|
||||
tf_map = {}
|
||||
counter = defaultdict(int)
|
||||
for stemmed_token in tokens:
|
||||
counter[stemmed_token] += 1
|
||||
|
||||
doc_len = len(tokens)
|
||||
for stemmed_token in counter:
|
||||
token_id = self.compute_token_id(stemmed_token)
|
||||
num_occurrences = counter[stemmed_token]
|
||||
tf_map[token_id] = num_occurrences * (self.k + 1)
|
||||
tf_map[token_id] /= num_occurrences + self.k * (
|
||||
1 - self.b + self.b * doc_len / self.avg_len
|
||||
)
|
||||
return tf_map
|
||||
|
||||
@classmethod
|
||||
def compute_token_id(cls, token: str) -> int:
|
||||
return abs(mmh3.hash(token))
|
||||
|
||||
def query_embed(self, query: Union[str, Iterable[str]], **kwargs) -> Iterable[SparseEmbedding]:
|
||||
"""To emulate BM25 behaviour, we don't need to use weights in the query, and
|
||||
it's enough to just hash the tokens and assign a weight of 1.0 to them.
|
||||
"""
|
||||
if isinstance(query, str):
|
||||
query = [query]
|
||||
|
||||
for text in query:
|
||||
tokens = self.tokenizer.tokenize(text)
|
||||
stemmed_tokens = self._stem(tokens)
|
||||
token_ids = np.array(
|
||||
[self.compute_token_id(token) for token in stemmed_tokens],
|
||||
dtype=np.float32,
|
||||
)
|
||||
values = np.ones_like(token_ids)
|
||||
yield SparseEmbedding(indices=token_ids, values=values)
|
||||
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type["Bm25Worker"]:
|
||||
return Bm25Worker
|
||||
|
||||
|
||||
class Bm25Worker(Worker):
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
**kwargs,
|
||||
):
|
||||
self.model = self.init_embedding(model_name, cache_dir, **kwargs)
|
||||
|
||||
@classmethod
|
||||
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]]:
|
||||
for idx, batch in items:
|
||||
onnx_output = self.model.raw_embed(batch)
|
||||
yield idx, onnx_output
|
||||
|
||||
@staticmethod
|
||||
def init_embedding(model_name: str, cache_dir: str, **kwargs) -> Bm25:
|
||||
return Bm25(model_name=model_name, cache_dir=cache_dir, **kwargs)
|
||||
@@ -0,0 +1,292 @@
|
||||
import math
|
||||
import string
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Iterable, List, Optional, Sequence, Tuple, Type, Union
|
||||
|
||||
import mmh3
|
||||
import numpy as np
|
||||
from snowballstemmer import stemmer as get_stemmer
|
||||
|
||||
from fastembed.common import OnnxProvider
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.common.utils import define_cache_dir
|
||||
from fastembed.sparse.sparse_embedding_base import (
|
||||
SparseEmbedding,
|
||||
SparseTextEmbeddingBase,
|
||||
)
|
||||
from fastembed.text.onnx_text_model import OnnxTextModel, TextEmbeddingWorker
|
||||
|
||||
supported_bm42_models = [
|
||||
{
|
||||
"model": "Qdrant/bm42-all-minilm-l6-v2-attentions",
|
||||
"vocab_size": 30522,
|
||||
"description": "Light sparse embedding model, which assigns an importance score to each token in the text",
|
||||
"size_in_GB": 0.09,
|
||||
"sources": {
|
||||
"hf": "Qdrant/all_miniLM_L6_v2_with_attentions",
|
||||
},
|
||||
"model_file": "model.onnx",
|
||||
"additional_files": ["stopwords.txt"],
|
||||
},
|
||||
]
|
||||
|
||||
MODEL_TO_LANGUAGE = {
|
||||
"Qdrant/bm42-all-minilm-l6-v2-attentions": "english",
|
||||
}
|
||||
|
||||
|
||||
class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
"""
|
||||
Bm42 is an extension of BM25, which tries to better evaluate importance of tokens in the documents,
|
||||
by extracting attention weights from the transformer model.
|
||||
|
||||
Traditional BM25 uses a count of tokens in the document to evaluate the importance of the token,
|
||||
but this approach doesn't work well with short documents or chunks of text, as almost all tokens
|
||||
there are unique.
|
||||
|
||||
BM42 addresses this issue by replacing the token count with the attention weights from the transformer model.
|
||||
This allows sparse embeddings to work well with short documents, handle rare tokens and leverage traditional NLP
|
||||
techniques like stemming and stopwords.
|
||||
|
||||
WARNING: This model is expected to be used with `modifier="idf"` in the sparse vector index of Qdrant.
|
||||
"""
|
||||
|
||||
ONNX_OUTPUT_NAMES = ["attention_6"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
alpha: float = 0.5,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
model_name (str): The name of the model to use.
|
||||
cache_dir (str, optional): The path to the cache directory.
|
||||
Can be set using the `FASTEMBED_CACHE_PATH` env variable.
|
||||
Defaults to `fastembed_cache` in the system's temp directory.
|
||||
threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.
|
||||
providers (Optional[Sequence[OnnxProvider]], optional): The providers to use for onnxruntime.
|
||||
alpha (float, optional): Parameter, that defines the importance of the token weight in the document
|
||||
versus the importance of the token frequency in the corpus. Defaults to 0.5, based on empirical testing.
|
||||
It is recommended to only change this parameter based on training data for a specific dataset.
|
||||
|
||||
Raises:
|
||||
ValueError: If the model_name is not in the format <org>/<model> e.g. BAAI/bge-base-en.
|
||||
"""
|
||||
|
||||
super().__init__(model_name, cache_dir, threads, **kwargs)
|
||||
|
||||
model_description = self._get_model_description(model_name)
|
||||
self.cache_dir = define_cache_dir(cache_dir)
|
||||
|
||||
model_dir = self.download_model(
|
||||
model_description, self.cache_dir, local_files_only=self._local_files_only
|
||||
)
|
||||
|
||||
self.load_onnx_model(
|
||||
model_dir=model_dir,
|
||||
model_file=model_description["model_file"],
|
||||
threads=threads,
|
||||
providers=providers,
|
||||
)
|
||||
|
||||
self.invert_vocab = {}
|
||||
|
||||
for token, idx in self.tokenizer.get_vocab().items():
|
||||
self.invert_vocab[idx] = token
|
||||
|
||||
self.special_tokens = set(self.special_token_to_id.keys())
|
||||
self.special_tokens_ids = set(self.special_token_to_id.values())
|
||||
self.punctuation = set(string.punctuation)
|
||||
self.stopwords = set(self._load_stopwords(model_dir))
|
||||
self.stemmer = get_stemmer(MODEL_TO_LANGUAGE[model_name])
|
||||
self.alpha = alpha
|
||||
|
||||
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:
|
||||
continue
|
||||
result.append((token, value))
|
||||
return result
|
||||
|
||||
def _stem_pair_tokens(self, tokens: List[Tuple[str, Any]]) -> List[Tuple[str, Any]]:
|
||||
result = []
|
||||
for token, value in tokens:
|
||||
processed_token = self.stemmer.stemWord(token)
|
||||
result.append((processed_token, value))
|
||||
return result
|
||||
|
||||
@classmethod
|
||||
def _aggregate_weights(
|
||||
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)
|
||||
result.append((token, sum_weight))
|
||||
return result
|
||||
|
||||
def _reconstruct_bpe(
|
||||
self, bpe_tokens: Iterable[Tuple[int, str]]
|
||||
) -> List[Tuple[str, List[int]]]:
|
||||
result = []
|
||||
acc = ""
|
||||
acc_idx = []
|
||||
|
||||
continuing_subword_prefix = self.tokenizer.model.continuing_subword_prefix
|
||||
continuing_subword_prefix_len = len(continuing_subword_prefix)
|
||||
|
||||
for idx, token in bpe_tokens:
|
||||
if token in self.special_tokens:
|
||||
continue
|
||||
|
||||
if token.startswith(continuing_subword_prefix):
|
||||
acc += token[continuing_subword_prefix_len:]
|
||||
acc_idx.append(idx)
|
||||
else:
|
||||
if acc:
|
||||
result.append((acc, acc_idx))
|
||||
acc_idx = []
|
||||
acc = token
|
||||
acc_idx.append(idx)
|
||||
|
||||
if acc:
|
||||
result.append((acc, acc_idx))
|
||||
|
||||
return result
|
||||
|
||||
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.
|
||||
"""
|
||||
|
||||
new_vector = {}
|
||||
|
||||
for token, value in vector.items():
|
||||
token_id = abs(mmh3.hash(token))
|
||||
# Examples:
|
||||
# Num 0: Log(1/1 + 1) = 0.6931471805599453
|
||||
# Num 1: Log(1/2 + 1) = 0.4054651081081644
|
||||
# Num 2: Log(1/3 + 1) = 0.28768207245178085
|
||||
new_vector[token_id] = math.log(1.0 + value) ** self.alpha # value
|
||||
|
||||
return new_vector
|
||||
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[SparseEmbedding]:
|
||||
token_ids_batch = output.input_ids
|
||||
|
||||
# attention_value shape: (batch_size, num_heads, num_tokens, num_tokens)
|
||||
pooled_attention = np.mean(output.model_output[:, :, 0], axis=1) * output.attention_mask
|
||||
|
||||
for document_token_ids, attention_value in zip(token_ids_batch, pooled_attention):
|
||||
document_tokens_with_ids = (
|
||||
(idx, self.invert_vocab[token_id])
|
||||
for idx, token_id in enumerate(document_token_ids)
|
||||
)
|
||||
|
||||
reconstructed = self._reconstruct_bpe(document_tokens_with_ids)
|
||||
|
||||
filtered = self._filter_pair_tokens(reconstructed)
|
||||
|
||||
stemmed = self._stem_pair_tokens(filtered)
|
||||
|
||||
weighted = self._aggregate_weights(stemmed, attention_value)
|
||||
|
||||
max_token_weight = {}
|
||||
|
||||
for token, weight in weighted:
|
||||
max_token_weight[token] = max(max_token_weight.get(token, 0), weight)
|
||||
|
||||
rescored = self._rescore_vector(max_token_weight)
|
||||
|
||||
yield SparseEmbedding.from_dict(rescored)
|
||||
|
||||
@classmethod
|
||||
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.
|
||||
"""
|
||||
return supported_bm42_models
|
||||
|
||||
@classmethod
|
||||
def _load_stopwords(cls, model_dir: Path) -> List[str]:
|
||||
stopwords_path = model_dir / "stopwords.txt"
|
||||
if not stopwords_path.exists():
|
||||
return []
|
||||
|
||||
with open(stopwords_path, "r") as f:
|
||||
return f.read().splitlines()
|
||||
|
||||
def embed(
|
||||
self,
|
||||
documents: Union[str, Iterable[str]],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs,
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
"""
|
||||
Encode a list of documents into list of embeddings.
|
||||
We use mean pooling with attention so that the model can handle variable-length inputs.
|
||||
|
||||
Args:
|
||||
documents: Iterator of documents or single document to embed
|
||||
batch_size: Batch size for encoding -- higher values will use more memory, but be faster
|
||||
parallel:
|
||||
If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.
|
||||
If 0, use all available cores.
|
||||
If None, don't use data-parallel processing, use default onnxruntime threading instead.
|
||||
|
||||
Returns:
|
||||
List of embeddings, one per document
|
||||
"""
|
||||
yield from self._embed_documents(
|
||||
model_name=self.model_name,
|
||||
cache_dir=str(self.cache_dir),
|
||||
documents=documents,
|
||||
batch_size=batch_size,
|
||||
parallel=parallel,
|
||||
alpha=self.alpha,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _query_rehash(cls, tokens: Iterable[str]) -> Dict[int, float]:
|
||||
result = {}
|
||||
for token in tokens:
|
||||
token_id = abs(mmh3.hash(token))
|
||||
result[token_id] = 1.0
|
||||
return result
|
||||
|
||||
def query_embed(self, query: Union[str, Iterable[str]], **kwargs) -> Iterable[SparseEmbedding]:
|
||||
"""
|
||||
To emulate BM25 behaviour, we don't need to use smart weights in the query, and
|
||||
it's enough to just hash the tokens and assign a weight of 1.0 to them.
|
||||
It is also faster, as we don't need to run the model for the query.
|
||||
"""
|
||||
if isinstance(query, str):
|
||||
query = [query]
|
||||
|
||||
for text in query:
|
||||
encoded = self.tokenizer.encode(text)
|
||||
document_tokens_with_ids = enumerate(encoded.tokens)
|
||||
reconstructed = self._reconstruct_bpe(document_tokens_with_ids)
|
||||
filtered = self._filter_pair_tokens(reconstructed)
|
||||
stemmed = self._stem_pair_tokens(filtered)
|
||||
|
||||
yield SparseEmbedding.from_dict(self._query_rehash(token for token, _ in stemmed))
|
||||
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type[TextEmbeddingWorker]:
|
||||
return Bm42TextEmbeddingWorker
|
||||
|
||||
|
||||
class Bm42TextEmbeddingWorker(TextEmbeddingWorker):
|
||||
def init_embedding(self, model_name: str, cache_dir: str, **kwargs) -> Bm42:
|
||||
return Bm42(model_name=model_name, cache_dir=cache_dir, **kwargs)
|
||||
@@ -20,6 +20,11 @@ class SparseEmbedding:
|
||||
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":
|
||||
indices, values = zip(*data.items())
|
||||
return cls(values=np.array(values), indices=np.array(indices))
|
||||
|
||||
|
||||
class SparseTextEmbeddingBase(ModelManagement):
|
||||
def __init__(
|
||||
@@ -32,6 +37,7 @@ class SparseTextEmbeddingBase(ModelManagement):
|
||||
self.model_name = model_name
|
||||
self.cache_dir = cache_dir
|
||||
self.threads = threads
|
||||
self._local_files_only = kwargs.pop("local_files_only", False)
|
||||
|
||||
def embed(
|
||||
self,
|
||||
@@ -41,3 +47,39 @@ class SparseTextEmbeddingBase(ModelManagement):
|
||||
**kwargs,
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
raise NotImplementedError()
|
||||
|
||||
def passage_embed(
|
||||
self, texts: Iterable[str], **kwargs
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
"""
|
||||
Embeds a list of text passages into a list of embeddings.
|
||||
|
||||
Args:
|
||||
texts (Iterable[str]): The list of texts to embed.
|
||||
**kwargs: Additional keyword argument to pass to the embed method.
|
||||
|
||||
Yields:
|
||||
Iterable[SparseEmbedding]: The sparse embeddings.
|
||||
"""
|
||||
|
||||
# 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]:
|
||||
"""
|
||||
Embeds queries
|
||||
|
||||
Args:
|
||||
query (Union[str, Iterable[str]]): The query to embed, or an iterable e.g. list of queries.
|
||||
|
||||
Returns:
|
||||
Iterable[SparseEmbedding]: The sparse embeddings.
|
||||
"""
|
||||
|
||||
# This is model-specific, so that different models can have specialized implementations
|
||||
if isinstance(query, str):
|
||||
yield from self.embed([query], **kwargs)
|
||||
if isinstance(query, Iterable):
|
||||
yield from self.embed(query, **kwargs)
|
||||
|
||||
@@ -1,13 +1,17 @@
|
||||
from typing import List, Type, Dict, Any, Union, Iterable, Optional
|
||||
from typing import Any, Dict, Iterable, List, Optional, Sequence, Type, Union
|
||||
|
||||
from fastembed.sparse.sparse_embedding_base import SparseTextEmbeddingBase, SparseEmbedding
|
||||
from fastembed.common import OnnxProvider
|
||||
from fastembed.sparse.bm25 import Bm25
|
||||
from fastembed.sparse.bm42 import Bm42
|
||||
from fastembed.sparse.sparse_embedding_base import (
|
||||
SparseEmbedding,
|
||||
SparseTextEmbeddingBase,
|
||||
)
|
||||
from fastembed.sparse.splade_pp import SpladePP
|
||||
|
||||
|
||||
class SparseTextEmbedding(SparseTextEmbeddingBase):
|
||||
EMBEDDINGS_REGISTRY: List[Type[SparseTextEmbeddingBase]] = [
|
||||
SpladePP,
|
||||
]
|
||||
EMBEDDINGS_REGISTRY: List[Type[SparseTextEmbeddingBase]] = [SpladePP, Bm42, Bm25]
|
||||
|
||||
@classmethod
|
||||
def list_supported_models(cls) -> List[Dict[str, Any]]:
|
||||
@@ -42,14 +46,24 @@ class SparseTextEmbedding(SparseTextEmbeddingBase):
|
||||
model_name: str,
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(model_name, cache_dir, threads, **kwargs)
|
||||
|
||||
for EMBEDDING_MODEL_TYPE in self.EMBEDDINGS_REGISTRY:
|
||||
supported_models = EMBEDDING_MODEL_TYPE.list_supported_models()
|
||||
if any(model_name.lower() == model["model"].lower() for model in supported_models):
|
||||
self.model = EMBEDDING_MODEL_TYPE(model_name, cache_dir, threads, **kwargs)
|
||||
if any(
|
||||
model_name.lower() == model["model"].lower()
|
||||
for model in supported_models
|
||||
):
|
||||
self.model = EMBEDDING_MODEL_TYPE(
|
||||
model_name,
|
||||
cache_dir,
|
||||
threads=threads,
|
||||
providers=providers,
|
||||
**kwargs,
|
||||
)
|
||||
return
|
||||
|
||||
raise ValueError(
|
||||
@@ -80,3 +94,17 @@ class SparseTextEmbedding(SparseTextEmbeddingBase):
|
||||
List of embeddings, one per document
|
||||
"""
|
||||
yield from self.model.embed(documents, batch_size, parallel, **kwargs)
|
||||
|
||||
def query_embed(
|
||||
self, query: Union[str, Iterable[str]], **kwargs
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
"""
|
||||
Embeds queries
|
||||
|
||||
Args:
|
||||
query (Union[str, Iterable[str]]): The query to embed, or an iterable e.g. list of queries.
|
||||
|
||||
Returns:
|
||||
Iterable[SparseEmbedding]: The sparse embeddings.
|
||||
"""
|
||||
yield from self.model.query_embed(query, **kwargs)
|
||||
|
||||
@@ -1,10 +1,15 @@
|
||||
from typing import Any, Dict, Iterable, List, Optional, Tuple, Union, Type
|
||||
from typing import Any, Dict, Iterable, List, Optional, Sequence, Type, Union
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastembed.common.onnx_model import EmbeddingWorker, OnnxModel
|
||||
from fastembed.common import OnnxProvider
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.common.utils import define_cache_dir
|
||||
from fastembed.sparse.sparse_embedding_base import SparseEmbedding, SparseTextEmbeddingBase
|
||||
from fastembed.sparse.sparse_embedding_base import (
|
||||
SparseEmbedding,
|
||||
SparseTextEmbeddingBase,
|
||||
)
|
||||
from fastembed.text.onnx_text_model import OnnxTextModel, TextEmbeddingWorker
|
||||
|
||||
supported_splade_models = [
|
||||
{
|
||||
@@ -30,15 +35,11 @@ supported_splade_models = [
|
||||
]
|
||||
|
||||
|
||||
class SpladePP(SparseTextEmbeddingBase, OnnxModel[SparseEmbedding]):
|
||||
@classmethod
|
||||
def _post_process_onnx_output(
|
||||
cls, output: Tuple[np.ndarray, np.ndarray]
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
logits, attention_mask = output
|
||||
relu_log = np.log(1 + np.maximum(logits, 0))
|
||||
class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[SparseEmbedding]:
|
||||
relu_log = np.log(1 + np.maximum(output.model_output, 0))
|
||||
|
||||
weighted_log = relu_log * np.expand_dims(attention_mask, axis=-1)
|
||||
weighted_log = relu_log * np.expand_dims(output.attention_mask, axis=-1)
|
||||
|
||||
scores = np.max(weighted_log, axis=1)
|
||||
|
||||
@@ -63,6 +64,7 @@ class SpladePP(SparseTextEmbeddingBase, OnnxModel[SparseEmbedding]):
|
||||
model_name: str,
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
@@ -80,14 +82,17 @@ class SpladePP(SparseTextEmbeddingBase, OnnxModel[SparseEmbedding]):
|
||||
super().__init__(model_name, cache_dir, threads, **kwargs)
|
||||
|
||||
model_description = self._get_model_description(model_name)
|
||||
cache_dir = define_cache_dir(cache_dir)
|
||||
self.cache_dir = define_cache_dir(cache_dir)
|
||||
|
||||
model_dir = self.download_model(model_description, cache_dir)
|
||||
model_dir = self.download_model(
|
||||
model_description, self.cache_dir, local_files_only=self._local_files_only
|
||||
)
|
||||
|
||||
self.load_onnx_model(
|
||||
model_dir=model_dir,
|
||||
model_file=model_description["model_file"],
|
||||
threads=threads,
|
||||
providers=providers,
|
||||
)
|
||||
|
||||
def embed(
|
||||
@@ -121,14 +126,10 @@ class SpladePP(SparseTextEmbeddingBase, OnnxModel[SparseEmbedding]):
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type[EmbeddingWorker]:
|
||||
def _get_worker_class(cls) -> Type[TextEmbeddingWorker]:
|
||||
return SpladePPEmbeddingWorker
|
||||
|
||||
|
||||
class SpladePPEmbeddingWorker(EmbeddingWorker):
|
||||
def init_embedding(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
) -> SpladePP:
|
||||
return SpladePP(model_name=model_name, cache_dir=cache_dir, threads=1)
|
||||
class SpladePPEmbeddingWorker(TextEmbeddingWorker):
|
||||
def init_embedding(self, model_name: str, cache_dir: str, **kwargs) -> SpladePP:
|
||||
return SpladePP(model_name=model_name, cache_dir=cache_dir, threads=1, **kwargs)
|
||||
|
||||
@@ -0,0 +1,113 @@
|
||||
# This code is a modified copy of the `NLTKWordTokenizer` class from `NLTK` library.
|
||||
|
||||
import re
|
||||
from typing import List
|
||||
|
||||
|
||||
class WordTokenizer:
|
||||
"""The tokenizer is "destructive" such that the regexes applied will munge the
|
||||
input string to a state beyond re-construction.
|
||||
"""
|
||||
|
||||
# Starting quotes.
|
||||
STARTING_QUOTES = [
|
||||
(re.compile("([«“‘„]|[`]+)", re.U), r" \1 "),
|
||||
(re.compile(r"^\""), r"``"),
|
||||
(re.compile(r"(``)"), r" \1 "),
|
||||
(re.compile(r"([ \(\[{<])(\"|\'{2})"), r"\1 `` "),
|
||||
(re.compile(r"(?i)(\')(?!re|ve|ll|m|t|s|d|n)(\w)\b", re.U), r"\1 \2"),
|
||||
]
|
||||
|
||||
# Ending quotes.
|
||||
ENDING_QUOTES = [
|
||||
(re.compile("([»”’])", re.U), r" \1 "),
|
||||
(re.compile(r"''"), " '' "),
|
||||
(re.compile(r'"'), " '' "),
|
||||
(re.compile(r"([^' ])('[sS]|'[mM]|'[dD]|') "), r"\1 \2 "),
|
||||
(re.compile(r"([^' ])('ll|'LL|'re|'RE|'ve|'VE|n't|N'T) "), r"\1 \2 "),
|
||||
]
|
||||
|
||||
# Punctuation.
|
||||
PUNCTUATION = [
|
||||
(re.compile(r'([^\.])(\.)([\]\)}>"\'' "»”’ " r"]*)\s*$", re.U), r"\1 \2 \3 "),
|
||||
(re.compile(r"([:,])([^\d])"), r" \1 \2"),
|
||||
(re.compile(r"([:,])$"), r" \1 "),
|
||||
(
|
||||
re.compile(r"\.{2,}", re.U),
|
||||
r" \g<0> ",
|
||||
),
|
||||
(re.compile(r"[;@#$%&]"), r" \g<0> "),
|
||||
(
|
||||
re.compile(r'([^\.])(\.)([\]\)}>"\']*)\s*$'),
|
||||
r"\1 \2\3 ",
|
||||
), # Handles the final period.
|
||||
(re.compile(r"[?!]"), r" \g<0> "),
|
||||
(re.compile(r"([^'])' "), r"\1 ' "),
|
||||
(
|
||||
re.compile(r"[*]", re.U),
|
||||
r" \g<0> ",
|
||||
),
|
||||
]
|
||||
|
||||
# Pads parentheses
|
||||
PARENS_BRACKETS = (re.compile(r"[\]\[\(\)\{\}\<\>]"), r" \g<0> ")
|
||||
DOUBLE_DASHES = (re.compile(r"--"), r" -- ")
|
||||
|
||||
# List of contractions adapted from Robert MacIntyre's tokenizer.
|
||||
CONTRACTIONS2 = [
|
||||
re.compile(pattern)
|
||||
for pattern in (
|
||||
r"(?i)\b(can)(?#X)(not)\b",
|
||||
r"(?i)\b(d)(?#X)('ye)\b",
|
||||
r"(?i)\b(gim)(?#X)(me)\b",
|
||||
r"(?i)\b(gon)(?#X)(na)\b",
|
||||
r"(?i)\b(got)(?#X)(ta)\b",
|
||||
r"(?i)\b(lem)(?#X)(me)\b",
|
||||
r"(?i)\b(more)(?#X)('n)\b",
|
||||
r"(?i)\b(wan)(?#X)(na)(?=\s)",
|
||||
)
|
||||
]
|
||||
CONTRACTIONS3 = [
|
||||
re.compile(pattern)
|
||||
for pattern in (r"(?i) ('t)(?#X)(is)\b", r"(?i) ('t)(?#X)(was)\b")
|
||||
]
|
||||
|
||||
@classmethod
|
||||
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.'''
|
||||
>>> WordTokenizer().tokenize(s)
|
||||
['Good', 'muffins', 'cost', '$', '3.88', '(', 'roughly', '3,36', 'euros', ')', 'in', 'New', 'York', '.']
|
||||
|
||||
Args:
|
||||
text: The text to be tokenized.
|
||||
|
||||
Returns:
|
||||
A list of tokens.
|
||||
"""
|
||||
for regexp, substitution in cls.STARTING_QUOTES:
|
||||
text = regexp.sub(substitution, text)
|
||||
|
||||
for regexp, substitution in cls.PUNCTUATION:
|
||||
text = regexp.sub(substitution, text)
|
||||
|
||||
# Handles parentheses.
|
||||
regexp, substitution = cls.PARENS_BRACKETS
|
||||
text = regexp.sub(substitution, text)
|
||||
|
||||
# Handles double dash.
|
||||
regexp, substitution = cls.DOUBLE_DASHES
|
||||
text = regexp.sub(substitution, text)
|
||||
|
||||
# add extra space to make things easier
|
||||
text = " " + text + " "
|
||||
|
||||
for regexp, substitution in cls.ENDING_QUOTES:
|
||||
text = regexp.sub(substitution, text)
|
||||
|
||||
for regexp in cls.CONTRACTIONS2:
|
||||
text = regexp.sub(r" \1 \2 ", text)
|
||||
for regexp in cls.CONTRACTIONS3:
|
||||
text = regexp.sub(r" \1 \2 ", text)
|
||||
return text.split()
|
||||
@@ -0,0 +1,49 @@
|
||||
from typing import Any, Dict, Iterable, List, Type
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.text.onnx_embedding import OnnxTextEmbedding, OnnxTextEmbeddingWorker
|
||||
from fastembed.text.onnx_text_model import TextEmbeddingWorker
|
||||
|
||||
supported_clip_models = [
|
||||
{
|
||||
"model": "Qdrant/clip-ViT-B-32-text",
|
||||
"dim": 512,
|
||||
"description": "CLIP text encoder",
|
||||
"size_in_GB": 0.25,
|
||||
"sources": {
|
||||
"hf": "Qdrant/clip-ViT-B-32-text",
|
||||
},
|
||||
"model_file": "model.onnx",
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
class CLIPOnnxEmbedding(OnnxTextEmbedding):
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type[TextEmbeddingWorker]:
|
||||
return CLIPEmbeddingWorker
|
||||
|
||||
@classmethod
|
||||
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.
|
||||
"""
|
||||
return supported_clip_models
|
||||
|
||||
def _post_process_onnx_output(
|
||||
self, output: OnnxOutputContext
|
||||
) -> Iterable[np.ndarray]:
|
||||
return output.model_output
|
||||
|
||||
|
||||
class CLIPEmbeddingWorker(OnnxTextEmbeddingWorker):
|
||||
def init_embedding(
|
||||
self, model_name: str, cache_dir: str, **kwargs
|
||||
) -> OnnxTextEmbedding:
|
||||
return CLIPOnnxEmbedding(
|
||||
model_name=model_name, cache_dir=cache_dir, threads=1, **kwargs
|
||||
)
|
||||
@@ -1,9 +1,9 @@
|
||||
from typing import Type, List, Dict, Any
|
||||
from typing import Any, Dict, List, Type
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastembed.common.onnx_model import EmbeddingWorker
|
||||
from fastembed.text.onnx_embedding import OnnxTextEmbedding, OnnxTextEmbeddingWorker
|
||||
from fastembed.text.onnx_text_model import TextEmbeddingWorker
|
||||
|
||||
supported_multilingual_e5_models = [
|
||||
{
|
||||
@@ -33,7 +33,7 @@ supported_multilingual_e5_models = [
|
||||
|
||||
class E5OnnxEmbedding(OnnxTextEmbedding):
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type["EmbeddingWorker"]:
|
||||
def _get_worker_class(cls) -> Type["TextEmbeddingWorker"]:
|
||||
return E5OnnxEmbeddingWorker
|
||||
|
||||
@classmethod
|
||||
@@ -45,7 +45,9 @@ class E5OnnxEmbedding(OnnxTextEmbedding):
|
||||
"""
|
||||
return supported_multilingual_e5_models
|
||||
|
||||
def _preprocess_onnx_input(self, onnx_input: Dict[str, np.ndarray]) -> Dict[str, np.ndarray]:
|
||||
def _preprocess_onnx_input(
|
||||
self, onnx_input: Dict[str, np.ndarray], **kwargs
|
||||
) -> Dict[str, np.ndarray]:
|
||||
"""
|
||||
Preprocess the onnx input.
|
||||
"""
|
||||
@@ -55,8 +57,8 @@ class E5OnnxEmbedding(OnnxTextEmbedding):
|
||||
|
||||
class E5OnnxEmbeddingWorker(OnnxTextEmbeddingWorker):
|
||||
def init_embedding(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
self, model_name: str, cache_dir: str, **kwargs
|
||||
) -> E5OnnxEmbedding:
|
||||
return E5OnnxEmbedding(model_name=model_name, cache_dir=cache_dir, threads=1)
|
||||
return E5OnnxEmbedding(
|
||||
model_name=model_name, cache_dir=cache_dir, threads=1, **kwargs
|
||||
)
|
||||
|
||||
@@ -1,10 +1,11 @@
|
||||
from typing import Type, List, Dict, Any, Tuple, Iterable
|
||||
from typing import Any, Dict, Iterable, List, Type
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastembed.common.models import normalize
|
||||
from fastembed.common.onnx_model import EmbeddingWorker
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.common.utils import normalize
|
||||
from fastembed.text.onnx_embedding import OnnxTextEmbedding, OnnxTextEmbeddingWorker
|
||||
from fastembed.text.onnx_text_model import TextEmbeddingWorker
|
||||
|
||||
supported_jina_models = [
|
||||
{
|
||||
@@ -23,12 +24,20 @@ supported_jina_models = [
|
||||
"sources": {"hf": "xenova/jina-embeddings-v2-small-en"},
|
||||
"model_file": "onnx/model.onnx",
|
||||
},
|
||||
{
|
||||
"model": "jinaai/jina-embeddings-v2-base-de",
|
||||
"dim": 768,
|
||||
"description": "German embedding model supporting 8192 sequence length",
|
||||
"size_in_GB": 0.32,
|
||||
"sources": {"hf": "jinaai/jina-embeddings-v2-base-de"},
|
||||
"model_file": "onnx/model_fp16.onnx",
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
class JinaOnnxEmbedding(OnnxTextEmbedding):
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type[EmbeddingWorker]:
|
||||
def _get_worker_class(cls) -> Type[TextEmbeddingWorker]:
|
||||
return JinaEmbeddingWorker
|
||||
|
||||
@classmethod
|
||||
@@ -50,18 +59,18 @@ class JinaOnnxEmbedding(OnnxTextEmbedding):
|
||||
"""
|
||||
return supported_jina_models
|
||||
|
||||
@classmethod
|
||||
def _post_process_onnx_output(
|
||||
cls, output: Tuple[np.ndarray, np.ndarray]
|
||||
self, output: OnnxOutputContext
|
||||
) -> Iterable[np.ndarray]:
|
||||
embeddings, attn_mask = output
|
||||
return normalize(cls.mean_pooling(embeddings, attn_mask)).astype(np.float32)
|
||||
embeddings = output.model_output
|
||||
attn_mask = output.attention_mask
|
||||
return normalize(self.mean_pooling(embeddings, attn_mask)).astype(np.float32)
|
||||
|
||||
|
||||
class JinaEmbeddingWorker(OnnxTextEmbeddingWorker):
|
||||
def init_embedding(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
self, model_name: str, cache_dir: str, **kwargs
|
||||
) -> OnnxTextEmbedding:
|
||||
return JinaOnnxEmbedding(model_name=model_name, cache_dir=cache_dir, threads=1)
|
||||
return JinaOnnxEmbedding(
|
||||
model_name=model_name, cache_dir=cache_dir, threads=1, **kwargs
|
||||
)
|
||||
|
||||
@@ -0,0 +1,58 @@
|
||||
from typing import Any, Dict, Iterable, List, Type
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.common.utils import normalize
|
||||
from fastembed.text.onnx_embedding import OnnxTextEmbedding, OnnxTextEmbeddingWorker
|
||||
from fastembed.text.onnx_text_model import TextEmbeddingWorker
|
||||
|
||||
supported_mini_lm_models = [
|
||||
{
|
||||
"model": "sentence-transformers/all-MiniLM-L6-v2",
|
||||
"dim": 384,
|
||||
"description": "Sentence Transformer model, MiniLM-L6-v2",
|
||||
"size_in_GB": 0.09,
|
||||
"sources": {
|
||||
"url": "https://storage.googleapis.com/qdrant-fastembed/sentence-transformers-all-MiniLM-L6-v2.tar.gz",
|
||||
"hf": "qdrant/all-MiniLM-L6-v2-onnx",
|
||||
},
|
||||
"model_file": "model.onnx",
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
class MiniLMOnnxEmbedding(OnnxTextEmbedding):
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type[TextEmbeddingWorker]:
|
||||
return MiniLMEmbeddingWorker
|
||||
|
||||
@classmethod
|
||||
def mean_pooling(cls, model_output: np.ndarray, attention_mask: np.ndarray) -> np.ndarray:
|
||||
token_embeddings = model_output
|
||||
input_mask_expanded = np.expand_dims(attention_mask, axis=-1)
|
||||
input_mask_expanded = np.tile(input_mask_expanded, (1, 1, token_embeddings.shape[-1]))
|
||||
input_mask_expanded = input_mask_expanded.astype(float)
|
||||
sum_embeddings = np.sum(token_embeddings * input_mask_expanded, axis=1)
|
||||
sum_mask = np.sum(input_mask_expanded, axis=1)
|
||||
pooled_embeddings = sum_embeddings / np.maximum(sum_mask, 1e-9)
|
||||
return pooled_embeddings
|
||||
|
||||
@classmethod
|
||||
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.
|
||||
"""
|
||||
return supported_mini_lm_models
|
||||
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[np.ndarray]:
|
||||
embeddings = output.model_output
|
||||
attn_mask = output.attention_mask
|
||||
return normalize(self.mean_pooling(embeddings, attn_mask)).astype(np.float32)
|
||||
|
||||
|
||||
class MiniLMEmbeddingWorker(OnnxTextEmbeddingWorker):
|
||||
def init_embedding(self, model_name: str, cache_dir: str, **kwargs) -> OnnxTextEmbedding:
|
||||
return MiniLMOnnxEmbedding(model_name=model_name, cache_dir=cache_dir, threads=1, **kwargs)
|
||||
@@ -1,10 +1,11 @@
|
||||
from typing import Dict, Optional, Tuple, Union, Iterable, Type, List, Any
|
||||
from typing import Any, Dict, Iterable, List, Optional, Sequence, Type, Union
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastembed.common.onnx_model import OnnxModel, EmbeddingWorker
|
||||
from fastembed.common.models import normalize
|
||||
from fastembed.common.utils import define_cache_dir
|
||||
from fastembed.common import OnnxProvider
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.common.utils import define_cache_dir, normalize
|
||||
from fastembed.text.onnx_text_model import OnnxTextModel, TextEmbeddingWorker
|
||||
from fastembed.text.text_embedding_base import TextEmbeddingBase
|
||||
|
||||
supported_onnx_models = [
|
||||
@@ -69,17 +70,6 @@ supported_onnx_models = [
|
||||
},
|
||||
"model_file": "model_optimized.onnx",
|
||||
},
|
||||
{
|
||||
"model": "sentence-transformers/all-MiniLM-L6-v2",
|
||||
"dim": 384,
|
||||
"description": "Sentence Transformer model, MiniLM-L6-v2",
|
||||
"size_in_GB": 0.09,
|
||||
"sources": {
|
||||
"url": "https://storage.googleapis.com/qdrant-fastembed/sentence-transformers-all-MiniLM-L6-v2.tar.gz",
|
||||
"hf": "qdrant/all-MiniLM-L6-v2-onnx",
|
||||
},
|
||||
"model_file": "model.onnx",
|
||||
},
|
||||
{
|
||||
"model": "sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2",
|
||||
"dim": 384,
|
||||
@@ -193,7 +183,7 @@ supported_onnx_models = [
|
||||
]
|
||||
|
||||
|
||||
class OnnxTextEmbedding(TextEmbeddingBase, OnnxModel[np.ndarray]):
|
||||
class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[np.ndarray]):
|
||||
"""Implementation of the Flag Embedding model."""
|
||||
|
||||
@classmethod
|
||||
@@ -211,6 +201,7 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxModel[np.ndarray]):
|
||||
model_name: str = "BAAI/bge-small-en-v1.5",
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
@@ -228,13 +219,16 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxModel[np.ndarray]):
|
||||
super().__init__(model_name, cache_dir, threads, **kwargs)
|
||||
|
||||
model_description = self._get_model_description(model_name)
|
||||
cache_dir = define_cache_dir(cache_dir)
|
||||
model_dir = self.download_model(model_description, cache_dir)
|
||||
self.cache_dir = define_cache_dir(cache_dir)
|
||||
model_dir = self.download_model(
|
||||
model_description, self.cache_dir, local_files_only=self._local_files_only
|
||||
)
|
||||
|
||||
self.load_onnx_model(
|
||||
model_dir=model_dir,
|
||||
model_file=model_description["model_file"],
|
||||
threads=threads,
|
||||
providers=providers,
|
||||
)
|
||||
|
||||
def embed(
|
||||
@@ -265,30 +259,31 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxModel[np.ndarray]):
|
||||
documents=documents,
|
||||
batch_size=batch_size,
|
||||
parallel=parallel,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type["EmbeddingWorker"]:
|
||||
def _get_worker_class(cls) -> Type["TextEmbeddingWorker"]:
|
||||
return OnnxTextEmbeddingWorker
|
||||
|
||||
def _preprocess_onnx_input(self, onnx_input: Dict[str, np.ndarray]) -> Dict[str, np.ndarray]:
|
||||
def _preprocess_onnx_input(
|
||||
self, onnx_input: Dict[str, np.ndarray], **kwargs
|
||||
) -> Dict[str, np.ndarray]:
|
||||
"""
|
||||
Preprocess the onnx input.
|
||||
"""
|
||||
return onnx_input
|
||||
|
||||
@classmethod
|
||||
def _post_process_onnx_output(
|
||||
cls, output: Tuple[np.ndarray, np.ndarray]
|
||||
) -> Iterable[np.ndarray]:
|
||||
embeddings, _ = output
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[np.ndarray]:
|
||||
embeddings = output.model_output
|
||||
return normalize(embeddings[:, 0]).astype(np.float32)
|
||||
|
||||
|
||||
class OnnxTextEmbeddingWorker(EmbeddingWorker):
|
||||
class OnnxTextEmbeddingWorker(TextEmbeddingWorker):
|
||||
def init_embedding(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
**kwargs,
|
||||
) -> OnnxTextEmbedding:
|
||||
return OnnxTextEmbedding(model_name=model_name, cache_dir=cache_dir, threads=1)
|
||||
return OnnxTextEmbedding(model_name=model_name, cache_dir=cache_dir, threads=1, **kwargs)
|
||||
|
||||
@@ -0,0 +1,126 @@
|
||||
import os
|
||||
from multiprocessing import get_all_start_methods
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Iterable, List, Optional, Sequence, Tuple, Type, Union
|
||||
|
||||
import numpy as np
|
||||
from tokenizers import Encoding
|
||||
|
||||
from fastembed.common import OnnxProvider
|
||||
from fastembed.common.onnx_model import EmbeddingWorker, OnnxModel, OnnxOutputContext, T
|
||||
from fastembed.common.preprocessor_utils import load_tokenizer
|
||||
from fastembed.common.utils import iter_batch
|
||||
from fastembed.parallel_processor import ParallelWorkerPool
|
||||
|
||||
|
||||
class OnnxTextModel(OnnxModel[T]):
|
||||
ONNX_OUTPUT_NAMES: Optional[List[str]] = None
|
||||
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type["TextEmbeddingWorker"]:
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[T]:
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.tokenizer = None
|
||||
self.special_token_to_id = {}
|
||||
|
||||
def _preprocess_onnx_input(
|
||||
self, onnx_input: Dict[str, np.ndarray], **kwargs
|
||||
) -> Dict[str, np.ndarray]:
|
||||
"""
|
||||
Preprocess the onnx input.
|
||||
"""
|
||||
return onnx_input
|
||||
|
||||
def load_onnx_model(
|
||||
self,
|
||||
model_dir: Path,
|
||||
model_file: str,
|
||||
threads: Optional[int],
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
) -> None:
|
||||
super().load_onnx_model(
|
||||
model_dir=model_dir,
|
||||
model_file=model_file,
|
||||
threads=threads,
|
||||
providers=providers,
|
||||
)
|
||||
self.tokenizer, self.special_token_to_id = load_tokenizer(model_dir=model_dir)
|
||||
|
||||
def tokenize(self, documents: List[str], **kwargs) -> List[Encoding]:
|
||||
return self.tokenizer.encode_batch(documents)
|
||||
|
||||
def onnx_embed(
|
||||
self,
|
||||
documents: List[str],
|
||||
**kwargs,
|
||||
) -> OnnxOutputContext:
|
||||
encoded = self.tokenize(documents, **kwargs)
|
||||
input_ids = np.array([e.ids for e in encoded])
|
||||
attention_mask = np.array([e.attention_mask for e in encoded])
|
||||
input_names = {node.name for node in self.model.get_inputs()}
|
||||
onnx_input = {
|
||||
"input_ids": np.array(input_ids, dtype=np.int64),
|
||||
}
|
||||
if "attention_mask" in input_names:
|
||||
onnx_input["attention_mask"] = np.array(attention_mask, dtype=np.int64)
|
||||
if "token_type_ids" in input_names:
|
||||
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)
|
||||
return OnnxOutputContext(
|
||||
model_output=model_output[0],
|
||||
attention_mask=onnx_input.get("attention_mask", attention_mask),
|
||||
input_ids=onnx_input.get("input_ids", input_ids),
|
||||
)
|
||||
|
||||
def _embed_documents(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
documents: Union[str, Iterable[str]],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs,
|
||||
) -> Iterable[T]:
|
||||
is_small = False
|
||||
|
||||
if isinstance(documents, str):
|
||||
documents = [documents]
|
||||
is_small = True
|
||||
|
||||
if isinstance(documents, list):
|
||||
if len(documents) < batch_size:
|
||||
is_small = True
|
||||
|
||||
if parallel == 0:
|
||||
parallel = os.cpu_count()
|
||||
|
||||
if parallel is None or is_small:
|
||||
for batch in iter_batch(documents, batch_size):
|
||||
yield from self._post_process_onnx_output(self.onnx_embed(batch))
|
||||
else:
|
||||
start_method = (
|
||||
"forkserver" if "forkserver" in get_all_start_methods() else "spawn"
|
||||
)
|
||||
params = {"model_name": model_name, "cache_dir": cache_dir, **kwargs}
|
||||
pool = ParallelWorkerPool(
|
||||
parallel, self._get_worker_class(), start_method=start_method
|
||||
)
|
||||
for batch in pool.ordered_map(iter_batch(documents, batch_size), **params):
|
||||
yield from self._post_process_onnx_output(batch)
|
||||
|
||||
|
||||
class TextEmbeddingWorker(EmbeddingWorker):
|
||||
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
|
||||
@@ -1,9 +1,12 @@
|
||||
from typing import Any, Dict, Iterable, List, Optional, Type, Union
|
||||
from typing import Any, Dict, Iterable, List, Optional, Sequence, Type, Union
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastembed.common import OnnxProvider
|
||||
from fastembed.text.clip_embedding import CLIPOnnxEmbedding
|
||||
from fastembed.text.e5_onnx_embedding import E5OnnxEmbedding
|
||||
from fastembed.text.jina_onnx_embedding import JinaOnnxEmbedding
|
||||
from fastembed.text.mini_lm_embedding import MiniLMOnnxEmbedding
|
||||
from fastembed.text.onnx_embedding import OnnxTextEmbedding
|
||||
from fastembed.text.text_embedding_base import TextEmbeddingBase
|
||||
|
||||
@@ -13,6 +16,8 @@ class TextEmbedding(TextEmbeddingBase):
|
||||
OnnxTextEmbedding,
|
||||
E5OnnxEmbedding,
|
||||
JinaOnnxEmbedding,
|
||||
CLIPOnnxEmbedding,
|
||||
MiniLMOnnxEmbedding,
|
||||
]
|
||||
|
||||
@classmethod
|
||||
@@ -49,14 +54,24 @@ class TextEmbedding(TextEmbeddingBase):
|
||||
model_name: str = "BAAI/bge-small-en-v1.5",
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(model_name, cache_dir, threads, **kwargs)
|
||||
|
||||
for EMBEDDING_MODEL_TYPE in self.EMBEDDINGS_REGISTRY:
|
||||
supported_models = EMBEDDING_MODEL_TYPE.list_supported_models()
|
||||
if any(model_name.lower() == model["model"].lower() for model in supported_models):
|
||||
self.model = EMBEDDING_MODEL_TYPE(model_name, cache_dir, threads, **kwargs)
|
||||
if any(
|
||||
model_name.lower() == model["model"].lower()
|
||||
for model in supported_models
|
||||
):
|
||||
self.model = EMBEDDING_MODEL_TYPE(
|
||||
model_name,
|
||||
cache_dir,
|
||||
threads=threads,
|
||||
providers=providers,
|
||||
**kwargs,
|
||||
)
|
||||
return
|
||||
|
||||
raise ValueError(
|
||||
|
||||
@@ -16,6 +16,7 @@ class TextEmbeddingBase(ModelManagement):
|
||||
self.model_name = model_name
|
||||
self.cache_dir = cache_dir
|
||||
self.threads = threads
|
||||
self._local_files_only = kwargs.pop("local_files_only", False)
|
||||
|
||||
def embed(
|
||||
self,
|
||||
@@ -41,7 +42,9 @@ class TextEmbeddingBase(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[np.ndarray]:
|
||||
def query_embed(
|
||||
self, query: Union[str, Iterable[str]], **kwargs
|
||||
) -> Iterable[np.ndarray]:
|
||||
"""
|
||||
Embeds queries
|
||||
|
||||
|
||||
Generated
-3449
File diff suppressed because it is too large
Load Diff
+13
-10
@@ -1,8 +1,8 @@
|
||||
[tool.poetry]
|
||||
name = "fastembed"
|
||||
version = "0.2.6"
|
||||
name = "fastembed-gpu"
|
||||
version = "0.3.1"
|
||||
description = "Fast, light, accurate library built for retrieval embedding generation"
|
||||
authors = ["NirantK <nirant.bits@gmail.com>"]
|
||||
authors = ["Qdrant Team <info@qdrant.tech>", "NirantK <nirant.bits@gmail.com>"]
|
||||
license = "Apache License"
|
||||
readme = "README.md"
|
||||
packages = [{include = "fastembed"}]
|
||||
@@ -12,21 +12,24 @@ keywords = ["vector", "embedding", "neural", "search", "qdrant", "sentence-trans
|
||||
|
||||
[tool.poetry.dependencies]
|
||||
python = ">=3.8.0,<3.13"
|
||||
onnx = "^1.15.0"
|
||||
onnxruntime = "^1.17.0"
|
||||
onnxruntime-gpu = "^1.17.0"
|
||||
tqdm = "^4.66"
|
||||
requests = "^2.31"
|
||||
tokenizers = "^0.15.1"
|
||||
huggingface-hub = "^0.20"
|
||||
tokenizers = ">=0.15,<1.0"
|
||||
huggingface-hub = ">=0.20,<1.0"
|
||||
loguru = "^0.7.2"
|
||||
numpy = [
|
||||
{ version = ">=1.21", python = "<3.12" },
|
||||
{ version = ">=1.26", python = ">=3.12" }
|
||||
{ version = ">=1.21, <2", python = "<3.12" },
|
||||
{ version = ">=1.26, <2", python = ">=3.12" }
|
||||
]
|
||||
pillow = "^10.3.0"
|
||||
snowballstemmer = "^2.2.0"
|
||||
PyStemmer = "^2.2.0"
|
||||
mmh3 = "^4.0"
|
||||
|
||||
[tool.poetry.group.dev.dependencies]
|
||||
pytest = "^7.4.2"
|
||||
ruff = "^0.3.1"
|
||||
ruff = ">=0.3.1,<1.0"
|
||||
notebook = ">=7.0.2"
|
||||
pre-commit = {version = "^3.6.2", python = ">=3.9,<3.12" }
|
||||
|
||||
|
||||
@@ -0,0 +1,4 @@
|
||||
from pathlib import Path
|
||||
|
||||
TEST_DIR = Path(__file__).parent
|
||||
TEST_MISC_DIR = TEST_DIR / "misc"
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 169 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 11 KiB |
+12
-4
@@ -57,7 +57,9 @@ class HF:
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(model_id)
|
||||
|
||||
def embed(self, texts: List[str]):
|
||||
encoded_input = self.tokenizer(texts, max_length=512, padding=True, truncation=True, return_tensors="pt")
|
||||
encoded_input = self.tokenizer(
|
||||
texts, max_length=512, padding=True, truncation=True, return_tensors="pt"
|
||||
)
|
||||
model_output = self.model(**encoded_input)
|
||||
sentence_embeddings = model_output[0][:, 0]
|
||||
sentence_embeddings = F.normalize(sentence_embeddings)
|
||||
@@ -84,7 +86,9 @@ embedding_model = DefaultEmbedding()
|
||||
|
||||
|
||||
# %%
|
||||
def calculate_time_stats(embed_func: Callable, documents: list, k: int) -> Tuple[float, float, float]:
|
||||
def calculate_time_stats(
|
||||
embed_func: Callable, documents: list, k: int
|
||||
) -> Tuple[float, float, float]:
|
||||
times = []
|
||||
for _ in range(k):
|
||||
# Timing the embed_func call
|
||||
@@ -101,13 +105,17 @@ def calculate_time_stats(embed_func: Callable, documents: list, k: int) -> Tuple
|
||||
# %%
|
||||
hf_stats = calculate_time_stats(hf.embed, documents, k=2)
|
||||
print(f"Huggingface Transformers (Average, Max, Min): {hf_stats}")
|
||||
fst_stats = calculate_time_stats(lambda x: list(embedding_model.embed(x)), documents, k=2)
|
||||
fst_stats = calculate_time_stats(
|
||||
lambda x: list(embedding_model.embed(x)), documents, k=2
|
||||
)
|
||||
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], documents: list
|
||||
hf_stats: Tuple[float, float, float],
|
||||
fst_stats: Tuple[float, float, float],
|
||||
documents: list,
|
||||
):
|
||||
# Calculating total characters in documents
|
||||
total_characters = sum(len(doc) for doc in documents)
|
||||
|
||||
@@ -0,0 +1,84 @@
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from fastembed import SparseTextEmbedding
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_name", ["Qdrant/bm42-all-minilm-l6-v2-attentions", "Qdrant/bm25"]
|
||||
)
|
||||
def test_attention_embeddings(model_name):
|
||||
model = SparseTextEmbedding(model_name=model_name)
|
||||
|
||||
output = list(
|
||||
model.query_embed(
|
||||
[
|
||||
"I must not fear. Fear is the mind-killer.",
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
assert len(output) == 1
|
||||
|
||||
for result in output:
|
||||
assert len(result.indices) == len(result.values)
|
||||
assert np.allclose(result.values, np.ones(len(result.values)))
|
||||
|
||||
quotes = [
|
||||
"I must not fear. Fear is the mind-killer.",
|
||||
"All animals are equal, but some animals are more equal than others.",
|
||||
"It was a pleasure to burn.",
|
||||
"The sky above the port was the color of television, tuned to a dead channel.",
|
||||
"In the beginning, the universe was created."
|
||||
" This has made a lot of people very angry and been widely regarded as a bad move.",
|
||||
"It's a truth universally acknowledged that a zombie in possession of brains must be in want of more brains.",
|
||||
"War is peace. Freedom is slavery. Ignorance is strength.",
|
||||
"We're not in Infinity; we're in the suburbs.",
|
||||
"I was a thousand times more evil than thou!",
|
||||
"History is merely a list of surprises... It can only prepare us to be surprised yet again.",
|
||||
]
|
||||
|
||||
output = list(model.embed(quotes))
|
||||
|
||||
assert len(output) == len(quotes)
|
||||
|
||||
for result in output:
|
||||
assert len(result.indices) == len(result.values)
|
||||
assert len(result.indices) > 0
|
||||
|
||||
# Test support for unknown languages
|
||||
output = list(
|
||||
model.query_embed(
|
||||
[
|
||||
"привет мир!",
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
assert len(output) == 1
|
||||
|
||||
for result in output:
|
||||
assert len(result.indices) == len(result.values)
|
||||
assert len(result.indices) == 2
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_name", ["Qdrant/bm42-all-minilm-l6-v2-attentions", "Qdrant/bm25"]
|
||||
)
|
||||
def test_parallel_processing(model_name):
|
||||
model = SparseTextEmbedding(model_name=model_name)
|
||||
|
||||
docs = ["hello world", "attention embedding"] * 100
|
||||
embeddings = list(model.embed(docs, batch_size=10, parallel=2))
|
||||
|
||||
embeddings_2 = list(model.embed(docs, batch_size=10, parallel=None))
|
||||
|
||||
embeddings_3 = list(model.embed(docs, batch_size=10, parallel=0))
|
||||
|
||||
assert len(embeddings) == len(docs)
|
||||
|
||||
for emb_1, emb_2, emb_3 in zip(embeddings, embeddings_2, embeddings_3):
|
||||
assert np.allclose(emb_1.indices, emb_2.indices)
|
||||
assert np.allclose(emb_1.indices, emb_3.indices)
|
||||
assert np.allclose(emb_1.values, emb_2.values)
|
||||
assert np.allclose(emb_1.values, emb_3.values)
|
||||
@@ -0,0 +1,73 @@
|
||||
import os
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from fastembed import ImageEmbedding
|
||||
from tests.config import TEST_MISC_DIR
|
||||
|
||||
CANONICAL_VECTOR_VALUES = {
|
||||
"Qdrant/clip-ViT-B-32-vision": np.array([-0.0098, 0.0128, -0.0274, 0.002, -0.0059]),
|
||||
"Qdrant/resnet50-onnx": np.array(
|
||||
[0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.01046245, 0.01171397, 0.00705971, 0.0]
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
def test_embedding():
|
||||
is_ci = os.getenv("CI")
|
||||
|
||||
for model_desc in ImageEmbedding.list_supported_models():
|
||||
if not is_ci and model_desc["size_in_GB"] > 1:
|
||||
continue
|
||||
|
||||
dim = model_desc["dim"]
|
||||
|
||||
model = ImageEmbedding(model_name=model_desc["model"])
|
||||
|
||||
images = [TEST_MISC_DIR / "image.jpeg", str(TEST_MISC_DIR / "small_image.jpeg")]
|
||||
embeddings = list(model.embed(images))
|
||||
embeddings = np.stack(embeddings, axis=0)
|
||||
assert embeddings.shape == (2, dim)
|
||||
|
||||
canonical_vector = CANONICAL_VECTOR_VALUES[model_desc["model"]]
|
||||
|
||||
assert np.allclose(
|
||||
embeddings[0, : canonical_vector.shape[0]], canonical_vector, atol=1e-3
|
||||
), model_desc["model"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("n_dims,model_name", [(512, "Qdrant/clip-ViT-B-32-vision")])
|
||||
def test_batch_embedding(n_dims, model_name):
|
||||
model = ImageEmbedding(model_name=model_name)
|
||||
n_images = 32
|
||||
images = [TEST_MISC_DIR / "image.jpeg", str(TEST_MISC_DIR / "small_image.jpeg")] * (
|
||||
n_images // 2
|
||||
)
|
||||
|
||||
embeddings = list(model.embed(images, batch_size=10))
|
||||
embeddings = np.stack(embeddings, axis=0)
|
||||
|
||||
assert embeddings.shape == (n_images, n_dims)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("n_dims,model_name", [(512, "Qdrant/clip-ViT-B-32-vision")])
|
||||
def test_parallel_processing(n_dims, model_name):
|
||||
model = ImageEmbedding(model_name=model_name)
|
||||
|
||||
n_images = 32
|
||||
images = [TEST_MISC_DIR / "image.jpeg", str(TEST_MISC_DIR / "small_image.jpeg")] * (
|
||||
n_images // 2
|
||||
)
|
||||
embeddings = list(model.embed(images, batch_size=10, parallel=2))
|
||||
embeddings = np.stack(embeddings, axis=0)
|
||||
|
||||
embeddings_2 = list(model.embed(images, batch_size=10, parallel=None))
|
||||
embeddings_2 = np.stack(embeddings_2, axis=0)
|
||||
|
||||
embeddings_3 = list(model.embed(images, batch_size=10, parallel=0))
|
||||
embeddings_3 = np.stack(embeddings_3, axis=0)
|
||||
|
||||
assert embeddings.shape == (n_images, n_dims)
|
||||
assert np.allclose(embeddings, embeddings_2, atol=1e-3)
|
||||
assert np.allclose(embeddings, embeddings_3, atol=1e-3)
|
||||
@@ -0,0 +1,112 @@
|
||||
import numpy as np
|
||||
|
||||
from fastembed.late_interaction.late_interaction_text_embedding import (
|
||||
LateInteractionTextEmbedding,
|
||||
)
|
||||
|
||||
# vectors are abridged and rounded for brevity
|
||||
CANONICAL_COLUMN_VALUES = {
|
||||
"colbert-ir/colbertv2.0": np.array(
|
||||
[
|
||||
[0.0759, 0.0841, -0.0299, 0.0374, 0.0254],
|
||||
[0.0005, -0.0163, -0.0127, 0.2165, 0.1517],
|
||||
[-0.0257, -0.0575, 0.0135, 0.2202, 0.1896],
|
||||
[0.0846, 0.0122, 0.0032, -0.0109, -0.1041],
|
||||
[0.0477, 0.1078, -0.0314, 0.016, 0.0156],
|
||||
]
|
||||
)
|
||||
}
|
||||
|
||||
CANONICAL_QUERY_VALUES = {
|
||||
"colbert-ir/colbertv2.0": np.array(
|
||||
[
|
||||
[0.0824, 0.0872, -0.0324, 0.0418, 0.024],
|
||||
[-0.0007, -0.0154, -0.0113, 0.2277, 0.1528],
|
||||
[-0.0251, -0.0565, 0.0136, 0.2236, 0.1838],
|
||||
[0.0848, 0.0056, 0.0041, -0.0036, -0.1032],
|
||||
[0.0574, 0.1072, -0.0332, 0.0233, 0.0209],
|
||||
[0.1041, 0.0364, -0.0058, -0.027, -0.0704],
|
||||
[0.106, 0.0371, -0.0055, -0.0339, -0.0719],
|
||||
[0.1063, 0.0363, 0.0014, -0.0334, -0.0698],
|
||||
[0.112, 0.036, 0.0026, -0.0355, -0.0675],
|
||||
[0.1184, 0.0441, 0.0166, -0.0169, -0.0244],
|
||||
[0.1033, 0.035, 0.0183, 0.0475, 0.0612],
|
||||
[-0.0028, -0.014, -0.016, 0.2175, 0.1537],
|
||||
[0.0547, 0.0219, -0.007, 0.1748, 0.1154],
|
||||
[-0.001, -0.0184, -0.0112, 0.2197, 0.1523],
|
||||
[-0.0012, -0.0149, -0.0119, 0.2147, 0.152],
|
||||
[-0.0186, -0.0239, -0.014, 0.2196, 0.156],
|
||||
[-0.017, -0.0232, -0.0108, 0.2212, 0.157],
|
||||
[-0.0109, -0.0024, -0.003, 0.1972, 0.1391],
|
||||
[0.0898, 0.0219, -0.0255, 0.0734, -0.0096],
|
||||
[0.1143, 0.015, -0.022, 0.0417, -0.0421],
|
||||
[0.1056, 0.0091, -0.0137, 0.0129, -0.0619],
|
||||
[0.0234, 0.004, -0.0285, 0.1565, 0.0883],
|
||||
[-0.0037, -0.0079, -0.0204, 0.1982, 0.1502],
|
||||
[0.0988, 0.0377, 0.0226, 0.0309, 0.0508],
|
||||
[-0.0103, -0.0128, -0.0035, 0.2114, 0.155],
|
||||
[-0.0103, -0.0184, -0.011, 0.2252, 0.157],
|
||||
[-0.0033, -0.0292, -0.0097, 0.2237, 0.1607],
|
||||
[-0.0198, -0.0257, -0.0193, 0.2265, 0.165],
|
||||
[-0.0227, -0.0028, -0.0084, 0.1995, 0.1306],
|
||||
[0.0916, 0.0185, -0.0186, 0.0173, -0.0577],
|
||||
[0.1022, 0.0228, -0.0174, -0.0102, -0.065],
|
||||
[0.1043, 0.0231, -0.0144, -0.0246, -0.067],
|
||||
]
|
||||
)
|
||||
}
|
||||
|
||||
docs = ["Hello World"]
|
||||
|
||||
|
||||
def test_batch_embedding():
|
||||
docs_to_embed = docs * 10
|
||||
|
||||
for model_name, expected_result in CANONICAL_COLUMN_VALUES.items():
|
||||
print("evaluating", model_name)
|
||||
model = LateInteractionTextEmbedding(model_name=model_name)
|
||||
result = list(model.embed(docs_to_embed, batch_size=6))
|
||||
|
||||
for value in result:
|
||||
token_num, abridged_dim = expected_result.shape
|
||||
assert np.allclose(value[:, :abridged_dim], expected_result, atol=10e-4)
|
||||
|
||||
|
||||
def test_single_embedding():
|
||||
docs_to_embed = docs
|
||||
|
||||
for model_name, expected_result in CANONICAL_COLUMN_VALUES.items():
|
||||
print("evaluating", model_name)
|
||||
model = LateInteractionTextEmbedding(model_name=model_name)
|
||||
result = next(iter(model.embed(docs_to_embed, batch_size=6)))
|
||||
token_num, abridged_dim = expected_result.shape
|
||||
assert np.allclose(result[:, :abridged_dim], expected_result, atol=10e-4)
|
||||
|
||||
|
||||
def test_single_embedding_query():
|
||||
queries_to_embed = docs
|
||||
|
||||
for model_name, expected_result in CANONICAL_QUERY_VALUES.items():
|
||||
print("evaluating", model_name)
|
||||
model = LateInteractionTextEmbedding(model_name=model_name)
|
||||
result = next(iter(model.query_embed(queries_to_embed)))
|
||||
token_num, abridged_dim = expected_result.shape
|
||||
assert np.allclose(result[:, :abridged_dim], expected_result, atol=10e-4)
|
||||
|
||||
|
||||
def test_parallel_processing():
|
||||
model = LateInteractionTextEmbedding(model_name="colbert-ir/colbertv2.0")
|
||||
token_dim = 128
|
||||
docs = ["hello world", "flag embedding"] * 100
|
||||
embeddings = list(model.embed(docs, batch_size=10, parallel=2))
|
||||
embeddings = np.stack(embeddings, axis=0)
|
||||
|
||||
embeddings_2 = list(model.embed(docs, batch_size=10, parallel=None))
|
||||
embeddings_2 = np.stack(embeddings_2, axis=0)
|
||||
|
||||
embeddings_3 = list(model.embed(docs, batch_size=10, parallel=0))
|
||||
embeddings_3 = np.stack(embeddings_3, axis=0)
|
||||
|
||||
assert embeddings.shape[0] == len(docs) and embeddings.shape[-1] == token_dim
|
||||
assert np.allclose(embeddings, embeddings_2, atol=1e-3)
|
||||
assert np.allclose(embeddings, embeddings_3, atol=1e-3)
|
||||
@@ -1,4 +1,5 @@
|
||||
import pytest
|
||||
|
||||
from fastembed.sparse.sparse_text_embedding import SparseTextEmbedding
|
||||
|
||||
CANONICAL_COLUMN_VALUES = {
|
||||
@@ -47,11 +48,8 @@ def test_batch_embedding():
|
||||
docs_to_embed = docs * 10
|
||||
|
||||
for model_name, expected_result in CANONICAL_COLUMN_VALUES.items():
|
||||
print("evaluating", model_name)
|
||||
model = SparseTextEmbedding(model_name=model_name)
|
||||
result = next(iter(model.embed(docs_to_embed, batch_size=6)))
|
||||
print(result.indices)
|
||||
|
||||
assert result.indices.tolist() == expected_result["indices"]
|
||||
|
||||
for i, value in enumerate(result.values):
|
||||
@@ -59,32 +57,33 @@ def test_batch_embedding():
|
||||
|
||||
|
||||
def test_single_embedding():
|
||||
docs_to_embed = docs
|
||||
|
||||
for model_name, expected_result in CANONICAL_COLUMN_VALUES.items():
|
||||
print("evaluating", model_name)
|
||||
model = SparseTextEmbedding(model_name=model_name)
|
||||
result = next(iter(model.embed(docs_to_embed, batch_size=6)))
|
||||
print(result.indices)
|
||||
|
||||
assert result.indices.tolist() == expected_result["indices"]
|
||||
passage_result = next(iter(model.embed(docs, batch_size=6)))
|
||||
query_result = next(iter(model.query_embed(docs)))
|
||||
for result in [passage_result, query_result]:
|
||||
assert result.indices.tolist() == expected_result["indices"]
|
||||
|
||||
for i, value in enumerate(result.values):
|
||||
assert pytest.approx(value, abs=0.001) == expected_result["values"][i]
|
||||
for i, value in enumerate(result.values):
|
||||
assert pytest.approx(value, abs=0.001) == expected_result["values"][i]
|
||||
|
||||
|
||||
def test_parallel_processing():
|
||||
import numpy as np
|
||||
|
||||
model = SparseTextEmbedding(
|
||||
model_name="prithivida/Splade_PP_en_v1",
|
||||
)
|
||||
model = SparseTextEmbedding(model_name="prithivida/Splade_PP_en_v1")
|
||||
docs = ["hello world", "flag embedding"] * 30
|
||||
sparse_embeddings_duo = list(model.embed(docs, batch_size=10, parallel=2))
|
||||
sparse_embeddings_all = list(model.embed(docs, batch_size=10, parallel=0))
|
||||
sparse_embeddings = list(model.embed(docs, batch_size=10, parallel=None))
|
||||
|
||||
assert len(sparse_embeddings) == len(sparse_embeddings_duo) == len(sparse_embeddings_all) == len(docs)
|
||||
assert (
|
||||
len(sparse_embeddings)
|
||||
== len(sparse_embeddings_duo)
|
||||
== len(sparse_embeddings_all)
|
||||
== len(docs)
|
||||
)
|
||||
|
||||
for sparse_embedding, sparse_embedding_duo, sparse_embedding_all in zip(
|
||||
sparse_embeddings, sparse_embeddings_duo, sparse_embeddings_all
|
||||
@@ -94,5 +93,9 @@ def test_parallel_processing():
|
||||
== sparse_embedding_duo.indices.tolist()
|
||||
== sparse_embedding_all.indices.tolist()
|
||||
)
|
||||
assert np.allclose(sparse_embedding.values, sparse_embedding_duo.values, atol=1e-3)
|
||||
assert np.allclose(sparse_embedding.values, sparse_embedding_all.values, atol=1e-3)
|
||||
assert np.allclose(
|
||||
sparse_embedding.values, sparse_embedding_duo.values, atol=1e-3
|
||||
)
|
||||
assert np.allclose(
|
||||
sparse_embedding.values, sparse_embedding_all.values, atol=1e-3
|
||||
)
|
||||
|
||||
@@ -7,37 +7,77 @@ from fastembed.text.text_embedding import TextEmbedding
|
||||
|
||||
CANONICAL_VECTOR_VALUES = {
|
||||
"BAAI/bge-small-en": np.array([-0.0232, -0.0255, 0.0174, -0.0639, -0.0006]),
|
||||
"BAAI/bge-small-en-v1.5": np.array([0.01522374, -0.02271799, 0.00860278, -0.07424029, 0.00386434]),
|
||||
"BAAI/bge-small-en-v1.5-quantized": np.array([0.01522374, -0.02271799, 0.00860278, -0.07424029, 0.00386434]),
|
||||
"BAAI/bge-small-zh-v1.5": np.array([-0.01023294, 0.07634465, 0.0691722, -0.04458365, -0.03160762]),
|
||||
"BAAI/bge-small-en-v1.5": np.array(
|
||||
[0.01522374, -0.02271799, 0.00860278, -0.07424029, 0.00386434]
|
||||
),
|
||||
"BAAI/bge-small-en-v1.5-quantized": np.array(
|
||||
[0.01522374, -0.02271799, 0.00860278, -0.07424029, 0.00386434]
|
||||
),
|
||||
"BAAI/bge-small-zh-v1.5": np.array(
|
||||
[-0.01023294, 0.07634465, 0.0691722, -0.04458365, -0.03160762]
|
||||
),
|
||||
"BAAI/bge-base-en": np.array([0.0115, 0.0372, 0.0295, 0.0121, 0.0346]),
|
||||
"BAAI/bge-base-en-v1.5": np.array([0.01129394, 0.05493144, 0.02615099, 0.00328772, 0.02996045]),
|
||||
"BAAI/bge-large-en-v1.5": np.array([0.03434538, 0.03316108, 0.02191251, -0.03713358, -0.01577825]),
|
||||
"BAAI/bge-large-en-v1.5-quantized": np.array([0.03434538, 0.03316108, 0.02191251, -0.03713358, -0.01577825]),
|
||||
"sentence-transformers/all-MiniLM-L6-v2": np.array([0.0259, 0.0058, 0.0114, 0.0380, -0.0233]),
|
||||
"sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2": np.array([0.0094, 0.0184, 0.0328, 0.0072, -0.0351]),
|
||||
"intfloat/multilingual-e5-large": np.array([0.0098, 0.0045, 0.0066, -0.0354, 0.0070]),
|
||||
"BAAI/bge-base-en-v1.5": np.array(
|
||||
[0.01129394, 0.05493144, 0.02615099, 0.00328772, 0.02996045]
|
||||
),
|
||||
"BAAI/bge-large-en-v1.5": np.array(
|
||||
[0.03434538, 0.03316108, 0.02191251, -0.03713358, -0.01577825]
|
||||
),
|
||||
"BAAI/bge-large-en-v1.5-quantized": np.array(
|
||||
[0.03434538, 0.03316108, 0.02191251, -0.03713358, -0.01577825]
|
||||
),
|
||||
"sentence-transformers/all-MiniLM-L6-v2": np.array(
|
||||
[-0.034478, 0.03102, 0.00673, 0.02611, -0.039362]
|
||||
),
|
||||
"sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2": np.array(
|
||||
[0.0094, 0.0184, 0.0328, 0.0072, -0.0351]
|
||||
),
|
||||
"intfloat/multilingual-e5-large": np.array(
|
||||
[0.0098, 0.0045, 0.0066, -0.0354, 0.0070]
|
||||
),
|
||||
"sentence-transformers/paraphrase-multilingual-mpnet-base-v2": np.array(
|
||||
[-0.01341097, 0.0416553, -0.00480805, 0.02844842, 0.0505299]
|
||||
),
|
||||
"jinaai/jina-embeddings-v2-small-en": np.array([-0.0455, -0.0428, -0.0122, 0.0613, 0.0015]),
|
||||
"jinaai/jina-embeddings-v2-base-en": np.array([-0.0332, -0.0509, 0.0287, -0.0043, -0.0077]),
|
||||
"nomic-ai/nomic-embed-text-v1": np.array([0.0061, 0.0103, -0.0296, -0.0242, -0.0170]),
|
||||
"jinaai/jina-embeddings-v2-small-en": np.array(
|
||||
[-0.0455, -0.0428, -0.0122, 0.0613, 0.0015]
|
||||
),
|
||||
"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]
|
||||
),
|
||||
"nomic-ai/nomic-embed-text-v1": np.array(
|
||||
[0.0061, 0.0103, -0.0296, -0.0242, -0.0170]
|
||||
),
|
||||
"nomic-ai/nomic-embed-text-v1.5": np.array(
|
||||
[-1.6531514e-02, 8.5380634e-05, -1.8171231e-01, -3.9333291e-03, 1.2763254e-02]
|
||||
),
|
||||
"nomic-ai/nomic-embed-text-v1.5-Q": np.array(
|
||||
[-0.01554983, 0.0129992 , -0.17909265, -0.01062993, 0.00512859]
|
||||
[-0.01554983, 0.0129992, -0.17909265, -0.01062993, 0.00512859]
|
||||
),
|
||||
"thenlper/gte-large": np.array(
|
||||
[-0.01920587, 0.00113156, -0.00708992, -0.00632304, -0.04025577]
|
||||
),
|
||||
"mixedbread-ai/mxbai-embed-large-v1": np.array(
|
||||
[0.02295546, 0.03196154, 0.016512, -0.04031524, -0.0219634]
|
||||
),
|
||||
"snowflake/snowflake-arctic-embed-xs": np.array(
|
||||
[0.0092, 0.0619, 0.0196, 0.009, -0.0114]
|
||||
),
|
||||
"snowflake/snowflake-arctic-embed-s": np.array(
|
||||
[-0.0416, -0.0867, 0.0209, 0.0554, -0.0272]
|
||||
),
|
||||
"snowflake/snowflake-arctic-embed-m": np.array(
|
||||
[-0.0329, 0.0364, 0.0481, 0.0016, 0.0328]
|
||||
),
|
||||
"thenlper/gte-large": np.array([-0.01920587, 0.00113156, -0.00708992, -0.00632304, -0.04025577]),
|
||||
"mixedbread-ai/mxbai-embed-large-v1": np.array([0.02295546, 0.03196154, 0.016512, -0.04031524, -0.0219634]),
|
||||
"snowflake/snowflake-arctic-embed-xs": np.array([0.0092, 0.0619, 0.0196, 0.009, -0.0114]),
|
||||
"snowflake/snowflake-arctic-embed-s": np.array([-0.0416, -0.0867, 0.0209, 0.0554, -0.0272]),
|
||||
"snowflake/snowflake-arctic-embed-m": np.array([-0.0329, 0.0364, 0.0481, 0.0016, 0.0328]),
|
||||
"snowflake/snowflake-arctic-embed-m-long": np.array(
|
||||
[0.0080, -0.0266, -0.0335, 0.0282, 0.0143]
|
||||
),
|
||||
"snowflake/snowflake-arctic-embed-l": np.array([0.0189, -0.0673, 0.0183, 0.0124, 0.0146]),
|
||||
"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]),
|
||||
}
|
||||
|
||||
|
||||
@@ -49,6 +89,7 @@ def test_embedding():
|
||||
continue
|
||||
|
||||
dim = model_desc["dim"]
|
||||
|
||||
model = TextEmbedding(model_name=model_desc["model"])
|
||||
|
||||
docs = ["hello world", "flag embedding"]
|
||||
@@ -57,11 +98,14 @@ def test_embedding():
|
||||
assert embeddings.shape == (2, dim)
|
||||
|
||||
canonical_vector = CANONICAL_VECTOR_VALUES[model_desc["model"]]
|
||||
assert np.allclose(embeddings[0, : canonical_vector.shape[0]], canonical_vector, atol=1e-3), model_desc["model"]
|
||||
assert np.allclose(
|
||||
embeddings[0, : canonical_vector.shape[0]], canonical_vector, atol=1e-3
|
||||
), model_desc["model"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"n_dims,model_name", [(384, "BAAI/bge-small-en-v1.5"), (768, "jinaai/jina-embeddings-v2-base-en")]
|
||||
"n_dims,model_name",
|
||||
[(384, "BAAI/bge-small-en-v1.5"), (768, "jinaai/jina-embeddings-v2-base-en")],
|
||||
)
|
||||
def test_batch_embedding(n_dims, model_name):
|
||||
model = TextEmbedding(model_name=model_name)
|
||||
@@ -74,7 +118,8 @@ def test_batch_embedding(n_dims, model_name):
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"n_dims,model_name", [(384, "BAAI/bge-small-en-v1.5"), (768, "jinaai/jina-embeddings-v2-base-en")]
|
||||
"n_dims,model_name",
|
||||
[(384, "BAAI/bge-small-en-v1.5"), (768, "jinaai/jina-embeddings-v2-base-en")],
|
||||
)
|
||||
def test_parallel_processing(n_dims, model_name):
|
||||
model = TextEmbedding(model_name=model_name)
|
||||
|
||||
Reference in New Issue
Block a user