mirror of
https://github.com/qdrant/fastembed.git
synced 2026-09-22 05:57:51 -05:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
40e5e96212 | ||
|
|
8d04b81782 | ||
|
|
b389798d8e | ||
|
|
6f43572373 | ||
|
|
de4ecb48c7 | ||
|
|
b9d605138f | ||
|
|
a931f143ef | ||
|
|
105ff19035 | ||
|
|
4599b933ca | ||
|
|
0fa1596c3d | ||
|
|
2fe33c5a62 | ||
|
|
969ea29923 | ||
|
|
a5b266e018 | ||
|
|
877d963bd1 | ||
|
|
37a66d9e16 | ||
|
|
b08febbf93 | ||
|
|
6dbdd6d171 | ||
|
|
f1a3a6d082 | ||
|
|
993dcd5f68 | ||
|
|
73e1e5ecb9 | ||
|
|
105d6cfb97 | ||
|
|
bb815405aa | ||
|
|
314842121d | ||
|
|
c2f6fd1c90 | ||
|
|
b05877de93 | ||
|
|
54f6cd9cbc | ||
|
|
ae37da3bd4 | ||
|
|
fa11d0f0c7 | ||
|
|
50289c62ed | ||
|
|
3c10b6625b | ||
|
|
e89654d435 | ||
|
|
cec8d54502 | ||
|
|
55b985c1ad | ||
|
|
c8b1a18cfc | ||
|
|
3b5e4c8722 | ||
|
|
516170cbaf | ||
|
|
0f79d3f9d8 | ||
|
|
2ef9c38b8b | ||
|
|
da30f934d8 | ||
|
|
adfc03ed0d | ||
|
|
e9dc3b1060 | ||
|
|
1343e55076 | ||
|
|
9841666bd5 | ||
|
|
860e2ad691 | ||
|
|
d141e29784 | ||
|
|
7c935717d2 | ||
|
|
8413066f7b | ||
|
|
faed7f1320 | ||
|
|
5b895420b1 | ||
|
|
e868bbaebc | ||
|
|
a5f3f11829 | ||
|
|
3fc0e2b382 | ||
|
|
12b06ece51 | ||
|
|
1deb830328 | ||
|
|
aba8fb43cf | ||
|
|
69ffaa0f48 | ||
|
|
cf67a80ff7 | ||
|
|
a21e925000 | ||
|
|
0638c011dc | ||
|
|
eaecf7d471 | ||
|
|
58b5a8ed9a | ||
|
|
519b310f22 | ||
|
|
e2e1f93685 | ||
|
|
dab4dcc99a | ||
|
|
40a03740ff | ||
|
|
65c2efd6f1 | ||
|
|
97f2fb278e | ||
|
|
fa2205115d | ||
|
|
9445f95a32 | ||
|
|
ab9ab73278 | ||
|
|
08925c9cfc | ||
|
|
3a1f468ef1 | ||
|
|
bfeeb28721 | ||
|
|
a6841a8bde | ||
|
|
62607c237b | ||
|
|
49762a6d19 | ||
|
|
782273f851 | ||
|
|
9c72d2f59f | ||
|
|
0e258ab875 | ||
|
|
e49789c129 | ||
|
|
1bf72922ce | ||
|
|
70566dff99 | ||
|
|
fd116dd507 | ||
|
|
0315c3b8c6 | ||
|
|
54e0f38914 | ||
|
|
3a8985b35c | ||
|
|
f0ff09c546 | ||
|
|
d09af55edd | ||
|
|
9387ca3205 | ||
|
|
9c74fb3cfb | ||
|
|
1fe42d8d18 | ||
|
|
f820c36656 |
@@ -1,6 +1,6 @@
|
||||
name: Bug/New Model Request
|
||||
description: File a bug report/Request a new Model
|
||||
title: "[Bug/Model Request]: "
|
||||
name: Bug
|
||||
description: File a bug report
|
||||
title: "[Bug]: "
|
||||
body:
|
||||
- type: markdown
|
||||
attributes:
|
||||
@@ -10,11 +10,22 @@ body:
|
||||
id: what-happened
|
||||
attributes:
|
||||
label: What happened?
|
||||
description: Also tell us, what did you expect to happen?
|
||||
placeholder: Tell us what you see!
|
||||
value: "A bug happened!"
|
||||
description: Describe the error you encountered.
|
||||
placeholder: <Description>
|
||||
validations:
|
||||
required: true
|
||||
- type: textarea
|
||||
id: expected
|
||||
attributes:
|
||||
label: What is the expected behaviour?
|
||||
description: Describe the way you expected the code to behave.
|
||||
placeholder: <Description>
|
||||
- type: textarea
|
||||
id: code-snippet
|
||||
attributes:
|
||||
label: A minimal reproducible example
|
||||
description: It would really help us to fix the problem if you could provide a code snippet that reproduces the issue.
|
||||
placeholder: <Code snippet>
|
||||
- type: textarea
|
||||
id: python-version
|
||||
attributes:
|
||||
@@ -23,21 +34,12 @@ body:
|
||||
placeholder: Python3.10
|
||||
validations:
|
||||
required: true
|
||||
- type: dropdown
|
||||
- type: textarea
|
||||
id: version
|
||||
attributes:
|
||||
label: Version
|
||||
label: FastEmbed 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.7 (Latest)
|
||||
- 0.2.6
|
||||
- 0.2.5
|
||||
- 0.2.4
|
||||
- 0.2.3
|
||||
- 0.2.2
|
||||
- 0.2.1
|
||||
- 0.1.x
|
||||
default: 0
|
||||
placeholder: v0.5.1
|
||||
validations:
|
||||
required: true
|
||||
- type: dropdown
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
blank_issues_enabled: false
|
||||
blank_issues_enabled: true
|
||||
contact_links:
|
||||
- name: GitHub Community Support
|
||||
url: https://github.com/qdrant/fastembed/discussions
|
||||
|
||||
@@ -0,0 +1,22 @@
|
||||
name: Feature
|
||||
description: New functionality request
|
||||
title: "[Feature]: "
|
||||
body:
|
||||
- type: markdown
|
||||
attributes:
|
||||
value: |
|
||||
Thanks for taking the time to fill out this report!
|
||||
- type: textarea
|
||||
id: feature-description
|
||||
attributes:
|
||||
label: What feature would you like to request?
|
||||
description: Please provide the description of the feature you would like to request.
|
||||
placeholder: <Description>
|
||||
validations:
|
||||
required: true
|
||||
- type: textarea
|
||||
id: additional-info
|
||||
attributes:
|
||||
label: Is there any additional information you would like to provide?
|
||||
description: Please provide any additional information that you think might be useful.
|
||||
placeholder: <Info>
|
||||
@@ -0,0 +1,22 @@
|
||||
name: Model
|
||||
description: Request a new model
|
||||
title: "[Model]: "
|
||||
body:
|
||||
- type: markdown
|
||||
attributes:
|
||||
value: |
|
||||
Thanks for taking the time to fill out this report!
|
||||
- type: textarea
|
||||
id: model-name
|
||||
attributes:
|
||||
label: Which model would you like to support?
|
||||
description: Please provide the name of the model you would like to see supported.
|
||||
placeholder: Link to the model (e.g. on HuggingFace)
|
||||
validations:
|
||||
required: true
|
||||
- type: textarea
|
||||
id: motivation
|
||||
attributes:
|
||||
label: What are the main advantages of this model?
|
||||
description: Please describe the main advantages of this model comparing to the existing ones and provide links to benchmarks if there are any.
|
||||
placeholder: <Description>
|
||||
@@ -0,0 +1,19 @@
|
||||
### All Submissions:
|
||||
|
||||
* [ ] Have you followed the guidelines in our Contributing document?
|
||||
* [ ] Have you checked to ensure there aren't other open [Pull Requests](../../../pulls) for the same update/change?
|
||||
|
||||
<!-- You can erase any parts of this template not applicable to your Pull Request. -->
|
||||
|
||||
### New Feature Submissions:
|
||||
|
||||
* [ ] Does your submission pass the existing tests?
|
||||
* [ ] Have you added tests for your feature?
|
||||
* [ ] Have you installed `pre-commit` with `pip3 install pre-commit` and set up hooks with `pre-commit install`?
|
||||
|
||||
### New models submission:
|
||||
|
||||
* [ ] Have you added an explanation of why it's important to include this model?
|
||||
* [ ] Have you added tests for the new model? Were canonical values for tests computed via the original model?
|
||||
* [ ] Have you added the code snippet for how canonical values were computed?
|
||||
* [ ] Have you successfully ran tests with your changes locally?
|
||||
@@ -21,5 +21,5 @@ jobs:
|
||||
path: .cache
|
||||
restore-keys: |
|
||||
mkdocs-material-
|
||||
- run: pip install mkdocs-material mkdocstrings pillow cairosvg mknotebooks
|
||||
- run: pip install mkdocs-material mkdocstrings==0.27.0 pillow cairosvg mknotebooks
|
||||
- run: mkdocs gh-deploy --force
|
||||
|
||||
@@ -3,8 +3,6 @@ name: Tests
|
||||
on:
|
||||
push:
|
||||
branches: [ master, main, gpu ]
|
||||
schedule:
|
||||
- cron: 0 0 * * *
|
||||
pull_request:
|
||||
|
||||
env:
|
||||
@@ -16,11 +14,11 @@ jobs:
|
||||
strategy:
|
||||
matrix:
|
||||
python-version:
|
||||
- '3.8.x'
|
||||
- '3.9.x'
|
||||
- '3.10.x'
|
||||
- '3.11.x'
|
||||
- '3.12.x'
|
||||
- '3.13.x'
|
||||
os:
|
||||
- ubuntu-latest
|
||||
- macos-latest
|
||||
@@ -40,15 +38,8 @@ jobs:
|
||||
run: |
|
||||
python -m pip install poetry
|
||||
poetry config virtualenvs.create false
|
||||
poetry install --no-interaction --no-ansi --without docs
|
||||
|
||||
- name: Install Test Dependencies
|
||||
run: pip install pytest pytest-md pytest-emoji
|
||||
poetry install --no-interaction --no-ansi --without dev,docs
|
||||
|
||||
- name: Run pytest
|
||||
uses: pavelzw/pytest-action@v2
|
||||
with:
|
||||
verbose: true
|
||||
emoji: true
|
||||
job-summary: true
|
||||
report-title: 'FastEmbed Test Report'
|
||||
run: |
|
||||
poetry run pytest
|
||||
|
||||
@@ -0,0 +1,40 @@
|
||||
name: type-checkers
|
||||
|
||||
on: [push]
|
||||
|
||||
jobs:
|
||||
build:
|
||||
runs-on: ${{ matrix.os }}
|
||||
strategy:
|
||||
fail-fast: true
|
||||
matrix:
|
||||
python-version: ["3.9", "3.10", "3.11", "3.12", "3.13"]
|
||||
os: [ubuntu-latest]
|
||||
|
||||
name: Python ${{ matrix.python-version }} test
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v1
|
||||
|
||||
- name: Set up Python ${{ matrix.python-version }}
|
||||
uses: actions/setup-python@v2
|
||||
with:
|
||||
python-version: ${{ matrix.python-version }}
|
||||
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip poetry
|
||||
poetry install --no-interaction --no-ansi --without dev,docs,test
|
||||
|
||||
poetry run pip install "numpy<2.0.0" # https://github.com/python/mypy/issues/17396
|
||||
|
||||
- name: mypy
|
||||
run: |
|
||||
poetry run mypy fastembed \
|
||||
--disallow-incomplete-defs \
|
||||
--disallow-untyped-defs \
|
||||
--disable-error-code=import-untyped
|
||||
|
||||
- name: pyright
|
||||
run: |
|
||||
poetry run pyright tests/type_stub.py
|
||||
@@ -0,0 +1,22 @@
|
||||
Copyright 2024 Qdrant
|
||||
|
||||
This product includes software developed by Qdrant
|
||||
|
||||
This distribution includes the following Jina AI models, each with its respective license:
|
||||
- jinaai/jina-colbert-v2
|
||||
- License: cc-by-nc-4.0
|
||||
- jinaai/jina-reranker-v2-base-multilingual
|
||||
- License: cc-by-nc-4.0
|
||||
- jinaai/jina-embeddings-v3
|
||||
- License: cc-by-nc-4.0
|
||||
|
||||
These models are developed by Jina (https://jina.ai/) and are subject to Jina AI's licensing terms.
|
||||
|
||||
This distribution includes the following Google models, each with its respective license:
|
||||
- vidore/colpali-v1.3
|
||||
- License: gemma
|
||||
|
||||
Gemma is provided under and subject to the Gemma Terms of Use found at ai.google.dev/gemma/terms
|
||||
|
||||
Additional Notes:
|
||||
This project also includes third-party libraries with their respective licenses. Please refer to the documentation of each library for details regarding its usage and licensing terms.
|
||||
@@ -8,9 +8,9 @@ The default text embedding (`TextEmbedding`) model is Flag Embedding, presented
|
||||
|
||||
1. Light: FastEmbed is a lightweight library with few external dependencies. We don't require a GPU and don't download GBs of PyTorch dependencies, and instead use the ONNX Runtime. This makes it a great candidate for serverless runtimes like AWS Lambda.
|
||||
|
||||
2. Fast: FastEmbed is designed for speed. We use the ONNX Runtime, which is faster than PyTorch. We also use data-parallelism for encoding large datasets.
|
||||
2. Fast: FastEmbed is designed for speed. We use the ONNX Runtime, which is faster than PyTorch. We also use data parallelism for encoding large datasets.
|
||||
|
||||
3. Accurate: FastEmbed is better than OpenAI Ada-002. We also [supported](https://qdrant.github.io/fastembed/examples/Supported_Models/) an ever expanding set of models, including a few multilingual models.
|
||||
3. Accurate: FastEmbed is better than OpenAI Ada-002. We also [support](https://qdrant.github.io/fastembed/examples/Supported_Models/) an ever-expanding set of models, including a few multilingual models.
|
||||
|
||||
## 🚀 Installation
|
||||
|
||||
@@ -28,10 +28,10 @@ pip install fastembed-gpu
|
||||
|
||||
```python
|
||||
from fastembed import TextEmbedding
|
||||
from typing import List
|
||||
|
||||
|
||||
# Example list of documents
|
||||
documents: List[str] = [
|
||||
documents: list[str] = [
|
||||
"This is built to be faster and lighter than other embedding libraries e.g. Transformers, Sentence-Transformers, etc.",
|
||||
"fastembed is supported by and maintained by Qdrant.",
|
||||
]
|
||||
@@ -54,7 +54,7 @@ The list of all the available models can be found [here](https://qdrant.github.i
|
||||
from fastembed import TextEmbedding
|
||||
|
||||
model = TextEmbedding(model_name="BAAI/bge-small-en-v1.5")
|
||||
embeddings = list(embedding_model.embed(documents))
|
||||
embeddings = list(model.embed(documents))
|
||||
|
||||
# [
|
||||
# array([-0.1115, 0.0097, 0.0052, 0.0195, ...], dtype=float32),
|
||||
@@ -73,7 +73,7 @@ embeddings = list(embedding_model.embed(documents))
|
||||
from fastembed import SparseTextEmbedding
|
||||
|
||||
model = SparseTextEmbedding(model_name="prithivida/Splade_PP_en_v1")
|
||||
embeddings = list(embedding_model.embed(documents))
|
||||
embeddings = list(model.embed(documents))
|
||||
|
||||
# [
|
||||
# SparseEmbedding(indices=[ 17, 123, 919, ... ], values=[0.71, 0.22, 0.39, ...]),
|
||||
@@ -88,7 +88,7 @@ embeddings = list(embedding_model.embed(documents))
|
||||
from fastembed import SparseTextEmbedding
|
||||
|
||||
model = SparseTextEmbedding(model_name="Qdrant/bm42-all-minilm-l6-v2-attentions")
|
||||
embeddings = list(embedding_model.embed(documents))
|
||||
embeddings = list(model.embed(documents))
|
||||
|
||||
# [
|
||||
# SparseEmbedding(indices=[ 17, 123, 919, ... ], values=[0.71, 0.22, 0.39, ...]),
|
||||
@@ -104,7 +104,7 @@ embeddings = list(embedding_model.embed(documents))
|
||||
from fastembed import LateInteractionTextEmbedding
|
||||
|
||||
model = LateInteractionTextEmbedding(model_name="colbert-ir/colbertv2.0")
|
||||
embeddings = list(embedding_model.embed(documents))
|
||||
embeddings = list(model.embed(documents))
|
||||
|
||||
# [
|
||||
# array([
|
||||
@@ -129,7 +129,7 @@ images = [
|
||||
]
|
||||
|
||||
model = ImageEmbedding(model_name="Qdrant/clip-ViT-B-32-vision")
|
||||
embeddings = list(embedding_model.embed(images))
|
||||
embeddings = list(model.embed(images))
|
||||
|
||||
# [
|
||||
# array([-0.1115, 0.0097, 0.0052, 0.0195, ...], dtype=float32),
|
||||
@@ -137,6 +137,20 @@ embeddings = list(embedding_model.embed(images))
|
||||
# ]
|
||||
```
|
||||
|
||||
### 🔄 Rerankers
|
||||
```python
|
||||
from fastembed.rerank.cross_encoder import TextCrossEncoder
|
||||
|
||||
query = "Who is maintaining Qdrant?"
|
||||
documents: list[str] = [
|
||||
"This is built to be faster and lighter than other embedding libraries e.g. Transformers, Sentence-Transformers, etc.",
|
||||
"fastembed is supported by and maintained by Qdrant.",
|
||||
]
|
||||
encoder = TextCrossEncoder(model_name="Xenova/ms-marco-MiniLM-L-6-v2")
|
||||
scores = list(encoder.rerank(query, documents))
|
||||
|
||||
# [-11.48061752319336, 5.472434997558594]
|
||||
```
|
||||
|
||||
## ⚡️ FastEmbed on a GPU
|
||||
|
||||
@@ -147,7 +161,7 @@ It requires installation of the `fastembed-gpu` package.
|
||||
pip install fastembed-gpu
|
||||
```
|
||||
|
||||
Check our [example](https://qdrant.github.io/fastembed/examples/FastEmbed_GPU/) for the detailed instructions and CUDA 12.x support.
|
||||
Check our [example](https://qdrant.github.io/fastembed/examples/FastEmbed_GPU/) for detailed instructions, CUDA 12.x support and troubleshooting of the common issues.
|
||||
|
||||
```python
|
||||
from fastembed import TextEmbedding
|
||||
@@ -209,4 +223,4 @@ search_result = client.query(
|
||||
query_text="This is a query document"
|
||||
)
|
||||
print(search_result)
|
||||
```
|
||||
```
|
||||
|
||||
+1
-1
@@ -12,7 +12,7 @@ This is a guide how to release `fastembed` and `fastembed-gpu` packages.
|
||||
```bash
|
||||
git checkout gpu
|
||||
git rebase main
|
||||
git push origin gpu
|
||||
git push -f origin gpu
|
||||
```
|
||||
|
||||
4. Draft release notes
|
||||
|
||||
@@ -65,15 +65,13 @@
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"from typing import List\n",
|
||||
"\n",
|
||||
"import numpy as np\n",
|
||||
"\n",
|
||||
"from fastembed import TextEmbedding\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"# Example list of documents\n",
|
||||
"documents: List[str] = [\n",
|
||||
"documents: list[str] = [\n",
|
||||
" \"This is built to be faster and lighter than other embedding libraries e.g. Transformers, Sentence-Transformers, etc.\",\n",
|
||||
" \"fastembed is supported by and maintained by Qdrant.\",\n",
|
||||
"]\n",
|
||||
|
||||
@@ -54,7 +54,14 @@
|
||||
},
|
||||
{
|
||||
"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'}]"
|
||||
"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": {},
|
||||
@@ -212,7 +219,9 @@
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": "((26, 128), (32, 128))"
|
||||
"text/plain": [
|
||||
"((26, 128), (32, 128))"
|
||||
]
|
||||
},
|
||||
"execution_count": 18,
|
||||
"metadata": {},
|
||||
@@ -271,7 +280,9 @@
|
||||
"import numpy as np\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"def compute_relevance_scores(query_embedding: np.array, document_embeddings: np.array, k: int):\n",
|
||||
"def compute_relevance_scores(\n",
|
||||
" query_embedding: np.array, document_embeddings: np.array, k: int\n",
|
||||
") -> list[int]:\n",
|
||||
" \"\"\"\n",
|
||||
" Compute relevance scores for top-k documents given a query.\n",
|
||||
"\n",
|
||||
|
||||
@@ -37,28 +37,7 @@
|
||||
"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"
|
||||
"**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`."
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -67,9 +46,65 @@
|
||||
"id": "3xx3r-9jgAMi"
|
||||
},
|
||||
"source": [
|
||||
"### CUDA 12.x support\n",
|
||||
"You can check your CUDA version using such commands as `nvidia-smi` or `nvcc --version`\n",
|
||||
"\n",
|
||||
"Google Colab notebooks have CUDA 12.x."
|
||||
"Starting from version 1.19.0, onnxruntime-gpu ships with support for CUDA 12.x by default.\n",
|
||||
"\n",
|
||||
"Google Colab notebooks have by default CUDA 12.x and CuDNN 8.x.\n",
|
||||
"\n",
|
||||
"Latest version of `onnxruntime-gpu` requires CuDNN 9.x, in order to install it you can run the following command: "
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"!sudo apt install cudnn9\n",
|
||||
"!pip install fastembed-gpu -qqq"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"If it necessary to work with CuDNN 8, you can consider locking `onnxruntime-gpu` to 1.18.0 with CUDA 12.x by this command:"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"!pip install onnxruntime-gpu==1.18.0 -i https://aiinfra.pkgs.visualstudio.com/PublicPackages/_packaging/onnxruntime-cuda-12/pypi/simple/ -qq\n",
|
||||
"!pip install fastembed-gpu -qqq"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### CUDA 11.x support\n",
|
||||
"To use latest version of `onnxruntime-gpu` with CUDA 11.x, you can run the following command:"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"!pip install onnxruntime-gpu -i https://aiinfra.pkgs.visualstudio.com/PublicPackages/_packaging/onnxruntime-cuda-11/pypi/simple/ -qq"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"**NOTE**: Ensure that CuDNN 9.x is installed when working with the latest `onnxruntime-gpu`, whether using CUDA 11.x or 12.x."
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -82,7 +117,80 @@
|
||||
"\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)"
|
||||
"The dependencies required for the chosen onnxruntime version are listed in the [CUDA Execution Provider requirements](https://onnxruntime.ai/docs/execution-providers/CUDA-ExecutionProvider.html#requirements)."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### Setting up fastembed-gpu on GCP\n",
|
||||
"\n",
|
||||
"#### CUDA drivers\n",
|
||||
"[CUDA 11.8 toolkit](https://developer.nvidia.com/cuda-11-8-0-download-archive) or [CUDA 12.x toolkit](https://developer.nvidia.com/cuda-downloads) has to be installed if they haven't yet been set up.\n",
|
||||
"\n",
|
||||
"#### Example of setting up CUDA 12.x on Ubuntu 22.04\n",
|
||||
"Make sure to download an archive which has been created for your particular platform, CPU architecture and OS distribution.\n",
|
||||
"\n",
|
||||
"For Ubuntu 22.04 with x86_64 CPU architecture the following [archive](https://developer.nvidia.com/cuda-downloads?target_os=Linux&target_arch=x86_64&Distribution=Ubuntu&target_version=22.04&target_type=deb_network) has to be downloaded.\n",
|
||||
"\n",
|
||||
"```bash\n",
|
||||
"wget https://developer.download.nvidia.com/compute/cuda/repos/ubuntu2204/x86_64/cuda-keyring_1.1-1_all.deb\n",
|
||||
"sudo dpkg -i cuda-keyring_1.1-1_all.deb\n",
|
||||
"sudo apt-get update\n",
|
||||
"sudo apt-get -y install cuda\n",
|
||||
"```\n",
|
||||
"**NOTE**: Specific CUDA libraries can be found in the [meta packages section](https://docs.nvidia.com/cuda/cuda-installation-guide-linux/#meta-packages) in the CUDA installation guide.\n",
|
||||
"\n",
|
||||
"**NOTE**: When installing CUDA, the environment variable might not be set by default. Make sure to add the following line to your environment variables:\n",
|
||||
"```bash\n",
|
||||
"LD_LIBRARY_PATH=/usr/local/cuda/lib64:$LD_LIBRARY_PATH\n",
|
||||
"```\n",
|
||||
"This will ensure that the CUDA libraries are properly linked.\n",
|
||||
"\n",
|
||||
"#### CuDNN 9.x\n",
|
||||
"CuDNN 9.x library can be installed via the following [archive](https://developer.nvidia.com/rdp/cudnn-archive).\n",
|
||||
"\n",
|
||||
"#### Example of setting up CuDNN 9.x on Ubuntu 22.04\n",
|
||||
"CuDNN 9.x for Ubuntu 22.04 x86_64 [archive](https://developer.nvidia.com/cudnn-downloads?target_os=Linux&target_arch=x86_64&Distribution=Ubuntu&target_version=22.04&target_type=deb_network) can be downloaded and installed in the following way:\n",
|
||||
"```bash\n",
|
||||
"wget https://developer.download.nvidia.com/compute/cuda/repos/ubuntu2204/x86_64/cuda-keyring_1.1-1_all.deb\n",
|
||||
"sudo dpkg -i cuda-keyring_1.1-1_all.deb\n",
|
||||
"sudo apt-get update\n",
|
||||
"sudo apt-get -y install cudnn\n",
|
||||
"```\n",
|
||||
"**NOTE**: When installing CuDNN, you can choose specific version, cudnn-cuda-11 or cudnn-cuda-12"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### Common issues\n",
|
||||
"\n",
|
||||
"The following are some common issues that may arise while using `fastembed-gpu` if not installed properly:\n",
|
||||
"\n",
|
||||
"CUDA library is not installed:\n",
|
||||
"```bash\n",
|
||||
"FAIL : Failed to load library libonnxruntime_providers_cuda.so with error: libcublasLt.so.x: cannot open shared object file: No such file or directory\n",
|
||||
"```\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"CuDNN library is not installed:\n",
|
||||
"```bash\n",
|
||||
"FAIL : Failed to load library libonnxruntime_providers_cuda.so with error: libcudnn.so.x: cannot open shared object file: No such file or directory\n",
|
||||
"```\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"CUDA library path is not set:\n",
|
||||
"```bash\n",
|
||||
"FAIL : Failed to load library libonnxruntime_providers_cuda.so with error: libcufft.so.x: failed to map segment from shared object\n",
|
||||
"```\n",
|
||||
"\n",
|
||||
"Make sure to add the following line to your environment variables:\n",
|
||||
"```bash\n",
|
||||
"LD_LIBRARY_PATH=/usr/local/cuda/lib64:$LD_LIBRARY_PATH\n",
|
||||
"```"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -280,8 +388,6 @@
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"from typing import List\n",
|
||||
"\n",
|
||||
"import numpy as np\n",
|
||||
"\n",
|
||||
"from fastembed import TextEmbedding\n",
|
||||
@@ -299,9 +405,7 @@
|
||||
"id": "iPtoHf7GeV-i"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"documents: List[str] = list(np.repeat(\"Demonstrating GPU acceleration in fastembed\", 500))"
|
||||
]
|
||||
"source": "documents: list[str] = list(np.repeat(\"Demonstrating GPU acceleration in fastembed\", 500))"
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
|
||||
@@ -0,0 +1,88 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### Fastembed Multi-GPU Tutorial\n",
|
||||
"This tutorial demonstrates how to leverage multi-GPU support in Fastembed. Fastembed supports embedding text and images utilizing modern GPUs for acceleration. Let's explore how to use Fastembed with multiple GPUs step by step."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"#### Prerequisites\n",
|
||||
"To get started, ensure you have the following installed:\n",
|
||||
"- Python 3.9 or later\n",
|
||||
"- Fastembed (`pip install fastembed-gpu`)\n",
|
||||
"- Refer to [this](https://github.com/qdrant/fastembed/blob/main/docs/examples/FastEmbed_GPU.ipynb) tutorial if you have issues with GPU dependencies\n",
|
||||
"- Access to a multi-GPU server"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### Multi-GPU using cuda argument with TextEmbedding Model"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from fastembed import TextEmbedding\n",
|
||||
"\n",
|
||||
"# define the documents to embed\n",
|
||||
"docs = [\"hello world\", \"flag embedding\"] * 100\n",
|
||||
"\n",
|
||||
"# define gpu ids\n",
|
||||
"device_ids = [0, 1]\n",
|
||||
"\n",
|
||||
"if __name__ == \"__main__\":\n",
|
||||
" # initialize a TextEmbedding model using CUDA\n",
|
||||
" text_model = TextEmbedding(\n",
|
||||
" model_name=\"sentence-transformers/all-MiniLM-L6-v2\",\n",
|
||||
" cuda=True,\n",
|
||||
" device_ids=device_ids,\n",
|
||||
" lazy_load=True,\n",
|
||||
" )\n",
|
||||
"\n",
|
||||
" # generate embeddings\n",
|
||||
" text_embeddings = list(text_model.embed(docs, batch_size=2, parallel=len(device_ids)))\n",
|
||||
" print(text_embeddings)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"In this snippet:\n",
|
||||
"- `cuda=True` enables GPU acceleration.\n",
|
||||
"- `device_ids=[0, 1]` specifies GPUs to use. Replace `[0, 1]` with available GPU IDs.\n",
|
||||
"- `lazy_load=True`\n",
|
||||
"\n",
|
||||
"**NOTE**: When using multi-GPU settings, it is important to configure `parallel` and `lazy_load` properly to avoid inefficiencies:\n",
|
||||
"\n",
|
||||
"`parallel`: This parameter enables multi-GPU support by spawning child processes for each GPU specified in device_ids. To ensure proper utilization, the value of `parallel` must match the number of GPUs in device_ids. If using a single GPU, this parameter is not necessary.\n",
|
||||
"\n",
|
||||
"`lazy_load`: Enabling `lazy_load` prevents redundant memory usage. Without `lazy_load`, the model is initially loaded into the memory of the first GPU by the main process. When child processes are spawned for each GPU, the model is reloaded on the first GPU, causing redundant memory consumption and inefficiencies."
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": ".venv",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"name": "python",
|
||||
"version": "3.10.15"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 2
|
||||
}
|
||||
@@ -36,7 +36,7 @@
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import time\n",
|
||||
"from typing import Callable, List, Tuple\n",
|
||||
"from typing import Callable\n",
|
||||
"\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import torch.nn.functional as F\n",
|
||||
@@ -64,6 +64,7 @@
|
||||
],
|
||||
"source": [
|
||||
"import fastembed\n",
|
||||
"\n",
|
||||
"fastembed.__version__"
|
||||
]
|
||||
},
|
||||
@@ -98,7 +99,7 @@
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"documents: List[str] = [\n",
|
||||
"documents: list[str] = [\n",
|
||||
" \"Chandrayaan-3 is India's third lunar mission\",\n",
|
||||
" \"It aimed to land a rover on the Moon's surface - joining the US, China and Russia\",\n",
|
||||
" \"The mission is a follow-up to Chandrayaan-2, which had partial success\",\n",
|
||||
@@ -151,11 +152,11 @@
|
||||
" HuggingFace Transformer implementation of FlagEmbedding\n",
|
||||
" \"\"\"\n",
|
||||
"\n",
|
||||
" def __init__(self, model_id: str):\n",
|
||||
" def __init__(self, model_id: str) -> None:\n",
|
||||
" self.model = AutoModel.from_pretrained(model_id)\n",
|
||||
" self.tokenizer = AutoTokenizer.from_pretrained(model_id)\n",
|
||||
"\n",
|
||||
" def embed(self, texts: List[str]):\n",
|
||||
" def embed(self, texts: list[str]):\n",
|
||||
" encoded_input = self.tokenizer(\n",
|
||||
" texts, max_length=512, padding=True, truncation=True, return_tensors=\"pt\"\n",
|
||||
" )\n",
|
||||
@@ -254,7 +255,7 @@
|
||||
"\n",
|
||||
"def calculate_time_stats(\n",
|
||||
" embed_func: Callable, documents: list, k: int\n",
|
||||
") -> Tuple[float, float, float]:\n",
|
||||
") -> tuple[float, float, float]:\n",
|
||||
" times = []\n",
|
||||
" for _ in range(k):\n",
|
||||
" # Timing the embed_func call\n",
|
||||
@@ -309,7 +310,7 @@
|
||||
],
|
||||
"source": [
|
||||
"def plot_character_per_second_comparison(\n",
|
||||
" hf_stats: Tuple[float, float, float], fst_stats: Tuple[float, float, float], documents: list\n",
|
||||
" hf_stats: tuple[float, float, float], fst_stats: tuple[float, float, float], documents: list\n",
|
||||
"):\n",
|
||||
" # Calculating total characters in documents\n",
|
||||
" total_characters = sum(len(doc) for doc in documents)\n",
|
||||
|
||||
@@ -44,7 +44,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 21,
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-03-30T00:45:24.814968Z",
|
||||
@@ -58,8 +58,6 @@
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from typing import List\n",
|
||||
"\n",
|
||||
"import numpy as np\n",
|
||||
"from datasets import load_dataset\n",
|
||||
"from peft import AutoPeftModelForCausalLM\n",
|
||||
@@ -72,11 +70,11 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 23,
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"hf_token = <YOUR_HF_TOKEN_HERE> # Get your token from https://huggingface.co/settings/token, needed for Gemma weights"
|
||||
"hf_token = \"<YOUR_HF_TOKEN_HERE>\" # Get your token from https://huggingface.co/settings/token, needed for Gemma weights"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -246,7 +244,7 @@
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"context_embeddings: List[np.ndarray] = list(\n",
|
||||
"context_embeddings: list[np.ndarray] = list(\n",
|
||||
" embedding_model.embed(contexts)\n",
|
||||
") # Note the list() call - this is a generator"
|
||||
]
|
||||
|
||||
@@ -50,7 +50,6 @@
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import json\n",
|
||||
"from typing import List, Tuple\n",
|
||||
"\n",
|
||||
"import numpy as np\n",
|
||||
"import pandas as pd\n",
|
||||
@@ -489,11 +488,11 @@
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"def make_sparse_embedding(texts: List[str]):\n",
|
||||
"def make_sparse_embedding(texts: list[str]) -> list[SparseEmbedding]:\n",
|
||||
" return list(sparse_model.embed(texts, batch_size=32))\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"sparse_embedding: List[SparseEmbedding] = make_sparse_embedding(\n",
|
||||
"sparse_embedding: list[SparseEmbedding] = make_sparse_embedding(\n",
|
||||
" [\"Fastembed is a great library for text embeddings!\"]\n",
|
||||
")\n",
|
||||
"sparse_embedding"
|
||||
@@ -616,7 +615,7 @@
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"def get_tokens_and_weights(sparse_embedding, model_name):\n",
|
||||
"def get_tokens_and_weights(sparse_embedding, model_name) -> dict[str, float]:\n",
|
||||
" # Find the tokenizer for the model\n",
|
||||
" tokenizer_source = None\n",
|
||||
" for model_info in SparseTextEmbedding.list_supported_models():\n",
|
||||
@@ -627,7 +626,7 @@
|
||||
" raise ValueError(f\"Model {model_name} not found in the supported models.\")\n",
|
||||
"\n",
|
||||
" tokenizer = AutoTokenizer.from_pretrained(tokenizer_source)\n",
|
||||
" token_weight_dict = {}\n",
|
||||
" token_weight_dict: dict[str, float] = {}\n",
|
||||
" for i in range(len(sparse_embedding.indices)):\n",
|
||||
" token = tokenizer.decode([sparse_embedding.indices[i]])\n",
|
||||
" weight = sparse_embedding.values[i]\n",
|
||||
@@ -662,7 +661,7 @@
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def make_dense_embedding(texts: List[str]):\n",
|
||||
"def make_dense_embedding(texts: list[str]):\n",
|
||||
" return list(dense_model.embed(texts))\n",
|
||||
"\n",
|
||||
"\n",
|
||||
@@ -872,7 +871,7 @@
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def make_points(df: pd.DataFrame) -> List[PointStruct]:\n",
|
||||
"def make_points(df: pd.DataFrame) -> list[PointStruct]:\n",
|
||||
" sparse_vectors = df[\"sparse_embedding\"].tolist()\n",
|
||||
" product_texts = df[\"combined_text\"].tolist()\n",
|
||||
" dense_vectors = df[\"dense_embedding\"].tolist()\n",
|
||||
@@ -899,7 +898,7 @@
|
||||
" return points\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"points: List[PointStruct] = make_points(df)"
|
||||
"points: list[PointStruct] = make_points(df)"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -942,8 +941,8 @@
|
||||
"source": [
|
||||
"def search(query_text: str):\n",
|
||||
" # # Compute sparse and dense vectors\n",
|
||||
" query_sparse_vectors: List[SparseEmbedding] = make_sparse_embedding([query_text])\n",
|
||||
" query_dense_vector: List[np.ndarray] = make_dense_embedding([query_text])\n",
|
||||
" query_sparse_vectors: list[SparseEmbedding] = make_sparse_embedding([query_text])\n",
|
||||
" query_dense_vector: list[np.ndarray] = make_dense_embedding([query_text])\n",
|
||||
"\n",
|
||||
" search_results = client.search_batch(\n",
|
||||
" collection_name=collection_name,\n",
|
||||
@@ -1075,7 +1074,7 @@
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def rank_list(search_result: List[ScoredPoint]):\n",
|
||||
"def rank_list(search_result: list[ScoredPoint]):\n",
|
||||
" return [(point.id, rank + 1) for rank, point in enumerate(search_result)]\n",
|
||||
"\n",
|
||||
"\n",
|
||||
@@ -1149,7 +1148,7 @@
|
||||
],
|
||||
"source": [
|
||||
"def find_point_by_id(\n",
|
||||
" client: QdrantClient, collection_name: str, rrf_rank_list: List[Tuple[int, float]]\n",
|
||||
" client: QdrantClient, collection_name: str, rrf_rank_list: list[tuple[int, float]]\n",
|
||||
"):\n",
|
||||
" return client.retrieve(\n",
|
||||
" collection_name=collection_name, ids=[item[0] for item in rrf_rank_list]\n",
|
||||
|
||||
@@ -47,7 +47,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-03-30T00:49:20.516644Z",
|
||||
@@ -56,8 +56,7 @@
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from fastembed import SparseTextEmbedding, SparseEmbedding\n",
|
||||
"from typing import List"
|
||||
"from fastembed import SparseTextEmbedding, SparseEmbedding"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -134,7 +133,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 5,
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-03-30T00:49:28.624109Z",
|
||||
@@ -143,7 +142,7 @@
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"documents: List[str] = [\n",
|
||||
"documents: list[str] = [\n",
|
||||
" \"Chandrayaan-3 is India's third lunar mission\",\n",
|
||||
" \"It aimed to land a rover on the Moon's surface - joining the US, China and Russia\",\n",
|
||||
" \"The mission is a follow-up to Chandrayaan-2, which had partial success\",\n",
|
||||
@@ -157,7 +156,7 @@
|
||||
" \"Chandrayaan-3 was launched from the Satish Dhawan Space Centre in Sriharikota\",\n",
|
||||
" \"Chandrayaan-3 was launched earlier in the year 2023\",\n",
|
||||
"]\n",
|
||||
"sparse_embeddings_list: List[SparseEmbedding] = list(\n",
|
||||
"sparse_embeddings_list: list[SparseEmbedding] = list(\n",
|
||||
" model.embed(documents, batch_size=6)\n",
|
||||
") # batch_size is optional, notice the generator"
|
||||
]
|
||||
@@ -235,7 +234,9 @@
|
||||
"source": [
|
||||
"# Let's print the first 5 features and their weights for better understanding.\n",
|
||||
"for i in range(5):\n",
|
||||
" print(f\"Token at index {sparse_embeddings_list[0].indices[i]} has weight {sparse_embeddings_list[0].values[i]}\")"
|
||||
" print(\n",
|
||||
" f\"Token at index {sparse_embeddings_list[0].indices[i]} has weight {sparse_embeddings_list[0].values[i]}\"\n",
|
||||
" )"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -261,7 +262,9 @@
|
||||
"import json\n",
|
||||
"from transformers import AutoTokenizer\n",
|
||||
"\n",
|
||||
"tokenizer = AutoTokenizer.from_pretrained(SparseTextEmbedding.list_supported_models()[0][\"sources\"][\"hf\"])"
|
||||
"tokenizer = AutoTokenizer.from_pretrained(\n",
|
||||
" SparseTextEmbedding.list_supported_models()[0][\"sources\"][\"hf\"]\n",
|
||||
")"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -326,7 +329,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",
|
||||
|
||||
@@ -2,13 +2,16 @@
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 4,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-05-31T18:13:23.806907Z",
|
||||
"start_time": "2024-05-31T18:13:23.797078Z"
|
||||
"end_time": "2024-11-13T09:01:03.324551Z",
|
||||
"start_time": "2024-11-13T09:01:03.234711Z"
|
||||
}
|
||||
},
|
||||
"source": [
|
||||
"%load_ext autoreload\n",
|
||||
"%autoreload 2"
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
@@ -19,21 +22,16 @@
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"%load_ext autoreload\n",
|
||||
"%autoreload 2"
|
||||
]
|
||||
"execution_count": 10
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-05-31T18:14:31.147674Z",
|
||||
"start_time": "2024-05-31T18:14:31.134015Z"
|
||||
"end_time": "2024-11-13T09:01:04.505772Z",
|
||||
"start_time": "2024-11-13T09:01:04.493296Z"
|
||||
}
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import pandas as pd\n",
|
||||
"\n",
|
||||
@@ -42,8 +40,11 @@
|
||||
" TextEmbedding,\n",
|
||||
" LateInteractionTextEmbedding,\n",
|
||||
" ImageEmbedding,\n",
|
||||
")"
|
||||
]
|
||||
")\n",
|
||||
"from fastembed.rerank.cross_encoder import TextCrossEncoder"
|
||||
],
|
||||
"outputs": [],
|
||||
"execution_count": 11
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
@@ -54,16 +55,79 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 6,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-05-31T18:13:25.863008Z",
|
||||
"start_time": "2024-05-31T18:13:25.837795Z"
|
||||
"end_time": "2024-11-13T09:01:05.812271Z",
|
||||
"start_time": "2024-11-13T09:01:05.795846Z"
|
||||
}
|
||||
},
|
||||
"source": [
|
||||
"supported_models = (\n",
|
||||
" pd.DataFrame(TextEmbedding.list_supported_models())\n",
|
||||
" .sort_values(\"size_in_GB\")\n",
|
||||
" .drop(columns=[\"sources\", \"model_file\", \"additional_files\"])\n",
|
||||
" .reset_index(drop=True)\n",
|
||||
")\n",
|
||||
"supported_models"
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
" model dim \\\n",
|
||||
"0 BAAI/bge-small-en-v1.5 384 \n",
|
||||
"1 BAAI/bge-small-zh-v1.5 512 \n",
|
||||
"2 snowflake/snowflake-arctic-embed-xs 384 \n",
|
||||
"3 sentence-transformers/all-MiniLM-L6-v2 384 \n",
|
||||
"4 jinaai/jina-embeddings-v2-small-en 512 \n",
|
||||
"5 BAAI/bge-small-en 384 \n",
|
||||
"6 snowflake/snowflake-arctic-embed-s 384 \n",
|
||||
"7 nomic-ai/nomic-embed-text-v1.5-Q 768 \n",
|
||||
"8 BAAI/bge-base-en-v1.5 768 \n",
|
||||
"9 sentence-transformers/paraphrase-multilingual-... 384 \n",
|
||||
"10 Qdrant/clip-ViT-B-32-text 512 \n",
|
||||
"11 jinaai/jina-embeddings-v2-base-de 768 \n",
|
||||
"12 BAAI/bge-base-en 768 \n",
|
||||
"13 snowflake/snowflake-arctic-embed-m 768 \n",
|
||||
"14 nomic-ai/nomic-embed-text-v1.5 768 \n",
|
||||
"15 jinaai/jina-embeddings-v2-base-en 768 \n",
|
||||
"16 nomic-ai/nomic-embed-text-v1 768 \n",
|
||||
"17 snowflake/snowflake-arctic-embed-m-long 768 \n",
|
||||
"18 mixedbread-ai/mxbai-embed-large-v1 1024 \n",
|
||||
"19 jinaai/jina-embeddings-v2-base-code 768 \n",
|
||||
"20 sentence-transformers/paraphrase-multilingual-... 768 \n",
|
||||
"21 snowflake/snowflake-arctic-embed-l 1024 \n",
|
||||
"22 thenlper/gte-large 1024 \n",
|
||||
"23 BAAI/bge-large-en-v1.5 1024 \n",
|
||||
"24 intfloat/multilingual-e5-large 1024 \n",
|
||||
"\n",
|
||||
" description license size_in_GB \n",
|
||||
"0 Text embeddings, Unimodal (text), English, 512... mit 0.067 \n",
|
||||
"1 Text embeddings, Unimodal (text), Chinese, 512... mit 0.090 \n",
|
||||
"2 Text embeddings, Unimodal (text), English, 512... apache-2.0 0.090 \n",
|
||||
"3 Text embeddings, Unimodal (text), English, 256... apache-2.0 0.090 \n",
|
||||
"4 Text embeddings, Unimodal (text), English, 819... apache-2.0 0.120 \n",
|
||||
"5 Text embeddings, Unimodal (text), English, 512... mit 0.130 \n",
|
||||
"6 Text embeddings, Unimodal (text), English, 512... apache-2.0 0.130 \n",
|
||||
"7 Text embeddings, Multimodal (text, image), Eng... apache-2.0 0.130 \n",
|
||||
"8 Text embeddings, Unimodal (text), English, 512... mit 0.210 \n",
|
||||
"9 Text embeddings, Unimodal (text), Multilingual... apache-2.0 0.220 \n",
|
||||
"10 Text embeddings, Multimodal (text&image), Engl... mit 0.250 \n",
|
||||
"11 Text embeddings, Unimodal (text), Multilingual... apache-2.0 0.320 \n",
|
||||
"12 Text embeddings, Unimodal (text), English, 512... mit 0.420 \n",
|
||||
"13 Text embeddings, Unimodal (text), English, 512... apache-2.0 0.430 \n",
|
||||
"14 Text embeddings, Multimodal (text, image), Eng... apache-2.0 0.520 \n",
|
||||
"15 Text embeddings, Unimodal (text), English, 819... apache-2.0 0.520 \n",
|
||||
"16 Text embeddings, Multimodal (text, image), Eng... apache-2.0 0.520 \n",
|
||||
"17 Text embeddings, Unimodal (text), English, 204... apache-2.0 0.540 \n",
|
||||
"18 Text embeddings, Unimodal (text), English, 512... apache-2.0 0.640 \n",
|
||||
"19 Text embeddings, Unimodal (text), Multilingual... apache-2.0 0.640 \n",
|
||||
"20 Text embeddings, Unimodal (text), Multilingual... apache-2.0 1.000 \n",
|
||||
"21 Text embeddings, Unimodal (text), English, 512... apache-2.0 1.020 \n",
|
||||
"22 Text embeddings, Unimodal (text), English, 512... mit 1.200 \n",
|
||||
"23 Text embeddings, Unimodal (text), English, 512... mit 1.200 \n",
|
||||
"24 Text embeddings, Unimodal (text), Multilingual... mit 2.240 "
|
||||
],
|
||||
"text/html": [
|
||||
"<div>\n",
|
||||
"<style scoped>\n",
|
||||
@@ -86,6 +150,7 @@
|
||||
" <th>model</th>\n",
|
||||
" <th>dim</th>\n",
|
||||
" <th>description</th>\n",
|
||||
" <th>license</th>\n",
|
||||
" <th>size_in_GB</th>\n",
|
||||
" </tr>\n",
|
||||
" </thead>\n",
|
||||
@@ -94,233 +159,213 @@
|
||||
" <th>0</th>\n",
|
||||
" <td>BAAI/bge-small-en-v1.5</td>\n",
|
||||
" <td>384</td>\n",
|
||||
" <td>Fast and Default English model</td>\n",
|
||||
" <td>Text embeddings, Unimodal (text), English, 512...</td>\n",
|
||||
" <td>mit</td>\n",
|
||||
" <td>0.067</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>1</th>\n",
|
||||
" <td>BAAI/bge-small-zh-v1.5</td>\n",
|
||||
" <td>512</td>\n",
|
||||
" <td>Fast and recommended Chinese model</td>\n",
|
||||
" <td>Text embeddings, Unimodal (text), Chinese, 512...</td>\n",
|
||||
" <td>mit</td>\n",
|
||||
" <td>0.090</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>2</th>\n",
|
||||
" <td>sentence-transformers/all-MiniLM-L6-v2</td>\n",
|
||||
" <td>snowflake/snowflake-arctic-embed-xs</td>\n",
|
||||
" <td>384</td>\n",
|
||||
" <td>Sentence Transformer model, MiniLM-L6-v2</td>\n",
|
||||
" <td>Text embeddings, Unimodal (text), English, 512...</td>\n",
|
||||
" <td>apache-2.0</td>\n",
|
||||
" <td>0.090</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>3</th>\n",
|
||||
" <td>snowflake/snowflake-arctic-embed-xs</td>\n",
|
||||
" <td>sentence-transformers/all-MiniLM-L6-v2</td>\n",
|
||||
" <td>384</td>\n",
|
||||
" <td>Based on all-MiniLM-L6-v2 model with only 22m ...</td>\n",
|
||||
" <td>Text embeddings, Unimodal (text), English, 256...</td>\n",
|
||||
" <td>apache-2.0</td>\n",
|
||||
" <td>0.090</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>4</th>\n",
|
||||
" <td>jinaai/jina-embeddings-v2-small-en</td>\n",
|
||||
" <td>512</td>\n",
|
||||
" <td>English embedding model supporting 8192 sequen...</td>\n",
|
||||
" <td>Text embeddings, Unimodal (text), English, 819...</td>\n",
|
||||
" <td>apache-2.0</td>\n",
|
||||
" <td>0.120</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>5</th>\n",
|
||||
" <td>snowflake/snowflake-arctic-embed-s</td>\n",
|
||||
" <td>BAAI/bge-small-en</td>\n",
|
||||
" <td>384</td>\n",
|
||||
" <td>Based on infloat/e5-small-unsupervised, does n...</td>\n",
|
||||
" <td>Text embeddings, Unimodal (text), English, 512...</td>\n",
|
||||
" <td>mit</td>\n",
|
||||
" <td>0.130</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>6</th>\n",
|
||||
" <td>BAAI/bge-small-en</td>\n",
|
||||
" <td>snowflake/snowflake-arctic-embed-s</td>\n",
|
||||
" <td>384</td>\n",
|
||||
" <td>Fast English model</td>\n",
|
||||
" <td>Text embeddings, Unimodal (text), English, 512...</td>\n",
|
||||
" <td>apache-2.0</td>\n",
|
||||
" <td>0.130</td>\n",
|
||||
" </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>Text embeddings, Multimodal (text, image), Eng...</td>\n",
|
||||
" <td>apache-2.0</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>Text embeddings, Unimodal (text), English, 512...</td>\n",
|
||||
" <td>mit</td>\n",
|
||||
" <td>0.210</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\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>Text embeddings, Unimodal (text), Multilingual...</td>\n",
|
||||
" <td>apache-2.0</td>\n",
|
||||
" <td>0.220</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\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>Text embeddings, Multimodal (text&image), Engl...</td>\n",
|
||||
" <td>mit</td>\n",
|
||||
" <td>0.250</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>11</th>\n",
|
||||
" <td>BAAI/bge-base-en</td>\n",
|
||||
" <td>jinaai/jina-embeddings-v2-base-de</td>\n",
|
||||
" <td>768</td>\n",
|
||||
" <td>Base English model</td>\n",
|
||||
" <td>0.420</td>\n",
|
||||
" <td>Text embeddings, Unimodal (text), Multilingual...</td>\n",
|
||||
" <td>apache-2.0</td>\n",
|
||||
" <td>0.320</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>12</th>\n",
|
||||
" <td>snowflake/snowflake-arctic-embed-m</td>\n",
|
||||
" <td>BAAI/bge-base-en</td>\n",
|
||||
" <td>768</td>\n",
|
||||
" <td>Based on intfloat/e5-base-unsupervised model, ...</td>\n",
|
||||
" <td>0.430</td>\n",
|
||||
" <td>Text embeddings, Unimodal (text), English, 512...</td>\n",
|
||||
" <td>mit</td>\n",
|
||||
" <td>0.420</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>13</th>\n",
|
||||
" <td>nomic-ai/nomic-embed-text-v1</td>\n",
|
||||
" <td>snowflake/snowflake-arctic-embed-m</td>\n",
|
||||
" <td>768</td>\n",
|
||||
" <td>8192 context length english model</td>\n",
|
||||
" <td>0.520</td>\n",
|
||||
" <td>Text embeddings, Unimodal (text), English, 512...</td>\n",
|
||||
" <td>apache-2.0</td>\n",
|
||||
" <td>0.430</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>14</th>\n",
|
||||
" <td>jinaai/jina-embeddings-v2-base-en</td>\n",
|
||||
" <td>nomic-ai/nomic-embed-text-v1.5</td>\n",
|
||||
" <td>768</td>\n",
|
||||
" <td>English embedding model supporting 8192 sequen...</td>\n",
|
||||
" <td>Text embeddings, Multimodal (text, image), Eng...</td>\n",
|
||||
" <td>apache-2.0</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>jinaai/jina-embeddings-v2-base-en</td>\n",
|
||||
" <td>768</td>\n",
|
||||
" <td>8192 context length english model</td>\n",
|
||||
" <td>Text embeddings, Unimodal (text), English, 819...</td>\n",
|
||||
" <td>apache-2.0</td>\n",
|
||||
" <td>0.520</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>16</th>\n",
|
||||
" <td>snowflake/snowflake-arctic-embed-m-long</td>\n",
|
||||
" <td>nomic-ai/nomic-embed-text-v1</td>\n",
|
||||
" <td>768</td>\n",
|
||||
" <td>Based on nomic-ai/nomic-embed-text-v1-unsuperv...</td>\n",
|
||||
" <td>0.540</td>\n",
|
||||
" <td>Text embeddings, Multimodal (text, image), Eng...</td>\n",
|
||||
" <td>apache-2.0</td>\n",
|
||||
" <td>0.520</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\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",
|
||||
" <td>snowflake/snowflake-arctic-embed-m-long</td>\n",
|
||||
" <td>768</td>\n",
|
||||
" <td>Text embeddings, Unimodal (text), English, 204...</td>\n",
|
||||
" <td>apache-2.0</td>\n",
|
||||
" <td>0.540</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\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",
|
||||
" <td>mixedbread-ai/mxbai-embed-large-v1</td>\n",
|
||||
" <td>1024</td>\n",
|
||||
" <td>Text embeddings, Unimodal (text), English, 512...</td>\n",
|
||||
" <td>apache-2.0</td>\n",
|
||||
" <td>0.640</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\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",
|
||||
" <td>jinaai/jina-embeddings-v2-base-code</td>\n",
|
||||
" <td>768</td>\n",
|
||||
" <td>Text embeddings, Unimodal (text), Multilingual...</td>\n",
|
||||
" <td>apache-2.0</td>\n",
|
||||
" <td>0.640</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\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",
|
||||
" <td>sentence-transformers/paraphrase-multilingual-...</td>\n",
|
||||
" <td>768</td>\n",
|
||||
" <td>Text embeddings, Unimodal (text), Multilingual...</td>\n",
|
||||
" <td>apache-2.0</td>\n",
|
||||
" <td>1.000</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>21</th>\n",
|
||||
" <td>thenlper/gte-large</td>\n",
|
||||
" <td>snowflake/snowflake-arctic-embed-l</td>\n",
|
||||
" <td>1024</td>\n",
|
||||
" <td>Large general text embeddings model</td>\n",
|
||||
" <td>1.200</td>\n",
|
||||
" <td>Text embeddings, Unimodal (text), English, 512...</td>\n",
|
||||
" <td>apache-2.0</td>\n",
|
||||
" <td>1.020</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>22</th>\n",
|
||||
" <td>thenlper/gte-large</td>\n",
|
||||
" <td>1024</td>\n",
|
||||
" <td>Text embeddings, Unimodal (text), English, 512...</td>\n",
|
||||
" <td>mit</td>\n",
|
||||
" <td>1.200</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>23</th>\n",
|
||||
" <td>BAAI/bge-large-en-v1.5</td>\n",
|
||||
" <td>1024</td>\n",
|
||||
" <td>Text embeddings, Unimodal (text), English, 512...</td>\n",
|
||||
" <td>mit</td>\n",
|
||||
" <td>1.200</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>24</th>\n",
|
||||
" <td>intfloat/multilingual-e5-large</td>\n",
|
||||
" <td>1024</td>\n",
|
||||
" <td>Multilingual model, e5-large. Recommend using ...</td>\n",
|
||||
" <td>Text embeddings, Unimodal (text), Multilingual...</td>\n",
|
||||
" <td>mit</td>\n",
|
||||
" <td>2.240</td>\n",
|
||||
" </tr>\n",
|
||||
" </tbody>\n",
|
||||
"</table>\n",
|
||||
"</div>"
|
||||
],
|
||||
"text/plain": [
|
||||
" model dim \\\n",
|
||||
"0 BAAI/bge-small-en-v1.5 384 \n",
|
||||
"1 BAAI/bge-small-zh-v1.5 512 \n",
|
||||
"2 sentence-transformers/all-MiniLM-L6-v2 384 \n",
|
||||
"3 snowflake/snowflake-arctic-embed-xs 384 \n",
|
||||
"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 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",
|
||||
"1 Fast and recommended Chinese model 0.090 \n",
|
||||
"2 Sentence Transformer model, MiniLM-L6-v2 0.090 \n",
|
||||
"3 Based on all-MiniLM-L6-v2 model with only 22m ... 0.090 \n",
|
||||
"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 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 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": 6,
|
||||
"execution_count": 12,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"supported_models = (\n",
|
||||
" pd.DataFrame(TextEmbedding.list_supported_models())\n",
|
||||
" .sort_values(\"size_in_GB\")\n",
|
||||
" .drop(columns=[\"sources\", \"model_file\", \"additional_files\"])\n",
|
||||
" .reset_index(drop=True)\n",
|
||||
")\n",
|
||||
"supported_models"
|
||||
]
|
||||
"execution_count": 12
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
@@ -331,16 +376,42 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 8,
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-05-31T18:13:27.124747Z",
|
||||
"start_time": "2024-05-31T18:13:27.096212Z"
|
||||
"end_time": "2024-11-13T09:01:07.038954Z",
|
||||
"start_time": "2024-11-13T09:01:07.019656Z"
|
||||
}
|
||||
},
|
||||
"source": [
|
||||
"(\n",
|
||||
" pd.DataFrame(SparseTextEmbedding.list_supported_models())\n",
|
||||
" .sort_values(\"size_in_GB\")\n",
|
||||
" .drop(columns=[\"sources\", \"model_file\", \"additional_files\"])\n",
|
||||
" .reset_index(drop=True)\n",
|
||||
")"
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
" model vocab_size \\\n",
|
||||
"0 Qdrant/bm25 NaN \n",
|
||||
"1 Qdrant/bm42-all-minilm-l6-v2-attentions 30522.0 \n",
|
||||
"2 prithivida/Splade_PP_en_v1 30522.0 \n",
|
||||
"3 prithvida/Splade_PP_en_v1 30522.0 \n",
|
||||
"\n",
|
||||
" description license size_in_GB \\\n",
|
||||
"0 BM25 as sparse embeddings meant to be used wit... apache-2.0 0.010 \n",
|
||||
"1 Light sparse embedding model, which assigns an... apache-2.0 0.090 \n",
|
||||
"2 Independent Implementation of SPLADE++ Model f... apache-2.0 0.532 \n",
|
||||
"3 Independent Implementation of SPLADE++ Model f... apache-2.0 0.532 \n",
|
||||
"\n",
|
||||
" requires_idf \n",
|
||||
"0 True \n",
|
||||
"1 True \n",
|
||||
"2 NaN \n",
|
||||
"3 NaN "
|
||||
],
|
||||
"text/html": [
|
||||
"<div>\n",
|
||||
"<style scoped>\n",
|
||||
@@ -363,60 +434,59 @@
|
||||
" <th>model</th>\n",
|
||||
" <th>vocab_size</th>\n",
|
||||
" <th>description</th>\n",
|
||||
" <th>license</th>\n",
|
||||
" <th>size_in_GB</th>\n",
|
||||
" <th>requires_idf</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",
|
||||
" <td>Qdrant/bm25</td>\n",
|
||||
" <td>NaN</td>\n",
|
||||
" <td>BM25 as sparse embeddings meant to be used wit...</td>\n",
|
||||
" <td>apache-2.0</td>\n",
|
||||
" <td>0.010</td>\n",
|
||||
" <td>True</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>Qdrant/bm42-all-minilm-l6-v2-attentions</td>\n",
|
||||
" <td>30522.0</td>\n",
|
||||
" <td>Light sparse embedding model, which assigns an...</td>\n",
|
||||
" <td>apache-2.0</td>\n",
|
||||
" <td>0.090</td>\n",
|
||||
" <td>True</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>2</th>\n",
|
||||
" <td>prithivida/Splade_PP_en_v1</td>\n",
|
||||
" <td>30522</td>\n",
|
||||
" <td>30522.0</td>\n",
|
||||
" <td>Independent Implementation of SPLADE++ Model f...</td>\n",
|
||||
" <td>apache-2.0</td>\n",
|
||||
" <td>0.532</td>\n",
|
||||
" <td>NaN</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>3</th>\n",
|
||||
" <td>prithvida/Splade_PP_en_v1</td>\n",
|
||||
" <td>30522.0</td>\n",
|
||||
" <td>Independent Implementation of SPLADE++ Model f...</td>\n",
|
||||
" <td>apache-2.0</td>\n",
|
||||
" <td>0.532</td>\n",
|
||||
" <td>NaN</td>\n",
|
||||
" </tr>\n",
|
||||
" </tbody>\n",
|
||||
"</table>\n",
|
||||
"</div>"
|
||||
],
|
||||
"text/plain": [
|
||||
" 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 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": 8,
|
||||
"execution_count": 13,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"(\n",
|
||||
" pd.DataFrame(SparseTextEmbedding.list_supported_models())\n",
|
||||
" .sort_values(\"size_in_GB\")\n",
|
||||
" .drop(columns=[\"sources\", \"model_file\", \"additional_files\"])\n",
|
||||
" .reset_index(drop=True)\n",
|
||||
")"
|
||||
]
|
||||
"execution_count": 13
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
@@ -429,17 +499,40 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 10,
|
||||
"metadata": {
|
||||
"collapsed": false,
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-05-31T18:14:34.370252Z",
|
||||
"start_time": "2024-05-31T18:14:34.354270Z"
|
||||
},
|
||||
"collapsed": false
|
||||
"end_time": "2024-11-13T09:01:08.074442Z",
|
||||
"start_time": "2024-11-13T09:01:08.056138Z"
|
||||
}
|
||||
},
|
||||
"source": [
|
||||
"(\n",
|
||||
" pd.DataFrame(LateInteractionTextEmbedding.list_supported_models())\n",
|
||||
" .sort_values(\"size_in_GB\")\n",
|
||||
" .drop(columns=[\"sources\", \"model_file\"])\n",
|
||||
" .reset_index(drop=True)\n",
|
||||
")"
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
" model dim \\\n",
|
||||
"0 answerdotai/answerai-colbert-small-v1 96 \n",
|
||||
"1 colbert-ir/colbertv2.0 128 \n",
|
||||
"2 jinaai/jina-colbert-v2 128 \n",
|
||||
"\n",
|
||||
" description license \\\n",
|
||||
"0 Text embeddings, Unimodal (text), Multilingual... apache-2.0 \n",
|
||||
"1 Late interaction model mit \n",
|
||||
"2 New model that expands capabilities of colbert... cc-by-nc-4.0 \n",
|
||||
"\n",
|
||||
" size_in_GB additional_files \n",
|
||||
"0 0.13 NaN \n",
|
||||
"1 0.44 NaN \n",
|
||||
"2 2.24 [onnx/model.onnx_data] "
|
||||
],
|
||||
"text/html": [
|
||||
"<div>\n",
|
||||
"<style scoped>\n",
|
||||
@@ -462,39 +555,50 @@
|
||||
" <th>model</th>\n",
|
||||
" <th>dim</th>\n",
|
||||
" <th>description</th>\n",
|
||||
" <th>license</th>\n",
|
||||
" <th>size_in_GB</th>\n",
|
||||
" <th>additional_files</th>\n",
|
||||
" </tr>\n",
|
||||
" </thead>\n",
|
||||
" <tbody>\n",
|
||||
" <tr>\n",
|
||||
" <th>0</th>\n",
|
||||
" <td>answerdotai/answerai-colbert-small-v1</td>\n",
|
||||
" <td>96</td>\n",
|
||||
" <td>Text embeddings, Unimodal (text), Multilingual...</td>\n",
|
||||
" <td>apache-2.0</td>\n",
|
||||
" <td>0.13</td>\n",
|
||||
" <td>NaN</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>1</th>\n",
|
||||
" <td>colbert-ir/colbertv2.0</td>\n",
|
||||
" <td>128</td>\n",
|
||||
" <td>Late interaction model</td>\n",
|
||||
" <td>mit</td>\n",
|
||||
" <td>0.44</td>\n",
|
||||
" <td>NaN</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>2</th>\n",
|
||||
" <td>jinaai/jina-colbert-v2</td>\n",
|
||||
" <td>128</td>\n",
|
||||
" <td>New model that expands capabilities of colbert...</td>\n",
|
||||
" <td>cc-by-nc-4.0</td>\n",
|
||||
" <td>2.24</td>\n",
|
||||
" <td>[onnx/model.onnx_data]</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,
|
||||
"execution_count": 14,
|
||||
"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",
|
||||
")"
|
||||
]
|
||||
"execution_count": 14
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
@@ -507,17 +611,37 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 12,
|
||||
"metadata": {
|
||||
"collapsed": false,
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-05-31T18:14:42.501881Z",
|
||||
"start_time": "2024-05-31T18:14:42.484726Z"
|
||||
},
|
||||
"collapsed": false
|
||||
"end_time": "2024-11-13T09:01:09.171647Z",
|
||||
"start_time": "2024-11-13T09:01:09.150940Z"
|
||||
}
|
||||
},
|
||||
"source": [
|
||||
"(\n",
|
||||
" pd.DataFrame(ImageEmbedding.list_supported_models())\n",
|
||||
" .sort_values(\"size_in_GB\")\n",
|
||||
" .drop(columns=[\"sources\", \"model_file\"])\n",
|
||||
" .reset_index(drop=True)\n",
|
||||
")"
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
" model dim \\\n",
|
||||
"0 Qdrant/resnet50-onnx 2048 \n",
|
||||
"1 Qdrant/clip-ViT-B-32-vision 512 \n",
|
||||
"2 Qdrant/Unicom-ViT-B-32 512 \n",
|
||||
"3 Qdrant/Unicom-ViT-B-16 768 \n",
|
||||
"\n",
|
||||
" description license size_in_GB \n",
|
||||
"0 Image embeddings, Unimodal (image), 2016 year apache-2.0 0.10 \n",
|
||||
"1 Image embeddings, Multimodal (text&image), 202... mit 0.34 \n",
|
||||
"2 Image embeddings, Multimodal (text&image), 202... apache-2.0 0.48 \n",
|
||||
"3 Image embeddings (more detailed than Unicom-Vi... apache-2.0 0.82 "
|
||||
],
|
||||
"text/html": [
|
||||
"<div>\n",
|
||||
"<style scoped>\n",
|
||||
@@ -540,6 +664,7 @@
|
||||
" <th>model</th>\n",
|
||||
" <th>dim</th>\n",
|
||||
" <th>description</th>\n",
|
||||
" <th>license</th>\n",
|
||||
" <th>size_in_GB</th>\n",
|
||||
" </tr>\n",
|
||||
" </thead>\n",
|
||||
@@ -548,42 +673,175 @@
|
||||
" <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>Image embeddings, Unimodal (image), 2016 year</td>\n",
|
||||
" <td>apache-2.0</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>Image embeddings, Multimodal (text&image), 202...</td>\n",
|
||||
" <td>mit</td>\n",
|
||||
" <td>0.34</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>2</th>\n",
|
||||
" <td>Qdrant/Unicom-ViT-B-32</td>\n",
|
||||
" <td>512</td>\n",
|
||||
" <td>Image embeddings, Multimodal (text&image), 202...</td>\n",
|
||||
" <td>apache-2.0</td>\n",
|
||||
" <td>0.48</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>3</th>\n",
|
||||
" <td>Qdrant/Unicom-ViT-B-16</td>\n",
|
||||
" <td>768</td>\n",
|
||||
" <td>Image embeddings (more detailed than Unicom-Vi...</td>\n",
|
||||
" <td>apache-2.0</td>\n",
|
||||
" <td>0.82</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,
|
||||
"execution_count": 15,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"execution_count": 15
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Supported Rerank Cross Encoder Models"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"metadata": {
|
||||
"ExecuteTime": {
|
||||
"end_time": "2024-11-13T09:01:10.313943Z",
|
||||
"start_time": "2024-11-13T09:01:10.298428Z"
|
||||
}
|
||||
},
|
||||
"source": [
|
||||
"(\n",
|
||||
" pd.DataFrame(ImageEmbedding.list_supported_models()).sort_values(\"size_in_GB\")\n",
|
||||
" pd.DataFrame(TextCrossEncoder.list_supported_models())\n",
|
||||
" .sort_values(\"size_in_GB\")\n",
|
||||
" .drop(columns=[\"sources\", \"model_file\"])\n",
|
||||
" .reset_index(drop=True)\n",
|
||||
")"
|
||||
]
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
" model size_in_GB \\\n",
|
||||
"0 Xenova/ms-marco-MiniLM-L-6-v2 0.08 \n",
|
||||
"1 Xenova/ms-marco-MiniLM-L-12-v2 0.12 \n",
|
||||
"2 jinaai/jina-reranker-v1-tiny-en 0.13 \n",
|
||||
"3 jinaai/jina-reranker-v1-turbo-en 0.15 \n",
|
||||
"4 BAAI/bge-reranker-base 1.04 \n",
|
||||
"5 jinaai/jina-reranker-v2-base-multilingual 1.11 \n",
|
||||
"\n",
|
||||
" description license \n",
|
||||
"0 MiniLM-L-6-v2 model optimized for re-ranking t... apache-2.0 \n",
|
||||
"1 MiniLM-L-12-v2 model optimized for re-ranking ... apache-2.0 \n",
|
||||
"2 Designed for blazing-fast re-ranking with 8K c... apache-2.0 \n",
|
||||
"3 Designed for blazing-fast re-ranking with 8K c... apache-2.0 \n",
|
||||
"4 BGE reranker base model for cross-encoder re-r... mit \n",
|
||||
"5 A multi-lingual reranker model for cross-encod... cc-by-nc-4.0 "
|
||||
],
|
||||
"text/html": [
|
||||
"<div>\n",
|
||||
"<style scoped>\n",
|
||||
" .dataframe tbody tr th:only-of-type {\n",
|
||||
" vertical-align: middle;\n",
|
||||
" }\n",
|
||||
"\n",
|
||||
" .dataframe tbody tr th {\n",
|
||||
" vertical-align: top;\n",
|
||||
" }\n",
|
||||
"\n",
|
||||
" .dataframe thead th {\n",
|
||||
" text-align: right;\n",
|
||||
" }\n",
|
||||
"</style>\n",
|
||||
"<table border=\"1\" class=\"dataframe\">\n",
|
||||
" <thead>\n",
|
||||
" <tr style=\"text-align: right;\">\n",
|
||||
" <th></th>\n",
|
||||
" <th>model</th>\n",
|
||||
" <th>size_in_GB</th>\n",
|
||||
" <th>description</th>\n",
|
||||
" <th>license</th>\n",
|
||||
" </tr>\n",
|
||||
" </thead>\n",
|
||||
" <tbody>\n",
|
||||
" <tr>\n",
|
||||
" <th>0</th>\n",
|
||||
" <td>Xenova/ms-marco-MiniLM-L-6-v2</td>\n",
|
||||
" <td>0.08</td>\n",
|
||||
" <td>MiniLM-L-6-v2 model optimized for re-ranking t...</td>\n",
|
||||
" <td>apache-2.0</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>1</th>\n",
|
||||
" <td>Xenova/ms-marco-MiniLM-L-12-v2</td>\n",
|
||||
" <td>0.12</td>\n",
|
||||
" <td>MiniLM-L-12-v2 model optimized for re-ranking ...</td>\n",
|
||||
" <td>apache-2.0</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>2</th>\n",
|
||||
" <td>jinaai/jina-reranker-v1-tiny-en</td>\n",
|
||||
" <td>0.13</td>\n",
|
||||
" <td>Designed for blazing-fast re-ranking with 8K c...</td>\n",
|
||||
" <td>apache-2.0</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>3</th>\n",
|
||||
" <td>jinaai/jina-reranker-v1-turbo-en</td>\n",
|
||||
" <td>0.15</td>\n",
|
||||
" <td>Designed for blazing-fast re-ranking with 8K c...</td>\n",
|
||||
" <td>apache-2.0</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>4</th>\n",
|
||||
" <td>BAAI/bge-reranker-base</td>\n",
|
||||
" <td>1.04</td>\n",
|
||||
" <td>BGE reranker base model for cross-encoder re-r...</td>\n",
|
||||
" <td>mit</td>\n",
|
||||
" </tr>\n",
|
||||
" <tr>\n",
|
||||
" <th>5</th>\n",
|
||||
" <td>jinaai/jina-reranker-v2-base-multilingual</td>\n",
|
||||
" <td>1.11</td>\n",
|
||||
" <td>A multi-lingual reranker model for cross-encod...</td>\n",
|
||||
" <td>cc-by-nc-4.0</td>\n",
|
||||
" </tr>\n",
|
||||
" </tbody>\n",
|
||||
"</table>\n",
|
||||
"</div>"
|
||||
]
|
||||
},
|
||||
"execution_count": 16,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"execution_count": 16
|
||||
},
|
||||
{
|
||||
"metadata": {},
|
||||
"cell_type": "code",
|
||||
"outputs": [],
|
||||
"execution_count": null,
|
||||
"source": ""
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
@@ -602,7 +860,7 @@
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.11.4"
|
||||
"version": "3.11.8"
|
||||
},
|
||||
"orig_nbformat": 4,
|
||||
"vscode": {
|
||||
|
||||
+2
-2
@@ -26,14 +26,14 @@ pip install fastembed
|
||||
```python
|
||||
from fastembed import TextEmbedding
|
||||
|
||||
documents: List[str] = [
|
||||
documents: list[str] = [
|
||||
"passage: Hello, World!",
|
||||
"query: Hello, World!",
|
||||
"passage: This is an example passage.",
|
||||
"fastembed is supported by and maintained by Qdrant."
|
||||
]
|
||||
embedding_model = TextEmbedding()
|
||||
embeddings: List[np.ndarray] = embedding_model.embed(documents)
|
||||
embeddings: list[np.ndarray] = embedding_model.embed(documents)
|
||||
```
|
||||
|
||||
## Usage with Qdrant
|
||||
|
||||
@@ -41,7 +41,6 @@
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from typing import List\n",
|
||||
"import numpy as np\n",
|
||||
"from fastembed import TextEmbedding"
|
||||
]
|
||||
@@ -71,7 +70,7 @@
|
||||
],
|
||||
"source": [
|
||||
"# Example list of documents\n",
|
||||
"documents: List[str] = [\n",
|
||||
"documents: list[str] = [\n",
|
||||
" \"Maharana Pratap was a Rajput warrior king from Mewar\",\n",
|
||||
" \"He fought against the Mughal Empire led by Akbar\",\n",
|
||||
" \"The Battle of Haldighati in 1576 was his most famous battle\",\n",
|
||||
@@ -87,7 +86,7 @@
|
||||
"embedding_model = TextEmbedding(model_name=\"BAAI/bge-small-en\")\n",
|
||||
"\n",
|
||||
"# We'll use the passage_embed method to get the embeddings for the documents\n",
|
||||
"embeddings: List[np.ndarray] = list(\n",
|
||||
"embeddings: list[np.ndarray] = list(\n",
|
||||
" embedding_model.passage_embed(documents)\n",
|
||||
") # notice that we are casting the generator to a list\n",
|
||||
"\n",
|
||||
|
||||
@@ -46,7 +46,6 @@
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from typing import List\n",
|
||||
"from qdrant_client import QdrantClient"
|
||||
]
|
||||
},
|
||||
@@ -67,7 +66,7 @@
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Example list of documents\n",
|
||||
"documents: List[str] = [\n",
|
||||
"documents: list[str] = [\n",
|
||||
" \"Maharana Pratap was a Rajput warrior king from Mewar\",\n",
|
||||
" \"He fought against the Mughal Empire led by Akbar\",\n",
|
||||
" \"The Battle of Haldighati in 1576 was his most famous battle\",\n",
|
||||
@@ -199,7 +198,9 @@
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"search_result = client.query(collection_name=\"demo_collection\", query_text=\"This is a query document\")\n",
|
||||
"search_result = client.query(\n",
|
||||
" collection_name=\"demo_collection\", query_text=\"This is a query document\"\n",
|
||||
")\n",
|
||||
"print(search_result)"
|
||||
]
|
||||
},
|
||||
|
||||
@@ -19,7 +19,7 @@
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from pathlib import Path\n",
|
||||
"from typing import List, Tuple, Any\n",
|
||||
"from typing import Any\n",
|
||||
"\n",
|
||||
"import numpy as np\n",
|
||||
"import time\n",
|
||||
@@ -91,9 +91,11 @@
|
||||
" return last_hidden.sum(dim=1) / attention_mask.sum(dim=1)[..., None]\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"def hf_embed(model_id: str, inputs: List[str]):\n",
|
||||
"def hf_embed(model_id: str, inputs: list[str]):\n",
|
||||
" # Tokenize the input texts\n",
|
||||
" batch_dict = hf_tokenizer(inputs, max_length=512, padding=True, truncation=True, return_tensors=\"pt\")\n",
|
||||
" batch_dict = hf_tokenizer(\n",
|
||||
" inputs, max_length=512, padding=True, truncation=True, return_tensors=\"pt\"\n",
|
||||
" )\n",
|
||||
"\n",
|
||||
" outputs = hf_model(**batch_dict)\n",
|
||||
" embeddings = average_pool(outputs.last_hidden_state, batch_dict[\"attention_mask\"])\n",
|
||||
@@ -133,7 +135,9 @@
|
||||
"optimization_config = AutoOptimizationConfig.O4()\n",
|
||||
"optimizer = ORTOptimizer.from_pretrained(model)\n",
|
||||
"\n",
|
||||
"optimizer.optimize(save_dir=save_dir, optimization_config=optimization_config, use_external_data_format=True)\n",
|
||||
"optimizer.optimize(\n",
|
||||
" save_dir=save_dir, optimization_config=optimization_config, use_external_data_format=True\n",
|
||||
")\n",
|
||||
"model = ORTModelForFeatureExtraction.from_pretrained(save_dir)\n",
|
||||
"\n",
|
||||
"tokenizer.save_pretrained(save_dir)\n",
|
||||
@@ -171,7 +175,9 @@
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def measure_pipeline_time(pipeline, input_texts: List[str], num_runs=10, **kwargs: Any) -> Tuple[float, float]:\n",
|
||||
"def measure_pipeline_time(\n",
|
||||
" pipeline, input_texts: list[str], num_runs=10, **kwargs: Any\n",
|
||||
") -> tuple[float, float]:\n",
|
||||
" \"\"\"Measures the time it takes to run the pipeline on the input texts.\"\"\"\n",
|
||||
" times = []\n",
|
||||
" total_chars = sum(len(text) for text in input_texts)\n",
|
||||
|
||||
File diff suppressed because one or more lines are too long
@@ -2,6 +2,7 @@ import importlib.metadata
|
||||
|
||||
from fastembed.image import ImageEmbedding
|
||||
from fastembed.late_interaction import LateInteractionTextEmbedding
|
||||
from fastembed.late_interaction_multimodal import LateInteractionMultimodalEmbedding
|
||||
from fastembed.sparse import SparseEmbedding, SparseTextEmbedding
|
||||
from fastembed.text import TextEmbedding
|
||||
|
||||
@@ -17,4 +18,5 @@ __all__ = [
|
||||
"SparseEmbedding",
|
||||
"ImageEmbedding",
|
||||
"LateInteractionTextEmbedding",
|
||||
"LateInteractionMultimodalEmbedding",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,40 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional, Any
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ModelSource:
|
||||
hf: Optional[str] = None
|
||||
url: Optional[str] = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if self.hf is None and self.url is None:
|
||||
raise ValueError(
|
||||
f"At least one source should be set, current sources: hf={self.hf}, url={self.url}"
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class BaseModelDescription:
|
||||
model: str
|
||||
sources: ModelSource
|
||||
model_file: str
|
||||
description: str
|
||||
license: str
|
||||
size_in_GB: float
|
||||
additional_files: list[str] = field(default_factory=list)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class DenseModelDescription(BaseModelDescription):
|
||||
dim: Optional[int] = None
|
||||
tasks: Optional[dict[str, Any]] = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
assert self.dim is not None, "dim is required for dense model description"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SparseModelDescription(BaseModelDescription):
|
||||
requires_idf: Optional[bool] = None
|
||||
vocab_size: Optional[int] = None
|
||||
@@ -1,28 +1,44 @@
|
||||
import os
|
||||
import time
|
||||
import json
|
||||
import shutil
|
||||
import tarfile
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional
|
||||
from typing import Any, Optional, Union, TypeVar, Generic
|
||||
|
||||
import requests
|
||||
from huggingface_hub import snapshot_download
|
||||
from huggingface_hub.utils import RepositoryNotFoundError
|
||||
from huggingface_hub import snapshot_download, model_info, list_repo_tree
|
||||
from huggingface_hub.hf_api import RepoFile
|
||||
from huggingface_hub.utils import (
|
||||
RepositoryNotFoundError,
|
||||
disable_progress_bars,
|
||||
enable_progress_bars,
|
||||
)
|
||||
from loguru import logger
|
||||
from tqdm import tqdm
|
||||
from fastembed.common.model_description import BaseModelDescription
|
||||
|
||||
T = TypeVar("T", bound=BaseModelDescription)
|
||||
|
||||
|
||||
class ModelManagement:
|
||||
class ModelManagement(Generic[T]):
|
||||
METADATA_FILE = "files_metadata.json"
|
||||
|
||||
@classmethod
|
||||
def list_supported_models(cls) -> List[Dict[str, Any]]:
|
||||
def list_supported_models(cls) -> list[dict[str, Any]]:
|
||||
"""Lists the supported models.
|
||||
|
||||
Returns:
|
||||
List[Dict[str, Any]]: A list of dictionaries containing the model information.
|
||||
list[T]: A list of dictionaries containing the model information.
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
@classmethod
|
||||
def _get_model_description(cls, model_name: str) -> Dict[str, Any]:
|
||||
def _list_supported_models(cls) -> list[T]:
|
||||
raise NotImplementedError()
|
||||
|
||||
@classmethod
|
||||
def _get_model_description(cls, model_name: str) -> T:
|
||||
"""
|
||||
Gets the model description from the model_name.
|
||||
|
||||
@@ -33,18 +49,16 @@ class ModelManagement:
|
||||
ValueError: If the model_name is not supported.
|
||||
|
||||
Returns:
|
||||
Dict[str, Any]: The model description.
|
||||
T: The model description.
|
||||
"""
|
||||
for model in cls.list_supported_models():
|
||||
if model_name.lower() == model["model"].lower():
|
||||
for model in cls._list_supported_models():
|
||||
if model_name.lower() == model.model.lower():
|
||||
return model
|
||||
|
||||
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.
|
||||
|
||||
@@ -73,11 +87,9 @@ 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
|
||||
show_progress = bool(total_size_in_bytes and show_progress)
|
||||
|
||||
with tqdm(
|
||||
total=total_size_in_bytes,
|
||||
@@ -96,20 +108,75 @@ class ModelManagement:
|
||||
def download_files_from_huggingface(
|
||||
cls,
|
||||
hf_source_repo: str,
|
||||
cache_dir: Optional[str] = None,
|
||||
extra_patterns: Optional[List[str]] = None,
|
||||
**kwargs,
|
||||
cache_dir: str,
|
||||
extra_patterns: list[str],
|
||||
local_files_only: bool = False,
|
||||
**kwargs: Any,
|
||||
) -> str:
|
||||
"""
|
||||
Downloads a model from HuggingFace Hub.
|
||||
Args:
|
||||
hf_source_repo (str): Name of the model on HuggingFace Hub, e.g. "qdrant/all-MiniLM-L6-v2-onnx".
|
||||
cache_dir (Optional[str]): The path to the cache directory.
|
||||
extra_patterns (Optional[List[str]]): extra patterns to allow in the snapshot download, typically
|
||||
extra_patterns (list[str]): extra patterns to allow in the snapshot download, typically
|
||||
includes the required model files.
|
||||
local_files_only (bool, optional): Whether to only use local files. Defaults to False.
|
||||
Returns:
|
||||
Path: The path to the model directory.
|
||||
"""
|
||||
|
||||
def _verify_files_from_metadata(
|
||||
model_dir: Path, stored_metadata: dict[str, Any], repo_files: list[RepoFile]
|
||||
) -> bool:
|
||||
try:
|
||||
for rel_path, meta in stored_metadata.items():
|
||||
file_path = model_dir / rel_path
|
||||
|
||||
if not file_path.exists():
|
||||
return False
|
||||
|
||||
if repo_files: # online verification
|
||||
file_info = next((f for f in repo_files if f.path == file_path.name), None)
|
||||
if (
|
||||
not file_info
|
||||
or file_info.size != meta["size"]
|
||||
or file_info.blob_id != meta["blob_id"]
|
||||
):
|
||||
return False
|
||||
|
||||
else: # offline verification
|
||||
if file_path.stat().st_size != meta["size"]:
|
||||
return False
|
||||
return True
|
||||
except (OSError, KeyError) as e:
|
||||
logger.error(f"Error verifying files: {str(e)}")
|
||||
return False
|
||||
|
||||
def _collect_file_metadata(
|
||||
model_dir: Path, repo_files: list[RepoFile]
|
||||
) -> dict[str, dict[str, Union[int, str]]]:
|
||||
meta: dict[str, dict[str, Union[int, str]]] = {}
|
||||
file_info_map = {f.path: f for f in repo_files}
|
||||
for file_path in model_dir.rglob("*"):
|
||||
if file_path.is_file() and file_path.name != cls.METADATA_FILE:
|
||||
repo_file = file_info_map.get(file_path.name)
|
||||
if repo_file:
|
||||
meta[str(file_path.relative_to(model_dir))] = {
|
||||
"size": repo_file.size,
|
||||
"blob_id": repo_file.blob_id,
|
||||
}
|
||||
return meta
|
||||
|
||||
def _save_file_metadata(
|
||||
model_dir: Path, meta: dict[str, dict[str, Union[int, str]]]
|
||||
) -> None:
|
||||
try:
|
||||
if not model_dir.exists():
|
||||
model_dir.mkdir(parents=True, exist_ok=True)
|
||||
(model_dir / cls.METADATA_FILE).write_text(json.dumps(meta))
|
||||
except (OSError, ValueError) as e:
|
||||
logger.warning(f"Error saving metadata: {str(e)}")
|
||||
|
||||
allow_patterns = [
|
||||
"config.json",
|
||||
"tokenizer.json",
|
||||
@@ -117,18 +184,86 @@ class ModelManagement:
|
||||
"special_tokens_map.json",
|
||||
"preprocessor_config.json",
|
||||
]
|
||||
if extra_patterns is not None:
|
||||
allow_patterns.extend(extra_patterns)
|
||||
|
||||
return snapshot_download(
|
||||
allow_patterns.extend(extra_patterns)
|
||||
|
||||
snapshot_dir = Path(cache_dir) / f"models--{hf_source_repo.replace('/', '--')}"
|
||||
metadata_file = snapshot_dir / cls.METADATA_FILE
|
||||
|
||||
if local_files_only:
|
||||
disable_progress_bars()
|
||||
if metadata_file.exists():
|
||||
metadata = json.loads(metadata_file.read_text())
|
||||
verified = _verify_files_from_metadata(snapshot_dir, metadata, repo_files=[])
|
||||
if not verified:
|
||||
logger.warning(
|
||||
"Local file sizes do not match the metadata."
|
||||
) # do not raise, still make an attempt to load the model
|
||||
else:
|
||||
logger.warning(
|
||||
"Metadata file not found. Proceeding without checking local files."
|
||||
) # if users have downloaded models from hf manually, or they're updating from previous versions of
|
||||
# fastembed
|
||||
result = snapshot_download(
|
||||
repo_id=hf_source_repo,
|
||||
allow_patterns=allow_patterns,
|
||||
cache_dir=cache_dir,
|
||||
local_files_only=local_files_only,
|
||||
**kwargs,
|
||||
)
|
||||
return result
|
||||
|
||||
repo_revision = model_info(hf_source_repo).sha
|
||||
repo_tree = list(list_repo_tree(hf_source_repo, revision=repo_revision, repo_type="model"))
|
||||
|
||||
allowed_extensions = {".json", ".onnx", ".txt"}
|
||||
repo_files = (
|
||||
[
|
||||
f
|
||||
for f in repo_tree
|
||||
if isinstance(f, RepoFile) and Path(f.path).suffix in allowed_extensions
|
||||
]
|
||||
if repo_tree
|
||||
else []
|
||||
)
|
||||
|
||||
verified_metadata = False
|
||||
|
||||
if snapshot_dir.exists() and metadata_file.exists():
|
||||
metadata = json.loads(metadata_file.read_text())
|
||||
verified_metadata = _verify_files_from_metadata(snapshot_dir, metadata, repo_files)
|
||||
|
||||
if verified_metadata:
|
||||
disable_progress_bars()
|
||||
|
||||
result = snapshot_download(
|
||||
repo_id=hf_source_repo,
|
||||
allow_patterns=allow_patterns,
|
||||
cache_dir=cache_dir,
|
||||
local_files_only=kwargs.get("local_files_only", False),
|
||||
local_files_only=local_files_only,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
if (
|
||||
not verified_metadata
|
||||
): # metadata is not up-to-date, update it and check whether the files have been
|
||||
# downloaded correctly
|
||||
metadata = _collect_file_metadata(snapshot_dir, repo_files)
|
||||
|
||||
download_successful = _verify_files_from_metadata(
|
||||
snapshot_dir, metadata, repo_files=[]
|
||||
) # offline verification
|
||||
if not download_successful:
|
||||
raise ValueError(
|
||||
"Files have been corrupted during downloading process. "
|
||||
"Please check your internet connection and try again."
|
||||
)
|
||||
_save_file_metadata(snapshot_dir, metadata)
|
||||
|
||||
return result
|
||||
|
||||
@classmethod
|
||||
def decompress_to_cache(cls, targz_path: str, cache_dir: str):
|
||||
def decompress_to_cache(cls, targz_path: str, cache_dir: str) -> str:
|
||||
"""
|
||||
Decompresses a .tar.gz file to a cache directory.
|
||||
|
||||
@@ -151,7 +286,9 @@ class ModelManagement:
|
||||
# Open the tar.gz file
|
||||
with tarfile.open(targz_path, "r:gz") as tar:
|
||||
# Extract all files into the cache directory
|
||||
tar.extractall(path=cache_dir)
|
||||
tar.extractall(
|
||||
path=cache_dir,
|
||||
)
|
||||
except tarfile.TarError as e:
|
||||
# If any error occurs while opening or extracting the tar.gz file,
|
||||
# delete the cache directory (if it was created in this function)
|
||||
@@ -164,10 +301,13 @@ class ModelManagement:
|
||||
|
||||
@classmethod
|
||||
def retrieve_model_gcs(
|
||||
cls, model_name: str, source_url: str, cache_dir: str
|
||||
cls,
|
||||
model_name: str,
|
||||
source_url: str,
|
||||
cache_dir: str,
|
||||
local_files_only: bool = False,
|
||||
) -> Path:
|
||||
fast_model_name = f"fast-{model_name.split('/')[-1]}"
|
||||
|
||||
cache_tmp_dir = Path(cache_dir) / "tmp"
|
||||
model_tmp_dir = cache_tmp_dir / fast_model_name
|
||||
model_dir = Path(cache_dir) / fast_model_name
|
||||
@@ -186,31 +326,35 @@ class ModelManagement:
|
||||
if model_tar_gz.exists():
|
||||
model_tar_gz.unlink()
|
||||
|
||||
cls.download_file_from_gcs(
|
||||
source_url,
|
||||
output_path=str(model_tar_gz),
|
||||
)
|
||||
if not local_files_only:
|
||||
cls.download_file_from_gcs(
|
||||
source_url,
|
||||
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
|
||||
model_tmp_dir.rename(model_dir)
|
||||
model_tar_gz.unlink()
|
||||
# Rename from tmp to final name is atomic
|
||||
model_tmp_dir.rename(model_dir)
|
||||
else:
|
||||
logger.error(
|
||||
f"Could not find the model tar.gz file at {model_dir} and local_files_only=True."
|
||||
)
|
||||
raise ValueError(
|
||||
f"Could not find the model tar.gz file at {model_dir} and local_files_only=True."
|
||||
)
|
||||
|
||||
return model_dir
|
||||
|
||||
@classmethod
|
||||
def download_model(cls, model: Dict[str, Any], cache_dir: Path, **kwargs) -> Path:
|
||||
def download_model(cls, model: T, cache_dir: str, retries: int = 3, **kwargs: Any) -> Path:
|
||||
"""
|
||||
Downloads a model from HuggingFace Hub or Google Cloud Storage.
|
||||
|
||||
Args:
|
||||
model (Dict[str, Any]): The model description.
|
||||
model (T): The model description.
|
||||
Example:
|
||||
```
|
||||
{
|
||||
@@ -225,34 +369,63 @@ class ModelManagement:
|
||||
}
|
||||
```
|
||||
cache_dir (str): The path to the cache directory.
|
||||
retries: (int): The number of times to retry (including the first attempt)
|
||||
|
||||
Returns:
|
||||
Path: The path to the downloaded model directory.
|
||||
"""
|
||||
local_files_only = kwargs.get("local_files_only", False)
|
||||
specific_model_path: Optional[str] = kwargs.pop("specific_model_path", None)
|
||||
if specific_model_path:
|
||||
return Path(specific_model_path)
|
||||
retries = 1 if local_files_only else retries
|
||||
hf_source = model.sources.hf
|
||||
url_source = model.sources.url
|
||||
|
||||
hf_source = model.get("sources", {}).get("hf")
|
||||
url_source = model.get("sources", {}).get("url")
|
||||
sleep = 3.0
|
||||
while retries > 0:
|
||||
retries -= 1
|
||||
|
||||
if hf_source:
|
||||
extra_patterns = [model["model_file"]]
|
||||
extra_patterns.extend(model.get("additional_files", []))
|
||||
if hf_source:
|
||||
extra_patterns = [model.model_file]
|
||||
extra_patterns.extend(model.additional_files)
|
||||
|
||||
try:
|
||||
return Path(
|
||||
cls.download_files_from_huggingface(
|
||||
hf_source,
|
||||
cache_dir=str(cache_dir),
|
||||
extra_patterns=extra_patterns,
|
||||
local_files_only=kwargs.get("local_files_only", False),
|
||||
try:
|
||||
return Path(
|
||||
cls.download_files_from_huggingface(
|
||||
hf_source,
|
||||
cache_dir=cache_dir,
|
||||
extra_patterns=extra_patterns,
|
||||
**kwargs,
|
||||
)
|
||||
)
|
||||
)
|
||||
except (EnvironmentError, RepositoryNotFoundError, ValueError) as e:
|
||||
except (EnvironmentError, RepositoryNotFoundError, ValueError) as e:
|
||||
if not local_files_only:
|
||||
logger.error(
|
||||
f"Could not download model from HuggingFace: {e} "
|
||||
"Falling back to other sources."
|
||||
)
|
||||
finally:
|
||||
enable_progress_bars()
|
||||
if url_source or local_files_only:
|
||||
try:
|
||||
return cls.retrieve_model_gcs(
|
||||
model.model,
|
||||
str(url_source),
|
||||
str(cache_dir),
|
||||
local_files_only=local_files_only,
|
||||
)
|
||||
except Exception:
|
||||
if not local_files_only:
|
||||
logger.error(f"Could not download model from url: {url_source}")
|
||||
|
||||
if local_files_only:
|
||||
logger.error("Could not find model in cache_dir")
|
||||
else:
|
||||
logger.error(
|
||||
f"Could not download model from HuggingFace: {e}"
|
||||
"Falling back to other sources."
|
||||
f"Could not download model from either source, sleeping for {sleep} seconds, {retries} retries left."
|
||||
)
|
||||
time.sleep(sleep)
|
||||
sleep *= 3
|
||||
|
||||
if url_source:
|
||||
return cls.retrieve_model_gcs(model["model"], url_source, str(cache_dir))
|
||||
|
||||
raise ValueError(f"Could not download model {model['model']} from any source.")
|
||||
raise ValueError(f"Could not load model {model.model} from any source.")
|
||||
|
||||
@@ -1,22 +1,15 @@
|
||||
import warnings
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import (
|
||||
Any,
|
||||
Dict,
|
||||
Generic,
|
||||
Iterable,
|
||||
Optional,
|
||||
Sequence,
|
||||
Tuple,
|
||||
Type,
|
||||
TypeVar,
|
||||
)
|
||||
from typing import Any, Generic, Iterable, Optional, Sequence, Type, TypeVar
|
||||
|
||||
import numpy as np
|
||||
import onnxruntime as ort
|
||||
|
||||
from fastembed.common.types import OnnxProvider
|
||||
from numpy.typing import NDArray
|
||||
from tokenizers import Tokenizer
|
||||
|
||||
from fastembed.common.types import OnnxProvider, NumpyArray
|
||||
from fastembed.parallel_processor import Worker
|
||||
|
||||
# Holds type of the embedding result
|
||||
@@ -25,46 +18,62 @@ T = TypeVar("T")
|
||||
|
||||
@dataclass
|
||||
class OnnxOutputContext:
|
||||
model_output: np.ndarray
|
||||
attention_mask: Optional[np.ndarray] = None
|
||||
input_ids: Optional[np.ndarray] = None
|
||||
model_output: NumpyArray
|
||||
attention_mask: Optional[NDArray[np.int64]] = None
|
||||
input_ids: Optional[NDArray[np.int64]] = None
|
||||
|
||||
|
||||
class OnnxModel(Generic[T]):
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type["EmbeddingWorker"]:
|
||||
def _get_worker_class(cls) -> Type["EmbeddingWorker[T]"]:
|
||||
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:
|
||||
self.model = None
|
||||
self.tokenizer = None
|
||||
self.model: Optional[ort.InferenceSession] = None
|
||||
self.tokenizer: Optional[Tokenizer] = None
|
||||
|
||||
def _preprocess_onnx_input(
|
||||
self, onnx_input: Dict[str, np.ndarray], **kwargs
|
||||
) -> Dict[str, np.ndarray]:
|
||||
self, onnx_input: dict[str, NumpyArray], **kwargs: Any
|
||||
) -> dict[str, NumpyArray]:
|
||||
"""
|
||||
Preprocess the onnx input.
|
||||
"""
|
||||
return onnx_input
|
||||
|
||||
def load_onnx_model(
|
||||
def _load_onnx_model(
|
||||
self,
|
||||
model_dir: Path,
|
||||
model_file: str,
|
||||
threads: Optional[int],
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
cuda: bool = False,
|
||||
device_id: Optional[int] = None,
|
||||
) -> None:
|
||||
model_path = model_dir / model_file
|
||||
# List of Execution Providers: https://onnxruntime.ai/docs/execution-providers
|
||||
|
||||
onnx_providers = (
|
||||
["CPUExecutionProvider"] if providers is None else list(providers)
|
||||
)
|
||||
if cuda and providers is not None:
|
||||
warnings.warn(
|
||||
f"`cuda` and `providers` are mutually exclusive parameters, cuda: {cuda}, providers: {providers}",
|
||||
category=UserWarning,
|
||||
stacklevel=6,
|
||||
)
|
||||
|
||||
if providers is not None:
|
||||
onnx_providers = list(providers)
|
||||
elif cuda:
|
||||
if device_id is None:
|
||||
onnx_providers = ["CUDAExecutionProvider"]
|
||||
else:
|
||||
onnx_providers = [("CUDAExecutionProvider", {"device_id": device_id})]
|
||||
else:
|
||||
onnx_providers = ["CPUExecutionProvider"]
|
||||
|
||||
available_providers = ort.get_available_providers()
|
||||
requested_provider_names = []
|
||||
requested_provider_names: list[str] = []
|
||||
for provider in onnx_providers:
|
||||
# check providers available
|
||||
provider_name = provider if isinstance(provider, str) else provider[0]
|
||||
@@ -85,6 +94,7 @@ class OnnxModel(Generic[T]):
|
||||
str(model_path), providers=onnx_providers, sess_options=so
|
||||
)
|
||||
if "CUDAExecutionProvider" in requested_provider_names:
|
||||
assert self.model is not None
|
||||
current_providers = self.model.get_providers()
|
||||
if "CUDAExecutionProvider" not in current_providers:
|
||||
warnings.warn(
|
||||
@@ -94,30 +104,33 @@ class OnnxModel(Generic[T]):
|
||||
RuntimeWarning,
|
||||
)
|
||||
|
||||
def onnx_embed(self, *args, **kwargs) -> OnnxOutputContext:
|
||||
def load_onnx_model(self) -> None:
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
def onnx_embed(self, *args: Any, **kwargs: Any) -> OnnxOutputContext:
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
|
||||
class EmbeddingWorker(Worker):
|
||||
class EmbeddingWorker(Worker, Generic[T]):
|
||||
def init_embedding(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
**kwargs,
|
||||
) -> OnnxModel:
|
||||
**kwargs: Any,
|
||||
) -> OnnxModel[T]:
|
||||
raise NotImplementedError()
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
**kwargs,
|
||||
**kwargs: Any,
|
||||
):
|
||||
self.model = self.init_embedding(model_name, cache_dir, **kwargs)
|
||||
|
||||
@classmethod
|
||||
def start(cls, model_name: str, cache_dir: str, **kwargs: Any) -> "EmbeddingWorker":
|
||||
def start(cls, model_name: str, cache_dir: str, **kwargs: Any) -> "EmbeddingWorker[T]":
|
||||
return cls(model_name=model_name, cache_dir=cache_dir, **kwargs)
|
||||
|
||||
def process(self, items: Iterable[Tuple[int, Any]]) -> Iterable[Tuple[int, Any]]:
|
||||
def process(self, items: Iterable[tuple[int, Any]]) -> Iterable[tuple[int, Any]]:
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
@@ -1,13 +1,13 @@
|
||||
import json
|
||||
from typing import Any
|
||||
from pathlib import Path
|
||||
from typing import Tuple
|
||||
|
||||
from tokenizers import AddedToken, Tokenizer
|
||||
|
||||
from fastembed.image.transform.operators import Compose
|
||||
|
||||
|
||||
def load_special_tokens(model_dir: Path) -> dict:
|
||||
def load_special_tokens(model_dir: Path) -> dict[str, Any]:
|
||||
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}")
|
||||
@@ -18,7 +18,7 @@ def load_special_tokens(model_dir: Path) -> dict:
|
||||
return tokens_map
|
||||
|
||||
|
||||
def load_tokenizer(model_dir: Path, max_length: int = 512) -> Tuple[Tokenizer, dict]:
|
||||
def load_tokenizer(model_dir: Path) -> tuple[Tokenizer, dict[str, int]]:
|
||||
config_path = model_dir / "config.json"
|
||||
if not config_path.exists():
|
||||
raise ValueError(f"Could not find config.json in {model_dir}")
|
||||
@@ -36,13 +36,20 @@ def load_tokenizer(model_dir: Path, max_length: int = 512) -> Tuple[Tokenizer, d
|
||||
|
||||
with open(str(tokenizer_config_path)) as tokenizer_config_file:
|
||||
tokenizer_config = json.load(tokenizer_config_file)
|
||||
assert (
|
||||
"model_max_length" in tokenizer_config or "max_length" in tokenizer_config
|
||||
), "Models without model_max_length or max_length are not supported."
|
||||
if "model_max_length" not in tokenizer_config:
|
||||
max_context = tokenizer_config["max_length"]
|
||||
elif "max_length" not in tokenizer_config:
|
||||
max_context = tokenizer_config["model_max_length"]
|
||||
else:
|
||||
max_context = min(tokenizer_config["model_max_length"], tokenizer_config["max_length"])
|
||||
|
||||
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=max_context)
|
||||
tokenizer.enable_padding(
|
||||
pad_id=config.get("pad_token_id", 0), pad_token=tokenizer_config["pad_token"]
|
||||
)
|
||||
@@ -53,7 +60,7 @@ def load_tokenizer(model_dir: Path, max_length: int = 512) -> Tuple[Tokenizer, d
|
||||
elif isinstance(token, dict):
|
||||
tokenizer.add_special_tokens([AddedToken(**token)])
|
||||
|
||||
special_token_to_id = {}
|
||||
special_token_to_id: dict[str, int] = {}
|
||||
|
||||
for token in tokens_map.values():
|
||||
if isinstance(token, str):
|
||||
|
||||
@@ -1,6 +1,9 @@
|
||||
import os
|
||||
from pathlib import Path
|
||||
import sys
|
||||
from typing import Any, Dict, Iterable, Tuple, Union
|
||||
from PIL import Image
|
||||
from typing import Any, Union
|
||||
import numpy as np
|
||||
from numpy.typing import NDArray
|
||||
|
||||
if sys.version_info >= (3, 10):
|
||||
from typing import TypeAlias
|
||||
@@ -8,7 +11,14 @@ else:
|
||||
from typing_extensions import TypeAlias
|
||||
|
||||
|
||||
PathInput: TypeAlias = Union[str, os.PathLike]
|
||||
ImageInput: TypeAlias = Union[PathInput, Iterable[PathInput]]
|
||||
PathInput: TypeAlias = Union[str, Path]
|
||||
ImageInput: TypeAlias = Union[PathInput, Image.Image]
|
||||
|
||||
OnnxProvider: TypeAlias = Union[str, Tuple[str, Dict[Any, Any]]]
|
||||
OnnxProvider: TypeAlias = Union[str, tuple[str, dict[Any, Any]]]
|
||||
NumpyArray = Union[
|
||||
NDArray[np.float32],
|
||||
NDArray[np.float16],
|
||||
NDArray[np.int8],
|
||||
NDArray[np.int64],
|
||||
NDArray[np.int32],
|
||||
]
|
||||
|
||||
@@ -1,13 +1,20 @@
|
||||
import os
|
||||
import sys
|
||||
import re
|
||||
import tempfile
|
||||
from itertools import islice
|
||||
import unicodedata
|
||||
from pathlib import Path
|
||||
from typing import Generator, Iterable, Optional, Union
|
||||
from itertools import islice
|
||||
from typing import Iterable, Optional, TypeVar
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastembed.common.types import NumpyArray
|
||||
|
||||
def normalize(input_array, p=2, dim=1, eps=1e-12) -> np.ndarray:
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
def normalize(input_array: NumpyArray, p: int = 2, dim: int = 1, eps: float = 1e-12) -> NumpyArray:
|
||||
# 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
|
||||
@@ -15,7 +22,7 @@ def normalize(input_array, p=2, dim=1, eps=1e-12) -> np.ndarray:
|
||||
return normalized_array
|
||||
|
||||
|
||||
def iter_batch(iterable: Union[Iterable, Generator], size: int) -> Iterable:
|
||||
def iter_batch(iterable: Iterable[T], size: int) -> Iterable[list[T]]:
|
||||
"""
|
||||
>>> list(iter_batch([1,2,3,4,5], 3))
|
||||
[[1, 2, 3], [4, 5]]
|
||||
@@ -37,7 +44,16 @@ def define_cache_dir(cache_dir: Optional[str] = None) -> Path:
|
||||
cache_path = Path(os.getenv("FASTEMBED_CACHE_PATH", default_cache_dir))
|
||||
else:
|
||||
cache_path = Path(cache_dir)
|
||||
|
||||
cache_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
return cache_path
|
||||
|
||||
|
||||
def get_all_punctuation() -> set[str]:
|
||||
return set(
|
||||
chr(i) for i in range(sys.maxunicode) if unicodedata.category(chr(i)).startswith("P")
|
||||
)
|
||||
|
||||
|
||||
def remove_non_alphanumeric(text: str) -> str:
|
||||
return re.sub(r"[^\w\s]", " ", text, flags=re.UNICODE)
|
||||
|
||||
@@ -0,0 +1,16 @@
|
||||
{
|
||||
"models": [
|
||||
{
|
||||
"model": "BAAI/bge-base-en",
|
||||
"dim": 768,
|
||||
"description": "Text embeddings, Unimodal (text), English...",
|
||||
"license": "mit",
|
||||
"size_in_GB": 0.42,
|
||||
"sources": {
|
||||
"hf": "Qdrant/fast-bge-base-en",
|
||||
"url": "https://storage.googleapis.com/qdrant-fastembed/fast-bge-base-en.tar.gz"
|
||||
},
|
||||
"model_file": "model_optimized.onnx"
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -1,4 +1,4 @@
|
||||
from typing import Optional
|
||||
from typing import Optional, Any
|
||||
|
||||
from loguru import logger
|
||||
|
||||
@@ -19,6 +19,6 @@ class JinaEmbedding(TextEmbedding):
|
||||
model_name: str = "jinaai/jina-embeddings-v2-base-en",
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
**kwargs,
|
||||
**kwargs: Any,
|
||||
):
|
||||
super().__init__(model_name, cache_dir, threads, **kwargs)
|
||||
|
||||
@@ -1,22 +1,23 @@
|
||||
from typing import Any, Dict, Iterable, List, Optional, Sequence, Type
|
||||
|
||||
import numpy as np
|
||||
from typing import Any, Iterable, Optional, Sequence, Type, Union
|
||||
from dataclasses import asdict
|
||||
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.common import ImageInput, OnnxProvider
|
||||
from fastembed.image.image_embedding_base import ImageEmbeddingBase
|
||||
from fastembed.image.onnx_embedding import OnnxImageEmbedding
|
||||
from fastembed.common.model_description import DenseModelDescription
|
||||
|
||||
|
||||
class ImageEmbedding(ImageEmbeddingBase):
|
||||
EMBEDDINGS_REGISTRY: List[Type[ImageEmbeddingBase]] = [OnnxImageEmbedding]
|
||||
EMBEDDINGS_REGISTRY: list[Type[ImageEmbeddingBase]] = [OnnxImageEmbedding]
|
||||
|
||||
@classmethod
|
||||
def list_supported_models(cls) -> List[Dict[str, Any]]:
|
||||
def list_supported_models(cls) -> list[dict[str, Any]]:
|
||||
"""
|
||||
Lists the supported models.
|
||||
|
||||
Returns:
|
||||
List[Dict[str, Any]]: A list of dictionaries containing the model information.
|
||||
list[dict[str, Any]]: A list of dictionaries containing the model information.
|
||||
|
||||
Example:
|
||||
```
|
||||
@@ -25,6 +26,7 @@ class ImageEmbedding(ImageEmbeddingBase):
|
||||
"model": "Qdrant/clip-ViT-B-32-vision",
|
||||
"dim": 512,
|
||||
"description": "CLIP vision encoder based on ViT-B/32",
|
||||
"license": "mit",
|
||||
"size_in_GB": 0.33,
|
||||
"sources": {
|
||||
"hf": "Qdrant/clip-ViT-B-32-vision",
|
||||
@@ -34,9 +36,13 @@ class ImageEmbedding(ImageEmbeddingBase):
|
||||
]
|
||||
```
|
||||
"""
|
||||
result = []
|
||||
return [asdict(model) for model in cls._list_supported_models()]
|
||||
|
||||
@classmethod
|
||||
def _list_supported_models(cls) -> list[DenseModelDescription]:
|
||||
result: list[DenseModelDescription] = []
|
||||
for embedding in cls.EMBEDDINGS_REGISTRY:
|
||||
result.extend(embedding.list_supported_models())
|
||||
result.extend(embedding._list_supported_models())
|
||||
return result
|
||||
|
||||
def __init__(
|
||||
@@ -45,40 +51,41 @@ class ImageEmbedding(ImageEmbeddingBase):
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
**kwargs,
|
||||
cuda: bool = False,
|
||||
device_ids: Optional[list[int]] = None,
|
||||
lazy_load: bool = False,
|
||||
**kwargs: Any,
|
||||
):
|
||||
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
|
||||
):
|
||||
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,
|
||||
cuda=cuda,
|
||||
device_ids=device_ids,
|
||||
lazy_load=lazy_load,
|
||||
**kwargs,
|
||||
)
|
||||
return
|
||||
|
||||
raise ValueError(
|
||||
f"Model {model_name} is not supported in TextEmbedding."
|
||||
"Please check the supported models using `TextEmbedding.list_supported_models()`"
|
||||
f"Model {model_name} is not supported in ImageEmbedding."
|
||||
"Please check the supported models using `ImageEmbedding.list_supported_models()`"
|
||||
)
|
||||
|
||||
def embed(
|
||||
self,
|
||||
images: ImageInput,
|
||||
images: Union[ImageInput, Iterable[ImageInput]],
|
||||
batch_size: int = 16,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs,
|
||||
) -> Iterable[np.ndarray]:
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
Encode a list of documents into list of embeddings.
|
||||
We use mean pooling with attention so that the model can handle variable-length inputs.
|
||||
Encode a list of images into list of embeddings.
|
||||
|
||||
Args:
|
||||
images: Iterator of image paths or single image path to embed
|
||||
|
||||
@@ -1,18 +1,18 @@
|
||||
from typing import Iterable, Optional
|
||||
|
||||
import numpy as np
|
||||
from typing import Iterable, Optional, Any, Union
|
||||
|
||||
from fastembed.common.model_description import DenseModelDescription
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.common.model_management import ModelManagement
|
||||
from fastembed.common.types import ImageInput
|
||||
|
||||
|
||||
class ImageEmbeddingBase(ModelManagement):
|
||||
class ImageEmbeddingBase(ModelManagement[DenseModelDescription]):
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
**kwargs,
|
||||
**kwargs: Any,
|
||||
):
|
||||
self.model_name = model_name
|
||||
self.cache_dir = cache_dir
|
||||
@@ -21,16 +21,16 @@ class ImageEmbeddingBase(ModelManagement):
|
||||
|
||||
def embed(
|
||||
self,
|
||||
images: ImageInput,
|
||||
images: Union[ImageInput, Iterable[ImageInput]],
|
||||
batch_size: int = 16,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs,
|
||||
) -> Iterable[np.ndarray]:
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
Embeds a list of images into a list of embeddings.
|
||||
|
||||
Args:
|
||||
images - The list of image paths to preprocess and embed.
|
||||
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.
|
||||
@@ -39,6 +39,6 @@ class ImageEmbeddingBase(ModelManagement):
|
||||
**kwargs: Additional keyword argument to pass to the embed method.
|
||||
|
||||
Yields:
|
||||
Iterable[np.ndarray]: The embeddings.
|
||||
Iterable[NdArray]: The embeddings.
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
@@ -1,45 +1,78 @@
|
||||
from typing import Any, Dict, Iterable, List, Optional, Sequence, Type
|
||||
from typing import Any, Iterable, Optional, Sequence, Type, Union
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastembed.common.types import NumpyArray
|
||||
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",
|
||||
},
|
||||
from fastembed.common.model_description import DenseModelDescription, ModelSource
|
||||
|
||||
supported_onnx_models: list[DenseModelDescription] = [
|
||||
DenseModelDescription(
|
||||
model="Qdrant/clip-ViT-B-32-vision",
|
||||
dim=512,
|
||||
description="Image embeddings, Multimodal (text&image), 2021 year",
|
||||
license="mit",
|
||||
size_in_GB=0.34,
|
||||
sources=ModelSource(hf="Qdrant/clip-ViT-B-32-vision"),
|
||||
model_file="model.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="Qdrant/resnet50-onnx",
|
||||
dim=2048,
|
||||
description="Image embeddings, Unimodal (image), 2016 year",
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.1,
|
||||
sources=ModelSource(hf="Qdrant/resnet50-onnx"),
|
||||
model_file="model.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="Qdrant/Unicom-ViT-B-16",
|
||||
dim=768,
|
||||
description="Image embeddings (more detailed than Unicom-ViT-B-32), Multimodal (text&image), 2023 year",
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.82,
|
||||
sources=ModelSource(hf="Qdrant/Unicom-ViT-B-16"),
|
||||
model_file="model.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="Qdrant/Unicom-ViT-B-32",
|
||||
dim=512,
|
||||
description="Image embeddings, Multimodal (text&image), 2023 year",
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.48,
|
||||
sources=ModelSource(hf="Qdrant/Unicom-ViT-B-32"),
|
||||
model_file="model.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="jinaai/jina-clip-v1",
|
||||
dim=768,
|
||||
description="Image embeddings, Multimodal (text&image), 2024 year",
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.34,
|
||||
sources=ModelSource(hf="jinaai/jina-clip-v1"),
|
||||
model_file="onnx/vision_model.onnx",
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[np.ndarray]):
|
||||
class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
**kwargs,
|
||||
cuda: bool = False,
|
||||
device_ids: Optional[list[int]] = None,
|
||||
lazy_load: bool = False,
|
||||
device_id: Optional[int] = None,
|
||||
specific_model_path: Optional[str] = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
@@ -48,43 +81,78 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[np.ndarray]):
|
||||
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 list of onnxruntime providers to use.
|
||||
Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
|
||||
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
|
||||
Defaults to False.
|
||||
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
|
||||
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
|
||||
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
|
||||
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
|
||||
device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
|
||||
specific_model_path (Optional[str], optional): The specific path to the onnx model dir if it should be imported from somewhere else
|
||||
|
||||
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)
|
||||
self.providers = providers
|
||||
self.lazy_load = lazy_load
|
||||
|
||||
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
|
||||
# List of device ids, that can be used for data parallel processing in workers
|
||||
self.device_ids = device_ids
|
||||
self.cuda = cuda
|
||||
|
||||
# This device_id will be used if we need to load model in current process
|
||||
self.device_id: Optional[int] = None
|
||||
if device_id is not None:
|
||||
self.device_id = device_id
|
||||
elif self.device_ids is not None:
|
||||
self.device_id = self.device_ids[0]
|
||||
|
||||
self.model_description = self._get_model_description(model_name)
|
||||
self.cache_dir = str(define_cache_dir(cache_dir))
|
||||
self._model_dir = self.download_model(
|
||||
self.model_description,
|
||||
self.cache_dir,
|
||||
local_files_only=self._local_files_only,
|
||||
specific_model_path=specific_model_path,
|
||||
)
|
||||
|
||||
self.load_onnx_model(
|
||||
model_dir=model_dir,
|
||||
model_file=model_description["model_file"],
|
||||
threads=threads,
|
||||
providers=providers,
|
||||
if not self.lazy_load:
|
||||
self.load_onnx_model()
|
||||
|
||||
def load_onnx_model(self) -> None:
|
||||
"""
|
||||
Load the onnx model.
|
||||
"""
|
||||
self._load_onnx_model(
|
||||
model_dir=self._model_dir,
|
||||
model_file=self.model_description.model_file,
|
||||
threads=self.threads,
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_id=self.device_id,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def list_supported_models(cls) -> List[Dict[str, Any]]:
|
||||
def _list_supported_models(cls) -> list[DenseModelDescription]:
|
||||
"""
|
||||
Lists the supported models.
|
||||
|
||||
Returns:
|
||||
List[Dict[str, Any]]: A list of dictionaries containing the model information.
|
||||
list[DenseModelDescription]: A list of DenseModelDescription objects containing the model information.
|
||||
"""
|
||||
return supported_onnx_models
|
||||
|
||||
def embed(
|
||||
self,
|
||||
images: ImageInput,
|
||||
images: Union[ImageInput, Iterable[ImageInput]],
|
||||
batch_size: int = 16,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs,
|
||||
) -> Iterable[np.ndarray]:
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
Encode a list of images into list of embeddings.
|
||||
We use mean pooling with attention so that the model can handle variable-length inputs.
|
||||
@@ -100,32 +168,41 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[np.ndarray]):
|
||||
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,
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_ids=self.device_ids,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type["ImageEmbeddingWorker"]:
|
||||
def _get_worker_class(cls) -> Type["ImageEmbeddingWorker[NumpyArray]"]:
|
||||
return OnnxImageEmbeddingWorker
|
||||
|
||||
def _preprocess_onnx_input(
|
||||
self, onnx_input: Dict[str, np.ndarray], **kwargs
|
||||
) -> Dict[str, np.ndarray]:
|
||||
self, onnx_input: dict[str, NumpyArray], **kwargs: Any
|
||||
) -> dict[str, NumpyArray]:
|
||||
"""
|
||||
Preprocess the onnx input.
|
||||
"""
|
||||
|
||||
return onnx_input
|
||||
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[np.ndarray]:
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[NumpyArray]:
|
||||
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)
|
||||
class OnnxImageEmbeddingWorker(ImageEmbeddingWorker[NumpyArray]):
|
||||
def init_embedding(self, model_name: str, cache_dir: str, **kwargs: Any) -> OnnxImageEmbedding:
|
||||
return OnnxImageEmbedding(
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
threads=1,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@@ -2,12 +2,14 @@ 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
|
||||
from typing import Any, Iterable, Optional, Sequence, Type, Union
|
||||
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
|
||||
from fastembed.common import ImageInput, OnnxProvider, PathInput
|
||||
from fastembed.image.transform.operators import Compose
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.common import ImageInput, OnnxProvider
|
||||
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
|
||||
@@ -18,7 +20,7 @@ from fastembed.parallel_processor import ParallelWorkerPool
|
||||
|
||||
class OnnxImageModel(OnnxModel[T]):
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type["ImageEmbeddingWorker"]:
|
||||
def _get_worker_class(cls) -> Type["ImageEmbeddingWorker[T]"]:
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[T]:
|
||||
@@ -26,41 +28,53 @@ class OnnxImageModel(OnnxModel[T]):
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.processor = None
|
||||
self.processor: Optional[Compose] = None
|
||||
|
||||
def _preprocess_onnx_input(
|
||||
self, onnx_input: Dict[str, np.ndarray], **kwargs
|
||||
) -> Dict[str, np.ndarray]:
|
||||
self, onnx_input: dict[str, NumpyArray], **kwargs: Any
|
||||
) -> dict[str, NumpyArray]:
|
||||
"""
|
||||
Preprocess the onnx input.
|
||||
"""
|
||||
return onnx_input
|
||||
|
||||
def load_onnx_model(
|
||||
def _load_onnx_model(
|
||||
self,
|
||||
model_dir: Path,
|
||||
model_file: str,
|
||||
threads: Optional[int],
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
cuda: bool = False,
|
||||
device_id: Optional[int] = None,
|
||||
) -> None:
|
||||
super().load_onnx_model(
|
||||
super()._load_onnx_model(
|
||||
model_dir=model_dir,
|
||||
model_file=model_file,
|
||||
threads=threads,
|
||||
providers=providers,
|
||||
cuda=cuda,
|
||||
device_id=device_id,
|
||||
)
|
||||
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 load_onnx_model(self) -> None:
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
def onnx_embed(self, images: List[PathInput], **kwargs) -> OnnxOutputContext:
|
||||
def _build_onnx_input(self, encoded: NumpyArray) -> dict[str, NumpyArray]:
|
||||
input_name = self.model.get_inputs()[0].name # type: ignore[union-attr]
|
||||
return {input_name: encoded}
|
||||
|
||||
def onnx_embed(self, images: list[ImageInput], **kwargs: Any) -> OnnxOutputContext:
|
||||
with contextlib.ExitStack():
|
||||
image_files = [Image.open(image) for image in images]
|
||||
encoded = self.processor(image_files)
|
||||
image_files = [
|
||||
Image.open(image) if not isinstance(image, Image.Image) else image
|
||||
for image in images
|
||||
]
|
||||
assert self.processor is not None, "Processor is not initialized"
|
||||
encoded = np.array(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)
|
||||
model_output = self.model.run(None, onnx_input) # type: ignore[union-attr]
|
||||
embeddings = model_output[0].reshape(len(images), -1)
|
||||
return OnnxOutputContext(model_output=embeddings)
|
||||
|
||||
@@ -68,41 +82,54 @@ class OnnxImageModel(OnnxModel[T]):
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
images: ImageInput,
|
||||
images: Union[ImageInput, Iterable[ImageInput]],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs,
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
cuda: bool = False,
|
||||
device_ids: Optional[list[int]] = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[T]:
|
||||
is_small = False
|
||||
|
||||
if isinstance(images, str) or isinstance(images, Path):
|
||||
if isinstance(images, (str, Path, Image.Image)):
|
||||
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 isinstance(images, list) and len(images) < batch_size:
|
||||
is_small = True
|
||||
|
||||
if parallel is None or is_small:
|
||||
if not hasattr(self, "model") or self.model is None:
|
||||
self.load_onnx_model()
|
||||
|
||||
for batch in iter_batch(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}
|
||||
if parallel == 0:
|
||||
parallel = os.cpu_count()
|
||||
|
||||
start_method = "forkserver" if "forkserver" in get_all_start_methods() else "spawn"
|
||||
params = {
|
||||
"model_name": model_name,
|
||||
"cache_dir": cache_dir,
|
||||
"providers": providers,
|
||||
**kwargs,
|
||||
}
|
||||
|
||||
pool = ParallelWorkerPool(
|
||||
parallel, self._get_worker_class(), start_method=start_method
|
||||
num_workers=parallel or 1,
|
||||
worker=self._get_worker_class(),
|
||||
cuda=cuda,
|
||||
device_ids=device_ids,
|
||||
start_method=start_method,
|
||||
)
|
||||
for batch in pool.ordered_map(iter_batch(images, batch_size), **params):
|
||||
yield from self._post_process_onnx_output(batch)
|
||||
yield from self._post_process_onnx_output(batch) # type: ignore
|
||||
|
||||
|
||||
class ImageEmbeddingWorker(EmbeddingWorker):
|
||||
def process(self, items: Iterable[Tuple[int, Any]]) -> Iterable[Tuple[int, Any]]:
|
||||
class ImageEmbeddingWorker(EmbeddingWorker[T]):
|
||||
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
|
||||
|
||||
@@ -1,8 +1,10 @@
|
||||
from typing import Sized, Tuple, Union
|
||||
from typing import Union
|
||||
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
|
||||
from fastembed.common.types import NumpyArray
|
||||
|
||||
|
||||
def convert_to_rgb(image: Image.Image) -> Image.Image:
|
||||
if image.mode == "RGB":
|
||||
@@ -13,9 +15,9 @@ def convert_to_rgb(image: Image.Image) -> Image.Image:
|
||||
|
||||
|
||||
def center_crop(
|
||||
image: Union[Image.Image, np.ndarray],
|
||||
size: Tuple[int, int],
|
||||
) -> np.ndarray:
|
||||
image: Union[Image.Image, NumpyArray],
|
||||
size: tuple[int, int],
|
||||
) -> NumpyArray:
|
||||
if isinstance(image, np.ndarray):
|
||||
_, orig_height, orig_width = image.shape
|
||||
else:
|
||||
@@ -40,7 +42,7 @@ def center_crop(
|
||||
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)
|
||||
new_image = np.zeros_like(image, shape=new_shape, dtype=np.float32)
|
||||
|
||||
top_pad = (new_height - orig_height) // 2
|
||||
bottom_pad = top_pad + orig_height
|
||||
@@ -61,45 +63,42 @@ def center_crop(
|
||||
|
||||
|
||||
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")
|
||||
|
||||
image: NumpyArray,
|
||||
mean: Union[float, list[float]],
|
||||
std: Union[float, list[float]],
|
||||
) -> NumpyArray:
|
||||
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)
|
||||
mean = mean if isinstance(mean, list) else [mean] * num_channels
|
||||
|
||||
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)
|
||||
if len(mean) != num_channels:
|
||||
raise ValueError(
|
||||
f"mean must have the same number of channels as the image, image has {num_channels} channels, got "
|
||||
f"{len(mean)}"
|
||||
)
|
||||
|
||||
image = ((image.T - mean) / std).T
|
||||
mean_arr = np.array(mean, dtype=np.float32)
|
||||
|
||||
std = std if isinstance(std, list) else [std] * num_channels
|
||||
if len(std) != num_channels:
|
||||
raise ValueError(
|
||||
f"std must have the same number of channels as the image, image has {num_channels} channels, got {len(std)}"
|
||||
)
|
||||
|
||||
std_arr = np.array(std, dtype=np.float32)
|
||||
|
||||
image = ((image.T - mean_arr) / std_arr).T
|
||||
return image
|
||||
|
||||
|
||||
def resize(
|
||||
image: Image,
|
||||
size: Union[int, Tuple[int, int]],
|
||||
resample: Image.Resampling = Image.Resampling.BILINEAR,
|
||||
) -> Image:
|
||||
image: Image.Image,
|
||||
size: Union[int, tuple[int, int]],
|
||||
resample: Union[int, Image.Resampling] = Image.Resampling.BILINEAR,
|
||||
) -> Image.Image:
|
||||
if isinstance(size, tuple):
|
||||
return image.resize(size, resample)
|
||||
|
||||
@@ -114,11 +113,37 @@ def resize(
|
||||
return image.resize(new_size, resample)
|
||||
|
||||
|
||||
def rescale(image: np.ndarray, scale: float, dtype=np.float32) -> np.ndarray:
|
||||
def rescale(image: NumpyArray, scale: float, dtype: type = np.float32) -> NumpyArray:
|
||||
return (image * scale).astype(dtype)
|
||||
|
||||
|
||||
def pil2ndarray(image: Union[Image.Image, np.ndarray]):
|
||||
def pil2ndarray(image: Union[Image.Image, NumpyArray]) -> NumpyArray:
|
||||
if isinstance(image, Image.Image):
|
||||
return np.asarray(image).transpose((2, 0, 1))
|
||||
return image
|
||||
|
||||
|
||||
def pad2square(
|
||||
image: Image.Image,
|
||||
size: int,
|
||||
fill_color: Union[str, int, tuple[int, ...]] = 0,
|
||||
) -> Image.Image:
|
||||
height, width = image.height, image.width
|
||||
|
||||
left, right = 0, width
|
||||
top, bottom = 0, height
|
||||
|
||||
crop_required = False
|
||||
if width > size:
|
||||
left = (width - size) // 2
|
||||
right = left + size
|
||||
crop_required = True
|
||||
|
||||
if height > size:
|
||||
top = (height - size) // 2
|
||||
bottom = top + size
|
||||
crop_required = True
|
||||
|
||||
new_image = Image.new(mode="RGB", size=(size, size), color=fill_color)
|
||||
new_image.paste(image.crop((left, top, right, bottom)) if crop_required else image)
|
||||
return new_image
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
from typing import Any, Dict, List, Tuple, Union
|
||||
from typing import Any, Union, Optional
|
||||
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.image.transform.functional import (
|
||||
center_crop,
|
||||
convert_to_rgb,
|
||||
@@ -10,93 +10,111 @@ from fastembed.image.transform.functional import (
|
||||
pil2ndarray,
|
||||
rescale,
|
||||
resize,
|
||||
pad2square,
|
||||
)
|
||||
|
||||
|
||||
class Transform:
|
||||
def __call__(self, images: List) -> Union[List[Image.Image], List[np.ndarray]]:
|
||||
def __call__(self, images: list[Any]) -> Union[list[Image.Image], list[NumpyArray]]:
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
|
||||
class ConvertToRGB(Transform):
|
||||
def __call__(self, images: List[Image.Image]) -> List[Image.Image]:
|
||||
def __call__(self, images: list[Image.Image]) -> list[Image.Image]:
|
||||
return [convert_to_rgb(image=image) for image in images]
|
||||
|
||||
|
||||
class CenterCrop(Transform):
|
||||
def __init__(self, size: Tuple[int, int]):
|
||||
def __init__(self, size: tuple[int, int]):
|
||||
self.size = size
|
||||
|
||||
def __call__(self, images: List[Image.Image]) -> List[np.ndarray]:
|
||||
def __call__(self, images: list[Image.Image]) -> list[NumpyArray]:
|
||||
return [center_crop(image=image, size=self.size) for image in images]
|
||||
|
||||
|
||||
class Normalize(Transform):
|
||||
def __init__(self, mean: Union[float, List[float]], std: Union[float, List[float]]):
|
||||
def __init__(self, mean: Union[float, list[float]], std: Union[float, list[float]]):
|
||||
self.mean = mean
|
||||
self.std = std
|
||||
|
||||
def __call__(self, images: List[np.ndarray]) -> List[np.ndarray]:
|
||||
def __call__(self, images: list[NumpyArray]) -> list[NumpyArray]:
|
||||
return [normalize(image, mean=self.mean, std=self.std) for image in images]
|
||||
|
||||
|
||||
class Resize(Transform):
|
||||
def __init__(
|
||||
self,
|
||||
size: Union[int, Tuple[int, int]],
|
||||
size: Union[int, tuple[int, int]],
|
||||
resample: Image.Resampling = Image.Resampling.BICUBIC,
|
||||
):
|
||||
self.size = size
|
||||
self.resample = resample
|
||||
|
||||
def __call__(self, images: List[Image.Image]) -> List[Image.Image]:
|
||||
return [
|
||||
resize(image, size=self.size, resample=self.resample) for image in images
|
||||
]
|
||||
def __call__(self, images: list[Image.Image]) -> list[Image.Image]:
|
||||
return [resize(image, size=self.size, resample=self.resample) for image in images]
|
||||
|
||||
|
||||
class Rescale(Transform):
|
||||
def __init__(self, scale: float = 1 / 255):
|
||||
self.scale = scale
|
||||
|
||||
def __call__(self, images: List[np.ndarray]) -> List[np.ndarray]:
|
||||
def __call__(self, images: list[NumpyArray]) -> list[NumpyArray]:
|
||||
return [rescale(image, scale=self.scale) for image in images]
|
||||
|
||||
|
||||
class PILtoNDarray(Transform):
|
||||
def __call__(
|
||||
self, images: List[Union[Image.Image, np.ndarray]]
|
||||
) -> List[np.ndarray]:
|
||||
def __call__(self, images: list[Union[Image.Image, NumpyArray]]) -> list[NumpyArray]:
|
||||
return [pil2ndarray(image) for image in images]
|
||||
|
||||
|
||||
class PadtoSquare(Transform):
|
||||
def __init__(
|
||||
self,
|
||||
size: int,
|
||||
fill_color: Union[str, int, tuple[int, ...]],
|
||||
):
|
||||
self.size = size
|
||||
self.fill_color = fill_color
|
||||
|
||||
def __call__(self, images: list[Image.Image]) -> list[Image.Image]:
|
||||
return [
|
||||
pad2square(image=image, size=self.size, fill_color=self.fill_color) for image in images
|
||||
]
|
||||
|
||||
|
||||
class Compose:
|
||||
def __init__(self, transforms: List[Transform]):
|
||||
def __init__(self, transforms: list[Transform]):
|
||||
self.transforms = transforms
|
||||
|
||||
def __call__(
|
||||
self, images: Union[List[Image.Image], List[np.ndarray]]
|
||||
) -> Union[List[np.ndarray], List[Image.Image]]:
|
||||
self, images: Union[list[Image.Image], list[NumpyArray]]
|
||||
) -> Union[list[NumpyArray], list[Image.Image]]:
|
||||
for transform in self.transforms:
|
||||
images = transform(images)
|
||||
return images
|
||||
|
||||
@classmethod
|
||||
def from_config(cls, config: Dict[str, Any]) -> "Compose":
|
||||
def from_config(cls, config: dict[str, Any]) -> "Compose":
|
||||
"""Creates processor from a config dict.
|
||||
Args:
|
||||
config (Dict[str, Any]): Configuration dictionary.
|
||||
config (dict[str, Any]): Configuration dictionary.
|
||||
|
||||
Valid keys:
|
||||
- do_resize
|
||||
- resize_mode
|
||||
- size
|
||||
- fill_color
|
||||
- do_center_crop
|
||||
- crop_size
|
||||
- do_rescale
|
||||
- rescale_factor
|
||||
- do_normalize
|
||||
- image_mean
|
||||
- mean
|
||||
- image_std
|
||||
- std
|
||||
- resample
|
||||
- interpolation
|
||||
Valid size keys (nested):
|
||||
- {"height", "width"}
|
||||
- {"shortest_edge"}
|
||||
@@ -104,9 +122,10 @@ class Compose:
|
||||
Returns:
|
||||
Compose: Image processor.
|
||||
"""
|
||||
transforms = []
|
||||
transforms: list[Transform] = []
|
||||
cls._get_convert_to_rgb(transforms, config)
|
||||
cls._get_resize(transforms, config)
|
||||
cls._get_pad2square(transforms, config)
|
||||
cls._get_center_crop(transforms, config)
|
||||
cls._get_pil2ndarray(transforms, config)
|
||||
cls._get_rescale(transforms, config)
|
||||
@@ -114,13 +133,13 @@ class Compose:
|
||||
return cls(transforms=transforms)
|
||||
|
||||
@staticmethod
|
||||
def _get_convert_to_rgb(transforms: List[Transform], config: Dict[str, Any]):
|
||||
def _get_convert_to_rgb(transforms: list[Transform], config: dict[str, Any]) -> None:
|
||||
transforms.append(ConvertToRGB())
|
||||
|
||||
@staticmethod
|
||||
def _get_resize(transforms: List[Transform], config: Dict[str, Any]):
|
||||
@classmethod
|
||||
def _get_resize(cls, transforms: list[Transform], config: dict[str, Any]) -> None:
|
||||
mode = config.get("image_processor_type", "CLIPImageProcessor")
|
||||
if mode == "CLIPImageProcessor":
|
||||
if mode in ("CLIPImageProcessor", "SiglipImageProcessor"):
|
||||
if config.get("do_resize", False):
|
||||
size = config["size"]
|
||||
if "shortest_edge" in size:
|
||||
@@ -161,38 +180,90 @@ class Compose:
|
||||
resample=config.get("resample", Image.Resampling.BICUBIC),
|
||||
)
|
||||
)
|
||||
elif mode == "JinaCLIPImageProcessor":
|
||||
interpolation = config.get("interpolation")
|
||||
if isinstance(interpolation, str):
|
||||
resample = cls._interpolation_resolver(interpolation)
|
||||
else:
|
||||
resample = interpolation or Image.Resampling.BICUBIC
|
||||
|
||||
if "size" in config:
|
||||
resize_mode = config.get("resize_mode", "shortest")
|
||||
if resize_mode == "shortest":
|
||||
transforms.append(
|
||||
Resize(
|
||||
size=config["size"],
|
||||
resample=resample,
|
||||
)
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Preprocessor {mode} is not supported")
|
||||
|
||||
@staticmethod
|
||||
def _get_center_crop(transforms: List[Transform], config: Dict[str, Any]):
|
||||
def _get_center_crop(transforms: list[Transform], config: dict[str, Any]) -> None:
|
||||
mode = config.get("image_processor_type", "CLIPImageProcessor")
|
||||
if mode == "CLIPImageProcessor":
|
||||
if mode in ("CLIPImageProcessor", "SiglipImageProcessor"):
|
||||
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"])
|
||||
crop_size_raw = config["crop_size"]
|
||||
crop_size: tuple[int, int]
|
||||
if isinstance(crop_size_raw, int):
|
||||
crop_size = (crop_size_raw, crop_size_raw)
|
||||
elif isinstance(crop_size_raw, dict):
|
||||
crop_size = (crop_size_raw["height"], crop_size_raw["width"])
|
||||
else:
|
||||
raise ValueError(f"Invalid crop size: {crop_size}")
|
||||
raise ValueError(f"Invalid crop size: {crop_size_raw}")
|
||||
transforms.append(CenterCrop(size=crop_size))
|
||||
elif mode == "ConvNextFeatureExtractor":
|
||||
pass
|
||||
elif mode == "JinaCLIPImageProcessor":
|
||||
pass
|
||||
else:
|
||||
raise ValueError(f"Preprocessor {mode} is not supported")
|
||||
|
||||
@staticmethod
|
||||
def _get_pil2ndarray(transforms: List[Transform], config: Dict[str, Any]):
|
||||
def _get_pil2ndarray(transforms: list[Transform], config: dict[str, Any]) -> None:
|
||||
transforms.append(PILtoNDarray())
|
||||
|
||||
@staticmethod
|
||||
def _get_rescale(transforms: List[Transform], config: Dict[str, Any]):
|
||||
def _get_rescale(transforms: list[Transform], config: dict[str, Any]) -> None:
|
||||
if config.get("do_rescale", True):
|
||||
rescale_factor = config.get("rescale_factor", 1 / 255)
|
||||
transforms.append(Rescale(scale=rescale_factor))
|
||||
|
||||
@staticmethod
|
||||
def _get_normalize(transforms: List[Transform], config: Dict[str, Any]):
|
||||
def _get_normalize(transforms: list[Transform], config: dict[str, Any]) -> None:
|
||||
if config.get("do_normalize", False):
|
||||
transforms.append(Normalize(mean=config["image_mean"], std=config["image_std"]))
|
||||
elif "mean" in config and "std" in config:
|
||||
transforms.append(Normalize(mean=config["mean"], std=config["std"]))
|
||||
|
||||
@staticmethod
|
||||
def _get_pad2square(transforms: list[Transform], config: dict[str, Any]) -> None:
|
||||
mode = config.get("image_processor_type", "CLIPImageProcessor")
|
||||
if mode == "CLIPImageProcessor":
|
||||
pass
|
||||
elif mode == "ConvNextFeatureExtractor":
|
||||
pass
|
||||
elif mode == "JinaCLIPImageProcessor":
|
||||
transforms.append(
|
||||
Normalize(mean=config["image_mean"], std=config["image_std"])
|
||||
PadtoSquare(
|
||||
size=config["size"],
|
||||
fill_color=config.get("fill_color", 0),
|
||||
)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _interpolation_resolver(resample: Optional[str] = None) -> Image.Resampling:
|
||||
interpolation_map = {
|
||||
"nearest": Image.Resampling.NEAREST,
|
||||
"lanczos": Image.Resampling.LANCZOS,
|
||||
"bilinear": Image.Resampling.BILINEAR,
|
||||
"bicubic": Image.Resampling.BICUBIC,
|
||||
"box": Image.Resampling.BOX,
|
||||
"hamming": Image.Resampling.HAMMING,
|
||||
}
|
||||
|
||||
if resample and (method := interpolation_map.get(resample.lower())):
|
||||
return method
|
||||
|
||||
raise ValueError(f"Unknown interpolation method: {resample}")
|
||||
|
||||
@@ -1,9 +1,10 @@
|
||||
import string
|
||||
from typing import Any, Dict, Iterable, List, Optional, Sequence, Type, Union
|
||||
from typing import Any, Iterable, Optional, Sequence, Type, Union
|
||||
|
||||
import numpy as np
|
||||
from tokenizers import Encoding
|
||||
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.common import OnnxProvider
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.common.utils import define_cache_dir
|
||||
@@ -11,35 +12,49 @@ from fastembed.late_interaction.late_interaction_embedding_base import (
|
||||
LateInteractionTextEmbeddingBase,
|
||||
)
|
||||
from fastembed.text.onnx_text_model import OnnxTextModel, TextEmbeddingWorker
|
||||
from fastembed.common.model_description import DenseModelDescription, ModelSource
|
||||
|
||||
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",
|
||||
}
|
||||
supported_colbert_models: list[DenseModelDescription] = [
|
||||
DenseModelDescription(
|
||||
model="colbert-ir/colbertv2.0",
|
||||
dim=128,
|
||||
description="Late interaction model",
|
||||
license="mit",
|
||||
size_in_GB=0.44,
|
||||
sources=ModelSource(hf="colbert-ir/colbertv2.0"),
|
||||
model_file="model.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="answerdotai/answerai-colbert-small-v1",
|
||||
dim=96,
|
||||
description="Text embeddings, Unimodal (text), Multilingual (~100 languages), 512 input tokens truncation, 2024 year",
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.13,
|
||||
sources=ModelSource(hf="answerdotai/answerai-colbert-small-v1"),
|
||||
model_file="vespa_colbert.onnx",
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[np.ndarray]):
|
||||
class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
|
||||
QUERY_MARKER_TOKEN_ID = 1
|
||||
DOCUMENT_MARKER_TOKEN_ID = 2
|
||||
MIN_QUERY_LENGTH = 32
|
||||
MIN_QUERY_LENGTH = 31 # it's 32, we add one additional special token in the beginning
|
||||
MASK_TOKEN = "[MASK]"
|
||||
|
||||
def _post_process_onnx_output(
|
||||
self, output: OnnxOutputContext, is_doc: bool = True
|
||||
) -> Iterable[np.ndarray]:
|
||||
) -> Iterable[NumpyArray]:
|
||||
if not is_doc:
|
||||
return output.model_output.astype(np.float32)
|
||||
|
||||
if output.input_ids is None or output.attention_mask is None:
|
||||
raise ValueError(
|
||||
"input_ids and attention_mask must be provided for document post-processing"
|
||||
)
|
||||
|
||||
for i, token_sequence in enumerate(output.input_ids):
|
||||
for j, token_id in enumerate(token_sequence):
|
||||
for j, token_id in enumerate(token_sequence): # type: ignore
|
||||
if token_id in self.skip_list or token_id == self.pad_token_id:
|
||||
output.attention_mask[i, j] = 0
|
||||
|
||||
@@ -50,25 +65,27 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[np.ndarray]):
|
||||
return output.model_output.astype(np.float32)
|
||||
|
||||
def _preprocess_onnx_input(
|
||||
self, onnx_input: Dict[str, np.ndarray], is_doc: bool = True
|
||||
) -> 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
|
||||
self, onnx_input: dict[str, NumpyArray], is_doc: bool = True, **kwargs: Any
|
||||
) -> dict[str, NumpyArray]:
|
||||
marker_token = self.DOCUMENT_MARKER_TOKEN_ID if is_doc else self.QUERY_MARKER_TOKEN_ID
|
||||
onnx_input["input_ids"] = np.insert(
|
||||
onnx_input["input_ids"].astype(np.int64), 1, marker_token, axis=1
|
||||
)
|
||||
onnx_input["attention_mask"] = np.insert(
|
||||
onnx_input["attention_mask"].astype(np.int64), 1, 1, axis=1
|
||||
)
|
||||
return onnx_input
|
||||
|
||||
def tokenize(self, documents: List[str], is_doc: bool = True) -> List[Encoding]:
|
||||
def tokenize(self, documents: list[str], is_doc: bool = True, **kwargs: Any) -> 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)
|
||||
def _tokenize_query(self, query: str) -> list[Encoding]:
|
||||
assert self.tokenizer is not None
|
||||
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
|
||||
@@ -79,25 +96,23 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[np.ndarray]):
|
||||
pad_id=self.mask_token_id,
|
||||
length=self.MIN_QUERY_LENGTH,
|
||||
)
|
||||
encoded = self.tokenizer.encode_batch(query)
|
||||
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)
|
||||
def _tokenize_documents(self, documents: list[str]) -> list[Encoding]:
|
||||
encoded = self.tokenizer.encode_batch(documents) # type: ignore[union-attr]
|
||||
return encoded
|
||||
|
||||
@classmethod
|
||||
def list_supported_models(cls) -> List[Dict[str, Any]]:
|
||||
def _list_supported_models(cls) -> list[DenseModelDescription]:
|
||||
"""Lists the supported models.
|
||||
|
||||
Returns:
|
||||
List[Dict[str, Any]]: A list of dictionaries containing the model information.
|
||||
list[DenseModelDescription]: A list of DenseModelDescription objects containing the model information.
|
||||
"""
|
||||
return supported_colbert_models
|
||||
|
||||
@@ -107,7 +122,12 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[np.ndarray]):
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
**kwargs,
|
||||
cuda: bool = False,
|
||||
device_ids: Optional[list[int]] = None,
|
||||
lazy_load: bool = False,
|
||||
device_id: Optional[int] = None,
|
||||
specific_model_path: Optional[str] = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
@@ -116,41 +136,79 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[np.ndarray]):
|
||||
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 list of onnxruntime providers to use.
|
||||
Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
|
||||
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
|
||||
Defaults to False.
|
||||
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
|
||||
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
|
||||
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
|
||||
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
|
||||
device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
|
||||
specific_model_path (Optional[str], optional): The specific path to the onnx model dir if it should be imported from somewhere else
|
||||
|
||||
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)
|
||||
self.providers = providers
|
||||
self.lazy_load = lazy_load
|
||||
|
||||
model_description = self._get_model_description(model_name)
|
||||
self.cache_dir = define_cache_dir(cache_dir)
|
||||
# List of device ids, that can be used for data parallel processing in workers
|
||||
self.device_ids = device_ids
|
||||
self.cuda = cuda
|
||||
|
||||
model_dir = self.download_model(
|
||||
model_description, self.cache_dir, local_files_only=self._local_files_only
|
||||
# This device_id will be used if we need to load model in current process
|
||||
self.device_id: Optional[int] = None
|
||||
if device_id is not None:
|
||||
self.device_id = device_id
|
||||
elif self.device_ids is not None:
|
||||
self.device_id = self.device_ids[0]
|
||||
|
||||
self.model_description = self._get_model_description(model_name)
|
||||
self.cache_dir = str(define_cache_dir(cache_dir))
|
||||
|
||||
self._model_dir = self.download_model(
|
||||
self.model_description,
|
||||
self.cache_dir,
|
||||
local_files_only=self._local_files_only,
|
||||
specific_model_path=specific_model_path,
|
||||
)
|
||||
self.mask_token_id: Optional[int] = None
|
||||
self.pad_token_id: Optional[int] = None
|
||||
self.skip_list: set[int] = set()
|
||||
|
||||
self.load_onnx_model(
|
||||
model_dir=model_dir,
|
||||
model_file=model_description["model_file"],
|
||||
threads=threads,
|
||||
providers=providers,
|
||||
if not self.lazy_load:
|
||||
self.load_onnx_model()
|
||||
|
||||
def load_onnx_model(self) -> None:
|
||||
self._load_onnx_model(
|
||||
model_dir=self._model_dir,
|
||||
model_file=self.model_description.model_file,
|
||||
threads=self.threads,
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_id=self.device_id,
|
||||
)
|
||||
self.mask_token_id = self.special_token_to_id["[MASK]"]
|
||||
assert self.tokenizer is not None
|
||||
self.mask_token_id = self.special_token_to_id[self.MASK_TOKEN]
|
||||
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
|
||||
}
|
||||
current_max_length = self.tokenizer.truncation["max_length"]
|
||||
# ensure not to overflow after adding document-marker
|
||||
self.tokenizer.enable_truncation(max_length=current_max_length - 1)
|
||||
|
||||
def embed(
|
||||
self,
|
||||
documents: Union[str, Iterable[str]],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs,
|
||||
) -> Iterable[np.ndarray]:
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
Encode a list of documents into list of embeddings.
|
||||
We use mean pooling with attention so that the model can handle variable-length inputs.
|
||||
@@ -172,23 +230,34 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[np.ndarray]):
|
||||
documents=documents,
|
||||
batch_size=batch_size,
|
||||
parallel=parallel,
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_ids=self.device_ids,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def query_embed(self, query: Union[str, List[str]], **kwargs) -> np.ndarray:
|
||||
def query_embed(self, query: Union[str, Iterable[str]], **kwargs: Any) -> Iterable[NumpyArray]:
|
||||
if isinstance(query, str):
|
||||
query = [query]
|
||||
|
||||
if not hasattr(self, "model") or self.model is None:
|
||||
self.load_onnx_model()
|
||||
|
||||
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]:
|
||||
def _get_worker_class(cls) -> Type[TextEmbeddingWorker[NumpyArray]]:
|
||||
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)
|
||||
class ColbertEmbeddingWorker(TextEmbeddingWorker[NumpyArray]):
|
||||
def init_embedding(self, model_name: str, cache_dir: str, **kwargs: Any) -> Colbert:
|
||||
return Colbert(
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
threads=1,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,58 @@
|
||||
from typing import Any, Type
|
||||
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.late_interaction.colbert import Colbert, ColbertEmbeddingWorker
|
||||
from fastembed.common.model_description import DenseModelDescription, ModelSource
|
||||
|
||||
supported_jina_colbert_models: list[DenseModelDescription] = [
|
||||
DenseModelDescription(
|
||||
model="jinaai/jina-colbert-v2",
|
||||
dim=128,
|
||||
description="New model that expands capabilities of colbert-v1 with multilingual and context length of 8192, 2024 year",
|
||||
license="cc-by-nc-4.0",
|
||||
size_in_GB=2.24,
|
||||
sources=ModelSource(hf="jinaai/jina-colbert-v2"),
|
||||
model_file="onnx/model.onnx",
|
||||
additional_files=["onnx/model.onnx_data"],
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
class JinaColbert(Colbert):
|
||||
QUERY_MARKER_TOKEN_ID = 250002
|
||||
DOCUMENT_MARKER_TOKEN_ID = 250003
|
||||
MIN_QUERY_LENGTH = 31 # it's 32, we add one additional special token in the beginning
|
||||
MASK_TOKEN = "<mask>"
|
||||
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type[ColbertEmbeddingWorker]:
|
||||
return JinaColbertEmbeddingWorker
|
||||
|
||||
@classmethod
|
||||
def _list_supported_models(cls) -> list[DenseModelDescription]:
|
||||
"""Lists the supported models.
|
||||
|
||||
Returns:
|
||||
list[DenseModelDescription]: A list of DenseModelDescription objects containing the model information.
|
||||
"""
|
||||
return supported_jina_colbert_models
|
||||
|
||||
def _preprocess_onnx_input(
|
||||
self, onnx_input: dict[str, NumpyArray], is_doc: bool = True, **kwargs: Any
|
||||
) -> dict[str, NumpyArray]:
|
||||
onnx_input = super()._preprocess_onnx_input(onnx_input, is_doc)
|
||||
|
||||
# the attention mask for jina-colbert-v2 is always 1 in queries
|
||||
if not is_doc:
|
||||
onnx_input["attention_mask"][:] = 1
|
||||
return onnx_input
|
||||
|
||||
|
||||
class JinaColbertEmbeddingWorker(ColbertEmbeddingWorker):
|
||||
def init_embedding(self, model_name: str, cache_dir: str, **kwargs: Any) -> JinaColbert:
|
||||
return JinaColbert(
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
threads=1,
|
||||
**kwargs,
|
||||
)
|
||||
@@ -1,17 +1,17 @@
|
||||
from typing import Iterable, Optional, Union
|
||||
|
||||
import numpy as np
|
||||
from typing import Iterable, Optional, Union, Any
|
||||
|
||||
from fastembed.common.model_description import DenseModelDescription
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.common.model_management import ModelManagement
|
||||
|
||||
|
||||
class LateInteractionTextEmbeddingBase(ModelManagement):
|
||||
class LateInteractionTextEmbeddingBase(ModelManagement[DenseModelDescription]):
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
**kwargs,
|
||||
**kwargs: Any,
|
||||
):
|
||||
self.model_name = model_name
|
||||
self.cache_dir = cache_dir
|
||||
@@ -23,11 +23,11 @@ class LateInteractionTextEmbeddingBase(ModelManagement):
|
||||
documents: Union[str, Iterable[str]],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs,
|
||||
) -> Iterable[np.ndarray]:
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
raise NotImplementedError()
|
||||
|
||||
def passage_embed(self, texts: Iterable[str], **kwargs) -> Iterable[np.ndarray]:
|
||||
def passage_embed(self, texts: Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
Embeds a list of text passages into a list of embeddings.
|
||||
|
||||
@@ -36,15 +36,13 @@ class LateInteractionTextEmbeddingBase(ModelManagement):
|
||||
**kwargs: Additional keyword argument to pass to the embed method.
|
||||
|
||||
Yields:
|
||||
Iterable[np.ndarray]: The embeddings.
|
||||
Iterable[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]:
|
||||
def query_embed(self, query: Union[str, Iterable[str]], **kwargs: Any) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
Embeds queries
|
||||
|
||||
@@ -52,11 +50,11 @@ class LateInteractionTextEmbeddingBase(ModelManagement):
|
||||
query (Union[str, Iterable[str]]): The query to embed, or an iterable e.g. list of queries.
|
||||
|
||||
Returns:
|
||||
Iterable[np.ndarray]: The embeddings.
|
||||
Iterable[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):
|
||||
else:
|
||||
yield from self.embed(query, **kwargs)
|
||||
|
||||
@@ -1,45 +1,51 @@
|
||||
from typing import Any, Dict, Iterable, List, Optional, Sequence, Type, Union
|
||||
|
||||
import numpy as np
|
||||
from typing import Any, Iterable, Optional, Sequence, Type, Union
|
||||
from dataclasses import asdict
|
||||
|
||||
from fastembed.common.model_description import DenseModelDescription
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.common import OnnxProvider
|
||||
from fastembed.late_interaction.colbert import Colbert
|
||||
from fastembed.late_interaction.jina_colbert import JinaColbert
|
||||
from fastembed.late_interaction.late_interaction_embedding_base import (
|
||||
LateInteractionTextEmbeddingBase,
|
||||
)
|
||||
|
||||
|
||||
class LateInteractionTextEmbedding(LateInteractionTextEmbeddingBase):
|
||||
EMBEDDINGS_REGISTRY: List[Type[LateInteractionTextEmbeddingBase]] = [
|
||||
Colbert,
|
||||
]
|
||||
EMBEDDINGS_REGISTRY: list[Type[LateInteractionTextEmbeddingBase]] = [Colbert, JinaColbert]
|
||||
|
||||
@classmethod
|
||||
def list_supported_models(cls) -> List[Dict[str, Any]]:
|
||||
def list_supported_models(cls) -> list[dict[str, Any]]:
|
||||
"""
|
||||
Lists the supported models.
|
||||
|
||||
Returns:
|
||||
List[Dict[str, Any]]: A list of dictionaries containing the model information.
|
||||
list[dict[str, Any]]: A list of dictionaries containing the model information.
|
||||
|
||||
Example:
|
||||
```
|
||||
[
|
||||
{
|
||||
"model": "prithvida/SPLADE_PP_en_v1",
|
||||
"vocab_size": 30522,
|
||||
"description": "Independent Implementation of SPLADE++ Model for English",
|
||||
"size_in_GB": 0.532,
|
||||
"model": "colbert-ir/colbertv2.0",
|
||||
"dim": 128,
|
||||
"description": "Late interaction model",
|
||||
"license": "mit",
|
||||
"size_in_GB": 0.44,
|
||||
"sources": {
|
||||
"hf": "qdrant/SPLADE_PP_en_v1",
|
||||
"hf": "colbert-ir/colbertv2.0",
|
||||
},
|
||||
}
|
||||
"model_file": "model.onnx",
|
||||
},
|
||||
]
|
||||
```
|
||||
"""
|
||||
result = []
|
||||
return [asdict(model) for model in cls._list_supported_models()]
|
||||
|
||||
@classmethod
|
||||
def _list_supported_models(cls) -> list[DenseModelDescription]:
|
||||
result: list[DenseModelDescription] = []
|
||||
for embedding in cls.EMBEDDINGS_REGISTRY:
|
||||
result.extend(embedding.list_supported_models())
|
||||
result.extend(embedding._list_supported_models())
|
||||
return result
|
||||
|
||||
def __init__(
|
||||
@@ -48,24 +54,30 @@ class LateInteractionTextEmbedding(LateInteractionTextEmbeddingBase):
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
**kwargs,
|
||||
cuda: bool = False,
|
||||
device_ids: Optional[list[int]] = None,
|
||||
lazy_load: bool = False,
|
||||
**kwargs: Any,
|
||||
):
|
||||
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
|
||||
):
|
||||
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
|
||||
model_name,
|
||||
cache_dir,
|
||||
threads=threads,
|
||||
providers=providers,
|
||||
cuda=cuda,
|
||||
device_ids=device_ids,
|
||||
lazy_load=lazy_load,
|
||||
**kwargs,
|
||||
)
|
||||
return
|
||||
|
||||
raise ValueError(
|
||||
f"Model {model_name} is not supported in SparseTextEmbedding."
|
||||
"Please check the supported models using `SparseTextEmbedding.list_supported_models()`"
|
||||
f"Model {model_name} is not supported in LateInteractionTextEmbedding."
|
||||
"Please check the supported models using `LateInteractionTextEmbedding.list_supported_models()`"
|
||||
)
|
||||
|
||||
def embed(
|
||||
@@ -73,8 +85,8 @@ class LateInteractionTextEmbedding(LateInteractionTextEmbeddingBase):
|
||||
documents: Union[str, Iterable[str]],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs,
|
||||
) -> Iterable[np.ndarray]:
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
Encode a list of documents into list of embeddings.
|
||||
We use mean pooling with attention so that the model can handle variable-length inputs.
|
||||
@@ -92,9 +104,7 @@ class LateInteractionTextEmbedding(LateInteractionTextEmbeddingBase):
|
||||
"""
|
||||
yield from self.model.embed(documents, batch_size, parallel, **kwargs)
|
||||
|
||||
def query_embed(
|
||||
self, query: Union[str, Iterable[str]], **kwargs
|
||||
) -> Iterable[np.ndarray]:
|
||||
def query_embed(self, query: Union[str, Iterable[str]], **kwargs: Any) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
Embeds queries
|
||||
|
||||
@@ -102,7 +112,7 @@ class LateInteractionTextEmbedding(LateInteractionTextEmbeddingBase):
|
||||
query (Union[str, Iterable[str]]): The query to embed, or an iterable e.g. list of queries.
|
||||
|
||||
Returns:
|
||||
Iterable[np.ndarray]: The embeddings.
|
||||
Iterable[NdArray]: The embeddings.
|
||||
"""
|
||||
|
||||
# This is model-specific, so that different models can have specialized implementations
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
from fastembed.late_interaction_multimodal.late_interaction_multimodal_embedding import (
|
||||
LateInteractionMultimodalEmbedding,
|
||||
)
|
||||
|
||||
__all__ = ["LateInteractionMultimodalEmbedding"]
|
||||
@@ -0,0 +1,301 @@
|
||||
from typing import Any, Iterable, Optional, Sequence, Type, Union
|
||||
|
||||
import numpy as np
|
||||
from tokenizers import Encoding
|
||||
|
||||
from fastembed.common import OnnxProvider, ImageInput
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.common.utils import define_cache_dir
|
||||
from fastembed.late_interaction_multimodal.late_interaction_multimodal_embedding_base import (
|
||||
LateInteractionMultimodalEmbeddingBase,
|
||||
)
|
||||
from fastembed.late_interaction_multimodal.onnx_multimodal_model import (
|
||||
OnnxMultimodalModel,
|
||||
TextEmbeddingWorker,
|
||||
ImageEmbeddingWorker,
|
||||
)
|
||||
from fastembed.common.model_description import DenseModelDescription, ModelSource
|
||||
|
||||
supported_colpali_models: list[DenseModelDescription] = [
|
||||
DenseModelDescription(
|
||||
model="Qdrant/colpali-v1.3-fp16",
|
||||
dim=128,
|
||||
description="Text embeddings, Multimodal (text&image), English, 50 tokens query length truncation, 2024.",
|
||||
license="mit",
|
||||
size_in_GB=6.5,
|
||||
sources=ModelSource(hf="Qdrant/colpali-v1.3-fp16"),
|
||||
additional_files=["model.onnx_data"],
|
||||
model_file="model.onnx",
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyArray]):
|
||||
QUERY_PREFIX = "Query: "
|
||||
BOS_TOKEN = "<s>"
|
||||
PAD_TOKEN = "<pad>"
|
||||
QUERY_MARKER_TOKEN_ID = [2, 5098]
|
||||
IMAGE_PLACEHOLDER_SIZE = (3, 448, 448)
|
||||
EMPTY_TEXT_PLACEHOLDER = np.array(
|
||||
[257152] * 1024 + [2, 50721, 573, 2416, 235265, 108]
|
||||
) # This is a tokenization of '<image>' * 1024 + '<bos>Describe the image.\n' line which is used as placeholder
|
||||
# while processing an image
|
||||
EVEN_ATTENTION_MASK = np.array([1] * 1030)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
cuda: bool = False,
|
||||
device_ids: Optional[list[int]] = None,
|
||||
lazy_load: bool = False,
|
||||
device_id: Optional[int] = None,
|
||||
specific_model_path: Optional[str] = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
"""
|
||||
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 list of onnxruntime providers to use.
|
||||
Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
|
||||
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
|
||||
Defaults to False.
|
||||
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
|
||||
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
|
||||
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
|
||||
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
|
||||
device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
|
||||
|
||||
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)
|
||||
self.providers = providers
|
||||
self.lazy_load = lazy_load
|
||||
|
||||
# List of device ids, that can be used for data parallel processing in workers
|
||||
self.device_ids = device_ids
|
||||
self.cuda = cuda
|
||||
|
||||
# This device_id will be used if we need to load model in current process
|
||||
self.device_id: Optional[int] = None
|
||||
if device_id is not None:
|
||||
self.device_id = device_id
|
||||
elif self.device_ids is not None:
|
||||
self.device_id = self.device_ids[0]
|
||||
|
||||
self.model_description = self._get_model_description(model_name)
|
||||
self.cache_dir = str(define_cache_dir(cache_dir))
|
||||
|
||||
self._model_dir = self.download_model(
|
||||
self.model_description,
|
||||
self.cache_dir,
|
||||
local_files_only=self._local_files_only,
|
||||
specific_model_path=specific_model_path,
|
||||
)
|
||||
self.mask_token_id = None
|
||||
self.pad_token_id = None
|
||||
|
||||
if not self.lazy_load:
|
||||
self.load_onnx_model()
|
||||
|
||||
@classmethod
|
||||
def _list_supported_models(cls) -> list[DenseModelDescription]:
|
||||
"""Lists the supported models.
|
||||
|
||||
Returns:
|
||||
list[DenseModelDescription]: A list of DenseModelDescription objects containing the model information.
|
||||
"""
|
||||
return supported_colpali_models
|
||||
|
||||
def load_onnx_model(self) -> None:
|
||||
self._load_onnx_model(
|
||||
model_dir=self._model_dir,
|
||||
model_file=self.model_description.model_file,
|
||||
threads=self.threads,
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_id=self.device_id,
|
||||
)
|
||||
|
||||
def _post_process_onnx_image_output(
|
||||
self,
|
||||
output: OnnxOutputContext,
|
||||
) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
Post-process the ONNX model output to convert it into a usable format.
|
||||
|
||||
Args:
|
||||
output (OnnxOutputContext): The raw output from the ONNX model.
|
||||
|
||||
Returns:
|
||||
Iterable[NumpyArray]: Post-processed output as NumPy arrays.
|
||||
"""
|
||||
assert self.model_description.dim is not None, "Model dim is not defined"
|
||||
return output.model_output.reshape(
|
||||
output.model_output.shape[0], -1, self.model_description.dim
|
||||
).astype(np.float32)
|
||||
|
||||
def _post_process_onnx_text_output(
|
||||
self,
|
||||
output: OnnxOutputContext,
|
||||
) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
Post-process the ONNX model output to convert it into a usable format.
|
||||
|
||||
Args:
|
||||
output (OnnxOutputContext): The raw output from the ONNX model.
|
||||
|
||||
Returns:
|
||||
Iterable[NumpyArray]: Post-processed output as NumPy arrays.
|
||||
"""
|
||||
return output.model_output.astype(np.float32)
|
||||
|
||||
def tokenize(self, documents: list[str], **kwargs: Any) -> list[Encoding]:
|
||||
texts_query: list[str] = []
|
||||
for query in documents:
|
||||
query = self.BOS_TOKEN + self.QUERY_PREFIX + query + self.PAD_TOKEN * 10
|
||||
query += "\n"
|
||||
|
||||
texts_query.append(query)
|
||||
encoded = self.tokenizer.encode_batch(texts_query) # type: ignore[union-attr]
|
||||
return encoded
|
||||
|
||||
def _preprocess_onnx_text_input(
|
||||
self, onnx_input: dict[str, NumpyArray], **kwargs: Any
|
||||
) -> dict[str, NumpyArray]:
|
||||
onnx_input["input_ids"] = np.array(
|
||||
[
|
||||
self.QUERY_MARKER_TOKEN_ID + input_ids[2:].tolist()
|
||||
for input_ids in onnx_input["input_ids"]
|
||||
]
|
||||
)
|
||||
empty_image_placeholder: NumpyArray = np.zeros(
|
||||
self.IMAGE_PLACEHOLDER_SIZE, dtype=np.float32
|
||||
)
|
||||
onnx_input["pixel_values"] = np.array(
|
||||
[empty_image_placeholder for _ in onnx_input["input_ids"]],
|
||||
)
|
||||
return onnx_input
|
||||
|
||||
def _preprocess_onnx_image_input(
|
||||
self, onnx_input: dict[str, np.ndarray], **kwargs: Any
|
||||
) -> dict[str, NumpyArray]:
|
||||
"""
|
||||
Add placeholders for text input when processing image data for ONNX.
|
||||
Args:
|
||||
onnx_input (Dict[str, NumpyArray]): Preprocessed image inputs.
|
||||
**kwargs: Additional arguments.
|
||||
Returns:
|
||||
Dict[str, NumpyArray]: ONNX input with text placeholders.
|
||||
"""
|
||||
|
||||
onnx_input["input_ids"] = np.array(
|
||||
[self.EMPTY_TEXT_PLACEHOLDER for _ in onnx_input["input_ids"]]
|
||||
)
|
||||
onnx_input["attention_mask"] = np.array(
|
||||
[self.EVEN_ATTENTION_MASK for _ in onnx_input["input_ids"]]
|
||||
)
|
||||
return onnx_input
|
||||
|
||||
def embed_text(
|
||||
self,
|
||||
documents: Union[str, Iterable[str]],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
Encode a list of documents into list of embeddings.
|
||||
|
||||
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,
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_ids=self.device_ids,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def embed_image(
|
||||
self,
|
||||
images: Union[ImageInput, Iterable[ImageInput]],
|
||||
batch_size: int = 16,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
Encode a list of images into list of embeddings.
|
||||
|
||||
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,
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_ids=self.device_ids,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _get_text_worker_class(cls) -> Type[TextEmbeddingWorker[NumpyArray]]:
|
||||
return ColPaliTextEmbeddingWorker
|
||||
|
||||
@classmethod
|
||||
def _get_image_worker_class(cls) -> Type[ImageEmbeddingWorker[NumpyArray]]:
|
||||
return ColPaliImageEmbeddingWorker
|
||||
|
||||
|
||||
class ColPaliTextEmbeddingWorker(TextEmbeddingWorker[NumpyArray]):
|
||||
def init_embedding(self, model_name: str, cache_dir: str, **kwargs: Any) -> ColPali:
|
||||
return ColPali(
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
threads=1,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
class ColPaliImageEmbeddingWorker(ImageEmbeddingWorker[NumpyArray]):
|
||||
def init_embedding(self, model_name: str, cache_dir: str, **kwargs: Any) -> ColPali:
|
||||
return ColPali(
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
threads=1,
|
||||
**kwargs,
|
||||
)
|
||||
@@ -0,0 +1,130 @@
|
||||
from typing import Any, Iterable, Optional, Sequence, Type, Union
|
||||
from dataclasses import asdict
|
||||
|
||||
from fastembed.common import OnnxProvider, ImageInput
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.late_interaction_multimodal.colpali import ColPali
|
||||
|
||||
from fastembed.late_interaction_multimodal.late_interaction_multimodal_embedding_base import (
|
||||
LateInteractionMultimodalEmbeddingBase,
|
||||
)
|
||||
from fastembed.common.model_description import DenseModelDescription
|
||||
|
||||
|
||||
class LateInteractionMultimodalEmbedding(LateInteractionMultimodalEmbeddingBase):
|
||||
EMBEDDINGS_REGISTRY: list[Type[LateInteractionMultimodalEmbeddingBase]] = [ColPali]
|
||||
|
||||
@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/colpali-v1.3-fp16",
|
||||
"dim": 128,
|
||||
"description": "Text embeddings, Unimodal (text), Aligned to image latent space, ColBERT-compatible, 512 tokens max, 2024.",
|
||||
"license": "mit",
|
||||
"size_in_GB": 6.06,
|
||||
"sources": {
|
||||
"hf": "Qdrant/colpali-v1.3-fp16",
|
||||
},
|
||||
"additional_files": [
|
||||
"model.onnx_data",
|
||||
],
|
||||
"model_file": "model.onnx",
|
||||
},
|
||||
]
|
||||
```
|
||||
"""
|
||||
return [asdict(model) for model in cls._list_supported_models()]
|
||||
|
||||
@classmethod
|
||||
def _list_supported_models(cls) -> list[DenseModelDescription]:
|
||||
result: list[DenseModelDescription] = []
|
||||
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,
|
||||
cuda: bool = False,
|
||||
device_ids: Optional[list[int]] = None,
|
||||
lazy_load: bool = False,
|
||||
**kwargs: Any,
|
||||
):
|
||||
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,
|
||||
cuda=cuda,
|
||||
device_ids=device_ids,
|
||||
lazy_load=lazy_load,
|
||||
**kwargs,
|
||||
)
|
||||
return
|
||||
|
||||
raise ValueError(
|
||||
f"Model {model_name} is not supported in LateInteractionMultimodalEmbedding."
|
||||
"Please check the supported models using `LateInteractionMultimodalEmbedding.list_supported_models()`"
|
||||
)
|
||||
|
||||
def embed_text(
|
||||
self,
|
||||
documents: Union[str, Iterable[str]],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
Encode a list of documents into list of embeddings.
|
||||
|
||||
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_text(documents, batch_size, parallel, **kwargs)
|
||||
|
||||
def embed_image(
|
||||
self,
|
||||
images: Union[ImageInput, Iterable[ImageInput]],
|
||||
batch_size: int = 16,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
Encode a list of images into list of embeddings.
|
||||
|
||||
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 image
|
||||
"""
|
||||
yield from self.model.embed_image(images, batch_size, parallel, **kwargs)
|
||||
@@ -0,0 +1,67 @@
|
||||
from typing import Iterable, Optional, Union, Any
|
||||
|
||||
|
||||
from fastembed.common import ImageInput
|
||||
from fastembed.common.model_description import DenseModelDescription
|
||||
from fastembed.common.model_management import ModelManagement
|
||||
from fastembed.common.types import NumpyArray
|
||||
|
||||
|
||||
class LateInteractionMultimodalEmbeddingBase(ModelManagement[DenseModelDescription]):
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
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_text(
|
||||
self,
|
||||
documents: Union[str, Iterable[str]],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
Embeds a list of documents into a list of embeddings.
|
||||
|
||||
Args:
|
||||
documents (Iterable[str]): The list of texts 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.
|
||||
**kwargs: Additional keyword argument to pass to the embed method.
|
||||
|
||||
Yields:
|
||||
Iterable[NumpyArray]: The embeddings.
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
def embed_image(
|
||||
self,
|
||||
images: Union[ImageInput, Iterable[ImageInput]],
|
||||
batch_size: int = 16,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
Encode a list of images into list of embeddings.
|
||||
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 image
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
@@ -0,0 +1,275 @@
|
||||
import contextlib
|
||||
import os
|
||||
from multiprocessing import get_all_start_methods
|
||||
from pathlib import Path
|
||||
from typing import Any, Iterable, Optional, Sequence, Type, Union
|
||||
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
from tokenizers import Encoding, Tokenizer
|
||||
|
||||
from fastembed.common import OnnxProvider, ImageInput
|
||||
from fastembed.common.onnx_model import EmbeddingWorker, OnnxModel, OnnxOutputContext, T
|
||||
from fastembed.common.preprocessor_utils import load_tokenizer, load_preprocessor
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.common.utils import iter_batch
|
||||
from fastembed.image.transform.operators import Compose
|
||||
from fastembed.parallel_processor import ParallelWorkerPool
|
||||
|
||||
|
||||
class OnnxMultimodalModel(OnnxModel[T]):
|
||||
ONNX_OUTPUT_NAMES: Optional[list[str]] = None
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.tokenizer: Optional[Tokenizer] = None
|
||||
self.processor: Optional[Compose] = None
|
||||
self.special_token_to_id: dict[str, int] = {}
|
||||
|
||||
def _preprocess_onnx_text_input(
|
||||
self, onnx_input: dict[str, NumpyArray], **kwargs: Any
|
||||
) -> dict[str, NumpyArray]:
|
||||
"""
|
||||
Preprocess the onnx input.
|
||||
"""
|
||||
return onnx_input
|
||||
|
||||
def _preprocess_onnx_image_input(
|
||||
self, onnx_input: dict[str, NumpyArray], **kwargs: Any
|
||||
) -> dict[str, NumpyArray]:
|
||||
"""
|
||||
Preprocess the onnx input.
|
||||
"""
|
||||
return onnx_input
|
||||
|
||||
@classmethod
|
||||
def _get_text_worker_class(cls) -> Type["TextEmbeddingWorker[T]"]:
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
@classmethod
|
||||
def _get_image_worker_class(cls) -> Type["ImageEmbeddingWorker[T]"]:
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
def _post_process_onnx_image_output(self, output: OnnxOutputContext) -> Iterable[T]:
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
def _post_process_onnx_text_output(self, output: OnnxOutputContext) -> Iterable[T]:
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
def _load_onnx_model(
|
||||
self,
|
||||
model_dir: Path,
|
||||
model_file: str,
|
||||
threads: Optional[int],
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
cuda: bool = False,
|
||||
device_id: Optional[int] = None,
|
||||
) -> None:
|
||||
super()._load_onnx_model(
|
||||
model_dir=model_dir,
|
||||
model_file=model_file,
|
||||
threads=threads,
|
||||
providers=providers,
|
||||
cuda=cuda,
|
||||
device_id=device_id,
|
||||
)
|
||||
assert self.tokenizer is not None
|
||||
self.tokenizer, self.special_token_to_id = load_tokenizer(model_dir=model_dir)
|
||||
self.processor = load_preprocessor(model_dir=model_dir)
|
||||
|
||||
def load_onnx_model(self) -> None:
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
def tokenize(self, documents: list[str], **kwargs: Any) -> list[Encoding]:
|
||||
return self.tokenizer.encode_batch(documents) # type: ignore[union-attr]
|
||||
|
||||
def onnx_embed_text(
|
||||
self,
|
||||
documents: list[str],
|
||||
**kwargs: Any,
|
||||
) -> 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]) # type: ignore[union-attr]
|
||||
input_names = {node.name for node in self.model.get_inputs()} # type: ignore[union-attr]
|
||||
onnx_input: dict[str, NumpyArray] = {
|
||||
"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_text_input(onnx_input, **kwargs)
|
||||
model_output = self.model.run(self.ONNX_OUTPUT_NAMES, onnx_input) # type: ignore[union-attr]
|
||||
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,
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
cuda: bool = False,
|
||||
device_ids: Optional[list[int]] = None,
|
||||
**kwargs: Any,
|
||||
) -> 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 is None or is_small:
|
||||
if not hasattr(self, "model") or self.model is None:
|
||||
self.load_onnx_model()
|
||||
for batch in iter_batch(documents, batch_size):
|
||||
yield from self._post_process_onnx_text_output(self.onnx_embed_text(batch))
|
||||
else:
|
||||
if parallel == 0:
|
||||
parallel = os.cpu_count()
|
||||
|
||||
start_method = "forkserver" if "forkserver" in get_all_start_methods() else "spawn"
|
||||
params = {
|
||||
"model_name": model_name,
|
||||
"cache_dir": cache_dir,
|
||||
"providers": providers,
|
||||
**kwargs,
|
||||
}
|
||||
|
||||
pool = ParallelWorkerPool(
|
||||
num_workers=parallel or 1,
|
||||
worker=self._get_text_worker_class(),
|
||||
cuda=cuda,
|
||||
device_ids=device_ids,
|
||||
start_method=start_method,
|
||||
)
|
||||
for batch in pool.ordered_map(iter_batch(documents, batch_size), **params):
|
||||
yield from self._post_process_onnx_text_output(batch) # type: ignore
|
||||
|
||||
def _build_onnx_image_input(self, encoded: NumpyArray) -> dict[str, NumpyArray]:
|
||||
input_name = self.model.get_inputs()[0].name # type: ignore[union-attr]
|
||||
return {input_name: encoded}
|
||||
|
||||
def onnx_embed_image(self, images: list[ImageInput], **kwargs: Any) -> OnnxOutputContext:
|
||||
with contextlib.ExitStack():
|
||||
image_files = [
|
||||
Image.open(image) if not isinstance(image, Image.Image) else image
|
||||
for image in images
|
||||
]
|
||||
assert self.processor is not None, "Processor is not initialized"
|
||||
encoded = np.array(self.processor(image_files))
|
||||
onnx_input = self._build_onnx_image_input(encoded)
|
||||
onnx_input = self._preprocess_onnx_image_input(onnx_input, **kwargs)
|
||||
model_output = self.model.run(None, onnx_input) # type: ignore[union-attr]
|
||||
embeddings = model_output[0].reshape(len(images), -1)
|
||||
return OnnxOutputContext(model_output=embeddings)
|
||||
|
||||
def _embed_images(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
images: Union[Iterable[ImageInput], ImageInput],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
cuda: bool = False,
|
||||
device_ids: Optional[list[int]] = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[T]:
|
||||
is_small = False
|
||||
|
||||
if isinstance(images, (str, Path, Image.Image)):
|
||||
images = [images]
|
||||
is_small = True
|
||||
|
||||
if isinstance(images, list) and len(images) < batch_size:
|
||||
is_small = True
|
||||
|
||||
if parallel is None or is_small:
|
||||
if not hasattr(self, "model") or self.model is None:
|
||||
self.load_onnx_model()
|
||||
|
||||
for batch in iter_batch(images, batch_size):
|
||||
yield from self._post_process_onnx_image_output(self.onnx_embed_image(batch))
|
||||
else:
|
||||
if parallel == 0:
|
||||
parallel = os.cpu_count()
|
||||
|
||||
start_method = "forkserver" if "forkserver" in get_all_start_methods() else "spawn"
|
||||
params = {
|
||||
"model_name": model_name,
|
||||
"cache_dir": cache_dir,
|
||||
"providers": providers,
|
||||
**kwargs,
|
||||
}
|
||||
|
||||
pool = ParallelWorkerPool(
|
||||
num_workers=parallel or 1,
|
||||
worker=self._get_image_worker_class(),
|
||||
cuda=cuda,
|
||||
device_ids=device_ids,
|
||||
start_method=start_method,
|
||||
)
|
||||
for batch in pool.ordered_map(iter_batch(images, batch_size), **params):
|
||||
yield from self._post_process_onnx_image_output(batch) # type: ignore
|
||||
|
||||
|
||||
class TextEmbeddingWorker(EmbeddingWorker[T]):
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
**kwargs: Any,
|
||||
):
|
||||
self.model: OnnxMultimodalModel
|
||||
super().__init__(model_name, cache_dir, **kwargs)
|
||||
|
||||
def init_embedding(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
**kwargs: Any,
|
||||
) -> OnnxMultimodalModel:
|
||||
raise NotImplementedError()
|
||||
|
||||
def process(self, items: Iterable[tuple[int, Any]]) -> Iterable[tuple[int, Any]]:
|
||||
for idx, batch in items:
|
||||
onnx_output = self.model.onnx_embed_text(batch)
|
||||
yield idx, onnx_output
|
||||
|
||||
|
||||
class ImageEmbeddingWorker(EmbeddingWorker[T]):
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
**kwargs: Any,
|
||||
):
|
||||
self.model: OnnxMultimodalModel
|
||||
super().__init__(model_name, cache_dir, **kwargs)
|
||||
|
||||
def init_embedding(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
**kwargs: Any,
|
||||
) -> OnnxMultimodalModel:
|
||||
raise NotImplementedError()
|
||||
|
||||
def process(self, items: Iterable[tuple[int, Any]]) -> Iterable[tuple[int, Any]]:
|
||||
for idx, batch in items:
|
||||
embeddings = self.model.onnx_embed_image(batch)
|
||||
yield idx, embeddings
|
||||
@@ -0,0 +1,16 @@
|
||||
from pathlib import Path
|
||||
import json
|
||||
from typing import Dict, List
|
||||
|
||||
|
||||
class ModelLoader:
|
||||
def __init__(self):
|
||||
self.config_dir = Path(__file__).parent / "configs"
|
||||
self._models: Dict[str, List[Dict]] = {}
|
||||
|
||||
def load_models(self, model_type: str) -> List[Dict]:
|
||||
if model_type not in self._models:
|
||||
config_path = self.config_dir / f"{model_type}_models.json"
|
||||
with open(config_path) as f:
|
||||
self._models[model_type] = json.load(f)["models"]
|
||||
return self._models[model_type]
|
||||
@@ -1,13 +1,15 @@
|
||||
import logging
|
||||
import os
|
||||
from collections import defaultdict
|
||||
from copy import deepcopy
|
||||
from enum import Enum
|
||||
from multiprocessing import Queue, get_context
|
||||
from multiprocessing.context import BaseContext
|
||||
from multiprocessing.process import BaseProcess
|
||||
from multiprocessing.sharedctypes import Synchronized as BaseValue
|
||||
from queue import Empty
|
||||
from typing import Any, Dict, Iterable, List, Optional, Tuple, Type
|
||||
from typing import Any, Iterable, Optional, Type
|
||||
|
||||
|
||||
# Single item should be processed in less than:
|
||||
processing_timeout = 10 * 60 # seconds
|
||||
@@ -23,10 +25,10 @@ class QueueSignals(str, Enum):
|
||||
|
||||
class Worker:
|
||||
@classmethod
|
||||
def start(cls, **kwargs: Any) -> "Worker":
|
||||
def start(cls, *args: Any, **kwargs: Any) -> "Worker":
|
||||
raise NotImplementedError()
|
||||
|
||||
def process(self, items: Iterable[Tuple[int, Any]]) -> Iterable[Tuple[int, Any]]:
|
||||
def process(self, items: Iterable[tuple[int, Any]]) -> Iterable[tuple[int, Any]]:
|
||||
raise NotImplementedError()
|
||||
|
||||
|
||||
@@ -36,7 +38,7 @@ def _worker(
|
||||
output_queue: Queue,
|
||||
num_active_workers: BaseValue,
|
||||
worker_id: int,
|
||||
kwargs: Optional[Dict[str, Any]] = None,
|
||||
kwargs: Optional[dict[str, Any]] = None,
|
||||
) -> None:
|
||||
"""
|
||||
A worker that pulls data pints off the input queue, and places the execution result on the output queue.
|
||||
@@ -47,7 +49,9 @@ def _worker(
|
||||
if kwargs is None:
|
||||
kwargs = {}
|
||||
|
||||
logging.info(f"Reader worker: {worker_id} PID: {os.getpid()}")
|
||||
logging.info(
|
||||
f"Reader worker: {worker_id} PID: {os.getpid()} Device: {kwargs.get('device_id', 'CPU')}"
|
||||
)
|
||||
try:
|
||||
worker = worker_class.start(**kwargs)
|
||||
|
||||
@@ -73,7 +77,9 @@ def _worker(
|
||||
# See:
|
||||
# https://docs.python.org/3.6/library/multiprocessing.html?highlight=process#pipes-and-queues
|
||||
# https://docs.python.org/3.6/library/multiprocessing.html?highlight=process#programming-guidelines
|
||||
input_queue.close()
|
||||
output_queue.close()
|
||||
input_queue.join_thread()
|
||||
output_queue.join_thread()
|
||||
|
||||
with num_active_workers.get_lock():
|
||||
@@ -84,16 +90,23 @@ def _worker(
|
||||
|
||||
class ParallelWorkerPool:
|
||||
def __init__(
|
||||
self, num_workers: int, worker: Type[Worker], start_method: Optional[str] = None
|
||||
self,
|
||||
num_workers: int,
|
||||
worker: Type[Worker],
|
||||
start_method: Optional[str] = None,
|
||||
device_ids: Optional[list[int]] = None,
|
||||
cuda: bool = False,
|
||||
):
|
||||
self.worker_class = worker
|
||||
self.num_workers = num_workers
|
||||
self.input_queue: Optional[Queue] = None
|
||||
self.output_queue: Optional[Queue] = None
|
||||
self.ctx: BaseContext = get_context(start_method)
|
||||
self.processes: List[BaseProcess] = []
|
||||
self.processes: list[BaseProcess] = []
|
||||
self.queue_size = self.num_workers * max_internal_batch_size
|
||||
|
||||
self.emergency_shutdown = False
|
||||
self.device_ids = device_ids
|
||||
self.cuda = cuda
|
||||
self.num_active_workers: Optional[BaseValue] = None
|
||||
|
||||
def start(self, **kwargs: Any) -> None:
|
||||
@@ -105,6 +118,12 @@ class ParallelWorkerPool:
|
||||
self.num_active_workers = ctx_value
|
||||
|
||||
for worker_id in range(0, self.num_workers):
|
||||
worker_kwargs = deepcopy(kwargs)
|
||||
if self.device_ids:
|
||||
device_id = self.device_ids[worker_id % len(self.device_ids)]
|
||||
worker_kwargs["device_id"] = device_id
|
||||
worker_kwargs["cuda"] = self.cuda
|
||||
|
||||
assert hasattr(self.ctx, "Process")
|
||||
process = self.ctx.Process(
|
||||
target=_worker,
|
||||
@@ -114,16 +133,14 @@ class ParallelWorkerPool:
|
||||
self.output_queue,
|
||||
self.num_active_workers,
|
||||
worker_id,
|
||||
kwargs.copy(),
|
||||
worker_kwargs,
|
||||
),
|
||||
)
|
||||
process.start()
|
||||
self.processes.append(process)
|
||||
|
||||
def ordered_map(
|
||||
self, stream: Iterable[Any], *args: Any, **kwargs: Any
|
||||
) -> Iterable[Any]:
|
||||
buffer = defaultdict(Any)
|
||||
def ordered_map(self, stream: Iterable[Any], *args: Any, **kwargs: Any) -> Iterable[Any]:
|
||||
buffer: defaultdict[int, Any] = defaultdict(Any) # type: ignore
|
||||
next_expected = 0
|
||||
|
||||
for idx, item in self.semi_ordered_map(stream, *args, **kwargs):
|
||||
@@ -134,7 +151,7 @@ class ParallelWorkerPool:
|
||||
|
||||
def semi_ordered_map(
|
||||
self, stream: Iterable[Any], *args: Any, **kwargs: Any
|
||||
) -> Iterable[Tuple[int, Any]]:
|
||||
) -> Iterable[tuple[int, Any]]:
|
||||
try:
|
||||
self.start(**kwargs)
|
||||
|
||||
@@ -144,6 +161,7 @@ class ParallelWorkerPool:
|
||||
pushed = 0
|
||||
read = 0
|
||||
for idx, item in enumerate(stream):
|
||||
self.check_worker_health()
|
||||
if pushed - read < self.queue_size:
|
||||
try:
|
||||
out_item = self.output_queue.get_nowait()
|
||||
@@ -170,6 +188,7 @@ class ParallelWorkerPool:
|
||||
self.input_queue.put(QueueSignals.stop)
|
||||
|
||||
while read < pushed:
|
||||
self.check_worker_health()
|
||||
out_item = self.output_queue.get(timeout=processing_timeout)
|
||||
if out_item == QueueSignals.error:
|
||||
self.join_or_terminate()
|
||||
@@ -179,8 +198,27 @@ class ParallelWorkerPool:
|
||||
finally:
|
||||
assert self.input_queue is not None, "Input queue is None"
|
||||
assert self.output_queue is not None, "Output queue is None"
|
||||
self.join()
|
||||
self.input_queue.close()
|
||||
self.output_queue.close()
|
||||
if self.emergency_shutdown:
|
||||
self.input_queue.cancel_join_thread()
|
||||
self.output_queue.cancel_join_thread()
|
||||
else:
|
||||
self.input_queue.join_thread()
|
||||
self.output_queue.join_thread()
|
||||
|
||||
def check_worker_health(self) -> None:
|
||||
"""
|
||||
Checks if any worker process has terminated unexpectedly
|
||||
"""
|
||||
for process in self.processes:
|
||||
if not process.is_alive() and process.exitcode != 0:
|
||||
self.emergency_shutdown = True
|
||||
self.join_or_terminate()
|
||||
raise RuntimeError(
|
||||
f"Worker PID: {process.pid} terminated unexpectedly with code {process.exitcode}"
|
||||
)
|
||||
|
||||
def join_or_terminate(self, timeout: Optional[int] = 1) -> None:
|
||||
"""
|
||||
@@ -210,4 +248,5 @@ class ParallelWorkerPool:
|
||||
https://eli.thegreenplace.net/2009/06/12/safely-using-destructors-in-python/.
|
||||
"""
|
||||
for process in self.processes:
|
||||
process.terminate()
|
||||
if process.is_alive():
|
||||
process.terminate()
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
partial
|
||||
@@ -0,0 +1,3 @@
|
||||
from fastembed.rerank.cross_encoder.text_cross_encoder import TextCrossEncoder
|
||||
|
||||
__all__ = ["TextCrossEncoder"]
|
||||
@@ -0,0 +1,215 @@
|
||||
from typing import Any, Iterable, Optional, Sequence, Type
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from fastembed.common import OnnxProvider
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.common.utils import define_cache_dir
|
||||
from fastembed.rerank.cross_encoder.onnx_text_model import (
|
||||
OnnxCrossEncoderModel,
|
||||
TextRerankerWorker,
|
||||
)
|
||||
from fastembed.rerank.cross_encoder.text_cross_encoder_base import TextCrossEncoderBase
|
||||
from fastembed.common.model_description import BaseModelDescription, ModelSource
|
||||
|
||||
supported_onnx_models: list[BaseModelDescription] = [
|
||||
BaseModelDescription(
|
||||
model="Xenova/ms-marco-MiniLM-L-6-v2",
|
||||
description="MiniLM-L-6-v2 model optimized for re-ranking tasks.",
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.08,
|
||||
sources=ModelSource(hf="Xenova/ms-marco-MiniLM-L-6-v2"),
|
||||
model_file="onnx/model.onnx",
|
||||
),
|
||||
BaseModelDescription(
|
||||
model="Xenova/ms-marco-MiniLM-L-12-v2",
|
||||
description="MiniLM-L-12-v2 model optimized for re-ranking tasks.",
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.12,
|
||||
sources=ModelSource(hf="Xenova/ms-marco-MiniLM-L-12-v2"),
|
||||
model_file="onnx/model.onnx",
|
||||
),
|
||||
BaseModelDescription(
|
||||
model="BAAI/bge-reranker-base",
|
||||
description="BGE reranker base model for cross-encoder re-ranking.",
|
||||
license="mit",
|
||||
size_in_GB=1.04,
|
||||
sources=ModelSource(hf="BAAI/bge-reranker-base"),
|
||||
model_file="onnx/model.onnx",
|
||||
),
|
||||
BaseModelDescription(
|
||||
model="jinaai/jina-reranker-v1-tiny-en",
|
||||
description="Designed for blazing-fast re-ranking with 8K context length and fewer parameters than jina-reranker-v1-turbo-en.",
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.13,
|
||||
sources=ModelSource(hf="jinaai/jina-reranker-v1-tiny-en"),
|
||||
model_file="onnx/model.onnx",
|
||||
),
|
||||
BaseModelDescription(
|
||||
model="jinaai/jina-reranker-v1-turbo-en",
|
||||
description="Designed for blazing-fast re-ranking with 8K context length.",
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.15,
|
||||
sources=ModelSource(hf="jinaai/jina-reranker-v1-turbo-en"),
|
||||
model_file="onnx/model.onnx",
|
||||
),
|
||||
BaseModelDescription(
|
||||
model="jinaai/jina-reranker-v2-base-multilingual",
|
||||
description="A multi-lingual reranker model for cross-encoder re-ranking with 1K context length and sliding window",
|
||||
license="cc-by-nc-4.0",
|
||||
size_in_GB=1.11,
|
||||
sources=ModelSource(hf="jinaai/jina-reranker-v2-base-multilingual"),
|
||||
model_file="onnx/model.onnx",
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
|
||||
@classmethod
|
||||
def _list_supported_models(cls) -> list[BaseModelDescription]:
|
||||
"""Lists the supported models.
|
||||
|
||||
Returns:
|
||||
list[BaseModelDescription]: A list of BaseModelDescription objects containing the model information.
|
||||
"""
|
||||
return supported_onnx_models
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
cuda: bool = False,
|
||||
device_ids: Optional[list[int]] = None,
|
||||
lazy_load: bool = False,
|
||||
device_id: Optional[int] = None,
|
||||
specific_model_path: Optional[str] = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
"""
|
||||
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 list of onnxruntime providers to use.
|
||||
Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
|
||||
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
|
||||
Defaults to False.
|
||||
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
|
||||
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
|
||||
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
|
||||
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
|
||||
device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
|
||||
specific_model_path (Optional[str], optional): The specific path to the onnx model dir if it should be imported from somewhere else
|
||||
|
||||
Raises:
|
||||
ValueError: If the model_name is not in the format <org>/<model> e.g. Xenova/ms-marco-MiniLM-L-6-v2.
|
||||
"""
|
||||
super().__init__(model_name, cache_dir, threads, **kwargs)
|
||||
self.providers = providers
|
||||
self.lazy_load = lazy_load
|
||||
|
||||
# List of device ids, that can be used for data parallel processing in workers
|
||||
self.device_ids = device_ids
|
||||
self.cuda = cuda
|
||||
|
||||
if self.device_ids is not None and len(self.device_ids) > 1:
|
||||
logger.warning(
|
||||
"Parallel execution is currently not supported for cross encoders, "
|
||||
f"only the first device will be used for inference: {self.device_ids[0]}."
|
||||
)
|
||||
|
||||
# This device_id will be used if we need to load model in current process
|
||||
self.device_id: Optional[int] = None
|
||||
if device_id is not None:
|
||||
self.device_id = device_id
|
||||
elif self.device_ids is not None:
|
||||
self.device_id = self.device_ids[0]
|
||||
|
||||
self.model_description = self._get_model_description(model_name)
|
||||
self.cache_dir = str(define_cache_dir(cache_dir))
|
||||
self._model_dir = self.download_model(
|
||||
self.model_description,
|
||||
self.cache_dir,
|
||||
local_files_only=self._local_files_only,
|
||||
specific_model_path=specific_model_path,
|
||||
)
|
||||
|
||||
if not self.lazy_load:
|
||||
self.load_onnx_model()
|
||||
|
||||
def load_onnx_model(self) -> None:
|
||||
self._load_onnx_model(
|
||||
model_dir=self._model_dir,
|
||||
model_file=self.model_description.model_file,
|
||||
threads=self.threads,
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_id=self.device_id,
|
||||
)
|
||||
|
||||
def rerank(
|
||||
self,
|
||||
query: str,
|
||||
documents: Iterable[str],
|
||||
batch_size: int = 64,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[float]:
|
||||
"""Reranks documents based on their relevance to a given query.
|
||||
|
||||
Args:
|
||||
query (str): The query string to which document relevance is calculated.
|
||||
documents (Iterable[str]): Iterable of documents to be reranked.
|
||||
batch_size (int, optional): The number of documents processed in each batch. Higher batch sizes improve speed
|
||||
but require more memory. Default is 64.
|
||||
Returns:
|
||||
Iterable[float]: An iterable of relevance scores for each document.
|
||||
"""
|
||||
|
||||
yield from self._rerank_documents(
|
||||
query=query, documents=documents, batch_size=batch_size, **kwargs
|
||||
)
|
||||
|
||||
def rerank_pairs(
|
||||
self,
|
||||
pairs: Iterable[tuple[str, str]],
|
||||
batch_size: int = 64,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[float]:
|
||||
yield from self._rerank_pairs(
|
||||
model_name=self.model_name,
|
||||
cache_dir=str(self.cache_dir),
|
||||
pairs=pairs,
|
||||
batch_size=batch_size,
|
||||
parallel=parallel,
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_ids=self.device_ids,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type[TextRerankerWorker]:
|
||||
return TextCrossEncoderWorker
|
||||
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[float]:
|
||||
return (float(elem) for elem in output.model_output)
|
||||
|
||||
|
||||
class TextCrossEncoderWorker(TextRerankerWorker):
|
||||
def init_embedding(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
**kwargs: Any,
|
||||
) -> OnnxTextCrossEncoder:
|
||||
return OnnxTextCrossEncoder(
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
threads=1,
|
||||
**kwargs,
|
||||
)
|
||||
@@ -0,0 +1,169 @@
|
||||
import os
|
||||
from multiprocessing import get_all_start_methods
|
||||
from pathlib import Path
|
||||
from typing import Any, Iterable, Optional, Sequence, Type
|
||||
|
||||
import numpy as np
|
||||
from tokenizers import Encoding
|
||||
|
||||
from fastembed.common.onnx_model import (
|
||||
EmbeddingWorker,
|
||||
OnnxModel,
|
||||
OnnxOutputContext,
|
||||
OnnxProvider,
|
||||
)
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.common.preprocessor_utils import load_tokenizer
|
||||
from fastembed.common.utils import iter_batch
|
||||
from fastembed.parallel_processor import ParallelWorkerPool
|
||||
|
||||
|
||||
class OnnxCrossEncoderModel(OnnxModel[float]):
|
||||
ONNX_OUTPUT_NAMES: Optional[list[str]] = None
|
||||
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type["TextRerankerWorker"]:
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
def _load_onnx_model(
|
||||
self,
|
||||
model_dir: Path,
|
||||
model_file: str,
|
||||
threads: Optional[int],
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
cuda: bool = False,
|
||||
device_id: Optional[int] = None,
|
||||
) -> None:
|
||||
super()._load_onnx_model(
|
||||
model_dir=model_dir,
|
||||
model_file=model_file,
|
||||
threads=threads,
|
||||
providers=providers,
|
||||
cuda=cuda,
|
||||
device_id=device_id,
|
||||
)
|
||||
self.tokenizer, _ = load_tokenizer(model_dir=model_dir)
|
||||
assert self.tokenizer is not None
|
||||
|
||||
def tokenize(self, pairs: list[tuple[str, str]], **_: Any) -> list[Encoding]:
|
||||
return self.tokenizer.encode_batch(pairs) # type: ignore[union-attr]
|
||||
|
||||
def _build_onnx_input(self, tokenized_input: list[Encoding]) -> dict[str, NumpyArray]:
|
||||
input_names: set[str] = {node.name for node in self.model.get_inputs()} # type: ignore[union-attr]
|
||||
inputs: dict[str, NumpyArray] = {
|
||||
"input_ids": np.array([enc.ids for enc in tokenized_input], dtype=np.int64),
|
||||
}
|
||||
if "token_type_ids" in input_names:
|
||||
inputs["token_type_ids"] = np.array(
|
||||
[enc.type_ids for enc in tokenized_input], dtype=np.int64
|
||||
)
|
||||
if "attention_mask" in input_names:
|
||||
inputs["attention_mask"] = np.array(
|
||||
[enc.attention_mask for enc in tokenized_input], dtype=np.int64
|
||||
)
|
||||
return inputs
|
||||
|
||||
def onnx_embed(self, query: str, documents: list[str], **kwargs: Any) -> OnnxOutputContext:
|
||||
pairs = [(query, doc) for doc in documents]
|
||||
return self.onnx_embed_pairs(pairs, **kwargs)
|
||||
|
||||
def onnx_embed_pairs(self, pairs: list[tuple[str, str]], **kwargs: Any) -> OnnxOutputContext:
|
||||
tokenized_input = self.tokenize(pairs, **kwargs)
|
||||
inputs = self._build_onnx_input(tokenized_input)
|
||||
onnx_input = self._preprocess_onnx_input(inputs, **kwargs)
|
||||
outputs = self.model.run(self.ONNX_OUTPUT_NAMES, onnx_input) # type: ignore[union-attr]
|
||||
relevant_output = outputs[0]
|
||||
scores: NumpyArray = relevant_output[:, 0]
|
||||
return OnnxOutputContext(model_output=scores)
|
||||
|
||||
def _rerank_documents(
|
||||
self, query: str, documents: Iterable[str], batch_size: int, **kwargs: Any
|
||||
) -> Iterable[float]:
|
||||
if not hasattr(self, "model") or self.model is None:
|
||||
self.load_onnx_model()
|
||||
for batch in iter_batch(documents, batch_size):
|
||||
yield from self._post_process_onnx_output(self.onnx_embed(query, batch, **kwargs))
|
||||
|
||||
def _rerank_pairs(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
pairs: Iterable[tuple[str, str]],
|
||||
batch_size: int,
|
||||
parallel: Optional[int] = None,
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
cuda: bool = False,
|
||||
device_ids: Optional[list[int]] = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[float]:
|
||||
is_small = False
|
||||
|
||||
if isinstance(pairs, tuple):
|
||||
pairs = [pairs]
|
||||
is_small = True
|
||||
|
||||
if isinstance(pairs, list):
|
||||
if len(pairs) < batch_size:
|
||||
is_small = True
|
||||
|
||||
if parallel is None or is_small:
|
||||
if not hasattr(self, "model") or self.model is None:
|
||||
self.load_onnx_model()
|
||||
for batch in iter_batch(pairs, batch_size):
|
||||
yield from self._post_process_onnx_output(self.onnx_embed_pairs(batch, **kwargs))
|
||||
else:
|
||||
if parallel == 0:
|
||||
parallel = os.cpu_count()
|
||||
|
||||
start_method = "forkserver" if "forkserver" in get_all_start_methods() else "spawn"
|
||||
params = {
|
||||
"model_name": model_name,
|
||||
"cache_dir": cache_dir,
|
||||
"providers": providers,
|
||||
**kwargs,
|
||||
}
|
||||
|
||||
pool = ParallelWorkerPool(
|
||||
num_workers=parallel or 1,
|
||||
worker=self._get_worker_class(),
|
||||
cuda=cuda,
|
||||
device_ids=device_ids,
|
||||
start_method=start_method,
|
||||
)
|
||||
for batch in pool.ordered_map(iter_batch(pairs, batch_size), **params):
|
||||
yield from self._post_process_onnx_output(batch) # type: ignore
|
||||
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[float]:
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
def _preprocess_onnx_input(
|
||||
self, onnx_input: dict[str, NumpyArray], **kwargs: Any
|
||||
) -> dict[str, NumpyArray]:
|
||||
"""
|
||||
Preprocess the onnx input.
|
||||
"""
|
||||
return onnx_input
|
||||
|
||||
|
||||
class TextRerankerWorker(EmbeddingWorker[float]):
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
**kwargs: Any,
|
||||
):
|
||||
self.model: OnnxCrossEncoderModel
|
||||
super().__init__(model_name, cache_dir, **kwargs)
|
||||
|
||||
def init_embedding(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
**kwargs: Any,
|
||||
) -> OnnxCrossEncoderModel:
|
||||
raise NotImplementedError()
|
||||
|
||||
def process(self, items: Iterable[tuple[int, Any]]) -> Iterable[tuple[int, Any]]:
|
||||
for idx, batch in items:
|
||||
onnx_output = self.model.onnx_embed_pairs(batch)
|
||||
yield idx, onnx_output
|
||||
@@ -0,0 +1,126 @@
|
||||
from typing import Any, Iterable, Optional, Sequence, Type
|
||||
from dataclasses import asdict
|
||||
|
||||
from fastembed.common import OnnxProvider
|
||||
from fastembed.rerank.cross_encoder.onnx_text_cross_encoder import OnnxTextCrossEncoder
|
||||
from fastembed.rerank.cross_encoder.text_cross_encoder_base import TextCrossEncoderBase
|
||||
from fastembed.common.model_description import BaseModelDescription
|
||||
|
||||
|
||||
class TextCrossEncoder(TextCrossEncoderBase):
|
||||
CROSS_ENCODER_REGISTRY: list[Type[TextCrossEncoderBase]] = [
|
||||
OnnxTextCrossEncoder,
|
||||
]
|
||||
|
||||
@classmethod
|
||||
def list_supported_models(cls) -> list[dict[str, Any]]:
|
||||
"""Lists the supported models.
|
||||
|
||||
Returns:
|
||||
list[BaseModelDescription]: A list of dictionaries containing the model information.
|
||||
|
||||
Example:
|
||||
```
|
||||
[
|
||||
{
|
||||
"model": "Xenova/ms-marco-MiniLM-L-6-v2",
|
||||
"size_in_GB": 0.08,
|
||||
"sources": {
|
||||
"hf": "Xenova/ms-marco-MiniLM-L-6-v2",
|
||||
},
|
||||
"model_file": "onnx/model.onnx",
|
||||
"description": "MiniLM-L-6-v2 model optimized for re-ranking tasks.",
|
||||
"license": "apache-2.0",
|
||||
}
|
||||
]
|
||||
```
|
||||
"""
|
||||
return [asdict(model) for model in cls._list_supported_models()]
|
||||
|
||||
@classmethod
|
||||
def _list_supported_models(cls) -> list[BaseModelDescription]:
|
||||
result: list[BaseModelDescription] = []
|
||||
for encoder in cls.CROSS_ENCODER_REGISTRY:
|
||||
result.extend(encoder._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,
|
||||
cuda: bool = False,
|
||||
device_ids: Optional[list[int]] = None,
|
||||
lazy_load: bool = False,
|
||||
**kwargs: Any,
|
||||
):
|
||||
super().__init__(model_name, cache_dir, threads, **kwargs)
|
||||
|
||||
for CROSS_ENCODER_TYPE in self.CROSS_ENCODER_REGISTRY:
|
||||
supported_models = CROSS_ENCODER_TYPE._list_supported_models()
|
||||
if any(model_name.lower() == model.model.lower() for model in supported_models):
|
||||
self.model = CROSS_ENCODER_TYPE(
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
threads=threads,
|
||||
providers=providers,
|
||||
cuda=cuda,
|
||||
device_ids=device_ids,
|
||||
lazy_load=lazy_load,
|
||||
**kwargs,
|
||||
)
|
||||
return
|
||||
|
||||
raise ValueError(
|
||||
f"Model {model_name} is not supported in TextCrossEncoder."
|
||||
"Please check the supported models using `TextCrossEncoder.list_supported_models()`"
|
||||
)
|
||||
|
||||
def rerank(
|
||||
self, query: str, documents: Iterable[str], batch_size: int = 64, **kwargs: Any
|
||||
) -> Iterable[float]:
|
||||
"""Rerank a list of documents based on a query.
|
||||
|
||||
Args:
|
||||
query: Query to rerank the documents against
|
||||
documents: Iterator of documents to rerank
|
||||
batch_size: Batch size for reranking
|
||||
|
||||
Returns:
|
||||
Iterable of scores for each document
|
||||
"""
|
||||
yield from self.model.rerank(query, documents, batch_size=batch_size, **kwargs)
|
||||
|
||||
def rerank_pairs(
|
||||
self,
|
||||
pairs: Iterable[tuple[str, str]],
|
||||
batch_size: int = 64,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[float]:
|
||||
"""
|
||||
Rerank a list of query-document pairs.
|
||||
|
||||
Args:
|
||||
pairs (Iterable[tuple[str, str]]): An iterable of tuples, where each tuple contains a query and a document
|
||||
to be scored together.
|
||||
batch_size (int, optional): The number of query-document pairs to process in a single batch. Defaults to 64.
|
||||
parallel (Optional[int], optional): The number of parallel processes to use for reranking.
|
||||
If None, parallelization is disabled. Defaults to None.
|
||||
**kwargs (Any): Additional arguments to pass to the underlying reranking model.
|
||||
|
||||
Returns:
|
||||
Iterable[float]: An iterable of scores corresponding to each query-document pair in the input.
|
||||
Higher scores indicate a stronger match between the query and the document.
|
||||
|
||||
Example:
|
||||
>>> encoder = TextCrossEncoder("Xenova/ms-marco-MiniLM-L-6-v2")
|
||||
>>> pairs = [("What is AI?", "Artificial intelligence is ..."), ("What is ML?", "Machine learning is ...")]
|
||||
>>> scores = list(encoder.rerank_pairs(pairs))
|
||||
>>> print(list(map(lambda x: round(x, 2), scores)))
|
||||
[-1.24, -10.6]
|
||||
"""
|
||||
yield from self.model.rerank_pairs(
|
||||
pairs, batch_size=batch_size, parallel=parallel, **kwargs
|
||||
)
|
||||
@@ -0,0 +1,59 @@
|
||||
from typing import Any, Iterable, Optional
|
||||
|
||||
from fastembed.common.model_description import BaseModelDescription
|
||||
from fastembed.common.model_management import ModelManagement
|
||||
|
||||
|
||||
class TextCrossEncoderBase(ModelManagement[BaseModelDescription]):
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
self.model_name = model_name
|
||||
self.cache_dir = cache_dir
|
||||
self.threads = threads
|
||||
self._local_files_only = kwargs.pop("local_files_only", False)
|
||||
|
||||
def rerank(
|
||||
self,
|
||||
query: str,
|
||||
documents: Iterable[str],
|
||||
batch_size: int = 64,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[float]:
|
||||
"""Rerank a list of documents given a query.
|
||||
|
||||
Args:
|
||||
query (str): The query to rerank the documents.
|
||||
documents (Iterable[str]): The list of texts to rerank.
|
||||
batch_size (int): The batch size to use for reranking.
|
||||
**kwargs: Additional keyword argument to pass to the rerank method.
|
||||
|
||||
Yields:
|
||||
Iterable[float]: The scores of the reranked the documents.
|
||||
"""
|
||||
raise NotImplementedError("This method should be overridden by subclasses")
|
||||
|
||||
def rerank_pairs(
|
||||
self,
|
||||
pairs: Iterable[tuple[str, str]],
|
||||
batch_size: int = 64,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[float]:
|
||||
"""Rerank query-document pairs.
|
||||
Args:
|
||||
pairs (Iterable[tuple[str, str]]): Query-document pairs to rerank
|
||||
batch_size (int): The batch size to use for reranking.
|
||||
parallel: parallel:
|
||||
If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.
|
||||
If 0, use all available cores.
|
||||
If None, don't use data-parallel processing, use default onnxruntime threading instead.
|
||||
**kwargs: Additional keyword argument to pass to the rerank method.
|
||||
Yields:
|
||||
Iterable[float]: Scores for each individual pair
|
||||
"""
|
||||
raise NotImplementedError("This method should be overridden by subclasses")
|
||||
+127
-56
@@ -1,38 +1,71 @@
|
||||
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
|
||||
from typing import Any, Iterable, Optional, 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 py_rust_stemmers import SnowballStemmer
|
||||
from fastembed.common.utils import (
|
||||
define_cache_dir,
|
||||
iter_batch,
|
||||
get_all_punctuation,
|
||||
remove_non_alphanumeric,
|
||||
)
|
||||
from fastembed.parallel_processor import ParallelWorkerPool, Worker
|
||||
from fastembed.sparse.sparse_embedding_base import (
|
||||
SparseEmbedding,
|
||||
SparseTextEmbeddingBase,
|
||||
)
|
||||
from fastembed.sparse.utils.tokenizer import WordTokenizer
|
||||
from fastembed.sparse.utils.tokenizer import SimpleTokenizer
|
||||
from fastembed.common.model_description import SparseModelDescription, ModelSource
|
||||
|
||||
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"],
|
||||
},
|
||||
supported_languages = [
|
||||
"arabic",
|
||||
"azerbaijani",
|
||||
"basque",
|
||||
"bengali",
|
||||
"catalan",
|
||||
"chinese",
|
||||
"danish",
|
||||
"dutch",
|
||||
"english",
|
||||
"finnish",
|
||||
"french",
|
||||
"german",
|
||||
"greek",
|
||||
"hebrew",
|
||||
"hinglish",
|
||||
"hungarian",
|
||||
"indonesian",
|
||||
"italian",
|
||||
"kazakh",
|
||||
"nepali",
|
||||
"norwegian",
|
||||
"portuguese",
|
||||
"romanian",
|
||||
"russian",
|
||||
"slovene",
|
||||
"spanish",
|
||||
"swedish",
|
||||
"tajik",
|
||||
"turkish",
|
||||
]
|
||||
|
||||
MODEL_TO_LANGUAGE = {
|
||||
"Qdrant/bm25": "english",
|
||||
}
|
||||
supported_bm25_models: list[SparseModelDescription] = [
|
||||
SparseModelDescription(
|
||||
model="Qdrant/bm25",
|
||||
vocab_size=0,
|
||||
description="BM25 as sparse embeddings meant to be used with Qdrant",
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.01,
|
||||
sources=ModelSource(hf="Qdrant/bm25"),
|
||||
additional_files=[f"{lang}.txt" for lang in supported_languages],
|
||||
requires_idf=True,
|
||||
model_file="mock.file",
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
class Bm25(SparseTextEmbeddingBase):
|
||||
@@ -59,6 +92,8 @@ class Bm25(SparseTextEmbeddingBase):
|
||||
b (float, optional): The b parameter in the BM25 formula. Defines the importance of the document length.
|
||||
Defaults to 0.75.
|
||||
avg_len (float, optional): The average length of the documents in the corpus. Defaults to 256.0.
|
||||
language (str): Specifies the language for the stemmer.
|
||||
disable_stemmer (bool): Disable the stemmer.
|
||||
Raises:
|
||||
ValueError: If the model_name is not in the format <org>/<model> e.g. BAAI/bge-base-en.
|
||||
"""
|
||||
@@ -70,38 +105,58 @@ class Bm25(SparseTextEmbeddingBase):
|
||||
k: float = 1.2,
|
||||
b: float = 0.75,
|
||||
avg_len: float = 256.0,
|
||||
**kwargs,
|
||||
language: str = "english",
|
||||
token_max_length: int = 40,
|
||||
disable_stemmer: bool = False,
|
||||
specific_model_path: Optional[str] = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
super().__init__(model_name, cache_dir, **kwargs)
|
||||
|
||||
if language not in supported_languages:
|
||||
raise ValueError(f"{language} language is not supported")
|
||||
else:
|
||||
self.language = language
|
||||
|
||||
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)
|
||||
self.cache_dir = str(define_cache_dir(cache_dir))
|
||||
|
||||
model_dir = self.download_model(
|
||||
model_description, self.cache_dir, local_files_only=self._local_files_only
|
||||
self._model_dir = self.download_model(
|
||||
model_description,
|
||||
self.cache_dir,
|
||||
local_files_only=self._local_files_only,
|
||||
specific_model_path=specific_model_path,
|
||||
)
|
||||
|
||||
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
|
||||
self.token_max_length = token_max_length
|
||||
self.punctuation = set(get_all_punctuation())
|
||||
self.disable_stemmer = disable_stemmer
|
||||
|
||||
if disable_stemmer:
|
||||
self.stopwords: set[str] = set()
|
||||
self.stemmer = None
|
||||
else:
|
||||
self.stopwords = set(self._load_stopwords(self._model_dir, self.language))
|
||||
self.stemmer = SnowballStemmer(language)
|
||||
|
||||
self.tokenizer = SimpleTokenizer
|
||||
|
||||
@classmethod
|
||||
def list_supported_models(cls) -> List[Dict[str, Any]]:
|
||||
def _list_supported_models(cls) -> list[SparseModelDescription]:
|
||||
"""Lists the supported models.
|
||||
|
||||
Returns:
|
||||
List[Dict[str, Any]]: A list of dictionaries containing the model information.
|
||||
list[SparseModelDescription]: A list of SparseModelDescription objects containing the model information.
|
||||
"""
|
||||
return supported_bm25_models
|
||||
|
||||
@classmethod
|
||||
def _load_stopwords(cls, model_dir: Path) -> List[str]:
|
||||
stopwords_path = model_dir / "stopwords.txt"
|
||||
def _load_stopwords(cls, model_dir: Path, language: str) -> list[str]:
|
||||
stopwords_path = model_dir / f"{language}.txt"
|
||||
if not stopwords_path.exists():
|
||||
return []
|
||||
|
||||
@@ -126,13 +181,13 @@ class Bm25(SparseTextEmbeddingBase):
|
||||
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:
|
||||
if parallel == 0:
|
||||
parallel = os.cpu_count()
|
||||
|
||||
start_method = "forkserver" if "forkserver" in get_all_start_methods() else "spawn"
|
||||
params = {
|
||||
"model_name": model_name,
|
||||
@@ -140,20 +195,25 @@ class Bm25(SparseTextEmbeddingBase):
|
||||
"k": self.k,
|
||||
"b": self.b,
|
||||
"avg_len": self.avg_len,
|
||||
"language": self.language,
|
||||
"token_max_length": self.token_max_length,
|
||||
"disable_stemmer": self.disable_stemmer,
|
||||
}
|
||||
pool = ParallelWorkerPool(
|
||||
parallel, self._get_worker_class(), start_method=start_method
|
||||
num_workers=parallel or 1,
|
||||
worker=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
|
||||
yield record # type: ignore
|
||||
|
||||
def embed(
|
||||
self,
|
||||
documents: Union[str, Iterable[str]],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
"""
|
||||
Encode a list of documents into list of embeddings.
|
||||
@@ -178,16 +238,21 @@ class Bm25(SparseTextEmbeddingBase):
|
||||
parallel=parallel,
|
||||
)
|
||||
|
||||
def _stem(self, tokens: List[str]) -> List[str]:
|
||||
stemmed_tokens = []
|
||||
def _stem(self, tokens: list[str]) -> list[str]:
|
||||
stemmed_tokens: list[str] = []
|
||||
for token in tokens:
|
||||
lower_token = token.lower()
|
||||
|
||||
if token in self.punctuation:
|
||||
continue
|
||||
|
||||
if token in self.stopwords:
|
||||
if lower_token in self.stopwords:
|
||||
continue
|
||||
|
||||
stemmed_token = self.stemmer.stemWord(token)
|
||||
if len(token) > self.token_max_length:
|
||||
continue
|
||||
|
||||
stemmed_token = self.stemmer.stem_word(lower_token) if self.stemmer else lower_token
|
||||
|
||||
if stemmed_token:
|
||||
stemmed_tokens.append(stemmed_token)
|
||||
@@ -195,17 +260,18 @@ class Bm25(SparseTextEmbeddingBase):
|
||||
|
||||
def raw_embed(
|
||||
self,
|
||||
documents: List[str],
|
||||
) -> List[SparseEmbedding]:
|
||||
embeddings = []
|
||||
documents: list[str],
|
||||
) -> list[SparseEmbedding]:
|
||||
embeddings: list[SparseEmbedding] = []
|
||||
for document in documents:
|
||||
document = remove_non_alphanumeric(document)
|
||||
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]:
|
||||
def _term_frequency(self, tokens: list[str]) -> dict[int, float]:
|
||||
"""Calculate the term frequency part of the BM25 formula.
|
||||
|
||||
(
|
||||
@@ -215,13 +281,13 @@ class Bm25(SparseTextEmbeddingBase):
|
||||
)
|
||||
|
||||
Args:
|
||||
tokens (List[str]): The list of tokens in the document.
|
||||
tokens (list[str]): The list of tokens in the document.
|
||||
|
||||
Returns:
|
||||
Dict[int, float]: The token_id to term frequency mapping.
|
||||
dict[int, float]: The token_id to term frequency mapping.
|
||||
"""
|
||||
tf_map = {}
|
||||
counter = defaultdict(int)
|
||||
tf_map: dict[int, float] = {}
|
||||
counter: defaultdict[str, int] = defaultdict(int)
|
||||
for stemmed_token in tokens:
|
||||
counter[stemmed_token] += 1
|
||||
|
||||
@@ -239,7 +305,9 @@ class Bm25(SparseTextEmbeddingBase):
|
||||
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]:
|
||||
def query_embed(
|
||||
self, query: Union[str, Iterable[str]], **kwargs: Any
|
||||
) -> 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.
|
||||
"""
|
||||
@@ -247,11 +315,12 @@ class Bm25(SparseTextEmbeddingBase):
|
||||
query = [query]
|
||||
|
||||
for text in query:
|
||||
text = remove_non_alphanumeric(text)
|
||||
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,
|
||||
list(set(self.compute_token_id(token) for token in stemmed_tokens)),
|
||||
dtype=np.int32,
|
||||
)
|
||||
values = np.ones_like(token_ids)
|
||||
yield SparseEmbedding(indices=token_ids, values=values)
|
||||
@@ -266,7 +335,7 @@ class Bm25Worker(Worker):
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
**kwargs,
|
||||
**kwargs: Any,
|
||||
):
|
||||
self.model = self.init_embedding(model_name, cache_dir, **kwargs)
|
||||
|
||||
@@ -274,11 +343,13 @@ class Bm25Worker(Worker):
|
||||
def start(cls, model_name: str, cache_dir: str, **kwargs: Any) -> "Bm25Worker":
|
||||
return cls(model_name=model_name, cache_dir=cache_dir, **kwargs)
|
||||
|
||||
def process(self, items: Iterable[Tuple[int, Any]]) -> Iterable[Tuple[int, Any]]:
|
||||
def process(
|
||||
self, items: Iterable[tuple[int, Any]]
|
||||
) -> Iterable[tuple[int, list[SparseEmbedding]]]:
|
||||
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:
|
||||
def init_embedding(model_name: str, cache_dir: str, **kwargs: Any) -> Bm25:
|
||||
return Bm25(model_name=model_name, cache_dir=cache_dir, **kwargs)
|
||||
|
||||
+118
-64
@@ -1,11 +1,11 @@
|
||||
import math
|
||||
import string
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Iterable, List, Optional, Sequence, Tuple, Type, Union
|
||||
from typing import Any, Iterable, Optional, Sequence, Type, Union
|
||||
|
||||
import mmh3
|
||||
import numpy as np
|
||||
from snowballstemmer import stemmer as get_stemmer
|
||||
from py_rust_stemmers import SnowballStemmer
|
||||
|
||||
from fastembed.common import OnnxProvider
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
@@ -15,19 +15,20 @@ from fastembed.sparse.sparse_embedding_base import (
|
||||
SparseTextEmbeddingBase,
|
||||
)
|
||||
from fastembed.text.onnx_text_model import OnnxTextModel, TextEmbeddingWorker
|
||||
from fastembed.common.model_description import SparseModelDescription, ModelSource
|
||||
|
||||
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"],
|
||||
},
|
||||
supported_bm42_models: list[SparseModelDescription] = [
|
||||
SparseModelDescription(
|
||||
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",
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.09,
|
||||
sources=ModelSource(hf="Qdrant/all_miniLM_L6_v2_with_attentions"),
|
||||
model_file="model.onnx",
|
||||
additional_files=["stopwords.txt"],
|
||||
requires_idf=True,
|
||||
),
|
||||
]
|
||||
|
||||
MODEL_TO_LANGUAGE = {
|
||||
@@ -60,7 +61,12 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
threads: Optional[int] = None,
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
alpha: float = 0.5,
|
||||
**kwargs,
|
||||
cuda: bool = False,
|
||||
device_ids: Optional[list[int]] = None,
|
||||
lazy_load: bool = False,
|
||||
device_id: Optional[int] = None,
|
||||
specific_model_path: Optional[str] = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
@@ -73,72 +79,105 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
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.
|
||||
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
|
||||
Defaults to False.
|
||||
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
|
||||
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
|
||||
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
|
||||
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
|
||||
device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
|
||||
specific_model_path (Optional[str], optional): The specific path to the onnx model dir if it should be imported from somewhere else
|
||||
|
||||
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)
|
||||
self.providers = providers
|
||||
self.lazy_load = lazy_load
|
||||
|
||||
model_description = self._get_model_description(model_name)
|
||||
self.cache_dir = define_cache_dir(cache_dir)
|
||||
# List of device ids, that can be used for data parallel processing in workers
|
||||
self.device_ids = device_ids
|
||||
self.cuda = cuda
|
||||
|
||||
model_dir = self.download_model(
|
||||
model_description, self.cache_dir, local_files_only=self._local_files_only
|
||||
# This device_id will be used if we need to load model in current process
|
||||
self.device_id: Optional[int] = None
|
||||
if device_id is not None:
|
||||
self.device_id = device_id
|
||||
elif self.device_ids is not None:
|
||||
self.device_id = self.device_ids[0]
|
||||
|
||||
self.model_description = self._get_model_description(model_name)
|
||||
self.cache_dir = str(define_cache_dir(cache_dir))
|
||||
|
||||
self._model_dir = self.download_model(
|
||||
self.model_description,
|
||||
self.cache_dir,
|
||||
local_files_only=self._local_files_only,
|
||||
specific_model_path=specific_model_path,
|
||||
)
|
||||
|
||||
self.load_onnx_model(
|
||||
model_dir=model_dir,
|
||||
model_file=model_description["model_file"],
|
||||
threads=threads,
|
||||
providers=providers,
|
||||
)
|
||||
self.invert_vocab: dict[int, str] = {}
|
||||
|
||||
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.special_tokens: set[str] = set()
|
||||
self.special_tokens_ids: set[int] = set()
|
||||
self.punctuation = set(string.punctuation)
|
||||
self.stopwords = set(self._load_stopwords(model_dir))
|
||||
self.stemmer = get_stemmer(MODEL_TO_LANGUAGE[model_name])
|
||||
self.stopwords = set(self._load_stopwords(self._model_dir))
|
||||
self.stemmer = SnowballStemmer(MODEL_TO_LANGUAGE[model_name])
|
||||
self.alpha = alpha
|
||||
|
||||
def _filter_pair_tokens(self, tokens: List[Tuple[str, Any]]) -> List[Tuple[str, Any]]:
|
||||
result = []
|
||||
if not self.lazy_load:
|
||||
self.load_onnx_model()
|
||||
|
||||
def load_onnx_model(self) -> None:
|
||||
self._load_onnx_model(
|
||||
model_dir=self._model_dir,
|
||||
model_file=self.model_description.model_file,
|
||||
threads=self.threads,
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_id=self.device_id,
|
||||
)
|
||||
|
||||
for token, idx in self.tokenizer.get_vocab().items(): # type: ignore[union-attr]
|
||||
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.stopwords = set(self._load_stopwords(self._model_dir))
|
||||
|
||||
def _filter_pair_tokens(self, tokens: list[tuple[str, Any]]) -> list[tuple[str, Any]]:
|
||||
result: list[tuple[str, Any]] = []
|
||||
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 = []
|
||||
def _stem_pair_tokens(self, tokens: list[tuple[str, Any]]) -> list[tuple[str, Any]]:
|
||||
result: list[tuple[str, Any]] = []
|
||||
for token, value in tokens:
|
||||
processed_token = self.stemmer.stemWord(token)
|
||||
processed_token = self.stemmer.stem_word(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 = []
|
||||
cls, tokens: list[tuple[str, list[int]]], weights: list[float]
|
||||
) -> list[tuple[str, float]]:
|
||||
result: list[tuple[str, float]] = []
|
||||
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 = []
|
||||
self, bpe_tokens: Iterable[tuple[int, str]]
|
||||
) -> list[tuple[str, list[int]]]:
|
||||
result: list[tuple[str, list[int]]] = []
|
||||
acc: str = ""
|
||||
acc_idx: list[int] = []
|
||||
|
||||
continuing_subword_prefix = self.tokenizer.model.continuing_subword_prefix
|
||||
continuing_subword_prefix = self.tokenizer.model.continuing_subword_prefix # type: ignore[union-attr]
|
||||
continuing_subword_prefix_len = len(continuing_subword_prefix)
|
||||
|
||||
for idx, token in bpe_tokens:
|
||||
@@ -160,13 +199,13 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
|
||||
return result
|
||||
|
||||
def _rescore_vector(self, vector: Dict[str, float]) -> Dict[int, float]:
|
||||
def _rescore_vector(self, vector: dict[str, float]) -> dict[int, float]:
|
||||
"""
|
||||
Orders all tokens in the vector by their importance and generates a new score based on the importance order.
|
||||
So that the scoring doesn't depend on absolute values assigned by the model, but on the relative importance.
|
||||
"""
|
||||
|
||||
new_vector = {}
|
||||
new_vector: dict[int, float] = {}
|
||||
|
||||
for token, value in vector.items():
|
||||
token_id = abs(mmh3.hash(token))
|
||||
@@ -179,7 +218,10 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
return new_vector
|
||||
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[SparseEmbedding]:
|
||||
token_ids_batch = output.input_ids
|
||||
if output.input_ids is None:
|
||||
raise ValueError("input_ids must be provided for document post-processing")
|
||||
|
||||
token_ids_batch = output.input_ids.astype(int)
|
||||
|
||||
# attention_value shape: (batch_size, num_heads, num_tokens, num_tokens)
|
||||
pooled_attention = np.mean(output.model_output[:, :, 0], axis=1) * output.attention_mask
|
||||
@@ -198,7 +240,7 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
|
||||
weighted = self._aggregate_weights(stemmed, attention_value)
|
||||
|
||||
max_token_weight = {}
|
||||
max_token_weight: dict[str, float] = {}
|
||||
|
||||
for token, weight in weighted:
|
||||
max_token_weight[token] = max(max_token_weight.get(token, 0), weight)
|
||||
@@ -208,16 +250,16 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
yield SparseEmbedding.from_dict(rescored)
|
||||
|
||||
@classmethod
|
||||
def list_supported_models(cls) -> List[Dict[str, Any]]:
|
||||
def _list_supported_models(cls) -> list[SparseModelDescription]:
|
||||
"""Lists the supported models.
|
||||
|
||||
Returns:
|
||||
List[Dict[str, Any]]: A list of dictionaries containing the model information.
|
||||
list[SparseModelDescription]: A list of SparseModelDescription objects containing the model information.
|
||||
"""
|
||||
return supported_bm42_models
|
||||
|
||||
@classmethod
|
||||
def _load_stopwords(cls, model_dir: Path) -> List[str]:
|
||||
def _load_stopwords(cls, model_dir: Path) -> list[str]:
|
||||
stopwords_path = model_dir / "stopwords.txt"
|
||||
if not stopwords_path.exists():
|
||||
return []
|
||||
@@ -230,7 +272,7 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
documents: Union[str, Iterable[str]],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
"""
|
||||
Encode a list of documents into list of embeddings.
|
||||
@@ -253,18 +295,23 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
documents=documents,
|
||||
batch_size=batch_size,
|
||||
parallel=parallel,
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_ids=self.device_ids,
|
||||
alpha=self.alpha,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _query_rehash(cls, tokens: Iterable[str]) -> Dict[int, float]:
|
||||
result = {}
|
||||
def _query_rehash(cls, tokens: Iterable[str]) -> dict[int, float]:
|
||||
result: dict[int, float] = {}
|
||||
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]:
|
||||
def query_embed(
|
||||
self, query: Union[str, Iterable[str]], **kwargs: Any
|
||||
) -> 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.
|
||||
@@ -273,8 +320,11 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
if isinstance(query, str):
|
||||
query = [query]
|
||||
|
||||
if not hasattr(self, "model") or self.model is None:
|
||||
self.load_onnx_model()
|
||||
|
||||
for text in query:
|
||||
encoded = self.tokenizer.encode(text)
|
||||
encoded = self.tokenizer.encode(text) # type: ignore[union-attr]
|
||||
document_tokens_with_ids = enumerate(encoded.tokens)
|
||||
reconstructed = self._reconstruct_bpe(document_tokens_with_ids)
|
||||
filtered = self._filter_pair_tokens(reconstructed)
|
||||
@@ -283,10 +333,14 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
yield SparseEmbedding.from_dict(self._query_rehash(token for token, _ in stemmed))
|
||||
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type[TextEmbeddingWorker]:
|
||||
def _get_worker_class(cls) -> Type[TextEmbeddingWorker[SparseEmbedding]]:
|
||||
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)
|
||||
class Bm42TextEmbeddingWorker(TextEmbeddingWorker[SparseEmbedding]):
|
||||
def init_embedding(self, model_name: str, cache_dir: str, **kwargs: Any) -> Bm42:
|
||||
return Bm42(
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@@ -1,38 +1,43 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Dict, Iterable, Optional, Union
|
||||
from typing import Iterable, Optional, Union, Any
|
||||
|
||||
import numpy as np
|
||||
from numpy.typing import NDArray
|
||||
|
||||
from fastembed.common.model_description import SparseModelDescription
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.common.model_management import ModelManagement
|
||||
|
||||
|
||||
@dataclass
|
||||
class SparseEmbedding:
|
||||
values: np.ndarray
|
||||
indices: np.ndarray
|
||||
values: NumpyArray
|
||||
indices: Union[NDArray[np.int64], NDArray[np.int32]]
|
||||
|
||||
def as_object(self) -> Dict[str, np.ndarray]:
|
||||
def as_object(self) -> dict[str, NumpyArray]:
|
||||
return {
|
||||
"values": self.values,
|
||||
"indices": self.indices,
|
||||
}
|
||||
|
||||
def as_dict(self) -> Dict[int, float]:
|
||||
return {i: v for i, v in zip(self.indices, self.values)}
|
||||
def as_dict(self) -> dict[int, float]:
|
||||
return {int(i): float(v) for i, v in zip(self.indices, self.values)} # type: ignore
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, data: Dict[int, float]) -> "SparseEmbedding":
|
||||
def from_dict(cls, data: dict[int, float]) -> "SparseEmbedding":
|
||||
if len(data) == 0:
|
||||
return cls(values=np.array([]), indices=np.array([]))
|
||||
indices, values = zip(*data.items())
|
||||
return cls(values=np.array(values), indices=np.array(indices))
|
||||
|
||||
|
||||
class SparseTextEmbeddingBase(ModelManagement):
|
||||
class SparseTextEmbeddingBase(ModelManagement[SparseModelDescription]):
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
**kwargs,
|
||||
**kwargs: Any,
|
||||
):
|
||||
self.model_name = model_name
|
||||
self.cache_dir = cache_dir
|
||||
@@ -44,13 +49,11 @@ class SparseTextEmbeddingBase(ModelManagement):
|
||||
documents: Union[str, Iterable[str]],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
raise NotImplementedError()
|
||||
|
||||
def passage_embed(
|
||||
self, texts: Iterable[str], **kwargs
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
def passage_embed(self, texts: Iterable[str], **kwargs: Any) -> Iterable[SparseEmbedding]:
|
||||
"""
|
||||
Embeds a list of text passages into a list of embeddings.
|
||||
|
||||
@@ -66,7 +69,7 @@ class SparseTextEmbeddingBase(ModelManagement):
|
||||
yield from self.embed(texts, **kwargs)
|
||||
|
||||
def query_embed(
|
||||
self, query: Union[str, Iterable[str]], **kwargs
|
||||
self, query: Union[str, Iterable[str]], **kwargs: Any
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
"""
|
||||
Embeds queries
|
||||
@@ -81,5 +84,5 @@ class SparseTextEmbeddingBase(ModelManagement):
|
||||
# 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):
|
||||
else:
|
||||
yield from self.embed(query, **kwargs)
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
from typing import Any, Dict, Iterable, List, Optional, Sequence, Type, Union
|
||||
from typing import Any, Iterable, Optional, Sequence, Type, Union
|
||||
from dataclasses import asdict
|
||||
|
||||
from fastembed.common import OnnxProvider
|
||||
from fastembed.sparse.bm25 import Bm25
|
||||
@@ -8,18 +9,20 @@ from fastembed.sparse.sparse_embedding_base import (
|
||||
SparseTextEmbeddingBase,
|
||||
)
|
||||
from fastembed.sparse.splade_pp import SpladePP
|
||||
import warnings
|
||||
from fastembed.common.model_description import SparseModelDescription
|
||||
|
||||
|
||||
class SparseTextEmbedding(SparseTextEmbeddingBase):
|
||||
EMBEDDINGS_REGISTRY: List[Type[SparseTextEmbeddingBase]] = [SpladePP, Bm42, Bm25]
|
||||
EMBEDDINGS_REGISTRY: list[Type[SparseTextEmbeddingBase]] = [SpladePP, Bm42, Bm25]
|
||||
|
||||
@classmethod
|
||||
def list_supported_models(cls) -> List[Dict[str, Any]]:
|
||||
def list_supported_models(cls) -> list[dict[str, Any]]:
|
||||
"""
|
||||
Lists the supported models.
|
||||
|
||||
Returns:
|
||||
List[Dict[str, Any]]: A list of dictionaries containing the model information.
|
||||
list[dict[str, Any]]: A list of dictionaries containing the model information.
|
||||
|
||||
Example:
|
||||
```
|
||||
@@ -28,6 +31,7 @@ class SparseTextEmbedding(SparseTextEmbeddingBase):
|
||||
"model": "prithvida/SPLADE_PP_en_v1",
|
||||
"vocab_size": 30522,
|
||||
"description": "Independent Implementation of SPLADE++ Model for English",
|
||||
"license": "apache-2.0",
|
||||
"size_in_GB": 0.532,
|
||||
"sources": {
|
||||
"hf": "qdrant/SPLADE_PP_en_v1",
|
||||
@@ -36,9 +40,13 @@ class SparseTextEmbedding(SparseTextEmbeddingBase):
|
||||
]
|
||||
```
|
||||
"""
|
||||
result = []
|
||||
return [asdict(model) for model in cls._list_supported_models()]
|
||||
|
||||
@classmethod
|
||||
def _list_supported_models(cls) -> list[SparseModelDescription]:
|
||||
result: list[SparseModelDescription] = []
|
||||
for embedding in cls.EMBEDDINGS_REGISTRY:
|
||||
result.extend(embedding.list_supported_models())
|
||||
result.extend(embedding._list_supported_models())
|
||||
return result
|
||||
|
||||
def __init__(
|
||||
@@ -47,21 +55,32 @@ class SparseTextEmbedding(SparseTextEmbeddingBase):
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
**kwargs,
|
||||
cuda: bool = False,
|
||||
device_ids: Optional[list[int]] = None,
|
||||
lazy_load: bool = False,
|
||||
**kwargs: Any,
|
||||
):
|
||||
super().__init__(model_name, cache_dir, threads, **kwargs)
|
||||
if model_name == "prithvida/Splade_PP_en_v1":
|
||||
warnings.warn(
|
||||
"The right spelling is prithivida/Splade_PP_en_v1. "
|
||||
"Support of this name will be removed soon, please fix the model_name",
|
||||
DeprecationWarning,
|
||||
stacklevel=2,
|
||||
)
|
||||
model_name = "prithivida/Splade_PP_en_v1"
|
||||
|
||||
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
|
||||
):
|
||||
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,
|
||||
cuda=cuda,
|
||||
device_ids=device_ids,
|
||||
lazy_load=lazy_load,
|
||||
**kwargs,
|
||||
)
|
||||
return
|
||||
@@ -76,7 +95,7 @@ class SparseTextEmbedding(SparseTextEmbeddingBase):
|
||||
documents: Union[str, Iterable[str]],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
"""
|
||||
Encode a list of documents into list of embeddings.
|
||||
@@ -96,7 +115,7 @@ class SparseTextEmbedding(SparseTextEmbeddingBase):
|
||||
yield from self.model.embed(documents, batch_size, parallel, **kwargs)
|
||||
|
||||
def query_embed(
|
||||
self, query: Union[str, Iterable[str]], **kwargs
|
||||
self, query: Union[str, Iterable[str]], **kwargs: Any
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
"""
|
||||
Embeds queries
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
from typing import Any, Dict, Iterable, List, Optional, Sequence, Type, Union
|
||||
from typing import Any, Iterable, Optional, Sequence, Type, Union
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastembed.common import OnnxProvider
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.common.utils import define_cache_dir
|
||||
@@ -10,33 +9,35 @@ from fastembed.sparse.sparse_embedding_base import (
|
||||
SparseTextEmbeddingBase,
|
||||
)
|
||||
from fastembed.text.onnx_text_model import OnnxTextModel, TextEmbeddingWorker
|
||||
from fastembed.common.model_description import SparseModelDescription, ModelSource
|
||||
|
||||
supported_splade_models = [
|
||||
{
|
||||
"model": "prithvida/Splade_PP_en_v1",
|
||||
"vocab_size": 30522,
|
||||
"description": "Misspelled version of the model. Retained for backward compatibility. Independent Implementation of SPLADE++ Model for English",
|
||||
"size_in_GB": 0.532,
|
||||
"sources": {
|
||||
"hf": "Qdrant/SPLADE_PP_en_v1",
|
||||
},
|
||||
"model_file": "model.onnx",
|
||||
},
|
||||
{
|
||||
"model": "prithivida/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",
|
||||
},
|
||||
"model_file": "model.onnx",
|
||||
},
|
||||
supported_splade_models: list[SparseModelDescription] = [
|
||||
SparseModelDescription(
|
||||
model="prithivida/Splade_PP_en_v1",
|
||||
vocab_size=30522,
|
||||
description="Independent Implementation of SPLADE++ Model for English.",
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.532,
|
||||
sources=ModelSource(hf="Qdrant/SPLADE_PP_en_v1"),
|
||||
model_file="model.onnx",
|
||||
),
|
||||
SparseModelDescription(
|
||||
model="prithvida/Splade_PP_en_v1",
|
||||
vocab_size=30522,
|
||||
description="Independent Implementation of SPLADE++ Model for English.",
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.532,
|
||||
sources=ModelSource(hf="Qdrant/SPLADE_PP_en_v1"),
|
||||
model_file="model.onnx",
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[SparseEmbedding]:
|
||||
if output.attention_mask is None:
|
||||
raise ValueError("attention_mask must be provided for document post-processing")
|
||||
|
||||
relu_log = np.log(1 + np.maximum(output.model_output, 0))
|
||||
|
||||
weighted_log = relu_log * np.expand_dims(output.attention_mask, axis=-1)
|
||||
@@ -51,11 +52,11 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
yield SparseEmbedding(values=scores, indices=indices)
|
||||
|
||||
@classmethod
|
||||
def list_supported_models(cls) -> List[Dict[str, Any]]:
|
||||
def _list_supported_models(cls) -> list[SparseModelDescription]:
|
||||
"""Lists the supported models.
|
||||
|
||||
Returns:
|
||||
List[Dict[str, Any]]: A list of dictionaries containing the model information.
|
||||
list[SparseModelDescription]: A list of SparseModelDescription objects containing the model information.
|
||||
"""
|
||||
return supported_splade_models
|
||||
|
||||
@@ -65,7 +66,12 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
**kwargs,
|
||||
cuda: bool = False,
|
||||
device_ids: Optional[list[int]] = None,
|
||||
lazy_load: bool = False,
|
||||
device_id: Optional[int] = None,
|
||||
specific_model_path: Optional[str] = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
@@ -74,25 +80,56 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
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 list of onnxruntime providers to use.
|
||||
Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
|
||||
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
|
||||
Defaults to False.
|
||||
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
|
||||
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
|
||||
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
|
||||
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
|
||||
device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
|
||||
specific_model_path (Optional[str], optional): The specific path to the onnx model dir if it should be imported from somewhere else
|
||||
|
||||
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)
|
||||
self.providers = providers
|
||||
self.lazy_load = lazy_load
|
||||
|
||||
model_description = self._get_model_description(model_name)
|
||||
self.cache_dir = define_cache_dir(cache_dir)
|
||||
# List of device ids, that can be used for data parallel processing in workers
|
||||
self.device_ids = device_ids
|
||||
self.cuda = cuda
|
||||
|
||||
model_dir = self.download_model(
|
||||
model_description, self.cache_dir, local_files_only=self._local_files_only
|
||||
# This device_id will be used if we need to load model in current process
|
||||
self.device_id: Optional[int] = None
|
||||
if device_id is not None:
|
||||
self.device_id = device_id
|
||||
elif self.device_ids is not None:
|
||||
self.device_id = self.device_ids[0]
|
||||
|
||||
self.model_description = self._get_model_description(model_name)
|
||||
self.cache_dir = str(define_cache_dir(cache_dir))
|
||||
|
||||
self._model_dir = self.download_model(
|
||||
self.model_description,
|
||||
self.cache_dir,
|
||||
local_files_only=self._local_files_only,
|
||||
specific_model_path=specific_model_path,
|
||||
)
|
||||
|
||||
self.load_onnx_model(
|
||||
model_dir=model_dir,
|
||||
model_file=model_description["model_file"],
|
||||
threads=threads,
|
||||
providers=providers,
|
||||
if not self.lazy_load:
|
||||
self.load_onnx_model()
|
||||
|
||||
def load_onnx_model(self) -> None:
|
||||
self._load_onnx_model(
|
||||
model_dir=self._model_dir,
|
||||
model_file=self.model_description.model_file,
|
||||
threads=self.threads,
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_id=self.device_id,
|
||||
)
|
||||
|
||||
def embed(
|
||||
@@ -100,7 +137,7 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
documents: Union[str, Iterable[str]],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
"""
|
||||
Encode a list of documents into list of embeddings.
|
||||
@@ -123,13 +160,22 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
documents=documents,
|
||||
batch_size=batch_size,
|
||||
parallel=parallel,
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_ids=self.device_ids,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type[TextEmbeddingWorker]:
|
||||
def _get_worker_class(cls) -> Type[TextEmbeddingWorker[SparseEmbedding]]:
|
||||
return SpladePPEmbeddingWorker
|
||||
|
||||
|
||||
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)
|
||||
class SpladePPEmbeddingWorker(TextEmbeddingWorker[SparseEmbedding]):
|
||||
def init_embedding(self, model_name: str, cache_dir: str, **kwargs: Any) -> SpladePP:
|
||||
return SpladePP(
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
threads=1,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@@ -1,7 +1,15 @@
|
||||
# This code is a modified copy of the `NLTKWordTokenizer` class from `NLTK` library.
|
||||
|
||||
import re
|
||||
from typing import List
|
||||
|
||||
|
||||
class SimpleTokenizer:
|
||||
@staticmethod
|
||||
def tokenize(text: str) -> list[str]:
|
||||
text = re.sub(r"[^\w]", " ", text.lower())
|
||||
text = re.sub(r"\s+", " ", text)
|
||||
|
||||
return text.strip().split()
|
||||
|
||||
|
||||
class WordTokenizer:
|
||||
@@ -68,12 +76,11 @@ class WordTokenizer:
|
||||
)
|
||||
]
|
||||
CONTRACTIONS3 = [
|
||||
re.compile(pattern)
|
||||
for pattern in (r"(?i) ('t)(?#X)(is)\b", r"(?i) ('t)(?#X)(was)\b")
|
||||
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]:
|
||||
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.'''
|
||||
|
||||
@@ -1,49 +1,54 @@
|
||||
from typing import Any, Dict, Iterable, List, Type
|
||||
|
||||
import numpy as np
|
||||
from typing import Any, Iterable, Type
|
||||
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.text.onnx_embedding import OnnxTextEmbedding, OnnxTextEmbeddingWorker
|
||||
from fastembed.text.onnx_text_model import TextEmbeddingWorker
|
||||
from fastembed.common.model_description import DenseModelDescription, ModelSource
|
||||
|
||||
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",
|
||||
},
|
||||
supported_clip_models: list[DenseModelDescription] = [
|
||||
DenseModelDescription(
|
||||
model="Qdrant/clip-ViT-B-32-text",
|
||||
dim=512,
|
||||
description=(
|
||||
"Text embeddings, Multimodal (text&image), English, 77 input tokens truncation, "
|
||||
"Prefixes for queries/documents: not necessary, 2021 year"
|
||||
),
|
||||
license="mit",
|
||||
size_in_GB=0.25,
|
||||
sources=ModelSource(hf="Qdrant/clip-ViT-B-32-text"),
|
||||
model_file="model.onnx",
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
class CLIPOnnxEmbedding(OnnxTextEmbedding):
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type[TextEmbeddingWorker]:
|
||||
def _get_worker_class(cls) -> Type[OnnxTextEmbeddingWorker]:
|
||||
return CLIPEmbeddingWorker
|
||||
|
||||
@classmethod
|
||||
def list_supported_models(cls) -> List[Dict[str, Any]]:
|
||||
def _list_supported_models(cls) -> list[DenseModelDescription]:
|
||||
"""Lists the supported models.
|
||||
|
||||
Returns:
|
||||
List[Dict[str, Any]]: A list of dictionaries containing the model information.
|
||||
list[DenseModelDescription]: A list of DenseModelDescription objects containing the model information.
|
||||
"""
|
||||
return supported_clip_models
|
||||
|
||||
def _post_process_onnx_output(
|
||||
self, output: OnnxOutputContext
|
||||
) -> Iterable[np.ndarray]:
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[NumpyArray]:
|
||||
return output.model_output
|
||||
|
||||
|
||||
class CLIPEmbeddingWorker(OnnxTextEmbeddingWorker):
|
||||
def init_embedding(
|
||||
self, model_name: str, cache_dir: str, **kwargs
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
**kwargs: Any,
|
||||
) -> OnnxTextEmbedding:
|
||||
return CLIPOnnxEmbedding(
|
||||
model_name=model_name, cache_dir=cache_dir, threads=1, **kwargs
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
threads=1,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@@ -1,64 +0,0 @@
|
||||
from typing import Any, Dict, List, Type
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastembed.text.onnx_embedding import OnnxTextEmbedding, OnnxTextEmbeddingWorker
|
||||
from fastembed.text.onnx_text_model import TextEmbeddingWorker
|
||||
|
||||
supported_multilingual_e5_models = [
|
||||
{
|
||||
"model": "intfloat/multilingual-e5-large",
|
||||
"dim": 1024,
|
||||
"description": "Multilingual model, e5-large. Recommend using this model for non-English languages",
|
||||
"size_in_GB": 2.24,
|
||||
"sources": {
|
||||
"url": "https://storage.googleapis.com/qdrant-fastembed/fast-multilingual-e5-large.tar.gz",
|
||||
"hf": "qdrant/multilingual-e5-large-onnx",
|
||||
},
|
||||
"model_file": "model.onnx",
|
||||
"additional_files": ["model.onnx_data"],
|
||||
},
|
||||
{
|
||||
"model": "sentence-transformers/paraphrase-multilingual-mpnet-base-v2",
|
||||
"dim": 768,
|
||||
"description": "Sentence-transformers model for tasks like clustering or semantic search",
|
||||
"size_in_GB": 1.00,
|
||||
"sources": {
|
||||
"hf": "xenova/paraphrase-multilingual-mpnet-base-v2",
|
||||
},
|
||||
"model_file": "onnx/model.onnx",
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
class E5OnnxEmbedding(OnnxTextEmbedding):
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type["TextEmbeddingWorker"]:
|
||||
return E5OnnxEmbeddingWorker
|
||||
|
||||
@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_multilingual_e5_models
|
||||
|
||||
def _preprocess_onnx_input(
|
||||
self, onnx_input: Dict[str, np.ndarray], **kwargs
|
||||
) -> Dict[str, np.ndarray]:
|
||||
"""
|
||||
Preprocess the onnx input.
|
||||
"""
|
||||
onnx_input.pop("token_type_ids", None)
|
||||
return onnx_input
|
||||
|
||||
|
||||
class E5OnnxEmbeddingWorker(OnnxTextEmbeddingWorker):
|
||||
def init_embedding(
|
||||
self, model_name: str, cache_dir: str, **kwargs
|
||||
) -> E5OnnxEmbedding:
|
||||
return E5OnnxEmbedding(
|
||||
model_name=model_name, cache_dir=cache_dir, threads=1, **kwargs
|
||||
)
|
||||
@@ -1,76 +0,0 @@
|
||||
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_jina_models = [
|
||||
{
|
||||
"model": "jinaai/jina-embeddings-v2-base-en",
|
||||
"dim": 768,
|
||||
"description": "English embedding model supporting 8192 sequence length",
|
||||
"size_in_GB": 0.52,
|
||||
"sources": {"hf": "xenova/jina-embeddings-v2-base-en"},
|
||||
"model_file": "onnx/model.onnx",
|
||||
},
|
||||
{
|
||||
"model": "jinaai/jina-embeddings-v2-small-en",
|
||||
"dim": 512,
|
||||
"description": "English embedding model supporting 8192 sequence length",
|
||||
"size_in_GB": 0.12,
|
||||
"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[TextEmbeddingWorker]:
|
||||
return JinaEmbeddingWorker
|
||||
|
||||
@classmethod
|
||||
def mean_pooling(cls, model_output, attention_mask) -> np.ndarray:
|
||||
token_embeddings = model_output
|
||||
input_mask_expanded = (np.expand_dims(attention_mask, axis=-1)).astype(float)
|
||||
|
||||
sum_embeddings = np.sum(token_embeddings * input_mask_expanded, axis=1)
|
||||
mask_sum = np.clip(np.sum(input_mask_expanded, axis=1), a_min=1e-9, a_max=None)
|
||||
|
||||
return sum_embeddings / mask_sum
|
||||
|
||||
@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_jina_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 JinaEmbeddingWorker(OnnxTextEmbeddingWorker):
|
||||
def init_embedding(
|
||||
self, model_name: str, cache_dir: str, **kwargs
|
||||
) -> OnnxTextEmbedding:
|
||||
return JinaOnnxEmbedding(
|
||||
model_name=model_name, cache_dir=cache_dir, threads=1, **kwargs
|
||||
)
|
||||
@@ -1,58 +0,0 @@
|
||||
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)
|
||||
@@ -0,0 +1,100 @@
|
||||
from enum import Enum
|
||||
from typing import Any, Type, Iterable, Union, Optional
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.text.pooled_normalized_embedding import PooledNormalizedEmbedding
|
||||
from fastembed.text.onnx_embedding import OnnxTextEmbeddingWorker
|
||||
from fastembed.common.model_description import DenseModelDescription, ModelSource
|
||||
|
||||
supported_multitask_models: list[DenseModelDescription] = [
|
||||
DenseModelDescription(
|
||||
model="jinaai/jina-embeddings-v3",
|
||||
dim=1024,
|
||||
tasks={
|
||||
"retrieval.query": 0,
|
||||
"retrieval.passage": 1,
|
||||
"separation": 2,
|
||||
"classification": 3,
|
||||
"text-matching": 4,
|
||||
},
|
||||
description=(
|
||||
"Multi-task unimodal (text) embedding model, multi-lingual (~100), "
|
||||
"1024 tokens truncation, and 8192 sequence length. Prefixes for queries/documents: not necessary, 2024 year."
|
||||
),
|
||||
license="cc-by-nc-4.0",
|
||||
size_in_GB=2.29,
|
||||
sources=ModelSource(hf="jinaai/jina-embeddings-v3"),
|
||||
model_file="onnx/model.onnx",
|
||||
additional_files=["onnx/model.onnx_data"],
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
class Task(int, Enum):
|
||||
RETRIEVAL_QUERY = 0
|
||||
RETRIEVAL_PASSAGE = 1
|
||||
SEPARATION = 2
|
||||
CLASSIFICATION = 3
|
||||
TEXT_MATCHING = 4
|
||||
|
||||
|
||||
class JinaEmbeddingV3(PooledNormalizedEmbedding):
|
||||
PASSAGE_TASK = Task.RETRIEVAL_PASSAGE
|
||||
QUERY_TASK = Task.RETRIEVAL_QUERY
|
||||
|
||||
def __init__(self, *args: Any, **kwargs: Any):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.current_task_id: Union[Task, int] = self.PASSAGE_TASK
|
||||
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type[OnnxTextEmbeddingWorker]:
|
||||
return JinaEmbeddingV3Worker
|
||||
|
||||
@classmethod
|
||||
def _list_supported_models(cls) -> list[DenseModelDescription]:
|
||||
return supported_multitask_models
|
||||
|
||||
def _preprocess_onnx_input(
|
||||
self, onnx_input: dict[str, NumpyArray], **kwargs: Any
|
||||
) -> dict[str, NumpyArray]:
|
||||
onnx_input["task_id"] = np.array(self.current_task_id, dtype=np.int64)
|
||||
return onnx_input
|
||||
|
||||
def embed(
|
||||
self,
|
||||
documents: Union[str, Iterable[str]],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
task_id: int = PASSAGE_TASK,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
self.current_task_id = task_id
|
||||
kwargs["task_id"] = task_id
|
||||
yield from super().embed(documents, batch_size, parallel, **kwargs)
|
||||
|
||||
def query_embed(self, query: Union[str, Iterable[str]], **kwargs: Any) -> Iterable[NumpyArray]:
|
||||
self.current_task_id = self.QUERY_TASK
|
||||
yield from super().embed(query, **kwargs)
|
||||
|
||||
def passage_embed(self, texts: Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:
|
||||
self.current_task_id = self.PASSAGE_TASK
|
||||
yield from super().embed(texts, **kwargs)
|
||||
|
||||
|
||||
class JinaEmbeddingV3Worker(OnnxTextEmbeddingWorker):
|
||||
def init_embedding(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
**kwargs: Any,
|
||||
) -> JinaEmbeddingV3:
|
||||
model = JinaEmbeddingV3(
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
threads=1,
|
||||
**kwargs,
|
||||
)
|
||||
model.current_task_id = kwargs["task_id"]
|
||||
return model
|
||||
+259
-200
@@ -1,198 +1,207 @@
|
||||
from typing import Any, Dict, Iterable, List, Optional, Sequence, Type, Union
|
||||
from typing import Any, Iterable, Optional, Sequence, Type, Union
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastembed.common import OnnxProvider
|
||||
from fastembed.common.types import NumpyArray, 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
|
||||
from fastembed.common.model_description import DenseModelDescription, ModelSource
|
||||
|
||||
supported_onnx_models = [
|
||||
{
|
||||
"model": "BAAI/bge-base-en",
|
||||
"dim": 768,
|
||||
"description": "Base English model",
|
||||
"size_in_GB": 0.42,
|
||||
"sources": {
|
||||
"url": "https://storage.googleapis.com/qdrant-fastembed/fast-bge-base-en.tar.gz",
|
||||
},
|
||||
"model_file": "model_optimized.onnx",
|
||||
},
|
||||
{
|
||||
"model": "BAAI/bge-base-en-v1.5",
|
||||
"dim": 768,
|
||||
"description": "Base English model, v1.5",
|
||||
"size_in_GB": 0.21,
|
||||
"sources": {
|
||||
"url": "https://storage.googleapis.com/qdrant-fastembed/fast-bge-base-en-v1.5.tar.gz",
|
||||
"hf": "qdrant/bge-base-en-v1.5-onnx-q",
|
||||
},
|
||||
"model_file": "model_optimized.onnx",
|
||||
},
|
||||
{
|
||||
"model": "BAAI/bge-large-en-v1.5",
|
||||
"dim": 1024,
|
||||
"description": "Large English model, v1.5",
|
||||
"size_in_GB": 1.20,
|
||||
"sources": {
|
||||
"hf": "qdrant/bge-large-en-v1.5-onnx",
|
||||
},
|
||||
"model_file": "model.onnx",
|
||||
},
|
||||
{
|
||||
"model": "BAAI/bge-small-en",
|
||||
"dim": 384,
|
||||
"description": "Fast English model",
|
||||
"size_in_GB": 0.13,
|
||||
"sources": {
|
||||
"url": "https://storage.googleapis.com/qdrant-fastembed/BAAI-bge-small-en.tar.gz",
|
||||
},
|
||||
"model_file": "model_optimized.onnx",
|
||||
},
|
||||
{
|
||||
"model": "BAAI/bge-small-en-v1.5",
|
||||
"dim": 384,
|
||||
"description": "Fast and Default English model",
|
||||
"size_in_GB": 0.067,
|
||||
"sources": {
|
||||
"hf": "qdrant/bge-small-en-v1.5-onnx-q",
|
||||
},
|
||||
"model_file": "model_optimized.onnx",
|
||||
},
|
||||
{
|
||||
"model": "BAAI/bge-small-zh-v1.5",
|
||||
"dim": 512,
|
||||
"description": "Fast and recommended Chinese model",
|
||||
"size_in_GB": 0.09,
|
||||
"sources": {
|
||||
"url": "https://storage.googleapis.com/qdrant-fastembed/fast-bge-small-zh-v1.5.tar.gz",
|
||||
},
|
||||
"model_file": "model_optimized.onnx",
|
||||
},
|
||||
{
|
||||
"model": "sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2",
|
||||
"dim": 384,
|
||||
"description": "Sentence Transformer model, paraphrase-multilingual-MiniLM-L12-v2",
|
||||
"size_in_GB": 0.22,
|
||||
"sources": {
|
||||
"hf": "qdrant/paraphrase-multilingual-MiniLM-L12-v2-onnx-Q",
|
||||
},
|
||||
"model_file": "model_optimized.onnx",
|
||||
},
|
||||
{
|
||||
"model": "nomic-ai/nomic-embed-text-v1",
|
||||
"dim": 768,
|
||||
"description": "8192 context length english model",
|
||||
"size_in_GB": 0.52,
|
||||
"sources": {
|
||||
"hf": "nomic-ai/nomic-embed-text-v1",
|
||||
},
|
||||
"model_file": "onnx/model.onnx",
|
||||
},
|
||||
{
|
||||
"model": "nomic-ai/nomic-embed-text-v1.5",
|
||||
"dim": 768,
|
||||
"description": "8192 context length english model",
|
||||
"size_in_GB": 0.52,
|
||||
"sources": {
|
||||
"hf": "nomic-ai/nomic-embed-text-v1.5",
|
||||
},
|
||||
"model_file": "onnx/model.onnx",
|
||||
},
|
||||
{
|
||||
"model": "nomic-ai/nomic-embed-text-v1.5-Q",
|
||||
"dim": 768,
|
||||
"description": "Quantized 8192 context length english model",
|
||||
"size_in_GB": 0.13,
|
||||
"sources": {
|
||||
"hf": "nomic-ai/nomic-embed-text-v1.5",
|
||||
},
|
||||
"model_file": "onnx/model_quantized.onnx",
|
||||
},
|
||||
{
|
||||
"model": "thenlper/gte-large",
|
||||
"dim": 1024,
|
||||
"description": "Large general text embeddings model",
|
||||
"size_in_GB": 1.20,
|
||||
"sources": {
|
||||
"hf": "qdrant/gte-large-onnx",
|
||||
},
|
||||
"model_file": "model.onnx",
|
||||
},
|
||||
{
|
||||
"model": "mixedbread-ai/mxbai-embed-large-v1",
|
||||
"dim": 1024,
|
||||
"description": "MixedBread Base sentence embedding model, does well on MTEB",
|
||||
"size_in_GB": 0.64,
|
||||
"sources": {
|
||||
"hf": "mixedbread-ai/mxbai-embed-large-v1",
|
||||
},
|
||||
"model_file": "onnx/model.onnx",
|
||||
},
|
||||
{
|
||||
"model": "snowflake/snowflake-arctic-embed-xs",
|
||||
"dim": 384,
|
||||
"description": "Based on all-MiniLM-L6-v2 model with only 22m parameters, ideal for latency/TCO budgets.",
|
||||
"size_in_GB": 0.09,
|
||||
"sources": {
|
||||
"hf": "snowflake/snowflake-arctic-embed-xs",
|
||||
},
|
||||
"model_file": "onnx/model.onnx",
|
||||
},
|
||||
{
|
||||
"model": "snowflake/snowflake-arctic-embed-s",
|
||||
"dim": 384,
|
||||
"description": "Based on infloat/e5-small-unsupervised, does not trade off retrieval accuracy for its small size.",
|
||||
"size_in_GB": 0.13,
|
||||
"sources": {
|
||||
"hf": "snowflake/snowflake-arctic-embed-s",
|
||||
},
|
||||
"model_file": "onnx/model.onnx",
|
||||
},
|
||||
{
|
||||
"model": "snowflake/snowflake-arctic-embed-m",
|
||||
"dim": 768,
|
||||
"description": "Based on intfloat/e5-base-unsupervised model, provides the best retrieval without slowing down inference.",
|
||||
"size_in_GB": 0.43,
|
||||
"sources": {
|
||||
"hf": "Snowflake/snowflake-arctic-embed-m",
|
||||
},
|
||||
"model_file": "onnx/model.onnx",
|
||||
},
|
||||
{
|
||||
"model": "snowflake/snowflake-arctic-embed-m-long",
|
||||
"dim": 768,
|
||||
"description": "Based on nomic-ai/nomic-embed-text-v1-unsupervised model, 8192 context-length model",
|
||||
"size_in_GB": 0.54,
|
||||
"sources": {
|
||||
"hf": "snowflake/snowflake-arctic-embed-m-long",
|
||||
},
|
||||
"model_file": "onnx/model.onnx",
|
||||
},
|
||||
{
|
||||
"model": "snowflake/snowflake-arctic-embed-l",
|
||||
"dim": 1024,
|
||||
"description": "Based on intfloat/e5-large-unsupervised, large model for most accurate retrieval.",
|
||||
"size_in_GB": 1.02,
|
||||
"sources": {
|
||||
"hf": "snowflake/snowflake-arctic-embed-l",
|
||||
},
|
||||
"model_file": "onnx/model.onnx",
|
||||
},
|
||||
supported_onnx_models: list[DenseModelDescription] = [
|
||||
DenseModelDescription(
|
||||
model="BAAI/bge-base-en",
|
||||
dim=768,
|
||||
description=(
|
||||
"Text embeddings, Unimodal (text), English, 512 input tokens truncation, "
|
||||
"Prefixes for queries/documents: necessary, 2023 year."
|
||||
),
|
||||
license="mit",
|
||||
size_in_GB=0.42,
|
||||
sources=ModelSource(
|
||||
hf="Qdrant/fast-bge-base-en",
|
||||
url="https://storage.googleapis.com/qdrant-fastembed/fast-bge-base-en.tar.gz",
|
||||
),
|
||||
model_file="model_optimized.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="BAAI/bge-base-en-v1.5",
|
||||
dim=768,
|
||||
description=(
|
||||
"Text embeddings, Unimodal (text), English, 512 input tokens truncation, "
|
||||
"Prefixes for queries/documents: not so necessary, 2023 year."
|
||||
),
|
||||
license="mit",
|
||||
size_in_GB=0.21,
|
||||
sources=ModelSource(
|
||||
hf="qdrant/bge-base-en-v1.5-onnx-q",
|
||||
url="https://storage.googleapis.com/qdrant-fastembed/fast-bge-base-en-v1.5.tar.gz",
|
||||
),
|
||||
model_file="model_optimized.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="BAAI/bge-large-en-v1.5",
|
||||
dim=1024,
|
||||
description=(
|
||||
"Text embeddings, Unimodal (text), English, 512 input tokens truncation, "
|
||||
"Prefixes for queries/documents: not so necessary, 2023 year."
|
||||
),
|
||||
license="mit",
|
||||
size_in_GB=1.20,
|
||||
sources=ModelSource(hf="qdrant/bge-large-en-v1.5-onnx"),
|
||||
model_file="model.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="BAAI/bge-small-en",
|
||||
dim=384,
|
||||
description=(
|
||||
"Text embeddings, Unimodal (text), English, 512 input tokens truncation, "
|
||||
"Prefixes for queries/documents: necessary, 2023 year."
|
||||
),
|
||||
license="mit",
|
||||
size_in_GB=0.13,
|
||||
sources=ModelSource(
|
||||
hf="Qdrant/bge-small-en",
|
||||
url="https://storage.googleapis.com/qdrant-fastembed/BAAI-bge-small-en.tar.gz",
|
||||
),
|
||||
model_file="model_optimized.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="BAAI/bge-small-en-v1.5",
|
||||
dim=384,
|
||||
description=(
|
||||
"Text embeddings, Unimodal (text), English, 512 input tokens truncation, "
|
||||
"Prefixes for queries/documents: not so necessary, 2023 year."
|
||||
),
|
||||
license="mit",
|
||||
size_in_GB=0.067,
|
||||
sources=ModelSource(hf="qdrant/bge-small-en-v1.5-onnx-q"),
|
||||
model_file="model_optimized.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="BAAI/bge-small-zh-v1.5",
|
||||
dim=512,
|
||||
description=(
|
||||
"Text embeddings, Unimodal (text), Chinese, 512 input tokens truncation, "
|
||||
"Prefixes for queries/documents: not so necessary, 2023 year."
|
||||
),
|
||||
license="mit",
|
||||
size_in_GB=0.09,
|
||||
sources=ModelSource(
|
||||
hf="Qdrant/bge-small-zh-v1.5",
|
||||
url="https://storage.googleapis.com/qdrant-fastembed/fast-bge-small-zh-v1.5.tar.gz",
|
||||
),
|
||||
model_file="model_optimized.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="thenlper/gte-large",
|
||||
dim=1024,
|
||||
description=(
|
||||
"Text embeddings, Unimodal (text), English, 512 input tokens truncation, "
|
||||
"Prefixes for queries/documents: not necessary, 2023 year."
|
||||
),
|
||||
license="mit",
|
||||
size_in_GB=1.20,
|
||||
sources=ModelSource(hf="qdrant/gte-large-onnx"),
|
||||
model_file="model.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="mixedbread-ai/mxbai-embed-large-v1",
|
||||
dim=1024,
|
||||
description=(
|
||||
"Text embeddings, Unimodal (text), English, 512 input tokens truncation, "
|
||||
"Prefixes for queries/documents: necessary, 2024 year."
|
||||
),
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.64,
|
||||
sources=ModelSource(hf="mixedbread-ai/mxbai-embed-large-v1"),
|
||||
model_file="onnx/model.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="snowflake/snowflake-arctic-embed-xs",
|
||||
dim=384,
|
||||
description=(
|
||||
"Text embeddings, Unimodal (text), English, 512 input tokens truncation, "
|
||||
"Prefixes for queries/documents: necessary, 2024 year."
|
||||
),
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.09,
|
||||
sources=ModelSource(hf="snowflake/snowflake-arctic-embed-xs"),
|
||||
model_file="onnx/model.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="snowflake/snowflake-arctic-embed-s",
|
||||
dim=384,
|
||||
description=(
|
||||
"Text embeddings, Unimodal (text), English, 512 input tokens truncation, "
|
||||
"Prefixes for queries/documents: necessary, 2024 year."
|
||||
),
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.13,
|
||||
sources=ModelSource(hf="snowflake/snowflake-arctic-embed-s"),
|
||||
model_file="onnx/model.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="snowflake/snowflake-arctic-embed-m",
|
||||
dim=768,
|
||||
description=(
|
||||
"Text embeddings, Unimodal (text), English, 512 input tokens truncation, "
|
||||
"Prefixes for queries/documents: necessary, 2024 year."
|
||||
),
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.43,
|
||||
sources=ModelSource(hf="Snowflake/snowflake-arctic-embed-m"),
|
||||
model_file="onnx/model.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="snowflake/snowflake-arctic-embed-m-long",
|
||||
dim=768,
|
||||
description=(
|
||||
"Text embeddings, Unimodal (text), English, 2048 input tokens truncation, "
|
||||
"Prefixes for queries/documents: necessary, 2024 year."
|
||||
),
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.54,
|
||||
sources=ModelSource(hf="snowflake/snowflake-arctic-embed-m-long"),
|
||||
model_file="onnx/model.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="snowflake/snowflake-arctic-embed-l",
|
||||
dim=1024,
|
||||
description=(
|
||||
"Text embeddings, Unimodal (text), English, 512 input tokens truncation, "
|
||||
"Prefixes for queries/documents: necessary, 2024 year."
|
||||
),
|
||||
license="apache-2.0",
|
||||
size_in_GB=1.02,
|
||||
sources=ModelSource(hf="snowflake/snowflake-arctic-embed-l"),
|
||||
model_file="onnx/model.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="jinaai/jina-clip-v1",
|
||||
dim=768,
|
||||
description=(
|
||||
"Text embeddings, Multimodal (text&image), English, Prefixes for queries/documents: "
|
||||
"not necessary, 2024 year"
|
||||
),
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.55,
|
||||
sources=ModelSource(hf="jinaai/jina-clip-v1"),
|
||||
model_file="onnx/text_model.onnx",
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[np.ndarray]):
|
||||
class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[NumpyArray]):
|
||||
"""Implementation of the Flag Embedding model."""
|
||||
|
||||
@classmethod
|
||||
def list_supported_models(cls) -> List[Dict[str, Any]]:
|
||||
def _list_supported_models(cls) -> list[DenseModelDescription]:
|
||||
"""
|
||||
Lists the supported models.
|
||||
|
||||
Returns:
|
||||
List[Dict[str, Any]]: A list of dictionaries containing the model information.
|
||||
list[DenseModelDescription]: A list of DenseModelDescription objects containing the model information.
|
||||
"""
|
||||
return supported_onnx_models
|
||||
|
||||
@@ -202,7 +211,12 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[np.ndarray]):
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
**kwargs,
|
||||
cuda: bool = False,
|
||||
device_ids: Optional[list[int]] = None,
|
||||
lazy_load: bool = False,
|
||||
device_id: Optional[int] = None,
|
||||
specific_model_path: Optional[str] = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
@@ -211,33 +225,54 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[np.ndarray]):
|
||||
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 list of onnxruntime providers to use.
|
||||
Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
|
||||
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
|
||||
Defaults to False.
|
||||
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
|
||||
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
|
||||
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
|
||||
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
|
||||
device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
|
||||
specific_model_path (Optional[str], optional): The specific path to the onnx model dir if it should be imported from somewhere else
|
||||
|
||||
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)
|
||||
self.providers = providers
|
||||
self.lazy_load = lazy_load
|
||||
|
||||
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
|
||||
# List of device ids, that can be used for data parallel processing in workers
|
||||
self.device_ids = device_ids
|
||||
self.cuda = cuda
|
||||
|
||||
# This device_id will be used if we need to load model in current process
|
||||
self.device_id: Optional[int] = None
|
||||
if device_id is not None:
|
||||
self.device_id = device_id
|
||||
elif self.device_ids is not None:
|
||||
self.device_id = self.device_ids[0]
|
||||
|
||||
self.model_description = self._get_model_description(model_name)
|
||||
self.cache_dir = str(define_cache_dir(cache_dir))
|
||||
self._model_dir = self.download_model(
|
||||
self.model_description,
|
||||
self.cache_dir,
|
||||
local_files_only=self._local_files_only,
|
||||
specific_model_path=specific_model_path,
|
||||
)
|
||||
|
||||
self.load_onnx_model(
|
||||
model_dir=model_dir,
|
||||
model_file=model_description["model_file"],
|
||||
threads=threads,
|
||||
providers=providers,
|
||||
)
|
||||
if not self.lazy_load:
|
||||
self.load_onnx_model()
|
||||
|
||||
def embed(
|
||||
self,
|
||||
documents: Union[str, Iterable[str]],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs,
|
||||
) -> Iterable[np.ndarray]:
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
Encode a list of documents into list of embeddings.
|
||||
We use mean pooling with attention so that the model can handle variable-length inputs.
|
||||
@@ -259,31 +294,55 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[np.ndarray]):
|
||||
documents=documents,
|
||||
batch_size=batch_size,
|
||||
parallel=parallel,
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_ids=self.device_ids,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type["TextEmbeddingWorker"]:
|
||||
def _get_worker_class(cls) -> Type["TextEmbeddingWorker[NumpyArray]"]:
|
||||
return OnnxTextEmbeddingWorker
|
||||
|
||||
def _preprocess_onnx_input(
|
||||
self, onnx_input: Dict[str, np.ndarray], **kwargs
|
||||
) -> Dict[str, np.ndarray]:
|
||||
self, onnx_input: dict[str, NumpyArray], **kwargs: Any
|
||||
) -> dict[str, NumpyArray]:
|
||||
"""
|
||||
Preprocess the onnx input.
|
||||
"""
|
||||
return onnx_input
|
||||
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[np.ndarray]:
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[NumpyArray]:
|
||||
embeddings = output.model_output
|
||||
return normalize(embeddings[:, 0]).astype(np.float32)
|
||||
if embeddings.ndim == 3: # (batch_size, seq_len, embedding_dim)
|
||||
processed_embeddings = embeddings[:, 0]
|
||||
elif embeddings.ndim == 2: # (batch_size, embedding_dim)
|
||||
processed_embeddings = embeddings
|
||||
else:
|
||||
raise ValueError(f"Unsupported embedding shape: {embeddings.shape}")
|
||||
return normalize(processed_embeddings).astype(np.float32)
|
||||
|
||||
def load_onnx_model(self) -> None:
|
||||
self._load_onnx_model(
|
||||
model_dir=self._model_dir,
|
||||
model_file=self.model_description.model_file,
|
||||
threads=self.threads,
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_id=self.device_id,
|
||||
)
|
||||
|
||||
|
||||
class OnnxTextEmbeddingWorker(TextEmbeddingWorker):
|
||||
class OnnxTextEmbeddingWorker(TextEmbeddingWorker[NumpyArray]):
|
||||
def init_embedding(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
**kwargs,
|
||||
**kwargs: Any,
|
||||
) -> OnnxTextEmbedding:
|
||||
return OnnxTextEmbedding(model_name=model_name, cache_dir=cache_dir, threads=1, **kwargs)
|
||||
return OnnxTextEmbedding(
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
threads=1,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@@ -1,12 +1,13 @@
|
||||
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
|
||||
from typing import Any, Iterable, Optional, Sequence, Type, Union
|
||||
|
||||
import numpy as np
|
||||
from tokenizers import Encoding
|
||||
from numpy.typing import NDArray
|
||||
from tokenizers import Encoding, Tokenizer
|
||||
|
||||
from fastembed.common import OnnxProvider
|
||||
from fastembed.common.types import NumpyArray, 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
|
||||
@@ -14,10 +15,10 @@ from fastembed.parallel_processor import ParallelWorkerPool
|
||||
|
||||
|
||||
class OnnxTextModel(OnnxModel[T]):
|
||||
ONNX_OUTPUT_NAMES: Optional[List[str]] = None
|
||||
ONNX_OUTPUT_NAMES: Optional[list[str]] = None
|
||||
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type["TextEmbeddingWorker"]:
|
||||
def _get_worker_class(cls) -> Type["TextEmbeddingWorker[T]"]:
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[T]:
|
||||
@@ -25,45 +26,52 @@ class OnnxTextModel(OnnxModel[T]):
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.tokenizer = None
|
||||
self.special_token_to_id = {}
|
||||
self.tokenizer: Optional[Tokenizer] = None
|
||||
self.special_token_to_id: dict[str, int] = {}
|
||||
|
||||
def _preprocess_onnx_input(
|
||||
self, onnx_input: Dict[str, np.ndarray], **kwargs
|
||||
) -> Dict[str, np.ndarray]:
|
||||
self, onnx_input: dict[str, NumpyArray], **kwargs: Any
|
||||
) -> dict[str, Union[NumpyArray, NDArray[np.int64]]]:
|
||||
"""
|
||||
Preprocess the onnx input.
|
||||
"""
|
||||
return onnx_input
|
||||
|
||||
def load_onnx_model(
|
||||
def _load_onnx_model(
|
||||
self,
|
||||
model_dir: Path,
|
||||
model_file: str,
|
||||
threads: Optional[int],
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
cuda: bool = False,
|
||||
device_id: Optional[int] = None,
|
||||
) -> None:
|
||||
super().load_onnx_model(
|
||||
super()._load_onnx_model(
|
||||
model_dir=model_dir,
|
||||
model_file=model_file,
|
||||
threads=threads,
|
||||
providers=providers,
|
||||
cuda=cuda,
|
||||
device_id=device_id,
|
||||
)
|
||||
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 load_onnx_model(self) -> None:
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
def tokenize(self, documents: list[str], **kwargs: Any) -> list[Encoding]:
|
||||
return self.tokenizer.encode_batch(documents) # type: ignore[union-attr]
|
||||
|
||||
def onnx_embed(
|
||||
self,
|
||||
documents: List[str],
|
||||
**kwargs,
|
||||
documents: list[str],
|
||||
**kwargs: Any,
|
||||
) -> 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_names = {node.name for node in self.model.get_inputs()} # type: ignore[union-attr]
|
||||
onnx_input: dict[str, NumpyArray] = {
|
||||
"input_ids": np.array(input_ids, dtype=np.int64),
|
||||
}
|
||||
if "attention_mask" in input_names:
|
||||
@@ -72,10 +80,9 @@ class OnnxTextModel(OnnxModel[T]):
|
||||
onnx_input["token_type_ids"] = np.array(
|
||||
[np.zeros(len(e), dtype=np.int64) for e in input_ids], dtype=np.int64
|
||||
)
|
||||
|
||||
onnx_input = self._preprocess_onnx_input(onnx_input, **kwargs)
|
||||
|
||||
model_output = self.model.run(self.ONNX_OUTPUT_NAMES, onnx_input)
|
||||
model_output = self.model.run(self.ONNX_OUTPUT_NAMES, onnx_input) # type: ignore[union-attr]
|
||||
return OnnxOutputContext(
|
||||
model_output=model_output[0],
|
||||
attention_mask=onnx_input.get("attention_mask", attention_mask),
|
||||
@@ -89,7 +96,10 @@ class OnnxTextModel(OnnxModel[T]):
|
||||
documents: Union[str, Iterable[str]],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs,
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
cuda: bool = False,
|
||||
device_ids: Optional[list[int]] = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[T]:
|
||||
is_small = False
|
||||
|
||||
@@ -101,26 +111,36 @@ class OnnxTextModel(OnnxModel[T]):
|
||||
if len(documents) < batch_size:
|
||||
is_small = True
|
||||
|
||||
if parallel == 0:
|
||||
parallel = os.cpu_count()
|
||||
|
||||
if parallel is None or is_small:
|
||||
if not hasattr(self, "model") or self.model is None:
|
||||
self.load_onnx_model()
|
||||
for batch in iter_batch(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}
|
||||
if parallel == 0:
|
||||
parallel = os.cpu_count()
|
||||
|
||||
start_method = "forkserver" if "forkserver" in get_all_start_methods() else "spawn"
|
||||
params = {
|
||||
"model_name": model_name,
|
||||
"cache_dir": cache_dir,
|
||||
"providers": providers,
|
||||
**kwargs,
|
||||
}
|
||||
|
||||
pool = ParallelWorkerPool(
|
||||
parallel, self._get_worker_class(), start_method=start_method
|
||||
num_workers=parallel or 1,
|
||||
worker=self._get_worker_class(),
|
||||
cuda=cuda,
|
||||
device_ids=device_ids,
|
||||
start_method=start_method,
|
||||
)
|
||||
for batch in pool.ordered_map(iter_batch(documents, batch_size), **params):
|
||||
yield from self._post_process_onnx_output(batch)
|
||||
yield from self._post_process_onnx_output(batch) # type: ignore
|
||||
|
||||
|
||||
class TextEmbeddingWorker(EmbeddingWorker):
|
||||
def process(self, items: Iterable[Tuple[int, Any]]) -> Iterable[Tuple[int, Any]]:
|
||||
class TextEmbeddingWorker(EmbeddingWorker[T]):
|
||||
def process(self, items: Iterable[tuple[int, Any]]) -> Iterable[tuple[int, OnnxOutputContext]]:
|
||||
for idx, batch in items:
|
||||
onnx_output = self.model.onnx_embed(batch)
|
||||
yield idx, onnx_output
|
||||
|
||||
@@ -0,0 +1,137 @@
|
||||
from typing import Any, Iterable, Type
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.text.onnx_embedding import OnnxTextEmbedding, OnnxTextEmbeddingWorker
|
||||
from fastembed.common.model_description import DenseModelDescription, ModelSource
|
||||
|
||||
supported_pooled_models: list[DenseModelDescription] = [
|
||||
DenseModelDescription(
|
||||
model="nomic-ai/nomic-embed-text-v1.5",
|
||||
dim=768,
|
||||
description=(
|
||||
"Text embeddings, Multimodal (text, image), English, 8192 input tokens truncation, "
|
||||
"Prefixes for queries/documents: necessary, 2024 year."
|
||||
),
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.52,
|
||||
sources=ModelSource(hf="nomic-ai/nomic-embed-text-v1.5"),
|
||||
model_file="onnx/model.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="nomic-ai/nomic-embed-text-v1.5-Q",
|
||||
dim=768,
|
||||
description=(
|
||||
"Text embeddings, Multimodal (text, image), English, 8192 input tokens truncation, "
|
||||
"Prefixes for queries/documents: necessary, 2024 year."
|
||||
),
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.13,
|
||||
sources=ModelSource(hf="nomic-ai/nomic-embed-text-v1.5"),
|
||||
model_file="onnx/model_quantized.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="nomic-ai/nomic-embed-text-v1",
|
||||
dim=768,
|
||||
description=(
|
||||
"Text embeddings, Multimodal (text, image), English, 8192 input tokens truncation, "
|
||||
"Prefixes for queries/documents: necessary, 2024 year."
|
||||
),
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.52,
|
||||
sources=ModelSource(hf="nomic-ai/nomic-embed-text-v1"),
|
||||
model_file="onnx/model.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2",
|
||||
dim=384,
|
||||
description=(
|
||||
"Text embeddings, Unimodal (text), Multilingual (~50 languages), 512 input tokens truncation, "
|
||||
"Prefixes for queries/documents: not necessary, 2019 year."
|
||||
),
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.22,
|
||||
sources=ModelSource(hf="qdrant/paraphrase-multilingual-MiniLM-L12-v2-onnx-Q"),
|
||||
model_file="model_optimized.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="sentence-transformers/paraphrase-multilingual-mpnet-base-v2",
|
||||
dim=768,
|
||||
description=(
|
||||
"Text embeddings, Unimodal (text), Multilingual (~50 languages), 384 input tokens truncation, "
|
||||
"Prefixes for queries/documents: not necessary, 2021 year."
|
||||
),
|
||||
license="apache-2.0",
|
||||
size_in_GB=1.00,
|
||||
sources=ModelSource(hf="xenova/paraphrase-multilingual-mpnet-base-v2"),
|
||||
model_file="onnx/model.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="intfloat/multilingual-e5-large",
|
||||
dim=1024,
|
||||
description=(
|
||||
"Text embeddings, Unimodal (text), Multilingual (~100 languages), 512 input tokens truncation, "
|
||||
"Prefixes for queries/documents: necessary, 2024 year."
|
||||
),
|
||||
license="mit",
|
||||
size_in_GB=2.24,
|
||||
sources=ModelSource(
|
||||
hf="qdrant/multilingual-e5-large-onnx",
|
||||
url="https://storage.googleapis.com/qdrant-fastembed/fast-multilingual-e5-large.tar.gz",
|
||||
),
|
||||
model_file="model.onnx",
|
||||
additional_files=["model.onnx_data"],
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
class PooledEmbedding(OnnxTextEmbedding):
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type[OnnxTextEmbeddingWorker]:
|
||||
return PooledEmbeddingWorker
|
||||
|
||||
@classmethod
|
||||
def mean_pooling(cls, model_output: NumpyArray, attention_mask: NumpyArray) -> NumpyArray:
|
||||
token_embeddings = model_output.astype(np.float32)
|
||||
attention_mask = attention_mask.astype(np.float32)
|
||||
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(np.float32)
|
||||
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[DenseModelDescription]:
|
||||
"""Lists the supported models.
|
||||
|
||||
Returns:
|
||||
list[DenseModelDescription]: A list of DenseModelDescription objects containing the model information.
|
||||
"""
|
||||
return supported_pooled_models
|
||||
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[NumpyArray]:
|
||||
if output.attention_mask is None:
|
||||
raise ValueError("attention_mask must be provided for document post-processing")
|
||||
|
||||
embeddings = output.model_output
|
||||
attn_mask = output.attention_mask
|
||||
return self.mean_pooling(embeddings, attn_mask).astype(np.float32)
|
||||
|
||||
|
||||
class PooledEmbeddingWorker(OnnxTextEmbeddingWorker):
|
||||
def init_embedding(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
**kwargs: Any,
|
||||
) -> OnnxTextEmbedding:
|
||||
return PooledEmbedding(
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
threads=1,
|
||||
**kwargs,
|
||||
)
|
||||
@@ -0,0 +1,150 @@
|
||||
from typing import Any, Iterable, Type
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.common.utils import normalize
|
||||
from fastembed.text.onnx_embedding import OnnxTextEmbedding, OnnxTextEmbeddingWorker
|
||||
from fastembed.text.pooled_embedding import PooledEmbedding
|
||||
from fastembed.common.model_description import DenseModelDescription, ModelSource
|
||||
|
||||
supported_pooled_normalized_models: list[DenseModelDescription] = [
|
||||
DenseModelDescription(
|
||||
model="sentence-transformers/all-MiniLM-L6-v2",
|
||||
dim=384,
|
||||
description=(
|
||||
"Text embeddings, Unimodal (text), English, 256 input tokens truncation, "
|
||||
"Prefixes for queries/documents: not necessary, 2021 year."
|
||||
),
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.09,
|
||||
sources=ModelSource(
|
||||
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",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="jinaai/jina-embeddings-v2-base-en",
|
||||
dim=768,
|
||||
description=(
|
||||
"Text embeddings, Unimodal (text), English, 8192 input tokens truncation, "
|
||||
"Prefixes for queries/documents: not necessary, 2023 year."
|
||||
),
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.52,
|
||||
sources=ModelSource(hf="xenova/jina-embeddings-v2-base-en"),
|
||||
model_file="onnx/model.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="jinaai/jina-embeddings-v2-small-en",
|
||||
dim=512,
|
||||
description=(
|
||||
"Text embeddings, Unimodal (text), English, 8192 input tokens truncation, "
|
||||
"Prefixes for queries/documents: not necessary, 2023 year."
|
||||
),
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.12,
|
||||
sources=ModelSource(hf="xenova/jina-embeddings-v2-small-en"),
|
||||
model_file="onnx/model.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="jinaai/jina-embeddings-v2-base-de",
|
||||
dim=768,
|
||||
description=(
|
||||
"Text embeddings, Unimodal (text), Multilingual (German, English), 8192 input tokens truncation, "
|
||||
"Prefixes for queries/documents: not necessary, 2024 year."
|
||||
),
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.32,
|
||||
sources=ModelSource(hf="jinaai/jina-embeddings-v2-base-de"),
|
||||
model_file="onnx/model_fp16.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="jinaai/jina-embeddings-v2-base-code",
|
||||
dim=768,
|
||||
description=(
|
||||
"Text embeddings, Unimodal (text), Multilingual (English, 30 programming languages), "
|
||||
"8192 input tokens truncation, Prefixes for queries/documents: not necessary, 2024 year."
|
||||
),
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.64,
|
||||
sources=ModelSource(hf="jinaai/jina-embeddings-v2-base-code"),
|
||||
model_file="onnx/model.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="jinaai/jina-embeddings-v2-base-zh",
|
||||
dim=768,
|
||||
description=(
|
||||
"Text embeddings, Unimodal (text), supports mixed Chinese-English input text, "
|
||||
"8192 input tokens truncation, Prefixes for queries/documents: not necessary, 2024 year."
|
||||
),
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.64,
|
||||
sources=ModelSource(hf="jinaai/jina-embeddings-v2-base-zh"),
|
||||
model_file="onnx/model.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="jinaai/jina-embeddings-v2-base-es",
|
||||
dim=768,
|
||||
description=(
|
||||
"Text embeddings, Unimodal (text), supports mixed Spanish-English input text, "
|
||||
"8192 input tokens truncation, Prefixes for queries/documents: not necessary, 2024 year."
|
||||
),
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.64,
|
||||
sources=ModelSource(hf="jinaai/jina-embeddings-v2-base-es"),
|
||||
model_file="onnx/model.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="thenlper/gte-base",
|
||||
dim=768,
|
||||
description=(
|
||||
"General text embeddings, Unimodal (text), supports English only input text, "
|
||||
"512 input tokens truncation, Prefixes for queries/documents: not necessary, 2024 year."
|
||||
),
|
||||
license="mit",
|
||||
size_in_GB=0.44,
|
||||
sources=ModelSource(hf="thenlper/gte-base"),
|
||||
model_file="onnx/model.onnx",
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
class PooledNormalizedEmbedding(PooledEmbedding):
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type[OnnxTextEmbeddingWorker]:
|
||||
return PooledNormalizedEmbeddingWorker
|
||||
|
||||
@classmethod
|
||||
def _list_supported_models(cls) -> list[DenseModelDescription]:
|
||||
"""Lists the supported models.
|
||||
|
||||
Returns:
|
||||
list[DenseModelDescription]: A list of DenseModelDescription objects containing the model information.
|
||||
"""
|
||||
return supported_pooled_normalized_models
|
||||
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[NumpyArray]:
|
||||
if output.attention_mask is None:
|
||||
raise ValueError("attention_mask must be provided for document post-processing")
|
||||
|
||||
embeddings = output.model_output
|
||||
attn_mask = output.attention_mask
|
||||
return normalize(self.mean_pooling(embeddings, attn_mask)).astype(np.float32)
|
||||
|
||||
|
||||
class PooledNormalizedEmbeddingWorker(OnnxTextEmbeddingWorker):
|
||||
def init_embedding(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
**kwargs: Any,
|
||||
) -> OnnxTextEmbedding:
|
||||
return PooledNormalizedEmbedding(
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
threads=1,
|
||||
**kwargs,
|
||||
)
|
||||
@@ -1,52 +1,40 @@
|
||||
from typing import Any, Dict, Iterable, List, Optional, Sequence, Type, Union
|
||||
import warnings
|
||||
from typing import Any, Iterable, Optional, Sequence, Type, Union
|
||||
from dataclasses import asdict
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastembed.common import OnnxProvider
|
||||
from fastembed.common.types import NumpyArray, 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.pooled_normalized_embedding import PooledNormalizedEmbedding
|
||||
from fastembed.text.pooled_embedding import PooledEmbedding
|
||||
from fastembed.text.multitask_embedding import JinaEmbeddingV3
|
||||
from fastembed.text.onnx_embedding import OnnxTextEmbedding
|
||||
from fastembed.text.text_embedding_base import TextEmbeddingBase
|
||||
from fastembed.common.model_description import DenseModelDescription
|
||||
|
||||
|
||||
class TextEmbedding(TextEmbeddingBase):
|
||||
EMBEDDINGS_REGISTRY: List[Type[TextEmbeddingBase]] = [
|
||||
EMBEDDINGS_REGISTRY: list[Type[TextEmbeddingBase]] = [
|
||||
OnnxTextEmbedding,
|
||||
E5OnnxEmbedding,
|
||||
JinaOnnxEmbedding,
|
||||
CLIPOnnxEmbedding,
|
||||
MiniLMOnnxEmbedding,
|
||||
PooledNormalizedEmbedding,
|
||||
PooledEmbedding,
|
||||
JinaEmbeddingV3,
|
||||
]
|
||||
|
||||
@classmethod
|
||||
def list_supported_models(cls) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Lists the supported models.
|
||||
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": "intfloat/multilingual-e5-large",
|
||||
"dim": 1024,
|
||||
"description": "Multilingual model, e5-large. Recommend using this model for non-English languages",
|
||||
"size_in_GB": 2.24,
|
||||
"sources": {
|
||||
"gcp": "https://storage.googleapis.com/qdrant-fastembed/fast-multilingual-e5-large.tar.gz",
|
||||
"hf": "qdrant/multilingual-e5-large-onnx",
|
||||
}
|
||||
}
|
||||
]
|
||||
```
|
||||
list[dict[str, Any]]: A list of dictionaries containing the model information.
|
||||
"""
|
||||
result = []
|
||||
return [asdict(model) for model in cls._list_supported_models()]
|
||||
|
||||
@classmethod
|
||||
def _list_supported_models(cls) -> list[DenseModelDescription]:
|
||||
result: list[DenseModelDescription] = []
|
||||
for embedding in cls.EMBEDDINGS_REGISTRY:
|
||||
result.extend(embedding.list_supported_models())
|
||||
result.extend(embedding._list_supported_models())
|
||||
return result
|
||||
|
||||
def __init__(
|
||||
@@ -55,27 +43,54 @@ class TextEmbedding(TextEmbeddingBase):
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
**kwargs,
|
||||
cuda: bool = False,
|
||||
device_ids: Optional[list[int]] = None,
|
||||
lazy_load: bool = False,
|
||||
**kwargs: Any,
|
||||
):
|
||||
super().__init__(model_name, cache_dir, threads, **kwargs)
|
||||
if model_name == "nomic-ai/nomic-embed-text-v1.5-Q":
|
||||
warnings.warn(
|
||||
"The model 'nomic-ai/nomic-embed-text-v1.5-Q' has been updated on HuggingFace. "
|
||||
"Please review the latest documentation and release notes to ensure compatibility with your workflow. ",
|
||||
UserWarning,
|
||||
stacklevel=2,
|
||||
)
|
||||
if model_name == "sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2":
|
||||
warnings.warn(
|
||||
"The model 'sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2' has been updated to "
|
||||
"include a mean pooling layer. Please ensure your usage aligns with the new functionality. "
|
||||
"Support for the previous version without mean pooling will be removed as of version 0.5.2.",
|
||||
UserWarning,
|
||||
stacklevel=2,
|
||||
)
|
||||
if model_name in {
|
||||
"sentence-transformers/paraphrase-multilingual-mpnet-base-v2",
|
||||
"intfloat/multilingual-e5-large",
|
||||
}:
|
||||
warnings.warn(
|
||||
f"{model_name} has been updated as of fastembed 0.5.2, outputs are now average pooled.",
|
||||
UserWarning,
|
||||
stacklevel=2,
|
||||
)
|
||||
|
||||
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
|
||||
):
|
||||
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,
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
threads=threads,
|
||||
providers=providers,
|
||||
cuda=cuda,
|
||||
device_ids=device_ids,
|
||||
lazy_load=lazy_load,
|
||||
**kwargs,
|
||||
)
|
||||
return
|
||||
|
||||
raise ValueError(
|
||||
f"Model {model_name} is not supported in TextEmbedding."
|
||||
f"Model {model_name} is not supported in TextEmbedding. "
|
||||
"Please check the supported models using `TextEmbedding.list_supported_models()`"
|
||||
)
|
||||
|
||||
@@ -84,8 +99,8 @@ class TextEmbedding(TextEmbeddingBase):
|
||||
documents: Union[str, Iterable[str]],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs,
|
||||
) -> Iterable[np.ndarray]:
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
Encode a list of documents into list of embeddings.
|
||||
We use mean pooling with attention so that the model can handle variable-length inputs.
|
||||
@@ -102,3 +117,30 @@ class TextEmbedding(TextEmbeddingBase):
|
||||
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: Any) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
Embeds queries
|
||||
|
||||
Args:
|
||||
query (Union[str, Iterable[str]]): The query to embed, or an iterable e.g. list of queries.
|
||||
|
||||
Returns:
|
||||
Iterable[NumpyArray]: The embeddings.
|
||||
"""
|
||||
# This is model-specific, so that different models can have specialized implementations
|
||||
yield from self.model.query_embed(query, **kwargs)
|
||||
|
||||
def passage_embed(self, texts: Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
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.model.passage_embed(texts, **kwargs)
|
||||
|
||||
@@ -1,17 +1,17 @@
|
||||
from typing import Iterable, Optional, Union
|
||||
|
||||
import numpy as np
|
||||
from typing import Iterable, Optional, Union, Any
|
||||
|
||||
from fastembed.common.model_description import DenseModelDescription
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.common.model_management import ModelManagement
|
||||
|
||||
|
||||
class TextEmbeddingBase(ModelManagement):
|
||||
class TextEmbeddingBase(ModelManagement[DenseModelDescription]):
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
**kwargs,
|
||||
**kwargs: Any,
|
||||
):
|
||||
self.model_name = model_name
|
||||
self.cache_dir = cache_dir
|
||||
@@ -23,11 +23,11 @@ class TextEmbeddingBase(ModelManagement):
|
||||
documents: Union[str, Iterable[str]],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
**kwargs,
|
||||
) -> Iterable[np.ndarray]:
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
raise NotImplementedError()
|
||||
|
||||
def passage_embed(self, texts: Iterable[str], **kwargs) -> Iterable[np.ndarray]:
|
||||
def passage_embed(self, texts: Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
Embeds a list of text passages into a list of embeddings.
|
||||
|
||||
@@ -36,15 +36,13 @@ class TextEmbeddingBase(ModelManagement):
|
||||
**kwargs: Additional keyword argument to pass to the embed method.
|
||||
|
||||
Yields:
|
||||
Iterable[np.ndarray]: The embeddings.
|
||||
Iterable[NumpyArray]: 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]:
|
||||
def query_embed(self, query: Union[str, Iterable[str]], **kwargs: Any) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
Embeds queries
|
||||
|
||||
@@ -52,11 +50,11 @@ class TextEmbeddingBase(ModelManagement):
|
||||
query (Union[str, Iterable[str]]): The query to embed, or an iterable e.g. list of queries.
|
||||
|
||||
Returns:
|
||||
Iterable[np.ndarray]: The embeddings.
|
||||
Iterable[NumpyArray]: 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):
|
||||
else:
|
||||
yield from self.embed(query, **kwargs)
|
||||
|
||||
+28
-15
@@ -1,6 +1,6 @@
|
||||
[tool.poetry]
|
||||
name = "fastembed"
|
||||
version = "0.3.1"
|
||||
version = "0.5.1"
|
||||
description = "Fast, light, accurate library built for retrieval embedding generation"
|
||||
authors = ["Qdrant Team <info@qdrant.tech>", "NirantK <nirant.bits@gmail.com>"]
|
||||
license = "Apache License"
|
||||
@@ -11,40 +11,53 @@ repository = "https://github.com/qdrant/fastembed"
|
||||
keywords = ["vector", "embedding", "neural", "search", "qdrant", "sentence-transformers"]
|
||||
|
||||
[tool.poetry.dependencies]
|
||||
python = ">=3.8.0,<3.13"
|
||||
onnx = "^1.15.0"
|
||||
onnxruntime = "^1.17.0"
|
||||
python = ">=3.9.0"
|
||||
numpy = [
|
||||
{ version = ">=1.21", python = ">=3.10,<3.12" },
|
||||
{ version = ">=1.26", python = ">=3.12,<3.13" },
|
||||
{ version = ">=2.1.0", python = ">=3.13" },
|
||||
{ version = ">=1.21,<2.1.0", python = "<3.10" },
|
||||
]
|
||||
onnxruntime = [
|
||||
{ version = ">1.20.0", python = ">=3.13" },
|
||||
{ version = ">=1.17.0,<1.20.0", python = "<3.10" },
|
||||
{ version = ">=1.17.0,!=1.20.0", python = ">=3.10,<3.13" },
|
||||
]
|
||||
tqdm = "^4.66"
|
||||
requests = "^2.31"
|
||||
tokenizers = ">=0.15,<1.0"
|
||||
huggingface-hub = ">=0.20,<1.0"
|
||||
loguru = "^0.7.2"
|
||||
numpy = [
|
||||
{ version = ">=1.21, <2", python = "<3.12" },
|
||||
{ version = ">=1.26, <2", python = ">=3.12" }
|
||||
]
|
||||
pillow = "^10.3.0"
|
||||
snowballstemmer = "^2.2.0"
|
||||
PyStemmer = "^2.2.0"
|
||||
mmh3 = "^4.0"
|
||||
pillow = ">=10.3.0,<12.0.0"
|
||||
mmh3 = "^4.1.0"
|
||||
py-rust-stemmers = "^0.1.0"
|
||||
|
||||
[tool.poetry.group.dev.dependencies]
|
||||
[tool.poetry.group.test.dependencies]
|
||||
pytest = "^7.4.2"
|
||||
ruff = ">=0.3.1,<1.0"
|
||||
|
||||
[tool.poetry.group.dev.dependencies]
|
||||
notebook = ">=7.0.2"
|
||||
pre-commit = {version = "^3.6.2", python = ">=3.9,<3.12" }
|
||||
pre-commit = "^3.6.2"
|
||||
onnx = ">=1.15.0"
|
||||
|
||||
[tool.poetry.group.docs.dependencies]
|
||||
mkdocs-material = "^9.5.10"
|
||||
mkdocstrings = "^0.24.0"
|
||||
pillow = "^10.2.0"
|
||||
pillow = ">=10.3.0,<12.0.0"
|
||||
cairosvg = "^2.7.1"
|
||||
mknotebooks = "^0.8.0"
|
||||
|
||||
[tool.poetry.group.types.dependencies]
|
||||
pyright = ">=1.1.293"
|
||||
mypy = "^1.0.0"
|
||||
|
||||
[build-system]
|
||||
requires = ["poetry-core"]
|
||||
build-backend = "poetry.core.masonry.api"
|
||||
|
||||
[tool.pyright]
|
||||
typeCheckingMode = "strict"
|
||||
|
||||
[tool.ruff]
|
||||
line-length = 99
|
||||
|
||||
@@ -0,0 +1,4 @@
|
||||
import os
|
||||
|
||||
# disable DeprecationWarning https://github.com/jupyter/jupyter_core/issues/398
|
||||
os.environ["JUPYTER_PLATFORM_DIRS"] = "1"
|
||||
|
||||
+7
-9
@@ -9,7 +9,7 @@
|
||||
|
||||
# %%
|
||||
import time
|
||||
from typing import Callable, List, Tuple
|
||||
from typing import Callable
|
||||
|
||||
import matplotlib.pyplot as plt
|
||||
import torch.nn.functional as F
|
||||
@@ -23,7 +23,7 @@ from fastembed.embedding import DefaultEmbedding
|
||||
# data is a list of strings, each string is a document.
|
||||
|
||||
# %%
|
||||
documents: List[str] = [
|
||||
documents: list[str] = [
|
||||
"Chandrayaan-3 is India's third lunar mission",
|
||||
"It aimed to land a rover on the Moon's surface - joining the US, China and Russia",
|
||||
"The mission is a follow-up to Chandrayaan-2, which had partial success",
|
||||
@@ -56,7 +56,7 @@ class HF:
|
||||
self.model = AutoModel.from_pretrained(model_id)
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(model_id)
|
||||
|
||||
def embed(self, texts: List[str]):
|
||||
def embed(self, texts: list[str]):
|
||||
encoded_input = self.tokenizer(
|
||||
texts, max_length=512, padding=True, truncation=True, return_tensors="pt"
|
||||
)
|
||||
@@ -88,7 +88,7 @@ embedding_model = DefaultEmbedding()
|
||||
# %%
|
||||
def calculate_time_stats(
|
||||
embed_func: Callable, documents: list, k: int
|
||||
) -> Tuple[float, float, float]:
|
||||
) -> tuple[float, float, float]:
|
||||
times = []
|
||||
for _ in range(k):
|
||||
# Timing the embed_func call
|
||||
@@ -105,16 +105,14 @@ def calculate_time_stats(
|
||||
# %%
|
||||
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],
|
||||
hf_stats: tuple[float, float, float],
|
||||
fst_stats: tuple[float, float, float],
|
||||
documents: list,
|
||||
):
|
||||
# Calculating total characters in documents
|
||||
|
||||
@@ -1,13 +1,15 @@
|
||||
import os
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from fastembed import SparseTextEmbedding
|
||||
from tests.utils import delete_model_cache
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_name", ["Qdrant/bm42-all-minilm-l6-v2-attentions", "Qdrant/bm25"]
|
||||
)
|
||||
def test_attention_embeddings(model_name):
|
||||
@pytest.mark.parametrize("model_name", ["Qdrant/bm42-all-minilm-l6-v2-attentions", "Qdrant/bm25"])
|
||||
def test_attention_embeddings(model_name: str) -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
model = SparseTextEmbedding(model_name=model_name)
|
||||
|
||||
output = list(
|
||||
@@ -36,16 +38,19 @@ def test_attention_embeddings(model_name):
|
||||
"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.",
|
||||
".", # Empty string
|
||||
]
|
||||
|
||||
output = list(model.embed(quotes))
|
||||
|
||||
assert len(output) == len(quotes)
|
||||
|
||||
for result in output:
|
||||
for result in output[:-1]:
|
||||
assert len(result.indices) == len(result.values)
|
||||
assert len(result.indices) > 0
|
||||
|
||||
assert len(output[-1].indices) == 0
|
||||
|
||||
# Test support for unknown languages
|
||||
output = list(
|
||||
model.query_embed(
|
||||
@@ -61,14 +66,17 @@ def test_attention_embeddings(model_name):
|
||||
assert len(result.indices) == len(result.values)
|
||||
assert len(result.indices) == 2
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_name", ["Qdrant/bm42-all-minilm-l6-v2-attentions", "Qdrant/bm25"])
|
||||
def test_parallel_processing(model_name: str) -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
|
||||
@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
|
||||
docs = ["hello world", "attention embedding", "Mangez-vous vraiment des grenouilles?"] * 100
|
||||
embeddings = list(model.embed(docs, batch_size=10, parallel=2))
|
||||
|
||||
embeddings_2 = list(model.embed(docs, batch_size=10, parallel=None))
|
||||
@@ -82,3 +90,70 @@ def test_parallel_processing(model_name):
|
||||
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)
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_name", ["Qdrant/bm25"])
|
||||
def test_multilanguage(model_name: str) -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
|
||||
docs = ["Mangez-vous vraiment des grenouilles?", "Je suis au lit"]
|
||||
|
||||
model = SparseTextEmbedding(model_name=model_name, language="french")
|
||||
embeddings = list(model.embed(docs))[:2]
|
||||
assert embeddings[0].values.shape == (3,)
|
||||
assert embeddings[0].indices.shape == (3,)
|
||||
|
||||
assert embeddings[1].values.shape == (1,)
|
||||
assert embeddings[1].indices.shape == (1,)
|
||||
|
||||
model = SparseTextEmbedding(model_name=model_name, language="english")
|
||||
embeddings = list(model.embed(docs))[:2]
|
||||
assert embeddings[0].values.shape == (5,)
|
||||
assert embeddings[0].indices.shape == (5,)
|
||||
|
||||
assert embeddings[1].values.shape == (4,)
|
||||
assert embeddings[1].indices.shape == (4,)
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_name", ["Qdrant/bm25"])
|
||||
def test_special_characters(model_name: str) -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
|
||||
docs = [
|
||||
"Über den größten Flüssen Österreichs äußern sich Experten häufig: Öko-Systeme müssen geschützt werden!",
|
||||
"L'élève français s'écrie : « Où est mon crayon ? J'ai besoin de finir cet exercice avant la récréation!",
|
||||
"Într-o zi însorită, Ștefan și Ioana au mâncat mămăligă cu brânză și au băut țuică la cabană.",
|
||||
"Üzgün öğretmen öğrencilere seslendi: Lütfen gürültü yapmayın, sınavınızı bitirmeye çalışıyorum!",
|
||||
"Ο Ξενοφών είπε: «Ψάχνω για ένα ωραίο δώρο για τη γιαγιά μου. Ίσως ένα φυτό ή ένα βιβλίο;»",
|
||||
"Hola! ¿Cómo estás? Estoy muy emocionado por el cumpleaños de mi hermano, ¡va a ser increíble! También quiero comprar un pastel de chocolate con fresas y un regalo especial: un libro titulado «Cien años de soledad",
|
||||
]
|
||||
|
||||
model = SparseTextEmbedding(model_name=model_name, language="english")
|
||||
embeddings = list(model.embed(docs))
|
||||
for idx, shape in enumerate([14, 18, 15, 10, 15]):
|
||||
assert embeddings[idx].values.shape == (shape,)
|
||||
assert embeddings[idx].indices.shape == (shape,)
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_name", ["Qdrant/bm42-all-minilm-l6-v2-attentions"])
|
||||
def test_lazy_load(model_name: str) -> None:
|
||||
model = SparseTextEmbedding(model_name=model_name, lazy_load=True)
|
||||
assert not hasattr(model.model, "model")
|
||||
docs = ["hello world", "flag embedding"]
|
||||
list(model.embed(docs))
|
||||
assert hasattr(model.model, "model")
|
||||
|
||||
model = SparseTextEmbedding(model_name=model_name, lazy_load=True)
|
||||
list(model.query_embed(docs))
|
||||
|
||||
model = SparseTextEmbedding(model_name=model_name, lazy_load=True)
|
||||
list(model.passage_embed(docs))
|
||||
|
||||
@@ -0,0 +1,30 @@
|
||||
from fastembed import (
|
||||
TextEmbedding,
|
||||
SparseTextEmbedding,
|
||||
ImageEmbedding,
|
||||
LateInteractionMultimodalEmbedding,
|
||||
LateInteractionTextEmbedding,
|
||||
)
|
||||
|
||||
|
||||
def test_text_list_supported_models():
|
||||
for model_type in [
|
||||
TextEmbedding,
|
||||
SparseTextEmbedding,
|
||||
ImageEmbedding,
|
||||
LateInteractionMultimodalEmbedding,
|
||||
LateInteractionTextEmbedding,
|
||||
]:
|
||||
supported_models = model_type.list_supported_models()
|
||||
assert isinstance(supported_models, list)
|
||||
description = supported_models[0]
|
||||
assert isinstance(description, dict)
|
||||
|
||||
assert "model" in description and description["model"]
|
||||
if model_type != SparseTextEmbedding:
|
||||
assert "dim" in description and description["dim"]
|
||||
assert "license" in description and description["license"]
|
||||
assert "size_in_GB" in description and description["size_in_GB"]
|
||||
assert "model_file" in description and description["model_file"]
|
||||
assert "sources" in description and description["sources"]
|
||||
assert "hf" in description["sources"] or "url" in description["sources"]
|
||||
@@ -1,64 +1,97 @@
|
||||
import os
|
||||
from io import BytesIO
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import requests
|
||||
from PIL import Image
|
||||
|
||||
from fastembed import ImageEmbedding
|
||||
from tests.config import TEST_MISC_DIR
|
||||
from tests.utils import delete_model_cache
|
||||
|
||||
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]
|
||||
),
|
||||
"Qdrant/Unicom-ViT-B-16": np.array(
|
||||
[0.0170, -0.0361, 0.0125, -0.0428, -0.0232, 0.0232, -0.0602, -0.0333, 0.0155, 0.0497]
|
||||
),
|
||||
"Qdrant/Unicom-ViT-B-32": np.array(
|
||||
[0.0418, 0.0550, 0.0003, 0.0253, -0.0185, 0.0016, -0.0368, -0.0402, -0.0891, -0.0186]
|
||||
),
|
||||
"jinaai/jina-clip-v1": np.array(
|
||||
[-0.029, 0.0216, 0.0396, 0.0283, -0.0023, 0.0151, 0.011, -0.0235, 0.0251, -0.0343]
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
def test_embedding():
|
||||
def test_embedding() -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
|
||||
for model_desc in ImageEmbedding.list_supported_models():
|
||||
if not is_ci and model_desc["size_in_GB"] > 1:
|
||||
for model_desc in ImageEmbedding._list_supported_models():
|
||||
if not is_ci and model_desc.size_in_GB > 1:
|
||||
continue
|
||||
|
||||
dim = model_desc["dim"]
|
||||
dim = model_desc.dim
|
||||
|
||||
model = ImageEmbedding(model_name=model_desc["model"])
|
||||
model = ImageEmbedding(model_name=model_desc.model)
|
||||
|
||||
images = [TEST_MISC_DIR / "image.jpeg", str(TEST_MISC_DIR / "small_image.jpeg")]
|
||||
images = [
|
||||
TEST_MISC_DIR / "image.jpeg",
|
||||
str(TEST_MISC_DIR / "small_image.jpeg"),
|
||||
Image.open((TEST_MISC_DIR / "small_image.jpeg")),
|
||||
Image.open(BytesIO(requests.get("https://qdrant.tech/img/logo.png").content)),
|
||||
]
|
||||
embeddings = list(model.embed(images))
|
||||
embeddings = np.stack(embeddings, axis=0)
|
||||
assert embeddings.shape == (2, dim)
|
||||
assert embeddings.shape == (len(images), dim)
|
||||
|
||||
canonical_vector = CANONICAL_VECTOR_VALUES[model_desc["model"]]
|
||||
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"]
|
||||
), model_desc.model
|
||||
|
||||
assert np.allclose(embeddings[1], embeddings[2]), model_desc.model
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("n_dims,model_name", [(512, "Qdrant/clip-ViT-B-32-vision")])
|
||||
def test_batch_embedding(n_dims, model_name):
|
||||
def test_batch_embedding(n_dims: int, model_name: str) -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
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
|
||||
)
|
||||
test_images = [
|
||||
TEST_MISC_DIR / "image.jpeg",
|
||||
str(TEST_MISC_DIR / "small_image.jpeg"),
|
||||
Image.open(TEST_MISC_DIR / "small_image.jpeg"),
|
||||
]
|
||||
images = test_images * n_images
|
||||
|
||||
embeddings = list(model.embed(images, batch_size=10))
|
||||
embeddings = np.stack(embeddings, axis=0)
|
||||
|
||||
assert embeddings.shape == (n_images, n_dims)
|
||||
assert embeddings.shape == (len(test_images) * n_images, n_dims)
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("n_dims,model_name", [(512, "Qdrant/clip-ViT-B-32-vision")])
|
||||
def test_parallel_processing(n_dims, model_name):
|
||||
def test_parallel_processing(n_dims: int, model_name: str) -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
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
|
||||
)
|
||||
test_images = [
|
||||
TEST_MISC_DIR / "image.jpeg",
|
||||
str(TEST_MISC_DIR / "small_image.jpeg"),
|
||||
Image.open(TEST_MISC_DIR / "small_image.jpeg"),
|
||||
]
|
||||
images = test_images * n_images
|
||||
embeddings = list(model.embed(images, batch_size=10, parallel=2))
|
||||
embeddings = np.stack(embeddings, axis=0)
|
||||
|
||||
@@ -68,6 +101,23 @@ def test_parallel_processing(n_dims, model_name):
|
||||
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 embeddings.shape == (n_images * len(test_images), n_dims)
|
||||
assert np.allclose(embeddings, embeddings_2, atol=1e-3)
|
||||
assert np.allclose(embeddings, embeddings_3, atol=1e-3)
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_name", ["Qdrant/clip-ViT-B-32-vision"])
|
||||
def test_lazy_load(model_name: str) -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
model = ImageEmbedding(model_name=model_name, lazy_load=True)
|
||||
assert not hasattr(model.model, "model")
|
||||
images = [
|
||||
TEST_MISC_DIR / "image.jpeg",
|
||||
str(TEST_MISC_DIR / "small_image.jpeg"),
|
||||
]
|
||||
list(model.embed(images))
|
||||
assert hasattr(model.model, "model")
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
@@ -1,8 +1,12 @@
|
||||
import os
|
||||
|
||||
import pytest
|
||||
import numpy as np
|
||||
|
||||
from fastembed.late_interaction.late_interaction_text_embedding import (
|
||||
LateInteractionTextEmbedding,
|
||||
)
|
||||
from tests.utils import delete_model_cache
|
||||
|
||||
# vectors are abridged and rounded for brevity
|
||||
CANONICAL_COLUMN_VALUES = {
|
||||
@@ -14,7 +18,25 @@ CANONICAL_COLUMN_VALUES = {
|
||||
[0.0846, 0.0122, 0.0032, -0.0109, -0.1041],
|
||||
[0.0477, 0.1078, -0.0314, 0.016, 0.0156],
|
||||
]
|
||||
)
|
||||
),
|
||||
"answerdotai/answerai-colbert-small-v1": np.array(
|
||||
[
|
||||
[-0.07281, 0.04632, -0.04711, 0.00762, -0.07374],
|
||||
[-0.04464, 0.04426, -0.074, 0.01801, -0.05233],
|
||||
[0.09936, -0.05123, -0.04925, -0.05276, -0.08944],
|
||||
[0.01644, 0.0203, -0.03789, 0.03165, -0.06501],
|
||||
[-0.07281, 0.04633, -0.04711, 0.00762, -0.07374],
|
||||
]
|
||||
),
|
||||
"jinaai/jina-colbert-v2": np.array(
|
||||
[
|
||||
[0.0742, 0.0591, -0.2403, -0.1774, 0.02],
|
||||
[0.1318, 0.0882, -0.1138, -0.2066, 0.146],
|
||||
[-0.0183, -0.1354, -0.0139, -0.1079, -0.051],
|
||||
[0.0003, -0.1184, -0.07, -0.0479, -0.0649],
|
||||
[0.0766, 0.0452, -0.2343, -0.183, 0.0058],
|
||||
]
|
||||
),
|
||||
}
|
||||
|
||||
CANONICAL_QUERY_VALUES = {
|
||||
@@ -53,13 +75,86 @@ CANONICAL_QUERY_VALUES = {
|
||||
[0.1022, 0.0228, -0.0174, -0.0102, -0.065],
|
||||
[0.1043, 0.0231, -0.0144, -0.0246, -0.067],
|
||||
]
|
||||
)
|
||||
),
|
||||
"answerdotai/answerai-colbert-small-v1": np.array(
|
||||
[
|
||||
[-0.07284, 0.04657, -0.04746, 0.00786, -0.07342],
|
||||
[-0.0473, 0.04615, -0.07551, 0.01591, -0.0517],
|
||||
[0.09658, -0.0506, -0.04593, -0.05225, -0.09086],
|
||||
[0.01815, 0.0165, -0.03366, 0.03214, -0.07019],
|
||||
[-0.07284, 0.04657, -0.04746, 0.00787, -0.07342],
|
||||
[-0.07748, 0.04493, -0.055, 0.00481, -0.0486],
|
||||
[-0.0803, 0.04229, -0.0589, 0.00379, -0.04506],
|
||||
[-0.08477, 0.03724, -0.06162, 0.00578, -0.04554],
|
||||
[-0.08392, 0.03805, -0.06202, 0.00899, -0.0409],
|
||||
[-0.07945, 0.04163, -0.06151, 0.00569, -0.04432],
|
||||
[-0.08469, 0.03985, -0.05765, 0.00485, -0.04485],
|
||||
[-0.08306, 0.04111, -0.05774, 0.00583, -0.04325],
|
||||
[-0.08244, 0.04597, -0.05842, 0.00433, -0.04025],
|
||||
[-0.08385, 0.04745, -0.05845, 0.00469, -0.04002],
|
||||
[-0.08402, 0.05014, -0.05941, 0.00692, -0.03452],
|
||||
[-0.08303, 0.05693, -0.05701, 0.00504, -0.03565],
|
||||
[-0.08216, 0.05516, -0.05687, 0.0057, -0.03748],
|
||||
[-0.08051, 0.05751, -0.05647, 0.00283, -0.03645],
|
||||
[-0.08172, 0.05608, -0.06064, 0.00252, -0.03533],
|
||||
[-0.08073, 0.06144, -0.06373, 0.00935, -0.03154],
|
||||
[-0.06651, 0.06697, -0.06769, 0.01717, -0.03369],
|
||||
[-0.06526, 0.06931, -0.06935, 0.0139, -0.03702],
|
||||
[-0.05435, 0.05829, -0.06593, 0.01708, -0.04559],
|
||||
[-0.03648, 0.05234, -0.06759, 0.02057, -0.05053],
|
||||
[-0.03461, 0.05032, -0.06747, 0.02216, -0.05209],
|
||||
[-0.03444, 0.04835, -0.06812, 0.02296, -0.05276],
|
||||
[-0.03292, 0.04853, -0.06811, 0.02348, -0.05303],
|
||||
[-0.03349, 0.04783, -0.06846, 0.02393, -0.05334],
|
||||
[-0.03485, 0.04677, -0.06826, 0.02362, -0.05326],
|
||||
[-0.03408, 0.04744, -0.06931, 0.02302, -0.05288],
|
||||
[-0.03444, 0.04838, -0.06945, 0.02133, -0.05277],
|
||||
[-0.03473, 0.04792, -0.07033, 0.02196, -0.05314],
|
||||
]
|
||||
),
|
||||
"jinaai/jina-colbert-v2": np.array(
|
||||
[
|
||||
[0.0477, 0.0255, -0.2224, -0.1085, -0.03],
|
||||
[0.0206, -0.0845, -0.0075, -0.1712, 0.0156],
|
||||
[-0.0056, -0.0957, -0.0147, -0.1277, -0.0225],
|
||||
[0.0486, -0.0499, -0.1609, 0.0194, 0.0274],
|
||||
[0.0481, 0.0253, -0.2278, -0.1126, -0.0294],
|
||||
[0.0599, -0.0678, -0.0956, -0.0757, 0.0236],
|
||||
[0.0592, -0.0862, -0.0621, -0.1084, 0.0155],
|
||||
[0.0874, -0.0714, -0.0772, -0.1414, 0.037],
|
||||
[0.1009, -0.0552, -0.0669, -0.163, 0.0493],
|
||||
[0.1135, -0.047, -0.0576, -0.1699, 0.0538],
|
||||
[0.1228, -0.0428, -0.0507, -0.1725, 0.0562],
|
||||
[0.1291, -0.0388, -0.042, -0.1753, 0.0569],
|
||||
[0.1365, -0.0337, -0.0326, -0.1786, 0.0574],
|
||||
[0.1439, -0.026, -0.024, -0.1831, 0.0574],
|
||||
[0.1527, -0.0099, -0.0179, -0.1874, 0.057],
|
||||
[0.1555, 0.0186, -0.023, -0.1801, 0.0539],
|
||||
[0.1389, 0.054, -0.0345, -0.1636, 0.0429],
|
||||
[0.1058, 0.0862, -0.0418, -0.1455, 0.0222],
|
||||
[0.0713, 0.1061, -0.0438, -0.1288, 0.0002],
|
||||
[0.0453, 0.1143, -0.0457, -0.1119, -0.019],
|
||||
[0.0346, 0.1131, -0.0487, -0.0952, -0.0338],
|
||||
[0.0355, 0.1073, -0.0493, -0.0823, -0.0438],
|
||||
[0.0424, 0.1041, -0.0459, -0.0761, -0.048],
|
||||
[0.048, 0.102, -0.0421, -0.0718, -0.0477],
|
||||
[0.0474, 0.0989, -0.0413, -0.0654, -0.0431],
|
||||
[0.0434, 0.095, -0.0415, -0.0589, -0.0345],
|
||||
[0.0408, 0.0897, -0.0405, -0.0554, -0.0197],
|
||||
[0.0433, 0.0811, -0.0407, -0.0545, 0.0055],
|
||||
[0.0514, 0.0629, -0.0446, -0.0549, 0.0368],
|
||||
[0.058, 0.048, -0.0527, -0.0607, 0.0568],
|
||||
[0.0561, 0.0447, -0.0661, -0.0702, 0.0764],
|
||||
[0.0204, -0.0856, -0.0386, -0.1232, -0.0332],
|
||||
]
|
||||
),
|
||||
}
|
||||
|
||||
docs = ["Hello World"]
|
||||
|
||||
|
||||
def test_batch_embedding():
|
||||
is_ci = os.getenv("CI")
|
||||
docs_to_embed = docs * 10
|
||||
|
||||
for model_name, expected_result in CANONICAL_COLUMN_VALUES.items():
|
||||
@@ -69,10 +164,14 @@ def test_batch_embedding():
|
||||
|
||||
for value in result:
|
||||
token_num, abridged_dim = expected_result.shape
|
||||
assert np.allclose(value[:, :abridged_dim], expected_result, atol=10e-4)
|
||||
assert np.allclose(value[:, :abridged_dim], expected_result, atol=2e-3)
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
def test_single_embedding():
|
||||
is_ci = os.getenv("CI")
|
||||
docs_to_embed = docs
|
||||
|
||||
for model_name, expected_result in CANONICAL_COLUMN_VALUES.items():
|
||||
@@ -80,10 +179,14 @@ def test_single_embedding():
|
||||
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)
|
||||
assert np.allclose(result[:, :abridged_dim], expected_result, atol=2e-3)
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
def test_single_embedding_query():
|
||||
is_ci = os.getenv("CI")
|
||||
queries_to_embed = docs
|
||||
|
||||
for model_name, expected_result in CANONICAL_QUERY_VALUES.items():
|
||||
@@ -91,10 +194,14 @@ def test_single_embedding_query():
|
||||
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)
|
||||
assert np.allclose(result[:, :abridged_dim], expected_result, atol=2e-3)
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
def test_parallel_processing():
|
||||
is_ci = os.getenv("CI")
|
||||
model = LateInteractionTextEmbedding(model_name="colbert-ir/colbertv2.0")
|
||||
token_dim = 128
|
||||
docs = ["hello world", "flag embedding"] * 100
|
||||
@@ -110,3 +217,30 @@ def test_parallel_processing():
|
||||
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)
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_name",
|
||||
["colbert-ir/colbertv2.0"],
|
||||
)
|
||||
def test_lazy_load(model_name: str):
|
||||
is_ci = os.getenv("CI")
|
||||
|
||||
model = LateInteractionTextEmbedding(model_name=model_name, lazy_load=True)
|
||||
assert not hasattr(model.model, "model")
|
||||
|
||||
docs = ["hello world", "flag embedding"]
|
||||
list(model.embed(docs))
|
||||
assert hasattr(model.model, "model")
|
||||
|
||||
model = LateInteractionTextEmbedding(model_name=model_name, lazy_load=True)
|
||||
list(model.query_embed(docs))
|
||||
|
||||
model = LateInteractionTextEmbedding(model_name=model_name, lazy_load=True)
|
||||
list(model.passage_embed(docs))
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
@@ -0,0 +1,84 @@
|
||||
import os
|
||||
|
||||
from PIL import Image
|
||||
import numpy as np
|
||||
|
||||
from fastembed import LateInteractionMultimodalEmbedding
|
||||
from tests.config import TEST_MISC_DIR
|
||||
|
||||
|
||||
# vectors are abridged and rounded for brevity
|
||||
CANONICAL_IMAGE_VALUES = {
|
||||
"Qdrant/colpali-v1.3-fp16": np.array(
|
||||
[
|
||||
[
|
||||
[-0.0345, -0.022, 0.0567, -0.0518, -0.0782, 0.1714, -0.1738],
|
||||
[-0.1181, -0.099, 0.0268, 0.0774, 0.0228, 0.0563, -0.1021],
|
||||
[-0.117, -0.0683, 0.0371, 0.0921, 0.0107, 0.0659, -0.0666],
|
||||
[-0.1393, -0.0948, 0.037, 0.0951, -0.0126, 0.0678, -0.087],
|
||||
[-0.0957, -0.081, 0.0404, 0.052, 0.0409, 0.0335, -0.064],
|
||||
[-0.0626, -0.0445, 0.056, 0.0592, -0.0229, 0.0409, -0.0301],
|
||||
[-0.1299, -0.0691, 0.1097, 0.0728, 0.0123, 0.0519, 0.0122],
|
||||
]
|
||||
]
|
||||
),
|
||||
}
|
||||
|
||||
CANONICAL_QUERY_VALUES = {
|
||||
"Qdrant/colpali-v1.3-fp16": np.array(
|
||||
[
|
||||
[-0.0023, 0.1477, 0.1594, 0.046, -0.0196, 0.0554, 0.1567],
|
||||
[-0.0139, -0.0057, 0.0932, 0.0052, -0.0678, 0.0131, 0.0537],
|
||||
[0.0054, 0.0364, 0.2078, -0.074, 0.0355, 0.061, 0.1593],
|
||||
[-0.0076, -0.0154, 0.2266, 0.0103, 0.0089, -0.024, 0.098],
|
||||
[-0.0274, 0.0098, 0.2106, -0.0634, 0.0616, -0.0021, 0.0708],
|
||||
[0.0074, 0.0025, 0.1631, -0.0802, 0.0418, -0.0219, 0.1022],
|
||||
[-0.0165, -0.0106, 0.1672, -0.0768, 0.0389, -0.0038, 0.1137],
|
||||
]
|
||||
),
|
||||
}
|
||||
|
||||
queries = ["hello world", "flag embedding"]
|
||||
images = [
|
||||
TEST_MISC_DIR / "image.jpeg",
|
||||
str(TEST_MISC_DIR / "image.jpeg"),
|
||||
Image.open((TEST_MISC_DIR / "image.jpeg")),
|
||||
]
|
||||
|
||||
|
||||
def test_batch_embedding():
|
||||
is_ci = os.getenv("CI")
|
||||
|
||||
if not is_ci:
|
||||
for model_name, expected_result in CANONICAL_IMAGE_VALUES.items():
|
||||
print("evaluating", model_name)
|
||||
model = LateInteractionMultimodalEmbedding(model_name=model_name)
|
||||
result = list(model.embed_image(images, batch_size=2))
|
||||
|
||||
for value in result:
|
||||
batch_size, token_num, abridged_dim = expected_result.shape
|
||||
assert np.allclose(value[:token_num, :abridged_dim], expected_result, atol=1e-3)
|
||||
|
||||
|
||||
def test_single_embedding():
|
||||
is_ci = os.getenv("CI")
|
||||
if not is_ci:
|
||||
for model_name, expected_result in CANONICAL_IMAGE_VALUES.items():
|
||||
print("evaluating", model_name)
|
||||
model = LateInteractionMultimodalEmbedding(model_name=model_name)
|
||||
result = next(iter(model.embed_image(images, batch_size=6)))
|
||||
batch_size, token_num, abridged_dim = expected_result.shape
|
||||
assert np.allclose(result[:token_num, :abridged_dim], expected_result, atol=2e-3)
|
||||
|
||||
|
||||
def test_single_embedding_query():
|
||||
is_ci = os.getenv("CI")
|
||||
if not is_ci:
|
||||
queries_to_embed = queries
|
||||
|
||||
for model_name, expected_result in CANONICAL_QUERY_VALUES.items():
|
||||
print("evaluating", model_name)
|
||||
model = LateInteractionMultimodalEmbedding(model_name=model_name)
|
||||
result = next(iter(model.embed_text(queries_to_embed)))
|
||||
token_num, abridged_dim = expected_result.shape
|
||||
assert np.allclose(result[:token_num, :abridged_dim], expected_result, atol=2e-3)
|
||||
@@ -0,0 +1,221 @@
|
||||
import pytest
|
||||
from typing import Optional
|
||||
from fastembed import (
|
||||
TextEmbedding,
|
||||
SparseTextEmbedding,
|
||||
LateInteractionTextEmbedding,
|
||||
ImageEmbedding,
|
||||
)
|
||||
from fastembed.rerank.cross_encoder import TextCrossEncoder
|
||||
from tests.config import TEST_MISC_DIR
|
||||
|
||||
CACHE_DIR = "../model_cache"
|
||||
|
||||
|
||||
@pytest.mark.skip(reason="Requires a multi-gpu server")
|
||||
@pytest.mark.parametrize("device_id", [None, 0, 1])
|
||||
def test_gpu_via_providers(device_id: Optional[int]) -> None:
|
||||
docs = ["hello world", "flag embedding"]
|
||||
|
||||
device_id = device_id if device_id is not None else 0
|
||||
providers = (
|
||||
["CUDAExecutionProvider"]
|
||||
if device_id is None
|
||||
else [("CUDAExecutionProvider", {"device_id": device_id})]
|
||||
)
|
||||
embedding_model = TextEmbedding(
|
||||
"sentence-transformers/all-MiniLM-L6-v2",
|
||||
providers=providers,
|
||||
cache_dir=CACHE_DIR,
|
||||
)
|
||||
list(embedding_model.embed(docs))
|
||||
options = embedding_model.model.model.get_provider_options()
|
||||
assert options["CUDAExecutionProvider"]["device_id"] == str(device_id)
|
||||
|
||||
embedding_model = SparseTextEmbedding(
|
||||
"prithvida/Splade_PP_en_v1",
|
||||
providers=providers,
|
||||
cache_dir=CACHE_DIR,
|
||||
)
|
||||
list(embedding_model.embed(docs))
|
||||
options = embedding_model.model.model.get_provider_options()
|
||||
assert options["CUDAExecutionProvider"]["device_id"] == str(device_id)
|
||||
|
||||
embedding_model = SparseTextEmbedding(
|
||||
"Qdrant/bm42-all-minilm-l6-v2-attentions",
|
||||
providers=providers,
|
||||
cache_dir=CACHE_DIR,
|
||||
)
|
||||
list(embedding_model.embed(docs))
|
||||
options = embedding_model.model.model.get_provider_options()
|
||||
assert options["CUDAExecutionProvider"]["device_id"] == str(device_id)
|
||||
|
||||
embedding_model = LateInteractionTextEmbedding(
|
||||
"colbert-ir/colbertv2.0",
|
||||
providers=providers,
|
||||
cache_dir=CACHE_DIR,
|
||||
)
|
||||
list(embedding_model.embed(docs))
|
||||
options = embedding_model.model.model.get_provider_options()
|
||||
assert options["CUDAExecutionProvider"]["device_id"] == str(device_id)
|
||||
|
||||
embedding_model = ImageEmbedding(
|
||||
model_name="Qdrant/clip-ViT-B-32-vision",
|
||||
providers=providers,
|
||||
cache_dir=CACHE_DIR,
|
||||
)
|
||||
images = [
|
||||
TEST_MISC_DIR / "image.jpeg",
|
||||
str(TEST_MISC_DIR / "small_image.jpeg"),
|
||||
]
|
||||
list(embedding_model.embed(images))
|
||||
options = embedding_model.model.model.get_provider_options()
|
||||
assert options["CUDAExecutionProvider"]["device_id"] == str(device_id)
|
||||
|
||||
model = TextCrossEncoder(
|
||||
model_name="Xenova/ms-marco-MiniLM-L-6-v2",
|
||||
providers=providers,
|
||||
cache_dir=CACHE_DIR,
|
||||
)
|
||||
query = "What is the capital of France?"
|
||||
documents = ["Paris is the capital of France.", "Berlin is the capital of Germany."]
|
||||
list(model.rerank(query, documents))
|
||||
options = embedding_model.model.model.get_provider_options()
|
||||
assert options["CUDAExecutionProvider"]["device_id"] == str(device_id)
|
||||
|
||||
|
||||
@pytest.mark.skip(reason="Requires a multi-gpu server")
|
||||
@pytest.mark.parametrize("device_ids", [None, [0], [1], [0, 1]])
|
||||
def test_gpu_cuda_device_ids(device_ids: Optional[list[int]]) -> None:
|
||||
docs = ["hello world", "flag embedding"]
|
||||
device_id = device_ids[0] if device_ids else 0
|
||||
embedding_model = TextEmbedding(
|
||||
"sentence-transformers/all-MiniLM-L6-v2",
|
||||
cuda=True,
|
||||
device_ids=device_ids,
|
||||
cache_dir=CACHE_DIR,
|
||||
)
|
||||
list(embedding_model.embed(docs))
|
||||
options = embedding_model.model.model.get_provider_options()
|
||||
assert options["CUDAExecutionProvider"]["device_id"] == str(
|
||||
device_id
|
||||
), f"Text embedding: {options}"
|
||||
|
||||
embedding_model = SparseTextEmbedding(
|
||||
"prithvida/Splade_PP_en_v1",
|
||||
cuda=True,
|
||||
device_ids=device_ids,
|
||||
cache_dir=CACHE_DIR,
|
||||
)
|
||||
list(embedding_model.embed(docs))
|
||||
options = embedding_model.model.model.get_provider_options()
|
||||
assert options["CUDAExecutionProvider"]["device_id"] == str(
|
||||
device_id
|
||||
), f"Sparse text embedding: {options}"
|
||||
|
||||
embedding_model = SparseTextEmbedding(
|
||||
"Qdrant/bm42-all-minilm-l6-v2-attentions",
|
||||
cuda=True,
|
||||
device_ids=device_ids,
|
||||
cache_dir=CACHE_DIR,
|
||||
)
|
||||
list(embedding_model.embed(docs))
|
||||
options = embedding_model.model.model.get_provider_options()
|
||||
assert options["CUDAExecutionProvider"]["device_id"] == str(device_id), f"Bm42: {options}"
|
||||
|
||||
embedding_model = LateInteractionTextEmbedding(
|
||||
"colbert-ir/colbertv2.0",
|
||||
cuda=True,
|
||||
device_ids=device_ids,
|
||||
cache_dir=CACHE_DIR,
|
||||
)
|
||||
list(embedding_model.embed(docs))
|
||||
options = embedding_model.model.model.get_provider_options()
|
||||
assert options["CUDAExecutionProvider"]["device_id"] == str(
|
||||
device_id
|
||||
), f"Late interaction text embedding: {options}"
|
||||
|
||||
embedding_model = ImageEmbedding(
|
||||
model_name="Qdrant/clip-ViT-B-32-vision",
|
||||
cuda=True,
|
||||
device_ids=device_ids,
|
||||
cache_dir=CACHE_DIR,
|
||||
)
|
||||
images = [
|
||||
TEST_MISC_DIR / "image.jpeg",
|
||||
str(TEST_MISC_DIR / "small_image.jpeg"),
|
||||
]
|
||||
list(embedding_model.embed(images))
|
||||
options = embedding_model.model.model.get_provider_options()
|
||||
assert options["CUDAExecutionProvider"]["device_id"] == str(
|
||||
device_id
|
||||
), f"Image embedding: {options}"
|
||||
|
||||
if device_ids is None or len(device_ids) == 1:
|
||||
model = TextCrossEncoder(
|
||||
model_name="Xenova/ms-marco-MiniLM-L-6-v2",
|
||||
cuda=True,
|
||||
device_ids=device_ids,
|
||||
cache_dir=CACHE_DIR,
|
||||
)
|
||||
query = "What is the capital of France?"
|
||||
documents = ["Paris is the capital of France.", "Berlin is the capital of Germany."]
|
||||
list(model.rerank(query, documents))
|
||||
options = embedding_model.model.model.get_provider_options()
|
||||
assert options["CUDAExecutionProvider"]["device_id"] == str(
|
||||
device_id
|
||||
), f"Text cross encoder: {options}"
|
||||
|
||||
|
||||
@pytest.mark.skip(reason="Requires a multi-gpu server")
|
||||
@pytest.mark.parametrize(
|
||||
"device_ids,parallel", [(None, None), (None, 2), ([1], None), ([1], 1), ([1], 2), ([0, 1], 2)]
|
||||
)
|
||||
def test_multi_gpu_parallel_inference(device_ids: Optional[list[int]], parallel: int) -> None:
|
||||
docs = ["hello world", "flag embedding"] * 100
|
||||
batch_size = 5
|
||||
|
||||
embedding_model = TextEmbedding(
|
||||
"sentence-transformers/all-MiniLM-L6-v2",
|
||||
cuda=True,
|
||||
device_ids=device_ids,
|
||||
cache_dir=CACHE_DIR,
|
||||
lazy_load=True,
|
||||
)
|
||||
list(embedding_model.embed(docs, batch_size=batch_size, parallel=parallel))
|
||||
|
||||
embedding_model = SparseTextEmbedding(
|
||||
"prithvida/Splade_PP_en_v1",
|
||||
cuda=True,
|
||||
device_ids=device_ids,
|
||||
cache_dir=CACHE_DIR,
|
||||
)
|
||||
list(embedding_model.embed(docs, batch_size=batch_size, parallel=parallel))
|
||||
|
||||
embedding_model = SparseTextEmbedding(
|
||||
"Qdrant/bm42-all-minilm-l6-v2-attentions",
|
||||
cuda=True,
|
||||
device_ids=device_ids,
|
||||
cache_dir=CACHE_DIR,
|
||||
)
|
||||
list(embedding_model.embed(docs, batch_size=batch_size, parallel=parallel))
|
||||
|
||||
embedding_model = LateInteractionTextEmbedding(
|
||||
"colbert-ir/colbertv2.0",
|
||||
cuda=True,
|
||||
device_ids=device_ids,
|
||||
cache_dir=CACHE_DIR,
|
||||
)
|
||||
list(embedding_model.embed(docs, batch_size=batch_size, parallel=parallel))
|
||||
|
||||
embedding_model = ImageEmbedding(
|
||||
model_name="Qdrant/clip-ViT-B-32-vision",
|
||||
cuda=True,
|
||||
device_ids=device_ids,
|
||||
cache_dir=CACHE_DIR,
|
||||
)
|
||||
images = [
|
||||
TEST_MISC_DIR / "image.jpeg",
|
||||
str(TEST_MISC_DIR / "small_image.jpeg"),
|
||||
] * 100
|
||||
list(embedding_model.embed(images, batch_size=batch_size, parallel=parallel))
|
||||
+106
-12
@@ -1,6 +1,11 @@
|
||||
import pytest
|
||||
import os
|
||||
|
||||
import pytest
|
||||
import numpy as np
|
||||
|
||||
from fastembed.sparse.bm25 import Bm25
|
||||
from fastembed.sparse.sparse_text_embedding import SparseTextEmbedding
|
||||
from tests.utils import delete_model_cache
|
||||
|
||||
CANONICAL_COLUMN_VALUES = {
|
||||
"prithvida/Splade_PP_en_v1": {
|
||||
@@ -44,7 +49,8 @@ CANONICAL_COLUMN_VALUES = {
|
||||
docs = ["Hello World"]
|
||||
|
||||
|
||||
def test_batch_embedding():
|
||||
def test_batch_embedding() -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
docs_to_embed = docs * 10
|
||||
|
||||
for model_name, expected_result in CANONICAL_COLUMN_VALUES.items():
|
||||
@@ -54,9 +60,12 @@ def test_batch_embedding():
|
||||
|
||||
for i, value in enumerate(result.values):
|
||||
assert pytest.approx(value, abs=0.001) == expected_result["values"][i]
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
def test_single_embedding():
|
||||
def test_single_embedding() -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
for model_name, expected_result in CANONICAL_COLUMN_VALUES.items():
|
||||
model = SparseTextEmbedding(model_name=model_name)
|
||||
|
||||
@@ -67,11 +76,12 @@ def test_single_embedding():
|
||||
|
||||
for i, value in enumerate(result.values):
|
||||
assert pytest.approx(value, abs=0.001) == expected_result["values"][i]
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
def test_parallel_processing():
|
||||
import numpy as np
|
||||
|
||||
def test_parallel_processing() -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
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))
|
||||
@@ -93,9 +103,93 @@ 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)
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def bm25_instance() -> None:
|
||||
ci = os.getenv("CI", True)
|
||||
model = Bm25("Qdrant/bm25", language="english")
|
||||
yield model
|
||||
if ci:
|
||||
delete_model_cache(model._model_dir)
|
||||
|
||||
|
||||
def test_stem_with_stopwords_and_punctuation(bm25_instance: Bm25) -> None:
|
||||
# Setup
|
||||
bm25_instance.stopwords = {"the", "is", "a"}
|
||||
bm25_instance.punctuation = {".", ",", "!"}
|
||||
|
||||
# Test data
|
||||
tokens = ["The", "quick", "brown", "fox", "is", "a", "test", "sentence", ".", "!"]
|
||||
|
||||
# Execute
|
||||
result = bm25_instance._stem(tokens)
|
||||
|
||||
# Assert
|
||||
expected = ["quick", "brown", "fox", "test", "sentenc"]
|
||||
assert result == expected, f"Expected {expected}, but got {result}"
|
||||
|
||||
|
||||
def test_stem_case_insensitive_stopwords(bm25_instance: Bm25) -> None:
|
||||
# Setup
|
||||
bm25_instance.stopwords = {"the", "is", "a"}
|
||||
bm25_instance.punctuation = {".", ",", "!"}
|
||||
|
||||
# Test data
|
||||
tokens = ["THE", "Quick", "Brown", "Fox", "IS", "A", "Test", "Sentence", ".", "!"]
|
||||
|
||||
# Execute
|
||||
result = bm25_instance._stem(tokens)
|
||||
|
||||
# Assert
|
||||
expected = ["quick", "brown", "fox", "test", "sentenc"]
|
||||
assert result == expected, f"Expected {expected}, but got {result}"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("disable_stemmer", [True, False])
|
||||
def test_disable_stemmer_behavior(disable_stemmer: bool) -> None:
|
||||
# Setup
|
||||
model = Bm25("Qdrant/bm25", language="english", disable_stemmer=disable_stemmer)
|
||||
model.stopwords = {"the", "is", "a"}
|
||||
model.punctuation = {".", ",", "!"}
|
||||
|
||||
# Test data
|
||||
tokens = ["The", "quick", "brown", "fox", "is", "a", "test", "sentence", ".", "!"]
|
||||
|
||||
# Execute
|
||||
result = model._stem(tokens)
|
||||
|
||||
# Assert
|
||||
if disable_stemmer:
|
||||
expected = ["quick", "brown", "fox", "test", "sentence"] # no stemming, lower case only
|
||||
else:
|
||||
expected = ["quick", "brown", "fox", "test", "sentenc"]
|
||||
assert result == expected, f"Expected {expected}, but got {result}"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_name",
|
||||
["prithivida/Splade_PP_en_v1"],
|
||||
)
|
||||
def test_lazy_load(model_name: str) -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
model = SparseTextEmbedding(model_name=model_name, lazy_load=True)
|
||||
assert not hasattr(model.model, "model")
|
||||
|
||||
docs = ["hello world", "flag embedding"]
|
||||
list(model.embed(docs))
|
||||
assert hasattr(model.model, "model")
|
||||
|
||||
model = SparseTextEmbedding(model_name=model_name, lazy_load=True)
|
||||
list(model.query_embed(docs))
|
||||
|
||||
model = SparseTextEmbedding(model_name=model_name, lazy_load=True)
|
||||
list(model.passage_embed(docs))
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
@@ -0,0 +1,119 @@
|
||||
import os
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from fastembed.rerank.cross_encoder import TextCrossEncoder
|
||||
from tests.utils import delete_model_cache
|
||||
|
||||
CANONICAL_SCORE_VALUES = {
|
||||
"Xenova/ms-marco-MiniLM-L-6-v2": np.array([8.500708, -2.541011]),
|
||||
"Xenova/ms-marco-MiniLM-L-12-v2": np.array([9.330912, -2.0380247]),
|
||||
"BAAI/bge-reranker-base": np.array([6.15733337, -3.65939403]),
|
||||
"jinaai/jina-reranker-v1-tiny-en": np.array([2.5911, 0.1122]),
|
||||
"jinaai/jina-reranker-v1-turbo-en": np.array([1.8295, -2.8908]),
|
||||
"jinaai/jina-reranker-v2-base-multilingual": np.array([1.6533, -1.6455]),
|
||||
}
|
||||
|
||||
SELECTED_MODELS = {
|
||||
"Xenova": "Xenova/ms-marco-MiniLM-L-6-v2",
|
||||
"BAAI": "BAAI/bge-reranker-base",
|
||||
"jinaai": "jinaai/jina-reranker-v1-tiny-en",
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_name",
|
||||
[model_name for model_name in CANONICAL_SCORE_VALUES],
|
||||
)
|
||||
def test_rerank(model_name: str) -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
|
||||
model = TextCrossEncoder(model_name=model_name)
|
||||
|
||||
query = "What is the capital of France?"
|
||||
documents = ["Paris is the capital of France.", "Berlin is the capital of Germany."]
|
||||
scores = np.array(list(model.rerank(query, documents)))
|
||||
|
||||
pairs = [(query, doc) for doc in documents]
|
||||
scores2 = np.array(list(model.rerank_pairs(pairs)))
|
||||
assert np.allclose(
|
||||
scores, scores2, atol=1e-5
|
||||
), f"Model: {model_name}, Scores: {scores}, Scores2: {scores2}"
|
||||
|
||||
canonical_scores = CANONICAL_SCORE_VALUES[model_name]
|
||||
assert np.allclose(
|
||||
scores, canonical_scores, atol=1e-3
|
||||
), f"Model: {model_name}, Scores: {scores}, Expected: {canonical_scores}"
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_name",
|
||||
[model_name for model_name in SELECTED_MODELS.values()],
|
||||
)
|
||||
def test_batch_rerank(model_name: str) -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
|
||||
model = TextCrossEncoder(model_name=model_name)
|
||||
|
||||
query = "What is the capital of France?"
|
||||
documents = ["Paris is the capital of France.", "Berlin is the capital of Germany."] * 50
|
||||
scores = np.array(list(model.rerank(query, documents, batch_size=10)))
|
||||
|
||||
pairs = [(query, doc) for doc in documents]
|
||||
scores2 = np.array(list(model.rerank_pairs(pairs)))
|
||||
assert np.allclose(
|
||||
scores, scores2, atol=1e-5
|
||||
), f"Model: {model_name}, Scores: {scores}, Scores2: {scores2}"
|
||||
|
||||
canonical_scores = np.tile(CANONICAL_SCORE_VALUES[model_name], 50)
|
||||
|
||||
assert scores.shape == canonical_scores.shape, f"Unexpected shape for model {model_name}"
|
||||
assert np.allclose(
|
||||
scores, canonical_scores, atol=1e-3
|
||||
), f"Model: {model_name}, Scores: {scores}, Expected: {canonical_scores}"
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_name",
|
||||
["Xenova/ms-marco-MiniLM-L-6-v2"],
|
||||
)
|
||||
def test_lazy_load(model_name: str) -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
model = TextCrossEncoder(model_name=model_name, lazy_load=True)
|
||||
assert not hasattr(model.model, "model")
|
||||
query = "What is the capital of France?"
|
||||
documents = ["Paris is the capital of France.", "Berlin is the capital of Germany."]
|
||||
list(model.rerank(query, documents))
|
||||
assert hasattr(model.model, "model")
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_name",
|
||||
[model_name for model_name in SELECTED_MODELS.values()],
|
||||
)
|
||||
def test_rerank_pairs_parallel(model_name: str) -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
|
||||
model = TextCrossEncoder(model_name=model_name)
|
||||
query = "What is the capital of France?"
|
||||
documents = ["Paris is the capital of France.", "Berlin is the capital of Germany."] * 10
|
||||
pairs = [(query, doc) for doc in documents]
|
||||
scores_parallel = np.array(list(model.rerank_pairs(pairs, parallel=2, batch_size=10)))
|
||||
scores_sequential = np.array(list(model.rerank_pairs(pairs, batch_size=10)))
|
||||
assert np.allclose(
|
||||
scores_parallel, scores_sequential, atol=1e-5
|
||||
), f"Model: {model_name}, Scores (Parallel): {scores_parallel}, Scores (Sequential): {scores_sequential}"
|
||||
canonical_scores = CANONICAL_SCORE_VALUES[model_name]
|
||||
assert np.allclose(
|
||||
scores_parallel[: len(canonical_scores)], canonical_scores, atol=1e-3
|
||||
), f"Model: {model_name}, Scores (Parallel): {scores_parallel}, Expected: {canonical_scores}"
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
@@ -0,0 +1,253 @@
|
||||
import os
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from fastembed import TextEmbedding
|
||||
from fastembed.text.multitask_embedding import Task
|
||||
from tests.utils import delete_model_cache
|
||||
|
||||
|
||||
CANONICAL_VECTOR_VALUES = {
|
||||
"jinaai/jina-embeddings-v3": [
|
||||
{
|
||||
"task_id": Task.RETRIEVAL_QUERY,
|
||||
"vectors": np.array(
|
||||
[
|
||||
[0.0623, -0.0402, 0.1706, -0.0143, 0.0617],
|
||||
[-0.1064, -0.0733, 0.0353, 0.0096, 0.0667],
|
||||
]
|
||||
),
|
||||
},
|
||||
{
|
||||
"task_id": Task.RETRIEVAL_PASSAGE,
|
||||
"vectors": np.array(
|
||||
[
|
||||
[0.0513, -0.0247, 0.1751, -0.0075, 0.0679],
|
||||
[-0.0987, -0.0786, 0.09, 0.0087, 0.0577],
|
||||
]
|
||||
),
|
||||
},
|
||||
{
|
||||
"task_id": Task.SEPARATION,
|
||||
"vectors": np.array(
|
||||
[
|
||||
[0.094, -0.1065, 0.1305, 0.0547, 0.0556],
|
||||
[0.0315, -0.1468, 0.065, 0.0568, 0.0546],
|
||||
]
|
||||
),
|
||||
},
|
||||
{
|
||||
"task_id": Task.CLASSIFICATION,
|
||||
"vectors": np.array(
|
||||
[
|
||||
[0.0606, -0.0877, 0.1384, 0.0065, 0.0722],
|
||||
[-0.0502, -0.119, 0.032, 0.0514, 0.0689],
|
||||
]
|
||||
),
|
||||
},
|
||||
{
|
||||
"task_id": Task.TEXT_MATCHING,
|
||||
"vectors": np.array(
|
||||
[
|
||||
[0.0911, -0.0341, 0.1305, -0.026, 0.0576],
|
||||
[-0.1432, -0.05, 0.0133, 0.0464, 0.0789],
|
||||
]
|
||||
),
|
||||
},
|
||||
]
|
||||
}
|
||||
docs = ["Hello World", "Follow the white rabbit."]
|
||||
|
||||
|
||||
def test_batch_embedding():
|
||||
is_ci = os.getenv("CI")
|
||||
docs_to_embed = docs * 10
|
||||
default_task = Task.RETRIEVAL_PASSAGE
|
||||
|
||||
for model_desc in TextEmbedding._list_supported_models():
|
||||
if not is_ci and model_desc.size_in_GB > 1:
|
||||
continue
|
||||
|
||||
model_name = model_desc.model
|
||||
dim = model_desc.dim
|
||||
|
||||
if model_name not in CANONICAL_VECTOR_VALUES.keys():
|
||||
continue
|
||||
|
||||
model = TextEmbedding(model_name=model_name)
|
||||
|
||||
print(f"evaluating {model_name} default task")
|
||||
|
||||
embeddings = list(model.embed(documents=docs_to_embed, batch_size=6))
|
||||
embeddings = np.stack(embeddings, axis=0)
|
||||
|
||||
assert embeddings.shape == (len(docs_to_embed), dim)
|
||||
|
||||
canonical_vector = CANONICAL_VECTOR_VALUES[model_name][default_task]["vectors"]
|
||||
assert np.allclose(
|
||||
embeddings[: len(docs), : canonical_vector.shape[1]], canonical_vector, atol=1e-4
|
||||
), model_desc.model
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
def test_single_embedding():
|
||||
is_ci = os.getenv("CI")
|
||||
|
||||
for model_desc in TextEmbedding._list_supported_models():
|
||||
if not is_ci and model_desc.size_in_GB > 1:
|
||||
continue
|
||||
|
||||
model_name = model_desc.model
|
||||
dim = model_desc.dim
|
||||
|
||||
if model_name not in CANONICAL_VECTOR_VALUES.keys():
|
||||
continue
|
||||
|
||||
model = TextEmbedding(model_name=model_name)
|
||||
|
||||
for task in CANONICAL_VECTOR_VALUES[model_name]:
|
||||
print(f"evaluating {model_name} task_id: {task['task_id']}")
|
||||
|
||||
embeddings = list(model.embed(documents=docs, task_id=task["task_id"]))
|
||||
embeddings = np.stack(embeddings, axis=0)
|
||||
|
||||
assert embeddings.shape == (len(docs), dim)
|
||||
|
||||
canonical_vector = task["vectors"]
|
||||
assert np.allclose(
|
||||
embeddings[: len(docs), : canonical_vector.shape[1]], canonical_vector, atol=1e-4
|
||||
), model_desc.model
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
def test_single_embedding_query():
|
||||
is_ci = os.getenv("CI")
|
||||
task_id = Task.RETRIEVAL_QUERY
|
||||
|
||||
for model_desc in TextEmbedding._list_supported_models():
|
||||
if not is_ci and model_desc.size_in_GB > 1:
|
||||
continue
|
||||
|
||||
model_name = model_desc.model
|
||||
dim = model_desc.dim
|
||||
|
||||
if model_name not in CANONICAL_VECTOR_VALUES.keys():
|
||||
continue
|
||||
|
||||
model = TextEmbedding(model_name=model_name)
|
||||
|
||||
print(f"evaluating {model_name} query_embed task_id: {task_id}")
|
||||
|
||||
embeddings = list(model.query_embed(query=docs))
|
||||
embeddings = np.stack(embeddings, axis=0)
|
||||
|
||||
assert embeddings.shape == (len(docs), dim)
|
||||
|
||||
canonical_vector = CANONICAL_VECTOR_VALUES[model_name][task_id]["vectors"]
|
||||
assert np.allclose(
|
||||
embeddings[: len(docs), : canonical_vector.shape[1]], canonical_vector, atol=1e-4
|
||||
), model_desc.model
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
def test_single_embedding_passage():
|
||||
is_ci = os.getenv("CI")
|
||||
task_id = Task.RETRIEVAL_PASSAGE
|
||||
|
||||
for model_desc in TextEmbedding._list_supported_models():
|
||||
if not is_ci and model_desc.size_in_GB > 1:
|
||||
continue
|
||||
|
||||
model_name = model_desc.model
|
||||
dim = model_desc.dim
|
||||
|
||||
if model_name not in CANONICAL_VECTOR_VALUES.keys():
|
||||
continue
|
||||
|
||||
model = TextEmbedding(model_name=model_name)
|
||||
|
||||
print(f"evaluating {model_name} passage_embed task_id: {task_id}")
|
||||
|
||||
embeddings = list(model.passage_embed(texts=docs))
|
||||
embeddings = np.stack(embeddings, axis=0)
|
||||
|
||||
assert embeddings.shape == (len(docs), dim)
|
||||
|
||||
canonical_vector = CANONICAL_VECTOR_VALUES[model_name][task_id]["vectors"]
|
||||
assert np.allclose(
|
||||
embeddings[: len(docs), : canonical_vector.shape[1]], canonical_vector, atol=1e-4
|
||||
), model_desc.model
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
def test_parallel_processing():
|
||||
is_ci = os.getenv("CI")
|
||||
|
||||
docs = ["Hello World", "Follow the white rabbit."] * 10
|
||||
|
||||
model_name = "jinaai/jina-embeddings-v3"
|
||||
dim = 1024
|
||||
|
||||
model = TextEmbedding(model_name=model_name)
|
||||
|
||||
task_id = Task.SEPARATION
|
||||
embeddings_1 = list(model.embed(docs, batch_size=10, parallel=None, task_id=task_id))
|
||||
embeddings_1 = np.stack(embeddings_1, axis=0)
|
||||
|
||||
embeddings_2 = list(model.embed(docs, batch_size=10, parallel=1, task_id=task_id))
|
||||
embeddings_2 = np.stack(embeddings_2, axis=0)
|
||||
|
||||
assert embeddings_1.shape[0] == len(docs) and embeddings_1.shape[-1] == dim
|
||||
assert np.allclose(embeddings_1, embeddings_2, atol=1e-4)
|
||||
|
||||
canonical_vector = CANONICAL_VECTOR_VALUES[model_name][task_id]["vectors"]
|
||||
assert np.allclose(embeddings_2[:2, : canonical_vector.shape[1]], canonical_vector, atol=1e-4)
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
def test_task_assignment():
|
||||
is_ci = os.getenv("CI")
|
||||
|
||||
for model_desc in TextEmbedding._list_supported_models():
|
||||
if not is_ci and model_desc.size_in_GB > 1:
|
||||
continue
|
||||
|
||||
model_name = model_desc.model
|
||||
if model_name not in CANONICAL_VECTOR_VALUES.keys():
|
||||
continue
|
||||
|
||||
model = TextEmbedding(model_name=model_name)
|
||||
|
||||
for i, task_id in enumerate(Task):
|
||||
_ = list(model.embed(documents=docs, batch_size=1, task_id=i))
|
||||
assert model.model.current_task_id == task_id
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_name",
|
||||
["jinaai/jina-embeddings-v3"],
|
||||
)
|
||||
def test_lazy_load(model_name: str):
|
||||
is_ci = os.getenv("CI")
|
||||
model = TextEmbedding(model_name=model_name, lazy_load=True)
|
||||
assert not hasattr(model.model, "model")
|
||||
|
||||
list(model.embed(docs))
|
||||
assert hasattr(model.model, "model")
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
@@ -1,9 +1,11 @@
|
||||
import os
|
||||
import platform
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from fastembed.text.text_embedding import TextEmbedding
|
||||
from tests.utils import delete_model_cache
|
||||
|
||||
CANONICAL_VECTOR_VALUES = {
|
||||
"BAAI/bge-small-en": np.array([-0.0232, -0.0255, 0.0174, -0.0639, -0.0006]),
|
||||
@@ -30,31 +32,24 @@ CANONICAL_VECTOR_VALUES = {
|
||||
[-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]
|
||||
[0.0361, 0.1862, 0.2776, 0.2461, -0.1904]
|
||||
),
|
||||
"intfloat/multilingual-e5-large": np.array([0.4544, -0.0968, 0.1054, -1.3753, 0.1500]),
|
||||
"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]
|
||||
),
|
||||
"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]
|
||||
[0.0047, 0.1334, -0.0102, 0.0714, 0.1930]
|
||||
),
|
||||
"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]),
|
||||
"jinaai/jina-embeddings-v2-base-code": np.array([0.0145, -0.0164, 0.0136, -0.0170, 0.0734]),
|
||||
"jinaai/jina-embeddings-v2-base-zh": np.array([0.0381, 0.0286, -0.0231, 0.0052, -0.0151]),
|
||||
"jinaai/jina-embeddings-v2-base-es": np.array([-0.0108, -0.0092, -0.0373, 0.0171, -0.0301]),
|
||||
"nomic-ai/nomic-embed-text-v1": np.array([0.3708, 0.2031, -0.3406, -0.2114, -0.3230]),
|
||||
"nomic-ai/nomic-embed-text-v1.5": np.array(
|
||||
[-1.6531514e-02, 8.5380634e-05, -1.8171231e-01, -3.9333291e-03, 1.2763254e-02]
|
||||
[-0.15407836, -0.03053198, -3.9138033, 0.1910364, 0.13224715]
|
||||
),
|
||||
"nomic-ai/nomic-embed-text-v1.5-Q": np.array(
|
||||
[-0.01554983, 0.0129992, -0.17909265, -0.01062993, 0.00512859]
|
||||
[0.0802303, 0.3700881, -4.3053818, 0.4431803, -0.271572]
|
||||
),
|
||||
"thenlper/gte-large": np.array(
|
||||
[-0.01920587, 0.00113156, -0.00708992, -0.00632304, -0.04025577]
|
||||
@@ -62,52 +57,55 @@ CANONICAL_VECTOR_VALUES = {
|
||||
"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-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]),
|
||||
"thenlper/gte-base": np.array([0.0038, 0.0355, 0.0181, 0.0092, 0.0654]),
|
||||
"jinaai/jina-clip-v1": np.array([-0.0862, -0.0101, -0.0056, 0.0375, -0.0472]),
|
||||
}
|
||||
|
||||
MULTI_TASK_MODELS = ["jinaai/jina-embeddings-v3"]
|
||||
|
||||
def test_embedding():
|
||||
|
||||
def test_embedding() -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
is_mac = platform.system() == "Darwin"
|
||||
|
||||
for model_desc in TextEmbedding.list_supported_models():
|
||||
if not is_ci and model_desc["size_in_GB"] > 1:
|
||||
for model_desc in TextEmbedding._list_supported_models():
|
||||
if (
|
||||
(not is_ci and model_desc.size_in_GB > 1)
|
||||
or model_desc.model in MULTI_TASK_MODELS
|
||||
or (is_mac and model_desc.model == "nomic-ai/nomic-embed-text-v1.5-Q")
|
||||
):
|
||||
continue
|
||||
|
||||
dim = model_desc["dim"]
|
||||
|
||||
model = TextEmbedding(model_name=model_desc["model"])
|
||||
dim = model_desc.dim
|
||||
|
||||
model = TextEmbedding(model_name=model_desc.model)
|
||||
docs = ["hello world", "flag embedding"]
|
||||
embeddings = list(model.embed(docs))
|
||||
embeddings = np.stack(embeddings, axis=0)
|
||||
assert embeddings.shape == (2, dim)
|
||||
|
||||
canonical_vector = CANONICAL_VECTOR_VALUES[model_desc["model"]]
|
||||
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"]
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"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):
|
||||
def test_batch_embedding(n_dims: int, model_name: str) -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
model = TextEmbedding(model_name=model_name)
|
||||
|
||||
docs = ["hello world", "flag embedding"] * 100
|
||||
@@ -115,13 +113,16 @@ def test_batch_embedding(n_dims, model_name):
|
||||
embeddings = np.stack(embeddings, axis=0)
|
||||
|
||||
assert embeddings.shape == (200, n_dims)
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"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):
|
||||
def test_parallel_processing(n_dims: int, model_name: str) -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
model = TextEmbedding(model_name=model_name)
|
||||
|
||||
docs = ["hello world", "flag embedding"] * 100
|
||||
@@ -137,3 +138,28 @@ def test_parallel_processing(n_dims, model_name):
|
||||
assert embeddings.shape == (200, n_dims)
|
||||
assert np.allclose(embeddings, embeddings_2, atol=1e-3)
|
||||
assert np.allclose(embeddings, embeddings_3, atol=1e-3)
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_name",
|
||||
["BAAI/bge-small-en-v1.5"],
|
||||
)
|
||||
def test_lazy_load(model_name: str) -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
model = TextEmbedding(model_name=model_name, lazy_load=True)
|
||||
assert not hasattr(model.model, "model")
|
||||
docs = ["hello world", "flag embedding"]
|
||||
list(model.embed(docs))
|
||||
assert hasattr(model.model, "model")
|
||||
|
||||
model = TextEmbedding(model_name=model_name, lazy_load=True)
|
||||
list(model.query_embed(docs))
|
||||
|
||||
model = TextEmbedding(model_name=model_name, lazy_load=True)
|
||||
list(model.passage_embed(docs))
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
@@ -0,0 +1,56 @@
|
||||
from fastembed import TextEmbedding, LateInteractionTextEmbedding, SparseTextEmbedding
|
||||
from fastembed.sparse.bm25 import Bm25
|
||||
from fastembed.rerank.cross_encoder import TextCrossEncoder
|
||||
|
||||
|
||||
text_embedder = TextEmbedding(cache_dir="models")
|
||||
late_interaction_embedder = LateInteractionTextEmbedding(model_name="", cache_dir="models")
|
||||
reranker = TextCrossEncoder(model_name="", cache_dir="models")
|
||||
sparse_embedder = SparseTextEmbedding(model_name="", cache_dir="models")
|
||||
bm25_embedder = Bm25(
|
||||
model_name="",
|
||||
k=1.0,
|
||||
b=1.0,
|
||||
avg_len=1.0,
|
||||
language="",
|
||||
token_max_length=1,
|
||||
disable_stemmer=False,
|
||||
specific_model_path="models",
|
||||
)
|
||||
|
||||
text_embedder.list_supported_models()
|
||||
text_embedder.embed(documents=[""], batch_size=1, parallel=1)
|
||||
text_embedder.embed(documents="", parallel=None, task_id=1)
|
||||
text_embedder.query_embed(query=[""], batch_size=1, parallel=1)
|
||||
text_embedder.query_embed(query="", parallel=None)
|
||||
text_embedder.passage_embed(texts=[""], batch_size=1, parallel=1)
|
||||
text_embedder.passage_embed(texts=[""], parallel=None)
|
||||
|
||||
late_interaction_embedder.list_supported_models()
|
||||
late_interaction_embedder.embed(documents=[""], batch_size=1, parallel=1)
|
||||
late_interaction_embedder.embed(documents="", parallel=None)
|
||||
late_interaction_embedder.query_embed(query=[""], batch_size=1, parallel=1)
|
||||
late_interaction_embedder.query_embed(query="", parallel=None)
|
||||
late_interaction_embedder.passage_embed(texts=[""], batch_size=1, parallel=1)
|
||||
late_interaction_embedder.passage_embed(texts=[""], parallel=None)
|
||||
|
||||
reranker.list_supported_models()
|
||||
reranker.rerank(query="", documents=[""], batch_size=1, parallel=1)
|
||||
reranker.rerank(query="", documents=[""], parallel=None)
|
||||
reranker.rerank_pairs(pairs=[("", "")], batch_size=1, parallel=1)
|
||||
reranker.rerank_pairs(pairs=[("", "")], parallel=None)
|
||||
|
||||
sparse_embedder.list_supported_models()
|
||||
sparse_embedder.embed(documents=[""], batch_size=1, parallel=1)
|
||||
sparse_embedder.embed(documents="", batch_size=1, parallel=None)
|
||||
sparse_embedder.query_embed(query=[""], batch_size=1, parallel=1)
|
||||
sparse_embedder.query_embed(query="", batch_size=1, parallel=None)
|
||||
sparse_embedder.passage_embed(texts=[""], batch_size=1, parallel=1)
|
||||
sparse_embedder.passage_embed(texts=[""], batch_size=1, parallel=None)
|
||||
|
||||
bm25_embedder.list_supported_models()
|
||||
bm25_embedder.embed(documents=[""], batch_size=1, parallel=1)
|
||||
bm25_embedder.embed(documents="", batch_size=1, parallel=None)
|
||||
bm25_embedder.query_embed(query=[""], batch_size=1, parallel=1)
|
||||
bm25_embedder.query_embed(query="", batch_size=1, parallel=None)
|
||||
bm25_embedder.raw_embed(documents=[""])
|
||||
@@ -0,0 +1,37 @@
|
||||
import shutil
|
||||
import traceback
|
||||
|
||||
from pathlib import Path
|
||||
from types import TracebackType
|
||||
from typing import Union, Callable, Any, Type
|
||||
|
||||
|
||||
def delete_model_cache(model_dir: Union[str, Path]) -> None:
|
||||
"""Delete the model cache directory.
|
||||
|
||||
If a model was downloaded from the HuggingFace model hub, then _model_dir is the dir to snapshots, removing
|
||||
it won't help to release the memory, because data is in blobs directory.
|
||||
If a model was downloaded from GCS, then we can just remove model_dir
|
||||
|
||||
Args:
|
||||
model_dir (Union[str, Path]): The path to the model cache directory.
|
||||
"""
|
||||
|
||||
def on_error(
|
||||
func: Callable[..., Any],
|
||||
path: str,
|
||||
exc_info: tuple[Type[BaseException], BaseException, TracebackType],
|
||||
) -> None:
|
||||
print("Failed to remove: ", path)
|
||||
print("Exception: ", exc_info)
|
||||
traceback.print_exception(*exc_info)
|
||||
|
||||
if isinstance(model_dir, str):
|
||||
model_dir = Path(model_dir)
|
||||
|
||||
if model_dir.parent.parent.name.startswith("models--"):
|
||||
model_dir = model_dir.parent.parent
|
||||
|
||||
if model_dir.exists():
|
||||
# todo: PermissionDenied is raised on blobs removal in Windows, with blobs > 2GB
|
||||
shutil.rmtree(model_dir, onerror=on_error)
|
||||
Reference in New Issue
Block a user