Compare commits

..
1 Commits
Author SHA1 Message Date
n0x29a 40e5e96212 draft 2025-02-18 17:57:33 +01:00
64 changed files with 779 additions and 8847 deletions
+1 -1
View File
@@ -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.7.4
placeholder: v0.5.1
validations:
required: true
- type: dropdown
+3 -20
View File
@@ -1,10 +1,9 @@
name: Tests
on:
pull_request:
push:
branches: [ master, main, gpu ]
workflow_dispatch:
pull_request:
env:
CARGO_TERM_COLOR: always
@@ -24,20 +23,6 @@ jobs:
- ubuntu-latest
- macos-latest
- windows-latest
exclude:
# Exclude 3.103.12 for macOS and Windows
- os: macos-latest
python-version: '3.10.x'
- os: macos-latest
python-version: '3.11.x'
- os: macos-latest
python-version: '3.12.x'
- os: windows-latest
python-version: '3.10.x'
- os: windows-latest
python-version: '3.11.x'
- os: windows-latest
python-version: '3.12.x'
runs-on: ${{ matrix.os }}
@@ -56,7 +41,5 @@ 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
+2
View File
@@ -26,6 +26,8 @@ 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 \
+20 -75
View File
@@ -63,23 +63,6 @@ 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
@@ -154,27 +137,6 @@ 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
@@ -190,23 +152,6 @@ 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.
@@ -246,36 +191,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, models
from qdrant_client import QdrantClient
# Initialize the client
client = QdrantClient("localhost", port=6333) # For production
# client = QdrantClient(":memory:") # For experimentation
# client = QdrantClient(":memory:") # For small experiments
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"},
# 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"},
]
docs = [models.Document(text=data["document"], model=model_name) for data in payload]
ids = [42, 2]
client.create_collection(
"demo_collection",
vectors_config=models.VectorParams(
size=client.get_embedding_size(model_name), distance=models.Distance.COSINE)
# 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.upload_collection(
search_result = client.query(
collection_name="demo_collection",
vectors=docs,
ids=ids,
payload=payload,
query_text="This is a query document"
)
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)
```
+1 -13
View File
@@ -1,5 +1,4 @@
from dataclasses import dataclass, field
from enum import Enum
from typing import Optional, Any
@@ -7,11 +6,6 @@ from typing import Optional, Any
class ModelSource:
hf: Optional[str] = None
url: Optional[str] = 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:
@@ -34,7 +28,7 @@ class BaseModelDescription:
@dataclass(frozen=True)
class DenseModelDescription(BaseModelDescription):
dim: Optional[int] = None
tasks: Optional[dict[str, Any]] = field(default_factory=dict)
tasks: Optional[dict[str, Any]] = None
def __post_init__(self) -> None:
assert self.dim is not None, "dim is required for dense model description"
@@ -44,9 +38,3 @@ class DenseModelDescription(BaseModelDescription):
class SparseModelDescription(BaseModelDescription):
requires_idf: Optional[bool] = None
vocab_size: Optional[int] = None
class PoolingType(str, Enum):
CLS = "CLS"
MEAN = "MEAN"
DISABLED = "DISABLED"
+12 -54
View File
@@ -3,7 +3,6 @@ import time
import json
import shutil
import tarfile
from copy import deepcopy
from pathlib import Path
from typing import Any, Optional, Union, TypeVar, Generic
@@ -34,31 +33,6 @@ 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()
@@ -225,6 +199,11 @@ 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,
@@ -326,10 +305,9 @@ class ModelManagement(Generic[T]):
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-' if deprecated_tar_struct else ''}{model_name.split('/')[-1]}"
fast_model_name = f"fast-{model_name.split('/')[-1]}"
cache_tmp_dir = Path(cache_dir) / "tmp"
model_tmp_dir = cache_tmp_dir / fast_model_name
model_dir = Path(cache_dir) / fast_model_name
@@ -404,32 +382,14 @@ class ModelManagement(Generic[T]):
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
return Path(
cls.download_files_from_huggingface(
hf_source,
cache_dir=cache_dir,
extra_patterns=extra_patterns,
**cache_kwargs,
)
)
except Exception:
pass
finally:
enable_progress_bars()
sleep = 3.0
while retries > 0:
retries -= 1
if hf_source and not local_files_only:
# we have already tried loading with `local_files_only=True` via hf and we failed
if hf_source:
extra_patterns = [model.model_file]
extra_patterns.extend(model.additional_files)
try:
return Path(
cls.download_files_from_huggingface(
@@ -453,7 +413,6 @@ 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:
@@ -462,12 +421,11 @@ 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 -48
View File
@@ -24,22 +24,11 @@ class OnnxOutputContext:
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, **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.
"""
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[T]:
raise NotImplementedError("Subclasses must implement this method")
def __init__(self) -> None:
@@ -62,7 +51,6 @@ class OnnxModel(Generic[T]):
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_id: Optional[int] = None,
extra_session_options: Optional[dict[str, Any]] = None,
) -> None:
model_path = model_dir / model_file
# List of Execution Providers: https://onnxruntime.ai/docs/execution-providers
@@ -102,9 +90,6 @@ 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
)
@@ -119,38 +104,6 @@ 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")
+3 -3
View File
@@ -36,9 +36,9 @@ def load_tokenizer(model_dir: Path) -> tuple[Tokenizer, dict[str, int]]:
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."
)
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:
-1
View File
@@ -16,7 +16,6 @@ ImageInput: TypeAlias = Union[PathInput, Image.Image]
OnnxProvider: TypeAlias = Union[str, tuple[str, dict[Any, Any]]]
NumpyArray = Union[
NDArray[np.float64],
NDArray[np.float32],
NDArray[np.float16],
NDArray[np.int8],
-10
View File
@@ -8,7 +8,6 @@ from itertools import islice
from typing import Iterable, Optional, TypeVar
import numpy as np
from numpy.typing import NDArray
from fastembed.common.types import NumpyArray
@@ -23,15 +22,6 @@ 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 iter_batch(iterable: Iterable[T], size: int) -> Iterable[list[T]]:
"""
>>> list(iter_batch([1,2,3,4,5], 3))
+16
View File
@@ -0,0 +1,16 @@
{
"models": [
{
"model": "BAAI/bge-base-en",
"dim": 768,
"description": "Text embeddings, Unimodal (text), English...",
"license": "mit",
"size_in_GB": 0.42,
"sources": {
"hf": "Qdrant/fast-bge-base-en",
"url": "https://storage.googleapis.com/qdrant-fastembed/fast-bge-base-en.tar.gz"
},
"model_file": "model_optimized.onnx"
}
]
}
View File
-34
View File
@@ -77,40 +77,6 @@ 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: Optional[int] = 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]],
-11
View File
@@ -18,7 +18,6 @@ class ImageEmbeddingBase(ModelManagement[DenseModelDescription]):
self.cache_dir = cache_dir
self.threads = threads
self._local_files_only = kwargs.pop("local_files_only", False)
self._embedding_size: Optional[int] = None
def embed(
self,
@@ -43,13 +42,3 @@ 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")
+4 -11
View File
@@ -1,5 +1,6 @@
from typing import Any, Iterable, Optional, Sequence, Type, Union
import numpy as np
from fastembed.common.types import NumpyArray
from fastembed.common import ImageInput, OnnxProvider
@@ -98,7 +99,6 @@ 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
@@ -113,12 +113,11 @@ 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=self._specific_model_path,
specific_model_path=specific_model_path,
)
if not self.lazy_load:
@@ -135,7 +134,6 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
providers=self.providers,
cuda=self.cuda,
device_id=self.device_id,
extra_session_options=self._extra_session_options,
)
@classmethod
@@ -180,9 +178,6 @@ 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,
)
@@ -199,10 +194,8 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
return onnx_input
def _post_process_onnx_output(
self, output: OnnxOutputContext, **kwargs: Any
) -> Iterable[NumpyArray]:
return normalize(output.model_output)
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[NumpyArray]:
return normalize(output.model_output).astype(np.float32)
class OnnxImageEmbeddingWorker(ImageEmbeddingWorker[NumpyArray]):
+3 -22
View File
@@ -23,16 +23,7 @@ class OnnxImageModel(OnnxModel[T]):
def _get_worker_class(cls) -> Type["ImageEmbeddingWorker[T]"]:
raise NotImplementedError("Subclasses must implement this method")
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.
"""
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[T]:
raise NotImplementedError("Subclasses must implement this method")
def __init__(self) -> None:
@@ -55,7 +46,6 @@ class OnnxImageModel(OnnxModel[T]):
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_id: Optional[int] = None,
extra_session_options: Optional[dict[str, Any]] = None,
) -> None:
super()._load_onnx_model(
model_dir=model_dir,
@@ -64,7 +54,6 @@ 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)
@@ -99,9 +88,6 @@ class OnnxImageModel(OnnxModel[T]):
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
local_files_only: bool = False,
specific_model_path: Optional[str] = None,
extra_session_options: Optional[dict[str, Any]] = None,
**kwargs: Any,
) -> Iterable[T]:
is_small = False
@@ -118,7 +104,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), **kwargs)
yield from self._post_process_onnx_output(self.onnx_embed(batch))
else:
if parallel == 0:
parallel = os.cpu_count()
@@ -128,14 +114,9 @@ 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(),
@@ -144,7 +125,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, **kwargs) # type: ignore
yield from self._post_process_onnx_output(batch) # type: ignore
class ImageEmbeddingWorker(EmbeddingWorker[T]):
+10 -10
View File
@@ -72,26 +72,26 @@ def normalize(
if not np.issubdtype(image.dtype, np.floating):
image = image.astype(np.float32)
mean_list = mean if isinstance(mean, list) else [mean] * num_channels
mean = mean if isinstance(mean, list) else [mean] * num_channels
if len(mean_list) != num_channels:
if len(mean) != num_channels:
raise ValueError(
f"mean must have the same number of channels as the image, image has {num_channels} channels, got "
f"{len(mean_list)}"
f"{len(mean)}"
)
mean_arr = np.array(mean_list, dtype=np.float32)
mean_arr = np.array(mean, dtype=np.float32)
std_list = std if isinstance(std, list) else [std] * num_channels
if len(std_list) != num_channels:
std = std if isinstance(std, list) else [std] * num_channels
if len(std) != num_channels:
raise ValueError(
f"std must have the same number of channels as the image, image has {num_channels} channels, got {len(std_list)}"
f"std must have the same number of channels as the image, image has {num_channels} channels, got {len(std)}"
)
std_arr = np.array(std_list, dtype=np.float32)
std_arr = np.array(std, dtype=np.float32)
image_upd = ((image.T - mean_arr) / std_arr).T
return image_upd
image = ((image.T - mean_arr) / std_arr).T
return image
def resize(
+35 -72
View File
@@ -2,13 +2,12 @@ import string
from typing import Any, Iterable, Optional, Sequence, Type, Union
import numpy as np
from tokenizers import Encoding, Tokenizer
from tokenizers import Encoding
from fastembed.common.preprocessor_utils import load_tokenizer
from fastembed.common.types import NumpyArray
from fastembed.common import OnnxProvider
from fastembed.common.onnx_model import OnnxOutputContext
from fastembed.common.utils import define_cache_dir, iter_batch
from fastembed.common.utils import define_cache_dir
from fastembed.late_interaction.late_interaction_embedding_base import (
LateInteractionTextEmbeddingBase,
)
@@ -44,29 +43,26 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
MASK_TOKEN = "[MASK]"
def _post_process_onnx_output(
self, output: OnnxOutputContext, is_doc: bool = True, **kwargs: Any
self, output: OnnxOutputContext, is_doc: bool = True
) -> Iterable[NumpyArray]:
if not is_doc:
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"
)
return output.model_output.astype(np.float32)
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
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"
)
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
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 embedding, attention_mask in zip(output.model_output, output.attention_mask):
yield embedding[attention_mask == 1]
output.model_output *= np.expand_dims(output.attention_mask, 2).astype(np.float32)
norm = np.linalg.norm(output.model_output, ord=2, axis=2, keepdims=True)
norm_clamped = np.maximum(norm, 1e-12)
output.model_output /= norm_clamped
return output.model_output.astype(np.float32)
def _preprocess_onnx_input(
self, onnx_input: dict[str, NumpyArray], is_doc: bool = True, **kwargs: Any
@@ -88,46 +84,29 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
)
def _tokenize_query(self, query: str) -> list[Encoding]:
assert self.query_tokenizer is not None
encoded = self.query_tokenizer.encode_batch([query])
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)
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: Union[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.
@@ -175,7 +154,6 @@ 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
@@ -191,19 +169,16 @@ 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=self._specific_model_path,
specific_model_path=specific_model_path,
)
self.mask_token_id: Optional[int] = None
self.pad_token_id: Optional[int] = None
self.skip_list: set[int] = set()
self.query_tokenizer: Optional[Tokenizer] = None
if not self.lazy_load:
self.load_onnx_model()
@@ -215,10 +190,7 @@ 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"]
@@ -229,12 +201,6 @@ 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,
@@ -267,9 +233,6 @@ 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,
)
@@ -17,7 +17,6 @@ class LateInteractionTextEmbeddingBase(ModelManagement[DenseModelDescription]):
self.cache_dir = cache_dir
self.threads = threads
self._local_files_only = kwargs.pop("local_files_only", False)
self._embedding_size: Optional[int] = None
def embed(
self,
@@ -59,22 +58,3 @@ 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: Union[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")
@@ -80,40 +80,6 @@ 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: Optional[int] = 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]],
@@ -151,30 +117,3 @@ 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: Union[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,
)
@@ -1,83 +0,0 @@
from dataclasses import asdict
from typing import Union, Iterable, Optional, 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: Union[str, Iterable[str]],
batch_size: int = 256,
parallel: Optional[int] = 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,
)
@@ -6,7 +6,7 @@ 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, iter_batch
from fastembed.common.utils import define_cache_dir
from fastembed.late_interaction_multimodal.late_interaction_multimodal_embedding_base import (
LateInteractionMultimodalEmbeddingBase,
)
@@ -80,7 +80,6 @@ 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
@@ -96,12 +95,11 @@ 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=self._specific_model_path,
specific_model_path=specific_model_path,
)
self.mask_token_id = None
self.pad_token_id = None
@@ -126,7 +124,6 @@ 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(
@@ -145,7 +142,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,
@@ -160,7 +157,7 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
Returns:
Iterable[NumpyArray]: Post-processed output as NumPy arrays.
"""
return output.model_output
return output.model_output.astype(np.float32)
def tokenize(self, documents: list[str], **kwargs: Any) -> list[Encoding]:
texts_query: list[str] = []
@@ -172,29 +169,12 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
encoded = self.tokenizer.encode_batch(texts_query) # type: ignore[union-attr]
return encoded
def token_count(
self,
texts: Union[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() # type: ignore[index]
self.QUERY_MARKER_TOKEN_ID + input_ids[2:].tolist()
for input_ids in onnx_input["input_ids"]
]
)
@@ -217,11 +197,12 @@ 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["pixel_values"]]
[self.EMPTY_TEXT_PLACEHOLDER for _ in onnx_input["input_ids"]]
)
onnx_input["attention_mask"] = np.array(
[self.EVEN_ATTENTION_MASK for _ in onnx_input["pixel_values"]]
[self.EVEN_ATTENTION_MASK for _ in onnx_input["input_ids"]]
)
return onnx_input
@@ -255,9 +236,6 @@ 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,
)
@@ -291,9 +269,6 @@ 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,
)
@@ -83,40 +83,6 @@ 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: Optional[int] = 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]],
@@ -162,24 +128,3 @@ 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: Union[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
)
@@ -19,7 +19,6 @@ class LateInteractionMultimodalEmbeddingBase(ModelManagement[DenseModelDescripti
self.cache_dir = cache_dir
self.threads = threads
self._local_files_only = kwargs.pop("local_files_only", False)
self._embedding_size: Optional[int] = None
def embed_text(
self,
@@ -66,21 +65,3 @@ 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: Union[str, Iterable[str]],
**kwargs: Any,
) -> int:
"""Returns the number of tokens in the texts."""
raise NotImplementedError("Subclasses must implement this method")
@@ -64,7 +64,6 @@ class OnnxMultimodalModel(OnnxModel[T]):
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_id: Optional[int] = None,
extra_session_options: Optional[dict[str, Any]] = None,
) -> None:
super()._load_onnx_model(
model_dir=model_dir,
@@ -73,10 +72,9 @@ class OnnxMultimodalModel(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)
assert self.tokenizer is not None
self.tokenizer, self.special_token_to_id = load_tokenizer(model_dir=model_dir)
self.processor = load_preprocessor(model_dir=model_dir)
def load_onnx_model(self) -> None:
@@ -122,9 +120,6 @@ class OnnxMultimodalModel(OnnxModel[T]):
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
local_files_only: bool = False,
specific_model_path: Optional[str] = None,
extra_session_options: Optional[dict[str, Any]] = None,
**kwargs: Any,
) -> Iterable[T]:
is_small = False
@@ -151,14 +146,9 @@ 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(),
@@ -169,6 +159,10 @@ 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():
image_files = [
@@ -177,7 +171,7 @@ class OnnxMultimodalModel(OnnxModel[T]):
]
assert self.processor is not None, "Processor is not initialized"
encoded = np.array(self.processor(image_files))
onnx_input = {"pixel_values": encoded}
onnx_input = self._build_onnx_image_input(encoded)
onnx_input = self._preprocess_onnx_image_input(onnx_input, **kwargs)
model_output = self.model.run(None, onnx_input) # type: ignore[union-attr]
embeddings = model_output[0].reshape(len(images), -1)
@@ -193,9 +187,6 @@ class OnnxMultimodalModel(OnnxModel[T]):
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
local_files_only: bool = False,
specific_model_path: Optional[str] = None,
extra_session_options: Optional[dict[str, Any]] = None,
**kwargs: Any,
) -> Iterable[T]:
is_small = False
@@ -222,14 +213,9 @@ 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(),
+16
View File
@@ -0,0 +1,16 @@
from pathlib import Path
import json
from typing import Dict, List
class ModelLoader:
def __init__(self):
self.config_dir = Path(__file__).parent / "configs"
self._models: Dict[str, List[Dict]] = {}
def load_models(self, model_type: str) -> List[Dict]:
if model_type not in self._models:
config_path = self.config_dir / f"{model_type}_models.json"
with open(config_path) as f:
self._models[model_type] = json.load(f)["models"]
return self._models[model_type]
-3
View File
@@ -1,3 +0,0 @@
from fastembed.postprocess.muvera import Muvera
__all__ = ["Muvera"]
-364
View File
@@ -1,364 +0,0 @@
from typing import Union
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 = Union[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
@@ -1,46 +0,0 @@
from typing import Optional, Sequence, Any
from fastembed.common import OnnxProvider
from fastembed.common.model_description import BaseModelDescription
from fastembed.rerank.cross_encoder.onnx_text_cross_encoder import OnnxTextCrossEncoder
class CustomTextCrossEncoder(OnnxTextCrossEncoder):
SUPPORTED_MODELS: list[BaseModelDescription] = []
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,
):
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 add_model(
cls,
model_description: BaseModelDescription,
) -> None:
cls.SUPPORTED_MODELS.append(model_description)
@@ -111,7 +111,6 @@ 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
@@ -132,12 +131,11 @@ 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=self._specific_model_path,
specific_model_path=specific_model_path,
)
if not self.lazy_load:
@@ -151,7 +149,6 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
providers=self.providers,
cuda=self.cuda,
device_id=self.device_id,
extra_session_options=self._extra_session_options,
)
def rerank(
@@ -192,9 +189,6 @@ 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,
)
@@ -202,25 +196,9 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
def _get_worker_class(cls) -> Type[TextRerankerWorker]:
return TextCrossEncoderWorker
def _post_process_onnx_output(
self, output: OnnxOutputContext, **kwargs: Any
) -> Iterable[float]:
def _post_process_onnx_output(self, output: OnnxOutputContext) -> 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(
@@ -33,7 +33,6 @@ class OnnxCrossEncoderModel(OnnxModel[float]):
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_id: Optional[int] = None,
extra_session_options: Optional[dict[str, Any]] = None,
) -> None:
super()._load_onnx_model(
model_dir=model_dir,
@@ -42,7 +41,6 @@ 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
@@ -96,9 +94,6 @@ class OnnxCrossEncoderModel(OnnxModel[float]):
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
local_files_only: bool = False,
specific_model_path: Optional[str] = None,
extra_session_options: Optional[dict[str, Any]] = None,
**kwargs: Any,
) -> Iterable[float]:
is_small = False
@@ -125,14 +120,9 @@ 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,
}
if extra_session_options is not None:
params.update(extra_session_options)
pool = ParallelWorkerPool(
num_workers=parallel or 1,
worker=self._get_worker_class(),
@@ -143,18 +133,7 @@ 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, **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.
"""
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[float]:
raise NotImplementedError("Subclasses must implement this method")
def _preprocess_onnx_input(
@@ -165,20 +144,6 @@ 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__(
@@ -3,19 +3,13 @@ from dataclasses import asdict
from fastembed.common import OnnxProvider
from fastembed.rerank.cross_encoder.onnx_text_cross_encoder import OnnxTextCrossEncoder
from fastembed.rerank.cross_encoder.custom_text_cross_encoder import CustomTextCrossEncoder
from fastembed.rerank.cross_encoder.text_cross_encoder_base import TextCrossEncoderBase
from fastembed.common.model_description import (
ModelSource,
BaseModelDescription,
)
from fastembed.common.model_description import BaseModelDescription
class TextCrossEncoder(TextCrossEncoderBase):
CROSS_ENCODER_REGISTRY: list[Type[TextCrossEncoderBase]] = [
OnnxTextCrossEncoder,
CustomTextCrossEncoder,
]
@classmethod
@@ -130,48 +124,3 @@ 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: Optional[list[str]] = 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)
@@ -57,7 +57,3 @@ 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")
+13 -19
View File
@@ -21,9 +21,13 @@ 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",
@@ -31,15 +35,21 @@ supported_languages = [
"french",
"german",
"greek",
"hebrew",
"hinglish",
"hungarian",
"indonesian",
"italian",
"kazakh",
"nepali",
"norwegian",
"portuguese",
"romanian",
"russian",
"slovene",
"spanish",
"swedish",
"tamil",
"tajik",
"turkish",
]
@@ -115,12 +125,11 @@ 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=self._specific_model_path,
specific_model_path=specific_model_path,
)
self.token_max_length = token_max_length
@@ -161,8 +170,6 @@ class Bm25(SparseTextEmbeddingBase):
documents: Union[str, Iterable[str]],
batch_size: int = 256,
parallel: Optional[int] = None,
local_files_only: bool = False,
specific_model_path: Optional[str] = None,
) -> Iterable[SparseEmbedding]:
is_small = False
@@ -191,8 +198,6 @@ 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,
@@ -231,8 +236,6 @@ 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]:
@@ -268,15 +271,6 @@ class Bm25(SparseTextEmbeddingBase):
embeddings.append(SparseEmbedding.from_dict(token_id2value))
return embeddings
def token_count(self, texts: Union[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.
+4 -27
View File
@@ -31,17 +31,9 @@ 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]):
@@ -103,7 +95,6 @@ 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
@@ -119,12 +110,11 @@ 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=self._specific_model_path,
specific_model_path=specific_model_path,
)
self.invert_vocab: dict[int, str] = {}
@@ -133,7 +123,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(get_language_by_model_name(self.model_name))
self.stemmer = SnowballStemmer(MODEL_TO_LANGUAGE[model_name])
self.alpha = alpha
if not self.lazy_load:
@@ -147,7 +137,6 @@ 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]
@@ -228,9 +217,7 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
return new_vector
def _post_process_onnx_output(
self, output: OnnxOutputContext, **kwargs: Any
) -> Iterable[SparseEmbedding]:
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[SparseEmbedding]:
if output.input_ids is None:
raise ValueError("input_ids must be provided for document post-processing")
@@ -312,9 +299,6 @@ 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
@@ -352,13 +336,6 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
def _get_worker_class(cls) -> Type[TextEmbeddingWorker[SparseEmbedding]]:
return Bm42TextEmbeddingWorker
def token_count(
self, texts: Union[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:
-372
View File
@@ -1,372 +0,0 @@
from pathlib import Path
from typing import Any, Optional, Sequence, Iterable, Union, 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.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: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
k: float = 1.2,
b: float = 0.75,
avg_len: float = 150.0,
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 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 (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to False.
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
specific_model_path (Optional[str], optional): The specific path to the onnx model dir if it should be imported from somewhere else
Raises:
ValueError: If the model_name is not in the format <org>/<model> e.g. BAAI/bge-base-en.
"""
super().__init__(model_name, cache_dir, threads, **kwargs)
self.providers = providers
self.lazy_load = lazy_load
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: Optional[Tokenizer] = 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: Optional[VocabResolver] = None
self.encoder: Optional[Encoder] = None
self.output_dim: Optional[int] = None
self.sparse_vector_converter: Optional[SparseVectorConverter] = 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: Union[str, Iterable[str]], batch_size: int = 1024, **kwargs: Any
) -> int:
return self._token_count(texts, batch_size=batch_size, **kwargs)
def embed(
self,
documents: Union[str, Iterable[str]],
batch_size: int = 256,
parallel: Optional[int] = 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: Union[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,
)
@@ -86,7 +86,3 @@ class SparseTextEmbeddingBase(ModelManagement[SparseModelDescription]):
yield from self.embed([query], **kwargs)
else:
yield from self.embed(query, **kwargs)
def token_count(self, texts: Union[str, Iterable[str]], **kwargs: Any) -> int:
"""Returns the number of tokens in the texts."""
raise NotImplementedError("Subclasses must implement this method")
+2 -17
View File
@@ -4,7 +4,6 @@ from dataclasses import asdict
from fastembed.common import OnnxProvider
from fastembed.sparse.bm25 import Bm25
from fastembed.sparse.bm42 import Bm42
from fastembed.sparse.minicoil import MiniCOIL
from fastembed.sparse.sparse_embedding_base import (
SparseEmbedding,
SparseTextEmbeddingBase,
@@ -15,7 +14,7 @@ from fastembed.common.model_description import SparseModelDescription
class SparseTextEmbedding(SparseTextEmbeddingBase):
EMBEDDINGS_REGISTRY: list[Type[SparseTextEmbeddingBase]] = [SpladePP, Bm42, Bm25, MiniCOIL]
EMBEDDINGS_REGISTRY: list[Type[SparseTextEmbeddingBase]] = [SpladePP, Bm42, Bm25]
@classmethod
def list_supported_models(cls) -> list[dict[str, Any]]:
@@ -62,7 +61,7 @@ class SparseTextEmbedding(SparseTextEmbeddingBase):
**kwargs: Any,
):
super().__init__(model_name, cache_dir, threads, **kwargs)
if model_name.lower() == "prithvida/Splade_PP_en_v1".lower():
if model_name == "prithvida/Splade_PP_en_v1":
warnings.warn(
"The right spelling is prithivida/Splade_PP_en_v1. "
"Support of this name will be removed soon, please fix the model_name",
@@ -128,17 +127,3 @@ class SparseTextEmbedding(SparseTextEmbeddingBase):
Iterable[SparseEmbedding]: The sparse embeddings.
"""
yield from self.model.query_embed(query, **kwargs)
def token_count(
self, texts: Union[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)
+4 -17
View File
@@ -18,7 +18,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,16 +27,14 @@ 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, **kwargs: Any
) -> Iterable[SparseEmbedding]:
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[SparseEmbedding]:
if output.attention_mask is None:
raise ValueError("attention_mask must be provided for document post-processing")
@@ -53,11 +51,6 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
scores = row_scores[indices]
yield SparseEmbedding(values=scores, indices=indices)
def token_count(
self, texts: Union[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.
@@ -104,7 +97,6 @@ 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
@@ -120,12 +112,11 @@ 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=self._specific_model_path,
specific_model_path=specific_model_path,
)
if not self.lazy_load:
@@ -139,7 +130,6 @@ 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(
@@ -173,9 +163,6 @@ 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,
)
-146
View File
@@ -1,146 +0,0 @@
"""
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
@@ -1,247 +0,0 @@
from typing import Dict, List, Set
from py_rust_stemmers import SnowballStemmer
from fastembed.common.utils import get_all_punctuation, remove_non_alphanumeric
import mmh3
import copy
from dataclasses import dataclass
import numpy as np
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 = {"": {"word": "", "word_id": -1, "count": 2, "embedding": [1], "forms": [""]},
"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"]},
"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),
)
-202
View File
@@ -1,202 +0,0 @@
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
+1 -3
View File
@@ -35,9 +35,7 @@ class CLIPOnnxEmbedding(OnnxTextEmbedding):
"""
return supported_clip_models
def _post_process_onnx_output(
self, output: OnnxOutputContext, **kwargs: Any
) -> Iterable[NumpyArray]:
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[NumpyArray]:
return output.model_output
-98
View File
@@ -1,98 +0,0 @@
from typing import Optional, Sequence, Any, Iterable
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
from fastembed.common.utils import normalize, mean_pooling
from fastembed.text.onnx_embedding import OnnxTextEmbedding
@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: 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,
):
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,
)
self._pooling = self.POSTPROCESSING_MAPPING[model_name].pooling
self._normalization = self.POSTPROCESSING_MAPPING[model_name].normalization
@classmethod
def _list_supported_models(cls) -> list[DenseModelDescription]:
return cls.SUPPORTED_MODELS
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: Optional[NDArray[np.int64]] = 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.DISABLED:
return embeddings
raise ValueError(
f"Unsupported pooling type {self._pooling}. "
f"Supported types are: {PoolingType.CLS}, {PoolingType.MEAN}, {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
)
+15 -26
View File
@@ -3,7 +3,6 @@ from typing import Any, Type, Iterable, Union, Optional
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
@@ -45,11 +44,9 @@ class JinaEmbeddingV3(PooledNormalizedEmbedding):
PASSAGE_TASK = Task.RETRIEVAL_PASSAGE
QUERY_TASK = Task.RETRIEVAL_QUERY
def __init__(self, *args: Any, task_id: Optional[int] = None, **kwargs: Any):
def __init__(self, *args: Any, **kwargs: Any):
super().__init__(*args, **kwargs)
self.default_task_id: Union[Task, int] = (
task_id if task_id is not None else self.PASSAGE_TASK
)
self.current_task_id: Union[Task, int] = self.PASSAGE_TASK
@classmethod
def _get_worker_class(cls) -> Type[OnnxTextEmbeddingWorker]:
@@ -60,14 +57,9 @@ class JinaEmbeddingV3(PooledNormalizedEmbedding):
return supported_multitask_models
def _preprocess_onnx_input(
self,
onnx_input: dict[str, NumpyArray],
task_id: Optional[Union[int, Task]] = None,
**kwargs: Any,
self, onnx_input: dict[str, NumpyArray], **kwargs: Any
) -> dict[str, NumpyArray]:
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)
onnx_input["task_id"] = np.array(self.current_task_id, dtype=np.int64)
return onnx_input
def embed(
@@ -75,19 +67,20 @@ class JinaEmbeddingV3(PooledNormalizedEmbedding):
documents: Union[str, Iterable[str]],
batch_size: int = 256,
parallel: Optional[int] = None,
task_id: Optional[int] = None,
task_id: int = PASSAGE_TASK,
**kwargs: Any,
) -> Iterable[NumpyArray]:
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)
self.current_task_id = task_id
kwargs["task_id"] = task_id
yield from super().embed(documents, batch_size, parallel, **kwargs)
def query_embed(self, query: Union[str, Iterable[str]], **kwargs: Any) -> Iterable[NumpyArray]:
yield from super().embed(query, task_id=self.QUERY_TASK, **kwargs)
self.current_task_id = self.QUERY_TASK
yield from super().embed(query, **kwargs)
def passage_embed(self, texts: Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:
yield from super().embed(texts, task_id=self.PASSAGE_TASK, **kwargs)
self.current_task_id = self.PASSAGE_TASK
yield from super().embed(texts, **kwargs)
class JinaEmbeddingV3Worker(OnnxTextEmbeddingWorker):
@@ -97,15 +90,11 @@ class JinaEmbeddingV3Worker(OnnxTextEmbeddingWorker):
cache_dir: str,
**kwargs: Any,
) -> JinaEmbeddingV3:
return JinaEmbeddingV3(
model = JinaEmbeddingV3(
model_name=model_name,
cache_dir=cache_dir,
threads=1,
**kwargs,
)
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
model.current_task_id = kwargs["task_id"]
return model
+17 -21
View File
@@ -1,5 +1,6 @@
from typing import Any, Iterable, Optional, Sequence, Type, Union
import numpy as np
from fastembed.common.types import NumpyArray, OnnxProvider
from fastembed.common.onnx_model import OnnxOutputContext
from fastembed.common.utils import define_cache_dir, normalize
@@ -20,7 +21,6 @@ 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",
),
@@ -36,7 +36,6 @@ supported_onnx_models: list[DenseModelDescription] = [
sources=ModelSource(
hf="qdrant/bge-base-en-v1.5-onnx-q",
url="https://storage.googleapis.com/qdrant-fastembed/fast-bge-base-en-v1.5.tar.gz",
_deprecated_tar_struct=True,
),
model_file="model_optimized.onnx",
),
@@ -64,7 +63,6 @@ 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",
),
@@ -92,10 +90,21 @@ 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,
@@ -233,7 +242,7 @@ 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
@@ -247,12 +256,11 @@ 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=self._specific_model_path,
specific_model_path=specific_model_path,
)
if not self.lazy_load:
@@ -289,9 +297,6 @@ 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,
)
@@ -307,18 +312,15 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[NumpyArray]):
"""
return onnx_input
def _post_process_onnx_output(
self, output: OnnxOutputContext, **kwargs: Any
) -> Iterable[NumpyArray]:
def _post_process_onnx_output(self, output: OnnxOutputContext) -> 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)
return normalize(processed_embeddings).astype(np.float32)
def load_onnx_model(self) -> None:
self._load_onnx_model(
@@ -328,14 +330,8 @@ 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: Union[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(
+3 -39
View File
@@ -21,16 +21,7 @@ class OnnxTextModel(OnnxModel[T]):
def _get_worker_class(cls) -> Type["TextEmbeddingWorker[T]"]:
raise NotImplementedError("Subclasses must implement this method")
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.
"""
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[T]:
raise NotImplementedError("Subclasses must implement this method")
def __init__(self) -> None:
@@ -54,7 +45,6 @@ class OnnxTextModel(OnnxModel[T]):
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_id: Optional[int] = None,
extra_session_options: Optional[dict[str, Any]] = None,
) -> None:
super()._load_onnx_model(
model_dir=model_dir,
@@ -63,7 +53,6 @@ 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)
@@ -110,9 +99,6 @@ class OnnxTextModel(OnnxModel[T]):
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
local_files_only: bool = False,
specific_model_path: Optional[str] = None,
extra_session_options: Optional[dict[str, Any]] = None,
**kwargs: Any,
) -> Iterable[T]:
is_small = False
@@ -129,9 +115,7 @@ 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, **kwargs), **kwargs
)
yield from self._post_process_onnx_output(self.onnx_embed(batch))
else:
if parallel == 0:
parallel = os.cpu_count()
@@ -141,14 +125,9 @@ 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,
}
if extra_session_options is not None:
params.update(extra_session_options)
pool = ParallelWorkerPool(
num_workers=parallel or 1,
worker=self._get_worker_class(),
@@ -157,22 +136,7 @@ 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, **kwargs) # type: ignore
def _token_count(
self, texts: Union[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
yield from self._post_process_onnx_output(batch) # type: ignore
class TextEmbeddingWorker(EmbeddingWorker[T]):
+12 -11
View File
@@ -1,11 +1,9 @@
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
@@ -82,7 +80,6 @@ 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"],
@@ -96,10 +93,16 @@ class PooledEmbedding(OnnxTextEmbedding):
return PooledEmbeddingWorker
@classmethod
def mean_pooling(
cls, model_output: NumpyArray, attention_mask: NDArray[np.int64]
) -> NumpyArray:
return mean_pooling(model_output, attention_mask)
def mean_pooling(cls, model_output: NumpyArray, attention_mask: NumpyArray) -> NumpyArray:
token_embeddings = model_output.astype(np.float32)
attention_mask = attention_mask.astype(np.float32)
input_mask_expanded = np.expand_dims(attention_mask, axis=-1)
input_mask_expanded = np.tile(input_mask_expanded, (1, 1, token_embeddings.shape[-1]))
input_mask_expanded = input_mask_expanded.astype(np.float32)
sum_embeddings = np.sum(token_embeddings * input_mask_expanded, axis=1)
sum_mask = np.sum(input_mask_expanded, axis=1)
pooled_embeddings = sum_embeddings / np.maximum(sum_mask, 1e-9)
return pooled_embeddings
@classmethod
def _list_supported_models(cls) -> list[DenseModelDescription]:
@@ -110,15 +113,13 @@ class PooledEmbedding(OnnxTextEmbedding):
"""
return supported_pooled_models
def _post_process_onnx_output(
self, output: OnnxOutputContext, **kwargs: Any
) -> Iterable[NumpyArray]:
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[NumpyArray]:
if output.attention_mask is None:
raise ValueError("attention_mask must be provided for document post-processing")
embeddings = output.model_output
attn_mask = output.attention_mask
return self.mean_pooling(embeddings, attn_mask)
return self.mean_pooling(embeddings, attn_mask).astype(np.float32)
class PooledEmbeddingWorker(OnnxTextEmbeddingWorker):
+3 -17
View File
@@ -1,5 +1,6 @@
from typing import Any, Iterable, Type
import numpy as np
from fastembed.common.types import NumpyArray
from fastembed.common.onnx_model import OnnxOutputContext
@@ -21,7 +22,6 @@ 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",
),
@@ -109,18 +109,6 @@ 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",
),
]
@@ -138,15 +126,13 @@ class PooledNormalizedEmbedding(PooledEmbedding):
"""
return supported_pooled_normalized_models
def _post_process_onnx_output(
self, output: OnnxOutputContext, **kwargs: Any
) -> Iterable[NumpyArray]:
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[NumpyArray]:
if output.attention_mask is None:
raise ValueError("attention_mask must be provided for document post-processing")
embeddings = output.model_output
attn_mask = output.attention_mask
return normalize(self.mean_pooling(embeddings, attn_mask))
return normalize(self.mean_pooling(embeddings, attn_mask)).astype(np.float32)
class PooledNormalizedEmbeddingWorker(OnnxTextEmbeddingWorker):
+17 -99
View File
@@ -4,13 +4,12 @@ from dataclasses import asdict
from fastembed.common.types import NumpyArray, OnnxProvider
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.onnx_embedding import OnnxTextEmbedding
from fastembed.text.text_embedding_base import TextEmbeddingBase
from fastembed.common.model_description import DenseModelDescription, ModelSource, PoolingType
from fastembed.common.model_description import DenseModelDescription
class TextEmbedding(TextEmbeddingBase):
@@ -20,7 +19,6 @@ class TextEmbedding(TextEmbeddingBase):
PooledNormalizedEmbedding,
PooledEmbedding,
JinaEmbeddingV3,
CustomTextEmbedding,
]
@classmethod
@@ -39,43 +37,6 @@ 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: Optional[list[str]] = 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",
@@ -88,26 +49,31 @@ class TextEmbedding(TextEmbeddingBase):
**kwargs: Any,
):
super().__init__(model_name, cache_dir, threads, **kwargs)
if model_name.lower() == "nomic-ai/nomic-embed-text-v1.5-Q".lower():
if model_name == "nomic-ai/nomic-embed-text-v1.5-Q":
warnings.warn(
"The model 'nomic-ai/nomic-embed-text-v1.5-Q' has been updated on HuggingFace. Please review "
"the latest documentation on HF and release notes to ensure compatibility with your workflow. ",
"The model 'nomic-ai/nomic-embed-text-v1.5-Q' has been updated on HuggingFace. "
"Please review the latest documentation and release notes to ensure compatibility with your workflow. ",
UserWarning,
stacklevel=2,
)
if model_name.lower() in {
"sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2".lower(),
"thenlper/gte-large".lower(),
"intfloat/multilingual-e5-large".lower(),
"sentence-transformers/paraphrase-multilingual-mpnet-base-v2".lower(),
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"The model {model_name} now uses mean pooling instead of CLS embedding. "
f"In order to preserve the previous behaviour, consider either pinning fastembed version to 0.5.1 or "
"using `add_custom_model` functionality.",
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):
@@ -128,40 +94,6 @@ 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: Optional[int] = 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]],
@@ -212,17 +144,3 @@ 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: Union[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)
-15
View File
@@ -17,7 +17,6 @@ class TextEmbeddingBase(ModelManagement[DenseModelDescription]):
self.cache_dir = cache_dir
self.threads = threads
self._local_files_only = kwargs.pop("local_files_only", False)
self._embedding_size: Optional[int] = None
def embed(
self,
@@ -59,17 +58,3 @@ 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: Union[str, Iterable[str]], **kwargs: Any) -> int:
"""Returns the number of tokens in the texts."""
raise NotImplementedError("Subclasses must implement this method")
Generated
-4997
View File
File diff suppressed because it is too large Load Diff
+10 -19
View File
@@ -1,6 +1,6 @@
[tool.poetry]
name = "fastembed"
version = "0.7.4"
version = "0.5.1"
description = "Fast, light, accurate library built for retrieval embedding generation"
authors = ["Qdrant Team <info@qdrant.tech>", "NirantK <nirant.bits@gmail.com>"]
license = "Apache License"
@@ -13,29 +13,23 @@ keywords = ["vector", "embedding", "neural", "search", "qdrant", "sentence-trans
[tool.poetry.dependencies]
python = ">=3.9.0"
numpy = [
{ version = ">=1.21,<2.1.0", python = "<3.10" },
{ version = ">=1.21,<2.3.0", python = ">=3.10,<3.11" },
{ version = ">=1.21", python = ">=3.11,<3.12" },
{ version = ">=1.21", python = ">=3.10,<3.12" },
{ version = ">=1.26", python = ">=3.12,<3.13" },
{ version = ">=2.1.0", python = ">=3.13,<3.14" },
{ version = ">=2.3.0", python = ">=3.14" },
{ version = ">=2.1.0", python = ">=3.13" },
{ version = ">=1.21,<2.1.0", python = "<3.10" },
]
onnxruntime = [
{ version = ">=1.17.0,<1.20.0", python = "<3.10" },
{ version = ">1.20.0", python = ">=3.13" },
{ version = ">=1.17.0,<1.20.0", python = "<3.10" },
{ version = ">=1.17.0,!=1.20.0", python = ">=3.10,<3.13" },
]
tqdm = "^4.66"
requests = "^2.31"
tokenizers = ">=0.15,<1.0"
huggingface-hub = ">=0.20,<2.0"
huggingface-hub = ">=0.20,<1.0"
loguru = "^0.7.2"
pillow = [
{ version = ">=10.3.0,<11.0", python = "<3.10" },
{ version = ">=10.3.0,<12.0", python = ">=3.10,<3.13" },
{ version = ">=11.0.0,<12.0", python = ">=3.13" },
]
mmh3 = ">=4.1.0,<6.0.0"
pillow = ">=10.3.0,<12.0.0"
mmh3 = "^4.1.0"
py-rust-stemmers = "^0.1.0"
[tool.poetry.group.test.dependencies]
@@ -45,15 +39,12 @@ ruff = ">=0.3.1,<1.0"
[tool.poetry.group.dev.dependencies]
notebook = ">=7.0.2"
pre-commit = "^3.6.2"
onnx = [
{ version = ">=1.15.0", python = "<3.13" },
{ version = ">=1.18.0", python = ">=3.13" },
]
onnx = ">=1.15.0"
[tool.poetry.group.docs.dependencies]
mkdocs-material = "^9.5.10"
mkdocstrings = "^0.24.0"
pillow = ">=10.3.0,<13.0.0"
pillow = ">=10.3.0,<12.0.0"
cairosvg = "^2.7.1"
mknotebooks = "^0.8.0"
+103 -116
View File
@@ -1,5 +1,4 @@
import os
from contextlib import contextmanager
import numpy as np
import pytest
@@ -8,119 +7,98 @@ from fastembed import SparseTextEmbedding
from tests.utils import delete_model_cache
_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():
@pytest.mark.parametrize("model_name", ["Qdrant/bm42-all-minilm-l6-v2-attentions", "Qdrant/bm25"])
def test_attention_embeddings(model_name: str) -> None:
is_ci = os.getenv("CI")
cache = {}
model = SparseTextEmbedding(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] = 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
output = list(
model.query_embed(
[
"I must not fear. Fear is the mind-killer.",
]
)
)
yield get_model
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
if is_ci:
for name, model in cache.items():
delete_model_cache(model.model._model_dir)
cache.clear()
delete_model_cache(model.model._model_dir)
@pytest.mark.parametrize("model_name", ["Qdrant/bm42-all-minilm-l6-v2-attentions", "Qdrant/bm25"])
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.",
]
)
)
def test_parallel_processing(model_name: str) -> None:
is_ci = os.getenv("CI")
assert len(output) == 1
model = SparseTextEmbedding(model_name=model_name)
for result in output:
assert len(result.indices) == len(result.values)
assert np.allclose(result.values, np.ones(len(result.values)))
docs = ["hello world", "attention embedding", "Mangez-vous vraiment des grenouilles?"] * 100
embeddings = list(model.embed(docs, batch_size=10, parallel=2))
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_2 = list(model.embed(docs, batch_size=10, parallel=None))
output = list(model.embed(quotes))
embeddings_3 = list(model.embed(docs, batch_size=10, parallel=0))
assert len(output) == len(quotes)
assert len(embeddings) == len(docs)
for result in output[:-1]:
assert len(result.indices) == len(result.values)
assert len(result.indices) > 0
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)
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)
if is_ci:
delete_model_cache(model.model._model_dir)
@pytest.mark.parametrize("model_name", ["Qdrant/bm25"])
def test_multilanguage(model_cache, model_name: str) -> None:
def test_multilanguage(model_name: str) -> None:
is_ci = os.getenv("CI")
docs = ["Mangez-vous vraiment des grenouilles?", "Je suis au lit"]
model = SparseTextEmbedding(model_name=model_name, language="french")
@@ -131,30 +109,39 @@ def test_multilanguage(model_cache, model_name: str) -> None:
assert embeddings[1].values.shape == (1,)
assert embeddings[1].indices.shape == (1,)
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,)
model = SparseTextEmbedding(model_name=model_name, language="english")
embeddings = list(model.embed(docs))[:2]
assert embeddings[0].values.shape == (5,)
assert embeddings[0].indices.shape == (5,)
assert embeddings[1].values.shape == (4,)
assert embeddings[1].indices.shape == (4,)
assert embeddings[1].values.shape == (4,)
assert embeddings[1].indices.shape == (4,)
if is_ci:
delete_model_cache(model.model._model_dir)
@pytest.mark.parametrize("model_name", ["Qdrant/bm25"])
def test_special_characters(model_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,)
def test_special_characters(model_name: str) -> None:
is_ci = os.getenv("CI")
docs = [
"Über den größten Flüssen Österreichs äußern sich Experten häufig: Öko-Systeme müssen geschützt werden!",
"L'élève français s'écrie : « Où est mon crayon ? J'ai besoin de finir cet exercice avant la récréation!",
"Într-o zi însorită, Ștefan și Ioana au mâncat mămăligă cu brânză și au băut țuică la cabană.",
"Üzgün öğretmen öğrencilere seslendi: Lütfen gürültü yapmayın, sınavınızı bitirmeye çalışıyorum!",
"Ο Ξενοφών είπε: «Ψάχνω για ένα ωραίο δώρο για τη γιαγιά μου. Ίσως ένα φυτό ή ένα βιβλίο;»",
"Hola! ¿Cómo estás? Estoy muy emocionado por el cumpleaños de mi hermano, ¡va a ser increíble! También quiero comprar un pastel de chocolate con fresas y un regalo especial: un libro titulado «Cien años de soledad",
]
model = SparseTextEmbedding(model_name=model_name, language="english")
embeddings = list(model.embed(docs))
for idx, shape in enumerate([14, 18, 15, 10, 15]):
assert embeddings[idx].values.shape == (shape,)
assert embeddings[idx].indices.shape == (shape,)
if is_ci:
delete_model_cache(model.model._model_dir)
@pytest.mark.parametrize("model_name", ["Qdrant/bm42-all-minilm-l6-v2-attentions"])
-243
View File
@@ -1,243 +0,0 @@
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
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 = []
CustomTextCrossEncoder.SUPPORTED_MODELS = []
yield
CustomTextEmbedding.SUPPORTED_MODELS = []
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)
CustomTextEmbedding.SUPPORTED_MODELS.clear()
CustomTextEmbedding.POSTPROCESSING_MAPPING.clear()
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)
CustomTextCrossEncoder.SUPPORTED_MODELS.clear()
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.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.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.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)
CustomTextEmbedding.SUPPORTED_MODELS.clear()
CustomTextEmbedding.POSTPROCESSING_MAPPING.clear()
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,
)
CustomTextEmbedding.SUPPORTED_MODELS.clear()
CustomTextEmbedding.POSTPROCESSING_MAPPING.clear()
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,
)
CustomTextCrossEncoder.SUPPORTED_MODELS.clear()
+59 -111
View File
@@ -1,5 +1,4 @@
import os
from contextlib import contextmanager
from io import BytesIO
import numpy as np
@@ -9,7 +8,7 @@ from PIL import Image
from fastembed import ImageEmbedding
from tests.config import TEST_MISC_DIR
from tests.utils import delete_model_cache, should_test_model
from tests.utils import delete_model_cache
CANONICAL_VECTOR_VALUES = {
"Qdrant/clip-ViT-B-32-vision": np.array([-0.0098, 0.0128, -0.0274, 0.002, -0.0059]),
@@ -27,109 +26,86 @@ CANONICAL_VECTOR_VALUES = {
),
}
_MODELS_TO_CACHE = ("Qdrant/clip-ViT-B-32-vision",)
MODELS_TO_CACHE = tuple([x.lower() for x in _MODELS_TO_CACHE])
@pytest.fixture(scope="module")
def model_cache():
def test_embedding() -> None:
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_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
for model_desc in ImageEmbedding._list_supported_models():
if not should_test_model(model_desc, model_name, is_ci, is_manual):
if not is_ci and model_desc.size_in_GB > 1:
continue
dim = model_desc.dim
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)
model = ImageEmbedding(model_name=model_desc.model)
canonical_vector = CANONICAL_VECTOR_VALUES[model_desc.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)
assert np.allclose(
embeddings[0, : canonical_vector.shape[0]], canonical_vector, atol=1e-3
), model_desc.model
canonical_vector = CANONICAL_VECTOR_VALUES[model_desc.model]
assert np.allclose(embeddings[1], embeddings[2]), 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
if is_ci:
delete_model_cache(model.model._model_dir)
@pytest.mark.parametrize("n_dims,model_name", [(512, "Qdrant/clip-ViT-B-32-vision")])
def test_batch_embedding(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
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
embeddings = list(model.embed(images, batch_size=10))
embeddings = np.stack(embeddings, axis=0)
assert np.allclose(embeddings[1], embeddings[2])
embeddings = list(model.embed(images, batch_size=10))
embeddings = np.stack(embeddings, axis=0)
canonical_vector = CANONICAL_VECTOR_VALUES[model_name]
assert embeddings.shape == (len(test_images) * n_images, n_dims)
assert np.allclose(embeddings[0, : canonical_vector.shape[0]], canonical_vector, atol=1e-3)
assert embeddings.shape == (len(test_images) * n_images, n_dims)
if is_ci:
delete_model_cache(model.model._model_dir)
@pytest.mark.parametrize("n_dims,model_name", [(512, "Qdrant/clip-ViT-B-32-vision")])
def test_parallel_processing(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)
def test_parallel_processing(n_dims: int, model_name: str) -> None:
is_ci = os.getenv("CI")
model = ImageEmbedding(model_name=model_name)
embeddings_2 = list(model.embed(images, batch_size=10, parallel=None))
embeddings_2 = np.stack(embeddings_2, axis=0)
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_3 = list(model.embed(images, batch_size=10, parallel=0))
embeddings_3 = np.stack(embeddings_3, axis=0)
embeddings_2 = list(model.embed(images, batch_size=10, parallel=None))
embeddings_2 = np.stack(embeddings_2, 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)
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)
@pytest.mark.parametrize("model_name", ["Qdrant/clip-ViT-B-32-vision"])
@@ -145,31 +121,3 @@ 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
+46 -144
View File
@@ -1,5 +1,4 @@
import os
from contextlib import contextmanager
import pytest
import numpy as np
@@ -7,7 +6,7 @@ import numpy as np
from fastembed.late_interaction.late_interaction_text_embedding import (
LateInteractionTextEmbedding,
)
from tests.utils import delete_model_cache, should_test_model
from tests.utils import delete_model_cache
# vectors are abridged and rounded for brevity
CANONICAL_COLUMN_VALUES = {
@@ -151,124 +150,82 @@ 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"]
@pytest.mark.parametrize("model_name", ["answerdotai/answerai-colbert-small-v1"])
def test_batch_embedding(model_cache, model_name: str):
def test_batch_embedding():
is_ci = os.getenv("CI")
docs_to_embed = docs * 10
with model_cache(model_name) as model:
for model_name, expected_result in CANONICAL_COLUMN_VALUES.items():
print("evaluating", model_name)
model = LateInteractionTextEmbedding(model_name=model_name)
result = list(model.embed(docs_to_embed, batch_size=6))
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)
@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])
if is_ci:
delete_model_cache(model.model._model_dir)
@pytest.mark.parametrize("model_name", ["answerdotai/answerai-colbert-small-v1"])
def test_single_embedding(model_cache, model_name: str):
def test_single_embedding():
is_ci = os.getenv("CI")
is_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
docs_to_embed = docs
for model_desc in LateInteractionTextEmbedding._list_supported_models():
if not should_test_model(model_desc, model_name, is_ci, is_manual):
continue
for model_name, expected_result in CANONICAL_COLUMN_VALUES.items():
print("evaluating", model_name)
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)
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)
@pytest.mark.parametrize("model_name", ["answerdotai/answerai-colbert-small-v1"])
def test_single_embedding_query(model_cache, model_name: str):
def test_single_embedding_query():
is_ci = os.getenv("CI")
is_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
queries_to_embed = docs
for model_desc in LateInteractionTextEmbedding._list_supported_models():
if not should_test_model(model_desc, model_name, is_ci, is_manual):
continue
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)
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)
if is_ci:
delete_model_cache(model.model._model_dir)
@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))
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)
embeddings_2 = list(model.embed(docs, batch_size=10, parallel=None))
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)) # 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
embeddings_3 = list(model.embed(docs, batch_size=10, parallel=0))
embeddings_3 = np.stack(embeddings_3, axis=0)
assert len(embeddings) == len(docs) and embeddings[0].shape[-1] == token_dim
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)
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)
if is_ci:
delete_model_cache(model.model._model_dir)
@pytest.mark.parametrize("model_name", ["answerdotai/answerai-colbert-small-v1"])
@pytest.mark.parametrize(
"model_name",
["colbert-ir/colbertv2.0"],
)
def test_lazy_load(model_name: str):
is_ci = os.getenv("CI")
@@ -287,58 +244,3 @@ 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
+35 -73
View File
@@ -1,6 +1,5 @@
import os
import pytest
from PIL import Image
import numpy as np
@@ -12,13 +11,15 @@ from tests.config import TEST_MISC_DIR
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],
]
]
),
}
@@ -46,77 +47,38 @@ images = [
def test_batch_embedding():
if os.getenv("CI"):
pytest.skip("Colpali is too large to test in CI")
is_ci = os.getenv("CI")
for model_name, expected_result in CANONICAL_IMAGE_VALUES.items():
print("evaluating", model_name)
model = LateInteractionMultimodalEmbedding(model_name=model_name)
result = list(model.embed_image(images, batch_size=2))
if not is_ci:
for model_name, expected_result in CANONICAL_IMAGE_VALUES.items():
print("evaluating", model_name)
model = LateInteractionMultimodalEmbedding(model_name=model_name)
result = list(model.embed_image(images, batch_size=2))
for value in result:
token_num, abridged_dim = expected_result.shape
assert np.allclose(value[:token_num, :abridged_dim], expected_result, atol=2e-3)
for value in result:
batch_size, token_num, abridged_dim = expected_result.shape
assert np.allclose(value[:token_num, :abridged_dim], expected_result, atol=1e-3)
def test_single_embedding():
if os.getenv("CI"):
pytest.skip("Colpali is too large to test in CI")
for model_name, expected_result in CANONICAL_IMAGE_VALUES.items():
print("evaluating", model_name)
model = LateInteractionMultimodalEmbedding(model_name=model_name)
result = next(iter(model.embed_image(images, batch_size=6)))
token_num, abridged_dim = expected_result.shape
assert np.allclose(result[:token_num, :abridged_dim], expected_result, atol=2e-3)
is_ci = os.getenv("CI")
if not is_ci:
for model_name, expected_result in CANONICAL_IMAGE_VALUES.items():
print("evaluating", model_name)
model = LateInteractionMultimodalEmbedding(model_name=model_name)
result = next(iter(model.embed_image(images, batch_size=6)))
batch_size, token_num, abridged_dim = expected_result.shape
assert np.allclose(result[:token_num, :abridged_dim], expected_result, atol=2e-3)
def test_single_embedding_query():
if os.getenv("CI"):
pytest.skip("Colpali is too large to test in CI")
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)))
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
def test_embedding_size():
if os.getenv("CI"):
pytest.skip("Colpali is too large to test in CI")
model_name = "Qdrant/colpali-v1.3-fp16"
model = LateInteractionMultimodalEmbedding(model_name=model_name, lazy_load=True)
assert model.embedding_size == 128
model_name = "Qdrant/ColPali-v1.3-fp16"
model = LateInteractionMultimodalEmbedding(model_name=model_name, lazy_load=True)
assert model.embedding_size == 128
def test_token_count() -> None:
if os.getenv("CI"):
pytest.skip("Colpali is too large to test in CI")
model_name = "Qdrant/colpali-v1.3-fp16"
model = LateInteractionMultimodalEmbedding(model_name=model_name, lazy_load=True)
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
)
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)
-38
View File
@@ -1,38 +0,0 @@
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)
+80 -207
View File
@@ -1,15 +1,14 @@
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, should_test_model
from tests.utils import delete_model_cache
CANONICAL_COLUMN_VALUES = {
"prithivida/Splade_PP_en_v1": {
"prithvida/Splade_PP_en_v1": {
"indices": [
2040,
2047,
@@ -44,202 +43,112 @@ 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,
],
},
}
}
CANONICAL_QUERY_VALUES = {
"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",
)
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"]
@pytest.mark.parametrize(
"model_name",
["prithivida/Splade_PP_en_v1", "Qdrant/minicoil-v1"],
)
def test_batch_embedding(model_cache, model_name: str) -> None:
def test_batch_embedding() -> None:
is_ci = os.getenv("CI")
docs_to_embed = docs * 10
with model_cache(model_name) as model:
for model_name, expected_result in CANONICAL_COLUMN_VALUES.items():
model = SparseTextEmbedding(model_name=model_name)
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(model_cache) -> None:
def test_single_embedding() -> None:
is_ci = os.getenv("CI")
is_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
for model_name, expected_result in CANONICAL_COLUMN_VALUES.items():
model = SparseTextEmbedding(model_name=model_name)
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
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"]
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)
assert passage_result.indices.tolist() == expected_result["indices"]
for i, value in enumerate(passage_result.values):
for i, value in enumerate(result.values):
assert pytest.approx(value, abs=0.001) == expected_result["values"][i]
assert query_result.indices.tolist() == expected_query_result["indices"]
for i, value in enumerate(query_result.values):
assert pytest.approx(value, abs=0.001) == expected_query_result["values"][i]
if is_ci:
delete_model_cache(model.model._model_dir)
@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))
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))
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 (
len(sparse_embeddings)
== len(sparse_embeddings_duo)
# == len(sparse_embeddings_all)
== len(docs)
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)
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)
if is_ci:
delete_model_cache(model.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
@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_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()
def test_stem_with_stopwords_and_punctuation(bm25_instance: Bm25) -> None:
# Setup
bm25_instance.stopwords = {"the", "is", "a"}
bm25_instance.punctuation = {".", ",", "!"}
# Setup
bm25_instance.stopwords = {"the", "is", "a"}
bm25_instance.punctuation = {".", ",", "!"}
# Test data
tokens = ["The", "quick", "brown", "fox", "is", "a", "test", "sentence", ".", "!"]
# Test data
tokens = ["THE", "Quick", "Brown", "Fox", "IS", "A", "Test", "Sentence", ".", "!"]
# Execute
result = bm25_instance._stem(tokens)
# 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
def test_stem_case_insensitive_stopwords(bm25_instance: Bm25) -> None:
# Setup
bm25_instance.stopwords = {"the", "is", "a"}
bm25_instance.punctuation = {".", ",", "!"}
# Test data
tokens = ["THE", "Quick", "Brown", "Fox", "IS", "A", "Test", "Sentence", ".", "!"]
# Execute
result = bm25_instance._stem(tokens)
# Assert
expected = ["quick", "brown", "fox", "test", "sentenc"]
assert result == expected, f"Expected {expected}, but got {result}"
@pytest.mark.parametrize("disable_stemmer", [True, False])
@@ -263,7 +172,10 @@ 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"])
@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)
@@ -281,42 +193,3 @@ 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)
+73 -105
View File
@@ -1,11 +1,10 @@
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, should_test_model
from tests.utils import delete_model_cache
CANONICAL_SCORE_VALUES = {
"Xenova/ms-marco-MiniLM-L-6-v2": np.array([8.500708, -2.541011]),
@@ -16,84 +15,73 @@ CANONICAL_SCORE_VALUES = {
"jinaai/jina-reranker-v2-base-multilingual": np.array([1.6533, -1.6455]),
}
_MODELS_TO_CACHE = ("Xenova/ms-marco-MiniLM-L-6-v2",)
MODELS_TO_CACHE = tuple([x.lower() for x in _MODELS_TO_CACHE])
SELECTED_MODELS = {
"Xenova": "Xenova/ms-marco-MiniLM-L-6-v2",
"BAAI": "BAAI/bge-reranker-base",
"jinaai": "jinaai/jina-reranker-v1-tiny-en",
}
@pytest.fixture(scope="module")
def model_cache():
@pytest.mark.parametrize(
"model_name",
[model_name for model_name in CANONICAL_SCORE_VALUES],
)
def test_rerank(model_name: str) -> None:
is_ci = os.getenv("CI")
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] = 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
model = TextCrossEncoder(model_name=model_name)
yield get_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}"
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:
for name, model in cache.items():
delete_model_cache(model.model._model_dir)
cache.clear()
delete_model_cache(model.model._model_dir)
@pytest.mark.parametrize("model_name", ["Xenova/ms-marco-MiniLM-L-6-v2"])
def test_rerank(model_cache, model_name: str) -> None:
@pytest.mark.parametrize(
"model_name",
[model_name for model_name in SELECTED_MODELS.values()],
)
def test_batch_rerank(model_name: str) -> None:
is_ci = os.getenv("CI")
is_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
for model_desc in TextCrossEncoder._list_supported_models():
if not should_test_model(model_desc, model_name, is_ci, is_manual):
continue
model = TextCrossEncoder(model_name=model_name)
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)))
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_desc.model}, 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_name}, Scores: {scores}, Scores2: {scores2}"
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}"
canonical_scores = np.tile(CANONICAL_SCORE_VALUES[model_name], 50)
assert scores.shape == canonical_scores.shape, f"Unexpected shape for model {model_name}"
assert np.allclose(
scores, canonical_scores, atol=1e-3
), f"Model: {model_name}, Scores: {scores}, Expected: {canonical_scores}"
if is_ci:
delete_model_cache(model.model._model_dir)
@pytest.mark.parametrize("model_name", ["Xenova/ms-marco-MiniLM-L-6-v2"])
def test_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"])
@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)
@@ -107,45 +95,25 @@ def test_lazy_load(model_name: str) -> None:
delete_model_cache(model.model._model_dir)
@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}"
@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_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
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)
+76 -64
View File
@@ -4,7 +4,7 @@ import numpy as np
import pytest
from fastembed import TextEmbedding
from fastembed.text.multitask_embedding import JinaEmbeddingV3, Task
from fastembed.text.multitask_embedding import Task
from tests.utils import delete_model_cache
@@ -60,43 +60,52 @@ CANONICAL_VECTOR_VALUES = {
docs = ["Hello World", "Follow the white rabbit."]
@pytest.mark.parametrize("dim,model_name", [(1024, "jinaai/jina-embeddings-v3")])
def test_batch_embedding(dim: int, model_name: str):
def test_batch_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")
docs_to_embed = docs * 10
default_task = Task.RETRIEVAL_PASSAGE
model = TextEmbedding(model_name=model_name)
for model_desc in TextEmbedding._list_supported_models():
if not is_ci and model_desc.size_in_GB > 1:
continue
embeddings = list(model.embed(documents=docs_to_embed, batch_size=6))
embeddings = np.stack(embeddings, axis=0)
model_name = model_desc.model
dim = model_desc.dim
assert embeddings.shape == (len(docs_to_embed), dim)
if model_name not in CANONICAL_VECTOR_VALUES.keys():
continue
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
model = TextEmbedding(model_name=model_name)
if is_ci:
delete_model_cache(model.model._model_dir)
print(f"evaluating {model_name} default task")
embeddings = list(model.embed(documents=docs_to_embed, batch_size=6))
embeddings = np.stack(embeddings, axis=0)
assert embeddings.shape == (len(docs_to_embed), dim)
canonical_vector = CANONICAL_VECTOR_VALUES[model_name][default_task]["vectors"]
assert np.allclose(
embeddings[: len(docs), : canonical_vector.shape[1]], canonical_vector, atol=1e-4
), model_desc.model
if is_ci:
delete_model_cache(model.model._model_dir)
def test_single_embedding():
is_ci = os.getenv("CI")
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 JinaEmbeddingV3._list_supported_models():
# todo: once we add more models, we should not test models >1GB size locally
for model_desc in TextEmbedding._list_supported_models():
if not is_ci and model_desc.size_in_GB > 1:
continue
model_name = model_desc.model
dim = model_desc.dim
if model_name not in CANONICAL_VECTOR_VALUES.keys():
continue
model = TextEmbedding(model_name=model_name)
for task in CANONICAL_VECTOR_VALUES[model_name]:
@@ -109,42 +118,27 @@ def test_single_embedding():
canonical_vector = task["vectors"]
assert np.allclose(
embeddings[:, : canonical_vector.shape[1]], canonical_vector, atol=1e-4
embeddings[: len(docs), : 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 JinaEmbeddingV3._list_supported_models():
# todo: once we add more models, we should not test models >1GB size locally
for model_desc in TextEmbedding._list_supported_models():
if not is_ci and model_desc.size_in_GB > 1:
continue
model_name = model_desc.model
dim = model_desc.dim
if model_name not in CANONICAL_VECTOR_VALUES.keys():
continue
model = TextEmbedding(model_name=model_name)
print(f"evaluating {model_name} query_embed task_id: {task_id}")
@@ -156,7 +150,7 @@ def test_single_embedding_query():
canonical_vector = CANONICAL_VECTOR_VALUES[model_name][task_id]["vectors"]
assert np.allclose(
embeddings[:, : canonical_vector.shape[1]], canonical_vector, atol=1e-4
embeddings[: len(docs), : canonical_vector.shape[1]], canonical_vector, atol=1e-4
), model_desc.model
if is_ci:
@@ -165,18 +159,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 JinaEmbeddingV3._list_supported_models():
# todo: once we add more models, we should not test models >1GB size locally
for model_desc in TextEmbedding._list_supported_models():
if not is_ci and model_desc.size_in_GB > 1:
continue
model_name = model_desc.model
dim = model_desc.dim
if model_name not in CANONICAL_VECTOR_VALUES.keys():
continue
model = TextEmbedding(model_name=model_name)
print(f"evaluating {model_name} passage_embed task_id: {task_id}")
@@ -188,22 +182,21 @@ def test_single_embedding_passage():
canonical_vector = CANONICAL_VECTOR_VALUES[model_name][task_id]["vectors"]
assert np.allclose(
embeddings[:, : canonical_vector.shape[1]], canonical_vector, atol=1e-4
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)
@pytest.mark.parametrize("dim,model_name", [(1024, "jinaai/jina-embeddings-v3")])
def test_parallel_processing(dim: int, model_name: str):
def test_parallel_processing():
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
@@ -223,14 +216,33 @@ def test_parallel_processing(dim: int, model_name: str):
delete_model_cache(model.model._model_dir)
@pytest.mark.parametrize("model_name", ["jinaai/jina-embeddings-v3"])
def test_task_assignment():
is_ci = os.getenv("CI")
for model_desc in TextEmbedding._list_supported_models():
if not is_ci and model_desc.size_in_GB > 1:
continue
model_name = model_desc.model
if model_name not in CANONICAL_VECTOR_VALUES.keys():
continue
model = TextEmbedding(model_name=model_name)
for i, task_id in enumerate(Task):
_ = list(model.embed(documents=docs, batch_size=1, task_id=i))
assert model.model.current_task_id == task_id
if is_ci:
delete_model_cache(model.model._model_dir)
@pytest.mark.parametrize(
"model_name",
["jinaai/jina-embeddings-v3"],
)
def test_lazy_load(model_name: str):
is_ci = os.getenv("CI")
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")
+60 -114
View File
@@ -1,12 +1,11 @@
import os
import platform
from contextlib import contextmanager
import numpy as np
import pytest
from fastembed.text.text_embedding import TextEmbedding
from tests.utils import delete_model_cache, should_test_model
from tests.utils import delete_model_cache
CANONICAL_VECTOR_VALUES = {
"BAAI/bge-small-en": np.array([-0.0232, -0.0255, 0.0174, -0.0639, -0.0006]),
@@ -53,7 +52,7 @@ CANONICAL_VECTOR_VALUES = {
[0.0802303, 0.3700881, -4.3053818, 0.4431803, -0.271572]
),
"thenlper/gte-large": np.array(
[-0.00986551, -0.00018734, 0.00605892, -0.03289612, -0.0387564],
[-0.01920587, 0.00113156, -0.00708992, -0.00632304, -0.04025577]
),
"mixedbread-ai/mxbai-embed-large-v1": np.array(
[0.02295546, 0.03196154, 0.016512, -0.04031524, -0.0219634]
@@ -72,92 +71,82 @@ CANONICAL_VECTOR_VALUES = {
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])
@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:
def test_embedding() -> 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"
if (
(not is_ci and model_desc.size_in_GB > 1)
or model_desc.model in MULTI_TASK_MODELS
or (is_mac and model_desc.model == "nomic-ai/nomic-embed-text-v1.5-Q")
):
continue
if not should_test_model(model_desc, model_name, is_ci, is_manual):
continue
dim = model_desc.dim
with model_cache(model_desc.model) as model:
docs = ["hello world", "flag embedding"]
embeddings = list(model.embed(docs))
embeddings = np.stack(embeddings, axis=0)
assert embeddings.shape == (2, dim)
canonical_vector = CANONICAL_VECTOR_VALUES[model_desc.model]
assert np.allclose(
embeddings[0, : canonical_vector.shape[0]], canonical_vector, atol=1e-3
), model_desc.model
@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))
model = TextEmbedding(model_name=model_desc.model)
docs = ["hello world", "flag embedding"]
embeddings = list(model.embed(docs))
embeddings = np.stack(embeddings, axis=0)
assert embeddings.shape == (2, dim)
assert embeddings.shape == (len(docs), n_dims)
canonical_vector = CANONICAL_VECTOR_VALUES[model_desc.model]
assert np.allclose(
embeddings[0, : canonical_vector.shape[0]], canonical_vector, atol=1e-3
), model_desc["model"]
if is_ci:
delete_model_cache(model.model._model_dir)
@pytest.mark.parametrize("n_dims,model_name", [(384, "BAAI/bge-small-en-v1.5")])
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)
@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)
embeddings_2 = list(model.embed(docs, batch_size=10, parallel=None))
embeddings_2 = np.stack(embeddings_2, axis=0)
docs = ["hello world", "flag embedding"] * 100
embeddings = list(model.embed(docs, batch_size=10))
embeddings = np.stack(embeddings, axis=0)
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)
assert embeddings.shape == (200, n_dims)
if is_ci:
delete_model_cache(model.model._model_dir)
@pytest.mark.parametrize("model_name", ["BAAI/bge-small-en-v1.5"])
@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"],
)
def test_lazy_load(model_name: str) -> None:
is_ci = os.getenv("CI")
model = TextEmbedding(model_name=model_name, lazy_load=True)
@@ -174,46 +163,3 @@ 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)
+1 -31
View File
@@ -3,9 +3,7 @@ import traceback
from pathlib import Path
from types import TracebackType
from typing import Union, Callable, Any, Type, Optional
from fastembed.common.model_description import BaseModelDescription
from typing import Union, Callable, Any, Type
def delete_model_cache(model_dir: Union[str, Path]) -> None:
@@ -37,31 +35,3 @@ 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: Optional[str],
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