mirror of
https://github.com/qdrant/fastembed.git
synced 2026-09-22 14:07:51 -05:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
9b04753f58 | ||
|
|
a107d5c994 | ||
|
|
20863c9ba9 | ||
|
|
833525a5fa | ||
|
|
05abc3744f | ||
|
|
8d590e3fb0 | ||
|
|
f4db4c5244 | ||
|
|
a2bef9821d | ||
|
|
fb68e86c23 | ||
|
|
0c63b6ab52 | ||
|
|
cbe60bf9dc | ||
|
|
40cca63d5f | ||
|
|
5c4d9b04bd | ||
|
|
a3a798f4f3 | ||
|
|
bdf6816da8 | ||
|
|
5dc53ebf49 | ||
|
|
b461314660 | ||
|
|
cc4d101828 | ||
|
|
0dab99c23e | ||
|
|
113bd565ec | ||
|
|
d5f552b6ad | ||
|
|
70c8f3cbc0 | ||
|
|
a34e7bcc42 | ||
|
|
c48247f15d | ||
|
|
f9d757ffc6 | ||
|
|
a5a702acad | ||
|
|
a0fe741532 | ||
|
|
525642bd74 | ||
|
|
39426c282e | ||
|
|
eff93e3cb3 | ||
|
|
5cbf4f8947 | ||
|
|
817fca3fab | ||
|
|
e60d5493bb | ||
|
|
b2494cb321 | ||
|
|
f613647297 | ||
|
|
50ae2dc088 | ||
|
|
0892291f75 | ||
|
|
8108e4467c | ||
|
|
8a8ea4f42f | ||
|
|
a499c313af | ||
|
|
fde1e0b361 | ||
|
|
adfc1aef59 | ||
|
|
1ec283c744 | ||
|
|
a6a4e375ca | ||
|
|
2b069d2fbd | ||
|
|
21df54e3a4 | ||
|
|
87678dd784 | ||
|
|
6fa442b960 | ||
|
|
52ebfba27c | ||
|
|
ea55268e01 | ||
|
|
800f3887b7 | ||
|
|
020d535f9c | ||
|
|
685fd9b5a1 | ||
|
|
b304a2aff0 | ||
|
|
c715416361 | ||
|
|
428381cb04 | ||
|
|
3511b08831 | ||
|
|
b718cc6a88 | ||
|
|
2ba8990260 | ||
|
|
dab185fd9d | ||
|
|
ec0e3128ee | ||
|
|
44e332999c | ||
|
|
533b54cee5 | ||
|
|
ba1f6053bd | ||
|
|
4dc76e3859 | ||
|
|
6efe06b172 | ||
|
|
887239239b | ||
|
|
ca023be0c0 | ||
|
|
d1ddc8142d | ||
|
|
faf3d9fe18 | ||
|
|
acec31277b | ||
|
|
cb902149c8 | ||
|
|
a260022ae6 | ||
|
|
d5da56299a | ||
|
|
4e5575f4a7 | ||
|
|
4736f46548 | ||
|
|
c85e8c278f | ||
|
|
0df605fc92 | ||
|
|
04bc7a3039 | ||
|
|
b785640bd5 | ||
|
|
5568a62c2f | ||
|
|
b34209dcfb | ||
|
|
c91d42dda7 | ||
|
|
aa0c475a1f | ||
|
|
4c239b11d5 | ||
|
|
6acfb001fb | ||
|
|
1729aab1ec | ||
|
|
42fca3b467 | ||
|
|
2082108baf | ||
|
|
6cda2ce7f0 | ||
|
|
5bd5c0a0f0 | ||
|
|
58ee7cc95c | ||
|
|
27eeb39473 | ||
|
|
4e527b1c63 |
@@ -39,7 +39,7 @@ body:
|
||||
attributes:
|
||||
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.
|
||||
placeholder: v0.5.1
|
||||
placeholder: v0.7.4
|
||||
validations:
|
||||
required: true
|
||||
- type: dropdown
|
||||
|
||||
@@ -0,0 +1,40 @@
|
||||
version: 2
|
||||
updates:
|
||||
- package-ecosystem: "pip"
|
||||
directory: "/"
|
||||
schedule:
|
||||
interval: "weekly"
|
||||
open-pull-requests-limit: 5
|
||||
groups:
|
||||
security-updates:
|
||||
applies-to: security-updates
|
||||
patterns:
|
||||
- "*"
|
||||
version-updates:
|
||||
applies-to: version-updates
|
||||
update-types:
|
||||
- "minor"
|
||||
- "patch"
|
||||
patterns:
|
||||
- "*"
|
||||
cooldown:
|
||||
default-days: 7
|
||||
- package-ecosystem: "github-actions"
|
||||
directory: "/"
|
||||
schedule:
|
||||
interval: "weekly"
|
||||
open-pull-requests-limit: 5
|
||||
groups:
|
||||
security-updates:
|
||||
applies-to: security-updates
|
||||
patterns:
|
||||
- "*"
|
||||
version-updates:
|
||||
applies-to: version-updates
|
||||
update-types:
|
||||
- "minor"
|
||||
- "patch"
|
||||
patterns:
|
||||
- "*"
|
||||
cooldown:
|
||||
default-days: 7
|
||||
@@ -10,12 +10,12 @@ jobs:
|
||||
deploy:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v3
|
||||
- uses: actions/setup-python@v4
|
||||
- uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
|
||||
- uses: actions/setup-python@5fda3b95a4ea91299a34e894583c3862153e4b97 # v7.0.0
|
||||
with:
|
||||
python-version: 3.x
|
||||
- run: echo "cache_id=$(date --utc '+%V')" >> $GITHUB_ENV
|
||||
- uses: actions/cache@v3
|
||||
- uses: actions/cache@55cc8345863c7cc4c66a329aec7e433d2d1c52a9 # v6.1.0
|
||||
with:
|
||||
key: mkdocs-material-${{ env.cache_id }}
|
||||
path: .cache
|
||||
|
||||
@@ -21,11 +21,11 @@ jobs:
|
||||
runs-on: ubuntu-latest
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v2
|
||||
- uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v2
|
||||
uses: actions/setup-python@5fda3b95a4ea91299a34e894583c3862153e4b97 # v7.0.0
|
||||
with:
|
||||
python-version: '3.9.x'
|
||||
python-version: '3.10.x'
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
python -m pip install poetry
|
||||
@@ -33,7 +33,7 @@ jobs:
|
||||
- name: Build package
|
||||
run: poetry build
|
||||
- name: Publish package
|
||||
uses: pypa/gh-action-pypi-publish@27b31702a0e7fc50959f5ad993c78deac1bdfc29
|
||||
uses: pypa/gh-action-pypi-publish@dc37677b2e1c63e2034f94d8a5b11f265b73ba33 # v1.14.2
|
||||
with:
|
||||
user: __token__
|
||||
password: ${{ secrets.PYPI_API_TOKEN }}
|
||||
|
||||
@@ -1,9 +1,11 @@
|
||||
name: Tests
|
||||
run-name: Tests (gpu)
|
||||
|
||||
on:
|
||||
push:
|
||||
branches: [ master, main, gpu ]
|
||||
pull_request:
|
||||
branches: [ master, main, gpu ]
|
||||
workflow_dispatch:
|
||||
|
||||
|
||||
env:
|
||||
CARGO_TERM_COLOR: always
|
||||
@@ -14,24 +16,21 @@ jobs:
|
||||
strategy:
|
||||
matrix:
|
||||
python-version:
|
||||
- '3.9.x'
|
||||
- '3.10.x'
|
||||
- '3.11.x'
|
||||
- '3.12.x'
|
||||
- '3.13.x'
|
||||
os:
|
||||
- ubuntu-latest
|
||||
- macos-latest
|
||||
- windows-latest
|
||||
|
||||
runs-on: ${{ matrix.os }}
|
||||
|
||||
name: Python ${{ matrix.python-version }} on ${{ matrix.os }} test
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v3
|
||||
- uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
uses: actions/setup-python@5fda3b95a4ea91299a34e894583c3862153e4b97 # v7.0.0
|
||||
with:
|
||||
python-version: ${{ matrix.python-version }}
|
||||
- name: Install dependencies
|
||||
@@ -41,5 +40,7 @@ jobs:
|
||||
poetry install --no-interaction --no-ansi --without dev,docs
|
||||
|
||||
- name: Run pytest
|
||||
env:
|
||||
HF_TOKEN: ${{ secrets.HF_TOKEN }}
|
||||
run: |
|
||||
poetry run pytest
|
||||
poetry run pytest
|
||||
@@ -8,16 +8,16 @@ jobs:
|
||||
strategy:
|
||||
fail-fast: true
|
||||
matrix:
|
||||
python-version: ["3.9", "3.10", "3.11", "3.12", "3.13"]
|
||||
python-version: ["3.10", "3.11", "3.12", "3.13"]
|
||||
os: [ubuntu-latest]
|
||||
|
||||
name: Python ${{ matrix.python-version }} test
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v1
|
||||
- uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
|
||||
|
||||
- name: Set up Python ${{ matrix.python-version }}
|
||||
uses: actions/setup-python@v2
|
||||
uses: actions/setup-python@5fda3b95a4ea91299a34e894583c3862153e4b97 # v7.0.0
|
||||
with:
|
||||
python-version: ${{ matrix.python-version }}
|
||||
|
||||
@@ -26,8 +26,6 @@ jobs:
|
||||
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 \
|
||||
|
||||
@@ -186,7 +186,7 @@
|
||||
same "printed page" as the copyright notice for easier
|
||||
identification within third-party archives.
|
||||
|
||||
Copyright [yyyy] [name of copyright owner]
|
||||
Copyright 2026 Qdrant Solutions GmbH
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
|
||||
@@ -15,6 +15,8 @@ These models are developed by Jina (https://jina.ai/) and are subject to Jina AI
|
||||
This distribution includes the following Google models, each with its respective license:
|
||||
- vidore/colpali-v1.3
|
||||
- License: gemma
|
||||
- google/embeddinggemma-300m
|
||||
- License: gemma
|
||||
|
||||
Gemma is provided under and subject to the Gemma Terms of Use found at ai.google.dev/gemma/terms
|
||||
|
||||
|
||||
@@ -63,6 +63,23 @@ embeddings = list(model.embed(documents))
|
||||
|
||||
```
|
||||
|
||||
Dense text embedding can also be extended with models which are not in the list of supported models.
|
||||
|
||||
```python
|
||||
from fastembed import TextEmbedding
|
||||
from fastembed.common.model_description import PoolingType, ModelSource
|
||||
|
||||
TextEmbedding.add_custom_model(
|
||||
model="intfloat/multilingual-e5-small",
|
||||
pooling=PoolingType.MEAN,
|
||||
normalization=True,
|
||||
sources=ModelSource(hf="intfloat/multilingual-e5-small"), # can be used with an `url` to load files from a private storage
|
||||
dim=384,
|
||||
model_file="onnx/model.onnx", # can be used to load an already supported model with another optimization or quantization, e.g. onnx/model_O4.onnx
|
||||
)
|
||||
model = TextEmbedding(model_name="intfloat/multilingual-e5-small")
|
||||
embeddings = list(model.embed(documents))
|
||||
```
|
||||
|
||||
|
||||
### 🔱 Sparse text embeddings
|
||||
@@ -137,6 +154,27 @@ embeddings = list(model.embed(images))
|
||||
# ]
|
||||
```
|
||||
|
||||
### Late interaction multimodal models (ColPali)
|
||||
|
||||
```python
|
||||
from fastembed import LateInteractionMultimodalEmbedding
|
||||
|
||||
doc_images = [
|
||||
"./path/to/qdrant_pdf_doc_1_screenshot.jpg",
|
||||
"./path/to/colpali_pdf_doc_2_screenshot.jpg",
|
||||
]
|
||||
|
||||
query = "What is Qdrant?"
|
||||
|
||||
model = LateInteractionMultimodalEmbedding(model_name="Qdrant/colpali-v1.3-fp16")
|
||||
doc_images_embeddings = list(model.embed_image(doc_images))
|
||||
# shape (2, 1030, 128)
|
||||
# [array([[-0.03353882, -0.02090454, ..., -0.15576172, -0.07678223]], dtype=float32)]
|
||||
query_embedding = model.embed_text(query)
|
||||
# shape (1, 20, 128)
|
||||
# [array([[-0.00218201, 0.14758301, ..., -0.02207947, 0.16833496]], dtype=float32)]
|
||||
```
|
||||
|
||||
### 🔄 Rerankers
|
||||
```python
|
||||
from fastembed.rerank.cross_encoder import TextCrossEncoder
|
||||
@@ -152,6 +190,23 @@ scores = list(encoder.rerank(query, documents))
|
||||
# [-11.48061752319336, 5.472434997558594]
|
||||
```
|
||||
|
||||
Text cross encoders can also be extended with models which are not in the list of supported models.
|
||||
|
||||
```python
|
||||
from fastembed.rerank.cross_encoder import TextCrossEncoder
|
||||
from fastembed.common.model_description import ModelSource
|
||||
|
||||
TextCrossEncoder.add_custom_model(
|
||||
model="Xenova/ms-marco-MiniLM-L-4-v2",
|
||||
model_file="onnx/model.onnx",
|
||||
sources=ModelSource(hf="Xenova/ms-marco-MiniLM-L-4-v2"),
|
||||
)
|
||||
model = TextCrossEncoder(model_name="Xenova/ms-marco-MiniLM-L-4-v2")
|
||||
scores = list(model.rerank_pairs(
|
||||
[("What is AI?", "Artificial intelligence is ..."), ("What is ML?", "Machine learning is ..."),]
|
||||
))
|
||||
```
|
||||
|
||||
## ⚡️ FastEmbed on a GPU
|
||||
|
||||
FastEmbed supports running on GPU devices.
|
||||
@@ -191,36 +246,36 @@ pip install qdrant-client[fastembed-gpu]
|
||||
You might have to use quotes ```pip install 'qdrant-client[fastembed]'``` on zsh.
|
||||
|
||||
```python
|
||||
from qdrant_client import QdrantClient
|
||||
from qdrant_client import QdrantClient, models
|
||||
|
||||
# Initialize the client
|
||||
client = QdrantClient("localhost", port=6333) # For production
|
||||
# client = QdrantClient(":memory:") # For small experiments
|
||||
# client = QdrantClient(":memory:") # For experimentation
|
||||
|
||||
# Prepare your documents, metadata, and IDs
|
||||
docs = ["Qdrant has Langchain integrations", "Qdrant also has Llama Index integrations"]
|
||||
metadata = [
|
||||
{"source": "Langchain-docs"},
|
||||
{"source": "Llama-index-docs"},
|
||||
model_name = "sentence-transformers/all-MiniLM-L6-v2"
|
||||
payload = [
|
||||
{"document": "Qdrant has Langchain integrations", "source": "Langchain-docs", },
|
||||
{"document": "Qdrant also has Llama Index integrations", "source": "LlamaIndex-docs"},
|
||||
]
|
||||
docs = [models.Document(text=data["document"], model=model_name) for data in payload]
|
||||
ids = [42, 2]
|
||||
|
||||
# If you want to change the model:
|
||||
# client.set_model("sentence-transformers/all-MiniLM-L6-v2")
|
||||
# List of supported models: https://qdrant.github.io/fastembed/examples/Supported_Models
|
||||
|
||||
# Use the new add() instead of upsert()
|
||||
# This internally calls embed() of the configured embedding model
|
||||
client.add(
|
||||
collection_name="demo_collection",
|
||||
documents=docs,
|
||||
metadata=metadata,
|
||||
ids=ids
|
||||
client.create_collection(
|
||||
"demo_collection",
|
||||
vectors_config=models.VectorParams(
|
||||
size=client.get_embedding_size(model_name), distance=models.Distance.COSINE)
|
||||
)
|
||||
|
||||
search_result = client.query(
|
||||
client.upload_collection(
|
||||
collection_name="demo_collection",
|
||||
query_text="This is a query document"
|
||||
vectors=docs,
|
||||
ids=ids,
|
||||
payload=payload,
|
||||
)
|
||||
|
||||
search_result = client.query_points(
|
||||
collection_name="demo_collection",
|
||||
query=models.Document(text="This is a query document", model=model_name)
|
||||
).points
|
||||
print(search_result)
|
||||
```
|
||||
|
||||
Binary file not shown.
|
Before Width: | Height: | Size: 1.8 KiB After Width: | Height: | Size: 2.0 KiB |
@@ -0,0 +1,134 @@
|
||||
"""Export an inference-free SPLADE document encoder to ONNX.
|
||||
|
||||
Converts `opensearch-project/opensearch-neural-sparse-encoding-doc-v3-gte` (an MLM head
|
||||
over a GTE backbone) into an onnx model producing token logits, and assembles a model dir
|
||||
with everything fastembed's `IfSplade` needs: model.onnx, tokenizer files and idf.json.
|
||||
|
||||
Usage:
|
||||
python experiments/if_splade_to_onnx.py --output-dir models/opensearch-neural-sparse-encoding-doc-v3-gte
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import shutil
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
from huggingface_hub import hf_hub_download
|
||||
from transformers import AutoModelForMaskedLM, AutoTokenizer
|
||||
|
||||
MODEL_ID = "opensearch-project/opensearch-neural-sparse-encoding-doc-v3-gte"
|
||||
# revision of the remote modeling code (Alibaba-NLP/new-impl), pinned in the model card
|
||||
CODE_REVISION = "40ced75c3017eb27626c9d4ea981bde21a2662f4"
|
||||
|
||||
TOKENIZER_FILES = [
|
||||
"config.json",
|
||||
"tokenizer.json",
|
||||
"tokenizer_config.json",
|
||||
"special_tokens_map.json",
|
||||
"vocab.txt",
|
||||
"idf.json",
|
||||
]
|
||||
|
||||
|
||||
class LogitsOnly(torch.nn.Module):
|
||||
def __init__(self, model: torch.nn.Module):
|
||||
super().__init__()
|
||||
self.model = model
|
||||
|
||||
def forward(self, input_ids: torch.Tensor, attention_mask: torch.Tensor) -> torch.Tensor:
|
||||
return self.model(input_ids=input_ids, attention_mask=attention_mask).logits
|
||||
|
||||
|
||||
def export(model_id: str, output_dir: Path, opset: int = 14) -> Path:
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
model = AutoModelForMaskedLM.from_pretrained(
|
||||
model_id, trust_remote_code=True, code_revision=CODE_REVISION
|
||||
)
|
||||
model.eval()
|
||||
wrapped = LogitsOnly(model)
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(model_id)
|
||||
dummy = tokenizer(
|
||||
["fastembed is a library", "onnx export"],
|
||||
padding=True,
|
||||
truncation=True,
|
||||
return_tensors="pt",
|
||||
return_token_type_ids=False,
|
||||
)
|
||||
|
||||
onnx_path = output_dir / "model.onnx"
|
||||
with torch.inference_mode():
|
||||
torch.onnx.export(
|
||||
wrapped,
|
||||
(dummy["input_ids"], dummy["attention_mask"]),
|
||||
f=onnx_path.as_posix(),
|
||||
input_names=["input_ids", "attention_mask"],
|
||||
output_names=["logits"],
|
||||
dynamic_axes={
|
||||
"input_ids": {0: "batch_size", 1: "sequence_length"},
|
||||
"attention_mask": {0: "batch_size", 1: "sequence_length"},
|
||||
"logits": {0: "batch_size", 1: "sequence_length"},
|
||||
},
|
||||
do_constant_folding=True,
|
||||
opset_version=opset,
|
||||
dynamo=False,
|
||||
)
|
||||
|
||||
for file_name in TOKENIZER_FILES:
|
||||
local_path = hf_hub_download(repo_id=model_id, filename=file_name)
|
||||
shutil.copy(local_path, output_dir / file_name)
|
||||
|
||||
return onnx_path
|
||||
|
||||
|
||||
def parity_check(model_id: str, output_dir: Path) -> None:
|
||||
import numpy as np
|
||||
import onnxruntime as ort
|
||||
|
||||
model = AutoModelForMaskedLM.from_pretrained(
|
||||
model_id, trust_remote_code=True, code_revision=CODE_REVISION
|
||||
)
|
||||
model.eval()
|
||||
tokenizer = AutoTokenizer.from_pretrained(model_id)
|
||||
|
||||
documents = [
|
||||
"Currently New York is rainy.",
|
||||
"fastembed is a lightweight library for generating embeddings",
|
||||
"hello world",
|
||||
]
|
||||
features = tokenizer(
|
||||
documents, padding=True, truncation=True, return_tensors="pt", return_token_type_ids=False
|
||||
)
|
||||
|
||||
with torch.inference_mode():
|
||||
torch_logits = model(**features).logits.numpy()
|
||||
|
||||
session = ort.InferenceSession(output_dir / "model.onnx")
|
||||
onnx_logits = session.run(
|
||||
["logits"],
|
||||
{
|
||||
"input_ids": features["input_ids"].numpy(),
|
||||
"attention_mask": features["attention_mask"].numpy(),
|
||||
},
|
||||
)[0]
|
||||
|
||||
max_diff = np.abs(torch_logits - onnx_logits).max()
|
||||
print(f"max |torch - onnx| logits diff: {max_diff}")
|
||||
assert max_diff < 1e-3, "onnx export does not match the torch model"
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--model-id", default=MODEL_ID)
|
||||
parser.add_argument("--output-dir", default=f"models/{MODEL_ID.replace('/', '_')}", type=Path)
|
||||
parser.add_argument("--opset", default=14, type=int)
|
||||
args = parser.parse_args()
|
||||
|
||||
onnx_path = export(args.model_id, args.output_dir, args.opset)
|
||||
print(f"Exported to {onnx_path}")
|
||||
parity_check(args.model_id, args.output_dir)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,11 +1,17 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional, Any
|
||||
from enum import Enum
|
||||
from typing import Any
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ModelSource:
|
||||
hf: Optional[str] = None
|
||||
url: Optional[str] = None
|
||||
hf: str | None = None
|
||||
url: str | None = None
|
||||
_deprecated_tar_struct: bool = False
|
||||
|
||||
@property
|
||||
def deprecated_tar_struct(self) -> bool:
|
||||
return self._deprecated_tar_struct
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if self.hf is None and self.url is None:
|
||||
@@ -27,8 +33,8 @@ class BaseModelDescription:
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class DenseModelDescription(BaseModelDescription):
|
||||
dim: Optional[int] = None
|
||||
tasks: Optional[dict[str, Any]] = None
|
||||
dim: int | None = None
|
||||
tasks: dict[str, Any] | None = field(default_factory=dict)
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
assert self.dim is not None, "dim is required for dense model description"
|
||||
@@ -36,5 +42,12 @@ class DenseModelDescription(BaseModelDescription):
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SparseModelDescription(BaseModelDescription):
|
||||
requires_idf: Optional[bool] = None
|
||||
vocab_size: Optional[int] = None
|
||||
requires_idf: bool | None = None
|
||||
vocab_size: int | None = None
|
||||
|
||||
|
||||
class PoolingType(str, Enum):
|
||||
CLS = "CLS"
|
||||
MEAN = "MEAN"
|
||||
LAST_TOKEN = "LAST_TOKEN"
|
||||
DISABLED = "DISABLED"
|
||||
|
||||
@@ -1,10 +1,14 @@
|
||||
import os
|
||||
import time
|
||||
import gzip
|
||||
import json
|
||||
import shutil
|
||||
import tarfile
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional, Union, TypeVar, Generic
|
||||
import tempfile
|
||||
import contextlib
|
||||
from copy import deepcopy
|
||||
from pathlib import Path, PureWindowsPath
|
||||
from typing import Any, TypeVar, Generic
|
||||
|
||||
import requests
|
||||
from huggingface_hub import snapshot_download, model_info, list_repo_tree
|
||||
@@ -20,6 +24,8 @@ from fastembed.common.model_description import BaseModelDescription
|
||||
|
||||
T = TypeVar("T", bound=BaseModelDescription)
|
||||
|
||||
_DOWNLOAD_CHUNK_SIZE = 256 * 1024
|
||||
|
||||
|
||||
class ModelManagement(Generic[T]):
|
||||
METADATA_FILE = "files_metadata.json"
|
||||
@@ -33,6 +39,31 @@ class ModelManagement(Generic[T]):
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
@classmethod
|
||||
def add_custom_model(
|
||||
cls,
|
||||
*args: Any,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Add a custom model to the existing embedding classes based on the passed model descriptions
|
||||
|
||||
Model description dict should contain the fields same as in one of the model descriptions presented
|
||||
in fastembed.common.model_description
|
||||
|
||||
E.g. for BaseModelDescription:
|
||||
model: str
|
||||
sources: ModelSource
|
||||
model_file: str
|
||||
description: str
|
||||
license: str
|
||||
size_in_GB: float
|
||||
additional_files: list[str]
|
||||
|
||||
Returns:
|
||||
None
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
@classmethod
|
||||
def _list_supported_models(cls) -> list[T]:
|
||||
raise NotImplementedError()
|
||||
@@ -71,9 +102,7 @@ class ModelManagement(Generic[T]):
|
||||
str: The path to the downloaded file.
|
||||
"""
|
||||
|
||||
if os.path.exists(output_path):
|
||||
return output_path
|
||||
response = requests.get(url, stream=True)
|
||||
response = requests.get(url, stream=True, timeout=(10, 120))
|
||||
|
||||
# Handle HTTP errors
|
||||
if response.status_code == 403:
|
||||
@@ -81,6 +110,8 @@ class ModelManagement(Generic[T]):
|
||||
"Authentication Error: You do not have permission to access this resource. "
|
||||
"Please check your credentials."
|
||||
)
|
||||
# Otherwise an error page gets written out as though it were the archive.
|
||||
response.raise_for_status()
|
||||
|
||||
# Get the total size of the file
|
||||
total_size_in_bytes = int(response.headers.get("content-length", 0))
|
||||
@@ -98,7 +129,7 @@ class ModelManagement(Generic[T]):
|
||||
disable=not show_progress,
|
||||
) as progress_bar:
|
||||
with open(output_path, "wb") as file:
|
||||
for chunk in response.iter_content(chunk_size=1024):
|
||||
for chunk in response.iter_content(chunk_size=_DOWNLOAD_CHUNK_SIZE):
|
||||
if chunk: # Filter out keep-alive new chunks
|
||||
progress_bar.update(len(chunk))
|
||||
file.write(chunk)
|
||||
@@ -154,8 +185,8 @@ class ModelManagement(Generic[T]):
|
||||
|
||||
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]]] = {}
|
||||
) -> dict[str, dict[str, int | str]]:
|
||||
meta: dict[str, dict[str, 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:
|
||||
@@ -167,9 +198,7 @@ class ModelManagement(Generic[T]):
|
||||
}
|
||||
return meta
|
||||
|
||||
def _save_file_metadata(
|
||||
model_dir: Path, meta: dict[str, dict[str, Union[int, str]]]
|
||||
) -> None:
|
||||
def _save_file_metadata(model_dir: Path, meta: dict[str, dict[str, int | str]]) -> None:
|
||||
try:
|
||||
if not model_dir.exists():
|
||||
model_dir.mkdir(parents=True, exist_ok=True)
|
||||
@@ -199,11 +228,6 @@ class ModelManagement(Generic[T]):
|
||||
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,
|
||||
@@ -267,12 +291,18 @@ class ModelManagement(Generic[T]):
|
||||
"""
|
||||
Decompresses a .tar.gz file to a cache directory.
|
||||
|
||||
Nothing is deleted on failure, since `cache_dir` may hold more than this archive.
|
||||
Cleaning up a partial extraction is the caller's job.
|
||||
|
||||
Args:
|
||||
targz_path (str): Path to the .tar.gz file.
|
||||
cache_dir (str): Path to the cache directory.
|
||||
|
||||
Returns:
|
||||
cache_dir (str): Path to the cache directory.
|
||||
|
||||
Raises:
|
||||
ValueError: If the archive is missing, corrupt, or holds an unsafe member.
|
||||
"""
|
||||
# Check if targz_path exists and is a file
|
||||
if not os.path.isfile(targz_path):
|
||||
@@ -285,60 +315,68 @@ class ModelManagement(Generic[T]):
|
||||
try:
|
||||
# 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,
|
||||
)
|
||||
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)
|
||||
# and raise the error again
|
||||
if "tmp" in cache_dir:
|
||||
shutil.rmtree(cache_dir)
|
||||
raise ValueError(f"An error occurred while decompressing {targz_path}: {e}")
|
||||
if hasattr(tarfile, "data_filter"):
|
||||
tar.extractall(path=cache_dir, filter="data")
|
||||
else:
|
||||
# No PEP 706 filter before 3.10.12, so vet the members by hand.
|
||||
members = tar.getmembers()
|
||||
for member in members:
|
||||
cls._validate_tar_member(member)
|
||||
tar.extractall(path=cache_dir, members=members)
|
||||
# tarfile stops at the end-of-archive marker, short of the gzip trailer, so
|
||||
# the CRC is only checked if the rest of the stream is read.
|
||||
while tar.fileobj.read(1 << 20):
|
||||
pass
|
||||
except (tarfile.TarError, ValueError, EOFError, gzip.BadGzipFile) as e:
|
||||
# gzip raises EOFError for a truncated stream and BadGzipFile for a corrupted one.
|
||||
raise ValueError(f"An error occurred while decompressing {targz_path}: {e}") from e
|
||||
|
||||
return cache_dir
|
||||
|
||||
@staticmethod
|
||||
def _is_unsafe_tar_path(path: str) -> bool:
|
||||
"""Checks whether a tar member name or link target may escape the extraction dir.
|
||||
|
||||
Lexical on purpose: resolving against the extraction directory is unsound before
|
||||
extraction, since `link/../escape` only escapes once an earlier member has been
|
||||
written as a symlink. Any `..` component is therefore rejected outright.
|
||||
"""
|
||||
# PureWindowsPath splits on both separators, so `root` covers POSIX "/evil" as
|
||||
# well as "\\evil", which escapes on Windows without being absolute.
|
||||
windows_path = PureWindowsPath(path)
|
||||
return bool(windows_path.drive or windows_path.root) or ".." in windows_path.parts
|
||||
|
||||
@classmethod
|
||||
def _validate_tar_member(cls, member: tarfile.TarInfo) -> None:
|
||||
"""Raises ValueError if a member could write outside the extraction directory."""
|
||||
if cls._is_unsafe_tar_path(member.name):
|
||||
raise ValueError(f"Unsafe tar member path: {member.name}")
|
||||
|
||||
if member.issym() or member.islnk():
|
||||
if cls._is_unsafe_tar_path(member.linkname):
|
||||
raise ValueError(f"Unsafe tar link target: {member.name} -> {member.linkname}")
|
||||
elif not (member.isfile() or member.isdir()):
|
||||
# Devices, fifos and the like have no place in a model archive.
|
||||
raise ValueError(f"Unsupported tar member type: {member.name}")
|
||||
|
||||
@classmethod
|
||||
def retrieve_model_gcs(
|
||||
cls,
|
||||
model_name: str,
|
||||
source_url: str,
|
||||
cache_dir: str,
|
||||
deprecated_tar_struct: bool = False,
|
||||
local_files_only: bool = False,
|
||||
) -> Path:
|
||||
fast_model_name = f"fast-{model_name.split('/')[-1]}"
|
||||
fast_model_name = f"{'fast-' if deprecated_tar_struct else ''}{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
|
||||
|
||||
# check if the model_dir and the model files are both present for macOS
|
||||
if model_dir.exists() and len(list(model_dir.glob("*"))) > 0:
|
||||
return model_dir
|
||||
|
||||
if model_tmp_dir.exists():
|
||||
shutil.rmtree(model_tmp_dir)
|
||||
|
||||
cache_tmp_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
model_tar_gz = Path(cache_dir) / f"{fast_model_name}.tar.gz"
|
||||
|
||||
if model_tar_gz.exists():
|
||||
model_tar_gz.unlink()
|
||||
|
||||
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}"
|
||||
|
||||
model_tar_gz.unlink()
|
||||
# Rename from tmp to final name is atomic
|
||||
model_tmp_dir.rename(model_dir)
|
||||
else:
|
||||
if local_files_only:
|
||||
logger.error(
|
||||
f"Could not find the model tar.gz file at {model_dir} and local_files_only=True."
|
||||
)
|
||||
@@ -346,6 +384,45 @@ class ModelManagement(Generic[T]):
|
||||
f"Could not find the model tar.gz file at {model_dir} and local_files_only=True."
|
||||
)
|
||||
|
||||
if cache_tmp_dir.is_symlink():
|
||||
raise ValueError(
|
||||
f"{cache_tmp_dir} is a symlink, refusing to stage downloads through it"
|
||||
)
|
||||
cache_tmp_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# The archive and everything extracted from it go in a directory of this attempt's own,
|
||||
# so removing it undoes the attempt without touching any other download of the model.
|
||||
staging_dir = Path(tempfile.mkdtemp(dir=cache_tmp_dir, prefix=f"{fast_model_name}-"))
|
||||
try:
|
||||
model_tar_gz = staging_dir / f"{fast_model_name}.tar.gz"
|
||||
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(staging_dir))
|
||||
|
||||
model_tmp_dir = staging_dir / fast_model_name
|
||||
if not model_tmp_dir.is_dir() or model_tmp_dir.is_symlink():
|
||||
raise ValueError(
|
||||
f"The archive from {source_url} has no {fast_model_name} directory"
|
||||
)
|
||||
|
||||
# Replace a stale empty model_dir, which Windows will not rename onto. rmdir leaves
|
||||
# anything else alone, including one another download has just filled.
|
||||
with contextlib.suppress(OSError):
|
||||
model_dir.rmdir()
|
||||
|
||||
try:
|
||||
# Rename from the staging dir to the final name is atomic
|
||||
model_tmp_dir.rename(model_dir)
|
||||
except OSError:
|
||||
# Another download of the same model finished first, so keep its copy.
|
||||
if not (model_dir.is_dir() and any(model_dir.iterdir())):
|
||||
raise
|
||||
finally:
|
||||
shutil.rmtree(staging_dir, ignore_errors=True)
|
||||
|
||||
return model_dir
|
||||
|
||||
@classmethod
|
||||
@@ -375,21 +452,47 @@ class ModelManagement(Generic[T]):
|
||||
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)
|
||||
hf_offline = os.environ.get("HF_HUB_OFFLINE", "").strip().upper()
|
||||
if not local_files_only and hf_offline in {"1", "TRUE", "YES", "ON"}:
|
||||
local_files_only = True
|
||||
kwargs["local_files_only"] = True
|
||||
specific_model_path: str | None = 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
|
||||
|
||||
extra_patterns = [model.model_file]
|
||||
extra_patterns.extend(model.additional_files)
|
||||
|
||||
if hf_source:
|
||||
try:
|
||||
cache_kwargs = deepcopy(kwargs)
|
||||
cache_kwargs["local_files_only"] = True
|
||||
resolved_path = Path(
|
||||
cls.download_files_from_huggingface(
|
||||
hf_source,
|
||||
cache_dir=cache_dir,
|
||||
extra_patterns=extra_patterns,
|
||||
**cache_kwargs,
|
||||
)
|
||||
)
|
||||
if (resolved_path / model.model_file).exists() and all(
|
||||
(resolved_path / file).exists() for file in extra_patterns
|
||||
):
|
||||
return resolved_path
|
||||
except Exception:
|
||||
pass
|
||||
finally:
|
||||
enable_progress_bars()
|
||||
|
||||
sleep = 3.0
|
||||
while retries > 0:
|
||||
retries -= 1
|
||||
|
||||
if hf_source:
|
||||
extra_patterns = [model.model_file]
|
||||
extra_patterns.extend(model.additional_files)
|
||||
|
||||
if hf_source and not local_files_only:
|
||||
# we have already tried loading with `local_files_only=True` via hf and we failed
|
||||
try:
|
||||
return Path(
|
||||
cls.download_files_from_huggingface(
|
||||
@@ -413,6 +516,7 @@ class ModelManagement(Generic[T]):
|
||||
model.model,
|
||||
str(url_source),
|
||||
str(cache_dir),
|
||||
deprecated_tar_struct=model.sources.deprecated_tar_struct,
|
||||
local_files_only=local_files_only,
|
||||
)
|
||||
except Exception:
|
||||
@@ -421,11 +525,12 @@ class ModelManagement(Generic[T]):
|
||||
|
||||
if local_files_only:
|
||||
logger.error("Could not find model in cache_dir")
|
||||
break
|
||||
else:
|
||||
logger.error(
|
||||
f"Could not download model from either source, sleeping for {sleep} seconds, {retries} retries left."
|
||||
)
|
||||
time.sleep(sleep)
|
||||
sleep *= 3
|
||||
time.sleep(sleep)
|
||||
sleep *= 3
|
||||
|
||||
raise ValueError(f"Could not load model {model.model} from any source.")
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import warnings
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any, Generic, Iterable, Optional, Sequence, Type, TypeVar
|
||||
from typing import Any, Generic, Iterable, Sequence, Type, TypeVar
|
||||
|
||||
import numpy as np
|
||||
import onnxruntime as ort
|
||||
@@ -9,7 +9,7 @@ import onnxruntime as ort
|
||||
from numpy.typing import NDArray
|
||||
from tokenizers import Tokenizer
|
||||
|
||||
from fastembed.common.types import OnnxProvider, NumpyArray
|
||||
from fastembed.common.types import OnnxProvider, NumpyArray, Device
|
||||
from fastembed.parallel_processor import Worker
|
||||
|
||||
# Holds type of the embedding result
|
||||
@@ -19,21 +19,45 @@ T = TypeVar("T")
|
||||
@dataclass
|
||||
class OnnxOutputContext:
|
||||
model_output: NumpyArray
|
||||
attention_mask: Optional[NDArray[np.int64]] = None
|
||||
input_ids: Optional[NDArray[np.int64]] = None
|
||||
attention_mask: NDArray[np.int64] | None = None
|
||||
input_ids: NDArray[np.int64] | None = None
|
||||
metadata: dict[str, Any] | None = None
|
||||
|
||||
|
||||
class OnnxModel(Generic[T]):
|
||||
EXPOSED_SESSION_OPTIONS = ("enable_cpu_mem_arena",)
|
||||
|
||||
@classmethod
|
||||
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]:
|
||||
def _get_worker_init_kwargs(self) -> dict[str, Any]:
|
||||
"""Additional kwargs a worker process needs to reconstruct this model.
|
||||
|
||||
Workers are started with `spawn`/`forkserver`, hence they don't inherit class-level state
|
||||
which has been set up in runtime, e.g. models registered via `add_custom_model`.
|
||||
Such state has to be shipped to the workers explicitly.
|
||||
|
||||
Returns:
|
||||
dict[str, Any]: kwargs to pass to `_get_worker_class().init_embedding`.
|
||||
"""
|
||||
return {}
|
||||
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext, **kwargs: Any) -> Iterable[T]:
|
||||
"""Post-process the ONNX model output to convert it into a usable format.
|
||||
|
||||
Args:
|
||||
output (OnnxOutputContext): The raw output from the ONNX model.
|
||||
**kwargs: Additional keyword arguments that may be needed by specific implementations.
|
||||
|
||||
Returns:
|
||||
Iterable[T]: Post-processed output as an iterable of type T.
|
||||
"""
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.model: Optional[ort.InferenceSession] = None
|
||||
self.tokenizer: Optional[Tokenizer] = None
|
||||
self.model: ort.InferenceSession | None = None
|
||||
self.tokenizer: Tokenizer | None = None
|
||||
|
||||
def _preprocess_onnx_input(
|
||||
self, onnx_input: dict[str, NumpyArray], **kwargs: Any
|
||||
@@ -47,24 +71,30 @@ class OnnxModel(Generic[T]):
|
||||
self,
|
||||
model_dir: Path,
|
||||
model_file: str,
|
||||
threads: Optional[int],
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
cuda: bool = False,
|
||||
device_id: Optional[int] = None,
|
||||
threads: int | None,
|
||||
providers: Sequence[OnnxProvider] | None = None,
|
||||
cuda: bool | Device = Device.AUTO,
|
||||
device_id: int | None = None,
|
||||
extra_session_options: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
model_path = model_dir / model_file
|
||||
# List of Execution Providers: https://onnxruntime.ai/docs/execution-providers
|
||||
available_providers = ort.get_available_providers()
|
||||
cuda_available = "CUDAExecutionProvider" in available_providers
|
||||
explicit_cuda = cuda is True or cuda == Device.CUDA
|
||||
|
||||
if cuda and providers is not None:
|
||||
if explicit_cuda and providers is not None:
|
||||
warnings.warn(
|
||||
f"`cuda` and `providers` are mutually exclusive parameters, cuda: {cuda}, providers: {providers}",
|
||||
f"`cuda` and `providers` are mutually exclusive parameters, "
|
||||
f"cuda: {cuda}, providers: {providers}. If you'd like to use providers, cuda should be one of "
|
||||
f"[False, Device.CPU, Device.AUTO].",
|
||||
category=UserWarning,
|
||||
stacklevel=6,
|
||||
)
|
||||
|
||||
if providers is not None:
|
||||
onnx_providers = list(providers)
|
||||
elif cuda:
|
||||
elif explicit_cuda or (cuda == Device.AUTO and cuda_available):
|
||||
if device_id is None:
|
||||
onnx_providers = ["CUDAExecutionProvider"]
|
||||
else:
|
||||
@@ -72,7 +102,6 @@ class OnnxModel(Generic[T]):
|
||||
else:
|
||||
onnx_providers = ["CPUExecutionProvider"]
|
||||
|
||||
available_providers = ort.get_available_providers()
|
||||
requested_provider_names: list[str] = []
|
||||
for provider in onnx_providers:
|
||||
# check providers available
|
||||
@@ -90,6 +119,9 @@ class OnnxModel(Generic[T]):
|
||||
so.intra_op_num_threads = threads
|
||||
so.inter_op_num_threads = threads
|
||||
|
||||
if extra_session_options is not None:
|
||||
self.add_extra_session_options(so, extra_session_options)
|
||||
|
||||
self.model = ort.InferenceSession(
|
||||
str(model_path), providers=onnx_providers, sess_options=so
|
||||
)
|
||||
@@ -104,6 +136,38 @@ class OnnxModel(Generic[T]):
|
||||
RuntimeWarning,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _select_exposed_session_options(cls, model_kwargs: dict[str, Any]) -> dict[str, Any]:
|
||||
"""A convenience method to select the exposed session options in models
|
||||
|
||||
Args:
|
||||
model_kwargs (dict[str, Any]): The model kwargs.
|
||||
|
||||
Returns:
|
||||
dict[str, Any]: a dict with filtered exposed session options.
|
||||
"""
|
||||
return {k: v for k, v in model_kwargs.items() if k in cls.EXPOSED_SESSION_OPTIONS}
|
||||
|
||||
@classmethod
|
||||
def add_extra_session_options(
|
||||
cls, session_options: ort.SessionOptions, extra_options: dict[str, Any]
|
||||
) -> None:
|
||||
"""Add extra session options to the existing options object in-place
|
||||
|
||||
Args:
|
||||
session_options (ort.SessionOptions): The existing session options object.
|
||||
extra_options (dict[str, Any]): The extra session options available in cls.EXPOSED_SESSION_OPTIONS.
|
||||
|
||||
Returns:
|
||||
None
|
||||
"""
|
||||
for option in extra_options:
|
||||
assert (
|
||||
option in cls.EXPOSED_SESSION_OPTIONS
|
||||
), f"{option} is unknown or not exposed (exposed options: {cls.EXPOSED_SESSION_OPTIONS})"
|
||||
if "enable_cpu_mem_arena" in extra_options:
|
||||
session_options.enable_cpu_mem_arena = extra_options["enable_cpu_mem_arena"]
|
||||
|
||||
def load_onnx_model(self) -> None:
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import json
|
||||
from typing import Any
|
||||
import sys
|
||||
from typing import Any, Iterator
|
||||
from pathlib import Path
|
||||
|
||||
from tokenizers import AddedToken, Tokenizer
|
||||
@@ -8,9 +9,10 @@ from fastembed.image.transform.operators import Compose
|
||||
|
||||
|
||||
def load_special_tokens(model_dir: Path) -> dict[str, Any]:
|
||||
"""Read special_tokens_map.json, treating an absent file as an empty map."""
|
||||
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}")
|
||||
return {}
|
||||
|
||||
with open(str(tokens_map_path)) as tokens_map_file:
|
||||
tokens_map = json.load(tokens_map_file)
|
||||
@@ -18,11 +20,58 @@ def load_special_tokens(model_dir: Path) -> dict[str, Any]:
|
||||
return tokens_map
|
||||
|
||||
|
||||
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}")
|
||||
def iter_special_tokens(tokens_map: dict[str, Any]) -> Iterator[str | dict[str, Any]]:
|
||||
"""Yield the individual tokens declared in a special tokens map.
|
||||
|
||||
Most keys hold one token, but `additional_special_tokens` holds a list of them.
|
||||
"""
|
||||
for value in tokens_map.values():
|
||||
if isinstance(value, list):
|
||||
yield from value
|
||||
else:
|
||||
yield value
|
||||
|
||||
|
||||
def _valid_context(value: Any) -> int | None:
|
||||
"""Return `value` if it can be used as a truncation limit, `None` otherwise.
|
||||
|
||||
Config files do not always carry a real limit: transformers writes `model_max_length` as
|
||||
1e30 when the value is unknown, and some repos ship a 0 or a null. `enable_truncation`
|
||||
raises an `OverflowError` on the former and silently produces empty encodings on the
|
||||
latter, so both are rejected here rather than passed through.
|
||||
"""
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
return None
|
||||
if not 0 < value <= sys.maxsize:
|
||||
return None
|
||||
return value
|
||||
|
||||
|
||||
def _resolve_max_context(tokenizer_config: dict[str, Any], model_dir: Path) -> int:
|
||||
"""Pick the truncation limit, preferring the stricter of the two tokenizer config keys.
|
||||
|
||||
`config.json:max_position_embeddings` deliberately is not used as a fallback: it is the size
|
||||
of the position table, not the usable context, and the two differ per architecture, e.g.
|
||||
roberta reports 514 for a usable 512.
|
||||
"""
|
||||
candidates = [
|
||||
context
|
||||
for context in (
|
||||
_valid_context(tokenizer_config.get("model_max_length")),
|
||||
_valid_context(tokenizer_config.get("max_length")),
|
||||
)
|
||||
if context is not None
|
||||
]
|
||||
if not candidates:
|
||||
raise ValueError(
|
||||
f"Could not determine the maximum context length for {model_dir}. Set a positive "
|
||||
"`model_max_length` or `max_length` in tokenizer_config.json."
|
||||
)
|
||||
|
||||
return min(candidates)
|
||||
|
||||
|
||||
def load_tokenizer(model_dir: Path) -> tuple[Tokenizer, dict[str, int]]:
|
||||
tokenizer_path = model_dir / "tokenizer.json"
|
||||
if not tokenizer_path.exists():
|
||||
raise ValueError(f"Could not find tokenizer.json in {model_dir}")
|
||||
@@ -31,43 +80,62 @@ def load_tokenizer(model_dir: Path) -> tuple[Tokenizer, dict[str, int]]:
|
||||
if not tokenizer_config_path.exists():
|
||||
raise ValueError(f"Could not find tokenizer_config.json in {model_dir}")
|
||||
|
||||
with open(str(config_path)) as config_file:
|
||||
config = json.load(config_file)
|
||||
# config.json is optional: transformers v5 no longer writes it for every model.
|
||||
config_path = model_dir / "config.json"
|
||||
config: dict[str, Any] = {}
|
||||
if config_path.exists():
|
||||
with open(str(config_path)) as config_file:
|
||||
config = json.load(config_file)
|
||||
|
||||
with open(str(tokenizer_config_path)) as tokenizer_config_file:
|
||||
tokenizer_config = json.load(tokenizer_config_file)
|
||||
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"])
|
||||
|
||||
max_context = _resolve_max_context(tokenizer_config, model_dir)
|
||||
|
||||
tokens_map = load_special_tokens(model_dir)
|
||||
|
||||
tokenizer = Tokenizer.from_file(str(tokenizer_path))
|
||||
tokenizer.enable_truncation(max_length=max_context)
|
||||
tokenizer.enable_padding(
|
||||
pad_id=config.get("pad_token_id", 0), pad_token=tokenizer_config["pad_token"]
|
||||
)
|
||||
|
||||
for token in tokens_map.values():
|
||||
# Registered before the padding is resolved: the map may name a pad token that
|
||||
# tokenizer.json does not carry, and it only gets an id once it is added.
|
||||
for token in iter_special_tokens(tokens_map):
|
||||
if isinstance(token, str):
|
||||
tokenizer.add_special_tokens([token])
|
||||
elif isinstance(token, dict):
|
||||
tokenizer.add_special_tokens([AddedToken(**token)])
|
||||
|
||||
special_token_to_id: dict[str, int] = {}
|
||||
# Padding is always normalized to batch-longest. A serialized fixed length shorter than the
|
||||
# truncation limit leaves longer encodings untouched, which produces ragged batches, and a
|
||||
# fixed length equal to it pads every batch to the maximum. Direction and pad token metadata
|
||||
# are taken from the serialized settings, since some models pad on the left.
|
||||
padding = tokenizer.padding or {}
|
||||
pad_token = padding.get("pad_token") or tokenizer_config.get("pad_token")
|
||||
if pad_token is None:
|
||||
raise ValueError(f"Could not find a pad token for {model_dir}")
|
||||
|
||||
for token in tokens_map.values():
|
||||
if isinstance(token, str):
|
||||
special_token_to_id[token] = tokenizer.token_to_id(token)
|
||||
elif isinstance(token, dict):
|
||||
token_str = token.get("content", "")
|
||||
special_token_to_id[token_str] = tokenizer.token_to_id(token_str)
|
||||
# The vocabulary is the last resort, not a hardcoded 0: that silently disagrees with
|
||||
# `pad_token` for every model whose pad token is not the first entry.
|
||||
pad_id = padding.get("pad_id", config.get("pad_token_id"))
|
||||
if pad_id is None:
|
||||
pad_id = tokenizer.token_to_id(pad_token)
|
||||
if pad_id is None:
|
||||
raise ValueError(f"Could not resolve an id for the pad token {pad_token!r} in {model_dir}")
|
||||
|
||||
tokenizer.enable_padding(
|
||||
direction=padding.get("direction", "right"),
|
||||
pad_id=pad_id,
|
||||
pad_type_id=padding.get("pad_type_id", 0),
|
||||
pad_token=pad_token,
|
||||
pad_to_multiple_of=padding.get("pad_to_multiple_of"),
|
||||
length=None,
|
||||
)
|
||||
|
||||
special_token_to_id = {
|
||||
token.content: token_id
|
||||
for token_id, token in tokenizer.get_added_tokens_decoder().items()
|
||||
if token.special
|
||||
}
|
||||
|
||||
return tokenizer, special_token_to_id
|
||||
|
||||
|
||||
+21
-18
@@ -1,24 +1,27 @@
|
||||
from enum import Enum
|
||||
from pathlib import Path
|
||||
import sys
|
||||
from PIL import Image
|
||||
from typing import Any, Union
|
||||
from typing import Any, TypeAlias
|
||||
|
||||
import numpy as np
|
||||
from numpy.typing import NDArray
|
||||
|
||||
if sys.version_info >= (3, 10):
|
||||
from typing import TypeAlias
|
||||
else:
|
||||
from typing_extensions import TypeAlias
|
||||
from PIL import Image
|
||||
|
||||
|
||||
PathInput: TypeAlias = Union[str, Path]
|
||||
ImageInput: TypeAlias = Union[PathInput, Image.Image]
|
||||
class Device(str, Enum):
|
||||
CPU = "cpu"
|
||||
CUDA = "cuda"
|
||||
AUTO = "auto"
|
||||
|
||||
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],
|
||||
]
|
||||
|
||||
PathInput: TypeAlias = str | Path
|
||||
ImageInput: TypeAlias = PathInput | Image.Image
|
||||
|
||||
OnnxProvider: TypeAlias = str | tuple[str, dict[Any, Any]]
|
||||
NumpyArray: TypeAlias = (
|
||||
NDArray[np.float64]
|
||||
| NDArray[np.float32]
|
||||
| NDArray[np.float16]
|
||||
| NDArray[np.int8]
|
||||
| NDArray[np.int64]
|
||||
| NDArray[np.int32]
|
||||
)
|
||||
|
||||
@@ -5,9 +5,10 @@ import tempfile
|
||||
import unicodedata
|
||||
from pathlib import Path
|
||||
from itertools import islice
|
||||
from typing import Iterable, Optional, TypeVar
|
||||
from typing import Iterable, TypeVar
|
||||
|
||||
import numpy as np
|
||||
from numpy.typing import NDArray
|
||||
|
||||
from fastembed.common.types import NumpyArray
|
||||
|
||||
@@ -22,6 +23,25 @@ def normalize(input_array: NumpyArray, p: int = 2, dim: int = 1, eps: float = 1e
|
||||
return normalized_array
|
||||
|
||||
|
||||
def mean_pooling(input_array: NumpyArray, attention_mask: NDArray[np.int64]) -> NumpyArray:
|
||||
input_mask_expanded = np.expand_dims(attention_mask, axis=-1).astype(np.int64)
|
||||
input_mask_expanded = np.tile(input_mask_expanded, (1, 1, input_array.shape[-1]))
|
||||
sum_embeddings = np.sum(input_array * 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
|
||||
|
||||
|
||||
def last_token_pooling(input_array: NumpyArray, attention_mask: NDArray[np.int64]) -> NumpyArray:
|
||||
"""Take the embedding of the last non-padding token of each sequence.
|
||||
|
||||
Locates the last position the attention mask marks as real, so it holds whichever
|
||||
side the tokenizer pads on.
|
||||
"""
|
||||
last_token_indices = attention_mask.shape[1] - 1 - np.argmax(attention_mask[:, ::-1], axis=1)
|
||||
return input_array[np.arange(input_array.shape[0]), last_token_indices]
|
||||
|
||||
|
||||
def iter_batch(iterable: Iterable[T], size: int) -> Iterable[list[T]]:
|
||||
"""
|
||||
>>> list(iter_batch([1,2,3,4,5], 3))
|
||||
@@ -35,7 +55,7 @@ def iter_batch(iterable: Iterable[T], size: int) -> Iterable[list[T]]:
|
||||
yield b
|
||||
|
||||
|
||||
def define_cache_dir(cache_dir: Optional[str] = None) -> Path:
|
||||
def define_cache_dir(cache_dir: str | None = None) -> Path:
|
||||
"""
|
||||
Define the cache directory for fastembed
|
||||
"""
|
||||
|
||||
@@ -1,16 +0,0 @@
|
||||
{
|
||||
"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, Any
|
||||
from typing import Any
|
||||
|
||||
from loguru import logger
|
||||
|
||||
@@ -17,8 +17,8 @@ class JinaEmbedding(TextEmbedding):
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str = "jinaai/jina-embeddings-v2-base-en",
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
cache_dir: str | None = None,
|
||||
threads: int | None = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
super().__init__(model_name, cache_dir, threads, **kwargs)
|
||||
|
||||
@@ -1,15 +1,21 @@
|
||||
from typing import Any, Iterable, Optional, Sequence, Type, Union
|
||||
from typing import Any, Iterable, Sequence, Type
|
||||
from dataclasses import asdict
|
||||
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.common.types import NumpyArray, Device
|
||||
from fastembed.common import ImageInput, OnnxProvider
|
||||
from fastembed.image.image_embedding_base import ImageEmbeddingBase
|
||||
from fastembed.image.onnx_embedding import OnnxImageEmbedding
|
||||
from fastembed.image.normalized_embedding import NormalizedEmbedding
|
||||
from fastembed.image.siglip_embedding import SiglipOnnxImageEmbedding
|
||||
from fastembed.common.model_description import DenseModelDescription
|
||||
|
||||
|
||||
class ImageEmbedding(ImageEmbeddingBase):
|
||||
EMBEDDINGS_REGISTRY: list[Type[ImageEmbeddingBase]] = [OnnxImageEmbedding]
|
||||
EMBEDDINGS_REGISTRY: list[Type[ImageEmbeddingBase]] = [
|
||||
OnnxImageEmbedding,
|
||||
NormalizedEmbedding,
|
||||
SiglipOnnxImageEmbedding,
|
||||
]
|
||||
|
||||
@classmethod
|
||||
def list_supported_models(cls) -> list[dict[str, Any]]:
|
||||
@@ -48,11 +54,11 @@ class ImageEmbedding(ImageEmbeddingBase):
|
||||
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,
|
||||
cache_dir: str | None = None,
|
||||
threads: int | None = None,
|
||||
providers: Sequence[OnnxProvider] | None = None,
|
||||
cuda: bool | Device = Device.AUTO,
|
||||
device_ids: list[int] | None = None,
|
||||
lazy_load: bool = False,
|
||||
**kwargs: Any,
|
||||
):
|
||||
@@ -77,11 +83,45 @@ class ImageEmbedding(ImageEmbeddingBase):
|
||||
"Please check the supported models using `ImageEmbedding.list_supported_models()`"
|
||||
)
|
||||
|
||||
@property
|
||||
def embedding_size(self) -> int:
|
||||
"""Get the embedding size of the current model"""
|
||||
if self._embedding_size is None:
|
||||
self._embedding_size = self.get_embedding_size(self.model_name)
|
||||
return self._embedding_size
|
||||
|
||||
@classmethod
|
||||
def get_embedding_size(cls, model_name: str) -> int:
|
||||
"""Get the embedding size of the passed model
|
||||
|
||||
Args:
|
||||
model_name (str): The name of the model to get embedding size for.
|
||||
|
||||
Returns:
|
||||
int: The size of the embedding.
|
||||
|
||||
Raises:
|
||||
ValueError: If the model name is not found in the supported models.
|
||||
"""
|
||||
descriptions = cls._list_supported_models()
|
||||
embedding_size: int | None = None
|
||||
for description in descriptions:
|
||||
if description.model.lower() == model_name.lower():
|
||||
embedding_size = description.dim
|
||||
break
|
||||
if embedding_size is None:
|
||||
model_names = [description.model for description in descriptions]
|
||||
raise ValueError(
|
||||
f"Embedding size for model {model_name} was None. "
|
||||
f"Available model names: {model_names}"
|
||||
)
|
||||
return embedding_size
|
||||
|
||||
def embed(
|
||||
self,
|
||||
images: Union[ImageInput, Iterable[ImageInput]],
|
||||
images: ImageInput | Iterable[ImageInput],
|
||||
batch_size: int = 16,
|
||||
parallel: Optional[int] = None,
|
||||
parallel: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from typing import Iterable, Optional, Any, Union
|
||||
from typing import Iterable, Any
|
||||
|
||||
from fastembed.common.model_description import DenseModelDescription
|
||||
from fastembed.common.types import NumpyArray
|
||||
@@ -10,20 +10,21 @@ class ImageEmbeddingBase(ModelManagement[DenseModelDescription]):
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
cache_dir: str | None = None,
|
||||
threads: int | None = 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)
|
||||
self._embedding_size: int | None = None
|
||||
|
||||
def embed(
|
||||
self,
|
||||
images: Union[ImageInput, Iterable[ImageInput]],
|
||||
images: ImageInput | Iterable[ImageInput],
|
||||
batch_size: int = 16,
|
||||
parallel: Optional[int] = None,
|
||||
parallel: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
@@ -42,3 +43,13 @@ class ImageEmbeddingBase(ModelManagement[DenseModelDescription]):
|
||||
Iterable[NdArray]: The embeddings.
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
@classmethod
|
||||
def get_embedding_size(cls, model_name: str) -> int:
|
||||
"""Returns embedding size of the chosen model."""
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
@property
|
||||
def embedding_size(self) -> int:
|
||||
"""Returns embedding size for the current model"""
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
@@ -0,0 +1,69 @@
|
||||
from typing import Any, Iterable, Type
|
||||
|
||||
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.common.utils import normalize
|
||||
from fastembed.image.onnx_embedding import OnnxImageEmbedding
|
||||
from fastembed.image.onnx_image_model import ImageEmbeddingWorker
|
||||
from fastembed.common.model_description import DenseModelDescription, ModelSource
|
||||
|
||||
supported_normalized_models: list[DenseModelDescription] = [
|
||||
DenseModelDescription(
|
||||
model="nomic-ai/nomic-embed-vision-v1.5",
|
||||
dim=768,
|
||||
description="Image embeddings, Multimodal (text&image), 2024 year",
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.37,
|
||||
sources=ModelSource(hf="nomic-ai/nomic-embed-vision-v1.5"),
|
||||
model_file="onnx/model.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="nomic-ai/nomic-embed-vision-v1.5-Q",
|
||||
dim=768,
|
||||
description="Image embeddings, Multimodal (text&image), 2024 year",
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.1,
|
||||
sources=ModelSource(hf="nomic-ai/nomic-embed-vision-v1.5"),
|
||||
model_file="onnx/model_quantized.onnx",
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
class NormalizedEmbedding(OnnxImageEmbedding):
|
||||
@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_normalized_models
|
||||
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type["ImageEmbeddingWorker[NumpyArray]"]:
|
||||
return NormalizedEmbeddingWorker
|
||||
|
||||
def _post_process_onnx_output(
|
||||
self, output: OnnxOutputContext, **kwargs: Any
|
||||
) -> Iterable[NumpyArray]:
|
||||
# The model emits last_hidden_state, which onnx_embed flattens to (batch, tokens * dim).
|
||||
# Recover the token axis, take the CLS token (index 0) and normalize, matching the reference
|
||||
# F.normalize(last_hidden_state[:, 0], p=2, dim=1).
|
||||
dim = self.model_description.dim
|
||||
assert dim is not None, "Model description is missing the embedding dim"
|
||||
hidden_states = output.model_output.reshape(output.model_output.shape[0], -1, dim)
|
||||
return normalize(hidden_states[:, 0])
|
||||
|
||||
|
||||
class NormalizedEmbeddingWorker(ImageEmbeddingWorker[NumpyArray]):
|
||||
def init_embedding(
|
||||
self, model_name: str, cache_dir: str, **kwargs: Any
|
||||
) -> NormalizedEmbedding:
|
||||
return NormalizedEmbedding(
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
threads=1,
|
||||
**kwargs,
|
||||
)
|
||||
@@ -1,8 +1,7 @@
|
||||
from typing import Any, Iterable, Optional, Sequence, Type, Union
|
||||
from typing import Any, Iterable, Sequence, Type
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.common.types import NumpyArray, Device
|
||||
from fastembed.common import ImageInput, OnnxProvider
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.common.utils import define_cache_dir, normalize
|
||||
@@ -64,14 +63,14 @@ 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,
|
||||
cuda: bool = False,
|
||||
device_ids: Optional[list[int]] = None,
|
||||
cache_dir: str | None = None,
|
||||
threads: int | None = None,
|
||||
providers: Sequence[OnnxProvider] | None = None,
|
||||
cuda: bool | Device = Device.AUTO,
|
||||
device_ids: list[int] | None = None,
|
||||
lazy_load: bool = False,
|
||||
device_id: Optional[int] = None,
|
||||
specific_model_path: Optional[str] = None,
|
||||
device_id: int | None = None,
|
||||
specific_model_path: str | None = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
"""
|
||||
@@ -83,10 +82,11 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
|
||||
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.
|
||||
cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`
|
||||
Defaults to Device.AUTO.
|
||||
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.
|
||||
workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, 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.
|
||||
@@ -99,13 +99,14 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
|
||||
super().__init__(model_name, cache_dir, threads, **kwargs)
|
||||
self.providers = providers
|
||||
self.lazy_load = lazy_load
|
||||
self._extra_session_options = self._select_exposed_session_options(kwargs)
|
||||
|
||||
# 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
|
||||
self.device_id: int | None = None
|
||||
if device_id is not None:
|
||||
self.device_id = device_id
|
||||
elif self.device_ids is not None:
|
||||
@@ -113,11 +114,12 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
|
||||
|
||||
self.model_description = self._get_model_description(model_name)
|
||||
self.cache_dir = str(define_cache_dir(cache_dir))
|
||||
self._specific_model_path = specific_model_path
|
||||
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,
|
||||
specific_model_path=self._specific_model_path,
|
||||
)
|
||||
|
||||
if not self.lazy_load:
|
||||
@@ -134,6 +136,7 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_id=self.device_id,
|
||||
extra_session_options=self._extra_session_options,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@@ -148,9 +151,9 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
|
||||
|
||||
def embed(
|
||||
self,
|
||||
images: Union[ImageInput, Iterable[ImageInput]],
|
||||
images: ImageInput | Iterable[ImageInput],
|
||||
batch_size: int = 16,
|
||||
parallel: Optional[int] = None,
|
||||
parallel: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
@@ -178,6 +181,9 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_ids=self.device_ids,
|
||||
local_files_only=self._local_files_only,
|
||||
specific_model_path=self._specific_model_path,
|
||||
extra_session_options=self._extra_session_options,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@@ -194,8 +200,10 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
|
||||
|
||||
return onnx_input
|
||||
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[NumpyArray]:
|
||||
return normalize(output.model_output).astype(np.float32)
|
||||
def _post_process_onnx_output(
|
||||
self, output: OnnxOutputContext, **kwargs: Any
|
||||
) -> Iterable[NumpyArray]:
|
||||
return normalize(output.model_output)
|
||||
|
||||
|
||||
class OnnxImageEmbeddingWorker(ImageEmbeddingWorker[NumpyArray]):
|
||||
|
||||
@@ -2,13 +2,13 @@ import contextlib
|
||||
import os
|
||||
from multiprocessing import get_all_start_methods
|
||||
from pathlib import Path
|
||||
from typing import Any, Iterable, Optional, Sequence, Type, Union
|
||||
from typing import Any, Iterable, Sequence, Type
|
||||
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
|
||||
from fastembed.image.transform.operators import Compose
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.common.types import NumpyArray, Device
|
||||
from fastembed.common import ImageInput, OnnxProvider
|
||||
from fastembed.common.onnx_model import EmbeddingWorker, OnnxModel, OnnxOutputContext, T
|
||||
from fastembed.common.preprocessor_utils import load_preprocessor
|
||||
@@ -19,16 +19,27 @@ from fastembed.parallel_processor import ParallelWorkerPool
|
||||
|
||||
|
||||
class OnnxImageModel(OnnxModel[T]):
|
||||
ONNX_OUTPUT_NAMES: list[str] | None = None
|
||||
|
||||
@classmethod
|
||||
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]:
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext, **kwargs: Any) -> Iterable[T]:
|
||||
"""Post-process the ONNX model output to convert it into a usable format.
|
||||
|
||||
Args:
|
||||
output (OnnxOutputContext): The raw output from the ONNX model.
|
||||
**kwargs: Additional keyword arguments that may be needed by specific implementations.
|
||||
|
||||
Returns:
|
||||
Iterable[T]: Post-processed output as an iterable of type T.
|
||||
"""
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.processor: Optional[Compose] = None
|
||||
self.processor: Compose | None = None
|
||||
|
||||
def _preprocess_onnx_input(
|
||||
self, onnx_input: dict[str, NumpyArray], **kwargs: Any
|
||||
@@ -42,10 +53,11 @@ class OnnxImageModel(OnnxModel[T]):
|
||||
self,
|
||||
model_dir: Path,
|
||||
model_file: str,
|
||||
threads: Optional[int],
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
cuda: bool = False,
|
||||
device_id: Optional[int] = None,
|
||||
threads: int | None,
|
||||
providers: Sequence[OnnxProvider] | None = None,
|
||||
cuda: bool | Device = Device.AUTO,
|
||||
device_id: int | None = None,
|
||||
extra_session_options: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
super()._load_onnx_model(
|
||||
model_dir=model_dir,
|
||||
@@ -54,6 +66,7 @@ class OnnxImageModel(OnnxModel[T]):
|
||||
providers=providers,
|
||||
cuda=cuda,
|
||||
device_id=device_id,
|
||||
extra_session_options=extra_session_options,
|
||||
)
|
||||
self.processor = load_preprocessor(model_dir=model_dir)
|
||||
|
||||
@@ -65,16 +78,18 @@ class OnnxImageModel(OnnxModel[T]):
|
||||
return {input_name: encoded}
|
||||
|
||||
def onnx_embed(self, images: list[ImageInput], **kwargs: Any) -> OnnxOutputContext:
|
||||
with contextlib.ExitStack():
|
||||
with contextlib.ExitStack() as stack:
|
||||
image_files = [
|
||||
Image.open(image) if not isinstance(image, Image.Image) else image
|
||||
stack.enter_context(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) # type: ignore[union-attr]
|
||||
model_output = self.model.run(self.ONNX_OUTPUT_NAMES, onnx_input) # type: ignore[union-attr]
|
||||
embeddings = model_output[0].reshape(len(images), -1)
|
||||
return OnnxOutputContext(model_output=embeddings)
|
||||
|
||||
@@ -82,12 +97,15 @@ class OnnxImageModel(OnnxModel[T]):
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
images: Union[ImageInput, Iterable[ImageInput]],
|
||||
images: ImageInput | Iterable[ImageInput],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
cuda: bool = False,
|
||||
device_ids: Optional[list[int]] = None,
|
||||
parallel: int | None = None,
|
||||
providers: Sequence[OnnxProvider] | None = None,
|
||||
cuda: bool | Device = Device.AUTO,
|
||||
device_ids: list[int] | None = None,
|
||||
local_files_only: bool = False,
|
||||
specific_model_path: str | None = None,
|
||||
extra_session_options: dict[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[T]:
|
||||
is_small = False
|
||||
@@ -104,7 +122,7 @@ class OnnxImageModel(OnnxModel[T]):
|
||||
self.load_onnx_model()
|
||||
|
||||
for batch in iter_batch(images, batch_size):
|
||||
yield from self._post_process_onnx_output(self.onnx_embed(batch))
|
||||
yield from self._post_process_onnx_output(self.onnx_embed(batch), **kwargs)
|
||||
else:
|
||||
if parallel == 0:
|
||||
parallel = os.cpu_count()
|
||||
@@ -114,9 +132,14 @@ class OnnxImageModel(OnnxModel[T]):
|
||||
"model_name": model_name,
|
||||
"cache_dir": cache_dir,
|
||||
"providers": providers,
|
||||
"local_files_only": local_files_only,
|
||||
"specific_model_path": specific_model_path,
|
||||
**kwargs,
|
||||
}
|
||||
|
||||
if extra_session_options is not None:
|
||||
params.update(extra_session_options)
|
||||
|
||||
pool = ParallelWorkerPool(
|
||||
num_workers=parallel or 1,
|
||||
worker=self._get_worker_class(),
|
||||
@@ -125,7 +148,7 @@ class OnnxImageModel(OnnxModel[T]):
|
||||
start_method=start_method,
|
||||
)
|
||||
for batch in pool.ordered_map(iter_batch(images, batch_size), **params):
|
||||
yield from self._post_process_onnx_output(batch) # type: ignore
|
||||
yield from self._post_process_onnx_output(batch, **kwargs) # type: ignore
|
||||
|
||||
|
||||
class ImageEmbeddingWorker(EmbeddingWorker[T]):
|
||||
|
||||
@@ -0,0 +1,44 @@
|
||||
from typing import Any, Type
|
||||
|
||||
from fastembed.common.model_description import DenseModelDescription, ModelSource
|
||||
from fastembed.image.onnx_embedding import OnnxImageEmbedding, OnnxImageEmbeddingWorker
|
||||
|
||||
supported_siglip_models: list[DenseModelDescription] = [
|
||||
DenseModelDescription(
|
||||
model="google/siglip2-base-patch16-224",
|
||||
dim=768,
|
||||
description="Image embeddings, Multimodal (text&image), 2025 year",
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.37,
|
||||
sources=ModelSource(hf="onnx-community/siglip2-base-patch16-224-ONNX"),
|
||||
model_file="onnx/vision_model.onnx",
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
class SiglipOnnxImageEmbedding(OnnxImageEmbedding):
|
||||
"""SigLIP vision tower.
|
||||
|
||||
The exported graph returns both the per-patch `last_hidden_state` and the pooled
|
||||
`pooler_output`; only the latter is the image embedding, so it must be selected explicitly.
|
||||
"""
|
||||
|
||||
ONNX_OUTPUT_NAMES = ["pooler_output"]
|
||||
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type["OnnxImageEmbeddingWorker"]:
|
||||
return SiglipImageEmbeddingWorker
|
||||
|
||||
@classmethod
|
||||
def _list_supported_models(cls) -> list[DenseModelDescription]:
|
||||
return supported_siglip_models
|
||||
|
||||
|
||||
class SiglipImageEmbeddingWorker(OnnxImageEmbeddingWorker):
|
||||
def init_embedding(self, model_name: str, cache_dir: str, **kwargs: Any) -> OnnxImageEmbedding:
|
||||
return SiglipOnnxImageEmbedding(
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
threads=1,
|
||||
**kwargs,
|
||||
)
|
||||
@@ -1,5 +1,3 @@
|
||||
from typing import Union
|
||||
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
|
||||
@@ -15,7 +13,7 @@ def convert_to_rgb(image: Image.Image) -> Image.Image:
|
||||
|
||||
|
||||
def center_crop(
|
||||
image: Union[Image.Image, NumpyArray],
|
||||
image: Image.Image | NumpyArray,
|
||||
size: tuple[int, int],
|
||||
) -> NumpyArray:
|
||||
if isinstance(image, np.ndarray):
|
||||
@@ -64,43 +62,56 @@ def center_crop(
|
||||
|
||||
def normalize(
|
||||
image: NumpyArray,
|
||||
mean: Union[float, list[float]],
|
||||
std: Union[float, list[float]],
|
||||
mean: float | list[float],
|
||||
std: float | list[float],
|
||||
) -> NumpyArray:
|
||||
num_channels = image.shape[1] if len(image.shape) == 4 else image.shape[0]
|
||||
if image.ndim < 3:
|
||||
raise ValueError(f"image must be (C, H, W) or (N, C, H, W), got shape {image.shape}")
|
||||
|
||||
# Channels sit on the third axis from the end, which covers (C, H, W) and
|
||||
# (N, C, H, W) alike. Transposing instead reversed every axis, which put the
|
||||
# batch dimension where the channels were meant to be.
|
||||
num_channels = image.shape[-3]
|
||||
|
||||
if not np.issubdtype(image.dtype, np.floating):
|
||||
image = image.astype(np.float32)
|
||||
|
||||
mean = mean if isinstance(mean, list) else [mean] * num_channels
|
||||
mean_list = mean if isinstance(mean, list) else [mean] * num_channels
|
||||
|
||||
if len(mean) != num_channels:
|
||||
if len(mean_list) != 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)}"
|
||||
f"{len(mean_list)}"
|
||||
)
|
||||
|
||||
mean_arr = np.array(mean, dtype=np.float32)
|
||||
# (C, 1, 1) lines the channels up with the trailing (C, H, W) axes under numpy
|
||||
# broadcasting, whatever batch dimensions lead them.
|
||||
mean_arr = np.array(mean_list, dtype=np.float32).reshape(-1, 1, 1)
|
||||
|
||||
std = std if isinstance(std, list) else [std] * num_channels
|
||||
if len(std) != num_channels:
|
||||
std_list = std if isinstance(std, list) else [std] * num_channels
|
||||
if len(std_list) != num_channels:
|
||||
raise ValueError(
|
||||
f"std must have the same number of channels as the image, image has {num_channels} channels, got {len(std)}"
|
||||
f"std must have the same number of channels as the image, image has {num_channels} channels, got {len(std_list)}"
|
||||
)
|
||||
|
||||
std_arr = np.array(std, dtype=np.float32)
|
||||
std_arr = np.array(std_list, dtype=np.float32).reshape(-1, 1, 1)
|
||||
|
||||
image = ((image.T - mean_arr) / std_arr).T
|
||||
return image
|
||||
image_upd = (image - mean_arr) / std_arr
|
||||
return image_upd
|
||||
|
||||
|
||||
def resize(
|
||||
image: Image.Image,
|
||||
size: Union[int, tuple[int, int]],
|
||||
resample: Union[int, Image.Resampling] = Image.Resampling.BILINEAR,
|
||||
size: int | tuple[int, int],
|
||||
resample: int | Image.Resampling = Image.Resampling.BILINEAR,
|
||||
) -> Image.Image:
|
||||
if isinstance(size, tuple):
|
||||
return image.resize(size, resample)
|
||||
# fastembed keeps sizes as (height, width) — `Compose.from_config` builds the
|
||||
# tuple as (size["height"], size["width"]) — while Pillow's resize takes
|
||||
# (width, height). The two agree for square sizes, so this only shows up on a
|
||||
# non-square image processor configuration.
|
||||
height, width = size
|
||||
return image.resize((width, height), resample)
|
||||
|
||||
height, width = image.height, image.width
|
||||
short, long = (width, height) if width <= height else (height, width)
|
||||
@@ -117,7 +128,7 @@ def rescale(image: NumpyArray, scale: float, dtype: type = np.float32) -> NumpyA
|
||||
return (image * scale).astype(dtype)
|
||||
|
||||
|
||||
def pil2ndarray(image: Union[Image.Image, NumpyArray]) -> NumpyArray:
|
||||
def pil2ndarray(image: Image.Image | NumpyArray) -> NumpyArray:
|
||||
if isinstance(image, Image.Image):
|
||||
return np.asarray(image).transpose((2, 0, 1))
|
||||
return image
|
||||
@@ -126,7 +137,7 @@ def pil2ndarray(image: Union[Image.Image, NumpyArray]) -> NumpyArray:
|
||||
def pad2square(
|
||||
image: Image.Image,
|
||||
size: int,
|
||||
fill_color: Union[str, int, tuple[int, ...]] = 0,
|
||||
fill_color: str | int | tuple[int, ...] = 0,
|
||||
) -> Image.Image:
|
||||
height, width = image.height, image.width
|
||||
|
||||
@@ -147,3 +158,77 @@ def pad2square(
|
||||
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
|
||||
|
||||
|
||||
def resize_longest_edge(
|
||||
image: Image.Image,
|
||||
max_size: int,
|
||||
resample: int | Image.Resampling = Image.Resampling.LANCZOS,
|
||||
) -> Image.Image:
|
||||
height, width = image.height, image.width
|
||||
aspect_ratio = width / height
|
||||
|
||||
if width >= height:
|
||||
# Width is longer
|
||||
new_width = max_size
|
||||
new_height = int(new_width / aspect_ratio)
|
||||
else:
|
||||
# Height is longer
|
||||
new_height = max_size
|
||||
new_width = int(new_height * aspect_ratio)
|
||||
|
||||
# Ensure even dimensions
|
||||
if new_height % 2 != 0:
|
||||
new_height += 1
|
||||
if new_width % 2 != 0:
|
||||
new_width += 1
|
||||
|
||||
return image.resize((new_width, new_height), resample)
|
||||
|
||||
|
||||
def crop_ndarray(
|
||||
image: NumpyArray,
|
||||
x1: int,
|
||||
y1: int,
|
||||
x2: int,
|
||||
y2: int,
|
||||
channel_first: bool = True,
|
||||
) -> NumpyArray:
|
||||
if channel_first:
|
||||
# (C, H, W) format
|
||||
return image[:, y1:y2, x1:x2]
|
||||
else:
|
||||
# (H, W, C) format
|
||||
return image[y1:y2, x1:x2, :]
|
||||
|
||||
|
||||
def resize_ndarray(
|
||||
image: NumpyArray,
|
||||
size: tuple[int, int],
|
||||
resample: int | Image.Resampling = Image.Resampling.LANCZOS,
|
||||
channel_first: bool = True,
|
||||
) -> NumpyArray:
|
||||
# Convert to PIL-friendly format (H, W, C)
|
||||
if channel_first:
|
||||
img_hwc = image.transpose((1, 2, 0))
|
||||
else:
|
||||
img_hwc = image
|
||||
|
||||
# Handle different dtypes
|
||||
if img_hwc.dtype == np.float32 or img_hwc.dtype == np.float64:
|
||||
# Assume normalized, scale to 0-255 for PIL
|
||||
img_hwc_scaled = (img_hwc * 255).astype(np.uint8)
|
||||
pil_img = Image.fromarray(img_hwc_scaled, mode="RGB")
|
||||
resized = pil_img.resize(size, resample)
|
||||
result = np.array(resized).astype(np.float32) / 255.0
|
||||
else:
|
||||
# uint8 or similar
|
||||
pil_img = Image.fromarray(img_hwc.astype(np.uint8), mode="RGB")
|
||||
resized = pil_img.resize(size, resample)
|
||||
result = np.array(resized)
|
||||
|
||||
# Convert back to original format
|
||||
if channel_first:
|
||||
result = result.transpose((2, 0, 1))
|
||||
|
||||
return result
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
from typing import Any, Union, Optional
|
||||
from typing import Any
|
||||
import math
|
||||
|
||||
from PIL import Image
|
||||
|
||||
@@ -6,16 +7,19 @@ from fastembed.common.types import NumpyArray
|
||||
from fastembed.image.transform.functional import (
|
||||
center_crop,
|
||||
convert_to_rgb,
|
||||
crop_ndarray,
|
||||
normalize,
|
||||
pil2ndarray,
|
||||
rescale,
|
||||
resize,
|
||||
resize_longest_edge,
|
||||
resize_ndarray,
|
||||
pad2square,
|
||||
)
|
||||
|
||||
|
||||
class Transform:
|
||||
def __call__(self, images: list[Any]) -> Union[list[Image.Image], list[NumpyArray]]:
|
||||
def __call__(self, images: list[Any]) -> list[Image.Image] | list[NumpyArray]:
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
|
||||
@@ -33,18 +37,28 @@ class CenterCrop(Transform):
|
||||
|
||||
|
||||
class Normalize(Transform):
|
||||
def __init__(self, mean: Union[float, list[float]], std: Union[float, list[float]]):
|
||||
def __init__(self, mean: float | list[float], std: float | list[float]):
|
||||
self.mean = mean
|
||||
self.std = std
|
||||
|
||||
def __call__(self, images: list[NumpyArray]) -> list[NumpyArray]:
|
||||
return [normalize(image, mean=self.mean, std=self.std) for image in images]
|
||||
def __call__( # type: ignore[override]
|
||||
self, images: list[NumpyArray] | list[list[NumpyArray]]
|
||||
) -> list[NumpyArray] | list[list[NumpyArray]]:
|
||||
if images and isinstance(images[0], list):
|
||||
# Nested structure from ImageSplitter
|
||||
return [
|
||||
[normalize(image, mean=self.mean, std=self.std) for image in img_patches] # type: ignore[arg-type]
|
||||
for img_patches in images
|
||||
]
|
||||
else:
|
||||
# Flat structure (backward compatibility)
|
||||
return [normalize(image, mean=self.mean, std=self.std) for image in images] # type: ignore[arg-type]
|
||||
|
||||
|
||||
class Resize(Transform):
|
||||
def __init__(
|
||||
self,
|
||||
size: Union[int, tuple[int, int]],
|
||||
size: int | tuple[int, int],
|
||||
resample: Image.Resampling = Image.Resampling.BICUBIC,
|
||||
):
|
||||
self.size = size
|
||||
@@ -58,12 +72,22 @@ class Rescale(Transform):
|
||||
def __init__(self, scale: float = 1 / 255):
|
||||
self.scale = scale
|
||||
|
||||
def __call__(self, images: list[NumpyArray]) -> list[NumpyArray]:
|
||||
return [rescale(image, scale=self.scale) for image in images]
|
||||
def __call__( # type: ignore[override]
|
||||
self, images: list[NumpyArray] | list[list[NumpyArray]]
|
||||
) -> list[NumpyArray] | list[list[NumpyArray]]:
|
||||
if images and isinstance(images[0], list):
|
||||
# Nested structure from ImageSplitter
|
||||
return [
|
||||
[rescale(image, scale=self.scale) for image in img_patches] # type: ignore[arg-type]
|
||||
for img_patches in images
|
||||
]
|
||||
else:
|
||||
# Flat structure (backward compatibility)
|
||||
return [rescale(image, scale=self.scale) for image in images] # type: ignore[arg-type]
|
||||
|
||||
|
||||
class PILtoNDarray(Transform):
|
||||
def __call__(self, images: list[Union[Image.Image, NumpyArray]]) -> list[NumpyArray]:
|
||||
def __call__(self, images: list[Image.Image | NumpyArray]) -> list[NumpyArray]:
|
||||
return [pil2ndarray(image) for image in images]
|
||||
|
||||
|
||||
@@ -71,7 +95,7 @@ class PadtoSquare(Transform):
|
||||
def __init__(
|
||||
self,
|
||||
size: int,
|
||||
fill_color: Union[str, int, tuple[int, ...]],
|
||||
fill_color: str | int | tuple[int, ...],
|
||||
):
|
||||
self.size = size
|
||||
self.fill_color = fill_color
|
||||
@@ -82,13 +106,174 @@ class PadtoSquare(Transform):
|
||||
]
|
||||
|
||||
|
||||
class ResizeLongestEdge(Transform):
|
||||
"""Resize images so the longest edge equals target size, preserving aspect ratio."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size: int,
|
||||
resample: Image.Resampling = Image.Resampling.LANCZOS,
|
||||
):
|
||||
self.size = size
|
||||
self.resample = resample
|
||||
|
||||
def __call__(self, images: list[Image.Image]) -> list[Image.Image]:
|
||||
return [resize_longest_edge(image, self.size, self.resample) for image in images]
|
||||
|
||||
|
||||
class ResizeForVisionEncoder(Transform):
|
||||
"""
|
||||
Resize both dimensions to be multiples of vision_encoder_max_size.
|
||||
Preserves aspect ratio approximately.
|
||||
Works on numpy arrays in (C, H, W) format.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
max_size: int,
|
||||
resample: Image.Resampling = Image.Resampling.LANCZOS,
|
||||
):
|
||||
self.max_size = max_size
|
||||
self.resample = resample
|
||||
|
||||
def __call__(self, images: list[NumpyArray]) -> list[NumpyArray]:
|
||||
result = []
|
||||
for image in images:
|
||||
# Assume (C, H, W) format
|
||||
_, height, width = image.shape
|
||||
|
||||
aspect_ratio = width / height
|
||||
|
||||
if width >= height:
|
||||
# Calculate new width as multiple of max_size
|
||||
new_width = math.ceil(width / self.max_size) * self.max_size
|
||||
new_height = int(new_width / aspect_ratio)
|
||||
new_height = math.ceil(new_height / self.max_size) * self.max_size
|
||||
else:
|
||||
# Calculate new height as multiple of max_size
|
||||
new_height = math.ceil(height / self.max_size) * self.max_size
|
||||
new_width = int(new_height * aspect_ratio)
|
||||
new_width = math.ceil(new_width / self.max_size) * self.max_size
|
||||
|
||||
# Resize using the ndarray resize function
|
||||
resized = resize_ndarray(
|
||||
image,
|
||||
size=(new_width, new_height), # PIL expects (width, height)
|
||||
resample=self.resample,
|
||||
channel_first=True,
|
||||
)
|
||||
result.append(resized)
|
||||
|
||||
return result
|
||||
|
||||
|
||||
class ImageSplitter(Transform):
|
||||
"""
|
||||
Split images into grid of patches plus a global view.
|
||||
|
||||
If image dimensions exceed max_size:
|
||||
- Divide into ceil(H/max_size) x ceil(W/max_size) patches
|
||||
- Each patch is cropped from the image
|
||||
- Add a global view (original resized to max_size x max_size)
|
||||
|
||||
If image is smaller than max_size:
|
||||
- Return single image unchanged
|
||||
|
||||
Works on numpy arrays in (C, H, W) format.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
max_size: int,
|
||||
resample: Image.Resampling = Image.Resampling.LANCZOS,
|
||||
):
|
||||
self.max_size = max_size
|
||||
self.resample = resample
|
||||
|
||||
def __call__(self, images: list[NumpyArray]) -> list[list[NumpyArray]]: # type: ignore[override]
|
||||
result = []
|
||||
|
||||
for image in images:
|
||||
# Assume (C, H, W) format
|
||||
_, height, width = image.shape
|
||||
max_height = max_width = self.max_size
|
||||
|
||||
frames = []
|
||||
|
||||
if height > max_height or width > max_width:
|
||||
# Calculate the number of splits needed
|
||||
num_splits_h = math.ceil(height / max_height)
|
||||
num_splits_w = math.ceil(width / max_width)
|
||||
|
||||
# Calculate optimal patch dimensions
|
||||
optimal_height = math.ceil(height / num_splits_h)
|
||||
optimal_width = math.ceil(width / num_splits_w)
|
||||
|
||||
# Generate patches in grid order (row by row)
|
||||
for r in range(num_splits_h):
|
||||
for c in range(num_splits_w):
|
||||
# Calculate crop coordinates
|
||||
start_x = c * optimal_width
|
||||
start_y = r * optimal_height
|
||||
end_x = min(start_x + optimal_width, width)
|
||||
end_y = min(start_y + optimal_height, height)
|
||||
|
||||
# Crop the patch
|
||||
cropped = crop_ndarray(
|
||||
image, x1=start_x, y1=start_y, x2=end_x, y2=end_y, channel_first=True
|
||||
)
|
||||
frames.append(cropped)
|
||||
|
||||
# Add global view (resized to max_size x max_size)
|
||||
global_view = resize_ndarray(
|
||||
image,
|
||||
size=(max_width, max_height), # PIL expects (width, height)
|
||||
resample=self.resample,
|
||||
channel_first=True,
|
||||
)
|
||||
frames.append(global_view)
|
||||
else:
|
||||
# Image is small enough, no splitting needed
|
||||
frames.append(image)
|
||||
|
||||
# Append (not extend) to preserve per-image grouping
|
||||
result.append(frames)
|
||||
|
||||
return result
|
||||
|
||||
|
||||
class SquareResize(Transform):
|
||||
"""
|
||||
Resize images to square dimensions (max_size x max_size).
|
||||
Works on numpy arrays in (C, H, W) format.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size: int,
|
||||
resample: Image.Resampling = Image.Resampling.LANCZOS,
|
||||
):
|
||||
self.size = size
|
||||
self.resample = resample
|
||||
|
||||
def __call__(self, images: list[NumpyArray]) -> list[list[NumpyArray]]: # type: ignore[override]
|
||||
return [
|
||||
[
|
||||
resize_ndarray(
|
||||
image, size=(self.size, self.size), resample=self.resample, channel_first=True
|
||||
)
|
||||
]
|
||||
for image in images
|
||||
]
|
||||
|
||||
|
||||
class Compose:
|
||||
def __init__(self, transforms: list[Transform]):
|
||||
self.transforms = transforms
|
||||
|
||||
def __call__(
|
||||
self, images: Union[list[Image.Image], list[NumpyArray]]
|
||||
) -> Union[list[NumpyArray], list[Image.Image]]:
|
||||
self, images: list[Image.Image] | list[NumpyArray]
|
||||
) -> list[NumpyArray] | list[Image.Image]:
|
||||
for transform in self.transforms:
|
||||
images = transform(images)
|
||||
return images
|
||||
@@ -118,6 +303,7 @@ class Compose:
|
||||
Valid size keys (nested):
|
||||
- {"height", "width"}
|
||||
- {"shortest_edge"}
|
||||
- {"longest_edge"}
|
||||
|
||||
Returns:
|
||||
Compose: Image processor.
|
||||
@@ -128,6 +314,7 @@ class Compose:
|
||||
cls._get_pad2square(transforms, config)
|
||||
cls._get_center_crop(transforms, config)
|
||||
cls._get_pil2ndarray(transforms, config)
|
||||
cls._get_image_splitting(transforms, config)
|
||||
cls._get_rescale(transforms, config)
|
||||
cls._get_normalize(transforms, config)
|
||||
return cls(transforms=transforms)
|
||||
@@ -196,6 +383,25 @@ class Compose:
|
||||
resample=resample,
|
||||
)
|
||||
)
|
||||
elif mode == "Idefics3ImageProcessor":
|
||||
if config.get("do_resize", False):
|
||||
size = config.get("size", {})
|
||||
if "longest_edge" not in size:
|
||||
raise ValueError(
|
||||
"Size dictionary must contain 'longest_edge' key for Idefics3ImageProcessor"
|
||||
)
|
||||
|
||||
# Handle resample parameter - can be int enum or PIL.Image.Resampling
|
||||
resample = config.get("resample", Image.Resampling.LANCZOS)
|
||||
if isinstance(resample, int):
|
||||
resample = Image.Resampling(resample)
|
||||
|
||||
transforms.append(
|
||||
ResizeLongestEdge(
|
||||
size=size["longest_edge"],
|
||||
resample=resample,
|
||||
)
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Preprocessor {mode} is not supported")
|
||||
|
||||
@@ -217,6 +423,8 @@ class Compose:
|
||||
pass
|
||||
elif mode == "JinaCLIPImageProcessor":
|
||||
pass
|
||||
elif mode == "Idefics3ImageProcessor":
|
||||
pass
|
||||
else:
|
||||
raise ValueError(f"Preprocessor {mode} is not supported")
|
||||
|
||||
@@ -224,6 +432,28 @@ class Compose:
|
||||
def _get_pil2ndarray(transforms: list[Transform], config: dict[str, Any]) -> None:
|
||||
transforms.append(PILtoNDarray())
|
||||
|
||||
@classmethod
|
||||
def _get_image_splitting(cls, transforms: list[Transform], config: dict[str, Any]) -> None:
|
||||
"""
|
||||
Add image splitting transforms for Idefics3.
|
||||
Handles conditional logic: splitting vs square resize.
|
||||
Must be called AFTER PILtoNDarray.
|
||||
"""
|
||||
mode = config.get("image_processor_type", "CLIPImageProcessor")
|
||||
|
||||
if mode == "Idefics3ImageProcessor":
|
||||
do_splitting = config.get("do_image_splitting", False)
|
||||
max_size = config.get("max_image_size", {}).get("longest_edge", 512)
|
||||
resample = config.get("resample", Image.Resampling.LANCZOS)
|
||||
if isinstance(resample, int):
|
||||
resample = Image.Resampling(resample)
|
||||
|
||||
if do_splitting:
|
||||
transforms.append(ResizeForVisionEncoder(max_size, resample))
|
||||
transforms.append(ImageSplitter(max_size, resample))
|
||||
else:
|
||||
transforms.append(SquareResize(max_size, resample))
|
||||
|
||||
@staticmethod
|
||||
def _get_rescale(transforms: list[Transform], config: dict[str, Any]) -> None:
|
||||
if config.get("do_rescale", True):
|
||||
@@ -253,7 +483,7 @@ class Compose:
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _interpolation_resolver(resample: Optional[str] = None) -> Image.Resampling:
|
||||
def _interpolation_resolver(resample: str | None = None) -> Image.Resampling:
|
||||
interpolation_map = {
|
||||
"nearest": Image.Resampling.NEAREST,
|
||||
"lanczos": Image.Resampling.LANCZOS,
|
||||
|
||||
@@ -1,13 +1,14 @@
|
||||
import string
|
||||
from typing import Any, Iterable, Optional, Sequence, Type, Union
|
||||
from typing import Any, Iterable, Sequence, Type
|
||||
|
||||
import numpy as np
|
||||
from tokenizers import Encoding
|
||||
from tokenizers import Encoding, Tokenizer
|
||||
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.common.preprocessor_utils import load_tokenizer
|
||||
from fastembed.common.types import NumpyArray, Device
|
||||
from fastembed.common import OnnxProvider
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.common.utils import define_cache_dir
|
||||
from fastembed.common.utils import define_cache_dir, iter_batch
|
||||
from fastembed.late_interaction.late_interaction_embedding_base import (
|
||||
LateInteractionTextEmbeddingBase,
|
||||
)
|
||||
@@ -18,7 +19,7 @@ supported_colbert_models: list[DenseModelDescription] = [
|
||||
DenseModelDescription(
|
||||
model="colbert-ir/colbertv2.0",
|
||||
dim=128,
|
||||
description="Late interaction model",
|
||||
description="Text embeddings, Unimodal (text), English, 512 input tokens truncation, 2023 year",
|
||||
license="mit",
|
||||
size_in_GB=0.44,
|
||||
sources=ModelSource(hf="colbert-ir/colbertv2.0"),
|
||||
@@ -27,7 +28,7 @@ supported_colbert_models: list[DenseModelDescription] = [
|
||||
DenseModelDescription(
|
||||
model="answerdotai/answerai-colbert-small-v1",
|
||||
dim=96,
|
||||
description="Text embeddings, Unimodal (text), Multilingual (~100 languages), 512 input tokens truncation, 2024 year",
|
||||
description="Text embeddings, Unimodal (text), English, 512 input tokens truncation, 2024 year",
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.13,
|
||||
sources=ModelSource(hf="answerdotai/answerai-colbert-small-v1"),
|
||||
@@ -43,26 +44,29 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
|
||||
MASK_TOKEN = "[MASK]"
|
||||
|
||||
def _post_process_onnx_output(
|
||||
self, output: OnnxOutputContext, is_doc: bool = True
|
||||
self, output: OnnxOutputContext, is_doc: bool = True, **kwargs: Any
|
||||
) -> Iterable[NumpyArray]:
|
||||
if not is_doc:
|
||||
return output.model_output.astype(np.float32)
|
||||
for embedding in output.model_output:
|
||||
yield embedding
|
||||
else:
|
||||
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"
|
||||
)
|
||||
|
||||
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): # type: ignore
|
||||
if token_id in self.skip_list or token_id == self.pad_token_id:
|
||||
output.attention_mask[i, j] = 0
|
||||
|
||||
for i, token_sequence in enumerate(output.input_ids):
|
||||
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
|
||||
output.model_output *= np.expand_dims(output.attention_mask, 2)
|
||||
norm = np.linalg.norm(output.model_output, ord=2, axis=2, keepdims=True)
|
||||
norm_clamped = np.maximum(norm, 1e-12)
|
||||
output.model_output /= norm_clamped
|
||||
|
||||
output.model_output *= np.expand_dims(output.attention_mask, 2).astype(np.float32)
|
||||
norm = np.linalg.norm(output.model_output, ord=2, axis=2, keepdims=True)
|
||||
norm_clamped = np.maximum(norm, 1e-12)
|
||||
output.model_output /= norm_clamped
|
||||
return output.model_output.astype(np.float32)
|
||||
for embedding, attention_mask in zip(output.model_output, output.attention_mask):
|
||||
yield embedding[attention_mask == 1]
|
||||
|
||||
def _preprocess_onnx_input(
|
||||
self, onnx_input: dict[str, NumpyArray], is_doc: bool = True, **kwargs: Any
|
||||
@@ -84,29 +88,46 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
|
||||
)
|
||||
|
||||
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
|
||||
if self.tokenizer.padding:
|
||||
prev_padding = self.tokenizer.padding
|
||||
self.tokenizer.enable_padding(
|
||||
pad_token=self.MASK_TOKEN,
|
||||
pad_id=self.mask_token_id,
|
||||
length=self.MIN_QUERY_LENGTH,
|
||||
)
|
||||
encoded = self.tokenizer.encode_batch([query])
|
||||
if prev_padding is None:
|
||||
self.tokenizer.no_padding()
|
||||
else:
|
||||
self.tokenizer.enable_padding(**prev_padding)
|
||||
assert self.query_tokenizer is not None
|
||||
encoded = self.query_tokenizer.encode_batch([query])
|
||||
return encoded
|
||||
|
||||
def _tokenize_documents(self, documents: list[str]) -> list[Encoding]:
|
||||
encoded = self.tokenizer.encode_batch(documents) # type: ignore[union-attr]
|
||||
return encoded
|
||||
|
||||
def token_count(
|
||||
self,
|
||||
texts: str | Iterable[str],
|
||||
batch_size: int = 1024,
|
||||
is_doc: bool = True,
|
||||
include_extension: bool = False,
|
||||
**kwargs: Any,
|
||||
) -> int:
|
||||
if not hasattr(self, "model") or self.model is None:
|
||||
self.load_onnx_model() # loads the tokenizer as well
|
||||
token_num = 0
|
||||
texts = [texts] if isinstance(texts, str) else texts
|
||||
tokenizer = self.tokenizer if is_doc else self.query_tokenizer
|
||||
assert tokenizer is not None
|
||||
for batch in iter_batch(texts, batch_size):
|
||||
for tokens in tokenizer.encode_batch(batch):
|
||||
if is_doc:
|
||||
token_num += sum(tokens.attention_mask)
|
||||
else:
|
||||
attend_count = sum(tokens.attention_mask)
|
||||
if include_extension:
|
||||
token_num += max(attend_count, self.MIN_QUERY_LENGTH)
|
||||
|
||||
else:
|
||||
token_num += attend_count
|
||||
if include_extension:
|
||||
token_num += len(
|
||||
batch
|
||||
) # add 1 for each cls.DOC_MARKER_TOKEN_ID or cls.QUERY_MARKER_TOKEN_ID
|
||||
|
||||
return token_num
|
||||
|
||||
@classmethod
|
||||
def _list_supported_models(cls) -> list[DenseModelDescription]:
|
||||
"""Lists the supported models.
|
||||
@@ -119,14 +140,14 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
|
||||
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,
|
||||
cache_dir: str | None = None,
|
||||
threads: int | None = None,
|
||||
providers: Sequence[OnnxProvider] | None = None,
|
||||
cuda: bool | Device = Device.AUTO,
|
||||
device_ids: list[int] | None = None,
|
||||
lazy_load: bool = False,
|
||||
device_id: Optional[int] = None,
|
||||
specific_model_path: Optional[str] = None,
|
||||
device_id: int | None = None,
|
||||
specific_model_path: str | None = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
"""
|
||||
@@ -138,10 +159,11 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
|
||||
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.
|
||||
cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`
|
||||
Defaults to Device.AUTO.
|
||||
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.
|
||||
workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, 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.
|
||||
@@ -154,13 +176,14 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
|
||||
super().__init__(model_name, cache_dir, threads, **kwargs)
|
||||
self.providers = providers
|
||||
self.lazy_load = lazy_load
|
||||
self._extra_session_options = self._select_exposed_session_options(kwargs)
|
||||
|
||||
# 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
|
||||
self.device_id: int | None = None
|
||||
if device_id is not None:
|
||||
self.device_id = device_id
|
||||
elif self.device_ids is not None:
|
||||
@@ -169,16 +192,19 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
|
||||
self.model_description = self._get_model_description(model_name)
|
||||
self.cache_dir = str(define_cache_dir(cache_dir))
|
||||
|
||||
self._specific_model_path = specific_model_path
|
||||
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,
|
||||
specific_model_path=self._specific_model_path,
|
||||
)
|
||||
self.mask_token_id: Optional[int] = None
|
||||
self.pad_token_id: Optional[int] = None
|
||||
self.mask_token_id: int | None = None
|
||||
self.pad_token_id: int | None = None
|
||||
self.skip_list: set[int] = set()
|
||||
|
||||
self.query_tokenizer: Tokenizer | None = None
|
||||
|
||||
if not self.lazy_load:
|
||||
self.load_onnx_model()
|
||||
|
||||
@@ -190,7 +216,10 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_id=self.device_id,
|
||||
extra_session_options=self._extra_session_options,
|
||||
)
|
||||
self.query_tokenizer, _ = load_tokenizer(model_dir=self._model_dir)
|
||||
|
||||
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"]
|
||||
@@ -201,12 +230,18 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
|
||||
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)
|
||||
self.query_tokenizer.enable_truncation(max_length=current_max_length - 1)
|
||||
self.query_tokenizer.enable_padding(
|
||||
pad_token=self.MASK_TOKEN,
|
||||
pad_id=self.mask_token_id,
|
||||
length=self.MIN_QUERY_LENGTH,
|
||||
)
|
||||
|
||||
def embed(
|
||||
self,
|
||||
documents: Union[str, Iterable[str]],
|
||||
documents: str | Iterable[str],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
parallel: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
@@ -233,10 +268,13 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_ids=self.device_ids,
|
||||
local_files_only=self._local_files_only,
|
||||
specific_model_path=self._specific_model_path,
|
||||
extra_session_options=self._extra_session_options,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def query_embed(self, query: Union[str, Iterable[str]], **kwargs: Any) -> Iterable[NumpyArray]:
|
||||
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:
|
||||
if isinstance(query, str):
|
||||
query = [query]
|
||||
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from typing import Iterable, Optional, Union, Any
|
||||
from typing import Iterable, Any
|
||||
|
||||
from fastembed.common.model_description import DenseModelDescription
|
||||
from fastembed.common.types import NumpyArray
|
||||
@@ -9,20 +9,21 @@ class LateInteractionTextEmbeddingBase(ModelManagement[DenseModelDescription]):
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
cache_dir: str | None = None,
|
||||
threads: int | None = 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)
|
||||
self._embedding_size: int | None = None
|
||||
|
||||
def embed(
|
||||
self,
|
||||
documents: Union[str, Iterable[str]],
|
||||
documents: str | Iterable[str],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
parallel: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
raise NotImplementedError()
|
||||
@@ -42,7 +43,7 @@ class LateInteractionTextEmbeddingBase(ModelManagement[DenseModelDescription]):
|
||||
# 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: Any) -> Iterable[NumpyArray]:
|
||||
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
Embeds queries
|
||||
|
||||
@@ -58,3 +59,22 @@ class LateInteractionTextEmbeddingBase(ModelManagement[DenseModelDescription]):
|
||||
yield from self.embed([query], **kwargs)
|
||||
else:
|
||||
yield from self.embed(query, **kwargs)
|
||||
|
||||
@classmethod
|
||||
def get_embedding_size(cls, model_name: str) -> int:
|
||||
"""Returns embedding size of the chosen model."""
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
@property
|
||||
def embedding_size(self) -> int:
|
||||
"""Returns embedding size for the current model"""
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
def token_count(
|
||||
self,
|
||||
texts: str | Iterable[str],
|
||||
batch_size: int = 1024,
|
||||
**kwargs: Any,
|
||||
) -> int:
|
||||
"""Returns the number of tokens in the texts."""
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
from typing import Any, Iterable, Optional, Sequence, Type, Union
|
||||
from typing import Any, Iterable, Sequence, Type
|
||||
from dataclasses import asdict
|
||||
|
||||
from fastembed.common.model_description import DenseModelDescription
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.common.types import NumpyArray, Device
|
||||
from fastembed.common import OnnxProvider
|
||||
from fastembed.late_interaction.colbert import Colbert
|
||||
from fastembed.late_interaction.jina_colbert import JinaColbert
|
||||
@@ -51,11 +51,11 @@ class LateInteractionTextEmbedding(LateInteractionTextEmbeddingBase):
|
||||
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,
|
||||
cache_dir: str | None = None,
|
||||
threads: int | None = None,
|
||||
providers: Sequence[OnnxProvider] | None = None,
|
||||
cuda: bool | Device = Device.AUTO,
|
||||
device_ids: list[int] | None = None,
|
||||
lazy_load: bool = False,
|
||||
**kwargs: Any,
|
||||
):
|
||||
@@ -80,11 +80,45 @@ class LateInteractionTextEmbedding(LateInteractionTextEmbeddingBase):
|
||||
"Please check the supported models using `LateInteractionTextEmbedding.list_supported_models()`"
|
||||
)
|
||||
|
||||
@property
|
||||
def embedding_size(self) -> int:
|
||||
"""Get the embedding size of the current model"""
|
||||
if self._embedding_size is None:
|
||||
self._embedding_size = self.get_embedding_size(self.model_name)
|
||||
return self._embedding_size
|
||||
|
||||
@classmethod
|
||||
def get_embedding_size(cls, model_name: str) -> int:
|
||||
"""Get the embedding size of the passed model
|
||||
|
||||
Args:
|
||||
model_name (str): The name of the model to get embedding size for.
|
||||
|
||||
Returns:
|
||||
int: The size of the embedding.
|
||||
|
||||
Raises:
|
||||
ValueError: If the model name is not found in the supported models.
|
||||
"""
|
||||
descriptions = cls._list_supported_models()
|
||||
embedding_size: int | None = None
|
||||
for description in descriptions:
|
||||
if description.model.lower() == model_name.lower():
|
||||
embedding_size = description.dim
|
||||
break
|
||||
if embedding_size is None:
|
||||
model_names = [description.model for description in descriptions]
|
||||
raise ValueError(
|
||||
f"Embedding size for model {model_name} was None. "
|
||||
f"Available model names: {model_names}"
|
||||
)
|
||||
return embedding_size
|
||||
|
||||
def embed(
|
||||
self,
|
||||
documents: Union[str, Iterable[str]],
|
||||
documents: str | Iterable[str],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
parallel: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
@@ -104,7 +138,7 @@ class LateInteractionTextEmbedding(LateInteractionTextEmbeddingBase):
|
||||
"""
|
||||
yield from self.model.embed(documents, batch_size, parallel, **kwargs)
|
||||
|
||||
def query_embed(self, query: Union[str, Iterable[str]], **kwargs: Any) -> Iterable[NumpyArray]:
|
||||
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
Embeds queries
|
||||
|
||||
@@ -117,3 +151,30 @@ class LateInteractionTextEmbedding(LateInteractionTextEmbeddingBase):
|
||||
|
||||
# This is model-specific, so that different models can have specialized implementations
|
||||
yield from self.model.query_embed(query, **kwargs)
|
||||
|
||||
def token_count(
|
||||
self,
|
||||
texts: str | Iterable[str],
|
||||
batch_size: int = 1024,
|
||||
is_doc: bool = True,
|
||||
include_extension: bool = False,
|
||||
**kwargs: Any,
|
||||
) -> int:
|
||||
"""Returns the number of tokens in the texts.
|
||||
|
||||
Args:
|
||||
texts (str | Iterable[str]): The list of texts to embed.
|
||||
batch_size (int): Batch size for encoding
|
||||
is_doc (bool): Whether the texts are documents (disable embedding a query with include_mask=True).
|
||||
include_extension (bool): Turn on to count DOC / QUERY marker tokens, and [MASK] token in query mode.
|
||||
|
||||
Returns:
|
||||
int: Sum of number of tokens in the texts.
|
||||
"""
|
||||
return self.model.token_count(
|
||||
texts,
|
||||
batch_size=batch_size,
|
||||
is_doc=is_doc,
|
||||
include_extension=include_extension,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,83 @@
|
||||
from dataclasses import asdict
|
||||
from typing import Iterable, Any, Type
|
||||
|
||||
from fastembed.common.model_description import DenseModelDescription, ModelSource
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.late_interaction.late_interaction_embedding_base import (
|
||||
LateInteractionTextEmbeddingBase,
|
||||
)
|
||||
from fastembed.text.onnx_embedding import OnnxTextEmbedding
|
||||
from fastembed.text.onnx_text_model import TextEmbeddingWorker
|
||||
|
||||
|
||||
supported_token_embeddings_models = [
|
||||
DenseModelDescription(
|
||||
model="jinaai/jina-embeddings-v2-small-en-tokens",
|
||||
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",
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
class TokenEmbeddingsModel(OnnxTextEmbedding, LateInteractionTextEmbeddingBase):
|
||||
@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_token_embeddings_models
|
||||
|
||||
@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 [asdict(model) for model in cls._list_supported_models()]
|
||||
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type[TextEmbeddingWorker[NumpyArray]]:
|
||||
return TokensEmbeddingWorker
|
||||
|
||||
def _post_process_onnx_output(
|
||||
self, output: OnnxOutputContext, **kwargs: Any
|
||||
) -> Iterable[NumpyArray]:
|
||||
# Size: (batch_size, sequence_length, hidden_size)
|
||||
embeddings = output.model_output
|
||||
# Size: (batch_size, sequence_length)
|
||||
assert output.attention_mask is not None
|
||||
masks = output.attention_mask
|
||||
|
||||
# For each document we only select those embeddings that are not masked out
|
||||
for i in range(embeddings.shape[0]):
|
||||
yield embeddings[i, masks[i] == 1]
|
||||
|
||||
def embed(
|
||||
self,
|
||||
documents: str | Iterable[str],
|
||||
batch_size: int = 256,
|
||||
parallel: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
yield from super().embed(documents, batch_size=batch_size, parallel=parallel, **kwargs)
|
||||
|
||||
|
||||
class TokensEmbeddingWorker(TextEmbeddingWorker[NumpyArray]):
|
||||
def init_embedding(
|
||||
self, model_name: str, cache_dir: str, **kwargs: Any
|
||||
) -> TokenEmbeddingsModel:
|
||||
return TokenEmbeddingsModel(
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
threads=1,
|
||||
**kwargs,
|
||||
)
|
||||
@@ -0,0 +1,532 @@
|
||||
import contextlib
|
||||
from typing import Any, Iterable, Type, Optional, Sequence
|
||||
import json
|
||||
|
||||
import numpy as np
|
||||
from tokenizers import Encoding
|
||||
from PIL import Image
|
||||
|
||||
from fastembed.common import ImageInput
|
||||
from fastembed.common.model_description import DenseModelDescription, ModelSource
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.common.types import NumpyArray, OnnxProvider
|
||||
from fastembed.common.utils import define_cache_dir, iter_batch
|
||||
from fastembed.late_interaction_multimodal.late_interaction_multimodal_embedding_base import (
|
||||
LateInteractionMultimodalEmbeddingBase,
|
||||
)
|
||||
from fastembed.late_interaction_multimodal.onnx_multimodal_model import (
|
||||
OnnxMultimodalModel,
|
||||
TextEmbeddingWorker,
|
||||
ImageEmbeddingWorker,
|
||||
)
|
||||
|
||||
supported_colmodernvbert_models: list[DenseModelDescription] = [
|
||||
DenseModelDescription(
|
||||
model="Qdrant/colmodernvbert",
|
||||
dim=128,
|
||||
description="The late-interaction version of ModernVBERT, CPU friendly, English, 2025.",
|
||||
license="mit",
|
||||
size_in_GB=1.0,
|
||||
sources=ModelSource(hf="Qdrant/colmodernvbert"),
|
||||
additional_files=["processor_config.json"],
|
||||
model_file="model.onnx",
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
class ColModernVBERT(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyArray]):
|
||||
"""
|
||||
The ModernVBERT/colmodernvbert model implementation. This model uses
|
||||
bidirectional attention, which proves to work better for retrieval.
|
||||
|
||||
See: https://huggingface.co/ModernVBERT/colmodernvbert
|
||||
"""
|
||||
|
||||
VISUAL_PROMPT_PREFIX = (
|
||||
"<|begin_of_text|>User:<image>Describe the image.<end_of_utterance>\nAssistant:"
|
||||
)
|
||||
QUERY_AUGMENTATION_TOKEN = "<end_of_utterance>"
|
||||
|
||||
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
|
||||
self._extra_session_options = self._select_exposed_session_options(kwargs)
|
||||
|
||||
# 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._specific_model_path = specific_model_path
|
||||
self._model_dir = self.download_model(
|
||||
self.model_description,
|
||||
self.cache_dir,
|
||||
local_files_only=self._local_files_only,
|
||||
specific_model_path=self._specific_model_path,
|
||||
)
|
||||
self.mask_token_id = None
|
||||
self.pad_token_id = None
|
||||
self.image_seq_len: Optional[int] = None
|
||||
self.max_image_size: Optional[int] = None
|
||||
self.image_size: Optional[int] = 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_colmodernvbert_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,
|
||||
extra_session_options=self._extra_session_options,
|
||||
)
|
||||
|
||||
# Load image processing configuration
|
||||
processor_config_path = self._model_dir / "processor_config.json"
|
||||
with open(processor_config_path) as f:
|
||||
processor_config = json.load(f)
|
||||
self.image_seq_len = processor_config.get("image_seq_len", 64)
|
||||
|
||||
preprocessor_config_path = self._model_dir / "preprocessor_config.json"
|
||||
with open(preprocessor_config_path) as f:
|
||||
preprocessor_config = json.load(f)
|
||||
self.max_image_size = preprocessor_config.get("max_image_size", {}).get(
|
||||
"longest_edge", 512
|
||||
)
|
||||
|
||||
# Load model configuration
|
||||
config_path = self._model_dir / "config.json"
|
||||
with open(config_path) as f:
|
||||
model_config = json.load(f)
|
||||
vision_config = model_config.get("vision_config", {})
|
||||
self.image_size = vision_config.get("image_size", 512)
|
||||
|
||||
def _preprocess_onnx_text_input(
|
||||
self, onnx_input: dict[str, NumpyArray], **kwargs: Any
|
||||
) -> dict[str, 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.
|
||||
"""
|
||||
batch_size, seq_length = onnx_input["input_ids"].shape
|
||||
empty_image_placeholder: NumpyArray = np.zeros(
|
||||
(batch_size, seq_length, 3, self.image_size, self.image_size),
|
||||
dtype=np.float32, # type: ignore[type-var,arg-type,assignment]
|
||||
)
|
||||
onnx_input["pixel_values"] = empty_image_placeholder
|
||||
return onnx_input
|
||||
|
||||
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
|
||||
|
||||
def tokenize(self, documents: list[str], **kwargs: Any) -> list[Encoding]:
|
||||
# Add query augmentation tokens (matching process_queries logic from colpali-engine)
|
||||
augmented_queries = [doc + self.QUERY_AUGMENTATION_TOKEN * 10 for doc in documents]
|
||||
encoded = self.tokenizer.encode_batch(augmented_queries) # type: ignore[union-attr]
|
||||
return encoded
|
||||
|
||||
def token_count(
|
||||
self,
|
||||
texts: str | Iterable[str],
|
||||
batch_size: int = 1024,
|
||||
include_extension: bool = False,
|
||||
**kwargs: Any,
|
||||
) -> int:
|
||||
if not hasattr(self, "model") or self.model is None:
|
||||
self.load_onnx_model() # loads the tokenizer as well
|
||||
token_num = 0
|
||||
texts = [texts] if isinstance(texts, str) else texts
|
||||
assert self.tokenizer is not None
|
||||
tokenize_func = self.tokenize if include_extension else self.tokenizer.encode_batch
|
||||
for batch in iter_batch(texts, batch_size):
|
||||
token_num += sum([sum(encoding.attention_mask) for encoding in tokenize_func(batch)])
|
||||
return token_num
|
||||
|
||||
def onnx_embed_image(self, images: list[ImageInput], **kwargs: Any) -> OnnxOutputContext:
|
||||
with contextlib.ExitStack() as stack:
|
||||
image_files = [
|
||||
stack.enter_context(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"
|
||||
processed = self.processor(image_files)
|
||||
encoded, attention_mask, metadata = self._process_nested_patches(processed) # type: ignore[arg-type]
|
||||
|
||||
onnx_input = {"pixel_values": encoded, "attention_mask": attention_mask}
|
||||
onnx_input = self._preprocess_onnx_image_input(onnx_input, **kwargs)
|
||||
model_output = self.model.run(None, onnx_input) # type: ignore[union-attr]
|
||||
|
||||
return OnnxOutputContext(
|
||||
model_output=model_output[0],
|
||||
attention_mask=attention_mask, # type: ignore[arg-type]
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _process_nested_patches(
|
||||
processed: list[list[NumpyArray]],
|
||||
) -> tuple[NumpyArray, NumpyArray, dict[str, Any]]:
|
||||
"""
|
||||
Process nested image patches (from ImageSplitter).
|
||||
|
||||
Args:
|
||||
processed: List of patch lists, one per image [[img1_patches], [img2_patches], ...]
|
||||
|
||||
Returns:
|
||||
tuple: (encoded array, attention_mask, metadata)
|
||||
- encoded: (batch_size, max_patches, C, H, W)
|
||||
- attention_mask: (batch_size, max_patches) with 1 for real patches, 0 for padding
|
||||
- metadata: Dict with 'patch_counts' key
|
||||
"""
|
||||
patch_counts = [len(patches) for patches in processed]
|
||||
max_patches = max(patch_counts)
|
||||
|
||||
# Get dimensions from first patch
|
||||
channels, height, width = processed[0][0].shape
|
||||
batch_size = len(processed)
|
||||
|
||||
# Create padded array
|
||||
encoded = np.zeros(
|
||||
(batch_size, max_patches, channels, height, width), dtype=processed[0][0].dtype
|
||||
)
|
||||
|
||||
# Create attention mask (1 for real patches, 0 for padding)
|
||||
attention_mask = np.zeros((batch_size, max_patches), dtype=np.int64)
|
||||
|
||||
# Fill in patches and attention mask
|
||||
for i, patches in enumerate(processed):
|
||||
for j, patch in enumerate(patches):
|
||||
encoded[i, j] = patch
|
||||
attention_mask[i, j] = 1
|
||||
|
||||
metadata = {"patch_counts": patch_counts}
|
||||
return encoded, attention_mask, metadata # type: ignore[return-value]
|
||||
|
||||
def _preprocess_onnx_image_input(
|
||||
self, onnx_input: dict[str, np.ndarray], **kwargs: Any
|
||||
) -> dict[str, NumpyArray]:
|
||||
"""
|
||||
Add text input placeholders for image data, following Idefics3 processing logic.
|
||||
|
||||
Constructs input_ids dynamically based on the actual number of image patches,
|
||||
using the same token expansion logic as Idefics3Processor.
|
||||
|
||||
Args:
|
||||
onnx_input: Dict with 'pixel_values' (batch, num_patches, C, H, W)
|
||||
and 'attention_mask' (batch, num_patches) indicating real patches
|
||||
**kwargs: Additional arguments
|
||||
|
||||
Returns:
|
||||
Updated onnx_input with 'input_ids' and updated 'attention_mask' for token sequence
|
||||
"""
|
||||
# The attention_mask in onnx_input has a shape of (batch_size, num_patches),
|
||||
# and should be used to create an attention mask matching the input_ids shape.
|
||||
patch_attention_mask = onnx_input["attention_mask"]
|
||||
pixel_values = onnx_input["pixel_values"]
|
||||
|
||||
batch_size = pixel_values.shape[0]
|
||||
batch_input_ids = []
|
||||
|
||||
# Build input_ids for each image based on its actual patch count
|
||||
for i in range(batch_size):
|
||||
# Count real patches (non-padded) from attention mask
|
||||
patch_count = int(np.sum(patch_attention_mask[i]))
|
||||
|
||||
# Compute rows/cols from patch count
|
||||
rows, cols = self._compute_rows_cols_from_patches(patch_count)
|
||||
|
||||
# Build input_ids for this image
|
||||
input_ids = self._build_input_ids_for_image(rows, cols)
|
||||
batch_input_ids.append(input_ids)
|
||||
|
||||
# Pad sequences to max length in batch
|
||||
max_len = max(len(ids) for ids in batch_input_ids)
|
||||
|
||||
# Get padding config from tokenizer
|
||||
padding_direction = self.tokenizer.padding["direction"] # type: ignore[index,union-attr]
|
||||
pad_token_id = self.tokenizer.padding["pad_id"] # type: ignore[index,union-attr]
|
||||
|
||||
# Initialize with pad token
|
||||
padded_input_ids = np.full((batch_size, max_len), pad_token_id, dtype=np.int64)
|
||||
attention_mask = np.zeros((batch_size, max_len), dtype=np.int64)
|
||||
|
||||
for i, input_ids in enumerate(batch_input_ids):
|
||||
seq_len = len(input_ids)
|
||||
if padding_direction == "left":
|
||||
# Left padding: place tokens at the END of the array
|
||||
start_idx = max_len - seq_len
|
||||
padded_input_ids[i, start_idx:] = input_ids
|
||||
attention_mask[i, start_idx:] = 1
|
||||
else:
|
||||
# Right padding: place tokens at the START of the array
|
||||
padded_input_ids[i, :seq_len] = input_ids
|
||||
attention_mask[i, :seq_len] = 1
|
||||
|
||||
onnx_input["input_ids"] = padded_input_ids
|
||||
# Update attention_mask with token-level data
|
||||
onnx_input["attention_mask"] = attention_mask
|
||||
return onnx_input
|
||||
|
||||
@staticmethod
|
||||
def _compute_rows_cols_from_patches(patch_count: int) -> tuple[int, int]:
|
||||
if patch_count <= 1:
|
||||
return 0, 0
|
||||
|
||||
# Subtract 1 for the global image
|
||||
grid_patches = patch_count - 1
|
||||
|
||||
# Find rows and cols (assume square or near-square grid)
|
||||
rows = int(grid_patches**0.5)
|
||||
cols = grid_patches // rows
|
||||
|
||||
# Verify the calculation
|
||||
if rows * cols + 1 != patch_count:
|
||||
# Handle non-square grids
|
||||
for r in range(1, grid_patches + 1):
|
||||
if grid_patches % r == 0:
|
||||
c = grid_patches // r
|
||||
if r * c + 1 == patch_count:
|
||||
return r, c
|
||||
# Fallback: treat as unsplit
|
||||
return 0, 0
|
||||
|
||||
return rows, cols
|
||||
|
||||
def _create_single_image_prompt_string(self) -> str:
|
||||
return (
|
||||
"<fake_token_around_image>"
|
||||
+ "<global-img>"
|
||||
+ "<image>" * self.image_seq_len # type: ignore[operator]
|
||||
+ "<fake_token_around_image>"
|
||||
)
|
||||
|
||||
def _create_split_image_prompt_string(self, rows: int, cols: int) -> str:
|
||||
text_split_images = ""
|
||||
|
||||
# Add tokens for each patch in the grid
|
||||
for n_h in range(rows):
|
||||
for n_w in range(cols):
|
||||
text_split_images += (
|
||||
"<fake_token_around_image>"
|
||||
+ f"<row_{n_h + 1}_col_{n_w + 1}>"
|
||||
+ "<image>" * self.image_seq_len # type: ignore[operator]
|
||||
)
|
||||
text_split_images += "\n"
|
||||
|
||||
# Add global image at the end
|
||||
text_split_images += (
|
||||
"\n<fake_token_around_image>"
|
||||
+ "<global-img>"
|
||||
+ "<image>" * self.image_seq_len # type: ignore[operator]
|
||||
+ "<fake_token_around_image>"
|
||||
)
|
||||
|
||||
return text_split_images
|
||||
|
||||
def _build_input_ids_for_image(self, rows: int, cols: int) -> np.ndarray:
|
||||
# Create the appropriate image prompt string
|
||||
if rows == 0 and cols == 0:
|
||||
image_prompt_tokens = self._create_single_image_prompt_string()
|
||||
else:
|
||||
image_prompt_tokens = self._create_split_image_prompt_string(rows, cols)
|
||||
|
||||
# Replace <image> in visual prompt with expanded tokens
|
||||
# The visual prompt is: "<|begin_of_text|>User:<image>Describe the image.<end_of_utterance>\nAssistant:"
|
||||
expanded_prompt = self.VISUAL_PROMPT_PREFIX.replace("<image>", image_prompt_tokens)
|
||||
|
||||
# Tokenize the complete prompt
|
||||
encoded = self.tokenizer.encode(expanded_prompt) # type: ignore[union-attr]
|
||||
|
||||
# Convert to numpy array
|
||||
return np.array(encoded.ids, dtype=np.int64)
|
||||
|
||||
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
|
||||
)
|
||||
|
||||
def embed_text(
|
||||
self,
|
||||
documents: 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,
|
||||
local_files_only=self._local_files_only,
|
||||
specific_model_path=self._specific_model_path,
|
||||
extra_session_options=self._extra_session_options,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def embed_image(
|
||||
self,
|
||||
images: 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,
|
||||
local_files_only=self._local_files_only,
|
||||
specific_model_path=self._specific_model_path,
|
||||
extra_session_options=self._extra_session_options,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _get_text_worker_class(cls) -> Type[TextEmbeddingWorker[NumpyArray]]:
|
||||
return ColModernVBERTTextEmbeddingWorker
|
||||
|
||||
@classmethod
|
||||
def _get_image_worker_class(cls) -> Type[ImageEmbeddingWorker[NumpyArray]]:
|
||||
return ColModernVBERTImageEmbeddingWorker
|
||||
|
||||
|
||||
class ColModernVBERTTextEmbeddingWorker(TextEmbeddingWorker[NumpyArray]):
|
||||
def init_embedding(self, model_name: str, cache_dir: str, **kwargs: Any) -> ColModernVBERT:
|
||||
return ColModernVBERT(
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
threads=1,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
class ColModernVBERTImageEmbeddingWorker(ImageEmbeddingWorker[NumpyArray]):
|
||||
def init_embedding(self, model_name: str, cache_dir: str, **kwargs: Any) -> ColModernVBERT:
|
||||
return ColModernVBERT(
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
threads=1,
|
||||
**kwargs,
|
||||
)
|
||||
@@ -1,12 +1,12 @@
|
||||
from typing import Any, Iterable, Optional, Sequence, Type, Union
|
||||
from typing import Any, Iterable, Sequence, Type
|
||||
|
||||
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.common.types import NumpyArray, Device
|
||||
from fastembed.common.utils import define_cache_dir, iter_batch
|
||||
from fastembed.late_interaction_multimodal.late_interaction_multimodal_embedding_base import (
|
||||
LateInteractionMultimodalEmbeddingBase,
|
||||
)
|
||||
@@ -46,14 +46,14 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
|
||||
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,
|
||||
cache_dir: str | None = None,
|
||||
threads: int | None = None,
|
||||
providers: Sequence[OnnxProvider] | None = None,
|
||||
cuda: bool | Device = Device.AUTO,
|
||||
device_ids: list[int] | None = None,
|
||||
lazy_load: bool = False,
|
||||
device_id: Optional[int] = None,
|
||||
specific_model_path: Optional[str] = None,
|
||||
device_id: int | None = None,
|
||||
specific_model_path: str | None = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
"""
|
||||
@@ -65,10 +65,11 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
|
||||
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.
|
||||
cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`
|
||||
Defaults to Device.AUTO.
|
||||
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.
|
||||
workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, 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.
|
||||
@@ -80,13 +81,14 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
|
||||
super().__init__(model_name, cache_dir, threads, **kwargs)
|
||||
self.providers = providers
|
||||
self.lazy_load = lazy_load
|
||||
self._extra_session_options = self._select_exposed_session_options(kwargs)
|
||||
|
||||
# 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
|
||||
self.device_id: int | None = None
|
||||
if device_id is not None:
|
||||
self.device_id = device_id
|
||||
elif self.device_ids is not None:
|
||||
@@ -95,11 +97,12 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
|
||||
self.model_description = self._get_model_description(model_name)
|
||||
self.cache_dir = str(define_cache_dir(cache_dir))
|
||||
|
||||
self._specific_model_path = specific_model_path
|
||||
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,
|
||||
specific_model_path=self._specific_model_path,
|
||||
)
|
||||
self.mask_token_id = None
|
||||
self.pad_token_id = None
|
||||
@@ -124,6 +127,7 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_id=self.device_id,
|
||||
extra_session_options=self._extra_session_options,
|
||||
)
|
||||
|
||||
def _post_process_onnx_image_output(
|
||||
@@ -142,7 +146,7 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
|
||||
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,
|
||||
@@ -157,7 +161,7 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
|
||||
Returns:
|
||||
Iterable[NumpyArray]: Post-processed output as NumPy arrays.
|
||||
"""
|
||||
return output.model_output.astype(np.float32)
|
||||
return output.model_output
|
||||
|
||||
def tokenize(self, documents: list[str], **kwargs: Any) -> list[Encoding]:
|
||||
texts_query: list[str] = []
|
||||
@@ -169,12 +173,29 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
|
||||
encoded = self.tokenizer.encode_batch(texts_query) # type: ignore[union-attr]
|
||||
return encoded
|
||||
|
||||
def token_count(
|
||||
self,
|
||||
texts: str | Iterable[str],
|
||||
batch_size: int = 1024,
|
||||
include_extension: bool = False,
|
||||
**kwargs: Any,
|
||||
) -> int:
|
||||
if not hasattr(self, "model") or self.model is None:
|
||||
self.load_onnx_model() # loads the tokenizer as well
|
||||
token_num = 0
|
||||
texts = [texts] if isinstance(texts, str) else texts
|
||||
assert self.tokenizer is not None
|
||||
tokenize_func = self.tokenize if include_extension else self.tokenizer.encode_batch
|
||||
for batch in iter_batch(texts, batch_size):
|
||||
token_num += sum([sum(encoding.attention_mask) for encoding in tokenize_func(batch)])
|
||||
return token_num
|
||||
|
||||
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()
|
||||
self.QUERY_MARKER_TOKEN_ID + input_ids[2:].tolist() # type: ignore[index]
|
||||
for input_ids in onnx_input["input_ids"]
|
||||
]
|
||||
)
|
||||
@@ -197,20 +218,19 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
|
||||
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"]]
|
||||
[self.EMPTY_TEXT_PLACEHOLDER for _ in onnx_input["pixel_values"]]
|
||||
)
|
||||
onnx_input["attention_mask"] = np.array(
|
||||
[self.EVEN_ATTENTION_MASK for _ in onnx_input["input_ids"]]
|
||||
[self.EVEN_ATTENTION_MASK for _ in onnx_input["pixel_values"]]
|
||||
)
|
||||
return onnx_input
|
||||
|
||||
def embed_text(
|
||||
self,
|
||||
documents: Union[str, Iterable[str]],
|
||||
documents: str | Iterable[str],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
parallel: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
@@ -236,14 +256,17 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_ids=self.device_ids,
|
||||
local_files_only=self._local_files_only,
|
||||
specific_model_path=self._specific_model_path,
|
||||
extra_session_options=self._extra_session_options,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def embed_image(
|
||||
self,
|
||||
images: Union[ImageInput, Iterable[ImageInput]],
|
||||
images: ImageInput | Iterable[ImageInput],
|
||||
batch_size: int = 16,
|
||||
parallel: Optional[int] = None,
|
||||
parallel: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
@@ -269,6 +292,9 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_ids=self.device_ids,
|
||||
local_files_only=self._local_files_only,
|
||||
specific_model_path=self._specific_model_path,
|
||||
extra_session_options=self._extra_session_options,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
@@ -1,9 +1,10 @@
|
||||
from typing import Any, Iterable, Optional, Sequence, Type, Union
|
||||
from typing import Any, Iterable, Sequence, Type
|
||||
from dataclasses import asdict
|
||||
|
||||
from fastembed.common import OnnxProvider, ImageInput
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.common.types import NumpyArray, Device
|
||||
from fastembed.late_interaction_multimodal.colpali import ColPali
|
||||
from fastembed.late_interaction_multimodal.colmodernvbert import ColModernVBERT
|
||||
|
||||
from fastembed.late_interaction_multimodal.late_interaction_multimodal_embedding_base import (
|
||||
LateInteractionMultimodalEmbeddingBase,
|
||||
@@ -12,7 +13,10 @@ from fastembed.common.model_description import DenseModelDescription
|
||||
|
||||
|
||||
class LateInteractionMultimodalEmbedding(LateInteractionMultimodalEmbeddingBase):
|
||||
EMBEDDINGS_REGISTRY: list[Type[LateInteractionMultimodalEmbeddingBase]] = [ColPali]
|
||||
EMBEDDINGS_REGISTRY: list[Type[LateInteractionMultimodalEmbeddingBase]] = [
|
||||
ColPali,
|
||||
ColModernVBERT,
|
||||
]
|
||||
|
||||
@classmethod
|
||||
def list_supported_models(cls) -> list[dict[str, Any]]:
|
||||
@@ -54,11 +58,11 @@ class LateInteractionMultimodalEmbedding(LateInteractionMultimodalEmbeddingBase)
|
||||
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,
|
||||
cache_dir: str | None = None,
|
||||
threads: int | None = None,
|
||||
providers: Sequence[OnnxProvider] | None = None,
|
||||
cuda: bool | Device = Device.AUTO,
|
||||
device_ids: list[int] | None = None,
|
||||
lazy_load: bool = False,
|
||||
**kwargs: Any,
|
||||
):
|
||||
@@ -83,11 +87,45 @@ class LateInteractionMultimodalEmbedding(LateInteractionMultimodalEmbeddingBase)
|
||||
"Please check the supported models using `LateInteractionMultimodalEmbedding.list_supported_models()`"
|
||||
)
|
||||
|
||||
@property
|
||||
def embedding_size(self) -> int:
|
||||
"""Get the embedding size of the current model"""
|
||||
if self._embedding_size is None:
|
||||
self._embedding_size = self.get_embedding_size(self.model_name)
|
||||
return self._embedding_size
|
||||
|
||||
@classmethod
|
||||
def get_embedding_size(cls, model_name: str) -> int:
|
||||
"""Get the embedding size of the passed model
|
||||
|
||||
Args:
|
||||
model_name (str): The name of the model to get embedding size for.
|
||||
|
||||
Returns:
|
||||
int: The size of the embedding.
|
||||
|
||||
Raises:
|
||||
ValueError: If the model name is not found in the supported models.
|
||||
"""
|
||||
descriptions = cls._list_supported_models()
|
||||
embedding_size: int | None = None
|
||||
for description in descriptions:
|
||||
if description.model.lower() == model_name.lower():
|
||||
embedding_size = description.dim
|
||||
break
|
||||
if embedding_size is None:
|
||||
model_names = [description.model for description in descriptions]
|
||||
raise ValueError(
|
||||
f"Embedding size for model {model_name} was None. "
|
||||
f"Available model names: {model_names}"
|
||||
)
|
||||
return embedding_size
|
||||
|
||||
def embed_text(
|
||||
self,
|
||||
documents: Union[str, Iterable[str]],
|
||||
documents: str | Iterable[str],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
parallel: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
@@ -108,9 +146,9 @@ class LateInteractionMultimodalEmbedding(LateInteractionMultimodalEmbeddingBase)
|
||||
|
||||
def embed_image(
|
||||
self,
|
||||
images: Union[ImageInput, Iterable[ImageInput]],
|
||||
images: ImageInput | Iterable[ImageInput],
|
||||
batch_size: int = 16,
|
||||
parallel: Optional[int] = None,
|
||||
parallel: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
@@ -128,3 +166,24 @@ class LateInteractionMultimodalEmbedding(LateInteractionMultimodalEmbeddingBase)
|
||||
List of embeddings, one per image
|
||||
"""
|
||||
yield from self.model.embed_image(images, batch_size, parallel, **kwargs)
|
||||
|
||||
def token_count(
|
||||
self,
|
||||
texts: str | Iterable[str],
|
||||
batch_size: int = 1024,
|
||||
include_extension: bool = False,
|
||||
**kwargs: Any,
|
||||
) -> int:
|
||||
"""Returns the number of tokens in the texts.
|
||||
|
||||
Args:
|
||||
texts (str | Iterable[str]): The list of texts to embed.
|
||||
batch_size (int): Batch size for encoding
|
||||
include_extension (bool): Whether to include tokens added by preprocessing
|
||||
|
||||
Returns:
|
||||
int: Sum of number of tokens in the texts.
|
||||
"""
|
||||
return self.model.token_count(
|
||||
texts, batch_size=batch_size, include_extension=include_extension, **kwargs
|
||||
)
|
||||
|
||||
+26
-7
@@ -1,4 +1,4 @@
|
||||
from typing import Iterable, Optional, Union, Any
|
||||
from typing import Iterable, Any
|
||||
|
||||
|
||||
from fastembed.common import ImageInput
|
||||
@@ -11,20 +11,21 @@ class LateInteractionMultimodalEmbeddingBase(ModelManagement[DenseModelDescripti
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
cache_dir: str | None = None,
|
||||
threads: int | None = 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)
|
||||
self._embedding_size: int | None = None
|
||||
|
||||
def embed_text(
|
||||
self,
|
||||
documents: Union[str, Iterable[str]],
|
||||
documents: str | Iterable[str],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
parallel: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
@@ -46,9 +47,9 @@ class LateInteractionMultimodalEmbeddingBase(ModelManagement[DenseModelDescripti
|
||||
|
||||
def embed_image(
|
||||
self,
|
||||
images: Union[ImageInput, Iterable[ImageInput]],
|
||||
images: ImageInput | Iterable[ImageInput],
|
||||
batch_size: int = 16,
|
||||
parallel: Optional[int] = None,
|
||||
parallel: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
@@ -65,3 +66,21 @@ class LateInteractionMultimodalEmbeddingBase(ModelManagement[DenseModelDescripti
|
||||
List of embeddings, one per image
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
@classmethod
|
||||
def get_embedding_size(cls, model_name: str) -> int:
|
||||
"""Returns embedding size of the chosen model."""
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
@property
|
||||
def embedding_size(self) -> int:
|
||||
"""Returns embedding size for the current model"""
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
def token_count(
|
||||
self,
|
||||
texts: str | Iterable[str],
|
||||
**kwargs: Any,
|
||||
) -> int:
|
||||
"""Returns the number of tokens in the texts."""
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
@@ -2,7 +2,7 @@ import contextlib
|
||||
import os
|
||||
from multiprocessing import get_all_start_methods
|
||||
from pathlib import Path
|
||||
from typing import Any, Iterable, Optional, Sequence, Type, Union
|
||||
from typing import Any, Iterable, Sequence, Type
|
||||
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
@@ -11,19 +11,19 @@ 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.types import NumpyArray, Device
|
||||
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
|
||||
ONNX_OUTPUT_NAMES: list[str] | None = None
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.tokenizer: Optional[Tokenizer] = None
|
||||
self.processor: Optional[Compose] = None
|
||||
self.tokenizer: Tokenizer | None = None
|
||||
self.processor: Compose | None = None
|
||||
self.special_token_to_id: dict[str, int] = {}
|
||||
|
||||
def _preprocess_onnx_text_input(
|
||||
@@ -60,10 +60,11 @@ class OnnxMultimodalModel(OnnxModel[T]):
|
||||
self,
|
||||
model_dir: Path,
|
||||
model_file: str,
|
||||
threads: Optional[int],
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
cuda: bool = False,
|
||||
device_id: Optional[int] = None,
|
||||
threads: int | None,
|
||||
providers: Sequence[OnnxProvider] | None = None,
|
||||
cuda: bool | Device = Device.AUTO,
|
||||
device_id: int | None = None,
|
||||
extra_session_options: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
super()._load_onnx_model(
|
||||
model_dir=model_dir,
|
||||
@@ -72,9 +73,10 @@ class OnnxMultimodalModel(OnnxModel[T]):
|
||||
providers=providers,
|
||||
cuda=cuda,
|
||||
device_id=device_id,
|
||||
extra_session_options=extra_session_options,
|
||||
)
|
||||
assert self.tokenizer is not None
|
||||
self.tokenizer, self.special_token_to_id = load_tokenizer(model_dir=model_dir)
|
||||
assert self.tokenizer is not None
|
||||
self.processor = load_preprocessor(model_dir=model_dir)
|
||||
|
||||
def load_onnx_model(self) -> None:
|
||||
@@ -114,12 +116,15 @@ class OnnxMultimodalModel(OnnxModel[T]):
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
documents: Union[str, Iterable[str]],
|
||||
documents: 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,
|
||||
parallel: int | None = None,
|
||||
providers: Sequence[OnnxProvider] | None = None,
|
||||
cuda: bool | Device = Device.AUTO,
|
||||
device_ids: list[int] | None = None,
|
||||
local_files_only: bool = False,
|
||||
specific_model_path: str | None = None,
|
||||
extra_session_options: dict[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[T]:
|
||||
is_small = False
|
||||
@@ -146,9 +151,14 @@ class OnnxMultimodalModel(OnnxModel[T]):
|
||||
"model_name": model_name,
|
||||
"cache_dir": cache_dir,
|
||||
"providers": providers,
|
||||
"local_files_only": local_files_only,
|
||||
"specific_model_path": specific_model_path,
|
||||
**kwargs,
|
||||
}
|
||||
|
||||
if extra_session_options is not None:
|
||||
params.update(extra_session_options)
|
||||
|
||||
pool = ParallelWorkerPool(
|
||||
num_workers=parallel or 1,
|
||||
worker=self._get_text_worker_class(),
|
||||
@@ -159,19 +169,17 @@ class OnnxMultimodalModel(OnnxModel[T]):
|
||||
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():
|
||||
with contextlib.ExitStack() as stack:
|
||||
image_files = [
|
||||
Image.open(image) if not isinstance(image, Image.Image) else image
|
||||
stack.enter_context(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 = {"pixel_values": 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)
|
||||
@@ -181,12 +189,15 @@ class OnnxMultimodalModel(OnnxModel[T]):
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
images: Union[Iterable[ImageInput], ImageInput],
|
||||
images: 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,
|
||||
parallel: int | None = None,
|
||||
providers: Sequence[OnnxProvider] | None = None,
|
||||
cuda: bool | Device = Device.AUTO,
|
||||
device_ids: list[int] | None = None,
|
||||
local_files_only: bool = False,
|
||||
specific_model_path: str | None = None,
|
||||
extra_session_options: dict[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[T]:
|
||||
is_small = False
|
||||
@@ -213,9 +224,14 @@ class OnnxMultimodalModel(OnnxModel[T]):
|
||||
"model_name": model_name,
|
||||
"cache_dir": cache_dir,
|
||||
"providers": providers,
|
||||
"local_files_only": local_files_only,
|
||||
"specific_model_path": specific_model_path,
|
||||
**kwargs,
|
||||
}
|
||||
|
||||
if extra_session_options is not None:
|
||||
params.update(extra_session_options)
|
||||
|
||||
pool = ParallelWorkerPool(
|
||||
num_workers=parallel or 1,
|
||||
worker=self._get_image_worker_class(),
|
||||
|
||||
@@ -1,16 +0,0 @@
|
||||
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]
|
||||
@@ -8,8 +8,9 @@ from multiprocessing.context import BaseContext
|
||||
from multiprocessing.process import BaseProcess
|
||||
from multiprocessing.sharedctypes import Synchronized as BaseValue
|
||||
from queue import Empty
|
||||
from typing import Any, Iterable, Optional, Type
|
||||
from typing import Any, Iterable, Type
|
||||
|
||||
from fastembed.common.types import Device
|
||||
|
||||
# Single item should be processed in less than:
|
||||
processing_timeout = 10 * 60 # seconds
|
||||
@@ -38,7 +39,7 @@ def _worker(
|
||||
output_queue: Queue,
|
||||
num_active_workers: BaseValue,
|
||||
worker_id: int,
|
||||
kwargs: Optional[dict[str, Any]] = None,
|
||||
kwargs: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
A worker that pulls data pints off the input queue, and places the execution result on the output queue.
|
||||
@@ -93,21 +94,21 @@ class ParallelWorkerPool:
|
||||
self,
|
||||
num_workers: int,
|
||||
worker: Type[Worker],
|
||||
start_method: Optional[str] = None,
|
||||
device_ids: Optional[list[int]] = None,
|
||||
cuda: bool = False,
|
||||
start_method: str | None = None,
|
||||
device_ids: list[int] | None = None,
|
||||
cuda: bool | Device = Device.AUTO,
|
||||
):
|
||||
self.worker_class = worker
|
||||
self.num_workers = num_workers
|
||||
self.input_queue: Optional[Queue] = None
|
||||
self.output_queue: Optional[Queue] = None
|
||||
self.input_queue: Queue | None = None
|
||||
self.output_queue: Queue | None = None
|
||||
self.ctx: BaseContext = get_context(start_method)
|
||||
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
|
||||
self.num_active_workers: BaseValue | None = None
|
||||
|
||||
def start(self, **kwargs: Any) -> None:
|
||||
self.input_queue = self.ctx.Queue(self.queue_size)
|
||||
@@ -220,7 +221,7 @@ class ParallelWorkerPool:
|
||||
f"Worker PID: {process.pid} terminated unexpectedly with code {process.exitcode}"
|
||||
)
|
||||
|
||||
def join_or_terminate(self, timeout: Optional[int] = 1) -> None:
|
||||
def join_or_terminate(self, timeout: int = 1) -> None:
|
||||
"""
|
||||
Emergency shutdown
|
||||
@param timeout:
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
from fastembed.postprocess.muvera import Muvera
|
||||
|
||||
__all__ = ["Muvera"]
|
||||
@@ -0,0 +1,362 @@
|
||||
import numpy as np
|
||||
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.late_interaction.late_interaction_embedding_base import (
|
||||
LateInteractionTextEmbeddingBase,
|
||||
)
|
||||
from fastembed.late_interaction_multimodal.late_interaction_multimodal_embedding_base import (
|
||||
LateInteractionMultimodalEmbeddingBase,
|
||||
)
|
||||
|
||||
|
||||
MultiVectorModel = LateInteractionTextEmbeddingBase | LateInteractionMultimodalEmbeddingBase
|
||||
MAX_HAMMING_DISTANCE = 65 # 64 bits + 1
|
||||
POPCOUNT_LUT = np.array([bin(x).count("1") for x in range(256)], dtype=np.uint8)
|
||||
|
||||
|
||||
def hamming_distance_matrix(ids: np.ndarray) -> np.ndarray:
|
||||
"""Compute full Hamming distance matrix
|
||||
|
||||
Args:
|
||||
ids: shape (n,) - array of ids, only size of the array matters
|
||||
|
||||
Return:
|
||||
np.ndarray (n, n) - hamming distance matrix
|
||||
"""
|
||||
n = len(ids)
|
||||
xor_vals = np.bitwise_xor(ids[:, None], ids[None, :]) # (n, n) uint64
|
||||
bytes_view = xor_vals.view(np.uint8).reshape(n, n, 8) # (n, n, 8)
|
||||
return POPCOUNT_LUT[bytes_view].sum(axis=2)
|
||||
|
||||
|
||||
class SimHashProjection:
|
||||
"""
|
||||
SimHash projection component for MUVERA clustering.
|
||||
|
||||
This class implements locality-sensitive hashing using random hyperplanes
|
||||
to partition the vector space into 2^k_sim clusters. Each vector is assigned
|
||||
to a cluster based on which side of k_sim random hyperplanes it falls on.
|
||||
|
||||
Attributes:
|
||||
k_sim (int): Number of SimHash functions (hyperplanes)
|
||||
dim (int): Dimensionality of input vectors
|
||||
simhash_vectors (np.ndarray): Random hyperplane normal vectors of shape (dim, k_sim)
|
||||
"""
|
||||
|
||||
def __init__(self, k_sim: int, dim: int, random_generator: np.random.Generator):
|
||||
"""
|
||||
Initialize SimHash projection with random hyperplanes.
|
||||
|
||||
Args:
|
||||
k_sim (int): Number of SimHash functions, determines 2^k_sim clusters
|
||||
dim (int): Dimensionality of input vectors
|
||||
random_generator (np.random.Generator): Random number generator for reproducibility
|
||||
"""
|
||||
self.k_sim = k_sim
|
||||
self.dim = dim
|
||||
# Generate k_sim random hyperplanes (normal vectors) from standard normal distribution
|
||||
self.simhash_vectors = random_generator.normal(size=(dim, k_sim))
|
||||
|
||||
def get_cluster_ids(self, vectors: np.ndarray) -> np.ndarray:
|
||||
"""
|
||||
Compute the cluster IDs for a given vector using SimHash.
|
||||
|
||||
The cluster ID is determined by computing the dot product of the vector
|
||||
with each hyperplane normal vector, taking the sign, and interpreting
|
||||
the resulting binary string as an integer.
|
||||
|
||||
Args:
|
||||
vectors (np.ndarray): Input vectors of shape (n, dim,)
|
||||
|
||||
Returns:
|
||||
np.ndarray: Cluster IDs in range [0, 2^k_sim - 1]
|
||||
|
||||
Raises:
|
||||
AssertionError: If a vector shape doesn't match expected dimensionality
|
||||
"""
|
||||
dot_product = (
|
||||
vectors @ self.simhash_vectors
|
||||
) # (token_num, dim) x (dim, k_sim) -> (token_num, k_sim)
|
||||
cluster_ids = (dot_product > 0) @ (1 << np.arange(self.k_sim))
|
||||
return cluster_ids
|
||||
|
||||
|
||||
class Muvera:
|
||||
"""
|
||||
MUVERA (Multi-Vector Retrieval Architecture) algorithm implementation.
|
||||
|
||||
This class creates Fixed Dimensional Encodings (FDEs) from variable-length
|
||||
sequences of vectors by using SimHash clustering and random projections.
|
||||
The process involves:
|
||||
1. Clustering vectors using multiple SimHash projections
|
||||
2. Computing cluster centers (with different strategies for docs vs queries)
|
||||
3. Applying random projections for dimensionality reduction
|
||||
4. Concatenating results from all projections
|
||||
|
||||
Attributes:
|
||||
k_sim (int): Number of SimHash functions per projection
|
||||
dim (int): Input vector dimensionality
|
||||
dim_proj (int): Output dimensionality after random projection
|
||||
r_reps (int): Number of random projection repetitions
|
||||
random_seed (int): Random seed for consistent random matrix generation
|
||||
simhash_projections (List[SimHashProjection]): SimHash instances for clustering
|
||||
dim_reduction_projections (np.ndarray): Random projection matrices of shape (R_reps, d, d_proj)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
k_sim: int = 5,
|
||||
dim_proj: int = 16,
|
||||
r_reps: int = 20,
|
||||
random_seed: int = 42,
|
||||
):
|
||||
"""
|
||||
Initialize MUVERA algorithm with specified parameters.
|
||||
|
||||
Args:
|
||||
dim (int): Dimensionality of individual input vectors
|
||||
k_sim (int, optional): Number of SimHash functions (creates 2^k_sim clusters).
|
||||
Defaults to 5.
|
||||
dim_proj (int, optional): Dimensionality after random projection (must be <= dim).
|
||||
Defaults to 16.
|
||||
r_reps (int, optional): Number of random projection repetitions for robustness.
|
||||
Defaults to 20.
|
||||
random_seed (int, optional): Seed for random number generator to ensure
|
||||
reproducible results. Defaults to 42.
|
||||
|
||||
Raises:
|
||||
ValueError: If dim_proj > dim (cannot project to higher dimensionality)
|
||||
"""
|
||||
if dim_proj > dim:
|
||||
raise ValueError(
|
||||
f"Cannot project to a higher dimensionality (dim_proj={dim_proj} > dim={dim})"
|
||||
)
|
||||
|
||||
self.k_sim = k_sim
|
||||
self.dim = dim
|
||||
self.dim_proj = dim_proj
|
||||
self.r_reps = r_reps
|
||||
# Create r_reps independent SimHash projections for robustness
|
||||
generator = np.random.default_rng(random_seed)
|
||||
self.simhash_projections = [
|
||||
SimHashProjection(k_sim=self.k_sim, dim=self.dim, random_generator=generator)
|
||||
for _ in range(r_reps)
|
||||
]
|
||||
# Random projection matrices with entries from {-1, +1} for each repetition
|
||||
self.dim_reduction_projections = generator.choice([-1, 1], size=(r_reps, dim, dim_proj))
|
||||
|
||||
@classmethod
|
||||
def from_multivector_model(
|
||||
cls,
|
||||
model: MultiVectorModel,
|
||||
k_sim: int = 5,
|
||||
dim_proj: int = 16,
|
||||
r_reps: int = 20, # noqa[naming]
|
||||
random_seed: int = 42,
|
||||
) -> "Muvera":
|
||||
"""
|
||||
Create a Muvera instance from a multi-vector embedding model.
|
||||
|
||||
This class method provides a convenient way to initialize a MUVERA
|
||||
that is compatible with a given multi-vector model by automatically extracting
|
||||
the embedding dimensionality from the model.
|
||||
|
||||
Args:
|
||||
model (MultiVectorModel): A late interaction text or multimodal embedding model
|
||||
that provides multi-vector embeddings. Must have an
|
||||
`embedding_size` attribute specifying the dimensionality
|
||||
of individual vectors.
|
||||
k_sim (int, optional): Number of SimHash functions (creates 2^k_sim clusters).
|
||||
Defaults to 5.
|
||||
dim_proj (int, optional): Dimensionality after random projection (must be <= model's
|
||||
embedding_size). Defaults to 16.
|
||||
r_reps (int, optional): Number of random projection repetitions for robustness.
|
||||
Defaults to 20.
|
||||
random_seed (int, optional): Seed for random number generator to ensure
|
||||
reproducible results. Defaults to 42.
|
||||
|
||||
Returns:
|
||||
Muvera: A configured MUVERA instance ready to process embeddings from the given model.
|
||||
|
||||
Raises:
|
||||
ValueError: If dim_proj > model.embedding_size (cannot project to higher dimensionality)
|
||||
|
||||
Example:
|
||||
>>> from fastembed import LateInteractionTextEmbedding
|
||||
>>> model = LateInteractionTextEmbedding(model_name="colbert-ir/colbertv2.0")
|
||||
>>> muvera = Muvera.from_multivector_model(
|
||||
... model=model,
|
||||
... k_sim=6,
|
||||
... dim_proj=32
|
||||
... )
|
||||
>>> # Now use postprocessor with embeddings from the model
|
||||
>>> embeddings = np.array(list(model.embed(["sample text"])))
|
||||
>>> fde = muvera.process_document(embeddings[0])
|
||||
"""
|
||||
return cls(
|
||||
dim=model.embedding_size,
|
||||
k_sim=k_sim,
|
||||
dim_proj=dim_proj,
|
||||
r_reps=r_reps,
|
||||
random_seed=random_seed,
|
||||
)
|
||||
|
||||
def _get_output_dimension(self) -> int:
|
||||
"""
|
||||
Get the output dimension of the MUVERA algorithm.
|
||||
|
||||
Returns:
|
||||
int: Output dimension (r_reps * num_partitions * dim_proj) where b = 2^k_sim
|
||||
"""
|
||||
num_partitions = 2**self.k_sim
|
||||
return self.r_reps * num_partitions * self.dim_proj
|
||||
|
||||
@property
|
||||
def embedding_size(self) -> int:
|
||||
return self._get_output_dimension()
|
||||
|
||||
def process_document(self, vectors: NumpyArray) -> NumpyArray:
|
||||
"""
|
||||
Encode a document's vectors into a Fixed Dimensional Encoding (FDE).
|
||||
|
||||
Uses document-specific settings: normalizes cluster centers by vector count
|
||||
and fills empty clusters using Hamming distance-based selection.
|
||||
|
||||
Args:
|
||||
vectors (NumpyArray): Document vectors of shape (n_tokens, dim)
|
||||
|
||||
Returns:
|
||||
NumpyArray: Fixed dimensional encodings of shape (r_reps * b * dim_proj,)
|
||||
"""
|
||||
return self.process(vectors, fill_empty_clusters=True, normalize_by_count=True)
|
||||
|
||||
def process_query(self, vectors: NumpyArray) -> NumpyArray:
|
||||
"""
|
||||
Encode a query's vectors into a Fixed Dimensional Encoding (FDE).
|
||||
|
||||
Uses query-specific settings: no normalization by count and no empty
|
||||
cluster filling to preserve query vector magnitudes.
|
||||
|
||||
Args:
|
||||
vectors (NumpyArray]): Query vectors of shape (n_tokens, dim)
|
||||
|
||||
Returns:
|
||||
NumpyArray: Fixed dimensional encoding of shape (r_reps * b * dim_proj,)
|
||||
"""
|
||||
return self.process(vectors, fill_empty_clusters=False, normalize_by_count=False)
|
||||
|
||||
def process(
|
||||
self,
|
||||
vectors: NumpyArray,
|
||||
fill_empty_clusters: bool = True,
|
||||
normalize_by_count: bool = True,
|
||||
) -> NumpyArray:
|
||||
"""
|
||||
Core encoding method that transforms variable-length vector sequences into FDEs.
|
||||
|
||||
The encoding process:
|
||||
1. For each of r_reps random projections:
|
||||
a. Assign vectors to clusters using SimHash
|
||||
b. Compute cluster centers (sum of vectors in each cluster)
|
||||
c. Optionally normalize by cluster size
|
||||
d. Fill empty clusters using Hamming distance if requested
|
||||
e. Apply random projection for dimensionality reduction
|
||||
f. Flatten cluster centers into a vector
|
||||
2. Concatenate all projection results
|
||||
|
||||
Args:
|
||||
vectors (np.ndarray): Input vectors of shape (n_vectors, dim)
|
||||
fill_empty_clusters (bool): Whether to fill empty clusters using nearest
|
||||
vectors based on Hamming distance of cluster IDs
|
||||
normalize_by_count (bool): Whether to normalize cluster centers by the
|
||||
number of vectors assigned to each cluster
|
||||
|
||||
Returns:
|
||||
np.ndarray: Fixed dimensional encoding of shape (r_reps * b * dim_proj)
|
||||
where B = 2^k_sim is the number of clusters
|
||||
|
||||
Raises:
|
||||
AssertionError: If input vectors don't have expected dimensionality
|
||||
"""
|
||||
assert (
|
||||
vectors.shape[1] == self.dim
|
||||
), f"Expected vectors of shape (n, {self.dim}), got {vectors.shape}"
|
||||
|
||||
# Store results from each random projection
|
||||
output_vectors = []
|
||||
|
||||
# num of space partitions in SimHash
|
||||
num_partitions = 2**self.k_sim
|
||||
cluster_center_ids = np.arange(num_partitions)
|
||||
precomputed_hamming_matrix = (
|
||||
hamming_distance_matrix(cluster_center_ids) if fill_empty_clusters else None
|
||||
)
|
||||
|
||||
for projection_index, simhash in enumerate(self.simhash_projections):
|
||||
# Initialize cluster centers and count vectors assigned to each cluster
|
||||
cluster_centers = np.zeros((num_partitions, self.dim))
|
||||
cluster_center_id_to_vectors: dict[int, list[int]] = {
|
||||
cluster_center_id: [] for cluster_center_id in cluster_center_ids
|
||||
}
|
||||
cluster_vector_counts = None
|
||||
empty_mask = None
|
||||
|
||||
# Assign each vector to its cluster and accumulate cluster centers
|
||||
vector_cluster_ids = simhash.get_cluster_ids(vectors)
|
||||
for cluster_id, (vec_idx, vec) in zip(vector_cluster_ids, enumerate(vectors)):
|
||||
cluster_centers[cluster_id] += vec
|
||||
cluster_center_id_to_vectors[cluster_id].append(vec_idx)
|
||||
|
||||
if normalize_by_count or fill_empty_clusters:
|
||||
cluster_vector_counts = np.bincount(vector_cluster_ids, minlength=num_partitions)
|
||||
empty_mask = cluster_vector_counts == 0
|
||||
|
||||
if normalize_by_count:
|
||||
assert empty_mask is not None
|
||||
assert cluster_vector_counts is not None
|
||||
non_empty_mask = ~empty_mask
|
||||
cluster_centers[non_empty_mask] /= cluster_vector_counts[non_empty_mask][:, None]
|
||||
|
||||
# Fill empty clusters using vectors with minimum Hamming distance
|
||||
if fill_empty_clusters:
|
||||
assert empty_mask is not None
|
||||
assert precomputed_hamming_matrix is not None
|
||||
masked_hamming = np.where(
|
||||
empty_mask[None, :], MAX_HAMMING_DISTANCE, precomputed_hamming_matrix
|
||||
)
|
||||
nearest_non_empty = np.argmin(masked_hamming, axis=1)
|
||||
fill_vectors = np.array(
|
||||
[
|
||||
vectors[cluster_center_id_to_vectors[cluster_id][0]]
|
||||
for cluster_id in nearest_non_empty[empty_mask]
|
||||
]
|
||||
).reshape(-1, self.dim)
|
||||
cluster_centers[empty_mask] = fill_vectors
|
||||
|
||||
# Apply random projection for dimensionality reduction if needed
|
||||
if self.dim_proj < self.dim:
|
||||
dim_reduction_projection = self.dim_reduction_projections[
|
||||
projection_index
|
||||
] # Get projection matrix for this repetition
|
||||
projected_centers = (1 / np.sqrt(self.dim_proj)) * (
|
||||
cluster_centers @ dim_reduction_projection
|
||||
)
|
||||
|
||||
# Flatten cluster centers into a single vector and add to output
|
||||
output_vectors.append(projected_centers.flatten())
|
||||
continue
|
||||
|
||||
# If no projection needed (dim_proj == dim), use original cluster centers
|
||||
output_vectors.append(cluster_centers.flatten())
|
||||
|
||||
# Concatenate results from all R_reps projections into final FDE
|
||||
return np.concatenate(output_vectors)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
v_arrs = np.random.randn(10, 100, 128)
|
||||
muvera = Muvera(128, 4, 8, 20, 42)
|
||||
|
||||
for v_arr in v_arrs:
|
||||
muvera.process(v_arr) # type: ignore
|
||||
@@ -0,0 +1,78 @@
|
||||
from typing import Sequence, Any, Type
|
||||
|
||||
from fastembed.common import OnnxProvider
|
||||
from fastembed.common.model_description import BaseModelDescription
|
||||
from fastembed.common.types import Device
|
||||
from fastembed.rerank.cross_encoder.onnx_text_cross_encoder import OnnxTextCrossEncoder
|
||||
from fastembed.rerank.cross_encoder.onnx_text_model import TextRerankerWorker
|
||||
|
||||
|
||||
class CustomTextCrossEncoder(OnnxTextCrossEncoder):
|
||||
SUPPORTED_MODELS: list[BaseModelDescription] = []
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str | None = None,
|
||||
threads: int | None = None,
|
||||
providers: Sequence[OnnxProvider] | None = None,
|
||||
cuda: bool | Device = Device.AUTO,
|
||||
device_ids: list[int] | None = None,
|
||||
lazy_load: bool = False,
|
||||
device_id: int | None = None,
|
||||
specific_model_path: str | None = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
super().__init__(
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
threads=threads,
|
||||
providers=providers,
|
||||
cuda=cuda,
|
||||
device_ids=device_ids,
|
||||
lazy_load=lazy_load,
|
||||
device_id=device_id,
|
||||
specific_model_path=specific_model_path,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _list_supported_models(cls) -> list[BaseModelDescription]:
|
||||
return cls.SUPPORTED_MODELS
|
||||
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type[TextRerankerWorker]:
|
||||
return CustomTextCrossEncoderWorker
|
||||
|
||||
def _get_worker_init_kwargs(self) -> dict[str, Any]:
|
||||
return {"model_description": self.model_description}
|
||||
|
||||
@classmethod
|
||||
def add_model(
|
||||
cls,
|
||||
model_description: BaseModelDescription,
|
||||
) -> None:
|
||||
cls.SUPPORTED_MODELS.append(model_description)
|
||||
|
||||
|
||||
class CustomTextCrossEncoderWorker(TextRerankerWorker):
|
||||
def init_embedding(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
model_description: BaseModelDescription | None = None,
|
||||
**kwargs: Any,
|
||||
) -> CustomTextCrossEncoder:
|
||||
if model_description is None:
|
||||
raise ValueError(
|
||||
"`model_description` is required to initialize a custom model in a worker "
|
||||
"process, it is provided by `CustomTextCrossEncoder._get_worker_init_kwargs`"
|
||||
)
|
||||
# custom models live in a class-level registry, which spawned workers don't inherit
|
||||
CustomTextCrossEncoder.add_model(model_description)
|
||||
return CustomTextCrossEncoder(
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
threads=1,
|
||||
**kwargs,
|
||||
)
|
||||
@@ -1,9 +1,10 @@
|
||||
from typing import Any, Iterable, Optional, Sequence, Type
|
||||
from typing import Any, Iterable, Sequence, Type
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from fastembed.common import OnnxProvider
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.common.types import Device
|
||||
from fastembed.common.utils import define_cache_dir
|
||||
from fastembed.rerank.cross_encoder.onnx_text_model import (
|
||||
OnnxCrossEncoderModel,
|
||||
@@ -77,14 +78,14 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
|
||||
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,
|
||||
cache_dir: str | None = None,
|
||||
threads: int | None = None,
|
||||
providers: Sequence[OnnxProvider] | None = None,
|
||||
cuda: bool | Device = Device.AUTO,
|
||||
device_ids: list[int] | None = None,
|
||||
lazy_load: bool = False,
|
||||
device_id: Optional[int] = None,
|
||||
specific_model_path: Optional[str] = None,
|
||||
device_id: int | None = None,
|
||||
specific_model_path: str | None = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
"""
|
||||
@@ -96,10 +97,11 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
|
||||
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.
|
||||
cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`
|
||||
Defaults to Device.AUTO.
|
||||
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.
|
||||
workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, 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.
|
||||
@@ -111,6 +113,7 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
|
||||
super().__init__(model_name, cache_dir, threads, **kwargs)
|
||||
self.providers = providers
|
||||
self.lazy_load = lazy_load
|
||||
self._extra_session_options = self._select_exposed_session_options(kwargs)
|
||||
|
||||
# List of device ids, that can be used for data parallel processing in workers
|
||||
self.device_ids = device_ids
|
||||
@@ -123,7 +126,7 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
|
||||
)
|
||||
|
||||
# This device_id will be used if we need to load model in current process
|
||||
self.device_id: Optional[int] = None
|
||||
self.device_id: int | None = None
|
||||
if device_id is not None:
|
||||
self.device_id = device_id
|
||||
elif self.device_ids is not None:
|
||||
@@ -131,11 +134,12 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
|
||||
|
||||
self.model_description = self._get_model_description(model_name)
|
||||
self.cache_dir = str(define_cache_dir(cache_dir))
|
||||
self._specific_model_path = specific_model_path
|
||||
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,
|
||||
specific_model_path=self._specific_model_path,
|
||||
)
|
||||
|
||||
if not self.lazy_load:
|
||||
@@ -149,6 +153,7 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_id=self.device_id,
|
||||
extra_session_options=self._extra_session_options,
|
||||
)
|
||||
|
||||
def rerank(
|
||||
@@ -177,7 +182,7 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
|
||||
self,
|
||||
pairs: Iterable[tuple[str, str]],
|
||||
batch_size: int = 64,
|
||||
parallel: Optional[int] = None,
|
||||
parallel: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[float]:
|
||||
yield from self._rerank_pairs(
|
||||
@@ -189,6 +194,9 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_ids=self.device_ids,
|
||||
local_files_only=self._local_files_only,
|
||||
specific_model_path=self._specific_model_path,
|
||||
extra_session_options=self._extra_session_options,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@@ -196,9 +204,25 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
|
||||
def _get_worker_class(cls) -> Type[TextRerankerWorker]:
|
||||
return TextCrossEncoderWorker
|
||||
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[float]:
|
||||
def _post_process_onnx_output(
|
||||
self, output: OnnxOutputContext, **kwargs: Any
|
||||
) -> Iterable[float]:
|
||||
return (float(elem) for elem in output.model_output)
|
||||
|
||||
def token_count(
|
||||
self, pairs: Iterable[tuple[str, str]], batch_size: int = 1024, **kwargs: Any
|
||||
) -> int:
|
||||
"""Returns the number of tokens in the pairs.
|
||||
|
||||
Args:
|
||||
pairs: Iterable of tuples, where each tuple contains a query and a document to be tokenized
|
||||
batch_size: Batch size for tokenizing
|
||||
|
||||
Returns:
|
||||
token count: overall number of tokens in the pairs
|
||||
"""
|
||||
return self._token_count(pairs, batch_size=batch_size, **kwargs)
|
||||
|
||||
|
||||
class TextCrossEncoderWorker(TextRerankerWorker):
|
||||
def init_embedding(
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import os
|
||||
from multiprocessing import get_all_start_methods
|
||||
from pathlib import Path
|
||||
from typing import Any, Iterable, Optional, Sequence, Type
|
||||
from typing import Any, Iterable, Sequence, Type
|
||||
|
||||
import numpy as np
|
||||
from tokenizers import Encoding
|
||||
@@ -12,14 +12,14 @@ from fastembed.common.onnx_model import (
|
||||
OnnxOutputContext,
|
||||
OnnxProvider,
|
||||
)
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.common.types import NumpyArray, Device
|
||||
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
|
||||
ONNX_OUTPUT_NAMES: list[str] | None = None
|
||||
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type["TextRerankerWorker"]:
|
||||
@@ -29,10 +29,11 @@ class OnnxCrossEncoderModel(OnnxModel[float]):
|
||||
self,
|
||||
model_dir: Path,
|
||||
model_file: str,
|
||||
threads: Optional[int],
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
cuda: bool = False,
|
||||
device_id: Optional[int] = None,
|
||||
threads: int | None,
|
||||
providers: Sequence[OnnxProvider] | None = None,
|
||||
cuda: bool | Device = Device.AUTO,
|
||||
device_id: int | None = None,
|
||||
extra_session_options: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
super()._load_onnx_model(
|
||||
model_dir=model_dir,
|
||||
@@ -41,6 +42,7 @@ class OnnxCrossEncoderModel(OnnxModel[float]):
|
||||
providers=providers,
|
||||
cuda=cuda,
|
||||
device_id=device_id,
|
||||
extra_session_options=extra_session_options,
|
||||
)
|
||||
self.tokenizer, _ = load_tokenizer(model_dir=model_dir)
|
||||
assert self.tokenizer is not None
|
||||
@@ -90,10 +92,13 @@ class OnnxCrossEncoderModel(OnnxModel[float]):
|
||||
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,
|
||||
parallel: int | None = None,
|
||||
providers: Sequence[OnnxProvider] | None = None,
|
||||
cuda: bool | Device = Device.AUTO,
|
||||
device_ids: list[int] | None = None,
|
||||
local_files_only: bool = False,
|
||||
specific_model_path: str | None = None,
|
||||
extra_session_options: dict[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[float]:
|
||||
is_small = False
|
||||
@@ -120,9 +125,15 @@ class OnnxCrossEncoderModel(OnnxModel[float]):
|
||||
"model_name": model_name,
|
||||
"cache_dir": cache_dir,
|
||||
"providers": providers,
|
||||
"local_files_only": local_files_only,
|
||||
"specific_model_path": specific_model_path,
|
||||
**kwargs,
|
||||
**self._get_worker_init_kwargs(),
|
||||
}
|
||||
|
||||
if extra_session_options is not None:
|
||||
params.update(extra_session_options)
|
||||
|
||||
pool = ParallelWorkerPool(
|
||||
num_workers=parallel or 1,
|
||||
worker=self._get_worker_class(),
|
||||
@@ -133,7 +144,18 @@ class OnnxCrossEncoderModel(OnnxModel[float]):
|
||||
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]:
|
||||
def _post_process_onnx_output(
|
||||
self, output: OnnxOutputContext, **kwargs: Any
|
||||
) -> Iterable[float]:
|
||||
"""Post-process the ONNX model output to convert it into a usable format.
|
||||
|
||||
Args:
|
||||
output (OnnxOutputContext): The raw output from the ONNX model.
|
||||
**kwargs: Additional keyword arguments that may be needed by specific implementations.
|
||||
|
||||
Returns:
|
||||
Iterable[float]: Post-processed output as an iterable of float values.
|
||||
"""
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
def _preprocess_onnx_input(
|
||||
@@ -144,6 +166,20 @@ class OnnxCrossEncoderModel(OnnxModel[float]):
|
||||
"""
|
||||
return onnx_input
|
||||
|
||||
def _token_count(
|
||||
self, pairs: Iterable[tuple[str, str]], batch_size: int = 1024, **_: Any
|
||||
) -> int:
|
||||
if not hasattr(self, "model") or self.model is None:
|
||||
self.load_onnx_model() # loads the tokenizer as well
|
||||
|
||||
token_num = 0
|
||||
assert self.tokenizer is not None
|
||||
for batch in iter_batch(pairs, batch_size):
|
||||
for tokens in self.tokenizer.encode_batch(batch):
|
||||
token_num += sum(tokens.attention_mask)
|
||||
|
||||
return token_num
|
||||
|
||||
|
||||
class TextRerankerWorker(EmbeddingWorker[float]):
|
||||
def __init__(
|
||||
|
||||
@@ -1,15 +1,22 @@
|
||||
from typing import Any, Iterable, Optional, Sequence, Type
|
||||
from typing import Any, Iterable, Sequence, Type
|
||||
from dataclasses import asdict
|
||||
|
||||
from fastembed.common import OnnxProvider
|
||||
from fastembed.common.types import Device
|
||||
from fastembed.rerank.cross_encoder.onnx_text_cross_encoder import OnnxTextCrossEncoder
|
||||
from fastembed.rerank.cross_encoder.custom_text_cross_encoder import CustomTextCrossEncoder
|
||||
|
||||
from fastembed.rerank.cross_encoder.text_cross_encoder_base import TextCrossEncoderBase
|
||||
from fastembed.common.model_description import BaseModelDescription
|
||||
from fastembed.common.model_description import (
|
||||
ModelSource,
|
||||
BaseModelDescription,
|
||||
)
|
||||
|
||||
|
||||
class TextCrossEncoder(TextCrossEncoderBase):
|
||||
CROSS_ENCODER_REGISTRY: list[Type[TextCrossEncoderBase]] = [
|
||||
OnnxTextCrossEncoder,
|
||||
CustomTextCrossEncoder,
|
||||
]
|
||||
|
||||
@classmethod
|
||||
@@ -47,11 +54,11 @@ class TextCrossEncoder(TextCrossEncoderBase):
|
||||
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,
|
||||
cache_dir: str | None = None,
|
||||
threads: int | None = None,
|
||||
providers: Sequence[OnnxProvider] | None = None,
|
||||
cuda: bool | Device = Device.AUTO,
|
||||
device_ids: list[int] | None = None,
|
||||
lazy_load: bool = False,
|
||||
**kwargs: Any,
|
||||
):
|
||||
@@ -96,7 +103,7 @@ class TextCrossEncoder(TextCrossEncoderBase):
|
||||
self,
|
||||
pairs: Iterable[tuple[str, str]],
|
||||
batch_size: int = 64,
|
||||
parallel: Optional[int] = None,
|
||||
parallel: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[float]:
|
||||
"""
|
||||
@@ -124,3 +131,48 @@ class TextCrossEncoder(TextCrossEncoderBase):
|
||||
yield from self.model.rerank_pairs(
|
||||
pairs, batch_size=batch_size, parallel=parallel, **kwargs
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def add_custom_model(
|
||||
cls,
|
||||
model: str,
|
||||
sources: ModelSource,
|
||||
model_file: str = "onnx/model.onnx",
|
||||
description: str = "",
|
||||
license: str = "",
|
||||
size_in_gb: float = 0.0,
|
||||
additional_files: list[str] | None = None,
|
||||
) -> None:
|
||||
registered_models = cls._list_supported_models()
|
||||
for registered_model in registered_models:
|
||||
if model == registered_model.model:
|
||||
raise ValueError(
|
||||
f"Model {model} is already registered in CrossEncoderModel, if you still want to add this model, "
|
||||
f"please use another model name"
|
||||
)
|
||||
|
||||
CustomTextCrossEncoder.add_model(
|
||||
BaseModelDescription(
|
||||
model=model,
|
||||
sources=sources,
|
||||
model_file=model_file,
|
||||
description=description,
|
||||
license=license,
|
||||
size_in_GB=size_in_gb,
|
||||
additional_files=additional_files or [],
|
||||
)
|
||||
)
|
||||
|
||||
def token_count(
|
||||
self, pairs: Iterable[tuple[str, str]], batch_size: int = 1024, **kwargs: Any
|
||||
) -> int:
|
||||
"""Returns the number of tokens in the pairs.
|
||||
|
||||
Args:
|
||||
pairs: Iterable of tuples, where each tuple contains a query and a document to be tokenized
|
||||
batch_size: Batch size for tokenizing
|
||||
|
||||
Returns:
|
||||
token count: overall number of tokens in the pairs
|
||||
"""
|
||||
return self.model.token_count(pairs, batch_size=batch_size, **kwargs)
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from typing import Any, Iterable, Optional
|
||||
from typing import Any, Iterable
|
||||
|
||||
from fastembed.common.model_description import BaseModelDescription
|
||||
from fastembed.common.model_management import ModelManagement
|
||||
@@ -8,8 +8,8 @@ class TextCrossEncoderBase(ModelManagement[BaseModelDescription]):
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
cache_dir: str | None = None,
|
||||
threads: int | None = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
self.model_name = model_name
|
||||
@@ -41,7 +41,7 @@ class TextCrossEncoderBase(ModelManagement[BaseModelDescription]):
|
||||
self,
|
||||
pairs: Iterable[tuple[str, str]],
|
||||
batch_size: int = 64,
|
||||
parallel: Optional[int] = None,
|
||||
parallel: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[float]:
|
||||
"""Rerank query-document pairs.
|
||||
@@ -57,3 +57,7 @@ class TextCrossEncoderBase(ModelManagement[BaseModelDescription]):
|
||||
Iterable[float]: Scores for each individual pair
|
||||
"""
|
||||
raise NotImplementedError("This method should be overridden by subclasses")
|
||||
|
||||
def token_count(self, pairs: Iterable[tuple[str, str]], **kwargs: Any) -> int:
|
||||
"""Returns the number of tokens in the pairs."""
|
||||
raise NotImplementedError("This method should be overridden by subclasses")
|
||||
|
||||
+27
-23
@@ -2,7 +2,7 @@ import os
|
||||
from collections import defaultdict
|
||||
from multiprocessing import get_all_start_methods
|
||||
from pathlib import Path
|
||||
from typing import Any, Iterable, Optional, Type, Union
|
||||
from typing import Any, Iterable, Type
|
||||
|
||||
import mmh3
|
||||
import numpy as np
|
||||
@@ -21,13 +21,9 @@ from fastembed.sparse.sparse_embedding_base import (
|
||||
from fastembed.sparse.utils.tokenizer import SimpleTokenizer
|
||||
from fastembed.common.model_description import SparseModelDescription, ModelSource
|
||||
|
||||
|
||||
supported_languages = [
|
||||
"arabic",
|
||||
"azerbaijani",
|
||||
"basque",
|
||||
"bengali",
|
||||
"catalan",
|
||||
"chinese",
|
||||
"danish",
|
||||
"dutch",
|
||||
"english",
|
||||
@@ -35,21 +31,15 @@ supported_languages = [
|
||||
"french",
|
||||
"german",
|
||||
"greek",
|
||||
"hebrew",
|
||||
"hinglish",
|
||||
"hungarian",
|
||||
"indonesian",
|
||||
"italian",
|
||||
"kazakh",
|
||||
"nepali",
|
||||
"norwegian",
|
||||
"portuguese",
|
||||
"romanian",
|
||||
"russian",
|
||||
"slovene",
|
||||
"spanish",
|
||||
"swedish",
|
||||
"tajik",
|
||||
"tamil",
|
||||
"turkish",
|
||||
]
|
||||
|
||||
@@ -101,14 +91,14 @@ class Bm25(SparseTextEmbeddingBase):
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: Optional[str] = None,
|
||||
cache_dir: str | None = None,
|
||||
k: float = 1.2,
|
||||
b: float = 0.75,
|
||||
avg_len: float = 256.0,
|
||||
language: str = "english",
|
||||
token_max_length: int = 40,
|
||||
disable_stemmer: bool = False,
|
||||
specific_model_path: Optional[str] = None,
|
||||
specific_model_path: str | None = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
super().__init__(model_name, cache_dir, **kwargs)
|
||||
@@ -125,11 +115,12 @@ class Bm25(SparseTextEmbeddingBase):
|
||||
model_description = self._get_model_description(model_name)
|
||||
self.cache_dir = str(define_cache_dir(cache_dir))
|
||||
|
||||
self._specific_model_path = specific_model_path
|
||||
self._model_dir = self.download_model(
|
||||
model_description,
|
||||
self.cache_dir,
|
||||
local_files_only=self._local_files_only,
|
||||
specific_model_path=specific_model_path,
|
||||
specific_model_path=self._specific_model_path,
|
||||
)
|
||||
|
||||
self.token_max_length = token_max_length
|
||||
@@ -167,9 +158,11 @@ class Bm25(SparseTextEmbeddingBase):
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
documents: Union[str, Iterable[str]],
|
||||
documents: str | Iterable[str],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
parallel: int | None = None,
|
||||
local_files_only: bool = False,
|
||||
specific_model_path: str | None = None,
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
is_small = False
|
||||
|
||||
@@ -198,6 +191,8 @@ class Bm25(SparseTextEmbeddingBase):
|
||||
"language": self.language,
|
||||
"token_max_length": self.token_max_length,
|
||||
"disable_stemmer": self.disable_stemmer,
|
||||
"local_files_only": local_files_only,
|
||||
"specific_model_path": specific_model_path,
|
||||
}
|
||||
pool = ParallelWorkerPool(
|
||||
num_workers=parallel or 1,
|
||||
@@ -210,9 +205,9 @@ class Bm25(SparseTextEmbeddingBase):
|
||||
|
||||
def embed(
|
||||
self,
|
||||
documents: Union[str, Iterable[str]],
|
||||
documents: str | Iterable[str],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
parallel: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
"""
|
||||
@@ -236,6 +231,8 @@ class Bm25(SparseTextEmbeddingBase):
|
||||
documents=documents,
|
||||
batch_size=batch_size,
|
||||
parallel=parallel,
|
||||
local_files_only=self._local_files_only,
|
||||
specific_model_path=self._specific_model_path,
|
||||
)
|
||||
|
||||
def _stem(self, tokens: list[str]) -> list[str]:
|
||||
@@ -271,6 +268,15 @@ class Bm25(SparseTextEmbeddingBase):
|
||||
embeddings.append(SparseEmbedding.from_dict(token_id2value))
|
||||
return embeddings
|
||||
|
||||
def token_count(self, texts: str | Iterable[str], **kwargs: Any) -> int:
|
||||
token_num = 0
|
||||
texts = [texts] if isinstance(texts, str) else texts
|
||||
for text in texts:
|
||||
document = remove_non_alphanumeric(text)
|
||||
tokens = self.tokenizer.tokenize(document)
|
||||
token_num += len(tokens)
|
||||
return token_num
|
||||
|
||||
def _term_frequency(self, tokens: list[str]) -> dict[int, float]:
|
||||
"""Calculate the term frequency part of the BM25 formula.
|
||||
|
||||
@@ -305,9 +311,7 @@ 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: Any
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
def query_embed(self, query: 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.
|
||||
"""
|
||||
|
||||
+44
-21
@@ -1,7 +1,7 @@
|
||||
import math
|
||||
import string
|
||||
from pathlib import Path
|
||||
from typing import Any, Iterable, Optional, Sequence, Type, Union
|
||||
from typing import Any, Iterable, Sequence, Type
|
||||
|
||||
import mmh3
|
||||
import numpy as np
|
||||
@@ -9,6 +9,7 @@ from py_rust_stemmers import SnowballStemmer
|
||||
|
||||
from fastembed.common import OnnxProvider
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.common.types import Device
|
||||
from fastembed.common.utils import define_cache_dir
|
||||
from fastembed.sparse.sparse_embedding_base import (
|
||||
SparseEmbedding,
|
||||
@@ -31,9 +32,17 @@ supported_bm42_models: list[SparseModelDescription] = [
|
||||
),
|
||||
]
|
||||
|
||||
MODEL_TO_LANGUAGE = {
|
||||
|
||||
_MODEL_TO_LANGUAGE = {
|
||||
"Qdrant/bm42-all-minilm-l6-v2-attentions": "english",
|
||||
}
|
||||
MODEL_TO_LANGUAGE = {
|
||||
model_name.lower(): language for model_name, language in _MODEL_TO_LANGUAGE.items()
|
||||
}
|
||||
|
||||
|
||||
def get_language_by_model_name(model_name: str) -> str:
|
||||
return MODEL_TO_LANGUAGE[model_name.lower()]
|
||||
|
||||
|
||||
class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
@@ -57,15 +66,15 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
cache_dir: str | None = None,
|
||||
threads: int | None = None,
|
||||
providers: Sequence[OnnxProvider] | None = None,
|
||||
alpha: float = 0.5,
|
||||
cuda: bool = False,
|
||||
device_ids: Optional[list[int]] = None,
|
||||
cuda: bool | Device = Device.AUTO,
|
||||
device_ids: list[int] | None = None,
|
||||
lazy_load: bool = False,
|
||||
device_id: Optional[int] = None,
|
||||
specific_model_path: Optional[str] = None,
|
||||
device_id: int | None = None,
|
||||
specific_model_path: str | None = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
"""
|
||||
@@ -79,10 +88,11 @@ 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.
|
||||
cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`
|
||||
Defaults to Device.AUTO.
|
||||
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.
|
||||
workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, 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.
|
||||
@@ -95,13 +105,14 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
super().__init__(model_name, cache_dir, threads, **kwargs)
|
||||
self.providers = providers
|
||||
self.lazy_load = lazy_load
|
||||
self._extra_session_options = self._select_exposed_session_options(kwargs)
|
||||
|
||||
# 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
|
||||
self.device_id: int | None = None
|
||||
if device_id is not None:
|
||||
self.device_id = device_id
|
||||
elif self.device_ids is not None:
|
||||
@@ -110,11 +121,12 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
self.model_description = self._get_model_description(model_name)
|
||||
self.cache_dir = str(define_cache_dir(cache_dir))
|
||||
|
||||
self._specific_model_path = specific_model_path
|
||||
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,
|
||||
specific_model_path=self._specific_model_path,
|
||||
)
|
||||
|
||||
self.invert_vocab: dict[int, str] = {}
|
||||
@@ -123,7 +135,7 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
self.special_tokens_ids: set[int] = set()
|
||||
self.punctuation = set(string.punctuation)
|
||||
self.stopwords = set(self._load_stopwords(self._model_dir))
|
||||
self.stemmer = SnowballStemmer(MODEL_TO_LANGUAGE[model_name])
|
||||
self.stemmer = SnowballStemmer(get_language_by_model_name(self.model_name))
|
||||
self.alpha = alpha
|
||||
|
||||
if not self.lazy_load:
|
||||
@@ -137,6 +149,7 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_id=self.device_id,
|
||||
extra_session_options=self._extra_session_options,
|
||||
)
|
||||
|
||||
for token, idx in self.tokenizer.get_vocab().items(): # type: ignore[union-attr]
|
||||
@@ -217,7 +230,9 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
|
||||
return new_vector
|
||||
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[SparseEmbedding]:
|
||||
def _post_process_onnx_output(
|
||||
self, output: OnnxOutputContext, **kwargs: Any
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
if output.input_ids is None:
|
||||
raise ValueError("input_ids must be provided for document post-processing")
|
||||
|
||||
@@ -269,9 +284,9 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
|
||||
def embed(
|
||||
self,
|
||||
documents: Union[str, Iterable[str]],
|
||||
documents: str | Iterable[str],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
parallel: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
"""
|
||||
@@ -299,6 +314,9 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
cuda=self.cuda,
|
||||
device_ids=self.device_ids,
|
||||
alpha=self.alpha,
|
||||
local_files_only=self._local_files_only,
|
||||
specific_model_path=self._specific_model_path,
|
||||
extra_session_options=self._extra_session_options,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@@ -309,9 +327,7 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
result[token_id] = 1.0
|
||||
return result
|
||||
|
||||
def query_embed(
|
||||
self, query: Union[str, Iterable[str]], **kwargs: Any
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
def query_embed(self, query: 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.
|
||||
@@ -336,6 +352,13 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
def _get_worker_class(cls) -> Type[TextEmbeddingWorker[SparseEmbedding]]:
|
||||
return Bm42TextEmbeddingWorker
|
||||
|
||||
def token_count(
|
||||
self, texts: str | Iterable[str], batch_size: int = 1024, **kwargs: Any
|
||||
) -> int:
|
||||
if not hasattr(self, "model") or self.model is None:
|
||||
self.load_onnx_model() # loads the tokenizer as well
|
||||
return self._token_count(texts, batch_size=batch_size, **kwargs)
|
||||
|
||||
|
||||
class Bm42TextEmbeddingWorker(TextEmbeddingWorker[SparseEmbedding]):
|
||||
def init_embedding(self, model_name: str, cache_dir: str, **kwargs: Any) -> Bm42:
|
||||
|
||||
@@ -0,0 +1,246 @@
|
||||
import json
|
||||
from typing import Any, Iterable, Sequence, Type
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastembed.common import OnnxProvider
|
||||
from fastembed.common.model_description import ModelSource, SparseModelDescription
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.common.preprocessor_utils import load_tokenizer
|
||||
from fastembed.common.types import Device
|
||||
from fastembed.common.utils import define_cache_dir, iter_batch
|
||||
from fastembed.sparse.sparse_embedding_base import (
|
||||
SparseEmbedding,
|
||||
SparseTextEmbeddingBase,
|
||||
)
|
||||
from fastembed.text.onnx_text_model import OnnxTextModel, TextEmbeddingWorker
|
||||
|
||||
IDF_FILE = "idf.json"
|
||||
|
||||
supported_if_splade_models: list[SparseModelDescription] = [
|
||||
SparseModelDescription(
|
||||
model="opensearch-project/opensearch-neural-sparse-encoding-doc-v3-gte",
|
||||
vocab_size=30522,
|
||||
description="Inference-free SPLADE model. Documents are expanded with an ONNX encoder at index "
|
||||
"time, queries are encoded with a tokenizer and an IDF lookup table only, "
|
||||
"without any model inference.",
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.55,
|
||||
sources=ModelSource(hf="Qdrant/opensearch-neural-sparse-encoding-doc-v3-gte"),
|
||||
model_file="model.onnx",
|
||||
additional_files=[IDF_FILE],
|
||||
requires_idf=None,
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
class IfSplade(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
"""Inference-free (asymmetric) SPLADE model.
|
||||
|
||||
Documents are encoded with a neural encoder which expands them into a sparse vocabulary-sized
|
||||
vector, while queries are encoded by tokenizing the text and looking up a precomputed IDF
|
||||
weight per token — no neural inference happens at query time.
|
||||
|
||||
Query and document embeddings are compared with a dot product.
|
||||
Special tokens are excluded from both document and query embeddings.
|
||||
"""
|
||||
|
||||
def _post_process_onnx_output(
|
||||
self, output: OnnxOutputContext, **kwargs: Any
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
if output.attention_mask is None:
|
||||
raise ValueError("attention_mask must be provided for document post-processing")
|
||||
|
||||
# Max-pool token logits over the sequence, masking out the padding
|
||||
pooled = np.max(
|
||||
output.model_output * np.expand_dims(output.attention_mask, axis=-1), axis=1
|
||||
)
|
||||
# v3 models of the opensearch-neural-sparse family use a double log activation,
|
||||
# log(1 + log(1 + relu(x))), to increase sparsity of document embeddings
|
||||
scores = np.log1p(np.log1p(np.maximum(pooled, 0.0)))
|
||||
|
||||
if self.special_tokens_ids:
|
||||
scores[:, list(self.special_tokens_ids)] = 0.0
|
||||
|
||||
for row_scores in scores:
|
||||
indices = row_scores.nonzero()[0]
|
||||
yield SparseEmbedding(values=row_scores[indices], indices=indices)
|
||||
|
||||
def token_count(
|
||||
self, texts: str | Iterable[str], batch_size: int = 1024, **kwargs: Any
|
||||
) -> int:
|
||||
# unlike `OnnxTextModel._token_count`, does not require the onnx model to be loaded
|
||||
token_num = 0
|
||||
texts = [texts] if isinstance(texts, str) else texts
|
||||
for batch in iter_batch(texts, batch_size):
|
||||
for tokens in self.tokenizer.encode_batch(batch): # type: ignore[union-attr]
|
||||
token_num += sum(tokens.attention_mask)
|
||||
return token_num
|
||||
|
||||
@classmethod
|
||||
def _list_supported_models(cls) -> list[SparseModelDescription]:
|
||||
"""Lists the supported models.
|
||||
|
||||
Returns:
|
||||
list[SparseModelDescription]: A list of SparseModelDescription objects containing the model information.
|
||||
"""
|
||||
return supported_if_splade_models
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str | None = None,
|
||||
threads: int | None = None,
|
||||
providers: Sequence[OnnxProvider] | None = None,
|
||||
cuda: bool | Device = Device.AUTO,
|
||||
device_ids: list[int] | None = None,
|
||||
lazy_load: bool = False,
|
||||
device_id: int | None = None,
|
||||
specific_model_path: str | None = 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 (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`
|
||||
Defaults to Device.
|
||||
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
|
||||
workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, 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
|
||||
self._extra_session_options = self._select_exposed_session_options(kwargs)
|
||||
|
||||
# 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: int | None = 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._specific_model_path = specific_model_path
|
||||
self._model_dir = self.download_model(
|
||||
self.model_description,
|
||||
self.cache_dir,
|
||||
local_files_only=self._local_files_only,
|
||||
specific_model_path=self._specific_model_path,
|
||||
)
|
||||
|
||||
# The tokenizer and the idf table are lightweight and are required for query embedding,
|
||||
# which does not involve any model inference, so they are loaded eagerly, while
|
||||
# `lazy_load` only defers the initialization of the onnx model
|
||||
self.tokenizer, self.special_token_to_id = load_tokenizer(model_dir=self._model_dir)
|
||||
self.special_tokens_ids: set[int] = set(self.special_token_to_id.values())
|
||||
self._token_id_to_idf = self._load_idf()
|
||||
|
||||
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,
|
||||
extra_session_options=self._extra_session_options,
|
||||
)
|
||||
|
||||
def _load_idf(self) -> dict[int, float]:
|
||||
with open(self._model_dir / IDF_FILE) as f:
|
||||
token_to_idf: dict[str, float] = json.load(f)
|
||||
|
||||
vocab: dict[str, int] = self.tokenizer.get_vocab() # type: ignore[union-attr]
|
||||
return {vocab[token]: idf for token, idf in token_to_idf.items() if token in vocab}
|
||||
|
||||
def embed(
|
||||
self,
|
||||
documents: str | Iterable[str],
|
||||
batch_size: int = 256,
|
||||
parallel: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
"""
|
||||
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,
|
||||
local_files_only=self._local_files_only,
|
||||
specific_model_path=self._specific_model_path,
|
||||
extra_session_options=self._extra_session_options,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[SparseEmbedding]:
|
||||
"""
|
||||
Encode a list of queries into list of sparse embeddings without any model inference.
|
||||
|
||||
A query is tokenized, and each unique token is assigned its IDF weight from
|
||||
a precomputed lookup table shipped with the model. Special tokens are ignored.
|
||||
"""
|
||||
if isinstance(query, str):
|
||||
query = [query]
|
||||
|
||||
for text in query:
|
||||
token_ids = set(self.tokenizer.encode(text).ids) - self.special_tokens_ids # type: ignore[union-attr]
|
||||
embedding = {
|
||||
token_id: self._token_id_to_idf[token_id]
|
||||
for token_id in sorted(token_ids)
|
||||
if token_id in self._token_id_to_idf
|
||||
}
|
||||
yield SparseEmbedding.from_dict(embedding)
|
||||
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type[TextEmbeddingWorker[SparseEmbedding]]:
|
||||
return IfSpladeEmbeddingWorker
|
||||
|
||||
|
||||
class IfSpladeEmbeddingWorker(TextEmbeddingWorker[SparseEmbedding]):
|
||||
def init_embedding(self, model_name: str, cache_dir: str, **kwargs: Any) -> IfSplade:
|
||||
return IfSplade(
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
threads=1,
|
||||
**kwargs,
|
||||
)
|
||||
@@ -0,0 +1,372 @@
|
||||
from pathlib import Path
|
||||
|
||||
from typing import Any, Sequence, Iterable, Type
|
||||
|
||||
import numpy as np
|
||||
from numpy.typing import NDArray
|
||||
from py_rust_stemmers import SnowballStemmer
|
||||
from tokenizers import Tokenizer
|
||||
|
||||
from fastembed.common.model_description import SparseModelDescription, ModelSource
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.common import OnnxProvider
|
||||
from fastembed.common.types import Device
|
||||
from fastembed.common.utils import define_cache_dir
|
||||
from fastembed.sparse.sparse_embedding_base import (
|
||||
SparseEmbedding,
|
||||
SparseTextEmbeddingBase,
|
||||
)
|
||||
from fastembed.sparse.utils.minicoil_encoder import Encoder
|
||||
from fastembed.sparse.utils.sparse_vectors_converter import SparseVectorConverter, WordEmbedding
|
||||
from fastembed.sparse.utils.vocab_resolver import VocabResolver, VocabTokenizer
|
||||
from fastembed.text.onnx_text_model import OnnxTextModel, TextEmbeddingWorker
|
||||
|
||||
|
||||
MINICOIL_MODEL_FILE = "minicoil.triplet.model.npy"
|
||||
MINICOIL_VOCAB_FILE = "minicoil.triplet.model.vocab"
|
||||
STOPWORDS_FILE = "stopwords.txt"
|
||||
|
||||
|
||||
supported_minicoil_models: list[SparseModelDescription] = [
|
||||
SparseModelDescription(
|
||||
model="Qdrant/minicoil-v1",
|
||||
vocab_size=19125,
|
||||
description="Sparse embedding model, that resolves semantic meaning of the words, "
|
||||
"while keeping exact keyword match behavior. "
|
||||
"Based on jinaai/jina-embeddings-v2-small-en-tokens",
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.09,
|
||||
sources=ModelSource(hf="Qdrant/minicoil-v1"),
|
||||
model_file="onnx/model.onnx",
|
||||
additional_files=[
|
||||
STOPWORDS_FILE,
|
||||
MINICOIL_MODEL_FILE,
|
||||
MINICOIL_VOCAB_FILE,
|
||||
],
|
||||
requires_idf=True,
|
||||
),
|
||||
]
|
||||
|
||||
_MODEL_TO_LANGUAGE = {
|
||||
"Qdrant/minicoil-v1": "english",
|
||||
}
|
||||
MODEL_TO_LANGUAGE = {
|
||||
model_name.lower(): language for model_name, language in _MODEL_TO_LANGUAGE.items()
|
||||
}
|
||||
|
||||
|
||||
def get_language_by_model_name(model_name: str) -> str:
|
||||
return MODEL_TO_LANGUAGE[model_name.lower()]
|
||||
|
||||
|
||||
class MiniCOIL(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
"""
|
||||
MiniCOIL is a sparse embedding model, that resolves semantic meaning of the words,
|
||||
while keeping exact keyword match behavior.
|
||||
|
||||
Each vocabulary token is converted into 4d component of a sparse vector, which is then weighted by the token frequency in the corpus.
|
||||
If the token is not found in the corpus, it is treated exactly like in BM25.
|
||||
`
|
||||
The model is based on `jinaai/jina-embeddings-v2-small-en-tokens`
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str | None = None,
|
||||
threads: int | None = None,
|
||||
providers: Sequence[OnnxProvider] | None = None,
|
||||
k: float = 1.2,
|
||||
b: float = 0.75,
|
||||
avg_len: float = 150.0,
|
||||
cuda: bool | Device = Device.AUTO,
|
||||
device_ids: list[int] | None = None,
|
||||
lazy_load: bool = False,
|
||||
device_id: int | None = None,
|
||||
specific_model_path: str | None = 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 providers to use for onnxruntime.
|
||||
k (float, optional): The k parameter in the BM25 formula. Defines the saturation of the term frequency.
|
||||
I.e. defines how fast the moment when additional terms stop to increase the score. Defaults to 1.2.
|
||||
b (float, optional): The b parameter in the BM25 formula. Defines the importance of the document length.
|
||||
Defaults to 0.75.
|
||||
avg_len (float, optional): The average length of the documents in the corpus. Defaults to 150.0.
|
||||
cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`
|
||||
Defaults to Device.AUTO.
|
||||
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
|
||||
workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, 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
|
||||
self.device_ids = device_ids
|
||||
self.cuda = cuda
|
||||
self.device_id = device_id
|
||||
self._extra_session_options = self._select_exposed_session_options(kwargs)
|
||||
|
||||
self.k = k
|
||||
self.b = b
|
||||
self.avg_len = avg_len
|
||||
|
||||
# Initialize class attributes
|
||||
self.tokenizer: Tokenizer | None = None
|
||||
self.invert_vocab: dict[int, str] = {}
|
||||
self.special_tokens: set[str] = set()
|
||||
self.special_tokens_ids: set[int] = set()
|
||||
self.stopwords: set[str] = set()
|
||||
self.vocab_resolver: VocabResolver | None = None
|
||||
self.encoder: Encoder | None = None
|
||||
self.output_dim: int | None = None
|
||||
self.sparse_vector_converter: SparseVectorConverter | None = None
|
||||
|
||||
self.model_description = self._get_model_description(model_name)
|
||||
self.cache_dir = str(define_cache_dir(cache_dir))
|
||||
self._specific_model_path = specific_model_path
|
||||
self._model_dir = self.download_model(
|
||||
self.model_description,
|
||||
self.cache_dir,
|
||||
local_files_only=self._local_files_only,
|
||||
specific_model_path=self._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,
|
||||
extra_session_options=self._extra_session_options,
|
||||
)
|
||||
|
||||
assert self.tokenizer is not None
|
||||
|
||||
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))
|
||||
|
||||
stemmer = SnowballStemmer(get_language_by_model_name(self.model_name))
|
||||
|
||||
self.vocab_resolver = VocabResolver(
|
||||
tokenizer=VocabTokenizer(self.tokenizer),
|
||||
stopwords=self.stopwords,
|
||||
stemmer=stemmer,
|
||||
)
|
||||
self.vocab_resolver.load_json_vocab(str(self._model_dir / MINICOIL_VOCAB_FILE))
|
||||
|
||||
weights = np.load(str(self._model_dir / MINICOIL_MODEL_FILE), mmap_mode="r")
|
||||
self.encoder = Encoder(weights)
|
||||
self.output_dim = self.encoder.output_dim
|
||||
|
||||
self.sparse_vector_converter = SparseVectorConverter(
|
||||
stopwords=self.stopwords,
|
||||
stemmer=stemmer,
|
||||
k=self.k,
|
||||
b=self.b,
|
||||
avg_len=self.avg_len,
|
||||
)
|
||||
|
||||
def token_count(
|
||||
self, texts: str | Iterable[str], batch_size: int = 1024, **kwargs: Any
|
||||
) -> int:
|
||||
return self._token_count(texts, batch_size=batch_size, **kwargs)
|
||||
|
||||
def embed(
|
||||
self,
|
||||
documents: str | Iterable[str],
|
||||
batch_size: int = 256,
|
||||
parallel: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
"""
|
||||
Encode a list of documents into list of embeddings.
|
||||
We use mean pooling with attention so that the model can handle variable-length inputs.
|
||||
|
||||
Args:
|
||||
documents: Iterator of documents or single document to embed
|
||||
batch_size: Batch size for encoding -- higher values will use more memory, but be faster
|
||||
parallel:
|
||||
If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.
|
||||
If 0, use all available cores.
|
||||
If None, don't use data-parallel processing, use default onnxruntime threading instead.
|
||||
|
||||
Returns:
|
||||
List of embeddings, one per document
|
||||
"""
|
||||
yield from self._embed_documents(
|
||||
model_name=self.model_name,
|
||||
cache_dir=str(self.cache_dir),
|
||||
documents=documents,
|
||||
batch_size=batch_size,
|
||||
parallel=parallel,
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_ids=self.device_ids,
|
||||
k=self.k,
|
||||
b=self.b,
|
||||
avg_len=self.avg_len,
|
||||
is_query=False,
|
||||
local_files_only=self._local_files_only,
|
||||
specific_model_path=self._specific_model_path,
|
||||
extra_session_options=self._extra_session_options,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[SparseEmbedding]:
|
||||
"""
|
||||
Encode a list of queries into list of embeddings.
|
||||
"""
|
||||
yield from self._embed_documents(
|
||||
model_name=self.model_name,
|
||||
cache_dir=str(self.cache_dir),
|
||||
documents=query,
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_ids=self.device_ids,
|
||||
k=self.k,
|
||||
b=self.b,
|
||||
avg_len=self.avg_len,
|
||||
is_query=True,
|
||||
local_files_only=self._local_files_only,
|
||||
specific_model_path=self._specific_model_path,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _load_stopwords(cls, model_dir: Path) -> list[str]:
|
||||
stopwords_path = model_dir / STOPWORDS_FILE
|
||||
if not stopwords_path.exists():
|
||||
return []
|
||||
|
||||
with open(stopwords_path, "r") as f:
|
||||
return f.read().splitlines()
|
||||
|
||||
@classmethod
|
||||
def _list_supported_models(cls) -> list[SparseModelDescription]:
|
||||
"""Lists the supported models.
|
||||
|
||||
Returns:
|
||||
list[SparseModelDescription]: A list of SparseModelDescription objects containing the model information.
|
||||
"""
|
||||
return supported_minicoil_models
|
||||
|
||||
def _post_process_onnx_output(
|
||||
self, output: OnnxOutputContext, is_query: bool = False, **kwargs: Any
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
if output.input_ids is None:
|
||||
raise ValueError("input_ids must be provided for document post-processing")
|
||||
|
||||
assert self.vocab_resolver is not None
|
||||
assert self.encoder is not None
|
||||
assert self.sparse_vector_converter is not None
|
||||
|
||||
# Size: (batch_size, sequence_length, hidden_size)
|
||||
embeddings = output.model_output
|
||||
# Size: (batch_size, sequence_length)
|
||||
assert output.attention_mask is not None
|
||||
masks = output.attention_mask
|
||||
|
||||
vocab_size = self.vocab_resolver.vocab_size()
|
||||
embedding_size = self.encoder.output_dim
|
||||
|
||||
# For each document we only select those embeddings that are not masked out
|
||||
|
||||
for i in range(embeddings.shape[0]):
|
||||
# Size: (sequence_length, hidden_size)
|
||||
token_embeddings = embeddings[i, masks[i] == 1]
|
||||
|
||||
# Size: (sequence_length)
|
||||
token_ids: NDArray[np.int64] = output.input_ids[i, masks[i] == 1]
|
||||
|
||||
word_ids_array, counts, oov, forms = self.vocab_resolver.resolve_tokens(token_ids)
|
||||
|
||||
# Size: (1, words)
|
||||
word_ids_array_expanded: NDArray[np.int64] = np.expand_dims(word_ids_array, axis=0)
|
||||
|
||||
# Size: (1, words, embedding_size)
|
||||
token_embeddings_array: NDArray[np.float32] = np.expand_dims(token_embeddings, axis=0)
|
||||
|
||||
assert word_ids_array_expanded.shape[1] == token_embeddings_array.shape[1]
|
||||
|
||||
# Size of word_ids_mapping: (unique_words, 2) - [vocab_id, batch_id]
|
||||
# Size of embeddings: (unique_words, embedding_size)
|
||||
ids_mapping, minicoil_embeddings = self.encoder.forward(
|
||||
word_ids_array_expanded, token_embeddings_array
|
||||
)
|
||||
|
||||
# Size of counts: (unique_words)
|
||||
words_ids: list[int] = ids_mapping[:, 0].tolist() # type: ignore[assignment]
|
||||
|
||||
sentence_result: dict[str, WordEmbedding] = {}
|
||||
|
||||
words = [self.vocab_resolver.lookup_word(word_id) for word_id in words_ids]
|
||||
|
||||
for word, word_id, emb in zip(words, words_ids, minicoil_embeddings.tolist()): # type: ignore[arg-type]
|
||||
if word_id == 0:
|
||||
continue
|
||||
|
||||
sentence_result[word] = WordEmbedding(
|
||||
word=word,
|
||||
forms=forms[word],
|
||||
count=int(counts[word_id]),
|
||||
word_id=int(word_id),
|
||||
embedding=emb, # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
for oov_word, count in oov.items():
|
||||
# {
|
||||
# "word": oov_word,
|
||||
# "forms": [oov_word],
|
||||
# "count": int(count),
|
||||
# "word_id": -1,
|
||||
# "embedding": [1]
|
||||
# }
|
||||
sentence_result[oov_word] = WordEmbedding(
|
||||
word=oov_word, forms=[oov_word], count=int(count), word_id=-1, embedding=[1]
|
||||
)
|
||||
|
||||
if not is_query:
|
||||
yield self.sparse_vector_converter.embedding_to_vector(
|
||||
sentence_result, vocab_size=vocab_size, embedding_size=embedding_size
|
||||
)
|
||||
else:
|
||||
yield self.sparse_vector_converter.embedding_to_vector_query(
|
||||
sentence_result, vocab_size=vocab_size, embedding_size=embedding_size
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type["MiniCoilTextEmbeddingWorker"]:
|
||||
return MiniCoilTextEmbeddingWorker
|
||||
|
||||
|
||||
class MiniCoilTextEmbeddingWorker(TextEmbeddingWorker[SparseEmbedding]):
|
||||
def init_embedding(self, model_name: str, cache_dir: str, **kwargs: Any) -> MiniCOIL:
|
||||
return MiniCOIL(
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
threads=1,
|
||||
**kwargs,
|
||||
)
|
||||
@@ -1,5 +1,5 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Iterable, Optional, Union, Any
|
||||
from typing import Iterable, Any
|
||||
|
||||
import numpy as np
|
||||
from numpy.typing import NDArray
|
||||
@@ -12,7 +12,7 @@ from fastembed.common.model_management import ModelManagement
|
||||
@dataclass
|
||||
class SparseEmbedding:
|
||||
values: NumpyArray
|
||||
indices: Union[NDArray[np.int64], NDArray[np.int32]]
|
||||
indices: NDArray[np.int64] | NDArray[np.int32]
|
||||
|
||||
def as_object(self) -> dict[str, NumpyArray]:
|
||||
return {
|
||||
@@ -35,8 +35,8 @@ class SparseTextEmbeddingBase(ModelManagement[SparseModelDescription]):
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
cache_dir: str | None = None,
|
||||
threads: int | None = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
self.model_name = model_name
|
||||
@@ -46,9 +46,9 @@ class SparseTextEmbeddingBase(ModelManagement[SparseModelDescription]):
|
||||
|
||||
def embed(
|
||||
self,
|
||||
documents: Union[str, Iterable[str]],
|
||||
documents: str | Iterable[str],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
parallel: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
raise NotImplementedError()
|
||||
@@ -68,9 +68,7 @@ class SparseTextEmbeddingBase(ModelManagement[SparseModelDescription]):
|
||||
# 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: Any
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[SparseEmbedding]:
|
||||
"""
|
||||
Embeds queries
|
||||
|
||||
@@ -86,3 +84,7 @@ class SparseTextEmbeddingBase(ModelManagement[SparseModelDescription]):
|
||||
yield from self.embed([query], **kwargs)
|
||||
else:
|
||||
yield from self.embed(query, **kwargs)
|
||||
|
||||
def token_count(self, texts: str | Iterable[str], **kwargs: Any) -> int:
|
||||
"""Returns the number of tokens in the texts."""
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
@@ -1,9 +1,12 @@
|
||||
from typing import Any, Iterable, Optional, Sequence, Type, Union
|
||||
from typing import Any, Iterable, Sequence, Type
|
||||
from dataclasses import asdict
|
||||
|
||||
from fastembed.common import OnnxProvider
|
||||
from fastembed.common.types import Device
|
||||
from fastembed.sparse.bm25 import Bm25
|
||||
from fastembed.sparse.bm42 import Bm42
|
||||
from fastembed.sparse.if_splade import IfSplade
|
||||
from fastembed.sparse.minicoil import MiniCOIL
|
||||
from fastembed.sparse.sparse_embedding_base import (
|
||||
SparseEmbedding,
|
||||
SparseTextEmbeddingBase,
|
||||
@@ -14,7 +17,13 @@ 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,
|
||||
MiniCOIL,
|
||||
IfSplade,
|
||||
]
|
||||
|
||||
@classmethod
|
||||
def list_supported_models(cls) -> list[dict[str, Any]]:
|
||||
@@ -52,16 +61,16 @@ class SparseTextEmbedding(SparseTextEmbeddingBase):
|
||||
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,
|
||||
cache_dir: str | None = None,
|
||||
threads: int | None = None,
|
||||
providers: Sequence[OnnxProvider] | None = None,
|
||||
cuda: bool | Device = Device.AUTO,
|
||||
device_ids: list[int] | None = None,
|
||||
lazy_load: bool = False,
|
||||
**kwargs: Any,
|
||||
):
|
||||
super().__init__(model_name, cache_dir, threads, **kwargs)
|
||||
if model_name == "prithvida/Splade_PP_en_v1":
|
||||
if model_name.lower() == "prithvida/Splade_PP_en_v1".lower():
|
||||
warnings.warn(
|
||||
"The right spelling is prithivida/Splade_PP_en_v1. "
|
||||
"Support of this name will be removed soon, please fix the model_name",
|
||||
@@ -92,9 +101,9 @@ class SparseTextEmbedding(SparseTextEmbeddingBase):
|
||||
|
||||
def embed(
|
||||
self,
|
||||
documents: Union[str, Iterable[str]],
|
||||
documents: str | Iterable[str],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
parallel: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
"""
|
||||
@@ -114,9 +123,7 @@ class SparseTextEmbedding(SparseTextEmbeddingBase):
|
||||
"""
|
||||
yield from self.model.embed(documents, batch_size, parallel, **kwargs)
|
||||
|
||||
def query_embed(
|
||||
self, query: Union[str, Iterable[str]], **kwargs: Any
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[SparseEmbedding]:
|
||||
"""
|
||||
Embeds queries
|
||||
|
||||
@@ -127,3 +134,17 @@ class SparseTextEmbedding(SparseTextEmbeddingBase):
|
||||
Iterable[SparseEmbedding]: The sparse embeddings.
|
||||
"""
|
||||
yield from self.model.query_embed(query, **kwargs)
|
||||
|
||||
def token_count(
|
||||
self, texts: str | Iterable[str], batch_size: int = 1024, **kwargs: Any
|
||||
) -> int:
|
||||
"""Returns the number of tokens in the texts.
|
||||
|
||||
Args:
|
||||
texts (str | Iterable[str]): The list of texts to embed.
|
||||
batch_size (int): Batch size for encoding
|
||||
|
||||
Returns:
|
||||
int: Sum of number of tokens in the texts.
|
||||
"""
|
||||
return self.model.token_count(texts, batch_size=batch_size, **kwargs)
|
||||
|
||||
@@ -1,8 +1,9 @@
|
||||
from typing import Any, Iterable, Optional, Sequence, Type, Union
|
||||
from typing import Any, Iterable, Sequence, Type
|
||||
|
||||
import numpy as np
|
||||
from fastembed.common import OnnxProvider
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.common.types import Device
|
||||
from fastembed.common.utils import define_cache_dir
|
||||
from fastembed.sparse.sparse_embedding_base import (
|
||||
SparseEmbedding,
|
||||
@@ -18,7 +19,7 @@ supported_splade_models: list[SparseModelDescription] = [
|
||||
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"),
|
||||
sources=ModelSource(hf="Qdrant/Splade_PP_en_v1"),
|
||||
model_file="model.onnx",
|
||||
),
|
||||
SparseModelDescription(
|
||||
@@ -27,14 +28,16 @@ supported_splade_models: list[SparseModelDescription] = [
|
||||
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"),
|
||||
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]:
|
||||
def _post_process_onnx_output(
|
||||
self, output: OnnxOutputContext, **kwargs: Any
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
if output.attention_mask is None:
|
||||
raise ValueError("attention_mask must be provided for document post-processing")
|
||||
|
||||
@@ -51,6 +54,11 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
scores = row_scores[indices]
|
||||
yield SparseEmbedding(values=scores, indices=indices)
|
||||
|
||||
def token_count(
|
||||
self, texts: str | Iterable[str], batch_size: int = 1024, **kwargs: Any
|
||||
) -> int:
|
||||
return self._token_count(texts, batch_size=batch_size, **kwargs)
|
||||
|
||||
@classmethod
|
||||
def _list_supported_models(cls) -> list[SparseModelDescription]:
|
||||
"""Lists the supported models.
|
||||
@@ -63,14 +71,14 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
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,
|
||||
cache_dir: str | None = None,
|
||||
threads: int | None = None,
|
||||
providers: Sequence[OnnxProvider] | None = None,
|
||||
cuda: bool | Device = Device.AUTO,
|
||||
device_ids: list[int] | None = None,
|
||||
lazy_load: bool = False,
|
||||
device_id: Optional[int] = None,
|
||||
specific_model_path: Optional[str] = None,
|
||||
device_id: int | None = None,
|
||||
specific_model_path: str | None = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
"""
|
||||
@@ -82,10 +90,11 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
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.
|
||||
cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`
|
||||
Defaults to Device.
|
||||
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.
|
||||
workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, 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.
|
||||
@@ -97,13 +106,14 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
super().__init__(model_name, cache_dir, threads, **kwargs)
|
||||
self.providers = providers
|
||||
self.lazy_load = lazy_load
|
||||
self._extra_session_options = self._select_exposed_session_options(kwargs)
|
||||
|
||||
# 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
|
||||
self.device_id: int | None = None
|
||||
if device_id is not None:
|
||||
self.device_id = device_id
|
||||
elif self.device_ids is not None:
|
||||
@@ -112,11 +122,12 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
self.model_description = self._get_model_description(model_name)
|
||||
self.cache_dir = str(define_cache_dir(cache_dir))
|
||||
|
||||
self._specific_model_path = specific_model_path
|
||||
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,
|
||||
specific_model_path=self._specific_model_path,
|
||||
)
|
||||
|
||||
if not self.lazy_load:
|
||||
@@ -130,13 +141,14 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_id=self.device_id,
|
||||
extra_session_options=self._extra_session_options,
|
||||
)
|
||||
|
||||
def embed(
|
||||
self,
|
||||
documents: Union[str, Iterable[str]],
|
||||
documents: str | Iterable[str],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
parallel: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[SparseEmbedding]:
|
||||
"""
|
||||
@@ -163,6 +175,9 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_ids=self.device_ids,
|
||||
local_files_only=self._local_files_only,
|
||||
specific_model_path=self._specific_model_path,
|
||||
extra_session_options=self._extra_session_options,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
@@ -0,0 +1,146 @@
|
||||
"""
|
||||
Pure numpy implementation of encoder model for a single word.
|
||||
|
||||
This model is not trainable, and should only be used for inference.
|
||||
"""
|
||||
|
||||
import numpy as np
|
||||
from fastembed.common.types import NumpyArray
|
||||
|
||||
|
||||
class Encoder:
|
||||
"""
|
||||
Encoder(768, 4, 10000)
|
||||
|
||||
Will look like this:
|
||||
|
||||
|
||||
Per-word
|
||||
Encoder Matrix
|
||||
┌─────────────────────┐
|
||||
│ Token Embedding(768)├──────┐ (10k, 768, 4)
|
||||
└─────────────────────┘ │ ┌─────────┐
|
||||
│ │ │
|
||||
┌─────────────────────┐ │ ┌─┴───────┐ │
|
||||
│ │ │ │ │ │
|
||||
└─────────────────────┘ │ ┌─┴───────┐ │ │ ┌─────────┐
|
||||
└────►│ │ │ ├─────►│Tanh │
|
||||
┌─────────────────────┐ │ │ │ │ └─────────┘
|
||||
│ │ │ │ ├─┘
|
||||
└─────────────────────┘ │ ├─┘
|
||||
│ │
|
||||
┌─────────────────────┐ └─────────┘
|
||||
│ │
|
||||
└─────────────────────┘
|
||||
|
||||
Final linear transformation is accompanied by a non-linear activation function: Tanh.
|
||||
|
||||
Tanh is used to ensure that the output is in the range [-1, 1].
|
||||
It would be easier to visually interpret the output of the model, assuming that each dimension
|
||||
would need to encode a type of semantic cluster.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
weights: NumpyArray,
|
||||
):
|
||||
self.weights = weights
|
||||
self.vocab_size, self.input_dim, self.output_dim = weights.shape
|
||||
|
||||
self.encoder_weights: NumpyArray = weights
|
||||
|
||||
# Activation function
|
||||
self.activation = np.tanh
|
||||
|
||||
@staticmethod
|
||||
def convert_vocab_ids(vocab_ids: NumpyArray) -> NumpyArray:
|
||||
"""
|
||||
Convert vocab_ids of shape (batch_size, seq_len) into (batch_size, seq_len, 2)
|
||||
by appending batch_id alongside each vocab_id.
|
||||
"""
|
||||
batch_size, seq_len = vocab_ids.shape
|
||||
batch_ids = np.arange(batch_size, dtype=vocab_ids.dtype).reshape(batch_size, 1)
|
||||
batch_ids = np.repeat(batch_ids, seq_len, axis=1)
|
||||
# Stack vocab_ids and batch_ids along the last dimension
|
||||
combined: NumpyArray = np.stack((vocab_ids, batch_ids), axis=2).astype(np.int32)
|
||||
return combined
|
||||
|
||||
@classmethod
|
||||
def avg_by_vocab_ids(
|
||||
cls, vocab_ids: NumpyArray, embeddings: NumpyArray
|
||||
) -> tuple[NumpyArray, NumpyArray]:
|
||||
"""
|
||||
Takes:
|
||||
vocab_ids: (batch_size, seq_len) int array
|
||||
embeddings: (batch_size, seq_len, input_dim) float array
|
||||
|
||||
Returns:
|
||||
unique_flattened_vocab_ids: (total_unique, 2) array of [vocab_id, batch_id]
|
||||
unique_flattened_embeddings: (total_unique, input_dim) averaged embeddings
|
||||
"""
|
||||
input_dim = embeddings.shape[2]
|
||||
|
||||
# Flatten vocab_ids and embeddings
|
||||
# flattened_vocab_ids: (batch_size*seq_len, 2)
|
||||
flattened_vocab_ids = cls.convert_vocab_ids(vocab_ids).reshape(-1, 2)
|
||||
|
||||
# flattened_embeddings: (batch_size*seq_len, input_dim)
|
||||
flattened_embeddings = embeddings.reshape(-1, input_dim)
|
||||
|
||||
# Find unique (vocab_id, batch_id) pairs
|
||||
unique_flattened_vocab_ids, inverse_indices = np.unique(
|
||||
flattened_vocab_ids, axis=0, return_inverse=True
|
||||
)
|
||||
|
||||
# Prepare arrays to accumulate sums
|
||||
unique_count = unique_flattened_vocab_ids.shape[0]
|
||||
unique_flattened_embeddings = np.zeros((unique_count, input_dim), dtype=np.float32)
|
||||
unique_flattened_count = np.zeros(unique_count, dtype=np.int32)
|
||||
|
||||
# Use np.add.at to accumulate sums based on inverse indices
|
||||
np.add.at(unique_flattened_embeddings, inverse_indices, flattened_embeddings)
|
||||
np.add.at(unique_flattened_count, inverse_indices, 1)
|
||||
|
||||
# Compute averages
|
||||
unique_flattened_embeddings /= unique_flattened_count[:, None]
|
||||
|
||||
return unique_flattened_vocab_ids.astype(np.int32), unique_flattened_embeddings.astype(
|
||||
np.float32
|
||||
)
|
||||
|
||||
def forward(
|
||||
self, vocab_ids: NumpyArray, embeddings: NumpyArray
|
||||
) -> tuple[NumpyArray, NumpyArray]:
|
||||
"""
|
||||
Args:
|
||||
vocab_ids: (batch_size, seq_len) int array
|
||||
embeddings: (batch_size, seq_len, input_dim) float array
|
||||
|
||||
Returns:
|
||||
unique_flattened_vocab_ids_and_batch_ids: (total_unique, 2)
|
||||
unique_flattened_encoded: (total_unique, output_dim)
|
||||
"""
|
||||
# Average embeddings for duplicate vocab_ids
|
||||
unique_flattened_vocab_ids_and_batch_ids, unique_flattened_embeddings = (
|
||||
self.avg_by_vocab_ids(vocab_ids, embeddings)
|
||||
)
|
||||
|
||||
# Select the encoder weights for each unique vocab_id
|
||||
unique_flattened_vocab_ids = unique_flattened_vocab_ids_and_batch_ids[:, 0].astype(
|
||||
np.int32
|
||||
)
|
||||
|
||||
# unique_encoder_weights: (total_unique, input_dim, output_dim)
|
||||
unique_encoder_weights = self.encoder_weights[unique_flattened_vocab_ids]
|
||||
|
||||
# Compute linear transform: (total_unique, output_dim)
|
||||
# Using Einstein summation for matrix multiplication:
|
||||
# 'bi,bio->bo' means: for each "b" (batch element), multiply embeddings (b,i) by weights (b,i,o) -> (b,o)
|
||||
unique_flattened_encoded = np.einsum(
|
||||
"bi,bio->bo", unique_flattened_embeddings, unique_encoder_weights
|
||||
)
|
||||
|
||||
# Apply Tanh activation and ensure float32 type
|
||||
unique_flattened_encoded = self.activation(unique_flattened_encoded).astype(np.float32)
|
||||
|
||||
return unique_flattened_vocab_ids_and_batch_ids.astype(np.int32), unique_flattened_encoded
|
||||
@@ -0,0 +1,244 @@
|
||||
import copy
|
||||
from dataclasses import dataclass
|
||||
|
||||
import mmh3
|
||||
import numpy as np
|
||||
from py_rust_stemmers import SnowballStemmer
|
||||
|
||||
from fastembed.common.utils import get_all_punctuation, remove_non_alphanumeric
|
||||
from fastembed.sparse.sparse_embedding_base import SparseEmbedding
|
||||
|
||||
GAP = 32000
|
||||
INT32_MAX = 2**31 - 1
|
||||
|
||||
|
||||
@dataclass
|
||||
class WordEmbedding:
|
||||
word: str
|
||||
forms: list[str]
|
||||
count: int
|
||||
word_id: int
|
||||
embedding: list[float]
|
||||
|
||||
|
||||
class SparseVectorConverter:
|
||||
def __init__(
|
||||
self,
|
||||
stopwords: set[str],
|
||||
stemmer: SnowballStemmer,
|
||||
k: float = 1.2,
|
||||
b: float = 0.75,
|
||||
avg_len: float = 150.0,
|
||||
):
|
||||
punctuation = set(get_all_punctuation())
|
||||
special_tokens = {"[CLS]", "[SEP]", "[PAD]", "[UNK]", "[MASK]"}
|
||||
|
||||
self.stemmer = stemmer
|
||||
self.unwanted_tokens = punctuation | special_tokens | stopwords
|
||||
|
||||
self.k = k
|
||||
self.b = b
|
||||
self.avg_len = avg_len
|
||||
|
||||
@classmethod
|
||||
def unkn_word_token_id(
|
||||
cls, word: str, shift: int
|
||||
) -> int: # 2-3 words can collide in 1 index with this mapping, not considering mm3 collisions
|
||||
token_hash = abs(mmh3.hash(word))
|
||||
|
||||
range_size = INT32_MAX - shift
|
||||
remapped_hash = shift + (token_hash % range_size)
|
||||
|
||||
return remapped_hash
|
||||
|
||||
def bm25_tf(self, num_occurrences: int, sentence_len: int) -> float:
|
||||
res = num_occurrences * (self.k + 1)
|
||||
res /= num_occurrences + self.k * (1 - self.b + self.b * sentence_len / self.avg_len)
|
||||
return res
|
||||
|
||||
@classmethod
|
||||
def normalize_vector(cls, vector: list[float]) -> list[float]:
|
||||
norm = sum([x**2 for x in vector]) ** 0.5
|
||||
if norm < 1e-8:
|
||||
return vector
|
||||
return [x / norm for x in vector]
|
||||
|
||||
def clean_words(
|
||||
self, sentence_embedding: dict[str, WordEmbedding], token_max_length: int = 40
|
||||
) -> dict[str, WordEmbedding]:
|
||||
"""
|
||||
Clean miniCOIL-produced sentence_embedding, as unknown to the miniCOIL's stemmer tokens should fully resemble
|
||||
our BM25 token representation.
|
||||
|
||||
sentence_embedding = {"9°": {"word": "9°", "word_id": -1, "count": 2, "embedding": [1], "forms": ["9°"]},
|
||||
"9": {"word": "9", "word_id": -1, "count": 2, "embedding": [1], "forms": ["9"]},
|
||||
"bat": {"word": "bat", "word_id": 2, "count": 3, "embedding": [0.2, 0.1, -0.2, -0.2], "forms": ["bats", "bat"]},
|
||||
"9°9": {"word": "9°9", "word_id": -1, "count": 1, "embedding": [1], "forms": ["9°9"]},
|
||||
"screech": {"word": "screech", "word_id": -1, "count": 1, "embedding": [1], "forms": ["screech"]},
|
||||
"screeched": {"word": "screeched", "word_id": -1, "count": 1, "embedding": [1], "forms": ["screeched"]}
|
||||
}
|
||||
cleaned_embedding_ground_truth = {
|
||||
"9": {"word": "9", "word_id": -1, "count": 6, "embedding": [1], "forms": ["9°", "9", "9°9", "9°9"]},
|
||||
"bat": {"word": "bat", "word_id": 2, "count": 3, "embedding": [0.2, 0.1, -0.2, -0.2], "forms": ["bats", "bat"]},
|
||||
"screech": {"word": "screech", "word_id": -1, "count": 2, "embedding": [1], "forms": ["screech", "screeched"]}
|
||||
}
|
||||
"""
|
||||
|
||||
new_sentence_embedding: dict[str, WordEmbedding] = {}
|
||||
|
||||
for word, embedding in sentence_embedding.items():
|
||||
# embedding = {
|
||||
# "word": "vector",
|
||||
# "forms": ["vector", "vectors"],
|
||||
# "count": 2,
|
||||
# "word_id": 1231,
|
||||
# "embedding": [0.1, 0.2, 0.3, 0.4]
|
||||
# }
|
||||
if embedding.word_id > 0:
|
||||
# Known word, no need to clean
|
||||
new_sentence_embedding[word] = embedding
|
||||
else:
|
||||
# Unknown word
|
||||
if word in self.unwanted_tokens:
|
||||
continue
|
||||
|
||||
# Example complex word split:
|
||||
# word = `word^vec`
|
||||
word_cleaned = remove_non_alphanumeric(word).strip()
|
||||
# word_cleaned = `word vec`
|
||||
|
||||
if len(word_cleaned) > 0:
|
||||
# Subwords: ['word', 'vec']
|
||||
for subword in word_cleaned.split():
|
||||
stemmed_subword: str = self.stemmer.stem_word(subword)
|
||||
if (
|
||||
len(stemmed_subword) <= token_max_length
|
||||
and stemmed_subword not in self.unwanted_tokens
|
||||
):
|
||||
if stemmed_subword not in new_sentence_embedding:
|
||||
new_sentence_embedding[stemmed_subword] = copy.deepcopy(embedding)
|
||||
new_sentence_embedding[stemmed_subword].word = stemmed_subword
|
||||
else:
|
||||
new_sentence_embedding[stemmed_subword].count += embedding.count
|
||||
new_sentence_embedding[stemmed_subword].forms += embedding.forms
|
||||
|
||||
return new_sentence_embedding
|
||||
|
||||
def embedding_to_vector(
|
||||
self,
|
||||
sentence_embedding: dict[str, WordEmbedding],
|
||||
embedding_size: int,
|
||||
vocab_size: int,
|
||||
) -> SparseEmbedding:
|
||||
"""
|
||||
Convert miniCOIL sentence embedding to Qdrant sparse vector
|
||||
|
||||
Example input:
|
||||
|
||||
```
|
||||
{
|
||||
"vector": WordEmbedding({ // Vocabulary word, encoded with miniCOIL normally
|
||||
"word": "vector",
|
||||
"forms": ["vector", "vectors"],
|
||||
"count": 2,
|
||||
"word_id": 1231,
|
||||
"embedding": [0.1, 0.2, 0.3, 0.4]
|
||||
}),
|
||||
"axiotic": WordEmbedding({ // Out-of-vocabulary word, fallback to BM25
|
||||
"word": "axiotic",
|
||||
"forms": ["axiotics"],
|
||||
"count": 1,
|
||||
"word_id": -1,
|
||||
})
|
||||
}
|
||||
```
|
||||
|
||||
"""
|
||||
|
||||
indices: list[int] = []
|
||||
values: list[float] = []
|
||||
|
||||
# Example:
|
||||
# vocab_size = 10000
|
||||
# embedding_size = 4
|
||||
# GAP = 32000
|
||||
#
|
||||
# We want to start random words section from the bucket, that is guaranteed to not
|
||||
# include any vocab words.
|
||||
# We need (vocab_size * embedding_size) slots for vocab words.
|
||||
# Therefore we need (vocab_size * embedding_size) // GAP + 1 buckets for vocab words.
|
||||
# Therefore, we can start random words from bucket (vocab_size * embedding_size) // GAP + 1 + 1
|
||||
|
||||
# ID at which the scope of OOV words starts
|
||||
unknown_words_shift = ((vocab_size * embedding_size) // GAP + 2) * GAP
|
||||
sentence_embedding_cleaned = self.clean_words(sentence_embedding)
|
||||
|
||||
# Calculate sentence length after cleaning
|
||||
sentence_len = 0
|
||||
for embedding in sentence_embedding_cleaned.values():
|
||||
sentence_len += embedding.count
|
||||
|
||||
for embedding in sentence_embedding_cleaned.values():
|
||||
word_id = embedding.word_id
|
||||
num_occurrences = embedding.count
|
||||
tf = self.bm25_tf(num_occurrences, sentence_len)
|
||||
if (
|
||||
word_id > 0
|
||||
): # miniCOIL starts with ID 1, we generally won't have word_id == 0 (UNK), as we don't add
|
||||
# these words to sentence_embedding
|
||||
embedding_values = embedding.embedding
|
||||
normalized_embedding = self.normalize_vector(embedding_values)
|
||||
|
||||
for val_id, value in enumerate(normalized_embedding):
|
||||
indices.append(
|
||||
word_id * embedding_size + val_id
|
||||
) # since miniCOIL IDs start with 1
|
||||
values.append(value * tf)
|
||||
else:
|
||||
indices.append(self.unkn_word_token_id(embedding.word, unknown_words_shift))
|
||||
values.append(tf)
|
||||
|
||||
return SparseEmbedding(
|
||||
indices=np.array(indices, dtype=np.int32),
|
||||
values=np.array(values, dtype=np.float32),
|
||||
)
|
||||
|
||||
def embedding_to_vector_query(
|
||||
self,
|
||||
sentence_embedding: dict[str, WordEmbedding],
|
||||
embedding_size: int,
|
||||
vocab_size: int,
|
||||
) -> SparseEmbedding:
|
||||
"""
|
||||
Same as `embedding_to_vector`, but no TF
|
||||
"""
|
||||
|
||||
indices: list[int] = []
|
||||
values: list[float] = []
|
||||
|
||||
# ID at which the scope of OOV words starts
|
||||
unknown_words_shift = ((vocab_size * embedding_size) // GAP + 2) * GAP
|
||||
|
||||
sentence_embedding_cleaned = self.clean_words(sentence_embedding)
|
||||
|
||||
for embedding in sentence_embedding_cleaned.values():
|
||||
word_id = embedding.word_id
|
||||
tf = 1.0
|
||||
|
||||
if word_id >= 0: # miniCOIL starts with ID 1
|
||||
embedding_values = embedding.embedding
|
||||
normalized_embedding = self.normalize_vector(embedding_values)
|
||||
|
||||
for val_id, value in enumerate(normalized_embedding):
|
||||
indices.append(
|
||||
word_id * embedding_size + val_id
|
||||
) # since miniCOIL IDs start with 1
|
||||
values.append(value * tf)
|
||||
else:
|
||||
indices.append(self.unkn_word_token_id(embedding.word, unknown_words_shift))
|
||||
values.append(tf)
|
||||
|
||||
return SparseEmbedding(
|
||||
indices=np.array(indices, dtype=np.int32),
|
||||
values=np.array(values, dtype=np.float32),
|
||||
)
|
||||
@@ -0,0 +1,202 @@
|
||||
from collections import defaultdict
|
||||
from typing import Iterable
|
||||
|
||||
from py_rust_stemmers import SnowballStemmer
|
||||
import numpy as np
|
||||
from tokenizers import Tokenizer
|
||||
from numpy.typing import NDArray
|
||||
|
||||
from fastembed.common.types import NumpyArray
|
||||
|
||||
|
||||
class VocabTokenizerBase:
|
||||
def tokenize(self, sentence: str) -> NumpyArray:
|
||||
raise NotImplementedError()
|
||||
|
||||
def convert_ids_to_tokens(self, token_ids: NumpyArray) -> list[str]:
|
||||
raise NotImplementedError()
|
||||
|
||||
|
||||
class VocabTokenizer(VocabTokenizerBase):
|
||||
def __init__(self, tokenizer: Tokenizer):
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
def tokenize(self, sentence: str) -> NumpyArray:
|
||||
return np.array(self.tokenizer.encode(sentence).ids)
|
||||
|
||||
def convert_ids_to_tokens(self, token_ids: NumpyArray) -> list[str]:
|
||||
return [self.tokenizer.id_to_token(token_id) for token_id in token_ids]
|
||||
|
||||
|
||||
class VocabResolver:
|
||||
def __init__(self, tokenizer: VocabTokenizerBase, stopwords: set[str], stemmer: SnowballStemmer):
|
||||
# Word to id mapping
|
||||
self.vocab: dict[str, int] = {}
|
||||
# Id to word mapping
|
||||
self.words: list[str] = []
|
||||
# Lemma to word mapping
|
||||
self.stem_mapping: dict[str, str] = {}
|
||||
self.tokenizer: VocabTokenizerBase = tokenizer
|
||||
self.stemmer = stemmer
|
||||
self.stopwords: set[str] = stopwords
|
||||
|
||||
def tokenize(self, sentence: str) -> NumpyArray:
|
||||
return self.tokenizer.tokenize(sentence)
|
||||
|
||||
def lookup_word(self, word_id: int) -> str:
|
||||
if word_id == 0:
|
||||
return "UNK"
|
||||
return self.words[word_id - 1]
|
||||
|
||||
def convert_ids_to_tokens(self, token_ids: NumpyArray) -> list[str]:
|
||||
return self.tokenizer.convert_ids_to_tokens(token_ids)
|
||||
|
||||
def vocab_size(self) -> int:
|
||||
# We need +1 for UNK token
|
||||
return len(self.vocab) + 1
|
||||
|
||||
def save_vocab(self, path: str) -> None:
|
||||
with open(path, "w") as f:
|
||||
for word in self.words:
|
||||
f.write(word + "\n")
|
||||
|
||||
def save_json_vocab(self, path: str) -> None:
|
||||
import json
|
||||
|
||||
with open(path, "w") as f:
|
||||
json.dump({"vocab": self.words, "stem_mapping": self.stem_mapping}, f, indent=2)
|
||||
|
||||
def load_json_vocab(self, path: str) -> None:
|
||||
import json
|
||||
|
||||
with open(path, "r") as f:
|
||||
data = json.load(f)
|
||||
self.words = data["vocab"]
|
||||
self.vocab = {word: idx + 1 for idx, word in enumerate(self.words)}
|
||||
self.stem_mapping = data["stem_mapping"]
|
||||
|
||||
def add_word(self, word: str) -> None:
|
||||
if word not in self.vocab:
|
||||
self.vocab[word] = len(self.vocab) + 1
|
||||
self.words.append(word)
|
||||
stem = self.stemmer.stem_word(word)
|
||||
if stem not in self.stem_mapping:
|
||||
self.stem_mapping[stem] = word
|
||||
else:
|
||||
existing_word = self.stem_mapping[stem]
|
||||
if len(existing_word) > len(word):
|
||||
# Prefer shorter words for the same stem
|
||||
# Example: "swim" is preferred over "swimming"
|
||||
self.stem_mapping[stem] = word
|
||||
|
||||
def load_vocab(self, path: str) -> None:
|
||||
with open(path, "r") as f:
|
||||
for line in f:
|
||||
self.add_word(line.strip())
|
||||
|
||||
@classmethod
|
||||
def _reconstruct_bpe(
|
||||
cls, 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 = "##"
|
||||
continuing_subword_prefix_len = len(continuing_subword_prefix)
|
||||
|
||||
for idx, token in bpe_tokens:
|
||||
if token.startswith(continuing_subword_prefix):
|
||||
acc += token[continuing_subword_prefix_len:]
|
||||
acc_idx.append(idx)
|
||||
else:
|
||||
if acc:
|
||||
result.append((acc, acc_idx))
|
||||
acc_idx = []
|
||||
acc = token
|
||||
acc_idx.append(idx)
|
||||
|
||||
if acc:
|
||||
result.append((acc, acc_idx))
|
||||
return result
|
||||
|
||||
def resolve_tokens(
|
||||
self, token_ids: NDArray[np.int64]
|
||||
) -> tuple[NDArray[np.int64], dict[int, int], dict[str, int], dict[str, list[str]]]:
|
||||
"""
|
||||
Mark known tokens (including composed tokens) with vocab ids.
|
||||
|
||||
Args:
|
||||
token_ids: (seq_len) - list of ids of tokens
|
||||
Example:
|
||||
[
|
||||
101, 3897, 19332, 12718, 23348,
|
||||
1010, 1996, 7151, 2296, 4845,
|
||||
2359, 2005, 4234, 1010, 4332,
|
||||
2871, 3191, 2062, 102
|
||||
]
|
||||
|
||||
returns:
|
||||
- token_ids with vocab ids
|
||||
[
|
||||
0, 151, 151, 0, 0,
|
||||
912, 0, 0, 0, 332,
|
||||
332, 332, 0, 7121, 191,
|
||||
0, 0, 332, 0
|
||||
]
|
||||
- counts of each token
|
||||
{
|
||||
151: 1,
|
||||
332: 3,
|
||||
7121: 1,
|
||||
191: 1,
|
||||
912: 1
|
||||
}
|
||||
- oov counts of each token
|
||||
{
|
||||
"the": 1,
|
||||
"a": 1,
|
||||
"[CLS]": 1,
|
||||
"[SEP]": 1,
|
||||
...
|
||||
}
|
||||
- forms of each token
|
||||
{
|
||||
"hello": ["hello"],
|
||||
"world": ["worlds", "world", "worlding"],
|
||||
}
|
||||
|
||||
"""
|
||||
tokens = self.convert_ids_to_tokens(token_ids)
|
||||
tokens_mapping = self._reconstruct_bpe(enumerate(tokens))
|
||||
|
||||
counts: dict[int, int] = defaultdict(int)
|
||||
oov_count: dict[str, int] = defaultdict(int)
|
||||
|
||||
forms: dict[str, list[str]] = defaultdict(list)
|
||||
|
||||
for token, mapped_token_ids in tokens_mapping:
|
||||
vocab_id = 0
|
||||
if token in self.stopwords:
|
||||
vocab_id = 0
|
||||
elif token in self.vocab:
|
||||
vocab_id = self.vocab[token]
|
||||
forms[token].append(token)
|
||||
elif token in self.stem_mapping:
|
||||
vocab_id = self.vocab[self.stem_mapping[token]]
|
||||
forms[self.stem_mapping[token]].append(token)
|
||||
else:
|
||||
stem = self.stemmer.stem_word(token)
|
||||
if stem in self.stem_mapping:
|
||||
vocab_id = self.vocab[self.stem_mapping[stem]]
|
||||
forms[self.stem_mapping[stem]].append(token)
|
||||
|
||||
for token_id in mapped_token_ids:
|
||||
token_ids[token_id] = vocab_id
|
||||
|
||||
if vocab_id == 0:
|
||||
oov_count[token] += 1
|
||||
else:
|
||||
counts[vocab_id] += 1
|
||||
return token_ids, counts, oov_count, forms
|
||||
|
||||
@@ -0,0 +1,69 @@
|
||||
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.common.model_description import DenseModelDescription, ModelSource
|
||||
|
||||
|
||||
supported_builtin_sentence_embedding_models: list[DenseModelDescription] = [
|
||||
DenseModelDescription(
|
||||
model="google/embeddinggemma-300m",
|
||||
dim=768,
|
||||
description=(
|
||||
"Text embeddings, Unimodal (text), multilingual, 2048 input tokens truncation, "
|
||||
"Prefixes for queries/documents: `task: search result | query: {content}` for query, "
|
||||
"`title: {title | 'none'} | text: {content}` for documents, 2025 year."
|
||||
),
|
||||
license="gemma",
|
||||
size_in_GB=1.24,
|
||||
sources=ModelSource(
|
||||
hf="onnx-community/embeddinggemma-300m-ONNX",
|
||||
),
|
||||
model_file="onnx/model.onnx",
|
||||
additional_files=["onnx/model.onnx_data"],
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
class BuiltinSentenceEmbedding(OnnxTextEmbedding):
|
||||
"""Builtin Sentence Embedding uses built-in pooling and normalization of underlying onnx models"""
|
||||
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type[OnnxTextEmbeddingWorker]:
|
||||
return BuiltinSentenceEmbeddingWorker
|
||||
|
||||
@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_builtin_sentence_embedding_models
|
||||
|
||||
def _post_process_onnx_output(
|
||||
self, output: OnnxOutputContext, **kwargs: Any
|
||||
) -> Iterable[NumpyArray]:
|
||||
return output.model_output
|
||||
|
||||
def _run_model(
|
||||
self, onnx_input: dict[str, Any], onnx_output_names: list[str] | None = None
|
||||
) -> NumpyArray:
|
||||
return self.model.run(onnx_output_names, onnx_input)[1] # type: ignore[union-attr]
|
||||
|
||||
|
||||
class BuiltinSentenceEmbeddingWorker(OnnxTextEmbeddingWorker):
|
||||
def init_embedding(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
**kwargs: Any,
|
||||
) -> OnnxTextEmbedding:
|
||||
return BuiltinSentenceEmbedding(
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
threads=1,
|
||||
**kwargs,
|
||||
)
|
||||
@@ -35,7 +35,9 @@ class CLIPOnnxEmbedding(OnnxTextEmbedding):
|
||||
"""
|
||||
return supported_clip_models
|
||||
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[NumpyArray]:
|
||||
def _post_process_onnx_output(
|
||||
self, output: OnnxOutputContext, **kwargs: Any
|
||||
) -> Iterable[NumpyArray]:
|
||||
return output.model_output
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,144 @@
|
||||
from typing import Sequence, Any, Iterable, Type
|
||||
from dataclasses import dataclass
|
||||
|
||||
import numpy as np
|
||||
from numpy.typing import NDArray
|
||||
|
||||
from fastembed.common import OnnxProvider
|
||||
from fastembed.common.model_description import (
|
||||
PoolingType,
|
||||
DenseModelDescription,
|
||||
)
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.common.types import NumpyArray, Device
|
||||
from fastembed.common.utils import normalize, mean_pooling, last_token_pooling
|
||||
from fastembed.text.onnx_embedding import OnnxTextEmbedding
|
||||
from fastembed.text.onnx_text_model import TextEmbeddingWorker
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class PostprocessingConfig:
|
||||
pooling: PoolingType
|
||||
normalization: bool
|
||||
|
||||
|
||||
class CustomTextEmbedding(OnnxTextEmbedding):
|
||||
SUPPORTED_MODELS: list[DenseModelDescription] = []
|
||||
POSTPROCESSING_MAPPING: dict[str, PostprocessingConfig] = {}
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str | None = None,
|
||||
threads: int | None = None,
|
||||
providers: Sequence[OnnxProvider] | None = None,
|
||||
cuda: bool | Device = Device.AUTO,
|
||||
device_ids: list[int] | None = None,
|
||||
lazy_load: bool = False,
|
||||
device_id: int | None = None,
|
||||
specific_model_path: str | None = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
super().__init__(
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
threads=threads,
|
||||
providers=providers,
|
||||
cuda=cuda,
|
||||
device_ids=device_ids,
|
||||
lazy_load=lazy_load,
|
||||
device_id=device_id,
|
||||
specific_model_path=specific_model_path,
|
||||
**kwargs,
|
||||
)
|
||||
postprocessing_config = self.POSTPROCESSING_MAPPING[self.model_description.model]
|
||||
self._pooling = postprocessing_config.pooling
|
||||
self._normalization = postprocessing_config.normalization
|
||||
|
||||
@classmethod
|
||||
def _list_supported_models(cls) -> list[DenseModelDescription]:
|
||||
return cls.SUPPORTED_MODELS
|
||||
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type["TextEmbeddingWorker[NumpyArray]"]:
|
||||
return CustomTextEmbeddingWorker
|
||||
|
||||
def _get_worker_init_kwargs(self) -> dict[str, Any]:
|
||||
return {
|
||||
"model_description": self.model_description,
|
||||
"postprocessing_config": self.POSTPROCESSING_MAPPING[self.model_description.model],
|
||||
}
|
||||
|
||||
def _post_process_onnx_output(
|
||||
self, output: OnnxOutputContext, **kwargs: Any
|
||||
) -> Iterable[NumpyArray]:
|
||||
return self._normalize(self._pool(output.model_output, output.attention_mask))
|
||||
|
||||
def _pool(
|
||||
self, embeddings: NumpyArray, attention_mask: NDArray[np.int64] | None = None
|
||||
) -> NumpyArray:
|
||||
if self._pooling == PoolingType.CLS:
|
||||
return embeddings[:, 0]
|
||||
|
||||
if self._pooling == PoolingType.MEAN:
|
||||
if attention_mask is None:
|
||||
raise ValueError("attention_mask must be provided for mean pooling")
|
||||
return mean_pooling(embeddings, attention_mask)
|
||||
|
||||
if self._pooling == PoolingType.LAST_TOKEN:
|
||||
if attention_mask is None:
|
||||
raise ValueError("attention_mask must be provided for last token pooling")
|
||||
return last_token_pooling(embeddings, attention_mask)
|
||||
|
||||
if self._pooling == PoolingType.DISABLED:
|
||||
return embeddings
|
||||
|
||||
raise ValueError(
|
||||
f"Unsupported pooling type {self._pooling}. "
|
||||
f"Supported types are: {PoolingType.CLS}, {PoolingType.MEAN}, "
|
||||
f"{PoolingType.LAST_TOKEN}, {PoolingType.DISABLED}."
|
||||
)
|
||||
|
||||
def _normalize(self, embeddings: NumpyArray) -> NumpyArray:
|
||||
return normalize(embeddings) if self._normalization else embeddings
|
||||
|
||||
@classmethod
|
||||
def add_model(
|
||||
cls,
|
||||
model_description: DenseModelDescription,
|
||||
pooling: PoolingType,
|
||||
normalization: bool,
|
||||
) -> None:
|
||||
cls.SUPPORTED_MODELS.append(model_description)
|
||||
cls.POSTPROCESSING_MAPPING[model_description.model] = PostprocessingConfig(
|
||||
pooling=pooling, normalization=normalization
|
||||
)
|
||||
|
||||
|
||||
class CustomTextEmbeddingWorker(TextEmbeddingWorker[NumpyArray]):
|
||||
def init_embedding(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
model_description: DenseModelDescription | None = None,
|
||||
postprocessing_config: PostprocessingConfig | None = None,
|
||||
**kwargs: Any,
|
||||
) -> CustomTextEmbedding:
|
||||
if model_description is None or postprocessing_config is None:
|
||||
raise ValueError(
|
||||
"`model_description` and `postprocessing_config` are required to initialize a "
|
||||
"custom model in a worker process, they are provided by "
|
||||
"`CustomTextEmbedding._get_worker_init_kwargs`"
|
||||
)
|
||||
# custom models live in a class-level registry, which spawned workers don't inherit
|
||||
CustomTextEmbedding.add_model(
|
||||
model_description,
|
||||
pooling=postprocessing_config.pooling,
|
||||
normalization=postprocessing_config.normalization,
|
||||
)
|
||||
return CustomTextEmbedding(
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
threads=1,
|
||||
**kwargs,
|
||||
)
|
||||
@@ -0,0 +1,94 @@
|
||||
from typing import Any, Iterable, Type
|
||||
|
||||
import onnxruntime as ort
|
||||
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.common.utils import last_token_pooling, normalize
|
||||
from fastembed.text.onnx_embedding import OnnxTextEmbedding, OnnxTextEmbeddingWorker
|
||||
from fastembed.common.model_description import DenseModelDescription, ModelSource
|
||||
|
||||
supported_last_token_normalized_models: list[DenseModelDescription] = [
|
||||
DenseModelDescription(
|
||||
model="Qwen/Qwen3-Embedding-0.6B",
|
||||
dim=1024,
|
||||
description=(
|
||||
"Text embeddings, Unimodal (text), multilingual, 32768 input tokens truncation, "
|
||||
"Prefixes for queries/documents: `Instruct: {task_description}\\nQuery:{query}` "
|
||||
"for queries, none for documents, 2025 year."
|
||||
),
|
||||
license="apache-2.0",
|
||||
size_in_GB=2.38,
|
||||
sources=ModelSource(hf="Qdrant/Qwen3-Embedding-0.6B-onnx"),
|
||||
model_file="onnx/model.onnx",
|
||||
additional_files=["onnx/model.onnx_data"],
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="Qwen/Qwen3-Embedding-0.6B-Q",
|
||||
dim=1024,
|
||||
description=(
|
||||
"Text embeddings, Unimodal (text), multilingual, 32768 input tokens truncation, "
|
||||
"Prefixes for queries/documents: `Instruct: {task_description}\\nQuery:{query}` "
|
||||
"for queries, none for documents, int8 weights, requires onnxruntime>=1.23, "
|
||||
"2025 year."
|
||||
),
|
||||
license="apache-2.0",
|
||||
size_in_GB=1.12,
|
||||
sources=ModelSource(hf="Qdrant/Qwen3-Embedding-0.6B-onnx"),
|
||||
model_file="onnx/model_quantized.onnx",
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
class LastTokenNormalizedEmbedding(OnnxTextEmbedding):
|
||||
"""Decoder-based embedding models, which pool the last non-padding token and normalize it"""
|
||||
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type[OnnxTextEmbeddingWorker]:
|
||||
return LastTokenNormalizedEmbeddingWorker
|
||||
|
||||
def load_onnx_model(self) -> None:
|
||||
try:
|
||||
super().load_onnx_model()
|
||||
except Exception as e:
|
||||
# int8 weights are stored as 8-bit MatMulNBits, which onnxruntime only
|
||||
# implements since 1.23; older versions fail with "nbits_ == 4 was false"
|
||||
if "nbits" not in str(e).lower():
|
||||
raise
|
||||
raise RuntimeError(
|
||||
f"Could not load {self.model_name}: its int8 weights require "
|
||||
f"onnxruntime>=1.23, but onnxruntime {ort.__version__} is installed. "
|
||||
f"Either upgrade onnxruntime or use a non-quantized model."
|
||||
) from e
|
||||
|
||||
@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_last_token_normalized_models
|
||||
|
||||
def _post_process_onnx_output(
|
||||
self, output: OnnxOutputContext, **kwargs: Any
|
||||
) -> Iterable[NumpyArray]:
|
||||
if output.attention_mask is None:
|
||||
raise ValueError("attention_mask must be provided for last token pooling")
|
||||
|
||||
return normalize(last_token_pooling(output.model_output, output.attention_mask))
|
||||
|
||||
|
||||
class LastTokenNormalizedEmbeddingWorker(OnnxTextEmbeddingWorker):
|
||||
def init_embedding(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
**kwargs: Any,
|
||||
) -> OnnxTextEmbedding:
|
||||
return LastTokenNormalizedEmbedding(
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
threads=1,
|
||||
**kwargs,
|
||||
)
|
||||
@@ -1,8 +1,9 @@
|
||||
from enum import Enum
|
||||
from typing import Any, Type, Iterable, Union, Optional
|
||||
from typing import Any, Type, Iterable
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.text.pooled_normalized_embedding import PooledNormalizedEmbedding
|
||||
from fastembed.text.onnx_embedding import OnnxTextEmbeddingWorker
|
||||
@@ -44,9 +45,9 @@ class JinaEmbeddingV3(PooledNormalizedEmbedding):
|
||||
PASSAGE_TASK = Task.RETRIEVAL_PASSAGE
|
||||
QUERY_TASK = Task.RETRIEVAL_QUERY
|
||||
|
||||
def __init__(self, *args: Any, **kwargs: Any):
|
||||
def __init__(self, *args: Any, task_id: int | None = None, **kwargs: Any):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.current_task_id: Union[Task, int] = self.PASSAGE_TASK
|
||||
self.default_task_id: Task | int = task_id if task_id is not None else self.PASSAGE_TASK
|
||||
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type[OnnxTextEmbeddingWorker]:
|
||||
@@ -57,30 +58,34 @@ class JinaEmbeddingV3(PooledNormalizedEmbedding):
|
||||
return supported_multitask_models
|
||||
|
||||
def _preprocess_onnx_input(
|
||||
self, onnx_input: dict[str, NumpyArray], **kwargs: Any
|
||||
self,
|
||||
onnx_input: dict[str, NumpyArray],
|
||||
task_id: int | Task | None = None,
|
||||
**kwargs: Any,
|
||||
) -> dict[str, NumpyArray]:
|
||||
onnx_input["task_id"] = np.array(self.current_task_id, dtype=np.int64)
|
||||
if task_id is None:
|
||||
raise ValueError(f"task_id must be provided for JinaEmbeddingV3, got <{task_id}>")
|
||||
onnx_input["task_id"] = np.array(task_id, dtype=np.int64)
|
||||
return onnx_input
|
||||
|
||||
def embed(
|
||||
self,
|
||||
documents: Union[str, Iterable[str]],
|
||||
documents: str | Iterable[str],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
task_id: int = PASSAGE_TASK,
|
||||
parallel: int | None = None,
|
||||
task_id: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
self.current_task_id = task_id
|
||||
kwargs["task_id"] = task_id
|
||||
yield from super().embed(documents, batch_size, parallel, **kwargs)
|
||||
task_id = (
|
||||
task_id if task_id is not None else self.default_task_id
|
||||
) # required for multiprocessing
|
||||
yield from super().embed(documents, batch_size, parallel, task_id=task_id, **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 query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:
|
||||
yield from super().embed(query, task_id=self.QUERY_TASK, **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)
|
||||
yield from super().embed(texts, task_id=self.PASSAGE_TASK, **kwargs)
|
||||
|
||||
|
||||
class JinaEmbeddingV3Worker(OnnxTextEmbeddingWorker):
|
||||
@@ -90,11 +95,15 @@ class JinaEmbeddingV3Worker(OnnxTextEmbeddingWorker):
|
||||
cache_dir: str,
|
||||
**kwargs: Any,
|
||||
) -> JinaEmbeddingV3:
|
||||
model = JinaEmbeddingV3(
|
||||
return JinaEmbeddingV3(
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
threads=1,
|
||||
**kwargs,
|
||||
)
|
||||
model.current_task_id = kwargs["task_id"]
|
||||
return model
|
||||
|
||||
def process(self, items: Iterable[tuple[int, Any]]) -> Iterable[tuple[int, OnnxOutputContext]]:
|
||||
self.model: JinaEmbeddingV3 # mypy complaints `self.model` does not have `default_task_id`
|
||||
for idx, batch in items:
|
||||
onnx_output = self.model.onnx_embed(batch, task_id=self.model.default_task_id)
|
||||
yield idx, onnx_output
|
||||
|
||||
@@ -1,12 +1,11 @@
|
||||
from typing import Any, Iterable, Optional, Sequence, Type, Union
|
||||
from typing import Any, Iterable, Sequence, Type
|
||||
|
||||
import numpy as np
|
||||
from fastembed.common.types import NumpyArray, OnnxProvider
|
||||
from fastembed.common.model_description import DenseModelDescription, ModelSource
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.common.types import Device, NumpyArray, OnnxProvider
|
||||
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: list[DenseModelDescription] = [
|
||||
DenseModelDescription(
|
||||
@@ -21,6 +20,7 @@ supported_onnx_models: list[DenseModelDescription] = [
|
||||
sources=ModelSource(
|
||||
hf="Qdrant/fast-bge-base-en",
|
||||
url="https://storage.googleapis.com/qdrant-fastembed/fast-bge-base-en.tar.gz",
|
||||
_deprecated_tar_struct=True,
|
||||
),
|
||||
model_file="model_optimized.onnx",
|
||||
),
|
||||
@@ -34,8 +34,9 @@ supported_onnx_models: list[DenseModelDescription] = [
|
||||
license="mit",
|
||||
size_in_GB=0.21,
|
||||
sources=ModelSource(
|
||||
hf="qdrant/bge-base-en-v1.5-onnx-q",
|
||||
hf="Qdrant/bge-base-en-v1.5-onnx-Q",
|
||||
url="https://storage.googleapis.com/qdrant-fastembed/fast-bge-base-en-v1.5.tar.gz",
|
||||
_deprecated_tar_struct=True,
|
||||
),
|
||||
model_file="model_optimized.onnx",
|
||||
),
|
||||
@@ -63,6 +64,7 @@ supported_onnx_models: list[DenseModelDescription] = [
|
||||
sources=ModelSource(
|
||||
hf="Qdrant/bge-small-en",
|
||||
url="https://storage.googleapis.com/qdrant-fastembed/BAAI-bge-small-en.tar.gz",
|
||||
_deprecated_tar_struct=True,
|
||||
),
|
||||
model_file="model_optimized.onnx",
|
||||
),
|
||||
@@ -75,7 +77,7 @@ supported_onnx_models: list[DenseModelDescription] = [
|
||||
),
|
||||
license="mit",
|
||||
size_in_GB=0.067,
|
||||
sources=ModelSource(hf="qdrant/bge-small-en-v1.5-onnx-q"),
|
||||
sources=ModelSource(hf="Qdrant/bge-small-en-v1.5-onnx-Q"),
|
||||
model_file="model_optimized.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
@@ -90,21 +92,10 @@ supported_onnx_models: list[DenseModelDescription] = [
|
||||
sources=ModelSource(
|
||||
hf="Qdrant/bge-small-zh-v1.5",
|
||||
url="https://storage.googleapis.com/qdrant-fastembed/fast-bge-small-zh-v1.5.tar.gz",
|
||||
_deprecated_tar_struct=True,
|
||||
),
|
||||
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,
|
||||
@@ -189,6 +180,42 @@ supported_onnx_models: list[DenseModelDescription] = [
|
||||
sources=ModelSource(hf="jinaai/jina-clip-v1"),
|
||||
model_file="onnx/text_model.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="minishlab/potion-base-8M",
|
||||
dim=256,
|
||||
description=(
|
||||
"Text embeddings, Unimodal (text), English, 512 input tokens truncation, "
|
||||
"Prefixes for queries/documents: not necessary, 2024 year."
|
||||
),
|
||||
license="mit",
|
||||
size_in_GB=0.030,
|
||||
sources=ModelSource(hf="minishlab/potion-base-8m-onnx"),
|
||||
model_file="model.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="minishlab/potion-retrieval-32M",
|
||||
dim=512,
|
||||
description=(
|
||||
"Text embeddings, Unimodal (text), English, 512 input tokens truncation, "
|
||||
"Prefixes for queries/documents: not necessary, 2025 year."
|
||||
),
|
||||
license="mit",
|
||||
size_in_GB=0.129,
|
||||
sources=ModelSource(hf="minishlab/potion-retrieval-32m-onnx"),
|
||||
model_file="model.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="minishlab/potion-multilingual-128M",
|
||||
dim=256,
|
||||
description=(
|
||||
"Text embeddings, Unimodal (text), Multilingual, 512 input tokens truncation, "
|
||||
"Prefixes for queries/documents: not necessary, 2025 year."
|
||||
),
|
||||
license="mit",
|
||||
size_in_GB=0.512,
|
||||
sources=ModelSource(hf="minishlab/potion-multilingual-128m-onnx"),
|
||||
model_file="model.onnx",
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
@@ -208,14 +235,14 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[NumpyArray]):
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str = "BAAI/bge-small-en-v1.5",
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
cuda: bool = False,
|
||||
device_ids: Optional[list[int]] = None,
|
||||
cache_dir: str | None = None,
|
||||
threads: int | None = None,
|
||||
providers: Sequence[OnnxProvider] | None = None,
|
||||
cuda: bool | Device = Device.AUTO,
|
||||
device_ids: list[int] | None = None,
|
||||
lazy_load: bool = False,
|
||||
device_id: Optional[int] = None,
|
||||
specific_model_path: Optional[str] = None,
|
||||
device_id: int | None = None,
|
||||
specific_model_path: str | None = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
"""
|
||||
@@ -227,10 +254,11 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[NumpyArray]):
|
||||
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.
|
||||
cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`
|
||||
Defaults to Device.AUTO.
|
||||
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.
|
||||
workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, 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.
|
||||
@@ -242,13 +270,13 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[NumpyArray]):
|
||||
super().__init__(model_name, cache_dir, threads, **kwargs)
|
||||
self.providers = providers
|
||||
self.lazy_load = lazy_load
|
||||
|
||||
self._extra_session_options = self._select_exposed_session_options(kwargs)
|
||||
# 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
|
||||
self.device_id: int | None = None
|
||||
if device_id is not None:
|
||||
self.device_id = device_id
|
||||
elif self.device_ids is not None:
|
||||
@@ -256,11 +284,12 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[NumpyArray]):
|
||||
|
||||
self.model_description = self._get_model_description(model_name)
|
||||
self.cache_dir = str(define_cache_dir(cache_dir))
|
||||
self._specific_model_path = specific_model_path
|
||||
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,
|
||||
specific_model_path=self._specific_model_path,
|
||||
)
|
||||
|
||||
if not self.lazy_load:
|
||||
@@ -268,9 +297,9 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[NumpyArray]):
|
||||
|
||||
def embed(
|
||||
self,
|
||||
documents: Union[str, Iterable[str]],
|
||||
documents: str | Iterable[str],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
parallel: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
@@ -297,6 +326,9 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[NumpyArray]):
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_ids=self.device_ids,
|
||||
local_files_only=self._local_files_only,
|
||||
specific_model_path=self._specific_model_path,
|
||||
extra_session_options=self._extra_session_options,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@@ -312,15 +344,18 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[NumpyArray]):
|
||||
"""
|
||||
return onnx_input
|
||||
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[NumpyArray]:
|
||||
def _post_process_onnx_output(
|
||||
self, output: OnnxOutputContext, **kwargs: Any
|
||||
) -> Iterable[NumpyArray]:
|
||||
embeddings = output.model_output
|
||||
|
||||
if embeddings.ndim == 3: # (batch_size, seq_len, embedding_dim)
|
||||
processed_embeddings = embeddings[:, 0]
|
||||
elif embeddings.ndim == 2: # (batch_size, embedding_dim)
|
||||
processed_embeddings = embeddings
|
||||
else:
|
||||
raise ValueError(f"Unsupported embedding shape: {embeddings.shape}")
|
||||
return normalize(processed_embeddings).astype(np.float32)
|
||||
return normalize(processed_embeddings)
|
||||
|
||||
def load_onnx_model(self) -> None:
|
||||
self._load_onnx_model(
|
||||
@@ -330,8 +365,14 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[NumpyArray]):
|
||||
providers=self.providers,
|
||||
cuda=self.cuda,
|
||||
device_id=self.device_id,
|
||||
extra_session_options=self._extra_session_options,
|
||||
)
|
||||
|
||||
def token_count(
|
||||
self, texts: str | Iterable[str], batch_size: int = 1024, **kwargs: Any
|
||||
) -> int:
|
||||
return self._token_count(texts, batch_size=batch_size, **kwargs)
|
||||
|
||||
|
||||
class OnnxTextEmbeddingWorker(TextEmbeddingWorker[NumpyArray]):
|
||||
def init_embedding(
|
||||
|
||||
@@ -1,13 +1,13 @@
|
||||
import os
|
||||
from multiprocessing import get_all_start_methods
|
||||
from pathlib import Path
|
||||
from typing import Any, Iterable, Optional, Sequence, Type, Union
|
||||
from typing import Any, Iterable, Sequence, Type
|
||||
|
||||
import numpy as np
|
||||
from numpy.typing import NDArray
|
||||
from tokenizers import Encoding, Tokenizer
|
||||
|
||||
from fastembed.common.types import NumpyArray, OnnxProvider
|
||||
from fastembed.common.types import NumpyArray, OnnxProvider, Device
|
||||
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
|
||||
@@ -15,23 +15,32 @@ from fastembed.parallel_processor import ParallelWorkerPool
|
||||
|
||||
|
||||
class OnnxTextModel(OnnxModel[T]):
|
||||
ONNX_OUTPUT_NAMES: Optional[list[str]] = None
|
||||
ONNX_OUTPUT_NAMES: list[str] | None = None
|
||||
|
||||
@classmethod
|
||||
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]:
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext, **kwargs: Any) -> Iterable[T]:
|
||||
"""Post-process the ONNX model output to convert it into a usable format.
|
||||
|
||||
Args:
|
||||
output (OnnxOutputContext): The raw output from the ONNX model.
|
||||
**kwargs: Additional keyword arguments that may be needed by specific implementations.
|
||||
|
||||
Returns:
|
||||
Iterable[T]: Post-processed output as an iterable of type T.
|
||||
"""
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.tokenizer: Optional[Tokenizer] = None
|
||||
self.tokenizer: Tokenizer | None = None
|
||||
self.special_token_to_id: dict[str, int] = {}
|
||||
|
||||
def _preprocess_onnx_input(
|
||||
self, onnx_input: dict[str, NumpyArray], **kwargs: Any
|
||||
) -> dict[str, Union[NumpyArray, NDArray[np.int64]]]:
|
||||
) -> dict[str, NumpyArray | NDArray[np.int64]]:
|
||||
"""
|
||||
Preprocess the onnx input.
|
||||
"""
|
||||
@@ -41,10 +50,11 @@ class OnnxTextModel(OnnxModel[T]):
|
||||
self,
|
||||
model_dir: Path,
|
||||
model_file: str,
|
||||
threads: Optional[int],
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
cuda: bool = False,
|
||||
device_id: Optional[int] = None,
|
||||
threads: int | None,
|
||||
providers: Sequence[OnnxProvider] | None = None,
|
||||
cuda: bool | Device = Device.AUTO,
|
||||
device_id: int | None = None,
|
||||
extra_session_options: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
super()._load_onnx_model(
|
||||
model_dir=model_dir,
|
||||
@@ -53,6 +63,7 @@ class OnnxTextModel(OnnxModel[T]):
|
||||
providers=providers,
|
||||
cuda=cuda,
|
||||
device_id=device_id,
|
||||
extra_session_options=extra_session_options,
|
||||
)
|
||||
self.tokenizer, self.special_token_to_id = load_tokenizer(model_dir=model_dir)
|
||||
|
||||
@@ -81,24 +92,34 @@ class OnnxTextModel(OnnxModel[T]):
|
||||
[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._run_model(
|
||||
onnx_input=onnx_input, onnx_output_names=self.ONNX_OUTPUT_NAMES
|
||||
)
|
||||
|
||||
model_output = self.model.run(self.ONNX_OUTPUT_NAMES, onnx_input) # type: ignore[union-attr]
|
||||
return OnnxOutputContext(
|
||||
model_output=model_output[0],
|
||||
model_output=model_output,
|
||||
attention_mask=onnx_input.get("attention_mask", attention_mask),
|
||||
input_ids=onnx_input.get("input_ids", input_ids),
|
||||
)
|
||||
|
||||
def _run_model(
|
||||
self, onnx_input: dict[str, Any], onnx_output_names: list[str] | None = None
|
||||
) -> NumpyArray:
|
||||
return self.model.run(onnx_output_names, onnx_input)[0] # type: ignore[union-attr]
|
||||
|
||||
def _embed_documents(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
documents: Union[str, Iterable[str]],
|
||||
documents: 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,
|
||||
parallel: int | None = None,
|
||||
providers: Sequence[OnnxProvider] | None = None,
|
||||
cuda: bool | Device = Device.AUTO,
|
||||
device_ids: list[int] | None = None,
|
||||
local_files_only: bool = False,
|
||||
specific_model_path: str | None = None,
|
||||
extra_session_options: dict[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[T]:
|
||||
is_small = False
|
||||
@@ -115,7 +136,9 @@ class OnnxTextModel(OnnxModel[T]):
|
||||
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))
|
||||
yield from self._post_process_onnx_output(
|
||||
self.onnx_embed(batch, **kwargs), **kwargs
|
||||
)
|
||||
else:
|
||||
if parallel == 0:
|
||||
parallel = os.cpu_count()
|
||||
@@ -125,9 +148,15 @@ class OnnxTextModel(OnnxModel[T]):
|
||||
"model_name": model_name,
|
||||
"cache_dir": cache_dir,
|
||||
"providers": providers,
|
||||
"local_files_only": local_files_only,
|
||||
"specific_model_path": specific_model_path,
|
||||
**kwargs,
|
||||
**self._get_worker_init_kwargs(),
|
||||
}
|
||||
|
||||
if extra_session_options is not None:
|
||||
params.update(extra_session_options)
|
||||
|
||||
pool = ParallelWorkerPool(
|
||||
num_workers=parallel or 1,
|
||||
worker=self._get_worker_class(),
|
||||
@@ -136,7 +165,20 @@ class OnnxTextModel(OnnxModel[T]):
|
||||
start_method=start_method,
|
||||
)
|
||||
for batch in pool.ordered_map(iter_batch(documents, batch_size), **params):
|
||||
yield from self._post_process_onnx_output(batch) # type: ignore
|
||||
yield from self._post_process_onnx_output(batch, **kwargs) # type: ignore
|
||||
|
||||
def _token_count(self, texts: str | Iterable[str], batch_size: int = 1024, **_: Any) -> int:
|
||||
if not hasattr(self, "model") or self.model is None:
|
||||
self.load_onnx_model() # loads the tokenizer as well
|
||||
|
||||
token_num = 0
|
||||
assert self.tokenizer is not None
|
||||
texts = [texts] if isinstance(texts, str) else texts
|
||||
for batch in iter_batch(texts, batch_size):
|
||||
for tokens in self.tokenizer.encode_batch(batch):
|
||||
token_num += sum(tokens.attention_mask)
|
||||
|
||||
return token_num
|
||||
|
||||
|
||||
class TextEmbeddingWorker(EmbeddingWorker[T]):
|
||||
|
||||
@@ -1,9 +1,11 @@
|
||||
from typing import Any, Iterable, Type
|
||||
|
||||
import numpy as np
|
||||
from numpy.typing import NDArray
|
||||
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.common.utils import mean_pooling
|
||||
from fastembed.text.onnx_embedding import OnnxTextEmbedding, OnnxTextEmbeddingWorker
|
||||
from fastembed.common.model_description import DenseModelDescription, ModelSource
|
||||
|
||||
@@ -80,6 +82,7 @@ supported_pooled_models: list[DenseModelDescription] = [
|
||||
sources=ModelSource(
|
||||
hf="qdrant/multilingual-e5-large-onnx",
|
||||
url="https://storage.googleapis.com/qdrant-fastembed/fast-multilingual-e5-large.tar.gz",
|
||||
_deprecated_tar_struct=True,
|
||||
),
|
||||
model_file="model.onnx",
|
||||
additional_files=["model.onnx_data"],
|
||||
@@ -93,16 +96,10 @@ class PooledEmbedding(OnnxTextEmbedding):
|
||||
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
|
||||
def mean_pooling(
|
||||
cls, model_output: NumpyArray, attention_mask: NDArray[np.int64]
|
||||
) -> NumpyArray:
|
||||
return mean_pooling(model_output, attention_mask)
|
||||
|
||||
@classmethod
|
||||
def _list_supported_models(cls) -> list[DenseModelDescription]:
|
||||
@@ -113,13 +110,15 @@ class PooledEmbedding(OnnxTextEmbedding):
|
||||
"""
|
||||
return supported_pooled_models
|
||||
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[NumpyArray]:
|
||||
def _post_process_onnx_output(
|
||||
self, output: OnnxOutputContext, **kwargs: Any
|
||||
) -> 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)
|
||||
return self.mean_pooling(embeddings, attn_mask)
|
||||
|
||||
|
||||
class PooledEmbeddingWorker(OnnxTextEmbeddingWorker):
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
from typing import Any, Iterable, Type
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastembed.common.types import NumpyArray
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
@@ -22,6 +21,7 @@ supported_pooled_normalized_models: list[DenseModelDescription] = [
|
||||
sources=ModelSource(
|
||||
url="https://storage.googleapis.com/qdrant-fastembed/sentence-transformers-all-MiniLM-L6-v2.tar.gz",
|
||||
hf="qdrant/all-MiniLM-L6-v2-onnx",
|
||||
_deprecated_tar_struct=True,
|
||||
),
|
||||
model_file="model.onnx",
|
||||
),
|
||||
@@ -57,9 +57,9 @@ supported_pooled_normalized_models: list[DenseModelDescription] = [
|
||||
"Prefixes for queries/documents: not necessary, 2024 year."
|
||||
),
|
||||
license="apache-2.0",
|
||||
size_in_GB=0.32,
|
||||
size_in_GB=0.64,
|
||||
sources=ModelSource(hf="jinaai/jina-embeddings-v2-base-de"),
|
||||
model_file="onnx/model_fp16.onnx",
|
||||
model_file="onnx/model.onnx",
|
||||
),
|
||||
DenseModelDescription(
|
||||
model="jinaai/jina-embeddings-v2-base-code",
|
||||
@@ -109,6 +109,18 @@ supported_pooled_normalized_models: list[DenseModelDescription] = [
|
||||
sources=ModelSource(hf="thenlper/gte-base"),
|
||||
model_file="onnx/model.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",
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
@@ -126,13 +138,15 @@ class PooledNormalizedEmbedding(PooledEmbedding):
|
||||
"""
|
||||
return supported_pooled_normalized_models
|
||||
|
||||
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[NumpyArray]:
|
||||
def _post_process_onnx_output(
|
||||
self, output: OnnxOutputContext, **kwargs: Any
|
||||
) -> 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)
|
||||
return normalize(self.mean_pooling(embeddings, attn_mask))
|
||||
|
||||
|
||||
class PooledNormalizedEmbeddingWorker(OnnxTextEmbeddingWorker):
|
||||
|
||||
@@ -0,0 +1,71 @@
|
||||
from typing import Any, Type
|
||||
|
||||
from fastembed.common.model_description import DenseModelDescription, ModelSource
|
||||
from fastembed.text.onnx_embedding import OnnxTextEmbedding, OnnxTextEmbeddingWorker
|
||||
|
||||
supported_siglip_models: list[DenseModelDescription] = [
|
||||
DenseModelDescription(
|
||||
model="google/siglip2-base-patch16-224",
|
||||
dim=768,
|
||||
description=(
|
||||
"Text embeddings, Multimodal (text&image), multilingual, 64 input tokens truncation, "
|
||||
"Prefixes for queries/documents: not necessary, 2025 year"
|
||||
),
|
||||
license="apache-2.0",
|
||||
size_in_GB=1.13,
|
||||
sources=ModelSource(hf="onnx-community/siglip2-base-patch16-224-ONNX"),
|
||||
model_file="onnx/text_model.onnx",
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
class SiglipOnnxTextEmbedding(OnnxTextEmbedding):
|
||||
"""SigLIP text tower.
|
||||
|
||||
SigLIP always pools the hidden state at the last sequence position, whether or not it is
|
||||
padding, so every batch must be padded to the model's fixed training length (rather than to
|
||||
the longest sequence in the batch) or the resulting embeddings become dependent on what else
|
||||
is in the batch.
|
||||
|
||||
The exported graph also returns both `last_hidden_state` and `pooler_output`; only the
|
||||
latter is the text embedding, so it must be selected explicitly.
|
||||
"""
|
||||
|
||||
ONNX_OUTPUT_NAMES = ["pooler_output"]
|
||||
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type[OnnxTextEmbeddingWorker]:
|
||||
return SiglipTextEmbeddingWorker
|
||||
|
||||
@classmethod
|
||||
def _list_supported_models(cls) -> list[DenseModelDescription]:
|
||||
return supported_siglip_models
|
||||
|
||||
def load_onnx_model(self) -> None:
|
||||
super().load_onnx_model()
|
||||
if self.tokenizer is not None:
|
||||
truncation = self.tokenizer.truncation
|
||||
padding = self.tokenizer.padding
|
||||
if truncation and padding and padding.get("length") is None:
|
||||
self.tokenizer.enable_padding(
|
||||
direction=padding["direction"],
|
||||
pad_id=padding["pad_id"],
|
||||
pad_type_id=padding["pad_type_id"],
|
||||
pad_token=padding["pad_token"],
|
||||
length=truncation["max_length"],
|
||||
)
|
||||
|
||||
|
||||
class SiglipTextEmbeddingWorker(OnnxTextEmbeddingWorker):
|
||||
def init_embedding(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: str,
|
||||
**kwargs: Any,
|
||||
) -> OnnxTextEmbedding:
|
||||
return SiglipOnnxTextEmbedding(
|
||||
model_name=model_name,
|
||||
cache_dir=cache_dir,
|
||||
threads=1,
|
||||
**kwargs,
|
||||
)
|
||||
@@ -1,24 +1,32 @@
|
||||
import warnings
|
||||
from typing import Any, Iterable, Optional, Sequence, Type, Union
|
||||
from typing import Any, Iterable, Sequence, Type
|
||||
from dataclasses import asdict
|
||||
|
||||
from fastembed.common.types import NumpyArray, OnnxProvider
|
||||
from fastembed.common.types import NumpyArray, OnnxProvider, Device
|
||||
from fastembed.text.clip_embedding import CLIPOnnxEmbedding
|
||||
from fastembed.text.custom_text_embedding import CustomTextEmbedding
|
||||
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.builtin_sentence_embedding import BuiltinSentenceEmbedding
|
||||
from fastembed.text.last_token_normalized_embedding import LastTokenNormalizedEmbedding
|
||||
from fastembed.text.onnx_embedding import OnnxTextEmbedding
|
||||
from fastembed.text.siglip_embedding import SiglipOnnxTextEmbedding
|
||||
from fastembed.text.text_embedding_base import TextEmbeddingBase
|
||||
from fastembed.common.model_description import DenseModelDescription
|
||||
from fastembed.common.model_description import DenseModelDescription, ModelSource, PoolingType
|
||||
|
||||
|
||||
class TextEmbedding(TextEmbeddingBase):
|
||||
EMBEDDINGS_REGISTRY: list[Type[TextEmbeddingBase]] = [
|
||||
OnnxTextEmbedding,
|
||||
CLIPOnnxEmbedding,
|
||||
SiglipOnnxTextEmbedding,
|
||||
PooledNormalizedEmbedding,
|
||||
PooledEmbedding,
|
||||
JinaEmbeddingV3,
|
||||
BuiltinSentenceEmbedding,
|
||||
LastTokenNormalizedEmbedding,
|
||||
CustomTextEmbedding,
|
||||
]
|
||||
|
||||
@classmethod
|
||||
@@ -37,43 +45,61 @@ class TextEmbedding(TextEmbeddingBase):
|
||||
result.extend(embedding._list_supported_models())
|
||||
return result
|
||||
|
||||
@classmethod
|
||||
def add_custom_model(
|
||||
cls,
|
||||
model: str,
|
||||
pooling: PoolingType,
|
||||
normalization: bool,
|
||||
sources: ModelSource,
|
||||
dim: int,
|
||||
model_file: str = "onnx/model.onnx",
|
||||
description: str = "",
|
||||
license: str = "",
|
||||
size_in_gb: float = 0.0,
|
||||
additional_files: list[str] | None = None,
|
||||
) -> None:
|
||||
registered_models = cls._list_supported_models()
|
||||
for registered_model in registered_models:
|
||||
if model.lower() == registered_model.model.lower():
|
||||
raise ValueError(
|
||||
f"Model {model} is already registered in TextEmbedding, if you still want to add this model, "
|
||||
f"please use another model name"
|
||||
)
|
||||
|
||||
CustomTextEmbedding.add_model(
|
||||
DenseModelDescription(
|
||||
model=model,
|
||||
sources=sources,
|
||||
dim=dim,
|
||||
model_file=model_file,
|
||||
description=description,
|
||||
license=license,
|
||||
size_in_GB=size_in_gb,
|
||||
additional_files=additional_files or [],
|
||||
),
|
||||
pooling=pooling,
|
||||
normalization=normalization,
|
||||
)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str = "BAAI/bge-small-en-v1.5",
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
providers: Optional[Sequence[OnnxProvider]] = None,
|
||||
cuda: bool = False,
|
||||
device_ids: Optional[list[int]] = None,
|
||||
cache_dir: str | None = None,
|
||||
threads: int | None = None,
|
||||
providers: Sequence[OnnxProvider] | None = None,
|
||||
cuda: bool | Device = Device.AUTO,
|
||||
device_ids: list[int] | None = 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":
|
||||
if model_name.lower() == "jinaai/jina-embeddings-v2-base-de":
|
||||
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. ",
|
||||
"The model 'jinaai/jina-embeddings-v2-base-de' used to run with fp16 model, but due to onnxruntime updates, now it runs with the original fp32 model.",
|
||||
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):
|
||||
@@ -94,11 +120,45 @@ class TextEmbedding(TextEmbeddingBase):
|
||||
"Please check the supported models using `TextEmbedding.list_supported_models()`"
|
||||
)
|
||||
|
||||
@property
|
||||
def embedding_size(self) -> int:
|
||||
"""Get the embedding size of the current model"""
|
||||
if self._embedding_size is None:
|
||||
self._embedding_size = self.get_embedding_size(self.model_name)
|
||||
return self._embedding_size
|
||||
|
||||
@classmethod
|
||||
def get_embedding_size(cls, model_name: str) -> int:
|
||||
"""Get the embedding size of the passed model
|
||||
|
||||
Args:
|
||||
model_name (str): The name of the model to get embedding size for.
|
||||
|
||||
Returns:
|
||||
int: The size of the embedding.
|
||||
|
||||
Raises:
|
||||
ValueError: If the model name is not found in the supported models.
|
||||
"""
|
||||
descriptions = cls._list_supported_models()
|
||||
embedding_size: int | None = None
|
||||
for description in descriptions:
|
||||
if description.model.lower() == model_name.lower():
|
||||
embedding_size = description.dim
|
||||
break
|
||||
if embedding_size is None:
|
||||
model_names = [description.model for description in descriptions]
|
||||
raise ValueError(
|
||||
f"Embedding size for model {model_name} was None. "
|
||||
f"Available model names: {model_names}"
|
||||
)
|
||||
return embedding_size
|
||||
|
||||
def embed(
|
||||
self,
|
||||
documents: Union[str, Iterable[str]],
|
||||
documents: str | Iterable[str],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
parallel: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
@@ -118,7 +178,7 @@ class TextEmbedding(TextEmbeddingBase):
|
||||
"""
|
||||
yield from self.model.embed(documents, batch_size, parallel, **kwargs)
|
||||
|
||||
def query_embed(self, query: Union[str, Iterable[str]], **kwargs: Any) -> Iterable[NumpyArray]:
|
||||
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
Embeds queries
|
||||
|
||||
@@ -144,3 +204,17 @@ class TextEmbedding(TextEmbeddingBase):
|
||||
"""
|
||||
# This is model-specific, so that different models can have specialized implementations
|
||||
yield from self.model.passage_embed(texts, **kwargs)
|
||||
|
||||
def token_count(
|
||||
self, texts: str | Iterable[str], batch_size: int = 1024, **kwargs: Any
|
||||
) -> int:
|
||||
"""Returns the number of tokens in the texts.
|
||||
|
||||
Args:
|
||||
texts (str | Iterable[str]): The list of texts to embed.
|
||||
batch_size (int): Batch size for encoding
|
||||
|
||||
Returns:
|
||||
int: Sum of number of tokens in the texts.
|
||||
"""
|
||||
return self.model.token_count(texts, batch_size=batch_size, **kwargs)
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from typing import Iterable, Optional, Union, Any
|
||||
from typing import Iterable, Any
|
||||
|
||||
from fastembed.common.model_description import DenseModelDescription
|
||||
from fastembed.common.types import NumpyArray
|
||||
@@ -9,20 +9,21 @@ class TextEmbeddingBase(ModelManagement[DenseModelDescription]):
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
cache_dir: Optional[str] = None,
|
||||
threads: Optional[int] = None,
|
||||
cache_dir: str | None = None,
|
||||
threads: int | None = 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)
|
||||
self._embedding_size: int | None = None
|
||||
|
||||
def embed(
|
||||
self,
|
||||
documents: Union[str, Iterable[str]],
|
||||
documents: str | Iterable[str],
|
||||
batch_size: int = 256,
|
||||
parallel: Optional[int] = None,
|
||||
parallel: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterable[NumpyArray]:
|
||||
raise NotImplementedError()
|
||||
@@ -42,7 +43,7 @@ class TextEmbeddingBase(ModelManagement[DenseModelDescription]):
|
||||
# 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: Any) -> Iterable[NumpyArray]:
|
||||
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:
|
||||
"""
|
||||
Embeds queries
|
||||
|
||||
@@ -58,3 +59,17 @@ class TextEmbeddingBase(ModelManagement[DenseModelDescription]):
|
||||
yield from self.embed([query], **kwargs)
|
||||
else:
|
||||
yield from self.embed(query, **kwargs)
|
||||
|
||||
@classmethod
|
||||
def get_embedding_size(cls, model_name: str) -> int:
|
||||
"""Returns embedding size of the passed model."""
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
@property
|
||||
def embedding_size(self) -> int:
|
||||
"""Returns embedding size for the current model"""
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
def token_count(self, texts: str | Iterable[str], **kwargs: Any) -> int:
|
||||
"""Returns the number of tokens in the texts."""
|
||||
raise NotImplementedError("Subclasses must implement this method")
|
||||
|
||||
@@ -13,6 +13,7 @@ copyright: |
|
||||
theme:
|
||||
name: material
|
||||
logo: assets/favicon.png
|
||||
favicon: assets/favicon.png
|
||||
custom_dir: docs/overrides
|
||||
icon:
|
||||
repo: fontawesome/brands/github
|
||||
|
||||
Generated
+4281
File diff suppressed because it is too large
Load Diff
+31
-21
@@ -1,9 +1,9 @@
|
||||
[tool.poetry]
|
||||
name = "fastembed"
|
||||
version = "0.5.1"
|
||||
name = "fastembed-gpu"
|
||||
version = "0.8.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"
|
||||
license = "Apache-2.0"
|
||||
readme = "README.md"
|
||||
packages = [{include = "fastembed"}]
|
||||
homepage = "https://github.com/qdrant/fastembed"
|
||||
@@ -11,46 +11,56 @@ repository = "https://github.com/qdrant/fastembed"
|
||||
keywords = ["vector", "embedding", "neural", "search", "qdrant", "sentence-transformers"]
|
||||
|
||||
[tool.poetry.dependencies]
|
||||
python = ">=3.9.0"
|
||||
python = ">=3.10.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" },
|
||||
{ version = ">=1.21,<2.3.0", python = "3.10" },
|
||||
{ version = ">=1.21", python = "3.11" },
|
||||
{ version = ">=1.26", python = "3.12" },
|
||||
{ version = ">=2.1.0", python = "3.13" },
|
||||
{ version = ">=2.3.0", python = ">=3.14" },
|
||||
]
|
||||
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" },
|
||||
onnxruntime-gpu = [
|
||||
{ version = ">=1.17.0,!=1.20.0,<1.24", python = "3.10" },
|
||||
{ version = ">=1.17.0,!=1.20.0,!=1.24.0,!=1.24.1", python = ">=3.11,<3.13" },
|
||||
{ version = ">1.21.0,!=1.24.0,!=1.24.1", python = "3.13" },
|
||||
{ version = ">=1.24.2", python = ">=3.14" },
|
||||
]
|
||||
tqdm = "^4.66"
|
||||
requests = "^2.31"
|
||||
tokenizers = ">=0.15,<1.0"
|
||||
huggingface-hub = ">=0.20,<1.0"
|
||||
huggingface-hub = ">=0.20,<2.0"
|
||||
loguru = "^0.7.2"
|
||||
pillow = ">=10.3.0,<12.0.0"
|
||||
mmh3 = "^4.1.0"
|
||||
pillow = [
|
||||
{ version = ">=10.3.0,<13.0", python = ">=3.10,<3.13" },
|
||||
{ version = ">=11.0.0,<13.0", python = "3.13" },
|
||||
{ version = ">=12.0.0,<13.0", python = ">=3.14" },
|
||||
]
|
||||
mmh3 = ">=4.1.0,<6.0.0"
|
||||
py-rust-stemmers = "^0.1.0"
|
||||
|
||||
[tool.poetry.group.test.dependencies]
|
||||
pytest = "^7.4.2"
|
||||
pytest = ">=7.4.2,<10.0.0"
|
||||
ruff = ">=0.3.1,<1.0"
|
||||
|
||||
[tool.poetry.group.dev.dependencies]
|
||||
notebook = ">=7.0.2"
|
||||
pre-commit = "^3.6.2"
|
||||
onnx = ">=1.15.0"
|
||||
pre-commit = ">=3.6.2,<5.0.0"
|
||||
onnx = [
|
||||
{ version = ">=1.15.0", python = ">=3.10,<3.13" },
|
||||
{ version = ">=1.18.0", python = "3.13" },
|
||||
{ version = ">=1.20.0", python = ">=3.14" },
|
||||
]
|
||||
|
||||
[tool.poetry.group.docs.dependencies]
|
||||
mkdocs-material = "^9.5.10"
|
||||
mkdocstrings = "^0.24.0"
|
||||
pillow = ">=10.3.0,<12.0.0"
|
||||
mkdocstrings = ">=0.24,<1.1"
|
||||
pillow = ">=10.3.0,<13.0.0"
|
||||
cairosvg = "^2.7.1"
|
||||
mknotebooks = "^0.8.0"
|
||||
|
||||
[tool.poetry.group.types.dependencies]
|
||||
pyright = ">=1.1.293"
|
||||
mypy = "^1.0.0"
|
||||
mypy = ">=1,<3"
|
||||
|
||||
[build-system]
|
||||
requires = ["poetry-core"]
|
||||
|
||||
+116
-103
@@ -1,4 +1,5 @@
|
||||
import os
|
||||
from contextlib import contextmanager
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
@@ -7,98 +8,119 @@ 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: str) -> None:
|
||||
_MODELS_TO_CACHE = ("Qdrant/bm42-all-minilm-l6-v2-attentions", "Qdrant/bm25")
|
||||
MODELS_TO_CACHE = tuple([x.lower() for x in _MODELS_TO_CACHE])
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def model_cache():
|
||||
is_ci = os.getenv("CI")
|
||||
model = SparseTextEmbedding(model_name=model_name)
|
||||
cache = {}
|
||||
|
||||
output = list(
|
||||
model.query_embed(
|
||||
[
|
||||
"I must not fear. Fear is the mind-killer.",
|
||||
]
|
||||
)
|
||||
)
|
||||
@contextmanager
|
||||
def get_model(model_name: str):
|
||||
lowercase_model_name = model_name.lower()
|
||||
if lowercase_model_name not in cache:
|
||||
cache[lowercase_model_name] = SparseTextEmbedding(lowercase_model_name)
|
||||
yield cache[lowercase_model_name]
|
||||
if lowercase_model_name not in MODELS_TO_CACHE:
|
||||
print("deleting model")
|
||||
model_inst = cache.pop(lowercase_model_name)
|
||||
if is_ci:
|
||||
delete_model_cache(model_inst.model._model_dir)
|
||||
del model_inst
|
||||
|
||||
assert len(output) == 1
|
||||
|
||||
for result in output:
|
||||
assert len(result.indices) == len(result.values)
|
||||
assert np.allclose(result.values, np.ones(len(result.values)))
|
||||
|
||||
quotes = [
|
||||
"I must not fear. Fear is the mind-killer.",
|
||||
"All animals are equal, but some animals are more equal than others.",
|
||||
"It was a pleasure to burn.",
|
||||
"The sky above the port was the color of television, tuned to a dead channel.",
|
||||
"In the beginning, the universe was created."
|
||||
" This has made a lot of people very angry and been widely regarded as a bad move.",
|
||||
"It's a truth universally acknowledged that a zombie in possession of brains must be in want of more brains.",
|
||||
"War is peace. Freedom is slavery. Ignorance is strength.",
|
||||
"We're not in Infinity; we're in the suburbs.",
|
||||
"I was a thousand times more evil than thou!",
|
||||
"History is merely a list of surprises... It can only prepare us to be surprised yet again.",
|
||||
".", # Empty string
|
||||
]
|
||||
|
||||
output = list(model.embed(quotes))
|
||||
|
||||
assert len(output) == len(quotes)
|
||||
|
||||
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(
|
||||
[
|
||||
"привет мир!",
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
assert len(output) == 1
|
||||
|
||||
for result in output:
|
||||
assert len(result.indices) == len(result.values)
|
||||
assert len(result.indices) == 2
|
||||
yield get_model
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
for name, model in cache.items():
|
||||
delete_model_cache(model.model._model_dir)
|
||||
cache.clear()
|
||||
|
||||
|
||||
@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")
|
||||
def test_attention_embeddings(model_cache, model_name: str) -> None:
|
||||
with model_cache(model_name) as model:
|
||||
output = list(
|
||||
model.query_embed(
|
||||
[
|
||||
"I must not fear. Fear is the mind-killer.",
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
model = SparseTextEmbedding(model_name=model_name)
|
||||
assert len(output) == 1
|
||||
|
||||
docs = ["hello world", "attention embedding", "Mangez-vous vraiment des grenouilles?"] * 100
|
||||
embeddings = list(model.embed(docs, batch_size=10, parallel=2))
|
||||
for result in output:
|
||||
assert len(result.indices) == len(result.values)
|
||||
assert np.allclose(result.values, np.ones(len(result.values)))
|
||||
|
||||
embeddings_2 = list(model.embed(docs, batch_size=10, parallel=None))
|
||||
quotes = [
|
||||
"I must not fear. Fear is the mind-killer.",
|
||||
"All animals are equal, but some animals are more equal than others.",
|
||||
"It was a pleasure to burn.",
|
||||
"The sky above the port was the color of television, tuned to a dead channel.",
|
||||
"In the beginning, the universe was created."
|
||||
" This has made a lot of people very angry and been widely regarded as a bad move.",
|
||||
"It's a truth universally acknowledged that a zombie in possession of brains must be in want of more brains.",
|
||||
"War is peace. Freedom is slavery. Ignorance is strength.",
|
||||
"We're not in Infinity; we're in the suburbs.",
|
||||
"I was a thousand times more evil than thou!",
|
||||
"History is merely a list of surprises... It can only prepare us to be surprised yet again.",
|
||||
".", # Empty string
|
||||
]
|
||||
|
||||
embeddings_3 = list(model.embed(docs, batch_size=10, parallel=0))
|
||||
output = list(model.embed(quotes))
|
||||
|
||||
assert len(embeddings) == len(docs)
|
||||
assert len(output) == len(quotes)
|
||||
|
||||
for emb_1, emb_2, emb_3 in zip(embeddings, embeddings_2, embeddings_3):
|
||||
assert np.allclose(emb_1.indices, emb_2.indices)
|
||||
assert np.allclose(emb_1.indices, emb_3.indices)
|
||||
assert np.allclose(emb_1.values, emb_2.values)
|
||||
assert np.allclose(emb_1.values, emb_3.values)
|
||||
for result in output[:-1]:
|
||||
assert len(result.indices) == len(result.values)
|
||||
assert len(result.indices) > 0
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
assert len(output[-1].indices) == 0
|
||||
|
||||
# Test support for unknown languages
|
||||
output = list(
|
||||
model.query_embed(
|
||||
[
|
||||
"привет мир!",
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
assert len(output) == 1
|
||||
|
||||
for result in output:
|
||||
assert len(result.indices) == len(result.values)
|
||||
assert len(result.indices) == 2
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_name", ["Qdrant/bm42-all-minilm-l6-v2-attentions", "Qdrant/bm25"])
|
||||
def test_parallel_processing(model_cache, model_name: str) -> None:
|
||||
with model_cache(model_name) as model:
|
||||
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))
|
||||
|
||||
embeddings_3 = list(model.embed(docs, batch_size=10, parallel=0))
|
||||
|
||||
assert len(embeddings) == len(docs)
|
||||
|
||||
for emb_1, emb_2, emb_3 in zip(embeddings, embeddings_2, embeddings_3):
|
||||
assert np.allclose(emb_1.indices, emb_2.indices)
|
||||
assert np.allclose(emb_1.indices, emb_3.indices)
|
||||
assert np.allclose(emb_1.values, emb_2.values)
|
||||
assert np.allclose(emb_1.values, emb_3.values)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_name", ["Qdrant/bm25"])
|
||||
def test_multilanguage(model_name: str) -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
|
||||
def test_multilanguage(model_cache, model_name: str) -> None:
|
||||
docs = ["Mangez-vous vraiment des grenouilles?", "Je suis au lit"]
|
||||
|
||||
model = SparseTextEmbedding(model_name=model_name, language="french")
|
||||
@@ -109,39 +131,30 @@ def test_multilanguage(model_name: str) -> None:
|
||||
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,)
|
||||
with model_cache(model_name) as model: # 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)
|
||||
assert embeddings[1].values.shape == (4,)
|
||||
assert embeddings[1].indices.shape == (4,)
|
||||
|
||||
|
||||
@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)
|
||||
def test_special_characters(model_cache, model_name: str) -> None:
|
||||
with model_cache(model_name) as model:
|
||||
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",
|
||||
]
|
||||
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,)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_name", ["Qdrant/bm42-all-minilm-l6-v2-attentions"])
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
import numpy as np
|
||||
|
||||
from fastembed import (
|
||||
TextEmbedding,
|
||||
SparseTextEmbedding,
|
||||
@@ -5,6 +7,7 @@ from fastembed import (
|
||||
LateInteractionMultimodalEmbedding,
|
||||
LateInteractionTextEmbedding,
|
||||
)
|
||||
from fastembed.common.utils import last_token_pooling
|
||||
|
||||
|
||||
def test_text_list_supported_models():
|
||||
@@ -28,3 +31,31 @@ def test_text_list_supported_models():
|
||||
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"]
|
||||
|
||||
|
||||
def test_last_token_pooling():
|
||||
token_embeddings = np.array(
|
||||
[
|
||||
[[1.0, 1.0], [2.0, 2.0], [9.0, 9.0], [9.0, 9.0]], # 2 real tokens, then padding
|
||||
[[3.0, 3.0], [4.0, 4.0], [5.0, 5.0], [6.0, 6.0]], # no padding
|
||||
]
|
||||
)
|
||||
attention_mask = np.array([[1, 1, 0, 0], [1, 1, 1, 1]], dtype=np.int64)
|
||||
|
||||
pooled = last_token_pooling(token_embeddings, attention_mask)
|
||||
|
||||
assert np.allclose(pooled, [[2.0, 2.0], [6.0, 6.0]])
|
||||
|
||||
|
||||
def test_last_token_pooling_with_left_padding():
|
||||
token_embeddings = np.array(
|
||||
[
|
||||
[[9.0, 9.0], [9.0, 9.0], [1.0, 1.0], [2.0, 2.0]], # padding, then 2 real tokens
|
||||
[[3.0, 3.0], [4.0, 4.0], [5.0, 5.0], [6.0, 6.0]], # no padding
|
||||
]
|
||||
)
|
||||
attention_mask = np.array([[0, 0, 1, 1], [1, 1, 1, 1]], dtype=np.int64)
|
||||
|
||||
pooled = last_token_pooling(token_embeddings, attention_mask)
|
||||
|
||||
assert np.allclose(pooled, [[2.0, 2.0], [6.0, 6.0]])
|
||||
|
||||
@@ -0,0 +1,305 @@
|
||||
import itertools
|
||||
import os
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from fastembed.common.model_description import (
|
||||
PoolingType,
|
||||
ModelSource,
|
||||
DenseModelDescription,
|
||||
BaseModelDescription,
|
||||
)
|
||||
from fastembed.common.onnx_model import OnnxOutputContext
|
||||
from fastembed.common.utils import normalize, mean_pooling, last_token_pooling
|
||||
from fastembed.text.custom_text_embedding import CustomTextEmbedding, PostprocessingConfig
|
||||
from fastembed.rerank.cross_encoder.custom_text_cross_encoder import CustomTextCrossEncoder
|
||||
from fastembed.rerank.cross_encoder import TextCrossEncoder
|
||||
from fastembed.text.text_embedding import TextEmbedding
|
||||
from tests.utils import delete_model_cache
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def restore_custom_models_fixture():
|
||||
CustomTextEmbedding.SUPPORTED_MODELS = []
|
||||
CustomTextEmbedding.POSTPROCESSING_MAPPING = {}
|
||||
CustomTextCrossEncoder.SUPPORTED_MODELS = []
|
||||
yield
|
||||
CustomTextEmbedding.SUPPORTED_MODELS = []
|
||||
CustomTextEmbedding.POSTPROCESSING_MAPPING = {}
|
||||
CustomTextCrossEncoder.SUPPORTED_MODELS = []
|
||||
|
||||
|
||||
def test_text_custom_model():
|
||||
is_ci = os.getenv("CI")
|
||||
custom_model_name = "intfloat/multilingual-e5-small"
|
||||
canonical_vector = np.array(
|
||||
[3.1317e-02, 3.0939e-02, -3.5117e-02, -6.7274e-02, 8.5084e-02], dtype=np.float32
|
||||
)
|
||||
pooling = PoolingType.MEAN
|
||||
normalization = True
|
||||
dim = 384
|
||||
size_in_gb = 0.47
|
||||
source = ModelSource(hf=custom_model_name)
|
||||
|
||||
TextEmbedding.add_custom_model(
|
||||
custom_model_name,
|
||||
pooling=pooling,
|
||||
normalization=normalization,
|
||||
sources=source,
|
||||
dim=dim,
|
||||
size_in_gb=size_in_gb,
|
||||
)
|
||||
|
||||
assert CustomTextEmbedding.SUPPORTED_MODELS[0] == DenseModelDescription(
|
||||
model=custom_model_name,
|
||||
sources=source,
|
||||
model_file="onnx/model.onnx",
|
||||
description="",
|
||||
license="",
|
||||
size_in_GB=size_in_gb,
|
||||
additional_files=[],
|
||||
dim=dim,
|
||||
tasks={},
|
||||
)
|
||||
assert CustomTextEmbedding.POSTPROCESSING_MAPPING[custom_model_name] == PostprocessingConfig(
|
||||
pooling=pooling, normalization=normalization
|
||||
)
|
||||
|
||||
model = TextEmbedding(custom_model_name)
|
||||
docs = ["hello world", "flag embedding"]
|
||||
embeddings = list(model.embed(docs))
|
||||
embeddings = np.stack(embeddings, axis=0)
|
||||
assert embeddings.shape == (2, dim)
|
||||
|
||||
assert np.allclose(embeddings[0, : canonical_vector.shape[0]], canonical_vector, atol=1e-3)
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
def test_text_custom_model_parallel_processing():
|
||||
is_ci = os.getenv("CI")
|
||||
custom_model_name = "intfloat/multilingual-e5-small"
|
||||
dim = 384
|
||||
|
||||
TextEmbedding.add_custom_model(
|
||||
custom_model_name,
|
||||
pooling=PoolingType.MEAN,
|
||||
normalization=True,
|
||||
sources=ModelSource(hf=custom_model_name),
|
||||
dim=dim,
|
||||
size_in_gb=0.47,
|
||||
)
|
||||
|
||||
model = TextEmbedding(custom_model_name)
|
||||
docs = ["hello world", "flag embedding"] * 50
|
||||
embeddings = np.stack(list(model.embed(docs, batch_size=10, parallel=2)), axis=0)
|
||||
|
||||
assert embeddings.shape == (len(docs), dim)
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
def test_cross_encoder_custom_model():
|
||||
is_ci = os.getenv("CI")
|
||||
custom_model_name = "Xenova/ms-marco-MiniLM-L-4-v2"
|
||||
size_in_gb = 0.08
|
||||
source = ModelSource(hf=custom_model_name)
|
||||
canonical_vector = np.array([-5.7170815, -11.112114], dtype=np.float32)
|
||||
|
||||
TextCrossEncoder.add_custom_model(
|
||||
custom_model_name,
|
||||
model_file="onnx/model.onnx",
|
||||
sources=source,
|
||||
size_in_gb=size_in_gb,
|
||||
)
|
||||
|
||||
assert CustomTextCrossEncoder.SUPPORTED_MODELS[0] == BaseModelDescription(
|
||||
model=custom_model_name,
|
||||
sources=source,
|
||||
model_file="onnx/model.onnx",
|
||||
description="",
|
||||
license="",
|
||||
size_in_GB=size_in_gb,
|
||||
)
|
||||
|
||||
model = TextCrossEncoder(custom_model_name)
|
||||
pairs = [
|
||||
("What is AI?", "Artificial intelligence is ..."),
|
||||
("What is ML?", "Machine learning is ..."),
|
||||
]
|
||||
scores = list(model.rerank_pairs(pairs))
|
||||
|
||||
embeddings = np.stack(scores, axis=0)
|
||||
assert embeddings.shape == (2,)
|
||||
assert np.allclose(embeddings, canonical_vector, atol=1e-3)
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
def test_cross_encoder_custom_model_parallel_processing():
|
||||
is_ci = os.getenv("CI")
|
||||
custom_model_name = "Xenova/ms-marco-MiniLM-L-4-v2"
|
||||
|
||||
TextCrossEncoder.add_custom_model(
|
||||
custom_model_name,
|
||||
model_file="onnx/model.onnx",
|
||||
sources=ModelSource(hf=custom_model_name),
|
||||
size_in_gb=0.08,
|
||||
)
|
||||
|
||||
model = TextCrossEncoder(custom_model_name)
|
||||
pairs = [("What is AI?", "Artificial intelligence is ...")] * 50
|
||||
scores = np.stack(list(model.rerank_pairs(pairs, batch_size=10, parallel=2)), axis=0)
|
||||
|
||||
assert scores.shape == (len(pairs),)
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
def test_mock_add_custom_models():
|
||||
dim = 5
|
||||
size_in_gb = 0.1
|
||||
source = ModelSource(hf="artificial")
|
||||
|
||||
num_tokens = 10
|
||||
dummy_pooled_embedding = np.random.random((1, dim)).astype(np.float32)
|
||||
dummy_token_embedding = np.random.random((1, num_tokens, dim)).astype(np.float32)
|
||||
dummy_attention_mask = np.ones((1, num_tokens)).astype(np.int64)
|
||||
|
||||
dummy_token_output = OnnxOutputContext(
|
||||
model_output=dummy_token_embedding, attention_mask=dummy_attention_mask
|
||||
)
|
||||
dummy_pooled_output = OnnxOutputContext(model_output=dummy_pooled_embedding)
|
||||
input_data = {
|
||||
f"{PoolingType.MEAN.lower()}-normalized": dummy_token_output,
|
||||
f"{PoolingType.MEAN.lower()}": dummy_token_output,
|
||||
f"{PoolingType.CLS.lower()}-normalized": dummy_token_output,
|
||||
f"{PoolingType.CLS.lower()}": dummy_token_output,
|
||||
f"{PoolingType.LAST_TOKEN.lower()}-normalized": dummy_token_output,
|
||||
f"{PoolingType.LAST_TOKEN.lower()}": dummy_token_output,
|
||||
f"{PoolingType.DISABLED.lower()}-normalized": dummy_pooled_output,
|
||||
f"{PoolingType.DISABLED.lower()}": dummy_pooled_output,
|
||||
}
|
||||
|
||||
expected_output = {
|
||||
f"{PoolingType.MEAN.lower()}-normalized": normalize(
|
||||
mean_pooling(dummy_token_embedding, dummy_attention_mask)
|
||||
),
|
||||
f"{PoolingType.MEAN.lower()}": mean_pooling(dummy_token_embedding, dummy_attention_mask),
|
||||
f"{PoolingType.CLS.lower()}-normalized": normalize(dummy_token_embedding[:, 0]),
|
||||
f"{PoolingType.CLS.lower()}": dummy_token_embedding[:, 0],
|
||||
f"{PoolingType.LAST_TOKEN.lower()}-normalized": normalize(
|
||||
last_token_pooling(dummy_token_embedding, dummy_attention_mask)
|
||||
),
|
||||
f"{PoolingType.LAST_TOKEN.lower()}": last_token_pooling(
|
||||
dummy_token_embedding, dummy_attention_mask
|
||||
),
|
||||
f"{PoolingType.DISABLED.lower()}-normalized": normalize(dummy_pooled_embedding),
|
||||
f"{PoolingType.DISABLED.lower()}": dummy_pooled_embedding,
|
||||
}
|
||||
|
||||
for pooling, normalization in itertools.product(
|
||||
(PoolingType.MEAN, PoolingType.CLS, PoolingType.LAST_TOKEN, PoolingType.DISABLED),
|
||||
(True, False),
|
||||
):
|
||||
model_name = f"{pooling.name.lower()}{'-normalized' if normalization else ''}"
|
||||
TextEmbedding.add_custom_model(
|
||||
model_name,
|
||||
pooling=pooling,
|
||||
normalization=normalization,
|
||||
sources=source,
|
||||
dim=dim,
|
||||
size_in_gb=size_in_gb,
|
||||
)
|
||||
|
||||
custom_text_embedding = CustomTextEmbedding(
|
||||
model_name,
|
||||
lazy_load=True,
|
||||
specific_model_path="./", # disable model downloading and loading
|
||||
)
|
||||
|
||||
post_processed_output = next(
|
||||
iter(custom_text_embedding._post_process_onnx_output(input_data[model_name]))
|
||||
)
|
||||
assert np.allclose(post_processed_output, expected_output[model_name], atol=1e-3)
|
||||
|
||||
|
||||
def test_custom_text_model_lookup_is_case_insensitive():
|
||||
model_name = "Org/Model"
|
||||
|
||||
TextEmbedding.add_custom_model(
|
||||
model_name,
|
||||
pooling=PoolingType.MEAN,
|
||||
normalization=True,
|
||||
sources=ModelSource(hf="artificial"),
|
||||
dim=5,
|
||||
size_in_gb=0.1,
|
||||
)
|
||||
|
||||
model = TextEmbedding("org/model", lazy_load=True, specific_model_path="./")
|
||||
|
||||
assert isinstance(model.model, CustomTextEmbedding)
|
||||
assert model.model._pooling == PoolingType.MEAN
|
||||
assert model.model._normalization is True
|
||||
|
||||
|
||||
def test_do_not_add_existing_model():
|
||||
existing_base_model = "sentence-transformers/all-MiniLM-L6-v2"
|
||||
custom_model_name = "intfloat/multilingual-e5-small"
|
||||
|
||||
with pytest.raises(ValueError, match=f"Model {existing_base_model} is already registered"):
|
||||
TextEmbedding.add_custom_model(
|
||||
existing_base_model,
|
||||
pooling=PoolingType.MEAN,
|
||||
normalization=True,
|
||||
sources=ModelSource(hf=existing_base_model),
|
||||
dim=384,
|
||||
size_in_gb=0.47,
|
||||
)
|
||||
|
||||
TextEmbedding.add_custom_model(
|
||||
custom_model_name,
|
||||
pooling=PoolingType.MEAN,
|
||||
normalization=False,
|
||||
sources=ModelSource(hf=existing_base_model),
|
||||
dim=384,
|
||||
size_in_gb=0.47,
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match=f"Model {custom_model_name} is already registered"):
|
||||
TextEmbedding.add_custom_model(
|
||||
custom_model_name,
|
||||
pooling=PoolingType.MEAN,
|
||||
normalization=True,
|
||||
sources=ModelSource(hf=custom_model_name),
|
||||
dim=384,
|
||||
size_in_gb=0.47,
|
||||
)
|
||||
|
||||
|
||||
def test_do_not_add_existing_cross_encoder():
|
||||
existing_base_model = "Xenova/ms-marco-MiniLM-L-6-v2"
|
||||
custom_model_name = "Xenova/ms-marco-MiniLM-L-4-v2"
|
||||
|
||||
with pytest.raises(ValueError, match=f"Model {existing_base_model} is already registered"):
|
||||
TextCrossEncoder.add_custom_model(
|
||||
existing_base_model,
|
||||
sources=ModelSource(hf=existing_base_model),
|
||||
size_in_gb=0.08,
|
||||
)
|
||||
|
||||
TextCrossEncoder.add_custom_model(
|
||||
custom_model_name,
|
||||
sources=ModelSource(hf=existing_base_model),
|
||||
size_in_gb=0.08,
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match=f"Model {custom_model_name} is already registered"):
|
||||
TextCrossEncoder.add_custom_model(
|
||||
custom_model_name,
|
||||
sources=ModelSource(hf=custom_model_name),
|
||||
size_in_gb=0.08,
|
||||
)
|
||||
@@ -1,4 +1,6 @@
|
||||
import os
|
||||
import platform
|
||||
from contextlib import contextmanager
|
||||
from io import BytesIO
|
||||
|
||||
import numpy as np
|
||||
@@ -8,7 +10,7 @@ from PIL import Image
|
||||
|
||||
from fastembed import ImageEmbedding
|
||||
from tests.config import TEST_MISC_DIR
|
||||
from tests.utils import delete_model_cache
|
||||
from tests.utils import delete_model_cache, should_test_model
|
||||
|
||||
CANONICAL_VECTOR_VALUES = {
|
||||
"Qdrant/clip-ViT-B-32-vision": np.array([-0.0098, 0.0128, -0.0274, 0.002, -0.0059]),
|
||||
@@ -24,88 +26,124 @@ CANONICAL_VECTOR_VALUES = {
|
||||
"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]
|
||||
),
|
||||
"nomic-ai/nomic-embed-vision-v1.5": np.array(
|
||||
[0.0048, -0.0254, 0.0067, -0.0296, -0.0435, -0.0123, 0.0024, -0.0361, -0.0703, -0.0186]
|
||||
),
|
||||
"nomic-ai/nomic-embed-vision-v1.5-Q": np.array(
|
||||
[-0.0011, -0.0477, 0.0024, -0.049, -0.0458, -0.0314, 0.017, -0.0383, -0.0537, -0.021]
|
||||
),
|
||||
"google/siglip2-base-patch16-224": np.array(
|
||||
[-0.02095927, -0.0075177, -0.00144479, -0.0080948, 0.05031789]
|
||||
),
|
||||
}
|
||||
|
||||
_MODELS_TO_CACHE = ("Qdrant/clip-ViT-B-32-vision",)
|
||||
MODELS_TO_CACHE = tuple([x.lower() for x in _MODELS_TO_CACHE])
|
||||
|
||||
def test_embedding() -> None:
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def model_cache():
|
||||
is_ci = os.getenv("CI")
|
||||
cache = {}
|
||||
|
||||
@contextmanager
|
||||
def get_model(model_name: str):
|
||||
lowercase_model_name = model_name.lower()
|
||||
if lowercase_model_name not in cache:
|
||||
cache[lowercase_model_name] = ImageEmbedding(lowercase_model_name)
|
||||
yield cache[lowercase_model_name]
|
||||
if lowercase_model_name not in MODELS_TO_CACHE:
|
||||
model_inst = cache.pop(lowercase_model_name)
|
||||
if is_ci:
|
||||
delete_model_cache(model_inst.model._model_dir)
|
||||
del model_inst
|
||||
|
||||
yield get_model
|
||||
|
||||
if is_ci:
|
||||
for name, model in cache.items():
|
||||
delete_model_cache(model.model._model_dir)
|
||||
cache.clear()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_name", ["Qdrant/clip-ViT-B-32-vision"])
|
||||
def test_embedding(model_cache, model_name: str) -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
is_mac = platform.system() == "Darwin"
|
||||
is_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
|
||||
|
||||
for model_desc in ImageEmbedding._list_supported_models():
|
||||
if not is_ci and model_desc.size_in_GB > 1:
|
||||
# quantized int8 ops diverge on macOS; canonical vector is generated on linux/amd64 (CI)
|
||||
if is_mac and model_desc.model == "nomic-ai/nomic-embed-vision-v1.5-Q":
|
||||
continue
|
||||
if not should_test_model(model_desc, model_name, is_ci, is_manual):
|
||||
continue
|
||||
|
||||
dim = model_desc.dim
|
||||
|
||||
model = ImageEmbedding(model_name=model_desc.model)
|
||||
with model_cache(model_desc.model) as model:
|
||||
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 == (len(images), dim)
|
||||
|
||||
images = [
|
||||
canonical_vector = CANONICAL_VECTOR_VALUES[model_desc.model]
|
||||
|
||||
assert np.allclose(
|
||||
embeddings[0, : canonical_vector.shape[0]], canonical_vector, atol=1e-3
|
||||
), model_desc.model
|
||||
|
||||
assert np.allclose(embeddings[1], embeddings[2]), model_desc.model
|
||||
|
||||
|
||||
@pytest.mark.parametrize("n_dims,model_name", [(512, "Qdrant/clip-ViT-B-32-vision")])
|
||||
def test_batch_embedding(model_cache, n_dims: int, model_name: str) -> None:
|
||||
with model_cache(model_name) as model:
|
||||
n_images = 32
|
||||
test_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)),
|
||||
Image.open(TEST_MISC_DIR / "small_image.jpeg"),
|
||||
]
|
||||
embeddings = list(model.embed(images))
|
||||
images = test_images * n_images
|
||||
|
||||
embeddings = list(model.embed(images, batch_size=10))
|
||||
embeddings = np.stack(embeddings, axis=0)
|
||||
assert embeddings.shape == (len(images), dim)
|
||||
assert np.allclose(embeddings[1], embeddings[2])
|
||||
|
||||
canonical_vector = CANONICAL_VECTOR_VALUES[model_desc.model]
|
||||
canonical_vector = CANONICAL_VECTOR_VALUES[model_name]
|
||||
|
||||
assert np.allclose(
|
||||
embeddings[0, : canonical_vector.shape[0]], canonical_vector, atol=1e-3
|
||||
), model_desc.model
|
||||
|
||||
assert np.allclose(embeddings[1], embeddings[2]), model_desc.model
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
assert embeddings.shape == (len(test_images) * n_images, n_dims)
|
||||
assert np.allclose(embeddings[0, : canonical_vector.shape[0]], canonical_vector, atol=1e-3)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("n_dims,model_name", [(512, "Qdrant/clip-ViT-B-32-vision")])
|
||||
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
|
||||
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
|
||||
def test_parallel_processing(model_cache, n_dims: int, model_name: str) -> None:
|
||||
with model_cache(model_name) as model:
|
||||
n_images = 32
|
||||
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)
|
||||
|
||||
embeddings = list(model.embed(images, batch_size=10))
|
||||
embeddings = np.stack(embeddings, axis=0)
|
||||
embeddings_2 = list(model.embed(images, batch_size=10, parallel=None))
|
||||
embeddings_2 = np.stack(embeddings_2, axis=0)
|
||||
|
||||
assert embeddings.shape == (len(test_images) * n_images, n_dims)
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
embeddings_3 = list(model.embed(images, batch_size=10, parallel=0))
|
||||
embeddings_3 = np.stack(embeddings_3, axis=0)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("n_dims,model_name", [(512, "Qdrant/clip-ViT-B-32-vision")])
|
||||
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
|
||||
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)
|
||||
|
||||
embeddings_2 = list(model.embed(images, batch_size=10, parallel=None))
|
||||
embeddings_2 = np.stack(embeddings_2, axis=0)
|
||||
|
||||
embeddings_3 = list(model.embed(images, batch_size=10, parallel=0))
|
||||
embeddings_3 = np.stack(embeddings_3, axis=0)
|
||||
|
||||
assert embeddings.shape == (n_images * 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)
|
||||
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)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_name", ["Qdrant/clip-ViT-B-32-vision"])
|
||||
@@ -121,3 +159,31 @@ def test_lazy_load(model_name: str) -> None:
|
||||
assert hasattr(model.model, "model")
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
def test_get_embedding_size() -> None:
|
||||
assert ImageEmbedding.get_embedding_size(model_name="Qdrant/clip-ViT-B-32-vision") == 512
|
||||
assert ImageEmbedding.get_embedding_size(model_name="Qdrant/clip-vit-b-32-vision") == 512
|
||||
|
||||
|
||||
def test_embedding_size() -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
model_name = "Qdrant/clip-ViT-B-32-vision"
|
||||
model = ImageEmbedding(model_name=model_name, lazy_load=True)
|
||||
assert model.embedding_size == 512
|
||||
|
||||
model_name = "Qdrant/clip-vit-b-32-vision"
|
||||
model = ImageEmbedding(model_name=model_name, lazy_load=True)
|
||||
assert model.embedding_size == 512
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_name", ["Qdrant/clip-ViT-B-32-vision"])
|
||||
def test_session_options(model_cache, model_name) -> None:
|
||||
with model_cache(model_name) as default_model:
|
||||
default_session_options = default_model.model.model.get_session_options()
|
||||
assert default_session_options.enable_cpu_mem_arena is True
|
||||
model = ImageEmbedding(model_name=model_name, enable_cpu_mem_arena=False)
|
||||
session_options = model.model.model.get_session_options()
|
||||
assert session_options.enable_cpu_mem_arena is False
|
||||
|
||||
@@ -0,0 +1,90 @@
|
||||
import numpy as np
|
||||
import pytest
|
||||
from PIL import Image
|
||||
|
||||
from fastembed.image.transform.functional import normalize, resize
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("size", "expected"),
|
||||
[
|
||||
((100, 200), (200, 100)), # the bug: a non-square size came back transposed
|
||||
((224, 224), (224, 224)), # the square path every shipped model takes
|
||||
],
|
||||
)
|
||||
def test_resize_tuple_is_height_width(size: tuple[int, int], expected: tuple[int, int]) -> None:
|
||||
"""A ``(height, width)`` size must reach Pillow as ``(width, height)``."""
|
||||
resized = resize(Image.new("RGB", (300, 300)), size=size)
|
||||
|
||||
assert resized.size == expected # PIL reports (width, height)
|
||||
|
||||
|
||||
def test_resize_int_keeps_shortest_edge_behaviour() -> None:
|
||||
"""The int branch already emitted Pillow order; it must not be disturbed."""
|
||||
landscape = Image.new("RGB", (400, 200))
|
||||
portrait = Image.new("RGB", (200, 400))
|
||||
|
||||
# size sets the shortest edge, and the aspect ratio is preserved.
|
||||
assert resize(landscape, size=100).size == (200, 100)
|
||||
assert resize(portrait, size=100).size == (100, 200)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("mean", "std"),
|
||||
[
|
||||
([0.1, 0.2, 0.3], [0.5, 0.6, 0.7]), # per-channel, as every model config gives it
|
||||
(0.5, 0.25), # scalar, expanded to one value per channel
|
||||
],
|
||||
)
|
||||
def test_normalize_chw_is_channel_wise(
|
||||
mean: list[float] | float, std: list[float] | float
|
||||
) -> None:
|
||||
"""Each channel must be normalized by its own mean/std, not by any other axis."""
|
||||
rng = np.random.default_rng(0)
|
||||
image = rng.random((3, 5, 7)).astype(np.float32)
|
||||
means = mean if isinstance(mean, list) else [mean] * 3
|
||||
stds = std if isinstance(std, list) else [std] * 3
|
||||
|
||||
result = normalize(image, mean=mean, std=std)
|
||||
|
||||
for c in range(3):
|
||||
assert np.allclose(result[c], (image[c] - means[c]) / stds[c], atol=1e-6)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("batch_size", [2, 3])
|
||||
def test_normalize_batched_matches_per_image(batch_size: int) -> None:
|
||||
"""A batch must give exactly what the (C, H, W) path gives image by image.
|
||||
|
||||
batch_size 2 used to raise, since transposing reversed every axis; batch_size 3
|
||||
matched the channel count and silently normalized along the batch axis instead.
|
||||
"""
|
||||
rng = np.random.default_rng(2)
|
||||
batch = rng.random((batch_size, 3, 4, 4)).astype(np.float32)
|
||||
mean, std = [0.1, 0.2, 0.3], [0.5, 0.6, 0.7]
|
||||
|
||||
result = normalize(batch, mean=mean, std=std)
|
||||
|
||||
per_image = np.stack([normalize(image, mean=mean, std=std) for image in batch])
|
||||
assert result.shape == batch.shape
|
||||
assert np.array_equal(result, per_image)
|
||||
|
||||
|
||||
def test_normalize_rejects_input_without_a_channel_axis() -> None:
|
||||
"""Every pipeline runs ConvertToRGB first, so normalize only ever sees (C, H, W)."""
|
||||
with pytest.raises(ValueError, match=r"must be \(C, H, W\)"):
|
||||
normalize(np.zeros((4, 6), dtype=np.float32), mean=0.5, std=0.25)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("mean", "std", "expected"),
|
||||
[
|
||||
([0.1, 0.2], [1.0, 1.0, 1.0], "mean must"),
|
||||
([0.1, 0.2, 0.3], [1.0, 1.0], "std must"),
|
||||
],
|
||||
)
|
||||
def test_normalize_channel_count_mismatch_raises(
|
||||
mean: list[float], std: list[float], expected: str
|
||||
) -> None:
|
||||
image = np.zeros((3, 4, 4), dtype=np.float32)
|
||||
with pytest.raises(ValueError, match=expected):
|
||||
normalize(image, mean=mean, std=std)
|
||||
@@ -1,4 +1,5 @@
|
||||
import os
|
||||
from contextlib import contextmanager
|
||||
|
||||
import pytest
|
||||
import numpy as np
|
||||
@@ -6,7 +7,7 @@ import numpy as np
|
||||
from fastembed.late_interaction.late_interaction_text_embedding import (
|
||||
LateInteractionTextEmbedding,
|
||||
)
|
||||
from tests.utils import delete_model_cache
|
||||
from tests.utils import delete_model_cache, should_test_model
|
||||
|
||||
# vectors are abridged and rounded for brevity
|
||||
CANONICAL_COLUMN_VALUES = {
|
||||
@@ -150,82 +151,124 @@ CANONICAL_QUERY_VALUES = {
|
||||
),
|
||||
}
|
||||
|
||||
_MODELS_TO_CACHE = ("answerdotai/answerai-colbert-small-v1",)
|
||||
MODELS_TO_CACHE = tuple([x.lower() for x in _MODELS_TO_CACHE])
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def model_cache():
|
||||
is_ci = os.getenv("CI")
|
||||
cache = {}
|
||||
|
||||
@contextmanager
|
||||
def get_model(model_name: str):
|
||||
lowercase_model_name = model_name.lower()
|
||||
if lowercase_model_name not in cache:
|
||||
cache[lowercase_model_name] = LateInteractionTextEmbedding(lowercase_model_name)
|
||||
yield cache[lowercase_model_name]
|
||||
if lowercase_model_name not in MODELS_TO_CACHE:
|
||||
model_inst = cache.pop(lowercase_model_name)
|
||||
if is_ci:
|
||||
delete_model_cache(model_inst.model._model_dir)
|
||||
del model_inst
|
||||
|
||||
yield get_model
|
||||
|
||||
if is_ci:
|
||||
for name, model in cache.items():
|
||||
delete_model_cache(model.model._model_dir)
|
||||
cache.clear()
|
||||
|
||||
|
||||
docs = ["Hello World"]
|
||||
|
||||
|
||||
def test_batch_embedding():
|
||||
is_ci = os.getenv("CI")
|
||||
@pytest.mark.parametrize("model_name", ["answerdotai/answerai-colbert-small-v1"])
|
||||
def test_batch_embedding(model_cache, model_name: str):
|
||||
docs_to_embed = docs * 10
|
||||
|
||||
for model_name, expected_result in CANONICAL_COLUMN_VALUES.items():
|
||||
print("evaluating", model_name)
|
||||
model = LateInteractionTextEmbedding(model_name=model_name)
|
||||
with model_cache(model_name) as model:
|
||||
result = list(model.embed(docs_to_embed, batch_size=6))
|
||||
expected_result = CANONICAL_COLUMN_VALUES[model_name]
|
||||
|
||||
for value in result:
|
||||
token_num, abridged_dim = expected_result.shape
|
||||
assert np.allclose(value[:, :abridged_dim], expected_result, atol=2e-3)
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
@pytest.mark.parametrize("model_name", ["answerdotai/answerai-colbert-small-v1"])
|
||||
def test_batch_inference_size_same_as_single_inference(model_cache, model_name: str):
|
||||
with model_cache(model_name) as model:
|
||||
docs_to_embed = [
|
||||
"short document",
|
||||
"A bit longer document, which should not affect the size",
|
||||
]
|
||||
result = list(model.embed(docs_to_embed, batch_size=1))
|
||||
result_2 = list(model.embed(docs_to_embed, batch_size=2))
|
||||
assert len(result[0]) == len(result_2[0])
|
||||
|
||||
|
||||
def test_single_embedding():
|
||||
@pytest.mark.parametrize("model_name", ["answerdotai/answerai-colbert-small-v1"])
|
||||
def test_single_embedding(model_cache, model_name: str):
|
||||
is_ci = os.getenv("CI")
|
||||
is_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
|
||||
docs_to_embed = docs
|
||||
|
||||
for model_name, expected_result in CANONICAL_COLUMN_VALUES.items():
|
||||
for model_desc in LateInteractionTextEmbedding._list_supported_models():
|
||||
if not should_test_model(model_desc, model_name, is_ci, is_manual):
|
||||
continue
|
||||
|
||||
print("evaluating", model_name)
|
||||
model = LateInteractionTextEmbedding(model_name=model_name)
|
||||
result = next(iter(model.embed(docs_to_embed, batch_size=6)))
|
||||
token_num, abridged_dim = expected_result.shape
|
||||
assert np.allclose(result[:, :abridged_dim], expected_result, atol=2e-3)
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
with model_cache(model_desc.model) as model:
|
||||
whole_result = list(model.embed(docs_to_embed, batch_size=6))
|
||||
assert len(whole_result) == 1
|
||||
result = whole_result[0]
|
||||
expected_result = CANONICAL_COLUMN_VALUES[model_desc.model]
|
||||
token_num, abridged_dim = expected_result.shape
|
||||
assert np.allclose(result[:, :abridged_dim], expected_result, atol=2e-3)
|
||||
|
||||
|
||||
def test_single_embedding_query():
|
||||
@pytest.mark.parametrize("model_name", ["answerdotai/answerai-colbert-small-v1"])
|
||||
def test_single_embedding_query(model_cache, model_name: str):
|
||||
is_ci = os.getenv("CI")
|
||||
is_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
|
||||
queries_to_embed = docs
|
||||
|
||||
for model_name, expected_result in CANONICAL_QUERY_VALUES.items():
|
||||
print("evaluating", model_name)
|
||||
model = LateInteractionTextEmbedding(model_name=model_name)
|
||||
result = next(iter(model.query_embed(queries_to_embed)))
|
||||
token_num, abridged_dim = expected_result.shape
|
||||
assert np.allclose(result[:, :abridged_dim], expected_result, atol=2e-3)
|
||||
for model_desc in LateInteractionTextEmbedding._list_supported_models():
|
||||
if not should_test_model(model_desc, model_name, is_ci, is_manual):
|
||||
continue
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
print("evaluating", model_desc.model)
|
||||
with model_cache(model_desc.model) as model:
|
||||
whole_result = list(model.query_embed(queries_to_embed))
|
||||
assert len(whole_result) == 1
|
||||
result = whole_result[0]
|
||||
expected_result = CANONICAL_QUERY_VALUES[model_desc.model]
|
||||
token_num, abridged_dim = expected_result.shape
|
||||
assert np.allclose(result[:, :abridged_dim], expected_result, atol=2e-3)
|
||||
|
||||
|
||||
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
|
||||
embeddings = list(model.embed(docs, batch_size=10, parallel=2))
|
||||
embeddings = np.stack(embeddings, axis=0)
|
||||
@pytest.mark.parametrize("token_dim,model_name", [(96, "answerdotai/answerai-colbert-small-v1")])
|
||||
def test_parallel_processing(model_cache, token_dim: int, model_name: str):
|
||||
with model_cache(model_name) as model:
|
||||
docs = ["hello world", "flag embedding"] * 100
|
||||
embeddings = list(model.embed(docs, batch_size=10, parallel=2))
|
||||
|
||||
embeddings_2 = list(model.embed(docs, batch_size=10, parallel=None))
|
||||
embeddings_2 = np.stack(embeddings_2, axis=0)
|
||||
embeddings_2 = list(model.embed(docs, batch_size=10, parallel=None))
|
||||
|
||||
embeddings_3 = list(model.embed(docs, batch_size=10, parallel=0))
|
||||
embeddings_3 = np.stack(embeddings_3, axis=0)
|
||||
# embeddings_3 = list(model.embed(docs, batch_size=10, parallel=0)) # inherits OnnxTextModel which
|
||||
# # is tested in TextEmbedding, disabling it here to reduce number of requests to hf
|
||||
# # multiprocessing is enough to test with `parallel=2`, and `parallel=None` is okay to tests since it reuses
|
||||
# # model from cache
|
||||
|
||||
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)
|
||||
assert len(embeddings) == len(docs) and embeddings[0].shape[-1] == token_dim
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
for i in range(len(embeddings)):
|
||||
assert np.allclose(embeddings[i], embeddings_2[i], atol=1e-3)
|
||||
# assert np.allclose(embeddings[i], embeddings_3[i], atol=1e-3)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_name",
|
||||
["colbert-ir/colbertv2.0"],
|
||||
)
|
||||
@pytest.mark.parametrize("model_name", ["answerdotai/answerai-colbert-small-v1"])
|
||||
def test_lazy_load(model_name: str):
|
||||
is_ci = os.getenv("CI")
|
||||
|
||||
@@ -244,3 +287,58 @@ def test_lazy_load(model_name: str):
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
def test_get_embedding_size():
|
||||
model_name = "answerdotai/answerai-colbert-small-v1"
|
||||
assert LateInteractionTextEmbedding.get_embedding_size(model_name) == 96
|
||||
|
||||
model_name = "answerdotai/answerai-ColBERT-small-v1"
|
||||
assert LateInteractionTextEmbedding.get_embedding_size(model_name) == 96
|
||||
|
||||
|
||||
def test_embedding_size():
|
||||
is_ci = os.getenv("CI")
|
||||
model_name = "answerdotai/answerai-colbert-small-v1"
|
||||
model = LateInteractionTextEmbedding(model_name=model_name, lazy_load=True)
|
||||
assert model.embedding_size == 96
|
||||
|
||||
model_name = "answerdotai/answerai-ColBERT-small-v1"
|
||||
model = LateInteractionTextEmbedding(model_name=model_name, lazy_load=True)
|
||||
assert model.embedding_size == 96
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_name", ["answerdotai/answerai-ColBERT-small-v1"])
|
||||
def test_session_options(model_cache, model_name) -> None:
|
||||
with model_cache(model_name) as default_model:
|
||||
default_session_options = default_model.model.model.get_session_options()
|
||||
assert default_session_options.enable_cpu_mem_arena is True
|
||||
model = LateInteractionTextEmbedding(model_name=model_name, enable_cpu_mem_arena=False)
|
||||
session_options = model.model.model.get_session_options()
|
||||
assert session_options.enable_cpu_mem_arena is False
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_name", ["answerdotai/answerai-colbert-small-v1"])
|
||||
def test_token_count(model_cache, model_name) -> None:
|
||||
with model_cache(model_name) as model:
|
||||
documents = ["short doc", "it is a long document to check attention mask for paddings"]
|
||||
short_doc_token_count = model.token_count(documents[0])
|
||||
long_doc_token_count = model.token_count(documents[1])
|
||||
documents_token_count = model.token_count(documents)
|
||||
assert short_doc_token_count + long_doc_token_count == documents_token_count
|
||||
# 2 is 2*DOC_MARKER_TOKEN_ID for each document
|
||||
assert short_doc_token_count + long_doc_token_count + 2 == model.token_count(
|
||||
documents, include_extension=True
|
||||
)
|
||||
assert short_doc_token_count + long_doc_token_count == model.token_count(
|
||||
documents, batch_size=1
|
||||
)
|
||||
assert short_doc_token_count + long_doc_token_count == model.token_count(
|
||||
documents, is_doc=False
|
||||
)
|
||||
# query min length is 32
|
||||
assert model.token_count(documents, is_doc=False, include_extension=True) == 64
|
||||
very_long_query = "It's a very long query which definitely contains more than 32 tokens and we're using it to check whether the method can handle large query properly without cutting it to 32 tokens"
|
||||
assert model.token_count(very_long_query, is_doc=False, include_extension=True) > 32
|
||||
|
||||
@@ -1,25 +1,36 @@
|
||||
import os
|
||||
from contextlib import contextmanager
|
||||
|
||||
import pytest
|
||||
from PIL import Image
|
||||
import numpy as np
|
||||
|
||||
from fastembed import LateInteractionMultimodalEmbedding
|
||||
from tests.config import TEST_MISC_DIR
|
||||
|
||||
from tests.utils import delete_model_cache
|
||||
|
||||
# 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],
|
||||
]
|
||||
[-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],
|
||||
]
|
||||
),
|
||||
"Qdrant/colmodernvbert": np.array(
|
||||
[
|
||||
[0.11614, -0.15793, -0.11194, 0.0688, 0.08001, 0.10575, -0.07871],
|
||||
[0.10094, -0.13301, -0.12069, 0.10932, 0.04645, 0.09884, 0.04048],
|
||||
[0.13106, -0.18613, -0.13469, 0.10566, 0.03659, 0.07712, -0.03916],
|
||||
[0.09754, -0.09596, -0.04839, 0.14991, 0.05692, 0.10569, -0.08349],
|
||||
[0.02576, -0.15651, -0.09977, 0.09707, 0.13412, 0.09994, -0.09931],
|
||||
[-0.06741, -0.1787, -0.19677, -0.07618, 0.13102, -0.02131, -0.02437],
|
||||
[-0.02776, -0.10187, -0.13793, 0.03835, 0.04766, 0.04701, -0.15635],
|
||||
]
|
||||
),
|
||||
}
|
||||
@@ -36,6 +47,17 @@ CANONICAL_QUERY_VALUES = {
|
||||
[-0.0165, -0.0106, 0.1672, -0.0768, 0.0389, -0.0038, 0.1137],
|
||||
]
|
||||
),
|
||||
"Qdrant/colmodernvbert": np.array(
|
||||
[
|
||||
[0.05, 0.06557, 0.04026, 0.14981, 0.1842, 0.0263, -0.18706],
|
||||
[-0.05664, -0.14028, 0.00649, -0.02849, 0.09034, -0.01494, 0.10693],
|
||||
[-0.10147, -0.00716, 0.09084, -0.08236, -0.01849, -0.00972, -0.00461],
|
||||
[-0.1233, -0.10814, -0.02337, -0.00329, 0.05984, 0.09934, 0.09846],
|
||||
[-0.07053, -0.13119, -0.06487, 0.01508, 0.07459, 0.07655, 0.14821],
|
||||
[0.00526, -0.13842, -0.05837, -0.02721, 0.13009, 0.05076, 0.17962],
|
||||
[0.00924, -0.14383, -0.03057, -0.03691, 0.11718, 0.037, 0.13344],
|
||||
]
|
||||
),
|
||||
}
|
||||
|
||||
queries = ["hello world", "flag embedding"]
|
||||
@@ -45,40 +67,99 @@ images = [
|
||||
Image.open((TEST_MISC_DIR / "image.jpeg")),
|
||||
]
|
||||
|
||||
_MODELS_TO_CACHE = ("Qdrant/colmodernvbert",)
|
||||
MODELS_TO_CACHE = tuple(model_name.lower() for model_name in _MODELS_TO_CACHE)
|
||||
|
||||
def test_batch_embedding():
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def model_cache():
|
||||
is_ci = os.getenv("CI")
|
||||
cache = {}
|
||||
|
||||
if not is_ci:
|
||||
for model_name, expected_result in CANONICAL_IMAGE_VALUES.items():
|
||||
print("evaluating", model_name)
|
||||
model = LateInteractionMultimodalEmbedding(model_name=model_name)
|
||||
@contextmanager
|
||||
def get_model(model_name: str):
|
||||
lowercase_model_name = model_name.lower()
|
||||
if lowercase_model_name not in cache:
|
||||
cache[lowercase_model_name] = LateInteractionMultimodalEmbedding(lowercase_model_name)
|
||||
yield cache[lowercase_model_name]
|
||||
if lowercase_model_name not in MODELS_TO_CACHE:
|
||||
model_inst = cache.pop(lowercase_model_name)
|
||||
if is_ci:
|
||||
delete_model_cache(model_inst.model._model_dir)
|
||||
del model_inst
|
||||
|
||||
yield get_model
|
||||
|
||||
if is_ci:
|
||||
for _, model in cache.items():
|
||||
delete_model_cache(model.model._model_dir)
|
||||
cache.clear()
|
||||
|
||||
|
||||
def test_batch_embedding(model_cache):
|
||||
for model_name, expected_result in CANONICAL_IMAGE_VALUES.items():
|
||||
if model_name.lower() == "Qdrant/colpali-v1.3-fp16".lower() and os.getenv("CI"):
|
||||
continue # colpali is too large for ci
|
||||
|
||||
print("evaluating", model_name)
|
||||
with model_cache(model_name) as model:
|
||||
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)
|
||||
token_num, abridged_dim = expected_result.shape
|
||||
assert np.allclose(value[:token_num, :abridged_dim], expected_result, atol=2e-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)
|
||||
def test_single_embedding(model_cache):
|
||||
for model_name, expected_result in CANONICAL_IMAGE_VALUES.items():
|
||||
if model_name.lower() == "Qdrant/colpali-v1.3-fp16".lower() and os.getenv("CI"):
|
||||
continue # colpali is too large for ci
|
||||
print("evaluating", model_name)
|
||||
with model_cache(model_name) as model:
|
||||
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)
|
||||
|
||||
|
||||
def test_single_embedding_query(model_cache):
|
||||
for model_name, expected_result in CANONICAL_QUERY_VALUES.items():
|
||||
if model_name.lower() == "Qdrant/colpali-v1.3-fp16".lower() and os.getenv("CI"):
|
||||
continue # colpali is too large for ci
|
||||
print("evaluating", model_name)
|
||||
with model_cache(model_name) as model:
|
||||
result = next(iter(model.embed_text(queries)))
|
||||
token_num, abridged_dim = expected_result.shape
|
||||
assert np.allclose(result[:token_num, :abridged_dim], expected_result, atol=2e-3)
|
||||
|
||||
|
||||
def test_get_embedding_size():
|
||||
model_name = "Qdrant/colpali-v1.3-fp16"
|
||||
assert LateInteractionMultimodalEmbedding.get_embedding_size(model_name) == 128
|
||||
|
||||
model_name = "Qdrant/ColPali-v1.3-fp16"
|
||||
assert LateInteractionMultimodalEmbedding.get_embedding_size(model_name) == 128
|
||||
|
||||
model_name = "Qdrant/colmodernvbert"
|
||||
assert LateInteractionMultimodalEmbedding.get_embedding_size(model_name) == 128
|
||||
|
||||
|
||||
def test_embedding_size():
|
||||
model_name = "Qdrant/colmodernvbert"
|
||||
model = LateInteractionMultimodalEmbedding(model_name=model_name, lazy_load=True)
|
||||
assert model.embedding_size == 128
|
||||
|
||||
|
||||
def test_token_count(model_cache) -> None:
|
||||
model_name = "Qdrant/colmodernvbert"
|
||||
with model_cache(model_name) as model:
|
||||
documents = ["short doc", "it is a long document to check attention mask for paddings"]
|
||||
short_doc_token_count = model.token_count(documents[0])
|
||||
long_doc_token_count = model.token_count(documents[1])
|
||||
documents_token_count = model.token_count(documents)
|
||||
assert short_doc_token_count + long_doc_token_count == documents_token_count
|
||||
assert short_doc_token_count + long_doc_token_count == model.token_count(
|
||||
documents, batch_size=1
|
||||
)
|
||||
assert short_doc_token_count + long_doc_token_count < model.token_count(
|
||||
documents, include_extension=True
|
||||
)
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import pytest
|
||||
from typing import Optional
|
||||
|
||||
from fastembed import (
|
||||
TextEmbedding,
|
||||
SparseTextEmbedding,
|
||||
@@ -14,7 +14,7 @@ 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:
|
||||
def test_gpu_via_providers(device_id: int | None) -> None:
|
||||
docs = ["hello world", "flag embedding"]
|
||||
|
||||
device_id = device_id if device_id is not None else 0
|
||||
@@ -86,7 +86,7 @@ def test_gpu_via_providers(device_id: Optional[int]) -> None:
|
||||
|
||||
@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:
|
||||
def test_gpu_cuda_device_ids(device_ids: list[int] | None) -> None:
|
||||
docs = ["hello world", "flag embedding"]
|
||||
device_id = device_ids[0] if device_ids else 0
|
||||
embedding_model = TextEmbedding(
|
||||
@@ -171,7 +171,7 @@ def test_gpu_cuda_device_ids(device_ids: Optional[list[int]]) -> None:
|
||||
@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:
|
||||
def test_multi_gpu_parallel_inference(device_ids: list[int] | None, parallel: int) -> None:
|
||||
docs = ["hello world", "flag embedding"] * 100
|
||||
batch_size = 5
|
||||
|
||||
|
||||
@@ -0,0 +1,38 @@
|
||||
import numpy as np
|
||||
|
||||
from fastembed import LateInteractionTextEmbedding
|
||||
from fastembed.postprocess import Muvera
|
||||
|
||||
CANONICAL_VALUES = [-2.61810007e-04, 1.89005750e00, -2.32070747e00]
|
||||
CANONICAL_QUERY_VALUES = [
|
||||
-0.85783903,
|
||||
1.1077204,
|
||||
-0.09522747,
|
||||
] # part of the values are zeros, should be compared with the result of nonzero mask
|
||||
|
||||
DIM = 128
|
||||
K_SIM = 5
|
||||
DIM_PROJ = 16
|
||||
R_REPS = 20
|
||||
|
||||
|
||||
def test_single_input():
|
||||
model = LateInteractionTextEmbedding("colbert-ir/colbertv2.0", lazy_load=True)
|
||||
random_generator = np.random.default_rng(42)
|
||||
multivector = random_generator.random((10, 128))
|
||||
|
||||
for muvera in (
|
||||
Muvera(dim=DIM, k_sim=K_SIM, dim_proj=DIM_PROJ, r_reps=R_REPS, random_seed=42),
|
||||
Muvera.from_multivector_model(model, k_sim=K_SIM, dim_proj=DIM_PROJ, r_reps=R_REPS),
|
||||
):
|
||||
fde = muvera.process(multivector)
|
||||
assert fde.shape[0] == muvera.embedding_size
|
||||
assert np.allclose(fde[:3], CANONICAL_VALUES)
|
||||
|
||||
fde_doc = muvera.process_document(multivector)
|
||||
assert fde_doc.shape[0] == muvera.embedding_size
|
||||
assert np.allclose(fde, fde_doc)
|
||||
|
||||
fde_query = muvera.process_query(multivector)
|
||||
assert fde_query.shape[0] == muvera.embedding_size
|
||||
assert np.allclose(fde_query[np.nonzero(fde_query)][:3], CANONICAL_QUERY_VALUES)
|
||||
@@ -0,0 +1,333 @@
|
||||
import itertools
|
||||
import json
|
||||
import os
|
||||
import shutil
|
||||
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from tokenizers import Tokenizer
|
||||
|
||||
from fastembed.common.preprocessor_utils import load_tokenizer
|
||||
from fastembed.text.text_embedding import TextEmbedding
|
||||
from tests.utils import delete_model_cache
|
||||
|
||||
# transformers writes its VERY_LARGE_INTEGER in place of `model_max_length` when the real
|
||||
# value is unknown, which is more than `enable_truncation` can accept
|
||||
HF_SENTINEL = int(1e30)
|
||||
|
||||
# a lightweight model whose config files serve as a realistic starting point for the cases below
|
||||
BASE_MODEL = "BAAI/bge-small-en-v1.5"
|
||||
TOKENIZER_FILES = (
|
||||
"config.json",
|
||||
"tokenizer.json",
|
||||
"tokenizer_config.json",
|
||||
"special_tokens_map.json",
|
||||
)
|
||||
|
||||
|
||||
def _patch_json(path: Path, overrides: dict[str, Any], drop: tuple[str, ...] = ()) -> None:
|
||||
with open(path) as source:
|
||||
content = json.load(source)
|
||||
|
||||
content.update(overrides)
|
||||
for key in drop:
|
||||
content.pop(key, None)
|
||||
|
||||
with open(path, "w") as target:
|
||||
json.dump(content, target)
|
||||
|
||||
|
||||
def _set_serialized_padding(path: Path, padding: dict[str, Any] | None) -> None:
|
||||
"""Rewrite tokenizer.json through the tokenizers library, so the format stays authoritative."""
|
||||
tokenizer = Tokenizer.from_file(str(path))
|
||||
if padding is None:
|
||||
tokenizer.no_padding()
|
||||
else:
|
||||
tokenizer.enable_padding(**padding)
|
||||
tokenizer.save(str(path))
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def make_model_dir(tmp_path_factory):
|
||||
"""Build model directories from a real model's config files, with targeted overrides.
|
||||
|
||||
`load_tokenizer` reads only the four files in `TOKENIZER_FILES`, so the onnx weights are
|
||||
never copied.
|
||||
"""
|
||||
is_ci = os.getenv("CI")
|
||||
base_model = TextEmbedding(BASE_MODEL)
|
||||
source_dir = Path(base_model.model._model_dir)
|
||||
counter = itertools.count()
|
||||
|
||||
def factory(
|
||||
tokenizer_config: dict[str, Any] | None = None,
|
||||
config: dict[str, Any] | None = None,
|
||||
padding: dict[str, Any] | None = None,
|
||||
drop_from_tokenizer_config: tuple[str, ...] = (),
|
||||
drop_from_config: tuple[str, ...] = (),
|
||||
drop_files: tuple[str, ...] = (),
|
||||
) -> Path:
|
||||
model_dir = tmp_path_factory.mktemp(f"model_dir_{next(counter)}")
|
||||
for file_name in TOKENIZER_FILES:
|
||||
if file_name in drop_files:
|
||||
continue
|
||||
shutil.copy(source_dir / file_name, model_dir / file_name)
|
||||
|
||||
if "tokenizer_config.json" not in drop_files:
|
||||
_patch_json(
|
||||
model_dir / "tokenizer_config.json",
|
||||
tokenizer_config or {},
|
||||
drop_from_tokenizer_config,
|
||||
)
|
||||
if "config.json" not in drop_files:
|
||||
_patch_json(model_dir / "config.json", config or {}, drop_from_config)
|
||||
if padding is not None:
|
||||
_set_serialized_padding(model_dir / "tokenizer.json", padding)
|
||||
|
||||
return model_dir
|
||||
|
||||
yield factory
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(base_model.model._model_dir)
|
||||
|
||||
|
||||
def test_fixed_padding_is_relaxed_to_batch_longest(make_model_dir) -> None:
|
||||
"""Fixed padding shorter than the truncation limit leaves longer encodings ragged."""
|
||||
model_dir = make_model_dir(
|
||||
tokenizer_config={"model_max_length": 512, "max_length": None},
|
||||
padding={"length": 128, "pad_id": 0, "pad_token": "[PAD]", "direction": "right"},
|
||||
)
|
||||
|
||||
tokenizer, _ = load_tokenizer(model_dir)
|
||||
|
||||
assert tokenizer.padding["length"] is None
|
||||
assert tokenizer.truncation["max_length"] == 512
|
||||
|
||||
encoded = tokenizer.encode_batch(["hello world", "retrieval " * 200])
|
||||
# ragged encodings make this raise, the same way onnx_embed does
|
||||
input_ids = np.array([encoding.ids for encoding in encoded])
|
||||
|
||||
assert input_ids.shape == (2, len(encoded[0].ids))
|
||||
assert input_ids.shape[1] > 128
|
||||
|
||||
|
||||
def test_batch_longest_padding_does_not_pad_to_the_truncation_limit(make_model_dir) -> None:
|
||||
model_dir = make_model_dir(
|
||||
tokenizer_config={"model_max_length": 512, "max_length": None},
|
||||
padding={"length": 128, "pad_id": 0, "pad_token": "[PAD]", "direction": "right"},
|
||||
)
|
||||
|
||||
tokenizer, _ = load_tokenizer(model_dir)
|
||||
encoded = tokenizer.encode_batch(["hello world", "hello"])
|
||||
|
||||
assert len(encoded[0].ids) == len(encoded[1].ids) < 128
|
||||
|
||||
|
||||
def test_serialized_left_padding_is_preserved(make_model_dir) -> None:
|
||||
"""ColModernVBERT pads on the left, normalizing the length must not reset the direction."""
|
||||
model_dir = make_model_dir(
|
||||
padding={"length": None, "pad_id": 0, "pad_token": "[PAD]", "direction": "left"},
|
||||
)
|
||||
|
||||
tokenizer, _ = load_tokenizer(model_dir)
|
||||
|
||||
assert tokenizer.padding["direction"] == "left"
|
||||
assert tokenizer.padding["length"] is None
|
||||
|
||||
encoded = tokenizer.encode_batch(["hello world and then some", "hello"])
|
||||
assert encoded[1].ids[0] == 0
|
||||
assert encoded[1].attention_mask[0] == 0
|
||||
|
||||
|
||||
def test_serialized_pad_to_multiple_of_is_preserved(make_model_dir) -> None:
|
||||
"""Everything the tokenizer declared is kept; only the fixed length is overridden."""
|
||||
model_dir = make_model_dir(
|
||||
padding={"length": 128, "pad_id": 0, "pad_token": "[PAD]", "pad_to_multiple_of": 8},
|
||||
)
|
||||
|
||||
tokenizer, _ = load_tokenizer(model_dir)
|
||||
|
||||
assert tokenizer.padding["length"] is None
|
||||
assert tokenizer.padding["pad_to_multiple_of"] == 8
|
||||
|
||||
encoded = tokenizer.encode_batch(["hello world", "hello"])
|
||||
assert len(encoded[0].ids) % 8 == 0
|
||||
|
||||
|
||||
def test_serialized_pad_id_takes_precedence_over_config(make_model_dir) -> None:
|
||||
mask_pad_id = 103 # [MASK] in the bert-base vocab, any id other than the config's works
|
||||
model_dir = make_model_dir(
|
||||
config={"pad_token_id": 7},
|
||||
padding={"length": None, "pad_id": mask_pad_id, "pad_token": "[MASK]"},
|
||||
)
|
||||
|
||||
tokenizer, _ = load_tokenizer(model_dir)
|
||||
|
||||
assert tokenizer.padding["pad_id"] == mask_pad_id
|
||||
assert tokenizer.padding["pad_token"] == "[MASK]"
|
||||
|
||||
|
||||
def test_pad_token_falls_back_to_tokenizer_config(make_model_dir) -> None:
|
||||
model_dir = make_model_dir(config={"pad_token_id": 3}, padding=None)
|
||||
|
||||
tokenizer, _ = load_tokenizer(model_dir)
|
||||
|
||||
assert tokenizer.padding["pad_token"] == "[PAD]"
|
||||
assert tokenizer.padding["pad_id"] == 3
|
||||
|
||||
|
||||
def test_missing_pad_token_raises(make_model_dir) -> None:
|
||||
model_dir = make_model_dir(drop_from_tokenizer_config=("pad_token",))
|
||||
|
||||
with pytest.raises(ValueError, match="Could not find a pad token"):
|
||||
load_tokenizer(model_dir)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_max_length,max_length,expected",
|
||||
[
|
||||
(512, 128, 128), # both usable, the stricter one wins
|
||||
(128, 512, 128),
|
||||
(512, None, 512),
|
||||
(None, 256, 256),
|
||||
(HF_SENTINEL, 128, 128), # qdrant/gte-large-onnx
|
||||
(0, 256, 256),
|
||||
(512, 0, 512),
|
||||
],
|
||||
)
|
||||
def test_max_context_resolution(make_model_dir, model_max_length, max_length, expected) -> None:
|
||||
model_dir = make_model_dir(
|
||||
tokenizer_config={"model_max_length": model_max_length, "max_length": max_length},
|
||||
)
|
||||
|
||||
tokenizer, _ = load_tokenizer(model_dir)
|
||||
|
||||
assert tokenizer.truncation["max_length"] == expected
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_max_length,max_length",
|
||||
[
|
||||
(HF_SENTINEL, None), # transformers' placeholder is not a limit
|
||||
(0, None), # a zero would truncate everything away
|
||||
(None, 0),
|
||||
(None, None),
|
||||
("512", None), # not an integer
|
||||
],
|
||||
)
|
||||
def test_unusable_max_context_raises(make_model_dir, model_max_length, max_length) -> None:
|
||||
model_dir = make_model_dir(
|
||||
tokenizer_config={"model_max_length": model_max_length, "max_length": max_length},
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="Could not determine the maximum context length"):
|
||||
load_tokenizer(model_dir)
|
||||
|
||||
|
||||
def test_absent_max_context_keys_raise(make_model_dir) -> None:
|
||||
model_dir = make_model_dir(
|
||||
drop_from_tokenizer_config=("model_max_length", "max_length"),
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="Could not determine the maximum context length"):
|
||||
load_tokenizer(model_dir)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def token_id(make_model_dir):
|
||||
"""Resolve vocabulary ids by name, so the cases below carry no magic numbers."""
|
||||
return Tokenizer.from_file(str(make_model_dir() / "tokenizer.json")).token_to_id
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"dropped",
|
||||
[("config.json",), ("special_tokens_map.json",), ("config.json", "special_tokens_map.json")],
|
||||
ids=["no-config", "no-special-tokens-map", "neither"],
|
||||
)
|
||||
def test_optional_files_do_not_change_what_is_loaded(make_model_dir, dropped) -> None:
|
||||
"""Both files are redundant: everything they carry is already in the tokenizer."""
|
||||
baseline, baseline_specials = load_tokenizer(make_model_dir())
|
||||
|
||||
tokenizer, specials = load_tokenizer(make_model_dir(drop_files=dropped))
|
||||
|
||||
assert specials == baseline_specials
|
||||
assert tokenizer.padding == baseline.padding
|
||||
assert tokenizer.encode("hello world").ids == baseline.encode("hello world").ids
|
||||
|
||||
|
||||
@pytest.mark.parametrize("missing", ("tokenizer.json", "tokenizer_config.json"))
|
||||
def test_the_remaining_files_are_still_required(make_model_dir, missing) -> None:
|
||||
"""Relaxing the optional two must not relax the two that carry irreplaceable data."""
|
||||
model_dir = make_model_dir(drop_files=(missing,))
|
||||
|
||||
with pytest.raises(ValueError, match=f"Could not find {missing}"):
|
||||
load_tokenizer(model_dir)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_files",
|
||||
[
|
||||
pytest.param({"drop_from_config": ("pad_token_id",)}, id="config-omits-pad-token-id"),
|
||||
pytest.param({"drop_files": ("config.json",)}, id="config-is-absent"),
|
||||
],
|
||||
)
|
||||
def test_pad_id_falls_back_to_the_vocabulary(make_model_dir, token_id, model_files) -> None:
|
||||
"""Last link of the chain; a hardcoded 0 would silently disagree with `pad_token`."""
|
||||
expected = token_id("[SEP]")
|
||||
assert expected != 0, "a pad token whose id is 0 would pass even without a lookup"
|
||||
model_dir = make_model_dir(tokenizer_config={"pad_token": "[SEP]"}, **model_files)
|
||||
|
||||
tokenizer, _ = load_tokenizer(model_dir)
|
||||
|
||||
assert tokenizer.padding["pad_id"] == expected
|
||||
|
||||
|
||||
def test_pad_token_that_resolves_nowhere_raises(make_model_dir) -> None:
|
||||
"""Without config.json a pad token outside the vocabulary has no id left to fall back on."""
|
||||
model_dir = make_model_dir(
|
||||
tokenizer_config={"pad_token": "[NOT_IN_VOCAB]"},
|
||||
drop_files=("config.json",),
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="Could not resolve an id for the pad token"):
|
||||
load_tokenizer(model_dir)
|
||||
|
||||
|
||||
def test_pad_token_named_only_in_the_map_resolves(make_model_dir) -> None:
|
||||
"""The map is read first, so it can name a pad token tokenizer.json does not carry."""
|
||||
model_dir = make_model_dir(
|
||||
tokenizer_config={"pad_token": "<|mypad|>"},
|
||||
drop_files=("config.json",),
|
||||
)
|
||||
_patch_json(model_dir / "special_tokens_map.json", {"pad_token": "<|mypad|>"})
|
||||
|
||||
tokenizer, specials = load_tokenizer(model_dir)
|
||||
|
||||
assert tokenizer.padding["pad_token"] == "<|mypad|>"
|
||||
assert tokenizer.padding["pad_id"] == specials["<|mypad|>"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"additional",
|
||||
[
|
||||
pytest.param(["<|list_str|>"], id="list-of-strings"),
|
||||
pytest.param([{"content": "<|list_str|>"}], id="list-of-added-token-dicts"),
|
||||
],
|
||||
)
|
||||
def test_list_valued_map_entries_are_registered(make_model_dir, additional) -> None:
|
||||
"""`additional_special_tokens` holds a list, which the str/dict dispatch alone drops.
|
||||
|
||||
Real repos ship both spellings, and their tokens are in tokenizer.json already, so
|
||||
only a token living nowhere else shows the drop.
|
||||
"""
|
||||
model_dir = make_model_dir()
|
||||
_patch_json(model_dir / "special_tokens_map.json", {"additional_special_tokens": additional})
|
||||
|
||||
_, specials = load_tokenizer(model_dir)
|
||||
|
||||
assert "<|list_str|>" in specials
|
||||
+271
-80
@@ -1,14 +1,16 @@
|
||||
import os
|
||||
from contextlib import contextmanager
|
||||
|
||||
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
|
||||
from tests.utils import delete_model_cache, should_test_model
|
||||
|
||||
|
||||
CANONICAL_COLUMN_VALUES = {
|
||||
"prithvida/Splade_PP_en_v1": {
|
||||
"prithivida/Splade_PP_en_v1": {
|
||||
"indices": [
|
||||
2040,
|
||||
2047,
|
||||
@@ -43,112 +45,249 @@ CANONICAL_COLUMN_VALUES = {
|
||||
2.1904349327087402,
|
||||
1.0531445741653442,
|
||||
],
|
||||
}
|
||||
},
|
||||
"Qdrant/minicoil-v1": {
|
||||
"indices": [80, 81, 82, 83, 6664, 6665, 6666, 6667],
|
||||
"values": [
|
||||
0.52634597,
|
||||
0.8711344,
|
||||
1.2264385,
|
||||
0.52123857,
|
||||
0.974713,
|
||||
-0.97803956,
|
||||
-0.94312465,
|
||||
-0.12508166,
|
||||
],
|
||||
},
|
||||
# first 15 non-zero dimensions of the embedding
|
||||
"opensearch-project/opensearch-neural-sparse-encoding-doc-v3-gte": {
|
||||
"indices": [
|
||||
999,
|
||||
1010,
|
||||
1011,
|
||||
1024,
|
||||
1028,
|
||||
1029,
|
||||
1045,
|
||||
1074,
|
||||
1993,
|
||||
2017,
|
||||
2033,
|
||||
2054,
|
||||
2073,
|
||||
2080,
|
||||
2088,
|
||||
],
|
||||
"values": [
|
||||
0.16544909,
|
||||
0.00529129,
|
||||
0.0392109,
|
||||
0.12337475,
|
||||
0.09640586,
|
||||
0.05325737,
|
||||
0.09611791,
|
||||
0.03159865,
|
||||
0.01349991,
|
||||
0.09392473,
|
||||
0.01928805,
|
||||
0.05238346,
|
||||
0.05515401,
|
||||
0.03156782,
|
||||
0.98263124,
|
||||
],
|
||||
},
|
||||
}
|
||||
|
||||
CANONICAL_QUERY_VALUES = {
|
||||
"opensearch-project/opensearch-neural-sparse-encoding-doc-v3-gte": {
|
||||
"indices": [2088, 7592],
|
||||
"values": [3.42086864, 6.93775654],
|
||||
},
|
||||
"Qdrant/minicoil-v1": {
|
||||
"indices": [80, 81, 82, 83, 6664, 6665, 6666, 6667],
|
||||
"values": [
|
||||
0.31389374,
|
||||
0.5195128,
|
||||
0.7314033,
|
||||
0.3108479,
|
||||
0.5812834,
|
||||
-0.5832673,
|
||||
-0.5624452,
|
||||
-0.0745942,
|
||||
],
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
_MODELS_TO_CACHE = (
|
||||
"prithivida/Splade_PP_en_v1",
|
||||
"Qdrant/minicoil-v1",
|
||||
"Qdrant/bm25",
|
||||
"Qdrant/bm42-all-minilm-l6-v2-attentions",
|
||||
"opensearch-project/opensearch-neural-sparse-encoding-doc-v3-gte",
|
||||
)
|
||||
MODELS_TO_CACHE = tuple([x.lower() for x in _MODELS_TO_CACHE])
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def model_cache():
|
||||
is_ci = os.getenv("CI")
|
||||
cache = {}
|
||||
|
||||
@contextmanager
|
||||
def get_model(model_name: str):
|
||||
lowercase_model_name = model_name.lower()
|
||||
if lowercase_model_name not in cache:
|
||||
cache[lowercase_model_name] = SparseTextEmbedding(lowercase_model_name)
|
||||
yield cache[lowercase_model_name]
|
||||
if lowercase_model_name not in MODELS_TO_CACHE:
|
||||
model_inst = cache.pop(lowercase_model_name)
|
||||
if is_ci:
|
||||
delete_model_cache(model_inst.model._model_dir)
|
||||
del model_inst
|
||||
|
||||
yield get_model
|
||||
|
||||
if is_ci:
|
||||
for name, model in cache.items():
|
||||
delete_model_cache(model.model._model_dir)
|
||||
cache.clear()
|
||||
|
||||
|
||||
docs = ["Hello World"]
|
||||
|
||||
|
||||
def test_batch_embedding() -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
@pytest.mark.parametrize(
|
||||
"model_name",
|
||||
["prithivida/Splade_PP_en_v1", "Qdrant/minicoil-v1"],
|
||||
)
|
||||
def test_batch_embedding(model_cache, model_name: str) -> None:
|
||||
docs_to_embed = docs * 10
|
||||
|
||||
for model_name, expected_result in CANONICAL_COLUMN_VALUES.items():
|
||||
model = SparseTextEmbedding(model_name=model_name)
|
||||
with model_cache(model_name) as model:
|
||||
result = next(iter(model.embed(docs_to_embed, batch_size=6)))
|
||||
expected_result = CANONICAL_COLUMN_VALUES[model_name]
|
||||
assert result.indices.tolist() == expected_result["indices"]
|
||||
|
||||
for i, value in enumerate(result.values):
|
||||
assert pytest.approx(value, abs=0.001) == expected_result["values"][i]
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
def test_single_embedding() -> None:
|
||||
def test_single_embedding(model_cache) -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
for model_name, expected_result in CANONICAL_COLUMN_VALUES.items():
|
||||
model = SparseTextEmbedding(model_name=model_name)
|
||||
is_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
|
||||
|
||||
passage_result = next(iter(model.embed(docs, batch_size=6)))
|
||||
query_result = next(iter(model.query_embed(docs)))
|
||||
for result in [passage_result, query_result]:
|
||||
assert result.indices.tolist() == expected_result["indices"]
|
||||
for model_desc in SparseTextEmbedding._list_supported_models():
|
||||
if (
|
||||
model_desc.model not in CANONICAL_COLUMN_VALUES
|
||||
): # attention models and bm25 are also parts of
|
||||
# SparseTextEmbedding, however, they have their own tests
|
||||
continue
|
||||
if not should_test_model(model_desc, model_desc.model, is_ci, is_manual):
|
||||
continue
|
||||
with model_cache(model_desc.model) as model:
|
||||
passage_result = next(iter(model.embed(docs, batch_size=6)))
|
||||
query_result = next(iter(model.query_embed(docs)))
|
||||
expected_result = CANONICAL_COLUMN_VALUES[model_desc.model]
|
||||
expected_query_result = CANONICAL_QUERY_VALUES.get(model_desc.model, expected_result)
|
||||
|
||||
for i, value in enumerate(result.values):
|
||||
# canonical values might contain only a prefix of the non-zero dimensions
|
||||
num_dims = len(expected_result["indices"])
|
||||
assert passage_result.indices.tolist()[:num_dims] == expected_result["indices"]
|
||||
for i, value in enumerate(passage_result.values[:num_dims]):
|
||||
assert pytest.approx(value, abs=0.001) == expected_result["values"][i]
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
num_query_dims = len(expected_query_result["indices"])
|
||||
assert (
|
||||
query_result.indices.tolist()[:num_query_dims] == expected_query_result["indices"]
|
||||
)
|
||||
for i, value in enumerate(query_result.values[:num_query_dims]):
|
||||
assert pytest.approx(value, abs=0.001) == expected_query_result["values"][i]
|
||||
|
||||
|
||||
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))
|
||||
sparse_embeddings_all = list(model.embed(docs, batch_size=10, parallel=0))
|
||||
sparse_embeddings = list(model.embed(docs, batch_size=10, parallel=None))
|
||||
@pytest.mark.parametrize(
|
||||
"model_name",
|
||||
["prithivida/Splade_PP_en_v1", "Qdrant/minicoil-v1"],
|
||||
)
|
||||
def test_parallel_processing(model_cache, model_name: str) -> None:
|
||||
with model_cache(model_name) as model:
|
||||
docs = ["hello world", "flag embedding"] * 30
|
||||
sparse_embeddings_duo = list(model.embed(docs, batch_size=10, parallel=2))
|
||||
# sparse_embeddings_all = list(model.embed(docs, batch_size=10, parallel=0)) # inherits OnnxTextModel which
|
||||
# is tested in TextEmbedding, disabling it here to reduce number of requests to hf
|
||||
# multiprocessing is enough to test with `parallel=2`, and `parallel=None` is okay to tests since it reuses
|
||||
# model from cache
|
||||
sparse_embeddings = list(model.embed(docs, batch_size=10, parallel=None))
|
||||
|
||||
assert (
|
||||
len(sparse_embeddings)
|
||||
== len(sparse_embeddings_duo)
|
||||
== len(sparse_embeddings_all)
|
||||
== len(docs)
|
||||
)
|
||||
|
||||
for sparse_embedding, sparse_embedding_duo, sparse_embedding_all in zip(
|
||||
sparse_embeddings, sparse_embeddings_duo, sparse_embeddings_all
|
||||
):
|
||||
assert (
|
||||
sparse_embedding.indices.tolist()
|
||||
== sparse_embedding_duo.indices.tolist()
|
||||
== sparse_embedding_all.indices.tolist()
|
||||
len(sparse_embeddings)
|
||||
== len(sparse_embeddings_duo)
|
||||
# == len(sparse_embeddings_all)
|
||||
== len(docs)
|
||||
)
|
||||
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)
|
||||
for (
|
||||
sparse_embedding,
|
||||
sparse_embedding_duo,
|
||||
# sparse_embedding_all
|
||||
) in zip(
|
||||
sparse_embeddings,
|
||||
sparse_embeddings_duo,
|
||||
# sparse_embeddings_all
|
||||
):
|
||||
assert (
|
||||
sparse_embedding.indices.tolist() == 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)
|
||||
|
||||
|
||||
@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(model_cache) -> None:
|
||||
with model_cache("Qdrant/bm25") as model:
|
||||
bm25_instance = model.model
|
||||
# Setup
|
||||
original_stopwords = bm25_instance.stopwords.copy()
|
||||
original_punctuation = bm25_instance.punctuation.copy()
|
||||
|
||||
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}"
|
||||
|
||||
bm25_instance.stopwords = original_stopwords
|
||||
bm25_instance.punctuation = original_punctuation
|
||||
|
||||
|
||||
def test_stem_with_stopwords_and_punctuation(bm25_instance: Bm25) -> None:
|
||||
# Setup
|
||||
bm25_instance.stopwords = {"the", "is", "a"}
|
||||
bm25_instance.punctuation = {".", ",", "!"}
|
||||
def test_stem_case_insensitive_stopwords(model_cache) -> None:
|
||||
with model_cache("Qdrant/bm25") as model:
|
||||
bm25_instance = model.model
|
||||
original_stopwords = bm25_instance.stopwords.copy()
|
||||
original_punctuation = bm25_instance.punctuation.copy()
|
||||
|
||||
# Test data
|
||||
tokens = ["The", "quick", "brown", "fox", "is", "a", "test", "sentence", ".", "!"]
|
||||
# Setup
|
||||
bm25_instance.stopwords = {"the", "is", "a"}
|
||||
bm25_instance.punctuation = {".", ",", "!"}
|
||||
|
||||
# Execute
|
||||
result = bm25_instance._stem(tokens)
|
||||
# Test data
|
||||
tokens = ["THE", "Quick", "Brown", "Fox", "IS", "A", "Test", "Sentence", ".", "!"]
|
||||
|
||||
# Assert
|
||||
expected = ["quick", "brown", "fox", "test", "sentenc"]
|
||||
assert result == expected, f"Expected {expected}, but got {result}"
|
||||
# Execute
|
||||
result = bm25_instance._stem(tokens)
|
||||
|
||||
|
||||
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}"
|
||||
# Assert
|
||||
expected = ["quick", "brown", "fox", "test", "sentenc"]
|
||||
assert result == expected, f"Expected {expected}, but got {result}"
|
||||
bm25_instance.stopwords = original_stopwords
|
||||
bm25_instance.punctuation = original_punctuation
|
||||
|
||||
|
||||
@pytest.mark.parametrize("disable_stemmer", [True, False])
|
||||
@@ -172,10 +311,23 @@ def test_disable_stemmer_behavior(disable_stemmer: bool) -> None:
|
||||
assert result == expected, f"Expected {expected}, but got {result}"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_name",
|
||||
["prithivida/Splade_PP_en_v1"],
|
||||
)
|
||||
def test_if_splade_query_embed_is_inference_free() -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
model = SparseTextEmbedding(
|
||||
model_name="opensearch-project/opensearch-neural-sparse-encoding-doc-v3-gte",
|
||||
lazy_load=True,
|
||||
)
|
||||
embeddings = list(model.query_embed(["hello world", "flag embedding"]))
|
||||
# queries are embedded with a tokenizer and an idf lookup table only,
|
||||
# the onnx model must stay unloaded
|
||||
assert not hasattr(model.model, "model")
|
||||
assert all(len(embedding.indices) > 0 for embedding in embeddings)
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
@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)
|
||||
@@ -193,3 +345,42 @@ def test_lazy_load(model_name: str) -> None:
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_name",
|
||||
[
|
||||
"prithivida/Splade_PP_en_v1",
|
||||
"Qdrant/minicoil-v1",
|
||||
"Qdrant/bm42-all-minilm-l6-v2-attentions",
|
||||
],
|
||||
)
|
||||
def test_session_options(model_cache, model_name) -> None:
|
||||
with model_cache(model_name) as default_model:
|
||||
default_session_options = default_model.model.model.get_session_options()
|
||||
assert default_session_options.enable_cpu_mem_arena is True
|
||||
model = SparseTextEmbedding(model_name=model_name, enable_cpu_mem_arena=False)
|
||||
session_options = model.model.model.get_session_options()
|
||||
assert session_options.enable_cpu_mem_arena is False
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_name",
|
||||
[
|
||||
"prithivida/Splade_PP_en_v1",
|
||||
"Qdrant/minicoil-v1",
|
||||
"Qdrant/bm42-all-minilm-l6-v2-attentions",
|
||||
"Qdrant/bm25",
|
||||
],
|
||||
)
|
||||
def test_token_count(model_cache, model_name) -> None:
|
||||
with model_cache(model_name) as model:
|
||||
documents = [
|
||||
"Name me a couple of cities were the capitals of Germany?",
|
||||
"Berlin is the current capital of Germany, Bonn is a former capital of Germany.",
|
||||
]
|
||||
first_doc_token_count = model.token_count(documents[0])
|
||||
second_doc_token_count = model.token_count(documents[1])
|
||||
doc_token_count = model.token_count(documents)
|
||||
assert first_doc_token_count + second_doc_token_count == doc_token_count
|
||||
assert doc_token_count == model.token_count(documents, batch_size=1)
|
||||
|
||||
@@ -1,10 +1,11 @@
|
||||
import os
|
||||
from contextlib import contextmanager
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from fastembed.rerank.cross_encoder import TextCrossEncoder
|
||||
from tests.utils import delete_model_cache
|
||||
from tests.utils import delete_model_cache, should_test_model
|
||||
|
||||
CANONICAL_SCORE_VALUES = {
|
||||
"Xenova/ms-marco-MiniLM-L-6-v2": np.array([8.500708, -2.541011]),
|
||||
@@ -15,73 +16,84 @@ CANONICAL_SCORE_VALUES = {
|
||||
"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",
|
||||
}
|
||||
|
||||
_MODELS_TO_CACHE = ("Xenova/ms-marco-MiniLM-L-6-v2",)
|
||||
MODELS_TO_CACHE = tuple([x.lower() for x in _MODELS_TO_CACHE])
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_name",
|
||||
[model_name for model_name in CANONICAL_SCORE_VALUES],
|
||||
)
|
||||
def test_rerank(model_name: str) -> None:
|
||||
@pytest.fixture(scope="module")
|
||||
def model_cache():
|
||||
is_ci = os.getenv("CI")
|
||||
cache = {}
|
||||
|
||||
model = TextCrossEncoder(model_name=model_name)
|
||||
@contextmanager
|
||||
def get_model(model_name: str):
|
||||
lowercase_model_name = model_name.lower()
|
||||
if lowercase_model_name not in cache:
|
||||
cache[lowercase_model_name] = TextCrossEncoder(lowercase_model_name)
|
||||
yield cache[lowercase_model_name]
|
||||
if lowercase_model_name not in MODELS_TO_CACHE:
|
||||
model_inst = cache.pop(lowercase_model_name)
|
||||
if is_ci:
|
||||
delete_model_cache(model_inst.model._model_dir)
|
||||
del model_inst
|
||||
|
||||
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)))
|
||||
yield get_model
|
||||
|
||||
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)
|
||||
for name, model in cache.items():
|
||||
delete_model_cache(model.model._model_dir)
|
||||
cache.clear()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_name",
|
||||
[model_name for model_name in SELECTED_MODELS.values()],
|
||||
)
|
||||
def test_batch_rerank(model_name: str) -> None:
|
||||
@pytest.mark.parametrize("model_name", ["Xenova/ms-marco-MiniLM-L-6-v2"])
|
||||
def test_rerank(model_cache, model_name: str) -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
is_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
|
||||
|
||||
model = TextCrossEncoder(model_name=model_name)
|
||||
for model_desc in TextCrossEncoder._list_supported_models():
|
||||
if not should_test_model(model_desc, model_name, is_ci, is_manual):
|
||||
continue
|
||||
|
||||
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)))
|
||||
with model_cache(model_desc.model) as model:
|
||||
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}"
|
||||
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_desc.model}, 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)
|
||||
canonical_scores = CANONICAL_SCORE_VALUES[model_desc.model]
|
||||
assert np.allclose(
|
||||
scores, canonical_scores, atol=1e-3
|
||||
), f"Model: {model_desc.model}, Scores: {scores}, Expected: {canonical_scores}"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_name",
|
||||
["Xenova/ms-marco-MiniLM-L-6-v2"],
|
||||
)
|
||||
@pytest.mark.parametrize("model_name", ["Xenova/ms-marco-MiniLM-L-6-v2"])
|
||||
def test_batch_rerank(model_cache, model_name: str) -> None:
|
||||
with model_cache(model_name) as model:
|
||||
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}"
|
||||
|
||||
|
||||
@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)
|
||||
@@ -95,25 +107,45 @@ def test_lazy_load(model_name: str) -> None:
|
||||
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")
|
||||
@pytest.mark.parametrize("model_name", ["Xenova/ms-marco-MiniLM-L-6-v2"])
|
||||
def test_rerank_pairs_parallel(model_cache, model_name: str) -> None:
|
||||
with model_cache(model_name) as model:
|
||||
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}"
|
||||
|
||||
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)
|
||||
|
||||
@pytest.mark.parametrize("model_name", ["Xenova/ms-marco-MiniLM-L-6-v2"])
|
||||
def test_token_count(model_cache, model_name: str) -> None:
|
||||
with model_cache(model_name) as model:
|
||||
pairs = [
|
||||
("What is the capital of France?", "Paris is the capital of France."),
|
||||
(
|
||||
"Name me a couple of cities were the capitals of Germany?",
|
||||
"Berlin is the current capital of Germany, Bonn is a former capital of Germany.",
|
||||
),
|
||||
]
|
||||
first_pair_token_count = model.token_count([pairs[0]])
|
||||
second_pair_token_count = model.token_count([pairs[1]])
|
||||
pairs_token_count = model.token_count(pairs)
|
||||
assert first_pair_token_count + second_pair_token_count == pairs_token_count
|
||||
assert pairs_token_count == model.token_count(pairs, batch_size=1)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_name", ["Xenova/ms-marco-MiniLM-L-6-v2"])
|
||||
def test_session_options(model_cache, model_name) -> None:
|
||||
with model_cache(model_name) as default_model:
|
||||
default_session_options = default_model.model.model.get_session_options()
|
||||
assert default_session_options.enable_cpu_mem_arena is True
|
||||
model = TextCrossEncoder(model_name=model_name, enable_cpu_mem_arena=False)
|
||||
session_options = model.model.model.get_session_options()
|
||||
assert session_options.enable_cpu_mem_arena is False
|
||||
|
||||
@@ -4,7 +4,7 @@ import numpy as np
|
||||
import pytest
|
||||
|
||||
from fastembed import TextEmbedding
|
||||
from fastembed.text.multitask_embedding import Task
|
||||
from fastembed.text.multitask_embedding import JinaEmbeddingV3, Task
|
||||
from tests.utils import delete_model_cache
|
||||
|
||||
|
||||
@@ -60,52 +60,43 @@ CANONICAL_VECTOR_VALUES = {
|
||||
docs = ["Hello World", "Follow the white rabbit."]
|
||||
|
||||
|
||||
def test_batch_embedding():
|
||||
@pytest.mark.parametrize("dim,model_name", [(1024, "jinaai/jina-embeddings-v3")])
|
||||
def test_batch_embedding(dim: int, model_name: str):
|
||||
is_ci = os.getenv("CI")
|
||||
is_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
|
||||
if is_ci and not is_manual:
|
||||
pytest.skip("Skipping multitask models in CI non-manual mode")
|
||||
|
||||
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 = TextEmbedding(model_name=model_name)
|
||||
|
||||
model_name = model_desc.model
|
||||
dim = model_desc.dim
|
||||
embeddings = list(model.embed(documents=docs_to_embed, batch_size=6))
|
||||
embeddings = np.stack(embeddings, axis=0)
|
||||
|
||||
if model_name not in CANONICAL_VECTOR_VALUES.keys():
|
||||
continue
|
||||
assert embeddings.shape == (len(docs_to_embed), dim)
|
||||
|
||||
model = TextEmbedding(model_name=model_name)
|
||||
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_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)
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
def test_single_embedding():
|
||||
is_ci = os.getenv("CI")
|
||||
is_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
|
||||
if is_ci and not is_manual:
|
||||
pytest.skip("Skipping multitask models in CI non-manual mode")
|
||||
|
||||
for model_desc in TextEmbedding._list_supported_models():
|
||||
if not is_ci and model_desc.size_in_GB > 1:
|
||||
continue
|
||||
|
||||
for model_desc in JinaEmbeddingV3._list_supported_models():
|
||||
# todo: once we add more models, we should not test models >1GB size locally
|
||||
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]:
|
||||
@@ -118,27 +109,42 @@ def test_single_embedding():
|
||||
|
||||
canonical_vector = task["vectors"]
|
||||
assert np.allclose(
|
||||
embeddings[: len(docs), : canonical_vector.shape[1]], canonical_vector, atol=1e-4
|
||||
embeddings[:, : canonical_vector.shape[1]], canonical_vector, atol=1e-4
|
||||
), model_desc.model
|
||||
|
||||
classification_embeddings = list(model.embed(documents=docs, task_id=Task.CLASSIFICATION))
|
||||
classification_embeddings = np.stack(classification_embeddings, axis=0)
|
||||
|
||||
assert classification_embeddings.shape == (len(docs), dim)
|
||||
|
||||
model = TextEmbedding(model_name=model_name, task_id=Task.CLASSIFICATION)
|
||||
default_embeddings = list(model.embed(documents=docs))
|
||||
default_embeddings = np.stack(default_embeddings, axis=0)
|
||||
|
||||
assert default_embeddings.shape == (len(docs), dim)
|
||||
|
||||
assert np.allclose(
|
||||
classification_embeddings,
|
||||
default_embeddings,
|
||||
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")
|
||||
is_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
|
||||
if is_ci and not is_manual:
|
||||
pytest.skip("Skipping multitask models in CI non-manual mode")
|
||||
|
||||
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
|
||||
|
||||
for model_desc in JinaEmbeddingV3._list_supported_models():
|
||||
# todo: once we add more models, we should not test models >1GB size locally
|
||||
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}")
|
||||
@@ -150,7 +156,7 @@ def test_single_embedding_query():
|
||||
|
||||
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
|
||||
embeddings[:, : canonical_vector.shape[1]], canonical_vector, atol=1e-4
|
||||
), model_desc.model
|
||||
|
||||
if is_ci:
|
||||
@@ -159,18 +165,18 @@ def test_single_embedding_query():
|
||||
|
||||
def test_single_embedding_passage():
|
||||
is_ci = os.getenv("CI")
|
||||
is_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
|
||||
if is_ci and not is_manual:
|
||||
pytest.skip("Skipping multitask models in CI non-manual mode")
|
||||
|
||||
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
|
||||
for model_desc in JinaEmbeddingV3._list_supported_models():
|
||||
# todo: once we add more models, we should not test models >1GB size locally
|
||||
|
||||
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}")
|
||||
@@ -182,21 +188,22 @@ def test_single_embedding_passage():
|
||||
|
||||
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
|
||||
embeddings[:, : 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():
|
||||
@pytest.mark.parametrize("dim,model_name", [(1024, "jinaai/jina-embeddings-v3")])
|
||||
def test_parallel_processing(dim: int, model_name: str):
|
||||
is_ci = os.getenv("CI")
|
||||
is_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
|
||||
if is_ci and not is_manual:
|
||||
pytest.skip("Skipping in CI non-manual mode")
|
||||
|
||||
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
|
||||
@@ -216,33 +223,14 @@ def test_parallel_processing():
|
||||
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"],
|
||||
)
|
||||
@pytest.mark.parametrize("model_name", ["jinaai/jina-embeddings-v3"])
|
||||
def test_lazy_load(model_name: str):
|
||||
is_ci = os.getenv("CI")
|
||||
is_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
|
||||
|
||||
if is_ci and not is_manual:
|
||||
pytest.skip("Skipping in CI non-manual mode")
|
||||
|
||||
model = TextEmbedding(model_name=model_name, lazy_load=True)
|
||||
assert not hasattr(model.model, "model")
|
||||
|
||||
|
||||
@@ -1,11 +1,14 @@
|
||||
import os
|
||||
import platform
|
||||
from contextlib import contextmanager
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from fastembed.text.last_token_normalized_embedding import LastTokenNormalizedEmbedding
|
||||
from fastembed.text.onnx_embedding import OnnxTextEmbedding
|
||||
from fastembed.text.text_embedding import TextEmbedding
|
||||
from tests.utils import delete_model_cache
|
||||
from tests.utils import delete_model_cache, should_test_model
|
||||
|
||||
CANONICAL_VECTOR_VALUES = {
|
||||
"BAAI/bge-small-en": np.array([-0.0232, -0.0255, 0.0174, -0.0639, -0.0006]),
|
||||
@@ -52,7 +55,7 @@ CANONICAL_VECTOR_VALUES = {
|
||||
[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]
|
||||
[-0.00986551, -0.00018734, 0.00605892, -0.03289612, -0.0387564],
|
||||
),
|
||||
"mixedbread-ai/mxbai-embed-large-v1": np.array(
|
||||
[0.02295546, 0.03196154, 0.016512, -0.04031524, -0.0219634]
|
||||
@@ -67,86 +70,194 @@ CANONICAL_VECTOR_VALUES = {
|
||||
"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]),
|
||||
"google/embeddinggemma-300m": np.array(
|
||||
[-0.08181356, 0.0214127, 0.05120273, -0.03690156, -0.0254504]
|
||||
),
|
||||
"Qwen/Qwen3-Embedding-0.6B": np.array(
|
||||
[-0.01476084, 0.01723184, -0.01195498, -0.07275258, 0.00281229]
|
||||
),
|
||||
"Qwen/Qwen3-Embedding-0.6B-Q": np.array(
|
||||
[-0.01599521, 0.01676456, -0.01195119, -0.07132675, 0.00346729]
|
||||
),
|
||||
"google/siglip2-base-patch16-224": np.array(
|
||||
[-0.01181389, 0.00737596, 0.01118064, 0.0103095, 0.3451049]
|
||||
),
|
||||
"minishlab/potion-base-8M": np.array(
|
||||
[-0.03432461, -0.08020256, -0.14396408, 0.08480079, 0.01958815]
|
||||
),
|
||||
"minishlab/potion-retrieval-32M": np.array(
|
||||
[0.019733, -0.01530093, -0.08678473, 0.0229059, 0.04700558]
|
||||
),
|
||||
"minishlab/potion-multilingual-128M": np.array(
|
||||
[0.02366836, 0.02973341, 0.05140258, -0.00745248, -0.06740689]
|
||||
),
|
||||
}
|
||||
|
||||
QWEN3_INSTRUCT_PREFIX = (
|
||||
"Instruct: Given a web search query, retrieve relevant passages that answer the query\nQuery:"
|
||||
)
|
||||
|
||||
DOC_PREFIXES = {
|
||||
"google/embeddinggemma-300m": "title: none | text: ",
|
||||
}
|
||||
QUERY_PREFIXES = {
|
||||
"google/embeddinggemma-300m": "task: search result | query: ",
|
||||
"Qwen/Qwen3-Embedding-0.6B": QWEN3_INSTRUCT_PREFIX,
|
||||
"Qwen/Qwen3-Embedding-0.6B-Q": QWEN3_INSTRUCT_PREFIX,
|
||||
}
|
||||
CANONICAL_QUERY_VECTOR_VALUES = {
|
||||
"google/embeddinggemma-300m": np.array(
|
||||
[-0.22990295, 0.03311195, 0.04290345, -0.03558498, -0.01399477]
|
||||
),
|
||||
"Qwen/Qwen3-Embedding-0.6B": np.array(
|
||||
[-0.01908712, 0.01635596, -0.00356586, -0.03947155, -0.01387356]
|
||||
),
|
||||
"Qwen/Qwen3-Embedding-0.6B-Q": np.array(
|
||||
[-0.02221339, 0.01932909, -0.00361797, -0.03888897, -0.01362813]
|
||||
),
|
||||
}
|
||||
|
||||
MULTI_TASK_MODELS = ["jinaai/jina-embeddings-v3"]
|
||||
|
||||
_MODELS_TO_CACHE = ("BAAI/bge-small-en-v1.5",)
|
||||
MODELS_TO_CACHE = tuple([x.lower() for x in _MODELS_TO_CACHE])
|
||||
|
||||
def test_embedding() -> None:
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def model_cache():
|
||||
is_ci = os.getenv("CI")
|
||||
cache = {}
|
||||
|
||||
@contextmanager
|
||||
def get_model(model_name: str):
|
||||
lowercase_model_name = model_name.lower()
|
||||
if lowercase_model_name not in cache:
|
||||
cache[lowercase_model_name] = TextEmbedding(lowercase_model_name)
|
||||
yield cache[lowercase_model_name]
|
||||
if lowercase_model_name not in MODELS_TO_CACHE:
|
||||
model_inst = cache.pop(lowercase_model_name)
|
||||
if is_ci:
|
||||
delete_model_cache(model_inst.model._model_dir)
|
||||
del model_inst
|
||||
|
||||
yield get_model
|
||||
|
||||
if is_ci:
|
||||
for name, model in cache.items():
|
||||
delete_model_cache(model.model._model_dir)
|
||||
cache.clear()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_name", ["BAAI/bge-small-en-v1.5"])
|
||||
def test_embedding(model_cache, model_name: str) -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
is_mac = platform.system() == "Darwin"
|
||||
is_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
|
||||
|
||||
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")
|
||||
if model_desc.model in MULTI_TASK_MODELS or (
|
||||
is_mac and model_desc.model == "nomic-ai/nomic-embed-text-v1.5-Q"
|
||||
):
|
||||
continue
|
||||
if not should_test_model(model_desc, model_name, is_ci, is_manual):
|
||||
continue
|
||||
|
||||
dim = model_desc.dim
|
||||
|
||||
model = TextEmbedding(model_name=model_desc.model)
|
||||
docs = ["hello world", "flag embedding"]
|
||||
embeddings = list(model.embed(docs))
|
||||
with model_cache(model_desc.model) as model:
|
||||
docs = ["hello world", "flag embedding"]
|
||||
if model_desc.model in DOC_PREFIXES:
|
||||
docs = [DOC_PREFIXES[model_desc.model] + doc for doc in docs]
|
||||
|
||||
embeddings = list(model.embed(docs))
|
||||
embeddings = np.stack(embeddings, axis=0)
|
||||
assert embeddings.shape == (2, dim)
|
||||
|
||||
canonical_vector = CANONICAL_VECTOR_VALUES[model_desc.model]
|
||||
assert np.allclose(
|
||||
embeddings[0, : canonical_vector.shape[0]], canonical_vector, atol=1e-3
|
||||
), model_desc.model
|
||||
|
||||
|
||||
def test_query_embedding(model_cache) -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
is_mac = platform.system() == "Darwin"
|
||||
is_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
|
||||
|
||||
for model_desc in TextEmbedding._list_supported_models():
|
||||
if model_desc.model in MULTI_TASK_MODELS or (
|
||||
is_mac and model_desc.model == "nomic-ai/nomic-embed-text-v1.5-Q"
|
||||
):
|
||||
continue
|
||||
|
||||
if model_desc.model not in CANONICAL_QUERY_VECTOR_VALUES:
|
||||
continue
|
||||
|
||||
if not should_test_model(model_desc, "", is_ci, is_manual):
|
||||
continue
|
||||
|
||||
dim = model_desc.dim
|
||||
with model_cache(model_desc.model) as model:
|
||||
queries = ["hello world", "flag embedding"]
|
||||
if model_desc.model in QUERY_PREFIXES:
|
||||
queries = [QUERY_PREFIXES[model_desc.model] + query for query in queries]
|
||||
|
||||
embeddings = list(model.query_embed(queries))
|
||||
embeddings = np.stack(embeddings, axis=0)
|
||||
assert embeddings.shape == (2, dim)
|
||||
|
||||
canonical_vector = CANONICAL_QUERY_VECTOR_VALUES[model_desc.model]
|
||||
assert np.allclose(
|
||||
embeddings[0, : canonical_vector.shape[0]], canonical_vector, atol=1e-3
|
||||
), model_desc.model
|
||||
|
||||
|
||||
def test_quantized_model_reports_onnxruntime_requirement(monkeypatch) -> None:
|
||||
"""Old onnxruntime only implements 4-bit MatMulNBits, the error should say so."""
|
||||
monkeypatch.setattr(
|
||||
OnnxTextEmbedding,
|
||||
"load_onnx_model",
|
||||
lambda self: (_ for _ in ()).throw(RuntimeError("nbits_ == 4 was false")),
|
||||
)
|
||||
model = LastTokenNormalizedEmbedding(
|
||||
"Qwen/Qwen3-Embedding-0.6B-Q",
|
||||
lazy_load=True,
|
||||
specific_model_path="./", # disable model downloading and loading
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError, match="onnxruntime>=1.23"):
|
||||
model.load_onnx_model()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("n_dims,model_name", [(384, "BAAI/bge-small-en-v1.5")])
|
||||
def test_batch_embedding(model_cache, n_dims: int, model_name: str) -> None:
|
||||
with model_cache(model_name) as model:
|
||||
docs = ["hello world", "flag embedding"] * 100
|
||||
embeddings = list(model.embed(docs, batch_size=10))
|
||||
embeddings = np.stack(embeddings, axis=0)
|
||||
assert embeddings.shape == (2, dim)
|
||||
|
||||
canonical_vector = CANONICAL_VECTOR_VALUES[model_desc.model]
|
||||
assert np.allclose(
|
||||
embeddings[0, : canonical_vector.shape[0]], canonical_vector, atol=1e-3
|
||||
), model_desc["model"]
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
assert embeddings.shape == (len(docs), n_dims)
|
||||
|
||||
|
||||
@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: int, model_name: str) -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
model = TextEmbedding(model_name=model_name)
|
||||
@pytest.mark.parametrize("n_dims,model_name", [(384, "BAAI/bge-small-en-v1.5")])
|
||||
def test_parallel_processing(model_cache, n_dims: int, model_name: str) -> None:
|
||||
with model_cache(model_name) as model:
|
||||
docs = ["hello world", "flag embedding"] * 100
|
||||
embeddings = list(model.embed(docs, batch_size=10, parallel=2))
|
||||
embeddings = np.stack(embeddings, axis=0)
|
||||
|
||||
docs = ["hello world", "flag embedding"] * 100
|
||||
embeddings = list(model.embed(docs, batch_size=10))
|
||||
embeddings = np.stack(embeddings, axis=0)
|
||||
embeddings_2 = list(model.embed(docs, batch_size=10, parallel=None))
|
||||
embeddings_2 = np.stack(embeddings_2, axis=0)
|
||||
|
||||
assert embeddings.shape == (200, n_dims)
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
embeddings_3 = list(model.embed(docs, batch_size=10, parallel=0))
|
||||
embeddings_3 = np.stack(embeddings_3, axis=0)
|
||||
|
||||
assert embeddings.shape == (len(docs), n_dims)
|
||||
assert np.allclose(embeddings, embeddings_2, atol=1e-3)
|
||||
assert np.allclose(embeddings, embeddings_3, atol=1e-3)
|
||||
|
||||
|
||||
@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: int, model_name: str) -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
model = TextEmbedding(model_name=model_name)
|
||||
|
||||
docs = ["hello world", "flag embedding"] * 100
|
||||
embeddings = list(model.embed(docs, batch_size=10, parallel=2))
|
||||
embeddings = np.stack(embeddings, axis=0)
|
||||
|
||||
embeddings_2 = list(model.embed(docs, batch_size=10, parallel=None))
|
||||
embeddings_2 = np.stack(embeddings_2, axis=0)
|
||||
|
||||
embeddings_3 = list(model.embed(docs, batch_size=10, parallel=0))
|
||||
embeddings_3 = np.stack(embeddings_3, axis=0)
|
||||
|
||||
assert embeddings.shape == (200, 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"],
|
||||
)
|
||||
@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)
|
||||
@@ -163,3 +274,60 @@ def test_lazy_load(model_name: str) -> None:
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
def test_get_embedding_size() -> None:
|
||||
assert TextEmbedding.get_embedding_size("sentence-transformers/all-MiniLM-L6-v2") == 384
|
||||
assert TextEmbedding.get_embedding_size("sentence-transformers/all-minilm-l6-v2") == 384
|
||||
|
||||
|
||||
def test_embedding_size() -> None:
|
||||
is_ci = os.getenv("CI")
|
||||
model_name = "sentence-transformers/all-MiniLM-L6-v2"
|
||||
model = TextEmbedding(model_name=model_name, lazy_load=True)
|
||||
assert model.embedding_size == 384
|
||||
|
||||
model_name = "sentence-transformers/all-minilm-l6-v2"
|
||||
model = TextEmbedding(model_name=model_name, lazy_load=True)
|
||||
assert model.embedding_size == 384
|
||||
|
||||
if is_ci:
|
||||
delete_model_cache(model.model._model_dir)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_name", ["sentence-transformers/all-MiniLM-L6-v2"])
|
||||
def test_session_options(model_cache, model_name) -> None:
|
||||
with model_cache(model_name) as default_model:
|
||||
default_session_options = default_model.model.model.get_session_options()
|
||||
assert default_session_options.enable_cpu_mem_arena is True
|
||||
model = TextEmbedding(model_name=model_name, enable_cpu_mem_arena=False)
|
||||
session_options = model.model.model.get_session_options()
|
||||
assert session_options.enable_cpu_mem_arena is False
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_name", ["sentence-transformers/all-MiniLM-L6-v2"])
|
||||
def test_token_count(model_cache, model_name) -> None:
|
||||
with model_cache(model_name) as model:
|
||||
documents = [
|
||||
"Name me a couple of cities were the capitals of Germany?",
|
||||
"Berlin is the current capital of Germany, Bonn is a former capital of Germany.",
|
||||
]
|
||||
first_doc_token_count = model.token_count(documents[0])
|
||||
second_doc_token_count = model.token_count(documents[1])
|
||||
doc_token_count = model.token_count(documents)
|
||||
assert first_doc_token_count + second_doc_token_count == doc_token_count
|
||||
assert doc_token_count == model.token_count(documents, batch_size=1)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_name,dim",
|
||||
[("sentence-transformers/all-MiniLM-L6-v2", 384), ("thenlper/gte-base", 768)],
|
||||
)
|
||||
def test_mixed_length_batch_with_fixed_padding(model_cache, model_name: str, dim: int) -> None:
|
||||
# both models serialize a fixed padding length of 128 in tokenizer.json; gte-base truncates
|
||||
# at 512, so a document longer than 128 makes the batch ragged unless the padding is relaxed
|
||||
with model_cache(model_name) as model:
|
||||
assert model.model.tokenizer.padding["length"] is None
|
||||
|
||||
embeddings = np.stack(list(model.embed(["hello world", "retrieval " * 200])), axis=0)
|
||||
assert embeddings.shape == (2, dim)
|
||||
|
||||
+32
-2
@@ -3,10 +3,12 @@ import traceback
|
||||
|
||||
from pathlib import Path
|
||||
from types import TracebackType
|
||||
from typing import Union, Callable, Any, Type
|
||||
from typing import Callable, Any, Type
|
||||
|
||||
from fastembed.common.model_description import BaseModelDescription
|
||||
|
||||
|
||||
def delete_model_cache(model_dir: Union[str, Path]) -> None:
|
||||
def delete_model_cache(model_dir: 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
|
||||
@@ -35,3 +37,31 @@ def delete_model_cache(model_dir: Union[str, Path]) -> None:
|
||||
if model_dir.exists():
|
||||
# todo: PermissionDenied is raised on blobs removal in Windows, with blobs > 2GB
|
||||
shutil.rmtree(model_dir, onerror=on_error)
|
||||
|
||||
|
||||
def should_test_model(
|
||||
model_desc: BaseModelDescription,
|
||||
autotest_model_name: str,
|
||||
is_ci: str | None,
|
||||
is_manual: bool,
|
||||
):
|
||||
"""Determine if a model should be tested based on environment
|
||||
|
||||
Tests can be run either in ci or locally.
|
||||
Testing all models each time in ci is too long.
|
||||
The testing scheme in ci and on a local machine are different, therefore, there are 3 possible scenarios.
|
||||
1) Run lightweight tests in ci:
|
||||
- test only one model that has been manually chosen as a representative for a certain class family
|
||||
2) Run heavyweight (manual) tests in ci:
|
||||
- test all models
|
||||
Running tests in ci each time is too expensive, however, it's fine to run it one time with a manual dispatch
|
||||
3) Run tests locally:
|
||||
- test all models, which are not too heavy, since network speed might be a bottleneck
|
||||
|
||||
"""
|
||||
if not is_ci:
|
||||
if model_desc.size_in_GB > 1:
|
||||
return False
|
||||
elif not is_manual and model_desc.model != autotest_model_name:
|
||||
return False
|
||||
return True
|
||||
|
||||
Reference in New Issue
Block a user