Compare commits

...
15 Commits
Author SHA1 Message Date
George d4121a5b73 new: gpu package (#224) 2026-03-23 23:30:41 +07:00
George Panchuk 4f7c82a7aa add eofl 2026-03-23 23:30:10 +07:00
George Panchuk bb9698d825 sync publih with main 2026-03-23 23:30:10 +07:00
George Panchuk e88b40d820 fix: workflow dispatch can only be triggered from the default branch 2026-03-23 23:30:10 +07:00
George Panchuk 3ea6fa67ce refactoring: alter workflow names 2026-03-23 23:30:10 +07:00
George Panchuk 5b6d9269f5 fix: do not run windows and mac os tests on gpu branch 2026-03-23 23:30:10 +07:00
George Panchuk 5e679579eb new: gpu package publish workflow 2026-03-23 23:30:10 +07:00
George Panchuk 6fa442b960 bump version to 0.8.0 2026-03-23 23:30:02 +07:00
Alexey Masolov 52ebfba27c fix: respect HF_HUB_OFFLINE in download_model to avoid network calls (#614)
When HF_HUB_OFFLINE is set to a truthy value (1, true, yes, on),
download_model() should treat local_files_only=True to avoid any
network calls. Currently, even with the local-cache-first pass (which
may fail due to missing metadata), the retry loop still calls
download_files_from_huggingface() without local_files_only, which
triggers model_info() — a network API call that immediately fails in
offline mode. This causes an unnecessary fallback to GCS download from
storage.googleapis.com.

By setting local_files_only=True when HF_HUB_OFFLINE is enabled:

1. The HF local cache pass works if the model is cached
2. The retry loop skips the network-dependent HF path entirely
3. retrieve_model_gcs() only checks for local fast-* directories
4. No network calls are attempted at all

The truthy value check aligns with huggingface_hub's own parsing of
HF_HUB_OFFLINE, which accepts "1", "true", "yes", "on" (case-insensitive).

This is critical for air-gapped / restricted environments where both
HuggingFace and Google Cloud Storage are unreachable.

Made-with: Cursor
2026-03-23 22:40:03 +07:00
George ea55268e01 fix: fix onnxruntime 1.24, uncap pillow (#611)
* fix: fix onnxruntime 1.24, uncap pillow

* fix: fix python3.10 onnxruntime version

* fix: fix onnxruntime for 3.14, update onnx dep
2026-03-13 00:50:10 +07:00
Kacper ŁukawskiandGeorge Panchuk 800f3887b7 Model: ModernVBERT/colmodernvbert (#588)
* Add ColModernVBERT to LateInteractionMultimodalEmbedding registry

* Implement image processing based on Idefics3ImageProcessor logic

* Fix padding support

* Implement ColModernVBERT logic

* Remove TODOs

* Handle empty pixel values with proper image_size

* Add ColModernVBERT tests

* Run pre-commit

* mypy fixes

* mypy fixes

* mypy fixes

* mypy fixes

* Fix typo in the class name

* Add processor_config.json to additional files

* Fix mypy errors

* Refactor onnx_embed_image

* Fix mypy errors

* fix: colmodernvbert tests and query processing

* fix: remove Union references

* fix: fix exit stack, update tests, implement token count

* fix: uncomment colpali in tests

* fix: lowercase models to cache

* fix: fix models to cache

* refactor: move colmodernvbert related onnx embed to its class

---------

Co-authored-by: George Panchuk <george.panchuk@qdrant.tech>
2026-01-09 18:52:50 +07:00
Bastian Hofmann 020d535f9c Update logo and favicon (#589) 2025-12-18 15:49:31 +01:00
George 685fd9b5a1 new: use cuda if available (#537)
* new: use cuda if available

* fix: fix warning msg

* fix: add missing import
2025-12-10 20:23:34 +07:00
George b304a2aff0 new: drop python3.9, replace optional and union with | (#574)
* new: drop python3.9, replace optional and union with |

* new: remove python 3.9 from pyproject

* refactor: replace remaining union and optional with |

* new: remove optional and union in dataclasses

* fix: add typealias to numpy type

* new: replace union with | in token count
2025-12-10 19:01:01 +07:00
George c715416361 fix: update colbert description (#534) 2025-12-09 19:22:54 +07:00
52 changed files with 2384 additions and 2111 deletions
+1 -1
View File
@@ -25,7 +25,7 @@ jobs:
- name: Set up Python
uses: actions/setup-python@v2
with:
python-version: '3.9.x'
python-version: '3.10.x'
- name: Install dependencies
run: |
python -m pip install poetry
+2 -18
View File
@@ -1,10 +1,11 @@
name: Tests
run-name: Tests (gpu)
on:
pull_request:
branches: [ master, main, gpu ]
workflow_dispatch:
env:
CARGO_TERM_COLOR: always
@@ -15,29 +16,12 @@ jobs:
strategy:
matrix:
python-version:
- '3.9.x'
- '3.10.x'
- '3.11.x'
- '3.12.x'
- '3.13.x'
os:
- ubuntu-latest
- macos-latest
- windows-latest
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 }}
+1 -1
View File
@@ -8,7 +8,7 @@ jobs:
strategy:
fail-fast: true
matrix:
python-version: ["3.9", "3.10", "3.11", "3.12", "3.13"]
python-version: ["3.10", "3.11", "3.12", "3.13"]
os: [ubuntu-latest]
name: Python ${{ matrix.python-version }} test
Binary file not shown.

Before

Width:  |  Height:  |  Size: 1.8 KiB

After

Width:  |  Height:  |  Size: 2.0 KiB

+7 -7
View File
@@ -1,12 +1,12 @@
from dataclasses import dataclass, field
from enum import Enum
from typing import Optional, Any
from typing import Any
@dataclass(frozen=True)
class ModelSource:
hf: Optional[str] = None
url: Optional[str] = None
hf: str | None = None
url: str | None = None
_deprecated_tar_struct: bool = False
@property
@@ -33,8 +33,8 @@ class BaseModelDescription:
@dataclass(frozen=True)
class DenseModelDescription(BaseModelDescription):
dim: Optional[int] = None
tasks: Optional[dict[str, Any]] = field(default_factory=dict)
dim: int | None = None
tasks: dict[str, Any] | None = field(default_factory=dict)
def __post_init__(self) -> None:
assert self.dim is not None, "dim is required for dense model description"
@@ -42,8 +42,8 @@ class DenseModelDescription(BaseModelDescription):
@dataclass(frozen=True)
class SparseModelDescription(BaseModelDescription):
requires_idf: Optional[bool] = None
vocab_size: Optional[int] = None
requires_idf: bool | None = None
vocab_size: int | None = None
class PoolingType(str, Enum):
+9 -7
View File
@@ -5,7 +5,7 @@ import shutil
import tarfile
from copy import deepcopy
from pathlib import Path
from typing import Any, Optional, Union, TypeVar, Generic
from typing import Any, TypeVar, Generic
import requests
from huggingface_hub import snapshot_download, model_info, list_repo_tree
@@ -180,8 +180,8 @@ class ModelManagement(Generic[T]):
def _collect_file_metadata(
model_dir: Path, repo_files: list[RepoFile]
) -> dict[str, dict[str, Union[int, str]]]:
meta: dict[str, dict[str, Union[int, str]]] = {}
) -> dict[str, dict[str, int | str]]:
meta: dict[str, dict[str, int | str]] = {}
file_info_map = {f.path: f for f in repo_files}
for file_path in model_dir.rglob("*"):
if file_path.is_file() and file_path.name != cls.METADATA_FILE:
@@ -193,9 +193,7 @@ class ModelManagement(Generic[T]):
}
return meta
def _save_file_metadata(
model_dir: Path, meta: dict[str, dict[str, Union[int, str]]]
) -> None:
def _save_file_metadata(model_dir: Path, meta: dict[str, dict[str, int | str]]) -> None:
try:
if not model_dir.exists():
model_dir.mkdir(parents=True, exist_ok=True)
@@ -397,7 +395,11 @@ class ModelManagement(Generic[T]):
Path: The path to the downloaded model directory.
"""
local_files_only = kwargs.get("local_files_only", False)
specific_model_path: Optional[str] = kwargs.pop("specific_model_path", None)
hf_offline = os.environ.get("HF_HUB_OFFLINE", "").strip().upper()
if not local_files_only and hf_offline in {"1", "TRUE", "YES", "ON"}:
local_files_only = True
kwargs["local_files_only"] = True
specific_model_path: str | None = kwargs.pop("specific_model_path", None)
if specific_model_path:
return Path(specific_model_path)
retries = 1 if local_files_only else retries
+20 -15
View File
@@ -1,7 +1,7 @@
import warnings
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Generic, Iterable, Optional, Sequence, Type, TypeVar
from typing import Any, Generic, Iterable, Sequence, Type, TypeVar
import numpy as np
import onnxruntime as ort
@@ -9,7 +9,7 @@ import onnxruntime as ort
from numpy.typing import NDArray
from tokenizers import Tokenizer
from fastembed.common.types import OnnxProvider, NumpyArray
from fastembed.common.types import OnnxProvider, NumpyArray, Device
from fastembed.parallel_processor import Worker
# Holds type of the embedding result
@@ -19,8 +19,9 @@ T = TypeVar("T")
@dataclass
class OnnxOutputContext:
model_output: NumpyArray
attention_mask: Optional[NDArray[np.int64]] = None
input_ids: Optional[NDArray[np.int64]] = None
attention_mask: NDArray[np.int64] | None = None
input_ids: NDArray[np.int64] | None = None
metadata: dict[str, Any] | None = None
class OnnxModel(Generic[T]):
@@ -43,8 +44,8 @@ class OnnxModel(Generic[T]):
raise NotImplementedError("Subclasses must implement this method")
def __init__(self) -> None:
self.model: Optional[ort.InferenceSession] = None
self.tokenizer: Optional[Tokenizer] = None
self.model: ort.InferenceSession | None = None
self.tokenizer: Tokenizer | None = None
def _preprocess_onnx_input(
self, onnx_input: dict[str, NumpyArray], **kwargs: Any
@@ -58,25 +59,30 @@ class OnnxModel(Generic[T]):
self,
model_dir: Path,
model_file: str,
threads: Optional[int],
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_id: Optional[int] = None,
extra_session_options: Optional[dict[str, Any]] = None,
threads: int | None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_id: int | None = None,
extra_session_options: dict[str, Any] | None = None,
) -> None:
model_path = model_dir / model_file
# List of Execution Providers: https://onnxruntime.ai/docs/execution-providers
available_providers = ort.get_available_providers()
cuda_available = "CUDAExecutionProvider" in available_providers
explicit_cuda = cuda is True or cuda == Device.CUDA
if cuda and providers is not None:
if explicit_cuda and providers is not None:
warnings.warn(
f"`cuda` and `providers` are mutually exclusive parameters, cuda: {cuda}, providers: {providers}",
f"`cuda` and `providers` are mutually exclusive parameters, "
f"cuda: {cuda}, providers: {providers}. If you'd like to use providers, cuda should be one of "
f"[False, Device.CPU, Device.AUTO].",
category=UserWarning,
stacklevel=6,
)
if providers is not None:
onnx_providers = list(providers)
elif cuda:
elif explicit_cuda or (cuda == Device.AUTO and cuda_available):
if device_id is None:
onnx_providers = ["CUDAExecutionProvider"]
else:
@@ -84,7 +90,6 @@ class OnnxModel(Generic[T]):
else:
onnx_providers = ["CPUExecutionProvider"]
available_providers = ort.get_available_providers()
requested_provider_names: list[str] = []
for provider in onnx_providers:
# check providers available
+4 -3
View File
@@ -50,9 +50,10 @@ def load_tokenizer(model_dir: Path) -> tuple[Tokenizer, dict[str, int]]:
tokenizer = Tokenizer.from_file(str(tokenizer_path))
tokenizer.enable_truncation(max_length=max_context)
tokenizer.enable_padding(
pad_id=config.get("pad_token_id", 0), pad_token=tokenizer_config["pad_token"]
)
if not tokenizer.padding:
tokenizer.enable_padding(
pad_id=config.get("pad_token_id", 0), pad_token=tokenizer_config["pad_token"]
)
for token in tokens_map.values():
if isinstance(token, str):
+21 -19
View File
@@ -1,25 +1,27 @@
from enum import Enum
from pathlib import Path
import sys
from PIL import Image
from typing import Any, Union
from typing import Any, TypeAlias
import numpy as np
from numpy.typing import NDArray
if sys.version_info >= (3, 10):
from typing import TypeAlias
else:
from typing_extensions import TypeAlias
from PIL import Image
PathInput: TypeAlias = Union[str, Path]
ImageInput: TypeAlias = Union[PathInput, Image.Image]
class Device(str, Enum):
CPU = "cpu"
CUDA = "cuda"
AUTO = "auto"
OnnxProvider: TypeAlias = Union[str, tuple[str, dict[Any, Any]]]
NumpyArray = Union[
NDArray[np.float64],
NDArray[np.float32],
NDArray[np.float16],
NDArray[np.int8],
NDArray[np.int64],
NDArray[np.int32],
]
PathInput: TypeAlias = str | Path
ImageInput: TypeAlias = PathInput | Image.Image
OnnxProvider: TypeAlias = str | tuple[str, dict[Any, Any]]
NumpyArray: TypeAlias = (
NDArray[np.float64]
| NDArray[np.float32]
| NDArray[np.float16]
| NDArray[np.int8]
| NDArray[np.int64]
| NDArray[np.int32]
)
+2 -2
View File
@@ -5,7 +5,7 @@ import tempfile
import unicodedata
from pathlib import Path
from itertools import islice
from typing import Iterable, Optional, TypeVar
from typing import Iterable, TypeVar
import numpy as np
from numpy.typing import NDArray
@@ -45,7 +45,7 @@ def iter_batch(iterable: Iterable[T], size: int) -> Iterable[list[T]]:
yield b
def define_cache_dir(cache_dir: Optional[str] = None) -> Path:
def define_cache_dir(cache_dir: str | None = None) -> Path:
"""
Define the cache directory for fastembed
"""
+3 -3
View File
@@ -1,4 +1,4 @@
from typing import Optional, Any
from typing import Any
from loguru import logger
@@ -17,8 +17,8 @@ class JinaEmbedding(TextEmbedding):
def __init__(
self,
model_name: str = "jinaai/jina-embeddings-v2-base-en",
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
cache_dir: str | None = None,
threads: int | None = None,
**kwargs: Any,
):
super().__init__(model_name, cache_dir, threads, **kwargs)
+10 -10
View File
@@ -1,7 +1,7 @@
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Iterable, Sequence, Type
from dataclasses import asdict
from fastembed.common.types import NumpyArray
from fastembed.common.types import NumpyArray, Device
from fastembed.common import ImageInput, OnnxProvider
from fastembed.image.image_embedding_base import ImageEmbeddingBase
from fastembed.image.onnx_embedding import OnnxImageEmbedding
@@ -48,11 +48,11 @@ class ImageEmbedding(ImageEmbeddingBase):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
**kwargs: Any,
):
@@ -98,7 +98,7 @@ class ImageEmbedding(ImageEmbeddingBase):
ValueError: If the model name is not found in the supported models.
"""
descriptions = cls._list_supported_models()
embedding_size: Optional[int] = None
embedding_size: int | None = None
for description in descriptions:
if description.model.lower() == model_name.lower():
embedding_size = description.dim
@@ -113,9 +113,9 @@ class ImageEmbedding(ImageEmbeddingBase):
def embed(
self,
images: Union[ImageInput, Iterable[ImageInput]],
images: ImageInput | Iterable[ImageInput],
batch_size: int = 16,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
"""
+6 -6
View File
@@ -1,4 +1,4 @@
from typing import Iterable, Optional, Any, Union
from typing import Iterable, Any
from fastembed.common.model_description import DenseModelDescription
from fastembed.common.types import NumpyArray
@@ -10,21 +10,21 @@ class ImageEmbeddingBase(ModelManagement[DenseModelDescription]):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
cache_dir: str | None = None,
threads: int | None = None,
**kwargs: Any,
):
self.model_name = model_name
self.cache_dir = cache_dir
self.threads = threads
self._local_files_only = kwargs.pop("local_files_only", False)
self._embedding_size: Optional[int] = None
self._embedding_size: int | None = None
def embed(
self,
images: Union[ImageInput, Iterable[ImageInput]],
images: ImageInput | Iterable[ImageInput],
batch_size: int = 16,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
"""
+17 -15
View File
@@ -1,7 +1,7 @@
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Iterable, Sequence, Type
from fastembed.common.types import NumpyArray
from fastembed.common.types import NumpyArray, Device
from fastembed.common import ImageInput, OnnxProvider
from fastembed.common.onnx_model import OnnxOutputContext
from fastembed.common.utils import define_cache_dir, normalize
@@ -63,14 +63,15 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
device_id: Optional[int] = None,
specific_model_path: Optional[str] = None,
device_id: int | None = None,
specific_model_path: str | None = None,
**kwargs: Any,
):
"""
@@ -82,10 +83,11 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.
providers (Optional[Sequence[OnnxProvider]], optional): The list of onnxruntime providers to use.
Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to False.
cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to Device.AUTO.
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, mutually exclusive
with `providers`. Defaults to None.
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
@@ -105,7 +107,7 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
self.cuda = cuda
# This device_id will be used if we need to load model in current process
self.device_id: Optional[int] = None
self.device_id: int | None = None
if device_id is not None:
self.device_id = device_id
elif self.device_ids is not None:
@@ -150,9 +152,9 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
def embed(
self,
images: Union[ImageInput, Iterable[ImageInput]],
images: ImageInput | Iterable[ImageInput],
batch_size: int = 16,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
"""
+19 -17
View File
@@ -2,13 +2,13 @@ import contextlib
import os
from multiprocessing import get_all_start_methods
from pathlib import Path
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Iterable, Sequence, Type
import numpy as np
from PIL import Image
from fastembed.image.transform.operators import Compose
from fastembed.common.types import NumpyArray
from fastembed.common.types import NumpyArray, Device
from fastembed.common import ImageInput, OnnxProvider
from fastembed.common.onnx_model import EmbeddingWorker, OnnxModel, OnnxOutputContext, T
from fastembed.common.preprocessor_utils import load_preprocessor
@@ -37,7 +37,7 @@ class OnnxImageModel(OnnxModel[T]):
def __init__(self) -> None:
super().__init__()
self.processor: Optional[Compose] = None
self.processor: Compose | None = None
def _preprocess_onnx_input(
self, onnx_input: dict[str, NumpyArray], **kwargs: Any
@@ -51,11 +51,11 @@ class OnnxImageModel(OnnxModel[T]):
self,
model_dir: Path,
model_file: str,
threads: Optional[int],
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_id: Optional[int] = None,
extra_session_options: Optional[dict[str, Any]] = None,
threads: int | None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_id: int | None = None,
extra_session_options: dict[str, Any] | None = None,
) -> None:
super()._load_onnx_model(
model_dir=model_dir,
@@ -76,9 +76,11 @@ class OnnxImageModel(OnnxModel[T]):
return {input_name: encoded}
def onnx_embed(self, images: list[ImageInput], **kwargs: Any) -> OnnxOutputContext:
with contextlib.ExitStack():
with contextlib.ExitStack() as stack:
image_files = [
Image.open(image) if not isinstance(image, Image.Image) else image
stack.enter_context(Image.open(image))
if not isinstance(image, Image.Image)
else image
for image in images
]
assert self.processor is not None, "Processor is not initialized"
@@ -93,15 +95,15 @@ class OnnxImageModel(OnnxModel[T]):
self,
model_name: str,
cache_dir: str,
images: Union[ImageInput, Iterable[ImageInput]],
images: ImageInput | Iterable[ImageInput],
batch_size: int = 256,
parallel: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
parallel: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
local_files_only: bool = False,
specific_model_path: Optional[str] = None,
extra_session_options: Optional[dict[str, Any]] = None,
specific_model_path: str | None = None,
extra_session_options: dict[str, Any] | None = None,
**kwargs: Any,
) -> Iterable[T]:
is_small = False
+81 -9
View File
@@ -1,5 +1,3 @@
from typing import Union
import numpy as np
from PIL import Image
@@ -15,7 +13,7 @@ def convert_to_rgb(image: Image.Image) -> Image.Image:
def center_crop(
image: Union[Image.Image, NumpyArray],
image: Image.Image | NumpyArray,
size: tuple[int, int],
) -> NumpyArray:
if isinstance(image, np.ndarray):
@@ -64,8 +62,8 @@ def center_crop(
def normalize(
image: NumpyArray,
mean: Union[float, list[float]],
std: Union[float, list[float]],
mean: float | list[float],
std: float | list[float],
) -> NumpyArray:
num_channels = image.shape[1] if len(image.shape) == 4 else image.shape[0]
@@ -96,8 +94,8 @@ def normalize(
def resize(
image: Image.Image,
size: Union[int, tuple[int, int]],
resample: Union[int, Image.Resampling] = Image.Resampling.BILINEAR,
size: int | tuple[int, int],
resample: int | Image.Resampling = Image.Resampling.BILINEAR,
) -> Image.Image:
if isinstance(size, tuple):
return image.resize(size, resample)
@@ -117,7 +115,7 @@ def rescale(image: NumpyArray, scale: float, dtype: type = np.float32) -> NumpyA
return (image * scale).astype(dtype)
def pil2ndarray(image: Union[Image.Image, NumpyArray]) -> NumpyArray:
def pil2ndarray(image: Image.Image | NumpyArray) -> NumpyArray:
if isinstance(image, Image.Image):
return np.asarray(image).transpose((2, 0, 1))
return image
@@ -126,7 +124,7 @@ def pil2ndarray(image: Union[Image.Image, NumpyArray]) -> NumpyArray:
def pad2square(
image: Image.Image,
size: int,
fill_color: Union[str, int, tuple[int, ...]] = 0,
fill_color: str | int | tuple[int, ...] = 0,
) -> Image.Image:
height, width = image.height, image.width
@@ -147,3 +145,77 @@ def pad2square(
new_image = Image.new(mode="RGB", size=(size, size), color=fill_color)
new_image.paste(image.crop((left, top, right, bottom)) if crop_required else image)
return new_image
def resize_longest_edge(
image: Image.Image,
max_size: int,
resample: int | Image.Resampling = Image.Resampling.LANCZOS,
) -> Image.Image:
height, width = image.height, image.width
aspect_ratio = width / height
if width >= height:
# Width is longer
new_width = max_size
new_height = int(new_width / aspect_ratio)
else:
# Height is longer
new_height = max_size
new_width = int(new_height * aspect_ratio)
# Ensure even dimensions
if new_height % 2 != 0:
new_height += 1
if new_width % 2 != 0:
new_width += 1
return image.resize((new_width, new_height), resample)
def crop_ndarray(
image: NumpyArray,
x1: int,
y1: int,
x2: int,
y2: int,
channel_first: bool = True,
) -> NumpyArray:
if channel_first:
# (C, H, W) format
return image[:, y1:y2, x1:x2]
else:
# (H, W, C) format
return image[y1:y2, x1:x2, :]
def resize_ndarray(
image: NumpyArray,
size: tuple[int, int],
resample: int | Image.Resampling = Image.Resampling.LANCZOS,
channel_first: bool = True,
) -> NumpyArray:
# Convert to PIL-friendly format (H, W, C)
if channel_first:
img_hwc = image.transpose((1, 2, 0))
else:
img_hwc = image
# Handle different dtypes
if img_hwc.dtype == np.float32 or img_hwc.dtype == np.float64:
# Assume normalized, scale to 0-255 for PIL
img_hwc_scaled = (img_hwc * 255).astype(np.uint8)
pil_img = Image.fromarray(img_hwc_scaled, mode="RGB")
resized = pil_img.resize(size, resample)
result = np.array(resized).astype(np.float32) / 255.0
else:
# uint8 or similar
pil_img = Image.fromarray(img_hwc.astype(np.uint8), mode="RGB")
resized = pil_img.resize(size, resample)
result = np.array(resized)
# Convert back to original format
if channel_first:
result = result.transpose((2, 0, 1))
return result
+243 -13
View File
@@ -1,4 +1,5 @@
from typing import Any, Union, Optional
from typing import Any
import math
from PIL import Image
@@ -6,16 +7,19 @@ from fastembed.common.types import NumpyArray
from fastembed.image.transform.functional import (
center_crop,
convert_to_rgb,
crop_ndarray,
normalize,
pil2ndarray,
rescale,
resize,
resize_longest_edge,
resize_ndarray,
pad2square,
)
class Transform:
def __call__(self, images: list[Any]) -> Union[list[Image.Image], list[NumpyArray]]:
def __call__(self, images: list[Any]) -> list[Image.Image] | list[NumpyArray]:
raise NotImplementedError("Subclasses must implement this method")
@@ -33,18 +37,28 @@ class CenterCrop(Transform):
class Normalize(Transform):
def __init__(self, mean: Union[float, list[float]], std: Union[float, list[float]]):
def __init__(self, mean: float | list[float], std: float | list[float]):
self.mean = mean
self.std = std
def __call__(self, images: list[NumpyArray]) -> list[NumpyArray]:
return [normalize(image, mean=self.mean, std=self.std) for image in images]
def __call__( # type: ignore[override]
self, images: list[NumpyArray] | list[list[NumpyArray]]
) -> list[NumpyArray] | list[list[NumpyArray]]:
if images and isinstance(images[0], list):
# Nested structure from ImageSplitter
return [
[normalize(image, mean=self.mean, std=self.std) for image in img_patches] # type: ignore[arg-type]
for img_patches in images
]
else:
# Flat structure (backward compatibility)
return [normalize(image, mean=self.mean, std=self.std) for image in images] # type: ignore[arg-type]
class Resize(Transform):
def __init__(
self,
size: Union[int, tuple[int, int]],
size: int | tuple[int, int],
resample: Image.Resampling = Image.Resampling.BICUBIC,
):
self.size = size
@@ -58,12 +72,22 @@ class Rescale(Transform):
def __init__(self, scale: float = 1 / 255):
self.scale = scale
def __call__(self, images: list[NumpyArray]) -> list[NumpyArray]:
return [rescale(image, scale=self.scale) for image in images]
def __call__( # type: ignore[override]
self, images: list[NumpyArray] | list[list[NumpyArray]]
) -> list[NumpyArray] | list[list[NumpyArray]]:
if images and isinstance(images[0], list):
# Nested structure from ImageSplitter
return [
[rescale(image, scale=self.scale) for image in img_patches] # type: ignore[arg-type]
for img_patches in images
]
else:
# Flat structure (backward compatibility)
return [rescale(image, scale=self.scale) for image in images] # type: ignore[arg-type]
class PILtoNDarray(Transform):
def __call__(self, images: list[Union[Image.Image, NumpyArray]]) -> list[NumpyArray]:
def __call__(self, images: list[Image.Image | NumpyArray]) -> list[NumpyArray]:
return [pil2ndarray(image) for image in images]
@@ -71,7 +95,7 @@ class PadtoSquare(Transform):
def __init__(
self,
size: int,
fill_color: Union[str, int, tuple[int, ...]],
fill_color: str | int | tuple[int, ...],
):
self.size = size
self.fill_color = fill_color
@@ -82,13 +106,174 @@ class PadtoSquare(Transform):
]
class ResizeLongestEdge(Transform):
"""Resize images so the longest edge equals target size, preserving aspect ratio."""
def __init__(
self,
size: int,
resample: Image.Resampling = Image.Resampling.LANCZOS,
):
self.size = size
self.resample = resample
def __call__(self, images: list[Image.Image]) -> list[Image.Image]:
return [resize_longest_edge(image, self.size, self.resample) for image in images]
class ResizeForVisionEncoder(Transform):
"""
Resize both dimensions to be multiples of vision_encoder_max_size.
Preserves aspect ratio approximately.
Works on numpy arrays in (C, H, W) format.
"""
def __init__(
self,
max_size: int,
resample: Image.Resampling = Image.Resampling.LANCZOS,
):
self.max_size = max_size
self.resample = resample
def __call__(self, images: list[NumpyArray]) -> list[NumpyArray]:
result = []
for image in images:
# Assume (C, H, W) format
_, height, width = image.shape
aspect_ratio = width / height
if width >= height:
# Calculate new width as multiple of max_size
new_width = math.ceil(width / self.max_size) * self.max_size
new_height = int(new_width / aspect_ratio)
new_height = math.ceil(new_height / self.max_size) * self.max_size
else:
# Calculate new height as multiple of max_size
new_height = math.ceil(height / self.max_size) * self.max_size
new_width = int(new_height * aspect_ratio)
new_width = math.ceil(new_width / self.max_size) * self.max_size
# Resize using the ndarray resize function
resized = resize_ndarray(
image,
size=(new_width, new_height), # PIL expects (width, height)
resample=self.resample,
channel_first=True,
)
result.append(resized)
return result
class ImageSplitter(Transform):
"""
Split images into grid of patches plus a global view.
If image dimensions exceed max_size:
- Divide into ceil(H/max_size) x ceil(W/max_size) patches
- Each patch is cropped from the image
- Add a global view (original resized to max_size x max_size)
If image is smaller than max_size:
- Return single image unchanged
Works on numpy arrays in (C, H, W) format.
"""
def __init__(
self,
max_size: int,
resample: Image.Resampling = Image.Resampling.LANCZOS,
):
self.max_size = max_size
self.resample = resample
def __call__(self, images: list[NumpyArray]) -> list[list[NumpyArray]]: # type: ignore[override]
result = []
for image in images:
# Assume (C, H, W) format
_, height, width = image.shape
max_height = max_width = self.max_size
frames = []
if height > max_height or width > max_width:
# Calculate the number of splits needed
num_splits_h = math.ceil(height / max_height)
num_splits_w = math.ceil(width / max_width)
# Calculate optimal patch dimensions
optimal_height = math.ceil(height / num_splits_h)
optimal_width = math.ceil(width / num_splits_w)
# Generate patches in grid order (row by row)
for r in range(num_splits_h):
for c in range(num_splits_w):
# Calculate crop coordinates
start_x = c * optimal_width
start_y = r * optimal_height
end_x = min(start_x + optimal_width, width)
end_y = min(start_y + optimal_height, height)
# Crop the patch
cropped = crop_ndarray(
image, x1=start_x, y1=start_y, x2=end_x, y2=end_y, channel_first=True
)
frames.append(cropped)
# Add global view (resized to max_size x max_size)
global_view = resize_ndarray(
image,
size=(max_width, max_height), # PIL expects (width, height)
resample=self.resample,
channel_first=True,
)
frames.append(global_view)
else:
# Image is small enough, no splitting needed
frames.append(image)
# Append (not extend) to preserve per-image grouping
result.append(frames)
return result
class SquareResize(Transform):
"""
Resize images to square dimensions (max_size x max_size).
Works on numpy arrays in (C, H, W) format.
"""
def __init__(
self,
size: int,
resample: Image.Resampling = Image.Resampling.LANCZOS,
):
self.size = size
self.resample = resample
def __call__(self, images: list[NumpyArray]) -> list[list[NumpyArray]]: # type: ignore[override]
return [
[
resize_ndarray(
image, size=(self.size, self.size), resample=self.resample, channel_first=True
)
]
for image in images
]
class Compose:
def __init__(self, transforms: list[Transform]):
self.transforms = transforms
def __call__(
self, images: Union[list[Image.Image], list[NumpyArray]]
) -> Union[list[NumpyArray], list[Image.Image]]:
self, images: list[Image.Image] | list[NumpyArray]
) -> list[NumpyArray] | list[Image.Image]:
for transform in self.transforms:
images = transform(images)
return images
@@ -118,6 +303,7 @@ class Compose:
Valid size keys (nested):
- {"height", "width"}
- {"shortest_edge"}
- {"longest_edge"}
Returns:
Compose: Image processor.
@@ -128,6 +314,7 @@ class Compose:
cls._get_pad2square(transforms, config)
cls._get_center_crop(transforms, config)
cls._get_pil2ndarray(transforms, config)
cls._get_image_splitting(transforms, config)
cls._get_rescale(transforms, config)
cls._get_normalize(transforms, config)
return cls(transforms=transforms)
@@ -196,6 +383,25 @@ class Compose:
resample=resample,
)
)
elif mode == "Idefics3ImageProcessor":
if config.get("do_resize", False):
size = config.get("size", {})
if "longest_edge" not in size:
raise ValueError(
"Size dictionary must contain 'longest_edge' key for Idefics3ImageProcessor"
)
# Handle resample parameter - can be int enum or PIL.Image.Resampling
resample = config.get("resample", Image.Resampling.LANCZOS)
if isinstance(resample, int):
resample = Image.Resampling(resample)
transforms.append(
ResizeLongestEdge(
size=size["longest_edge"],
resample=resample,
)
)
else:
raise ValueError(f"Preprocessor {mode} is not supported")
@@ -217,6 +423,8 @@ class Compose:
pass
elif mode == "JinaCLIPImageProcessor":
pass
elif mode == "Idefics3ImageProcessor":
pass
else:
raise ValueError(f"Preprocessor {mode} is not supported")
@@ -224,6 +432,28 @@ class Compose:
def _get_pil2ndarray(transforms: list[Transform], config: dict[str, Any]) -> None:
transforms.append(PILtoNDarray())
@classmethod
def _get_image_splitting(cls, transforms: list[Transform], config: dict[str, Any]) -> None:
"""
Add image splitting transforms for Idefics3.
Handles conditional logic: splitting vs square resize.
Must be called AFTER PILtoNDarray.
"""
mode = config.get("image_processor_type", "CLIPImageProcessor")
if mode == "Idefics3ImageProcessor":
do_splitting = config.get("do_image_splitting", False)
max_size = config.get("max_image_size", {}).get("longest_edge", 512)
resample = config.get("resample", Image.Resampling.LANCZOS)
if isinstance(resample, int):
resample = Image.Resampling(resample)
if do_splitting:
transforms.append(ResizeForVisionEncoder(max_size, resample))
transforms.append(ImageSplitter(max_size, resample))
else:
transforms.append(SquareResize(max_size, resample))
@staticmethod
def _get_rescale(transforms: list[Transform], config: dict[str, Any]) -> None:
if config.get("do_rescale", True):
@@ -253,7 +483,7 @@ class Compose:
)
@staticmethod
def _interpolation_resolver(resample: Optional[str] = None) -> Image.Resampling:
def _interpolation_resolver(resample: str | None = None) -> Image.Resampling:
interpolation_map = {
"nearest": Image.Resampling.NEAREST,
"lanczos": Image.Resampling.LANCZOS,
+23 -22
View File
@@ -1,11 +1,11 @@
import string
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Iterable, Sequence, Type
import numpy as np
from tokenizers import Encoding, Tokenizer
from fastembed.common.preprocessor_utils import load_tokenizer
from fastembed.common.types import NumpyArray
from fastembed.common.types import NumpyArray, Device
from fastembed.common import OnnxProvider
from fastembed.common.onnx_model import OnnxOutputContext
from fastembed.common.utils import define_cache_dir, iter_batch
@@ -19,7 +19,7 @@ supported_colbert_models: list[DenseModelDescription] = [
DenseModelDescription(
model="colbert-ir/colbertv2.0",
dim=128,
description="Late interaction model",
description="Text embeddings, Unimodal (text), English, 512 input tokens truncation, 2023 year",
license="mit",
size_in_GB=0.44,
sources=ModelSource(hf="colbert-ir/colbertv2.0"),
@@ -28,7 +28,7 @@ supported_colbert_models: list[DenseModelDescription] = [
DenseModelDescription(
model="answerdotai/answerai-colbert-small-v1",
dim=96,
description="Text embeddings, Unimodal (text), Multilingual (~100 languages), 512 input tokens truncation, 2024 year",
description="Text embeddings, Unimodal (text), English, 512 input tokens truncation, 2024 year",
license="apache-2.0",
size_in_GB=0.13,
sources=ModelSource(hf="answerdotai/answerai-colbert-small-v1"),
@@ -98,7 +98,7 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
def token_count(
self,
texts: Union[str, Iterable[str]],
texts: str | Iterable[str],
batch_size: int = 1024,
is_doc: bool = True,
include_extension: bool = False,
@@ -140,14 +140,14 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
device_id: Optional[int] = None,
specific_model_path: Optional[str] = None,
device_id: int | None = None,
specific_model_path: str | None = None,
**kwargs: Any,
):
"""
@@ -159,10 +159,11 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.
providers (Optional[Sequence[OnnxProvider]], optional): The list of onnxruntime providers to use.
Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to False.
cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to Device.AUTO.
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, mutually exclusive
with `providers`. Defaults to None.
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
@@ -182,7 +183,7 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
self.cuda = cuda
# This device_id will be used if we need to load model in current process
self.device_id: Optional[int] = None
self.device_id: int | None = None
if device_id is not None:
self.device_id = device_id
elif self.device_ids is not None:
@@ -198,11 +199,11 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
local_files_only=self._local_files_only,
specific_model_path=self._specific_model_path,
)
self.mask_token_id: Optional[int] = None
self.pad_token_id: Optional[int] = None
self.mask_token_id: int | None = None
self.pad_token_id: int | None = None
self.skip_list: set[int] = set()
self.query_tokenizer: Optional[Tokenizer] = None
self.query_tokenizer: Tokenizer | None = None
if not self.lazy_load:
self.load_onnx_model()
@@ -238,9 +239,9 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
def embed(
self,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
"""
@@ -273,7 +274,7 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
**kwargs,
)
def query_embed(self, query: Union[str, Iterable[str]], **kwargs: Any) -> Iterable[NumpyArray]:
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:
if isinstance(query, str):
query = [query]
@@ -1,4 +1,4 @@
from typing import Iterable, Optional, Union, Any
from typing import Iterable, Any
from fastembed.common.model_description import DenseModelDescription
from fastembed.common.types import NumpyArray
@@ -9,21 +9,21 @@ class LateInteractionTextEmbeddingBase(ModelManagement[DenseModelDescription]):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
cache_dir: str | None = None,
threads: int | None = None,
**kwargs: Any,
):
self.model_name = model_name
self.cache_dir = cache_dir
self.threads = threads
self._local_files_only = kwargs.pop("local_files_only", False)
self._embedding_size: Optional[int] = None
self._embedding_size: int | None = None
def embed(
self,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
raise NotImplementedError()
@@ -43,7 +43,7 @@ class LateInteractionTextEmbeddingBase(ModelManagement[DenseModelDescription]):
# This is model-specific, so that different models can have specialized implementations
yield from self.embed(texts, **kwargs)
def query_embed(self, query: Union[str, Iterable[str]], **kwargs: Any) -> Iterable[NumpyArray]:
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:
"""
Embeds queries
@@ -72,7 +72,7 @@ class LateInteractionTextEmbeddingBase(ModelManagement[DenseModelDescription]):
def token_count(
self,
texts: Union[str, Iterable[str]],
texts: str | Iterable[str],
batch_size: int = 1024,
**kwargs: Any,
) -> int:
@@ -1,8 +1,8 @@
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Iterable, Sequence, Type
from dataclasses import asdict
from fastembed.common.model_description import DenseModelDescription
from fastembed.common.types import NumpyArray
from fastembed.common.types import NumpyArray, Device
from fastembed.common import OnnxProvider
from fastembed.late_interaction.colbert import Colbert
from fastembed.late_interaction.jina_colbert import JinaColbert
@@ -51,11 +51,11 @@ class LateInteractionTextEmbedding(LateInteractionTextEmbeddingBase):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
**kwargs: Any,
):
@@ -101,7 +101,7 @@ class LateInteractionTextEmbedding(LateInteractionTextEmbeddingBase):
ValueError: If the model name is not found in the supported models.
"""
descriptions = cls._list_supported_models()
embedding_size: Optional[int] = None
embedding_size: int | None = None
for description in descriptions:
if description.model.lower() == model_name.lower():
embedding_size = description.dim
@@ -116,9 +116,9 @@ class LateInteractionTextEmbedding(LateInteractionTextEmbeddingBase):
def embed(
self,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
"""
@@ -138,7 +138,7 @@ class LateInteractionTextEmbedding(LateInteractionTextEmbeddingBase):
"""
yield from self.model.embed(documents, batch_size, parallel, **kwargs)
def query_embed(self, query: Union[str, Iterable[str]], **kwargs: Any) -> Iterable[NumpyArray]:
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:
"""
Embeds queries
@@ -154,7 +154,7 @@ class LateInteractionTextEmbedding(LateInteractionTextEmbeddingBase):
def token_count(
self,
texts: Union[str, Iterable[str]],
texts: str | Iterable[str],
batch_size: int = 1024,
is_doc: bool = True,
include_extension: bool = False,
@@ -1,5 +1,5 @@
from dataclasses import asdict
from typing import Union, Iterable, Optional, Any, Type
from typing import Iterable, Any, Type
from fastembed.common.model_description import DenseModelDescription, ModelSource
from fastembed.common.onnx_model import OnnxOutputContext
@@ -63,9 +63,9 @@ class TokenEmbeddingsModel(OnnxTextEmbedding, LateInteractionTextEmbeddingBase):
def embed(
self,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
yield from super().embed(documents, batch_size=batch_size, parallel=parallel, **kwargs)
@@ -0,0 +1,532 @@
import contextlib
from typing import Any, Iterable, Type, Optional, Sequence
import json
import numpy as np
from tokenizers import Encoding
from PIL import Image
from fastembed.common import ImageInput
from fastembed.common.model_description import DenseModelDescription, ModelSource
from fastembed.common.onnx_model import OnnxOutputContext
from fastembed.common.types import NumpyArray, OnnxProvider
from fastembed.common.utils import define_cache_dir, iter_batch
from fastembed.late_interaction_multimodal.late_interaction_multimodal_embedding_base import (
LateInteractionMultimodalEmbeddingBase,
)
from fastembed.late_interaction_multimodal.onnx_multimodal_model import (
OnnxMultimodalModel,
TextEmbeddingWorker,
ImageEmbeddingWorker,
)
supported_colmodernvbert_models: list[DenseModelDescription] = [
DenseModelDescription(
model="Qdrant/colmodernvbert",
dim=128,
description="The late-interaction version of ModernVBERT, CPU friendly, English, 2025.",
license="mit",
size_in_GB=1.0,
sources=ModelSource(hf="Qdrant/colmodernvbert"),
additional_files=["processor_config.json"],
model_file="model.onnx",
),
]
class ColModernVBERT(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyArray]):
"""
The ModernVBERT/colmodernvbert model implementation. This model uses
bidirectional attention, which proves to work better for retrieval.
See: https://huggingface.co/ModernVBERT/colmodernvbert
"""
VISUAL_PROMPT_PREFIX = (
"<|begin_of_text|>User:<image>Describe the image.<end_of_utterance>\nAssistant:"
)
QUERY_AUGMENTATION_TOKEN = "<end_of_utterance>"
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
lazy_load: bool = False,
device_id: Optional[int] = None,
specific_model_path: Optional[str] = None,
**kwargs: Any,
):
"""
Args:
model_name (str): The name of the model to use.
cache_dir (str, optional): The path to the cache directory.
Can be set using the `FASTEMBED_CACHE_PATH` env variable.
Defaults to `fastembed_cache` in the system's temp directory.
threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.
providers (Optional[Sequence[OnnxProvider]], optional): The list of onnxruntime providers to use.
Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to False.
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
Raises:
ValueError: If the model_name is not in the format <org>/<model> e.g. BAAI/bge-base-en.
"""
super().__init__(model_name, cache_dir, threads, **kwargs)
self.providers = providers
self.lazy_load = lazy_load
self._extra_session_options = self._select_exposed_session_options(kwargs)
# List of device ids, that can be used for data parallel processing in workers
self.device_ids = device_ids
self.cuda = cuda
# This device_id will be used if we need to load model in current process
self.device_id: Optional[int] = None
if device_id is not None:
self.device_id = device_id
elif self.device_ids is not None:
self.device_id = self.device_ids[0]
self.model_description = self._get_model_description(model_name)
self.cache_dir = str(define_cache_dir(cache_dir))
self._specific_model_path = specific_model_path
self._model_dir = self.download_model(
self.model_description,
self.cache_dir,
local_files_only=self._local_files_only,
specific_model_path=self._specific_model_path,
)
self.mask_token_id = None
self.pad_token_id = None
self.image_seq_len: Optional[int] = None
self.max_image_size: Optional[int] = None
self.image_size: Optional[int] = None
if not self.lazy_load:
self.load_onnx_model()
@classmethod
def _list_supported_models(cls) -> list[DenseModelDescription]:
"""Lists the supported models.
Returns:
list[DenseModelDescription]: A list of DenseModelDescription objects containing the model information.
"""
return supported_colmodernvbert_models
def load_onnx_model(self) -> None:
self._load_onnx_model(
model_dir=self._model_dir,
model_file=self.model_description.model_file,
threads=self.threads,
providers=self.providers,
cuda=self.cuda,
device_id=self.device_id,
extra_session_options=self._extra_session_options,
)
# Load image processing configuration
processor_config_path = self._model_dir / "processor_config.json"
with open(processor_config_path) as f:
processor_config = json.load(f)
self.image_seq_len = processor_config.get("image_seq_len", 64)
preprocessor_config_path = self._model_dir / "preprocessor_config.json"
with open(preprocessor_config_path) as f:
preprocessor_config = json.load(f)
self.max_image_size = preprocessor_config.get("max_image_size", {}).get(
"longest_edge", 512
)
# Load model configuration
config_path = self._model_dir / "config.json"
with open(config_path) as f:
model_config = json.load(f)
vision_config = model_config.get("vision_config", {})
self.image_size = vision_config.get("image_size", 512)
def _preprocess_onnx_text_input(
self, onnx_input: dict[str, NumpyArray], **kwargs: Any
) -> dict[str, NumpyArray]:
"""
Post-process the ONNX model output to convert it into a usable format.
Args:
output (OnnxOutputContext): The raw output from the ONNX model.
Returns:
Iterable[NumpyArray]: Post-processed output as NumPy arrays.
"""
batch_size, seq_length = onnx_input["input_ids"].shape
empty_image_placeholder: NumpyArray = np.zeros(
(batch_size, seq_length, 3, self.image_size, self.image_size),
dtype=np.float32, # type: ignore[type-var,arg-type,assignment]
)
onnx_input["pixel_values"] = empty_image_placeholder
return onnx_input
def _post_process_onnx_text_output(
self,
output: OnnxOutputContext,
) -> Iterable[NumpyArray]:
"""
Post-process the ONNX model output to convert it into a usable format.
Args:
output (OnnxOutputContext): The raw output from the ONNX model.
Returns:
Iterable[NumpyArray]: Post-processed output as NumPy arrays.
"""
return output.model_output
def tokenize(self, documents: list[str], **kwargs: Any) -> list[Encoding]:
# Add query augmentation tokens (matching process_queries logic from colpali-engine)
augmented_queries = [doc + self.QUERY_AUGMENTATION_TOKEN * 10 for doc in documents]
encoded = self.tokenizer.encode_batch(augmented_queries) # type: ignore[union-attr]
return encoded
def token_count(
self,
texts: str | Iterable[str],
batch_size: int = 1024,
include_extension: bool = False,
**kwargs: Any,
) -> int:
if not hasattr(self, "model") or self.model is None:
self.load_onnx_model() # loads the tokenizer as well
token_num = 0
texts = [texts] if isinstance(texts, str) else texts
assert self.tokenizer is not None
tokenize_func = self.tokenize if include_extension else self.tokenizer.encode_batch
for batch in iter_batch(texts, batch_size):
token_num += sum([sum(encoding.attention_mask) for encoding in tokenize_func(batch)])
return token_num
def onnx_embed_image(self, images: list[ImageInput], **kwargs: Any) -> OnnxOutputContext:
with contextlib.ExitStack() as stack:
image_files = [
stack.enter_context(Image.open(image))
if not isinstance(image, Image.Image)
else image
for image in images
]
assert self.processor is not None, "Processor is not initialized"
processed = self.processor(image_files)
encoded, attention_mask, metadata = self._process_nested_patches(processed) # type: ignore[arg-type]
onnx_input = {"pixel_values": encoded, "attention_mask": attention_mask}
onnx_input = self._preprocess_onnx_image_input(onnx_input, **kwargs)
model_output = self.model.run(None, onnx_input) # type: ignore[union-attr]
return OnnxOutputContext(
model_output=model_output[0],
attention_mask=attention_mask, # type: ignore[arg-type]
metadata=metadata,
)
@staticmethod
def _process_nested_patches(
processed: list[list[NumpyArray]],
) -> tuple[NumpyArray, NumpyArray, dict[str, Any]]:
"""
Process nested image patches (from ImageSplitter).
Args:
processed: List of patch lists, one per image [[img1_patches], [img2_patches], ...]
Returns:
tuple: (encoded array, attention_mask, metadata)
- encoded: (batch_size, max_patches, C, H, W)
- attention_mask: (batch_size, max_patches) with 1 for real patches, 0 for padding
- metadata: Dict with 'patch_counts' key
"""
patch_counts = [len(patches) for patches in processed]
max_patches = max(patch_counts)
# Get dimensions from first patch
channels, height, width = processed[0][0].shape
batch_size = len(processed)
# Create padded array
encoded = np.zeros(
(batch_size, max_patches, channels, height, width), dtype=processed[0][0].dtype
)
# Create attention mask (1 for real patches, 0 for padding)
attention_mask = np.zeros((batch_size, max_patches), dtype=np.int64)
# Fill in patches and attention mask
for i, patches in enumerate(processed):
for j, patch in enumerate(patches):
encoded[i, j] = patch
attention_mask[i, j] = 1
metadata = {"patch_counts": patch_counts}
return encoded, attention_mask, metadata # type: ignore[return-value]
def _preprocess_onnx_image_input(
self, onnx_input: dict[str, np.ndarray], **kwargs: Any
) -> dict[str, NumpyArray]:
"""
Add text input placeholders for image data, following Idefics3 processing logic.
Constructs input_ids dynamically based on the actual number of image patches,
using the same token expansion logic as Idefics3Processor.
Args:
onnx_input: Dict with 'pixel_values' (batch, num_patches, C, H, W)
and 'attention_mask' (batch, num_patches) indicating real patches
**kwargs: Additional arguments
Returns:
Updated onnx_input with 'input_ids' and updated 'attention_mask' for token sequence
"""
# The attention_mask in onnx_input has a shape of (batch_size, num_patches),
# and should be used to create an attention mask matching the input_ids shape.
patch_attention_mask = onnx_input["attention_mask"]
pixel_values = onnx_input["pixel_values"]
batch_size = pixel_values.shape[0]
batch_input_ids = []
# Build input_ids for each image based on its actual patch count
for i in range(batch_size):
# Count real patches (non-padded) from attention mask
patch_count = int(np.sum(patch_attention_mask[i]))
# Compute rows/cols from patch count
rows, cols = self._compute_rows_cols_from_patches(patch_count)
# Build input_ids for this image
input_ids = self._build_input_ids_for_image(rows, cols)
batch_input_ids.append(input_ids)
# Pad sequences to max length in batch
max_len = max(len(ids) for ids in batch_input_ids)
# Get padding config from tokenizer
padding_direction = self.tokenizer.padding["direction"] # type: ignore[index,union-attr]
pad_token_id = self.tokenizer.padding["pad_id"] # type: ignore[index,union-attr]
# Initialize with pad token
padded_input_ids = np.full((batch_size, max_len), pad_token_id, dtype=np.int64)
attention_mask = np.zeros((batch_size, max_len), dtype=np.int64)
for i, input_ids in enumerate(batch_input_ids):
seq_len = len(input_ids)
if padding_direction == "left":
# Left padding: place tokens at the END of the array
start_idx = max_len - seq_len
padded_input_ids[i, start_idx:] = input_ids
attention_mask[i, start_idx:] = 1
else:
# Right padding: place tokens at the START of the array
padded_input_ids[i, :seq_len] = input_ids
attention_mask[i, :seq_len] = 1
onnx_input["input_ids"] = padded_input_ids
# Update attention_mask with token-level data
onnx_input["attention_mask"] = attention_mask
return onnx_input
@staticmethod
def _compute_rows_cols_from_patches(patch_count: int) -> tuple[int, int]:
if patch_count <= 1:
return 0, 0
# Subtract 1 for the global image
grid_patches = patch_count - 1
# Find rows and cols (assume square or near-square grid)
rows = int(grid_patches**0.5)
cols = grid_patches // rows
# Verify the calculation
if rows * cols + 1 != patch_count:
# Handle non-square grids
for r in range(1, grid_patches + 1):
if grid_patches % r == 0:
c = grid_patches // r
if r * c + 1 == patch_count:
return r, c
# Fallback: treat as unsplit
return 0, 0
return rows, cols
def _create_single_image_prompt_string(self) -> str:
return (
"<fake_token_around_image>"
+ "<global-img>"
+ "<image>" * self.image_seq_len # type: ignore[operator]
+ "<fake_token_around_image>"
)
def _create_split_image_prompt_string(self, rows: int, cols: int) -> str:
text_split_images = ""
# Add tokens for each patch in the grid
for n_h in range(rows):
for n_w in range(cols):
text_split_images += (
"<fake_token_around_image>"
+ f"<row_{n_h + 1}_col_{n_w + 1}>"
+ "<image>" * self.image_seq_len # type: ignore[operator]
)
text_split_images += "\n"
# Add global image at the end
text_split_images += (
"\n<fake_token_around_image>"
+ "<global-img>"
+ "<image>" * self.image_seq_len # type: ignore[operator]
+ "<fake_token_around_image>"
)
return text_split_images
def _build_input_ids_for_image(self, rows: int, cols: int) -> np.ndarray:
# Create the appropriate image prompt string
if rows == 0 and cols == 0:
image_prompt_tokens = self._create_single_image_prompt_string()
else:
image_prompt_tokens = self._create_split_image_prompt_string(rows, cols)
# Replace <image> in visual prompt with expanded tokens
# The visual prompt is: "<|begin_of_text|>User:<image>Describe the image.<end_of_utterance>\nAssistant:"
expanded_prompt = self.VISUAL_PROMPT_PREFIX.replace("<image>", image_prompt_tokens)
# Tokenize the complete prompt
encoded = self.tokenizer.encode(expanded_prompt) # type: ignore[union-attr]
# Convert to numpy array
return np.array(encoded.ids, dtype=np.int64)
def _post_process_onnx_image_output(
self,
output: OnnxOutputContext,
) -> Iterable[NumpyArray]:
"""
Post-process the ONNX model output to convert it into a usable format.
Args:
output (OnnxOutputContext): The raw output from the ONNX model.
Returns:
Iterable[NumpyArray]: Post-processed output as NumPy arrays.
"""
assert self.model_description.dim is not None, "Model dim is not defined"
return output.model_output.reshape(
output.model_output.shape[0], -1, self.model_description.dim
)
def embed_text(
self,
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
"""
Encode a list of documents into list of embeddings.
Args:
documents: Iterator of documents or single document to embed
batch_size: Batch size for encoding -- higher values will use more memory, but be faster
parallel:
If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.
If 0, use all available cores.
If None, don't use data-parallel processing, use default onnxruntime threading instead.
Returns:
List of embeddings, one per document
"""
yield from self._embed_documents(
model_name=self.model_name,
cache_dir=str(self.cache_dir),
documents=documents,
batch_size=batch_size,
parallel=parallel,
providers=self.providers,
cuda=self.cuda,
device_ids=self.device_ids,
local_files_only=self._local_files_only,
specific_model_path=self._specific_model_path,
extra_session_options=self._extra_session_options,
**kwargs,
)
def embed_image(
self,
images: ImageInput | Iterable[ImageInput],
batch_size: int = 16,
parallel: Optional[int] = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
"""
Encode a list of images into list of embeddings.
Args:
images: Iterator of image paths or single image path to embed
batch_size: Batch size for encoding -- higher values will use more memory, but be faster
parallel:
If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.
If 0, use all available cores.
If None, don't use data-parallel processing, use default onnxruntime threading instead.
Returns:
List of embeddings, one per document
"""
yield from self._embed_images(
model_name=self.model_name,
cache_dir=str(self.cache_dir),
images=images,
batch_size=batch_size,
parallel=parallel,
providers=self.providers,
cuda=self.cuda,
device_ids=self.device_ids,
local_files_only=self._local_files_only,
specific_model_path=self._specific_model_path,
extra_session_options=self._extra_session_options,
**kwargs,
)
@classmethod
def _get_text_worker_class(cls) -> Type[TextEmbeddingWorker[NumpyArray]]:
return ColModernVBERTTextEmbeddingWorker
@classmethod
def _get_image_worker_class(cls) -> Type[ImageEmbeddingWorker[NumpyArray]]:
return ColModernVBERTImageEmbeddingWorker
class ColModernVBERTTextEmbeddingWorker(TextEmbeddingWorker[NumpyArray]):
def init_embedding(self, model_name: str, cache_dir: str, **kwargs: Any) -> ColModernVBERT:
return ColModernVBERT(
model_name=model_name,
cache_dir=cache_dir,
threads=1,
**kwargs,
)
class ColModernVBERTImageEmbeddingWorker(ImageEmbeddingWorker[NumpyArray]):
def init_embedding(self, model_name: str, cache_dir: str, **kwargs: Any) -> ColModernVBERT:
return ColModernVBERT(
model_name=model_name,
cache_dir=cache_dir,
threads=1,
**kwargs,
)
@@ -1,11 +1,11 @@
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Iterable, Sequence, Type
import numpy as np
from tokenizers import Encoding
from fastembed.common import OnnxProvider, ImageInput
from fastembed.common.onnx_model import OnnxOutputContext
from fastembed.common.types import NumpyArray
from fastembed.common.types import NumpyArray, Device
from fastembed.common.utils import define_cache_dir, iter_batch
from fastembed.late_interaction_multimodal.late_interaction_multimodal_embedding_base import (
LateInteractionMultimodalEmbeddingBase,
@@ -46,14 +46,14 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
device_id: Optional[int] = None,
specific_model_path: Optional[str] = None,
device_id: int | None = None,
specific_model_path: str | None = None,
**kwargs: Any,
):
"""
@@ -65,10 +65,11 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.
providers (Optional[Sequence[OnnxProvider]], optional): The list of onnxruntime providers to use.
Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to False.
cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to Device.AUTO.
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, mutually exclusive
with `providers`. Defaults to None.
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
@@ -87,7 +88,7 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
self.cuda = cuda
# This device_id will be used if we need to load model in current process
self.device_id: Optional[int] = None
self.device_id: int | None = None
if device_id is not None:
self.device_id = device_id
elif self.device_ids is not None:
@@ -174,7 +175,7 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
def token_count(
self,
texts: Union[str, Iterable[str]],
texts: str | Iterable[str],
batch_size: int = 1024,
include_extension: bool = False,
**kwargs: Any,
@@ -227,9 +228,9 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
def embed_text(
self,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
"""
@@ -263,9 +264,9 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
def embed_image(
self,
images: Union[ImageInput, Iterable[ImageInput]],
images: ImageInput | Iterable[ImageInput],
batch_size: int = 16,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
"""
@@ -1,9 +1,10 @@
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Iterable, Sequence, Type
from dataclasses import asdict
from fastembed.common import OnnxProvider, ImageInput
from fastembed.common.types import NumpyArray
from fastembed.common.types import NumpyArray, Device
from fastembed.late_interaction_multimodal.colpali import ColPali
from fastembed.late_interaction_multimodal.colmodernvbert import ColModernVBERT
from fastembed.late_interaction_multimodal.late_interaction_multimodal_embedding_base import (
LateInteractionMultimodalEmbeddingBase,
@@ -12,7 +13,10 @@ from fastembed.common.model_description import DenseModelDescription
class LateInteractionMultimodalEmbedding(LateInteractionMultimodalEmbeddingBase):
EMBEDDINGS_REGISTRY: list[Type[LateInteractionMultimodalEmbeddingBase]] = [ColPali]
EMBEDDINGS_REGISTRY: list[Type[LateInteractionMultimodalEmbeddingBase]] = [
ColPali,
ColModernVBERT,
]
@classmethod
def list_supported_models(cls) -> list[dict[str, Any]]:
@@ -54,11 +58,11 @@ class LateInteractionMultimodalEmbedding(LateInteractionMultimodalEmbeddingBase)
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
**kwargs: Any,
):
@@ -104,7 +108,7 @@ class LateInteractionMultimodalEmbedding(LateInteractionMultimodalEmbeddingBase)
ValueError: If the model name is not found in the supported models.
"""
descriptions = cls._list_supported_models()
embedding_size: Optional[int] = None
embedding_size: int | None = None
for description in descriptions:
if description.model.lower() == model_name.lower():
embedding_size = description.dim
@@ -119,9 +123,9 @@ class LateInteractionMultimodalEmbedding(LateInteractionMultimodalEmbeddingBase)
def embed_text(
self,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
"""
@@ -142,9 +146,9 @@ class LateInteractionMultimodalEmbedding(LateInteractionMultimodalEmbeddingBase)
def embed_image(
self,
images: Union[ImageInput, Iterable[ImageInput]],
images: ImageInput | Iterable[ImageInput],
batch_size: int = 16,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
"""
@@ -165,7 +169,7 @@ class LateInteractionMultimodalEmbedding(LateInteractionMultimodalEmbeddingBase)
def token_count(
self,
texts: Union[str, Iterable[str]],
texts: str | Iterable[str],
batch_size: int = 1024,
include_extension: bool = False,
**kwargs: Any,
@@ -1,4 +1,4 @@
from typing import Iterable, Optional, Union, Any
from typing import Iterable, Any
from fastembed.common import ImageInput
@@ -11,21 +11,21 @@ class LateInteractionMultimodalEmbeddingBase(ModelManagement[DenseModelDescripti
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
cache_dir: str | None = None,
threads: int | None = None,
**kwargs: Any,
):
self.model_name = model_name
self.cache_dir = cache_dir
self.threads = threads
self._local_files_only = kwargs.pop("local_files_only", False)
self._embedding_size: Optional[int] = None
self._embedding_size: int | None = None
def embed_text(
self,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
"""
@@ -47,9 +47,9 @@ class LateInteractionMultimodalEmbeddingBase(ModelManagement[DenseModelDescripti
def embed_image(
self,
images: Union[ImageInput, Iterable[ImageInput]],
images: ImageInput | Iterable[ImageInput],
batch_size: int = 16,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
"""
@@ -79,7 +79,7 @@ class LateInteractionMultimodalEmbeddingBase(ModelManagement[DenseModelDescripti
def token_count(
self,
texts: Union[str, Iterable[str]],
texts: str | Iterable[str],
**kwargs: Any,
) -> int:
"""Returns the number of tokens in the texts."""
@@ -2,7 +2,7 @@ import contextlib
import os
from multiprocessing import get_all_start_methods
from pathlib import Path
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Iterable, Sequence, Type
import numpy as np
from PIL import Image
@@ -11,19 +11,19 @@ from tokenizers import Encoding, Tokenizer
from fastembed.common import OnnxProvider, ImageInput
from fastembed.common.onnx_model import EmbeddingWorker, OnnxModel, OnnxOutputContext, T
from fastembed.common.preprocessor_utils import load_tokenizer, load_preprocessor
from fastembed.common.types import NumpyArray
from fastembed.common.types import NumpyArray, Device
from fastembed.common.utils import iter_batch
from fastembed.image.transform.operators import Compose
from fastembed.parallel_processor import ParallelWorkerPool
class OnnxMultimodalModel(OnnxModel[T]):
ONNX_OUTPUT_NAMES: Optional[list[str]] = None
ONNX_OUTPUT_NAMES: list[str] | None = None
def __init__(self) -> None:
super().__init__()
self.tokenizer: Optional[Tokenizer] = None
self.processor: Optional[Compose] = None
self.tokenizer: Tokenizer | None = None
self.processor: Compose | None = None
self.special_token_to_id: dict[str, int] = {}
def _preprocess_onnx_text_input(
@@ -60,11 +60,11 @@ class OnnxMultimodalModel(OnnxModel[T]):
self,
model_dir: Path,
model_file: str,
threads: Optional[int],
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_id: Optional[int] = None,
extra_session_options: Optional[dict[str, Any]] = None,
threads: int | None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_id: int | None = None,
extra_session_options: dict[str, Any] | None = None,
) -> None:
super()._load_onnx_model(
model_dir=model_dir,
@@ -116,15 +116,15 @@ class OnnxMultimodalModel(OnnxModel[T]):
self,
model_name: str,
cache_dir: str,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
parallel: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
local_files_only: bool = False,
specific_model_path: Optional[str] = None,
extra_session_options: Optional[dict[str, Any]] = None,
specific_model_path: str | None = None,
extra_session_options: dict[str, Any] | None = None,
**kwargs: Any,
) -> Iterable[T]:
is_small = False
@@ -170,9 +170,11 @@ class OnnxMultimodalModel(OnnxModel[T]):
yield from self._post_process_onnx_text_output(batch) # type: ignore
def onnx_embed_image(self, images: list[ImageInput], **kwargs: Any) -> OnnxOutputContext:
with contextlib.ExitStack():
with contextlib.ExitStack() as stack:
image_files = [
Image.open(image) if not isinstance(image, Image.Image) else image
stack.enter_context(Image.open(image))
if not isinstance(image, Image.Image)
else image
for image in images
]
assert self.processor is not None, "Processor is not initialized"
@@ -187,15 +189,15 @@ class OnnxMultimodalModel(OnnxModel[T]):
self,
model_name: str,
cache_dir: str,
images: Union[Iterable[ImageInput], ImageInput],
images: Iterable[ImageInput] | ImageInput,
batch_size: int = 256,
parallel: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
parallel: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
local_files_only: bool = False,
specific_model_path: Optional[str] = None,
extra_session_options: Optional[dict[str, Any]] = None,
specific_model_path: str | None = None,
extra_session_options: dict[str, Any] | None = None,
**kwargs: Any,
) -> Iterable[T]:
is_small = False
+10 -9
View File
@@ -8,8 +8,9 @@ from multiprocessing.context import BaseContext
from multiprocessing.process import BaseProcess
from multiprocessing.sharedctypes import Synchronized as BaseValue
from queue import Empty
from typing import Any, Iterable, Optional, Type
from typing import Any, Iterable, Type
from fastembed.common.types import Device
# Single item should be processed in less than:
processing_timeout = 10 * 60 # seconds
@@ -38,7 +39,7 @@ def _worker(
output_queue: Queue,
num_active_workers: BaseValue,
worker_id: int,
kwargs: Optional[dict[str, Any]] = None,
kwargs: dict[str, Any] | None = None,
) -> None:
"""
A worker that pulls data pints off the input queue, and places the execution result on the output queue.
@@ -93,21 +94,21 @@ class ParallelWorkerPool:
self,
num_workers: int,
worker: Type[Worker],
start_method: Optional[str] = None,
device_ids: Optional[list[int]] = None,
cuda: bool = False,
start_method: str | None = None,
device_ids: list[int] | None = None,
cuda: bool | Device = Device.AUTO,
):
self.worker_class = worker
self.num_workers = num_workers
self.input_queue: Optional[Queue] = None
self.output_queue: Optional[Queue] = None
self.input_queue: Queue | None = None
self.output_queue: Queue | None = None
self.ctx: BaseContext = get_context(start_method)
self.processes: list[BaseProcess] = []
self.queue_size = self.num_workers * max_internal_batch_size
self.emergency_shutdown = False
self.device_ids = device_ids
self.cuda = cuda
self.num_active_workers: Optional[BaseValue] = None
self.num_active_workers: BaseValue | None = None
def start(self, **kwargs: Any) -> None:
self.input_queue = self.ctx.Queue(self.queue_size)
@@ -220,7 +221,7 @@ class ParallelWorkerPool:
f"Worker PID: {process.pid} terminated unexpectedly with code {process.exitcode}"
)
def join_or_terminate(self, timeout: Optional[int] = 1) -> None:
def join_or_terminate(self, timeout: int = 1) -> None:
"""
Emergency shutdown
@param timeout:
+1 -3
View File
@@ -1,5 +1,3 @@
from typing import Union
import numpy as np
from fastembed.common.types import NumpyArray
@@ -11,7 +9,7 @@ from fastembed.late_interaction_multimodal.late_interaction_multimodal_embedding
)
MultiVectorModel = Union[LateInteractionTextEmbeddingBase, LateInteractionMultimodalEmbeddingBase]
MultiVectorModel = LateInteractionTextEmbeddingBase | LateInteractionMultimodalEmbeddingBase
MAX_HAMMING_DISTANCE = 65 # 64 bits + 1
POPCOUNT_LUT = np.array([bin(x).count("1") for x in range(256)], dtype=np.uint8)
@@ -1,7 +1,8 @@
from typing import Optional, Sequence, Any
from typing import Sequence, Any
from fastembed.common import OnnxProvider
from fastembed.common.model_description import BaseModelDescription
from fastembed.common.types import Device
from fastembed.rerank.cross_encoder.onnx_text_cross_encoder import OnnxTextCrossEncoder
@@ -11,14 +12,14 @@ class CustomTextCrossEncoder(OnnxTextCrossEncoder):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
device_id: Optional[int] = None,
specific_model_path: Optional[str] = None,
device_id: int | None = None,
specific_model_path: str | None = None,
**kwargs: Any,
):
super().__init__(
@@ -1,9 +1,10 @@
from typing import Any, Iterable, Optional, Sequence, Type
from typing import Any, Iterable, Sequence, Type
from loguru import logger
from fastembed.common import OnnxProvider
from fastembed.common.onnx_model import OnnxOutputContext
from fastembed.common.types import Device
from fastembed.common.utils import define_cache_dir
from fastembed.rerank.cross_encoder.onnx_text_model import (
OnnxCrossEncoderModel,
@@ -77,14 +78,14 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
device_id: Optional[int] = None,
specific_model_path: Optional[str] = None,
device_id: int | None = None,
specific_model_path: str | None = None,
**kwargs: Any,
):
"""
@@ -96,10 +97,11 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.
providers (Optional[Sequence[OnnxProvider]], optional): The list of onnxruntime providers to use.
Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to False.
cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to Device.AUTO.
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, mutually exclusive
with `providers`. Defaults to None.
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
@@ -124,7 +126,7 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
)
# This device_id will be used if we need to load model in current process
self.device_id: Optional[int] = None
self.device_id: int | None = None
if device_id is not None:
self.device_id = device_id
elif self.device_ids is not None:
@@ -180,7 +182,7 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
self,
pairs: Iterable[tuple[str, str]],
batch_size: int = 64,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[float]:
yield from self._rerank_pairs(
@@ -1,7 +1,7 @@
import os
from multiprocessing import get_all_start_methods
from pathlib import Path
from typing import Any, Iterable, Optional, Sequence, Type
from typing import Any, Iterable, Sequence, Type
import numpy as np
from tokenizers import Encoding
@@ -12,14 +12,14 @@ from fastembed.common.onnx_model import (
OnnxOutputContext,
OnnxProvider,
)
from fastembed.common.types import NumpyArray
from fastembed.common.types import NumpyArray, Device
from fastembed.common.preprocessor_utils import load_tokenizer
from fastembed.common.utils import iter_batch
from fastembed.parallel_processor import ParallelWorkerPool
class OnnxCrossEncoderModel(OnnxModel[float]):
ONNX_OUTPUT_NAMES: Optional[list[str]] = None
ONNX_OUTPUT_NAMES: list[str] | None = None
@classmethod
def _get_worker_class(cls) -> Type["TextRerankerWorker"]:
@@ -29,11 +29,11 @@ class OnnxCrossEncoderModel(OnnxModel[float]):
self,
model_dir: Path,
model_file: str,
threads: Optional[int],
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_id: Optional[int] = None,
extra_session_options: Optional[dict[str, Any]] = None,
threads: int | None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_id: int | None = None,
extra_session_options: dict[str, Any] | None = None,
) -> None:
super()._load_onnx_model(
model_dir=model_dir,
@@ -92,13 +92,13 @@ class OnnxCrossEncoderModel(OnnxModel[float]):
cache_dir: str,
pairs: Iterable[tuple[str, str]],
batch_size: int,
parallel: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
parallel: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
local_files_only: bool = False,
specific_model_path: Optional[str] = None,
extra_session_options: Optional[dict[str, Any]] = None,
specific_model_path: str | None = None,
extra_session_options: dict[str, Any] | None = None,
**kwargs: Any,
) -> Iterable[float]:
is_small = False
@@ -1,7 +1,8 @@
from typing import Any, Iterable, Optional, Sequence, Type
from typing import Any, Iterable, Sequence, Type
from dataclasses import asdict
from fastembed.common import OnnxProvider
from fastembed.common.types import Device
from fastembed.rerank.cross_encoder.onnx_text_cross_encoder import OnnxTextCrossEncoder
from fastembed.rerank.cross_encoder.custom_text_cross_encoder import CustomTextCrossEncoder
@@ -53,11 +54,11 @@ class TextCrossEncoder(TextCrossEncoderBase):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
**kwargs: Any,
):
@@ -102,7 +103,7 @@ class TextCrossEncoder(TextCrossEncoderBase):
self,
pairs: Iterable[tuple[str, str]],
batch_size: int = 64,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[float]:
"""
@@ -140,7 +141,7 @@ class TextCrossEncoder(TextCrossEncoderBase):
description: str = "",
license: str = "",
size_in_gb: float = 0.0,
additional_files: Optional[list[str]] = None,
additional_files: list[str] | None = None,
) -> None:
registered_models = cls._list_supported_models()
for registered_model in registered_models:
@@ -1,4 +1,4 @@
from typing import Any, Iterable, Optional
from typing import Any, Iterable
from fastembed.common.model_description import BaseModelDescription
from fastembed.common.model_management import ModelManagement
@@ -8,8 +8,8 @@ class TextCrossEncoderBase(ModelManagement[BaseModelDescription]):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
cache_dir: str | None = None,
threads: int | None = None,
**kwargs: Any,
):
self.model_name = model_name
@@ -41,7 +41,7 @@ class TextCrossEncoderBase(ModelManagement[BaseModelDescription]):
self,
pairs: Iterable[tuple[str, str]],
batch_size: int = 64,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[float]:
"""Rerank query-document pairs.
+10 -12
View File
@@ -2,7 +2,7 @@ import os
from collections import defaultdict
from multiprocessing import get_all_start_methods
from pathlib import Path
from typing import Any, Iterable, Optional, Type, Union
from typing import Any, Iterable, Type
import mmh3
import numpy as np
@@ -91,14 +91,14 @@ class Bm25(SparseTextEmbeddingBase):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
cache_dir: str | None = None,
k: float = 1.2,
b: float = 0.75,
avg_len: float = 256.0,
language: str = "english",
token_max_length: int = 40,
disable_stemmer: bool = False,
specific_model_path: Optional[str] = None,
specific_model_path: str | None = None,
**kwargs: Any,
):
super().__init__(model_name, cache_dir, **kwargs)
@@ -158,11 +158,11 @@ class Bm25(SparseTextEmbeddingBase):
self,
model_name: str,
cache_dir: str,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
parallel: int | None = None,
local_files_only: bool = False,
specific_model_path: Optional[str] = None,
specific_model_path: str | None = None,
) -> Iterable[SparseEmbedding]:
is_small = False
@@ -205,9 +205,9 @@ class Bm25(SparseTextEmbeddingBase):
def embed(
self,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[SparseEmbedding]:
"""
@@ -268,7 +268,7 @@ class Bm25(SparseTextEmbeddingBase):
embeddings.append(SparseEmbedding.from_dict(token_id2value))
return embeddings
def token_count(self, texts: Union[str, Iterable[str]], **kwargs: Any) -> int:
def token_count(self, texts: str | Iterable[str], **kwargs: Any) -> int:
token_num = 0
texts = [texts] if isinstance(texts, str) else texts
for text in texts:
@@ -311,9 +311,7 @@ class Bm25(SparseTextEmbeddingBase):
def compute_token_id(cls, token: str) -> int:
return abs(mmh3.hash(token))
def query_embed(
self, query: Union[str, Iterable[str]], **kwargs: Any
) -> Iterable[SparseEmbedding]:
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[SparseEmbedding]:
"""To emulate BM25 behaviour, we don't need to use weights in the query, and
it's enough to just hash the tokens and assign a weight of 1.0 to them.
"""
+18 -18
View File
@@ -1,7 +1,7 @@
import math
import string
from pathlib import Path
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Iterable, Sequence, Type
import mmh3
import numpy as np
@@ -9,6 +9,7 @@ from py_rust_stemmers import SnowballStemmer
from fastembed.common import OnnxProvider
from fastembed.common.onnx_model import OnnxOutputContext
from fastembed.common.types import Device
from fastembed.common.utils import define_cache_dir
from fastembed.sparse.sparse_embedding_base import (
SparseEmbedding,
@@ -65,15 +66,15 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
alpha: float = 0.5,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
device_id: Optional[int] = None,
specific_model_path: Optional[str] = None,
device_id: int | None = None,
specific_model_path: str | None = None,
**kwargs: Any,
):
"""
@@ -87,10 +88,11 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
alpha (float, optional): Parameter, that defines the importance of the token weight in the document
versus the importance of the token frequency in the corpus. Defaults to 0.5, based on empirical testing.
It is recommended to only change this parameter based on training data for a specific dataset.
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to False.
cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to Device.AUTO.
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, mutually exclusive
with `providers`. Defaults to None.
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
@@ -110,7 +112,7 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
self.cuda = cuda
# This device_id will be used if we need to load model in current process
self.device_id: Optional[int] = None
self.device_id: int | None = None
if device_id is not None:
self.device_id = device_id
elif self.device_ids is not None:
@@ -282,9 +284,9 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
def embed(
self,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[SparseEmbedding]:
"""
@@ -325,9 +327,7 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
result[token_id] = 1.0
return result
def query_embed(
self, query: Union[str, Iterable[str]], **kwargs: Any
) -> Iterable[SparseEmbedding]:
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[SparseEmbedding]:
"""
To emulate BM25 behaviour, we don't need to use smart weights in the query, and
it's enough to just hash the tokens and assign a weight of 1.0 to them.
@@ -353,7 +353,7 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
return Bm42TextEmbeddingWorker
def token_count(
self, texts: Union[str, Iterable[str]], batch_size: int = 1024, **kwargs: Any
self, texts: str | Iterable[str], batch_size: int = 1024, **kwargs: Any
) -> int:
if not hasattr(self, "model") or self.model is None:
self.load_onnx_model() # loads the tokenizer as well
+22 -22
View File
@@ -1,6 +1,6 @@
from pathlib import Path
from typing import Any, Optional, Sequence, Iterable, Union, Type
from typing import Any, Sequence, Iterable, Type
import numpy as np
from numpy.typing import NDArray
@@ -10,6 +10,7 @@ from tokenizers import Tokenizer
from fastembed.common.model_description import SparseModelDescription, ModelSource
from fastembed.common.onnx_model import OnnxOutputContext
from fastembed.common import OnnxProvider
from fastembed.common.types import Device
from fastembed.common.utils import define_cache_dir
from fastembed.sparse.sparse_embedding_base import (
SparseEmbedding,
@@ -72,17 +73,17 @@ class MiniCOIL(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
k: float = 1.2,
b: float = 0.75,
avg_len: float = 150.0,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
device_id: Optional[int] = None,
specific_model_path: Optional[str] = None,
device_id: int | None = None,
specific_model_path: str | None = None,
**kwargs: Any,
):
"""
@@ -98,10 +99,11 @@ class MiniCOIL(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
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.
cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to Device.AUTO.
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, mutually exclusive
with `providers`. Defaults to None.
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
@@ -124,15 +126,15 @@ class MiniCOIL(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
self.avg_len = avg_len
# Initialize class attributes
self.tokenizer: Optional[Tokenizer] = None
self.tokenizer: Tokenizer | None = None
self.invert_vocab: dict[int, str] = {}
self.special_tokens: set[str] = set()
self.special_tokens_ids: set[int] = set()
self.stopwords: set[str] = set()
self.vocab_resolver: Optional[VocabResolver] = None
self.encoder: Optional[Encoder] = None
self.output_dim: Optional[int] = None
self.sparse_vector_converter: Optional[SparseVectorConverter] = None
self.vocab_resolver: VocabResolver | None = None
self.encoder: Encoder | None = None
self.output_dim: int | None = None
self.sparse_vector_converter: SparseVectorConverter | None = None
self.model_description = self._get_model_description(model_name)
self.cache_dir = str(define_cache_dir(cache_dir))
@@ -188,15 +190,15 @@ class MiniCOIL(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
)
def token_count(
self, texts: Union[str, Iterable[str]], batch_size: int = 1024, **kwargs: Any
self, texts: str | Iterable[str], batch_size: int = 1024, **kwargs: Any
) -> int:
return self._token_count(texts, batch_size=batch_size, **kwargs)
def embed(
self,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[SparseEmbedding]:
"""
@@ -233,9 +235,7 @@ class MiniCOIL(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
**kwargs,
)
def query_embed(
self, query: Union[str, Iterable[str]], **kwargs: Any
) -> Iterable[SparseEmbedding]:
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[SparseEmbedding]:
"""
Encode a list of queries into list of embeddings.
"""
+8 -10
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass
from typing import Iterable, Optional, Union, Any
from typing import Iterable, Any
import numpy as np
from numpy.typing import NDArray
@@ -12,7 +12,7 @@ from fastembed.common.model_management import ModelManagement
@dataclass
class SparseEmbedding:
values: NumpyArray
indices: Union[NDArray[np.int64], NDArray[np.int32]]
indices: NDArray[np.int64] | NDArray[np.int32]
def as_object(self) -> dict[str, NumpyArray]:
return {
@@ -35,8 +35,8 @@ class SparseTextEmbeddingBase(ModelManagement[SparseModelDescription]):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
cache_dir: str | None = None,
threads: int | None = None,
**kwargs: Any,
):
self.model_name = model_name
@@ -46,9 +46,9 @@ class SparseTextEmbeddingBase(ModelManagement[SparseModelDescription]):
def embed(
self,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[SparseEmbedding]:
raise NotImplementedError()
@@ -68,9 +68,7 @@ class SparseTextEmbeddingBase(ModelManagement[SparseModelDescription]):
# This is model-specific, so that different models can have specialized implementations
yield from self.embed(texts, **kwargs)
def query_embed(
self, query: Union[str, Iterable[str]], **kwargs: Any
) -> Iterable[SparseEmbedding]:
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[SparseEmbedding]:
"""
Embeds queries
@@ -87,6 +85,6 @@ class SparseTextEmbeddingBase(ModelManagement[SparseModelDescription]):
else:
yield from self.embed(query, **kwargs)
def token_count(self, texts: Union[str, Iterable[str]], **kwargs: Any) -> int:
def token_count(self, texts: str | Iterable[str], **kwargs: Any) -> int:
"""Returns the number of tokens in the texts."""
raise NotImplementedError("Subclasses must implement this method")
+11 -12
View File
@@ -1,7 +1,8 @@
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Iterable, Sequence, Type
from dataclasses import asdict
from fastembed.common import OnnxProvider
from fastembed.common.types import Device
from fastembed.sparse.bm25 import Bm25
from fastembed.sparse.bm42 import Bm42
from fastembed.sparse.minicoil import MiniCOIL
@@ -53,11 +54,11 @@ class SparseTextEmbedding(SparseTextEmbeddingBase):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
**kwargs: Any,
):
@@ -93,9 +94,9 @@ class SparseTextEmbedding(SparseTextEmbeddingBase):
def embed(
self,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[SparseEmbedding]:
"""
@@ -115,9 +116,7 @@ class SparseTextEmbedding(SparseTextEmbeddingBase):
"""
yield from self.model.embed(documents, batch_size, parallel, **kwargs)
def query_embed(
self, query: Union[str, Iterable[str]], **kwargs: Any
) -> Iterable[SparseEmbedding]:
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[SparseEmbedding]:
"""
Embeds queries
@@ -130,7 +129,7 @@ class SparseTextEmbedding(SparseTextEmbeddingBase):
yield from self.model.query_embed(query, **kwargs)
def token_count(
self, texts: Union[str, Iterable[str]], batch_size: int = 1024, **kwargs: Any
self, texts: str | Iterable[str], batch_size: int = 1024, **kwargs: Any
) -> int:
"""Returns the number of tokens in the texts.
+17 -15
View File
@@ -1,8 +1,9 @@
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Iterable, Sequence, Type
import numpy as np
from fastembed.common import OnnxProvider
from fastembed.common.onnx_model import OnnxOutputContext
from fastembed.common.types import Device
from fastembed.common.utils import define_cache_dir
from fastembed.sparse.sparse_embedding_base import (
SparseEmbedding,
@@ -54,7 +55,7 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
yield SparseEmbedding(values=scores, indices=indices)
def token_count(
self, texts: Union[str, Iterable[str]], batch_size: int = 1024, **kwargs: Any
self, texts: str | Iterable[str], batch_size: int = 1024, **kwargs: Any
) -> int:
return self._token_count(texts, batch_size=batch_size, **kwargs)
@@ -70,14 +71,14 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
device_id: Optional[int] = None,
specific_model_path: Optional[str] = None,
device_id: int | None = None,
specific_model_path: str | None = None,
**kwargs: Any,
):
"""
@@ -89,10 +90,11 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.
providers (Optional[Sequence[OnnxProvider]], optional): The list of onnxruntime providers to use.
Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to False.
cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to Device.
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, mutually exclusive
with `providers`. Defaults to None.
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
@@ -111,7 +113,7 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
self.cuda = cuda
# This device_id will be used if we need to load model in current process
self.device_id: Optional[int] = None
self.device_id: int | None = None
if device_id is not None:
self.device_id = device_id
elif self.device_ids is not None:
@@ -144,9 +146,9 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
def embed(
self,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[SparseEmbedding]:
"""
@@ -1,12 +1,11 @@
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 mmh3
import numpy as np
from py_rust_stemmers import SnowballStemmer
from fastembed.common.utils import get_all_punctuation, remove_non_alphanumeric
from fastembed.sparse.sparse_embedding_base import SparseEmbedding
GAP = 32000
@@ -16,16 +15,16 @@ INT32_MAX = 2**31 - 1
@dataclass
class WordEmbedding:
word: str
forms: List[str]
forms: list[str]
count: int
word_id: int
embedding: List[float]
embedding: list[float]
class SparseVectorConverter:
def __init__(
self,
stopwords: Set[str],
stopwords: set[str],
stemmer: SnowballStemmer,
k: float = 1.2,
b: float = 0.75,
@@ -58,15 +57,15 @@ class SparseVectorConverter:
return res
@classmethod
def normalize_vector(cls, vector: List[float]) -> List[float]:
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]:
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.
@@ -85,7 +84,7 @@ class SparseVectorConverter:
}
"""
new_sentence_embedding: Dict[str, WordEmbedding] = {}
new_sentence_embedding: dict[str, WordEmbedding] = {}
for word, embedding in sentence_embedding.items():
# embedding = {
@@ -127,7 +126,7 @@ class SparseVectorConverter:
def embedding_to_vector(
self,
sentence_embedding: Dict[str, WordEmbedding],
sentence_embedding: dict[str, WordEmbedding],
embedding_size: int,
vocab_size: int,
) -> SparseEmbedding:
@@ -156,14 +155,14 @@ class SparseVectorConverter:
"""
indices: List[int] = []
values: List[float] = []
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.
@@ -171,9 +170,7 @@ class SparseVectorConverter:
# 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
unknown_words_shift = ((vocab_size * embedding_size) // GAP + 2) * GAP
sentence_embedding_cleaned = self.clean_words(sentence_embedding)
# Calculate sentence length after cleaning
@@ -208,7 +205,7 @@ class SparseVectorConverter:
def embedding_to_vector_query(
self,
sentence_embedding: Dict[str, WordEmbedding],
sentence_embedding: dict[str, WordEmbedding],
embedding_size: int,
vocab_size: int,
) -> SparseEmbedding:
@@ -216,8 +213,8 @@ class SparseVectorConverter:
Same as `embedding_to_vector`, but no TF
"""
indices: List[int] = []
values: List[float] = []
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
+10 -11
View File
@@ -1,5 +1,4 @@
from typing import Optional, Sequence, Any, Iterable
from typing import Sequence, Any, Iterable
from dataclasses import dataclass
import numpy as np
@@ -11,7 +10,7 @@ from fastembed.common.model_description import (
DenseModelDescription,
)
from fastembed.common.onnx_model import OnnxOutputContext
from fastembed.common.types import NumpyArray
from fastembed.common.types import NumpyArray, Device
from fastembed.common.utils import normalize, mean_pooling
from fastembed.text.onnx_embedding import OnnxTextEmbedding
@@ -29,14 +28,14 @@ class CustomTextEmbedding(OnnxTextEmbedding):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
device_id: Optional[int] = None,
specific_model_path: Optional[str] = None,
device_id: int | None = None,
specific_model_path: str | None = None,
**kwargs: Any,
):
super().__init__(
@@ -64,7 +63,7 @@ class CustomTextEmbedding(OnnxTextEmbedding):
return self._normalize(self._pool(output.model_output, output.attention_mask))
def _pool(
self, embeddings: NumpyArray, attention_mask: Optional[NDArray[np.int64]] = None
self, embeddings: NumpyArray, attention_mask: NDArray[np.int64] | None = None
) -> NumpyArray:
if self._pooling == PoolingType.CLS:
return embeddings[:, 0]
+8 -10
View File
@@ -1,5 +1,5 @@
from enum import Enum
from typing import Any, Type, Iterable, Union, Optional
from typing import Any, Type, Iterable
import numpy as np
@@ -45,11 +45,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, task_id: int | None = None, **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.default_task_id: Task | int = task_id if task_id is not None else self.PASSAGE_TASK
@classmethod
def _get_worker_class(cls) -> Type[OnnxTextEmbeddingWorker]:
@@ -62,7 +60,7 @@ class JinaEmbeddingV3(PooledNormalizedEmbedding):
def _preprocess_onnx_input(
self,
onnx_input: dict[str, NumpyArray],
task_id: Optional[Union[int, Task]] = None,
task_id: int | Task | None = None,
**kwargs: Any,
) -> dict[str, NumpyArray]:
if task_id is None:
@@ -72,10 +70,10 @@ class JinaEmbeddingV3(PooledNormalizedEmbedding):
def embed(
self,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
task_id: Optional[int] = None,
parallel: int | None = None,
task_id: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
task_id = (
@@ -83,7 +81,7 @@ class JinaEmbeddingV3(PooledNormalizedEmbedding):
) # required for multiprocessing
yield from super().embed(documents, batch_size, parallel, task_id=task_id, **kwargs)
def query_embed(self, query: Union[str, Iterable[str]], **kwargs: Any) -> Iterable[NumpyArray]:
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:
yield from super().embed(query, task_id=self.QUERY_TASK, **kwargs)
def passage_embed(self, texts: Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:
+17 -16
View File
@@ -1,6 +1,6 @@
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Iterable, Sequence, Type
from fastembed.common.types import NumpyArray, OnnxProvider
from fastembed.common.types import NumpyArray, OnnxProvider, Device
from fastembed.common.onnx_model import OnnxOutputContext
from fastembed.common.utils import define_cache_dir, normalize
from fastembed.text.onnx_text_model import OnnxTextModel, TextEmbeddingWorker
@@ -199,14 +199,14 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[NumpyArray]):
def __init__(
self,
model_name: str = "BAAI/bge-small-en-v1.5",
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
device_id: Optional[int] = None,
specific_model_path: Optional[str] = None,
device_id: int | None = None,
specific_model_path: str | None = None,
**kwargs: Any,
):
"""
@@ -218,10 +218,11 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[NumpyArray]):
threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.
providers (Optional[Sequence[OnnxProvider]], optional): The list of onnxruntime providers to use.
Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to False.
cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to Device.AUTO.
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, mutually exclusive
with `providers`. Defaults to None.
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
@@ -239,7 +240,7 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[NumpyArray]):
self.cuda = cuda
# This device_id will be used if we need to load model in current process
self.device_id: Optional[int] = None
self.device_id: int | None = None
if device_id is not None:
self.device_id = device_id
elif self.device_ids is not None:
@@ -260,9 +261,9 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[NumpyArray]):
def embed(
self,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
"""
@@ -332,7 +333,7 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[NumpyArray]):
)
def token_count(
self, texts: Union[str, Iterable[str]], batch_size: int = 1024, **kwargs: Any
self, texts: str | Iterable[str], batch_size: int = 1024, **kwargs: Any
) -> int:
return self._token_count(texts, batch_size=batch_size, **kwargs)
+18 -20
View File
@@ -1,13 +1,13 @@
import os
from multiprocessing import get_all_start_methods
from pathlib import Path
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Iterable, Sequence, Type
import numpy as np
from numpy.typing import NDArray
from tokenizers import Encoding, Tokenizer
from fastembed.common.types import NumpyArray, OnnxProvider
from fastembed.common.types import NumpyArray, OnnxProvider, Device
from fastembed.common.onnx_model import EmbeddingWorker, OnnxModel, OnnxOutputContext, T
from fastembed.common.preprocessor_utils import load_tokenizer
from fastembed.common.utils import iter_batch
@@ -15,7 +15,7 @@ from fastembed.parallel_processor import ParallelWorkerPool
class OnnxTextModel(OnnxModel[T]):
ONNX_OUTPUT_NAMES: Optional[list[str]] = None
ONNX_OUTPUT_NAMES: list[str] | None = None
@classmethod
def _get_worker_class(cls) -> Type["TextEmbeddingWorker[T]"]:
@@ -35,12 +35,12 @@ class OnnxTextModel(OnnxModel[T]):
def __init__(self) -> None:
super().__init__()
self.tokenizer: Optional[Tokenizer] = None
self.tokenizer: Tokenizer | None = None
self.special_token_to_id: dict[str, int] = {}
def _preprocess_onnx_input(
self, onnx_input: dict[str, NumpyArray], **kwargs: Any
) -> dict[str, Union[NumpyArray, NDArray[np.int64]]]:
) -> dict[str, NumpyArray | NDArray[np.int64]]:
"""
Preprocess the onnx input.
"""
@@ -50,11 +50,11 @@ class OnnxTextModel(OnnxModel[T]):
self,
model_dir: Path,
model_file: str,
threads: Optional[int],
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_id: Optional[int] = None,
extra_session_options: Optional[dict[str, Any]] = None,
threads: int | None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_id: int | None = None,
extra_session_options: dict[str, Any] | None = None,
) -> None:
super()._load_onnx_model(
model_dir=model_dir,
@@ -104,15 +104,15 @@ class OnnxTextModel(OnnxModel[T]):
self,
model_name: str,
cache_dir: str,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
parallel: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
local_files_only: bool = False,
specific_model_path: Optional[str] = None,
extra_session_options: Optional[dict[str, Any]] = None,
specific_model_path: str | None = None,
extra_session_options: dict[str, Any] | None = None,
**kwargs: Any,
) -> Iterable[T]:
is_small = False
@@ -159,9 +159,7 @@ class OnnxTextModel(OnnxModel[T]):
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:
def _token_count(self, texts: str | Iterable[str], batch_size: int = 1024, **_: Any) -> int:
if not hasattr(self, "model") or self.model is None:
self.load_onnx_model() # loads the tokenizer as well
+13 -13
View File
@@ -1,8 +1,8 @@
import warnings
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Iterable, Sequence, Type
from dataclasses import asdict
from fastembed.common.types import NumpyArray, OnnxProvider
from fastembed.common.types import NumpyArray, OnnxProvider, Device
from fastembed.text.clip_embedding import CLIPOnnxEmbedding
from fastembed.text.custom_text_embedding import CustomTextEmbedding
from fastembed.text.pooled_normalized_embedding import PooledNormalizedEmbedding
@@ -51,7 +51,7 @@ class TextEmbedding(TextEmbeddingBase):
description: str = "",
license: str = "",
size_in_gb: float = 0.0,
additional_files: Optional[list[str]] = None,
additional_files: list[str] | None = None,
) -> None:
registered_models = cls._list_supported_models()
for registered_model in registered_models:
@@ -79,11 +79,11 @@ class TextEmbedding(TextEmbeddingBase):
def __init__(
self,
model_name: str = "BAAI/bge-small-en-v1.5",
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
**kwargs: Any,
):
@@ -149,7 +149,7 @@ class TextEmbedding(TextEmbeddingBase):
ValueError: If the model name is not found in the supported models.
"""
descriptions = cls._list_supported_models()
embedding_size: Optional[int] = None
embedding_size: int | None = None
for description in descriptions:
if description.model.lower() == model_name.lower():
embedding_size = description.dim
@@ -164,9 +164,9 @@ class TextEmbedding(TextEmbeddingBase):
def embed(
self,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
"""
@@ -186,7 +186,7 @@ class TextEmbedding(TextEmbeddingBase):
"""
yield from self.model.embed(documents, batch_size, parallel, **kwargs)
def query_embed(self, query: Union[str, Iterable[str]], **kwargs: Any) -> Iterable[NumpyArray]:
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:
"""
Embeds queries
@@ -214,7 +214,7 @@ class TextEmbedding(TextEmbeddingBase):
yield from self.model.passage_embed(texts, **kwargs)
def token_count(
self, texts: Union[str, Iterable[str]], batch_size: int = 1024, **kwargs: Any
self, texts: str | Iterable[str], batch_size: int = 1024, **kwargs: Any
) -> int:
"""Returns the number of tokens in the texts.
+8 -8
View File
@@ -1,4 +1,4 @@
from typing import Iterable, Optional, Union, Any
from typing import Iterable, Any
from fastembed.common.model_description import DenseModelDescription
from fastembed.common.types import NumpyArray
@@ -9,21 +9,21 @@ class TextEmbeddingBase(ModelManagement[DenseModelDescription]):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
cache_dir: str | None = None,
threads: int | None = None,
**kwargs: Any,
):
self.model_name = model_name
self.cache_dir = cache_dir
self.threads = threads
self._local_files_only = kwargs.pop("local_files_only", False)
self._embedding_size: Optional[int] = None
self._embedding_size: int | None = None
def embed(
self,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
raise NotImplementedError()
@@ -43,7 +43,7 @@ class TextEmbeddingBase(ModelManagement[DenseModelDescription]):
# This is model-specific, so that different models can have specialized implementations
yield from self.embed(texts, **kwargs)
def query_embed(self, query: Union[str, Iterable[str]], **kwargs: Any) -> Iterable[NumpyArray]:
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:
"""
Embeds queries
@@ -70,6 +70,6 @@ class TextEmbeddingBase(ModelManagement[DenseModelDescription]):
"""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:
def token_count(self, texts: str | Iterable[str], **kwargs: Any) -> int:
"""Returns the number of tokens in the texts."""
raise NotImplementedError("Subclasses must implement this method")
+1
View File
@@ -13,6 +13,7 @@ copyright: |
theme:
name: material
logo: assets/favicon.png
favicon: assets/favicon.png
custom_dir: docs/overrides
icon:
repo: fontawesome/brands/github
Generated
+923 -1528
View File
File diff suppressed because it is too large Load Diff
+18 -17
View File
@@ -1,6 +1,6 @@
[tool.poetry]
name = "fastembed"
version = "0.7.4"
name = "fastembed-gpu"
version = "0.8.0"
description = "Fast, light, accurate library built for retrieval embedding generation"
authors = ["Qdrant Team <info@qdrant.tech>", "NirantK <nirant.bits@gmail.com>"]
license = "Apache License"
@@ -11,19 +11,19 @@ repository = "https://github.com/qdrant/fastembed"
keywords = ["vector", "embedding", "neural", "search", "qdrant", "sentence-transformers"]
[tool.poetry.dependencies]
python = ">=3.9.0"
python = ">=3.10.0"
numpy = [
{ version = ">=1.21,<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.26", python = ">=3.12,<3.13" },
{ version = ">=2.1.0", python = ">=3.13,<3.14" },
{ version = ">=1.21,<2.3.0", python = "3.10" },
{ version = ">=1.21", python = "3.11" },
{ version = ">=1.26", python = "3.12" },
{ version = ">=2.1.0", python = "3.13" },
{ version = ">=2.3.0", python = ">=3.14" },
]
onnxruntime = [
{ version = ">=1.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,<3.13" },
onnxruntime-gpu = [
{ version = ">=1.17.0,!=1.20.0,<1.24", python = "3.10" },
{ version = ">=1.17.0,!=1.20.0,!=1.24.0,!=1.24.1", python = ">=3.11,<3.13" },
{ version = ">1.21.0,!=1.24.0,!=1.24.1", python = "3.13" },
{ version = ">=1.24.2", python = ">=3.14" },
]
tqdm = "^4.66"
requests = "^2.31"
@@ -31,9 +31,9 @@ tokenizers = ">=0.15,<1.0"
huggingface-hub = ">=0.20,<2.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" },
{ version = ">=10.3.0,<13.0", python = ">=3.10,<3.13" },
{ version = ">=11.0.0,<13.0", python = "3.13" },
{ version = ">=12.0.0,<13.0", python = ">=3.14" },
]
mmh3 = ">=4.1.0,<6.0.0"
py-rust-stemmers = "^0.1.0"
@@ -46,8 +46,9 @@ ruff = ">=0.3.1,<1.0"
notebook = ">=7.0.2"
pre-commit = "^3.6.2"
onnx = [
{ version = ">=1.15.0", python = "<3.13" },
{ version = ">=1.18.0", python = ">=3.13" },
{ version = ">=1.15.0", python = ">=3.10,<3.13" },
{ version = ">=1.18.0", python = "3.13" },
{ version = ">=1.20.0", python = ">=3.14" },
]
[tool.poetry.group.docs.dependencies]
+96 -53
View File
@@ -1,4 +1,5 @@
import os
from contextlib import contextmanager
import pytest
from PIL import Image
@@ -6,7 +7,7 @@ import numpy as np
from fastembed import LateInteractionMultimodalEmbedding
from tests.config import TEST_MISC_DIR
from tests.utils import delete_model_cache
# vectors are abridged and rounded for brevity
CANONICAL_IMAGE_VALUES = {
@@ -21,6 +22,17 @@ CANONICAL_IMAGE_VALUES = {
[-0.1299, -0.0691, 0.1097, 0.0728, 0.0123, 0.0519, 0.0122],
]
),
"Qdrant/colmodernvbert": np.array(
[
[0.11614, -0.15793, -0.11194, 0.0688, 0.08001, 0.10575, -0.07871],
[0.10094, -0.13301, -0.12069, 0.10932, 0.04645, 0.09884, 0.04048],
[0.13106, -0.18613, -0.13469, 0.10566, 0.03659, 0.07712, -0.03916],
[0.09754, -0.09596, -0.04839, 0.14991, 0.05692, 0.10569, -0.08349],
[0.02576, -0.15651, -0.09977, 0.09707, 0.13412, 0.09994, -0.09931],
[-0.06741, -0.1787, -0.19677, -0.07618, 0.13102, -0.02131, -0.02437],
[-0.02776, -0.10187, -0.13793, 0.03835, 0.04766, 0.04701, -0.15635],
]
),
}
CANONICAL_QUERY_VALUES = {
@@ -35,6 +47,17 @@ CANONICAL_QUERY_VALUES = {
[-0.0165, -0.0106, 0.1672, -0.0768, 0.0389, -0.0038, 0.1137],
]
),
"Qdrant/colmodernvbert": np.array(
[
[0.05, 0.06557, 0.04026, 0.14981, 0.1842, 0.0263, -0.18706],
[-0.05664, -0.14028, 0.00649, -0.02849, 0.09034, -0.01494, 0.10693],
[-0.10147, -0.00716, 0.09084, -0.08236, -0.01849, -0.00972, -0.00461],
[-0.1233, -0.10814, -0.02337, -0.00329, 0.05984, 0.09934, 0.09846],
[-0.07053, -0.13119, -0.06487, 0.01508, 0.07459, 0.07655, 0.14821],
[0.00526, -0.13842, -0.05837, -0.02721, 0.13009, 0.05076, 0.17962],
[0.00924, -0.14383, -0.03057, -0.03691, 0.11718, 0.037, 0.13344],
]
),
}
queries = ["hello world", "flag embedding"]
@@ -44,43 +67,69 @@ images = [
Image.open((TEST_MISC_DIR / "image.jpeg")),
]
_MODELS_TO_CACHE = ("Qdrant/colmodernvbert",)
MODELS_TO_CACHE = tuple(model_name.lower() for model_name in _MODELS_TO_CACHE)
def test_batch_embedding():
if os.getenv("CI"):
pytest.skip("Colpali is too large to test in CI")
@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] = LateInteractionMultimodalEmbedding(lowercase_model_name)
yield cache[lowercase_model_name]
if lowercase_model_name not in MODELS_TO_CACHE:
model_inst = cache.pop(lowercase_model_name)
if is_ci:
delete_model_cache(model_inst.model._model_dir)
del model_inst
yield get_model
if is_ci:
for _, model in cache.items():
delete_model_cache(model.model._model_dir)
cache.clear()
def test_batch_embedding(model_cache):
for model_name, expected_result in CANONICAL_IMAGE_VALUES.items():
print("evaluating", model_name)
model = LateInteractionMultimodalEmbedding(model_name=model_name)
result = list(model.embed_image(images, batch_size=2))
if model_name.lower() == "Qdrant/colpali-v1.3-fp16".lower() and os.getenv("CI"):
continue # colpali is too large for ci
for value in result:
print("evaluating", model_name)
with model_cache(model_name) as model:
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)
def test_single_embedding(model_cache):
for model_name, expected_result in CANONICAL_IMAGE_VALUES.items():
if model_name.lower() == "Qdrant/colpali-v1.3-fp16".lower() and os.getenv("CI"):
continue # colpali is too large for ci
print("evaluating", model_name)
with model_cache(model_name) as model:
result = next(iter(model.embed_image(images, batch_size=6)))
token_num, abridged_dim = expected_result.shape
assert np.allclose(value[:token_num, :abridged_dim], expected_result, atol=2e-3)
assert np.allclose(result[:token_num, :abridged_dim], expected_result, atol=2e-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)
def test_single_embedding_query():
if os.getenv("CI"):
pytest.skip("Colpali is too large to test in CI")
def test_single_embedding_query(model_cache):
for model_name, expected_result in CANONICAL_QUERY_VALUES.items():
if model_name.lower() == "Qdrant/colpali-v1.3-fp16".lower() and os.getenv("CI"):
continue # colpali is too large for ci
print("evaluating", model_name)
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)
with model_cache(model_name) as model:
result = next(iter(model.embed_text(queries)))
token_num, abridged_dim = expected_result.shape
assert np.allclose(result[:token_num, :abridged_dim], expected_result, atol=2e-3)
def test_get_embedding_size():
@@ -90,33 +139,27 @@ def test_get_embedding_size():
model_name = "Qdrant/ColPali-v1.3-fp16"
assert LateInteractionMultimodalEmbedding.get_embedding_size(model_name) == 128
model_name = "Qdrant/colmodernvbert"
assert LateInteractionMultimodalEmbedding.get_embedding_size(model_name) == 128
def test_embedding_size():
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_name = "Qdrant/colmodernvbert"
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
)
def test_token_count(model_cache) -> None:
model_name = "Qdrant/colmodernvbert"
with model_cache(model_name) as model:
documents = ["short doc", "it is a long document to check attention mask for paddings"]
short_doc_token_count = model.token_count(documents[0])
long_doc_token_count = model.token_count(documents[1])
documents_token_count = model.token_count(documents)
assert short_doc_token_count + long_doc_token_count == documents_token_count
assert short_doc_token_count + long_doc_token_count == model.token_count(
documents, batch_size=1
)
assert short_doc_token_count + long_doc_token_count < model.token_count(
documents, include_extension=True
)
+4 -4
View File
@@ -1,5 +1,5 @@
import pytest
from typing import Optional
from fastembed import (
TextEmbedding,
SparseTextEmbedding,
@@ -14,7 +14,7 @@ CACHE_DIR = "../model_cache"
@pytest.mark.skip(reason="Requires a multi-gpu server")
@pytest.mark.parametrize("device_id", [None, 0, 1])
def test_gpu_via_providers(device_id: Optional[int]) -> None:
def test_gpu_via_providers(device_id: int | None) -> None:
docs = ["hello world", "flag embedding"]
device_id = device_id if device_id is not None else 0
@@ -86,7 +86,7 @@ def test_gpu_via_providers(device_id: Optional[int]) -> None:
@pytest.mark.skip(reason="Requires a multi-gpu server")
@pytest.mark.parametrize("device_ids", [None, [0], [1], [0, 1]])
def test_gpu_cuda_device_ids(device_ids: Optional[list[int]]) -> None:
def test_gpu_cuda_device_ids(device_ids: list[int] | None) -> None:
docs = ["hello world", "flag embedding"]
device_id = device_ids[0] if device_ids else 0
embedding_model = TextEmbedding(
@@ -171,7 +171,7 @@ def test_gpu_cuda_device_ids(device_ids: Optional[list[int]]) -> None:
@pytest.mark.parametrize(
"device_ids,parallel", [(None, None), (None, 2), ([1], None), ([1], 1), ([1], 2), ([0, 1], 2)]
)
def test_multi_gpu_parallel_inference(device_ids: Optional[list[int]], parallel: int) -> None:
def test_multi_gpu_parallel_inference(device_ids: list[int] | None, parallel: int) -> None:
docs = ["hello world", "flag embedding"] * 100
batch_size = 5
+3 -3
View File
@@ -3,12 +3,12 @@ import traceback
from pathlib import Path
from types import TracebackType
from typing import Union, Callable, Any, Type, Optional
from typing import Callable, Any, Type
from fastembed.common.model_description import BaseModelDescription
def delete_model_cache(model_dir: Union[str, Path]) -> None:
def delete_model_cache(model_dir: str | Path) -> None:
"""Delete the model cache directory.
If a model was downloaded from the HuggingFace model hub, then _model_dir is the dir to snapshots, removing
@@ -42,7 +42,7 @@ def delete_model_cache(model_dir: Union[str, Path]) -> None:
def should_test_model(
model_desc: BaseModelDescription,
autotest_model_name: str,
is_ci: Optional[str],
is_ci: str | None,
is_manual: bool,
):
"""Determine if a model should be tested based on environment