Compare commits

...
79 Commits
Author SHA1 Message Date
Yufeng HeandGeorge Panchuk 0c63b6ab52 fix: block unsafe tar extraction paths (#647)
* fix: block unsafe tar extraction paths

* fix: correct the version gate and fallback in tar extraction

* fix: fix windows vulnerability

---------

Co-authored-by: George Panchuk <george.panchuk@qdrant.tech>
2026-09-22 15:59:52 +07:00
Serhii ZghamaandGeorge Panchuk cbe60bf9dc fix(image): normalize batched (N, C, H, W) input along the channel axis (#682)
* fix(image): normalize batched input along the channel axis

normalize() advertises 4D (N, C, H, W) support via its num_channels
branch and the channel-count validation, but the actual math used
((image.T - mean) / std).T. Transpose reverses every axis, so on 4D
input the channels no longer line up with mean/std: it raises when
N != C and silently normalizes along the batch axis when N == C.
Reshape mean/std to broadcast on the real channel axis instead; the
(C, H, W) path is unchanged.

* test(image): cover channel-wise normalize for 3D and batched input

* refactor

---------

Co-authored-by: George Panchuk <george.panchuk@qdrant.tech>
2026-09-22 15:59:21 +07:00
Baojiang LeeandGeorge Panchuk 40cca63d5f fix: make tokenizer metadata files optional (#693)
* fix: support optional tokenizer metadata files

* fix: pad id fallback chain and additional_special_tokens lists

* refactor: remove redundant tests and comments

---------

Co-authored-by: George Panchuk <george.panchuk@qdrant.tech>
2026-09-22 01:37:52 +07:00
5c4d9b04bd fix: pass (width, height) to Pillow in the Resize transform (#697)
* fix: pass (width, height) to Pillow in the Resize transform

`resize()` handed a tuple size straight to `PIL.Image.resize()`. fastembed
keeps sizes as (height, width) — `Transform.from_config` builds the tuple
as `(size["height"], size["width"])` — while Pillow takes (width, height),
so a non-square image processor configuration produced a transposed image:

    Resize(size=(100, 200))(Image.new("RGB", (300, 300)))[0].size
    # (100, 200), expected (200, 100)

Square sizes are unaffected, which is why this went unnoticed. The int
branch of `resize()` already emits Pillow order and is untouched, as are
`resize_ndarray()`'s callers, which pass (width, height) explicitly.

`Resize.__call__` is the only caller of this function and always supplies
fastembed's height-first order, so converting here is safe.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>

* tests: simplify tests

---------

Co-authored-by: George Panchuk <george.panchuk@qdrant.tech>
2026-09-21 23:07:15 +07:00
GeorgeandMohammed Alshyakh a3a798f4f3 fix: preserve pad_to_multiple_of when normalizing padding (#717)
Co-authored-by: Mohammed Alshyakh <zzzzmmmm298@gmail.com>
2026-09-21 22:13:43 +07:00
George bdf6816da8 fix: normalize tokenizer padding to batch-longest (#716)
* fix: normalize tokenizer padding to batch-longest

* fix: don't use max position embeddings as max len
2026-09-21 21:05:57 +07:00
dependabot[bot] 5dc53ebf49 chore(deps-dev): bump the security-updates group across 1 directory with 2 updates (#705)
Bumps the security-updates group with 2 updates in the / directory: [mkdocs-material](https://github.com/squidfunk/mkdocs-material) and [mistune](https://github.com/lepture/mistune).


Updates `mkdocs-material` from 9.7.4 to 9.7.7
- [Release notes](https://github.com/squidfunk/mkdocs-material/releases)
- [Changelog](https://github.com/squidfunk/mkdocs-material/blob/master/CHANGELOG)
- [Commits](https://github.com/squidfunk/mkdocs-material/compare/9.7.4...9.7.7)

Updates `mistune` from 3.3.0 to 3.3.3
- [Release notes](https://github.com/lepture/mistune/releases)
- [Changelog](https://github.com/lepture/mistune/blob/main/docs/changes.rst)
- [Commits](https://github.com/lepture/mistune/compare/v3.3.0...v3.3.3)

---
updated-dependencies:
- dependency-name: mistune
  dependency-version: 3.3.3
  dependency-type: indirect
  dependency-group: security-updates
- dependency-name: mkdocs-material
  dependency-version: 9.7.7
  dependency-type: direct:development
  dependency-group: security-updates
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-09-21 15:06:59 +07:00
GeorgeandS0rryHorizon b461314660 fix: fix registering custom models in child workers (#714)
* fix: fix registering custom models in child workers

Co-authored-by: S0rryHorizon <151612757+S0rryHorizon@users.noreply.github.com>

* fix: fix mypy

---------

Co-authored-by: S0rryHorizon <151612757+S0rryHorizon@users.noreply.github.com>
2026-09-21 15:05:47 +07:00
DarshandGeorge Panchuk cc4d101828 fix case insensitive lookup for custom text models (#645)
* fix case insensitive lookup for custom text models

* refactor: refactor a bit

---------

Co-authored-by: George Panchuk <george.panchuk@qdrant.tech>
2026-09-16 17:46:21 +07:00
0dab99c23e fix: use the canonical Hugging Face source for BGE-small (#707)
* fix: use the canonical Hugging Face source for BGE-small

* tests: remove excess test

---------

Co-authored-by: Basil Chen <192173459+rastagan-git@users.noreply.github.com>
Co-authored-by: George <george.panchuk@qdrant.tech>
2026-09-09 17:16:00 +07:00
Harnas 113bd565ec Correct Qdrant/bge-base-en-v1.5-onnx-Q model name (#593)
Wrong name causes HTTP redirection, what may be wrongly handled by proxies.
2026-09-09 17:13:43 +07:00
Stephan TulkensandDylan Couzon d5f552b6ad new: add minish models (#692)
* new: add minish models

* add canonical values to test

* amend description

* Apply suggestions from code review

Co-authored-by: Dylan Couzon <dylancouzon@gmail.com>

* Update fastembed/text/onnx_embedding.py

Co-authored-by: Dylan Couzon <dylancouzon@gmail.com>

* fix typo

* fix keys in tests

---------

Co-authored-by: Dylan Couzon <dylancouzon@gmail.com>
2026-09-07 17:40:07 +07:00
dependabot[bot] 70c8f3cbc0 chore(deps-dev): bump tornado (#698)
Bumps the security-updates group with 1 update in the / directory: [tornado](https://github.com/tornadoweb/tornado).


Updates `tornado` from 6.5.7 to 6.5.8
- [Changelog](https://github.com/tornadoweb/tornado/blob/master/docs/releases.rst)
- [Commits](https://github.com/tornadoweb/tornado/compare/v6.5.7...v6.5.8)

---
updated-dependencies:
- dependency-name: tornado
  dependency-version: 6.5.8
  dependency-type: indirect
  dependency-group: security-updates
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-09-07 16:34:37 +07:00
Bastian Hofmann a34e7bcc42 Set copyright holder in LICENSE files (#696)
Replace the Apache 2.0 placeholder with Qdrant Solutions GmbH and the
year 2026.
2026-09-01 14:54:52 +02:00
George c48247f15d new: add siglip (#683) 2026-08-19 23:02:56 +07:00
George f9d757ffc6 fix: fix qwen (#680) 2026-08-18 23:47:02 +07:00
George a5a702acad new: add qwen embedding (#678) 2026-08-18 19:26:16 +07:00
dependabot[bot] a0fe741532 chore(deps-dev): bump mypy from 1.19.1 to 2.3.0 (#663)
Bumps [mypy](https://github.com/python/mypy) from 1.19.1 to 2.3.0.
- [Changelog](https://github.com/python/mypy/blob/master/CHANGELOG.md)
- [Commits](https://github.com/python/mypy/compare/v1.19.1...v2.3.0)

---
updated-dependencies:
- dependency-name: mypy
  dependency-version: 2.3.0
  dependency-type: direct:development
  update-type: version-update:semver-major
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-08-12 20:53:31 +07:00
dependabot[bot] 525642bd74 chore(deps): bump actions/setup-python from 6.2.0 to 7.0.0 (#658)
Bumps [actions/setup-python](https://github.com/actions/setup-python) from 6.2.0 to 7.0.0.
- [Release notes](https://github.com/actions/setup-python/releases)
- [Commits](https://github.com/actions/setup-python/compare/a309ff8b426b58ec0e2a45f0f869d46889d02405...5fda3b95a4ea91299a34e894583c3862153e4b97)

---
updated-dependencies:
- dependency-name: actions/setup-python
  dependency-version: 7.0.0
  dependency-type: direct:production
  update-type: version-update:semver-major
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-08-12 20:39:23 +07:00
dependabot[bot] 39426c282e chore(deps): bump pypa/gh-action-pypi-publish (#657)
Bumps the version-updates group with 1 update in the / directory: [pypa/gh-action-pypi-publish](https://github.com/pypa/gh-action-pypi-publish).


Updates `pypa/gh-action-pypi-publish` from 1.14.0 to 1.14.2
- [Release notes](https://github.com/pypa/gh-action-pypi-publish/releases)
- [Commits](https://github.com/pypa/gh-action-pypi-publish/compare/cef221092ed1bacb1cc03d23a2d87d1d172e277b...dc37677b2e1c63e2034f94d8a5b11f265b73ba33)

---
updated-dependencies:
- dependency-name: pypa/gh-action-pypi-publish
  dependency-version: 1.14.1
  dependency-type: direct:production
  update-type: version-update:semver-patch
  dependency-group: version-updates
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-08-12 20:39:10 +07:00
dependabot[bot] eff93e3cb3 chore(deps-dev): bump pytest from 7.4.4 to 9.1.1 (#664)
Bumps [pytest](https://github.com/pytest-dev/pytest) from 7.4.4 to 9.1.1.
- [Release notes](https://github.com/pytest-dev/pytest/releases)
- [Changelog](https://github.com/pytest-dev/pytest/blob/main/CHANGELOG.rst)
- [Commits](https://github.com/pytest-dev/pytest/compare/7.4.4...9.1.1)

---
updated-dependencies:
- dependency-name: pytest
  dependency-version: 9.1.1
  dependency-type: direct:development
  update-type: version-update:semver-major
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-08-12 20:36:05 +07:00
dependabot[bot] 5cbf4f8947 chore(deps): bump actions/cache from 5.0.5 to 6.1.0 (#659)
Bumps [actions/cache](https://github.com/actions/cache) from 5.0.5 to 6.1.0.
- [Release notes](https://github.com/actions/cache/releases)
- [Changelog](https://github.com/actions/cache/blob/main/RELEASES.md)
- [Commits](https://github.com/actions/cache/compare/27d5ce7f107fe9357f9df03efb73ab90386fccae...55cc8345863c7cc4c66a329aec7e433d2d1c52a9)

---
updated-dependencies:
- dependency-name: actions/cache
  dependency-version: 6.1.0
  dependency-type: direct:production
  update-type: version-update:semver-major
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-08-12 20:34:47 +07:00
dependabot[bot] 817fca3fab chore(deps): bump actions/checkout from 6.0.2 to 7.0.1 (#660)
Bumps [actions/checkout](https://github.com/actions/checkout) from 6.0.2 to 7.0.1.
- [Release notes](https://github.com/actions/checkout/releases)
- [Changelog](https://github.com/actions/checkout/blob/main/CHANGELOG.md)
- [Commits](https://github.com/actions/checkout/compare/de0fac2e4500dabe0009e67214ff5f5447ce83dd...3d3c42e5aac5ba805825da76410c181273ba90b1)

---
updated-dependencies:
- dependency-name: actions/checkout
  dependency-version: 7.0.1
  dependency-type: direct:production
  update-type: version-update:semver-major
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-08-12 20:34:31 +07:00
dependabot[bot] e60d5493bb chore(deps-dev): bump mkdocstrings from 0.24.3 to 1.0.6 (#665)
Bumps [mkdocstrings](https://github.com/mkdocstrings/mkdocstrings) from 0.24.3 to 1.0.6.
- [Release notes](https://github.com/mkdocstrings/mkdocstrings/releases)
- [Changelog](https://github.com/mkdocstrings/mkdocstrings/blob/main/CHANGELOG.md)
- [Commits](https://github.com/mkdocstrings/mkdocstrings/compare/0.24.3...1.0.6)

---
updated-dependencies:
- dependency-name: mkdocstrings
  dependency-version: 1.0.6
  dependency-type: direct:development
  update-type: version-update:semver-major
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-08-12 20:33:11 +07:00
dependabot[bot] b2494cb321 chore(deps): bump the security-updates group across 1 directory with 15 updates (#672)
Bumps the security-updates group with 14 updates in the / directory:

| Package | From | To |
| --- | --- | --- |
| [pillow](https://github.com/python-pillow/Pillow) | `12.1.1` | `12.3.0` |
| [notebook](https://github.com/jupyter/notebook) | `7.5.4` | `7.5.6` |
| [onnx](https://github.com/onnx/onnx) | `1.20.1` | `1.22.0` |
| [bleach](https://github.com/mozilla/bleach) | `6.3.0` | `6.4.0` |
| [gitpython](https://github.com/gitpython-developers/GitPython) | `3.1.46` | `3.1.58` |
| [idna](https://github.com/kjd/idna) | `3.11` | `3.15` |
| [jupyter-server](https://github.com/jupyter-server/jupyter_server) | `2.17.0` | `2.20.0` |
| [mistune](https://github.com/lepture/mistune) | `3.2.0` | `3.3.0` |
| [nbconvert](https://github.com/jupyter/nbconvert) | `7.17.0` | `7.17.1` |
| [pymdown-extensions](https://github.com/facelessuser/pymdown-extensions) | `10.21` | `11.0.1` |
| [setuptools](https://github.com/pypa/setuptools) | `82.0.0` | `83.0.0` |
| [soupsieve](https://github.com/facelessuser/soupsieve) | `2.8.3` | `2.8.4` |
| [tornado](https://github.com/tornadoweb/tornado) | `6.5.4` | `6.5.7` |
| [urllib3](https://github.com/urllib3/urllib3) | `2.6.3` | `2.7.0` |



Updates `pillow` from 12.1.1 to 12.3.0
- [Release notes](https://github.com/python-pillow/Pillow/releases)
- [Changelog](https://github.com/python-pillow/Pillow/blob/main/CHANGES.rst)
- [Commits](https://github.com/python-pillow/Pillow/compare/12.1.1...12.3.0)

Updates `notebook` from 7.5.4 to 7.5.6
- [Release notes](https://github.com/jupyter/notebook/releases)
- [Changelog](https://github.com/jupyter/notebook/blob/@jupyter-notebook/tree@7.5.6/CHANGELOG.md)
- [Commits](https://github.com/jupyter/notebook/compare/@jupyter-notebook/tree@7.5.4...@jupyter-notebook/tree@7.5.6)

Updates `onnx` from 1.20.1 to 1.22.0
- [Release notes](https://github.com/onnx/onnx/releases)
- [Changelog](https://github.com/onnx/onnx/blob/main/docs/Changelog-ml.md)
- [Commits](https://github.com/onnx/onnx/compare/v1.20.1...v1.22.0)

Updates `bleach` from 6.3.0 to 6.4.0
- [Changelog](https://github.com/mozilla/bleach/blob/main/CHANGES)
- [Commits](https://github.com/mozilla/bleach/compare/v6.3.0...v6.4.0)

Updates `gitpython` from 3.1.46 to 3.1.58
- [Release notes](https://github.com/gitpython-developers/GitPython/releases)
- [Changelog](https://github.com/gitpython-developers/GitPython/blob/main/CHANGES)
- [Commits](https://github.com/gitpython-developers/GitPython/compare/3.1.46...3.1.58)

Updates `idna` from 3.11 to 3.15
- [Release notes](https://github.com/kjd/idna/releases)
- [Changelog](https://github.com/kjd/idna/blob/master/HISTORY.md)
- [Commits](https://github.com/kjd/idna/compare/v3.11...v3.15)

Updates `jupyter-server` from 2.17.0 to 2.20.0
- [Release notes](https://github.com/jupyter-server/jupyter_server/releases)
- [Changelog](https://github.com/jupyter-server/jupyter_server/blob/main/CHANGELOG.md)
- [Commits](https://github.com/jupyter-server/jupyter_server/compare/v2.17.0...v2.20.0)

Updates `jupyterlab` from 4.5.5 to 4.5.10
- [Release notes](https://github.com/jupyterlab/jupyterlab/releases)
- [Changelog](https://github.com/jupyterlab/jupyterlab/blob/main/RELEASE.md)
- [Commits](https://github.com/jupyterlab/jupyterlab/compare/@jupyterlab/lsp@4.5.5...@jupyterlab/lsp@4.5.10)

Updates `mistune` from 3.2.0 to 3.3.0
- [Release notes](https://github.com/lepture/mistune/releases)
- [Changelog](https://github.com/lepture/mistune/blob/main/docs/changes.rst)
- [Commits](https://github.com/lepture/mistune/compare/v3.2.0...v3.3.0)

Updates `nbconvert` from 7.17.0 to 7.17.1
- [Release notes](https://github.com/jupyter/nbconvert/releases)
- [Changelog](https://github.com/jupyter/nbconvert/blob/main/CHANGELOG.md)
- [Commits](https://github.com/jupyter/nbconvert/compare/v7.17.0...v7.17.1)

Updates `pymdown-extensions` from 10.21 to 11.0.1
- [Release notes](https://github.com/facelessuser/pymdown-extensions/releases)
- [Commits](https://github.com/facelessuser/pymdown-extensions/compare/10.21...11.0.1)

Updates `setuptools` from 82.0.0 to 83.0.0
- [Release notes](https://github.com/pypa/setuptools/releases)
- [Changelog](https://github.com/pypa/setuptools/blob/main/NEWS.rst)
- [Commits](https://github.com/pypa/setuptools/compare/v82.0.0...v83.0.0)

Updates `soupsieve` from 2.8.3 to 2.8.4
- [Release notes](https://github.com/facelessuser/soupsieve/releases)
- [Commits](https://github.com/facelessuser/soupsieve/compare/2.8.3...2.8.4)

Updates `tornado` from 6.5.4 to 6.5.7
- [Changelog](https://github.com/tornadoweb/tornado/blob/master/docs/releases.rst)
- [Commits](https://github.com/tornadoweb/tornado/compare/v6.5.4...v6.5.7)

Updates `urllib3` from 2.6.3 to 2.7.0
- [Release notes](https://github.com/urllib3/urllib3/releases)
- [Changelog](https://github.com/urllib3/urllib3/blob/main/CHANGES.rst)
- [Commits](https://github.com/urllib3/urllib3/compare/2.6.3...2.7.0)

---
updated-dependencies:
- dependency-name: pillow
  dependency-version: 12.3.0
  dependency-type: direct:production
  dependency-group: security-updates
- dependency-name: notebook
  dependency-version: 7.5.6
  dependency-type: direct:development
  dependency-group: security-updates
- dependency-name: onnx
  dependency-version: 1.22.0
  dependency-type: direct:development
  dependency-group: security-updates
- dependency-name: bleach
  dependency-version: 6.4.0
  dependency-type: indirect
  dependency-group: security-updates
- dependency-name: gitpython
  dependency-version: 3.1.58
  dependency-type: indirect
  dependency-group: security-updates
- dependency-name: idna
  dependency-version: '3.15'
  dependency-type: indirect
  dependency-group: security-updates
- dependency-name: jupyter-server
  dependency-version: 2.20.0
  dependency-type: indirect
  dependency-group: security-updates
- dependency-name: jupyterlab
  dependency-version: 4.5.10
  dependency-type: indirect
  dependency-group: security-updates
- dependency-name: mistune
  dependency-version: 3.3.0
  dependency-type: indirect
  dependency-group: security-updates
- dependency-name: nbconvert
  dependency-version: 7.17.1
  dependency-type: indirect
  dependency-group: security-updates
- dependency-name: pymdown-extensions
  dependency-version: 11.0.1
  dependency-type: indirect
  dependency-group: security-updates
- dependency-name: setuptools
  dependency-version: 83.0.0
  dependency-type: indirect
  dependency-group: security-updates
- dependency-name: soupsieve
  dependency-version: 2.8.4
  dependency-type: indirect
  dependency-group: security-updates
- dependency-name: tornado
  dependency-version: 6.5.7
  dependency-type: indirect
  dependency-group: security-updates
- dependency-name: urllib3
  dependency-version: 2.7.0
  dependency-type: indirect
  dependency-group: security-updates
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-08-12 20:32:35 +07:00
dependabot[bot] f613647297 chore(deps-dev): bump pre-commit from 3.8.0 to 4.6.1 (#662)
Bumps [pre-commit](https://github.com/pre-commit/pre-commit) from 3.8.0 to 4.6.1.
- [Release notes](https://github.com/pre-commit/pre-commit/releases)
- [Changelog](https://github.com/pre-commit/pre-commit/blob/main/CHANGELOG.md)
- [Commits](https://github.com/pre-commit/pre-commit/compare/v3.8.0...v4.6.1)

---
updated-dependencies:
- dependency-name: pre-commit
  dependency-version: 4.6.1
  dependency-type: direct:development
  update-type: version-update:semver-major
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-08-05 00:24:23 +07:00
andres-qd 50ae2dc088 chore: add Dependabot configuration (#644)
* chore: add Dependabot configuration

* chore: add dependabot cooldown configuration
2026-08-04 23:36:37 +07:00
George 0892291f75 new: add inference free splade (#652) 2026-07-22 22:32:30 +07:00
Dylan Couzon 8108e4467c add nomic-embed-vision-v1.5 (#651) 2026-07-20 17:55:48 -04:00
Esteban Yusunguaira 8a8ea4f42f ci: update all actions to node 24 (#636) 2026-05-25 09:52:43 -05:00
Esteban Yusunguaira a499c313af ci: fix remaining actions on node 20 (#634) 2026-05-22 11:30:12 +07:00
Esteban Yusunguaira fde1e0b361 ci: update github actions to node 24 (#633) 2026-05-20 16:33:49 +07:00
George adfc1aef59 fix: add timeout for download from gcs (#629) 2026-04-21 15:32:52 +07:00
George 1ec283c744 fix: fix license (#625) 2026-04-15 14:59:45 +07:00
George a6a4e375ca fix: use original jina de model instead of fp16 due to onnxruntime up… (#623)
* fix: use original jina de model instead of fp16 due to onnxruntime updates

* fix: update size in description
2026-04-15 12:53:10 +07:00
George 2b069d2fbd fix: check if model files exist before returning cached path (#624) 2026-04-10 14:11:02 +07:00
estebany-qd 21df54e3a4 ci: Pin all gh actions to commit SHAs (#617) 2026-03-30 20:31:46 +02:00
George 87678dd784 new: add gemmaembedding-300m (#592)
* new: add gemma embed

* refactor: rename builtin pooling normalized embedding to builtin sentence embedding
2026-03-25 15:22:16 +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
George 428381cb04 bump version to v0.7.4 2025-12-05 18:38:45 +07:00
George 3511b08831 fix: numpy dropped 3.10 as of 2.3.0 (#585)
* fix: numpy dropped 3.10 as of 2.3.0

* fix: update poetry lock
2025-12-05 18:08:20 +07:00
George b718cc6a88 new: add token count method (#583)
* new: add token count method

* fix: fix mypy

* fix: load model in token_count

* fix: remove debug code
2025-12-04 11:29:36 +07:00
George 2ba8990260 new: try unlocking huggingface hub and pillow (#582)
* new: try unlocking huggingface hub and pillow

* new: adjust pillow version for 3.9

* fix: fix pillow for python3.13
2025-12-01 18:21:28 +07:00
George dab185fd9d fix: fix onnx version for python3.13 (#580) 2025-11-25 21:42:10 +07:00
George ec0e3128ee new: expose some onnx session options (#578)
* new: expose some onnx session options

* fix: fix extra session options is None case

* fix: fix missing params

* new: add tests
2025-11-25 17:49:02 +07:00
George 44e332999c new: try loading models from cache before making any network calls (#577) 2025-11-25 12:07:50 +07:00
George 533b54cee5 tests: introduce model cache to tests (#573)
* tests: introduce model cache to tests

* fix: fix not cached model deletion

* new: do not run CI tests on mac os and windows on python 3.10-3.12

* fix: lowercase cache keys, bm25 caching

* tests: do not run parallel processing on all cpus in sparse text embed

* fix: fix models to cache names, do not run parallel=0

* fix: fix sparse embedding tests

* fix: bm42 language by lower case model name
2025-11-12 18:54:33 +07:00
George Panchuk ba1f6053bd bump version to v0.7.3 2025-08-29 14:15:25 +03:00
George 4dc76e3859 fix: fix colbert query postprocessing (#557)
* fix: fix colbert query postprocessing

* fix: improve colbert single embedding tests
2025-08-29 13:25:43 +03:00
George 6efe06b172 new: decouple colbert query and document tokenizer (#556) 2025-08-29 13:25:18 +03:00
George Panchuk 887239239b bump version to v0.7.2 2025-08-25 16:29:23 +03:00
Kacper ŁukawskiandGeorge ca023be0c0 feat: MUVERA embeddings (#542)
* Implement MuveraEmbedding

* Add random generator parameter for reproducibility in MuveraEmbedding

* Document random_seed parameter

* Remove unnecessary module docstring from muvera_embedding.py

* refactor: clean up constructor parameters and improve formatting in MuveraEmbedding

* refactor: rename muvera_embedding.py to muvera.py and update related references

* feat: enhance MuveraEmbedding with multi-vector model support and improve parameter defaults

* feat: add embedding_size property to MuveraEmbedding

* feat: update MuveraPostprocessor to use model description for embedding size and add Jupyter notebook for MUVERA usage

* fix: fix types, doctest, rename variables, refactor (#545)

* fix: fix types, doctest, rename variables, refactor

* fix: fix python3.9 compatibility

* fix: make get_output_dimension protected

* Optimize muvera (#551)

* vectorize operations

* fix: fill empty clusters with dataset vectors

* rollback get_output_dimension

* fix: fix type hints

* fix: review comments

* tests: add tests

---------

Co-authored-by: George <george.panchuk@qdrant.tech>
2025-08-21 02:59:02 +03:00
Anush d1ddc8142d docs: Updated README.md (#550) 2025-08-11 13:29:45 +05:30
Andrey Vasnetsov faf3d9fe18 batch inference should return same shape as individual inference (#547) 2025-08-06 22:44:28 +02:00
George acec31277b fix: fix mypy colpali (#533) 2025-06-16 11:55:29 +03:00
George Panchuk cb902149c8 bump version to 0.7.1 2025-06-16 12:23:35 +04:00
George a260022ae6 fix: propagate local files only and specific model path into embed parallel (#524) 2025-05-20 14:46:42 +04:00
George d5da56299a new: add embedding size property (#521)
* new: add embedding size property

* fix: format exception message

* new: replace embedding size property with get_embedding_size classmethod

* fix: fix missed parts

* chore: fix docstrings

* new: add embedding_size property
2025-05-20 14:46:32 +04:00
George 4e5575f4a7 fix: raise exception if pooling is incorrect in custom model (#522) 2025-05-19 13:49:22 +03:00
George 4736f46548 fix: check lowercase model name for warnings (#523) 2025-05-19 13:49:12 +03:00
Andrey Vasnetsov c85e8c278f remove unused function (#516) 2025-05-13 17:04:51 +03:00
George 0df605fc92 fix: remove onnxruntime cap as not needed anymore (#517)
* fix: remove onnxruntime cap as not needed anymore

* fix: fix mypy complaints in python3.13
2025-05-13 17:54:54 +04:00
Andrey VasnetsovandGeorge 04bc7a3039 MiniCOIL v1 (#513)
* add token embeddings

* fix parallel worker init

* implement minicoil

* fix mypy

* fix mypy

* register minicoil

* rollback suggested "fix"

* add minicoil test

* some pr issues (#514)

* some pr issues

* revert query embed refactor

* test: add query embed tests

* nit

* Update tests/test_sparse_embeddings.py

---------

Co-authored-by: Andrey Vasnetsov <andrey@vasnetsov.com>

* review

* fix: revert change to colbert query_embed

---------

Co-authored-by: George <george.panchuk@qdrant.tech>
2025-05-13 14:22:50 +02:00
George b785640bd5 fix: fix list of supported bm25 languages (#506) 2025-04-15 14:16:09 +03:00
Hossam Hagag 5568a62c2f ci: Unlock numpy in ci (#504) 2025-04-11 13:00:54 +03:00
George Panchuk b34209dcfb bump version to 0.6.1 2025-04-10 16:23:54 +03:00
George c91d42dda7 Update setting jina v3 tasks (#503)
* new: improve task setter in jina v3

* refactor

* new: add hf_token secret

* fix: cross platform env propagation
2025-04-10 16:19:41 +03:00
George aa0c475a1f fix: fix splade name (#499) 2025-03-17 11:07:14 +03:00
Dmitrii OgnandGeorge Panchuk 4c239b11d5 Custom rerankers support (#496)
* Custom rerankers support

* Test for reranker_custom_model

* test fix

* Model description type fix

* Test fix

* fix: fix naming

* fix: remove redundant arg from tests

* new: update readme

---------

Co-authored-by: George Panchuk <george.panchuk@qdrant.tech>
2025-03-16 20:02:58 +03:00
Hossam HagagandGeorge Panchuk 6acfb001fb Speedup ci (#489)
* chore: Trigger CI test

* chore: Trigger CI test

* chore: Trigger CI test

* chore: Trigger CI test

* chore: Trigger CI test

* chore: Trigger CI test

* chore: Trigger CI test

* chore: Trigger CI test

* chore: Trigger CI test

* Trigger CI

* Trigger CI

* Trigger CI

* Trigger CI test

* Trigger CI test

* Trigger CI test

* Trigger CI test

* Trigger CI test

* Trigger CI test

* Trigger CI test

* Trigger CI test

* Trigger CI test

* Trigger CI test

* Trigger CI test

* Trigger CI test

* Trigger CI test

* Trigger CI test

* Trigger CI test

* Trigger CI test

* Trigger CI test

* Trigger CI test

* Trigger CI test

* Trigger CI test

* Trigger CI test

* Trigger CI test

* Trigger CI test

* Trigger CI test

* Trigger CI test

* Trigger CI test

* new: Added on workflow dispatch

* tests: Updated tests

* fix: Fix CI

* fix: Fix CI

* fix: Fix CI

* improve: Prevent stop iteration error caused by next

* fix: Fix variable might be referenced before assignment

* refactor: Revised the way of getting models to test

* fix: Fix test in image model

* refactor: Call one model

* fix: Fix ci

* fix: Fix splade model name

* tests: Updated tests

* chore: Remove cache

* tests: Update multi task tests

* tests: Update multi task tests

* tests: Updated tests

* refactor: refactor utils func, add comments, conditions refactor

---------

Co-authored-by: George Panchuk <george.panchuk@qdrant.tech>
2025-03-06 04:39:58 +02:00
George 1729aab1ec new: preserve embeddings in a type set by their model (#492)
* new: preserve embeddings in a type set by their model

* fix: remove type coercion

* fix: remove redundant type

* fix: fix random data type in tests
2025-03-03 17:46:35 +01:00
George 42fca3b467 Deprecate archive struct (#490)
* bump version to 0.6.0

* new: deprecate gcp archive structure
2025-03-03 17:12:47 +01:00
83 changed files with 11053 additions and 1161 deletions
+1 -1
View File
@@ -39,7 +39,7 @@ body:
attributes:
label: FastEmbed version
description: What version of FastEmbed are you running? python -c "import fastembed; print(fastembed.__version__)". If you're not on the latest, please upgrade and see if the problem persists.
placeholder: v0.5.1
placeholder: v0.7.4
validations:
required: true
- type: dropdown
+40
View File
@@ -0,0 +1,40 @@
version: 2
updates:
- package-ecosystem: "pip"
directory: "/"
schedule:
interval: "weekly"
open-pull-requests-limit: 5
groups:
security-updates:
applies-to: security-updates
patterns:
- "*"
version-updates:
applies-to: version-updates
update-types:
- "minor"
- "patch"
patterns:
- "*"
cooldown:
default-days: 7
- package-ecosystem: "github-actions"
directory: "/"
schedule:
interval: "weekly"
open-pull-requests-limit: 5
groups:
security-updates:
applies-to: security-updates
patterns:
- "*"
version-updates:
applies-to: version-updates
update-types:
- "minor"
- "patch"
patterns:
- "*"
cooldown:
default-days: 7
+3 -3
View File
@@ -10,12 +10,12 @@ jobs:
deploy:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v3
- uses: actions/setup-python@v4
- uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
- uses: actions/setup-python@5fda3b95a4ea91299a34e894583c3862153e4b97 # v7.0.0
with:
python-version: 3.x
- run: echo "cache_id=$(date --utc '+%V')" >> $GITHUB_ENV
- uses: actions/cache@v3
- uses: actions/cache@55cc8345863c7cc4c66a329aec7e433d2d1c52a9 # v6.1.0
with:
key: mkdocs-material-${{ env.cache_id }}
path: .cache
+4 -4
View File
@@ -21,11 +21,11 @@ jobs:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v2
- uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
- name: Set up Python
uses: actions/setup-python@v2
uses: actions/setup-python@5fda3b95a4ea91299a34e894583c3862153e4b97 # v7.0.0
with:
python-version: '3.9.x'
python-version: '3.10.x'
- name: Install dependencies
run: |
python -m pip install poetry
@@ -33,7 +33,7 @@ jobs:
- name: Build package
run: poetry build
- name: Publish package
uses: pypa/gh-action-pypi-publish@27b31702a0e7fc50959f5ad993c78deac1bdfc29
uses: pypa/gh-action-pypi-publish@dc37677b2e1c63e2034f94d8a5b11f265b73ba33 # v1.14.2
with:
user: __token__
password: ${{ secrets.PYPI_API_TOKEN }}
+22 -6
View File
@@ -1,9 +1,10 @@
name: Tests
on:
push:
branches: [ master, main, gpu ]
pull_request:
branches: [ master, main, gpu ]
workflow_dispatch:
env:
CARGO_TERM_COLOR: always
@@ -14,7 +15,6 @@ jobs:
strategy:
matrix:
python-version:
- '3.9.x'
- '3.10.x'
- '3.11.x'
- '3.12.x'
@@ -23,15 +23,29 @@ jobs:
- ubuntu-latest
- macos-latest
- windows-latest
exclude:
# Exclude 3.103.12 for macOS and Windows
- os: macos-latest
python-version: '3.10.x'
- os: macos-latest
python-version: '3.11.x'
- os: macos-latest
python-version: '3.12.x'
- os: windows-latest
python-version: '3.10.x'
- os: windows-latest
python-version: '3.11.x'
- os: windows-latest
python-version: '3.12.x'
runs-on: ${{ matrix.os }}
name: Python ${{ matrix.python-version }} on ${{ matrix.os }} test
steps:
- uses: actions/checkout@v3
- uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
- name: Set up Python
uses: actions/setup-python@v5
uses: actions/setup-python@5fda3b95a4ea91299a34e894583c3862153e4b97 # v7.0.0
with:
python-version: ${{ matrix.python-version }}
- name: Install dependencies
@@ -41,5 +55,7 @@ jobs:
poetry install --no-interaction --no-ansi --without dev,docs
- name: Run pytest
env:
HF_TOKEN: ${{ secrets.HF_TOKEN }}
run: |
poetry run pytest
poetry run pytest
+3 -5
View File
@@ -8,16 +8,16 @@ jobs:
strategy:
fail-fast: true
matrix:
python-version: ["3.9", "3.10", "3.11", "3.12", "3.13"]
python-version: ["3.10", "3.11", "3.12", "3.13"]
os: [ubuntu-latest]
name: Python ${{ matrix.python-version }} test
steps:
- uses: actions/checkout@v1
- uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
- name: Set up Python ${{ matrix.python-version }}
uses: actions/setup-python@v2
uses: actions/setup-python@5fda3b95a4ea91299a34e894583c3862153e4b97 # v7.0.0
with:
python-version: ${{ matrix.python-version }}
@@ -26,8 +26,6 @@ jobs:
python -m pip install --upgrade pip poetry
poetry install --no-interaction --no-ansi --without dev,docs,test
poetry run pip install "numpy<2.0.0" # https://github.com/python/mypy/issues/17396
- name: mypy
run: |
poetry run mypy fastembed \
+1 -1
View File
@@ -186,7 +186,7 @@
same "printed page" as the copyright notice for easier
identification within third-party archives.
Copyright [yyyy] [name of copyright owner]
Copyright 2026 Qdrant Solutions GmbH
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
+2
View File
@@ -15,6 +15,8 @@ These models are developed by Jina (https://jina.ai/) and are subject to Jina AI
This distribution includes the following Google models, each with its respective license:
- vidore/colpali-v1.3
- License: gemma
- google/embeddinggemma-300m
- License: gemma
Gemma is provided under and subject to the Gemma Terms of Use found at ai.google.dev/gemma/terms
+37 -20
View File
@@ -190,6 +190,23 @@ scores = list(encoder.rerank(query, documents))
# [-11.48061752319336, 5.472434997558594]
```
Text cross encoders can also be extended with models which are not in the list of supported models.
```python
from fastembed.rerank.cross_encoder import TextCrossEncoder
from fastembed.common.model_description import ModelSource
TextCrossEncoder.add_custom_model(
model="Xenova/ms-marco-MiniLM-L-4-v2",
model_file="onnx/model.onnx",
sources=ModelSource(hf="Xenova/ms-marco-MiniLM-L-4-v2"),
)
model = TextCrossEncoder(model_name="Xenova/ms-marco-MiniLM-L-4-v2")
scores = list(model.rerank_pairs(
[("What is AI?", "Artificial intelligence is ..."), ("What is ML?", "Machine learning is ..."),]
))
```
## ⚡️ FastEmbed on a GPU
FastEmbed supports running on GPU devices.
@@ -229,36 +246,36 @@ pip install qdrant-client[fastembed-gpu]
You might have to use quotes ```pip install 'qdrant-client[fastembed]'``` on zsh.
```python
from qdrant_client import QdrantClient
from qdrant_client import QdrantClient, models
# Initialize the client
client = QdrantClient("localhost", port=6333) # For production
# client = QdrantClient(":memory:") # For small experiments
# client = QdrantClient(":memory:") # For experimentation
# Prepare your documents, metadata, and IDs
docs = ["Qdrant has Langchain integrations", "Qdrant also has Llama Index integrations"]
metadata = [
{"source": "Langchain-docs"},
{"source": "Llama-index-docs"},
model_name = "sentence-transformers/all-MiniLM-L6-v2"
payload = [
{"document": "Qdrant has Langchain integrations", "source": "Langchain-docs", },
{"document": "Qdrant also has Llama Index integrations", "source": "LlamaIndex-docs"},
]
docs = [models.Document(text=data["document"], model=model_name) for data in payload]
ids = [42, 2]
# If you want to change the model:
# client.set_model("sentence-transformers/all-MiniLM-L6-v2")
# List of supported models: https://qdrant.github.io/fastembed/examples/Supported_Models
# Use the new add() instead of upsert()
# This internally calls embed() of the configured embedding model
client.add(
collection_name="demo_collection",
documents=docs,
metadata=metadata,
ids=ids
client.create_collection(
"demo_collection",
vectors_config=models.VectorParams(
size=client.get_embedding_size(model_name), distance=models.Distance.COSINE)
)
search_result = client.query(
client.upload_collection(
collection_name="demo_collection",
query_text="This is a query document"
vectors=docs,
ids=ids,
payload=payload,
)
search_result = client.query_points(
collection_name="demo_collection",
query=models.Document(text="This is a query document", model=model_name)
).points
print(search_result)
```
Binary file not shown.

Before

Width:  |  Height:  |  Size: 1.8 KiB

After

Width:  |  Height:  |  Size: 2.0 KiB

+134
View File
@@ -0,0 +1,134 @@
"""Export an inference-free SPLADE document encoder to ONNX.
Converts `opensearch-project/opensearch-neural-sparse-encoding-doc-v3-gte` (an MLM head
over a GTE backbone) into an onnx model producing token logits, and assembles a model dir
with everything fastembed's `IfSplade` needs: model.onnx, tokenizer files and idf.json.
Usage:
python experiments/if_splade_to_onnx.py --output-dir models/opensearch-neural-sparse-encoding-doc-v3-gte
"""
import argparse
import shutil
from pathlib import Path
import torch
from huggingface_hub import hf_hub_download
from transformers import AutoModelForMaskedLM, AutoTokenizer
MODEL_ID = "opensearch-project/opensearch-neural-sparse-encoding-doc-v3-gte"
# revision of the remote modeling code (Alibaba-NLP/new-impl), pinned in the model card
CODE_REVISION = "40ced75c3017eb27626c9d4ea981bde21a2662f4"
TOKENIZER_FILES = [
"config.json",
"tokenizer.json",
"tokenizer_config.json",
"special_tokens_map.json",
"vocab.txt",
"idf.json",
]
class LogitsOnly(torch.nn.Module):
def __init__(self, model: torch.nn.Module):
super().__init__()
self.model = model
def forward(self, input_ids: torch.Tensor, attention_mask: torch.Tensor) -> torch.Tensor:
return self.model(input_ids=input_ids, attention_mask=attention_mask).logits
def export(model_id: str, output_dir: Path, opset: int = 14) -> Path:
output_dir.mkdir(parents=True, exist_ok=True)
model = AutoModelForMaskedLM.from_pretrained(
model_id, trust_remote_code=True, code_revision=CODE_REVISION
)
model.eval()
wrapped = LogitsOnly(model)
tokenizer = AutoTokenizer.from_pretrained(model_id)
dummy = tokenizer(
["fastembed is a library", "onnx export"],
padding=True,
truncation=True,
return_tensors="pt",
return_token_type_ids=False,
)
onnx_path = output_dir / "model.onnx"
with torch.inference_mode():
torch.onnx.export(
wrapped,
(dummy["input_ids"], dummy["attention_mask"]),
f=onnx_path.as_posix(),
input_names=["input_ids", "attention_mask"],
output_names=["logits"],
dynamic_axes={
"input_ids": {0: "batch_size", 1: "sequence_length"},
"attention_mask": {0: "batch_size", 1: "sequence_length"},
"logits": {0: "batch_size", 1: "sequence_length"},
},
do_constant_folding=True,
opset_version=opset,
dynamo=False,
)
for file_name in TOKENIZER_FILES:
local_path = hf_hub_download(repo_id=model_id, filename=file_name)
shutil.copy(local_path, output_dir / file_name)
return onnx_path
def parity_check(model_id: str, output_dir: Path) -> None:
import numpy as np
import onnxruntime as ort
model = AutoModelForMaskedLM.from_pretrained(
model_id, trust_remote_code=True, code_revision=CODE_REVISION
)
model.eval()
tokenizer = AutoTokenizer.from_pretrained(model_id)
documents = [
"Currently New York is rainy.",
"fastembed is a lightweight library for generating embeddings",
"hello world",
]
features = tokenizer(
documents, padding=True, truncation=True, return_tensors="pt", return_token_type_ids=False
)
with torch.inference_mode():
torch_logits = model(**features).logits.numpy()
session = ort.InferenceSession(output_dir / "model.onnx")
onnx_logits = session.run(
["logits"],
{
"input_ids": features["input_ids"].numpy(),
"attention_mask": features["attention_mask"].numpy(),
},
)[0]
max_diff = np.abs(torch_logits - onnx_logits).max()
print(f"max |torch - onnx| logits diff: {max_diff}")
assert max_diff < 1e-3, "onnx export does not match the torch model"
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--model-id", default=MODEL_ID)
parser.add_argument("--output-dir", default=f"models/{MODEL_ID.replace('/', '_')}", type=Path)
parser.add_argument("--opset", default=14, type=int)
args = parser.parse_args()
onnx_path = export(args.model_id, args.output_dir, args.opset)
print(f"Exported to {onnx_path}")
parity_check(args.model_id, args.output_dir)
if __name__ == "__main__":
main()
+13 -7
View File
@@ -1,12 +1,17 @@
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
def deprecated_tar_struct(self) -> bool:
return self._deprecated_tar_struct
def __post_init__(self) -> None:
if self.hf is None and self.url is None:
@@ -28,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"
@@ -37,11 +42,12 @@ class DenseModelDescription(BaseModelDescription):
@dataclass(frozen=True)
class SparseModelDescription(BaseModelDescription):
requires_idf: Optional[bool] = None
vocab_size: Optional[int] = None
requires_idf: bool | None = None
vocab_size: int | None = None
class PoolingType(str, Enum):
CLS = "CLS"
MEAN = "MEAN"
LAST_TOKEN = "LAST_TOKEN"
DISABLED = "DISABLED"
+81 -28
View File
@@ -3,8 +3,9 @@ import time
import json
import shutil
import tarfile
from pathlib import Path
from typing import Any, Optional, Union, TypeVar, Generic
from copy import deepcopy
from pathlib import Path, PureWindowsPath
from typing import Any, TypeVar, Generic
import requests
from huggingface_hub import snapshot_download, model_info, list_repo_tree
@@ -98,7 +99,7 @@ class ModelManagement(Generic[T]):
if os.path.exists(output_path):
return output_path
response = requests.get(url, stream=True)
response = requests.get(url, stream=True, timeout=(10, 120))
# Handle HTTP errors
if response.status_code == 403:
@@ -179,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:
@@ -192,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)
@@ -224,11 +223,6 @@ class ModelManagement(Generic[T]):
logger.warning(
"Local file sizes do not match the metadata."
) # do not raise, still make an attempt to load the model
else:
logger.warning(
"Metadata file not found. Proceeding without checking local files."
) # if users have downloaded models from hf manually, or they're updating from previous versions of
# fastembed
result = snapshot_download(
repo_id=hf_source_repo,
allow_patterns=allow_patterns,
@@ -310,29 +304,60 @@ class ModelManagement(Generic[T]):
try:
# Open the tar.gz file
with tarfile.open(targz_path, "r:gz") as tar:
# Extract all files into the cache directory
tar.extractall(
path=cache_dir,
)
except tarfile.TarError as e:
if hasattr(tarfile, "data_filter"):
tar.extractall(path=cache_dir, filter="data")
else:
# No PEP 706 filter before 3.10.12, so vet the members by hand.
members = tar.getmembers()
for member in members:
cls._validate_tar_member(member)
tar.extractall(path=cache_dir, members=members)
except (tarfile.TarError, ValueError) as e:
# If any error occurs while opening or extracting the tar.gz file,
# delete the cache directory (if it was created in this function)
# and raise the error again
if "tmp" in cache_dir:
if "tmp" in cache_dir and os.path.exists(cache_dir):
shutil.rmtree(cache_dir)
raise ValueError(f"An error occurred while decompressing {targz_path}: {e}")
raise ValueError(f"An error occurred while decompressing {targz_path}: {e}") from e
return cache_dir
@staticmethod
def _is_unsafe_tar_path(path: str) -> bool:
"""Checks whether a tar member name or link target may escape the extraction dir.
Lexical on purpose: resolving against the extraction directory is unsound before
extraction, since `link/../escape` only escapes once an earlier member has been
written as a symlink. Any `..` component is therefore rejected outright.
"""
# PureWindowsPath splits on both separators, so `root` covers POSIX "/evil" as
# well as "\\evil", which escapes on Windows without being absolute.
windows_path = PureWindowsPath(path)
return bool(windows_path.drive or windows_path.root) or ".." in windows_path.parts
@classmethod
def _validate_tar_member(cls, member: tarfile.TarInfo) -> None:
"""Raises ValueError if a member could write outside the extraction directory."""
if cls._is_unsafe_tar_path(member.name):
raise ValueError(f"Unsafe tar member path: {member.name}")
if member.issym() or member.islnk():
if cls._is_unsafe_tar_path(member.linkname):
raise ValueError(f"Unsafe tar link target: {member.name} -> {member.linkname}")
elif not (member.isfile() or member.isdir()):
# Devices, fifos and the like have no place in a model archive.
raise ValueError(f"Unsupported tar member type: {member.name}")
@classmethod
def retrieve_model_gcs(
cls,
model_name: str,
source_url: str,
cache_dir: str,
deprecated_tar_struct: bool = False,
local_files_only: bool = False,
) -> Path:
fast_model_name = f"fast-{model_name.split('/')[-1]}"
fast_model_name = f"{'fast-' if deprecated_tar_struct else ''}{model_name.split('/')[-1]}"
cache_tmp_dir = Path(cache_dir) / "tmp"
model_tmp_dir = cache_tmp_dir / fast_model_name
model_dir = Path(cache_dir) / fast_model_name
@@ -400,21 +425,47 @@ class ModelManagement(Generic[T]):
Path: The path to the downloaded model directory.
"""
local_files_only = kwargs.get("local_files_only", False)
specific_model_path: Optional[str] = kwargs.pop("specific_model_path", None)
hf_offline = os.environ.get("HF_HUB_OFFLINE", "").strip().upper()
if not local_files_only and hf_offline in {"1", "TRUE", "YES", "ON"}:
local_files_only = True
kwargs["local_files_only"] = True
specific_model_path: str | None = kwargs.pop("specific_model_path", None)
if specific_model_path:
return Path(specific_model_path)
retries = 1 if local_files_only else retries
hf_source = model.sources.hf
url_source = model.sources.url
extra_patterns = [model.model_file]
extra_patterns.extend(model.additional_files)
if hf_source:
try:
cache_kwargs = deepcopy(kwargs)
cache_kwargs["local_files_only"] = True
resolved_path = Path(
cls.download_files_from_huggingface(
hf_source,
cache_dir=cache_dir,
extra_patterns=extra_patterns,
**cache_kwargs,
)
)
if (resolved_path / model.model_file).exists() and all(
(resolved_path / file).exists() for file in extra_patterns
):
return resolved_path
except Exception:
pass
finally:
enable_progress_bars()
sleep = 3.0
while retries > 0:
retries -= 1
if hf_source:
extra_patterns = [model.model_file]
extra_patterns.extend(model.additional_files)
if hf_source and not local_files_only:
# we have already tried loading with `local_files_only=True` via hf and we failed
try:
return Path(
cls.download_files_from_huggingface(
@@ -438,6 +489,7 @@ class ModelManagement(Generic[T]):
model.model,
str(url_source),
str(cache_dir),
deprecated_tar_struct=model.sources.deprecated_tar_struct,
local_files_only=local_files_only,
)
except Exception:
@@ -446,11 +498,12 @@ class ModelManagement(Generic[T]):
if local_files_only:
logger.error("Could not find model in cache_dir")
break
else:
logger.error(
f"Could not download model from either source, sleeping for {sleep} seconds, {retries} retries left."
)
time.sleep(sleep)
sleep *= 3
time.sleep(sleep)
sleep *= 3
raise ValueError(f"Could not load model {model.model} from any source.")
+79 -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,21 +19,45 @@ T = TypeVar("T")
@dataclass
class OnnxOutputContext:
model_output: NumpyArray
attention_mask: Optional[NDArray[np.int64]] = None
input_ids: Optional[NDArray[np.int64]] = None
attention_mask: NDArray[np.int64] | None = None
input_ids: NDArray[np.int64] | None = None
metadata: dict[str, Any] | None = None
class OnnxModel(Generic[T]):
EXPOSED_SESSION_OPTIONS = ("enable_cpu_mem_arena",)
@classmethod
def _get_worker_class(cls) -> Type["EmbeddingWorker[T]"]:
raise NotImplementedError("Subclasses must implement this method")
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[T]:
def _get_worker_init_kwargs(self) -> dict[str, Any]:
"""Additional kwargs a worker process needs to reconstruct this model.
Workers are started with `spawn`/`forkserver`, hence they don't inherit class-level state
which has been set up in runtime, e.g. models registered via `add_custom_model`.
Such state has to be shipped to the workers explicitly.
Returns:
dict[str, Any]: kwargs to pass to `_get_worker_class().init_embedding`.
"""
return {}
def _post_process_onnx_output(self, output: OnnxOutputContext, **kwargs: Any) -> Iterable[T]:
"""Post-process the ONNX model output to convert it into a usable format.
Args:
output (OnnxOutputContext): The raw output from the ONNX model.
**kwargs: Additional keyword arguments that may be needed by specific implementations.
Returns:
Iterable[T]: Post-processed output as an iterable of type T.
"""
raise NotImplementedError("Subclasses must implement this method")
def __init__(self) -> None:
self.model: Optional[ort.InferenceSession] = None
self.tokenizer: Optional[Tokenizer] = None
self.model: ort.InferenceSession | None = None
self.tokenizer: Tokenizer | None = None
def _preprocess_onnx_input(
self, onnx_input: dict[str, NumpyArray], **kwargs: Any
@@ -47,24 +71,30 @@ class OnnxModel(Generic[T]):
self,
model_dir: Path,
model_file: str,
threads: Optional[int],
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_id: Optional[int] = None,
threads: int | None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_id: int | None = None,
extra_session_options: dict[str, Any] | None = None,
) -> None:
model_path = model_dir / model_file
# List of Execution Providers: https://onnxruntime.ai/docs/execution-providers
available_providers = ort.get_available_providers()
cuda_available = "CUDAExecutionProvider" in available_providers
explicit_cuda = cuda is True or cuda == Device.CUDA
if cuda and providers is not None:
if explicit_cuda and providers is not None:
warnings.warn(
f"`cuda` and `providers` are mutually exclusive parameters, cuda: {cuda}, providers: {providers}",
f"`cuda` and `providers` are mutually exclusive parameters, "
f"cuda: {cuda}, providers: {providers}. If you'd like to use providers, cuda should be one of "
f"[False, Device.CPU, Device.AUTO].",
category=UserWarning,
stacklevel=6,
)
if providers is not None:
onnx_providers = list(providers)
elif cuda:
elif explicit_cuda or (cuda == Device.AUTO and cuda_available):
if device_id is None:
onnx_providers = ["CUDAExecutionProvider"]
else:
@@ -72,7 +102,6 @@ class OnnxModel(Generic[T]):
else:
onnx_providers = ["CPUExecutionProvider"]
available_providers = ort.get_available_providers()
requested_provider_names: list[str] = []
for provider in onnx_providers:
# check providers available
@@ -90,6 +119,9 @@ class OnnxModel(Generic[T]):
so.intra_op_num_threads = threads
so.inter_op_num_threads = threads
if extra_session_options is not None:
self.add_extra_session_options(so, extra_session_options)
self.model = ort.InferenceSession(
str(model_path), providers=onnx_providers, sess_options=so
)
@@ -104,6 +136,38 @@ class OnnxModel(Generic[T]):
RuntimeWarning,
)
@classmethod
def _select_exposed_session_options(cls, model_kwargs: dict[str, Any]) -> dict[str, Any]:
"""A convenience method to select the exposed session options in models
Args:
model_kwargs (dict[str, Any]): The model kwargs.
Returns:
dict[str, Any]: a dict with filtered exposed session options.
"""
return {k: v for k, v in model_kwargs.items() if k in cls.EXPOSED_SESSION_OPTIONS}
@classmethod
def add_extra_session_options(
cls, session_options: ort.SessionOptions, extra_options: dict[str, Any]
) -> None:
"""Add extra session options to the existing options object in-place
Args:
session_options (ort.SessionOptions): The existing session options object.
extra_options (dict[str, Any]): The extra session options available in cls.EXPOSED_SESSION_OPTIONS.
Returns:
None
"""
for option in extra_options:
assert (
option in cls.EXPOSED_SESSION_OPTIONS
), f"{option} is unknown or not exposed (exposed options: {cls.EXPOSED_SESSION_OPTIONS})"
if "enable_cpu_mem_arena" in extra_options:
session_options.enable_cpu_mem_arena = extra_options["enable_cpu_mem_arena"]
def load_onnx_model(self) -> None:
raise NotImplementedError("Subclasses must implement this method")
+96 -28
View File
@@ -1,5 +1,6 @@
import json
from typing import Any
import sys
from typing import Any, Iterator
from pathlib import Path
from tokenizers import AddedToken, Tokenizer
@@ -8,9 +9,10 @@ from fastembed.image.transform.operators import Compose
def load_special_tokens(model_dir: Path) -> dict[str, Any]:
"""Read special_tokens_map.json, treating an absent file as an empty map."""
tokens_map_path = model_dir / "special_tokens_map.json"
if not tokens_map_path.exists():
raise ValueError(f"Could not find special_tokens_map.json in {model_dir}")
return {}
with open(str(tokens_map_path)) as tokens_map_file:
tokens_map = json.load(tokens_map_file)
@@ -18,11 +20,58 @@ def load_special_tokens(model_dir: Path) -> dict[str, Any]:
return tokens_map
def load_tokenizer(model_dir: Path) -> tuple[Tokenizer, dict[str, int]]:
config_path = model_dir / "config.json"
if not config_path.exists():
raise ValueError(f"Could not find config.json in {model_dir}")
def iter_special_tokens(tokens_map: dict[str, Any]) -> Iterator[str | dict[str, Any]]:
"""Yield the individual tokens declared in a special tokens map.
Most keys hold one token, but `additional_special_tokens` holds a list of them.
"""
for value in tokens_map.values():
if isinstance(value, list):
yield from value
else:
yield value
def _valid_context(value: Any) -> int | None:
"""Return `value` if it can be used as a truncation limit, `None` otherwise.
Config files do not always carry a real limit: transformers writes `model_max_length` as
1e30 when the value is unknown, and some repos ship a 0 or a null. `enable_truncation`
raises an `OverflowError` on the former and silently produces empty encodings on the
latter, so both are rejected here rather than passed through.
"""
if isinstance(value, bool) or not isinstance(value, int):
return None
if not 0 < value <= sys.maxsize:
return None
return value
def _resolve_max_context(tokenizer_config: dict[str, Any], model_dir: Path) -> int:
"""Pick the truncation limit, preferring the stricter of the two tokenizer config keys.
`config.json:max_position_embeddings` deliberately is not used as a fallback: it is the size
of the position table, not the usable context, and the two differ per architecture, e.g.
roberta reports 514 for a usable 512.
"""
candidates = [
context
for context in (
_valid_context(tokenizer_config.get("model_max_length")),
_valid_context(tokenizer_config.get("max_length")),
)
if context is not None
]
if not candidates:
raise ValueError(
f"Could not determine the maximum context length for {model_dir}. Set a positive "
"`model_max_length` or `max_length` in tokenizer_config.json."
)
return min(candidates)
def load_tokenizer(model_dir: Path) -> tuple[Tokenizer, dict[str, int]]:
tokenizer_path = model_dir / "tokenizer.json"
if not tokenizer_path.exists():
raise ValueError(f"Could not find tokenizer.json in {model_dir}")
@@ -31,43 +80,62 @@ def load_tokenizer(model_dir: Path) -> tuple[Tokenizer, dict[str, int]]:
if not tokenizer_config_path.exists():
raise ValueError(f"Could not find tokenizer_config.json in {model_dir}")
with open(str(config_path)) as config_file:
config = json.load(config_file)
# config.json is optional: transformers v5 no longer writes it for every model.
config_path = model_dir / "config.json"
config: dict[str, Any] = {}
if config_path.exists():
with open(str(config_path)) as config_file:
config = json.load(config_file)
with open(str(tokenizer_config_path)) as tokenizer_config_file:
tokenizer_config = json.load(tokenizer_config_file)
assert (
"model_max_length" in tokenizer_config or "max_length" in tokenizer_config
), "Models without model_max_length or max_length are not supported."
if "model_max_length" not in tokenizer_config:
max_context = tokenizer_config["max_length"]
elif "max_length" not in tokenizer_config:
max_context = tokenizer_config["model_max_length"]
else:
max_context = min(tokenizer_config["model_max_length"], tokenizer_config["max_length"])
max_context = _resolve_max_context(tokenizer_config, model_dir)
tokens_map = load_special_tokens(model_dir)
tokenizer = Tokenizer.from_file(str(tokenizer_path))
tokenizer.enable_truncation(max_length=max_context)
tokenizer.enable_padding(
pad_id=config.get("pad_token_id", 0), pad_token=tokenizer_config["pad_token"]
)
for token in tokens_map.values():
# Registered before the padding is resolved: the map may name a pad token that
# tokenizer.json does not carry, and it only gets an id once it is added.
for token in iter_special_tokens(tokens_map):
if isinstance(token, str):
tokenizer.add_special_tokens([token])
elif isinstance(token, dict):
tokenizer.add_special_tokens([AddedToken(**token)])
special_token_to_id: dict[str, int] = {}
# Padding is always normalized to batch-longest. A serialized fixed length shorter than the
# truncation limit leaves longer encodings untouched, which produces ragged batches, and a
# fixed length equal to it pads every batch to the maximum. Direction and pad token metadata
# are taken from the serialized settings, since some models pad on the left.
padding = tokenizer.padding or {}
pad_token = padding.get("pad_token") or tokenizer_config.get("pad_token")
if pad_token is None:
raise ValueError(f"Could not find a pad token for {model_dir}")
for token in tokens_map.values():
if isinstance(token, str):
special_token_to_id[token] = tokenizer.token_to_id(token)
elif isinstance(token, dict):
token_str = token.get("content", "")
special_token_to_id[token_str] = tokenizer.token_to_id(token_str)
# The vocabulary is the last resort, not a hardcoded 0: that silently disagrees with
# `pad_token` for every model whose pad token is not the first entry.
pad_id = padding.get("pad_id", config.get("pad_token_id"))
if pad_id is None:
pad_id = tokenizer.token_to_id(pad_token)
if pad_id is None:
raise ValueError(f"Could not resolve an id for the pad token {pad_token!r} in {model_dir}")
tokenizer.enable_padding(
direction=padding.get("direction", "right"),
pad_id=pad_id,
pad_type_id=padding.get("pad_type_id", 0),
pad_token=pad_token,
pad_to_multiple_of=padding.get("pad_to_multiple_of"),
length=None,
)
special_token_to_id = {
token.content: token_id
for token_id, token in tokenizer.get_added_tokens_decoder().items()
if token.special
}
return tokenizer, special_token_to_id
+21 -18
View File
@@ -1,24 +1,27 @@
from enum import Enum
from pathlib import Path
import sys
from PIL import Image
from typing import Any, Union
from typing import Any, TypeAlias
import numpy as np
from numpy.typing import NDArray
if sys.version_info >= (3, 10):
from typing import TypeAlias
else:
from typing_extensions import TypeAlias
from PIL import Image
PathInput: TypeAlias = Union[str, Path]
ImageInput: TypeAlias = Union[PathInput, Image.Image]
class Device(str, Enum):
CPU = "cpu"
CUDA = "cuda"
AUTO = "auto"
OnnxProvider: TypeAlias = Union[str, tuple[str, dict[Any, Any]]]
NumpyArray = Union[
NDArray[np.float32],
NDArray[np.float16],
NDArray[np.int8],
NDArray[np.int64],
NDArray[np.int32],
]
PathInput: TypeAlias = str | Path
ImageInput: TypeAlias = PathInput | Image.Image
OnnxProvider: TypeAlias = str | tuple[str, dict[Any, Any]]
NumpyArray: TypeAlias = (
NDArray[np.float64]
| NDArray[np.float32]
| NDArray[np.float16]
| NDArray[np.int8]
| NDArray[np.int64]
| NDArray[np.int32]
)
+12 -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
@@ -32,6 +32,16 @@ def mean_pooling(input_array: NumpyArray, attention_mask: NDArray[np.int64]) ->
return pooled_embeddings
def last_token_pooling(input_array: NumpyArray, attention_mask: NDArray[np.int64]) -> NumpyArray:
"""Take the embedding of the last non-padding token of each sequence.
Locates the last position the attention mask marks as real, so it holds whichever
side the tokenizer pads on.
"""
last_token_indices = attention_mask.shape[1] - 1 - np.argmax(attention_mask[:, ::-1], axis=1)
return input_array[np.arange(input_array.shape[0]), last_token_indices]
def iter_batch(iterable: Iterable[T], size: int) -> Iterable[list[T]]:
"""
>>> list(iter_batch([1,2,3,4,5], 3))
@@ -45,7 +55,7 @@ def iter_batch(iterable: Iterable[T], size: int) -> Iterable[list[T]]:
yield b
def define_cache_dir(cache_dir: Optional[str] = None) -> Path:
def define_cache_dir(cache_dir: str | None = None) -> Path:
"""
Define the cache directory for fastembed
"""
+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)
+50 -10
View File
@@ -1,15 +1,21 @@
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Iterable, Sequence, Type
from dataclasses import asdict
from fastembed.common.types import NumpyArray
from fastembed.common.types import NumpyArray, Device
from fastembed.common import ImageInput, OnnxProvider
from fastembed.image.image_embedding_base import ImageEmbeddingBase
from fastembed.image.onnx_embedding import OnnxImageEmbedding
from fastembed.image.normalized_embedding import NormalizedEmbedding
from fastembed.image.siglip_embedding import SiglipOnnxImageEmbedding
from fastembed.common.model_description import DenseModelDescription
class ImageEmbedding(ImageEmbeddingBase):
EMBEDDINGS_REGISTRY: list[Type[ImageEmbeddingBase]] = [OnnxImageEmbedding]
EMBEDDINGS_REGISTRY: list[Type[ImageEmbeddingBase]] = [
OnnxImageEmbedding,
NormalizedEmbedding,
SiglipOnnxImageEmbedding,
]
@classmethod
def list_supported_models(cls) -> list[dict[str, Any]]:
@@ -48,11 +54,11 @@ class ImageEmbedding(ImageEmbeddingBase):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
**kwargs: Any,
):
@@ -77,11 +83,45 @@ class ImageEmbedding(ImageEmbeddingBase):
"Please check the supported models using `ImageEmbedding.list_supported_models()`"
)
@property
def embedding_size(self) -> int:
"""Get the embedding size of the current model"""
if self._embedding_size is None:
self._embedding_size = self.get_embedding_size(self.model_name)
return self._embedding_size
@classmethod
def get_embedding_size(cls, model_name: str) -> int:
"""Get the embedding size of the passed model
Args:
model_name (str): The name of the model to get embedding size for.
Returns:
int: The size of the embedding.
Raises:
ValueError: If the model name is not found in the supported models.
"""
descriptions = cls._list_supported_models()
embedding_size: int | None = None
for description in descriptions:
if description.model.lower() == model_name.lower():
embedding_size = description.dim
break
if embedding_size is None:
model_names = [description.model for description in descriptions]
raise ValueError(
f"Embedding size for model {model_name} was None. "
f"Available model names: {model_names}"
)
return embedding_size
def embed(
self,
images: Union[ImageInput, Iterable[ImageInput]],
images: ImageInput | Iterable[ImageInput],
batch_size: int = 16,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
"""
+16 -5
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,20 +10,21 @@ class ImageEmbeddingBase(ModelManagement[DenseModelDescription]):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
cache_dir: str | None = None,
threads: int | None = None,
**kwargs: Any,
):
self.model_name = model_name
self.cache_dir = cache_dir
self.threads = threads
self._local_files_only = kwargs.pop("local_files_only", False)
self._embedding_size: int | None = None
def embed(
self,
images: Union[ImageInput, Iterable[ImageInput]],
images: ImageInput | Iterable[ImageInput],
batch_size: int = 16,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
"""
@@ -42,3 +43,13 @@ class ImageEmbeddingBase(ModelManagement[DenseModelDescription]):
Iterable[NdArray]: The embeddings.
"""
raise NotImplementedError()
@classmethod
def get_embedding_size(cls, model_name: str) -> int:
"""Returns embedding size of the chosen model."""
raise NotImplementedError("Subclasses must implement this method")
@property
def embedding_size(self) -> int:
"""Returns embedding size for the current model"""
raise NotImplementedError("Subclasses must implement this method")
+69
View File
@@ -0,0 +1,69 @@
from typing import Any, Iterable, Type
from fastembed.common.types import NumpyArray
from fastembed.common.onnx_model import OnnxOutputContext
from fastembed.common.utils import normalize
from fastembed.image.onnx_embedding import OnnxImageEmbedding
from fastembed.image.onnx_image_model import ImageEmbeddingWorker
from fastembed.common.model_description import DenseModelDescription, ModelSource
supported_normalized_models: list[DenseModelDescription] = [
DenseModelDescription(
model="nomic-ai/nomic-embed-vision-v1.5",
dim=768,
description="Image embeddings, Multimodal (text&image), 2024 year",
license="apache-2.0",
size_in_GB=0.37,
sources=ModelSource(hf="nomic-ai/nomic-embed-vision-v1.5"),
model_file="onnx/model.onnx",
),
DenseModelDescription(
model="nomic-ai/nomic-embed-vision-v1.5-Q",
dim=768,
description="Image embeddings, Multimodal (text&image), 2024 year",
license="apache-2.0",
size_in_GB=0.1,
sources=ModelSource(hf="nomic-ai/nomic-embed-vision-v1.5"),
model_file="onnx/model_quantized.onnx",
),
]
class NormalizedEmbedding(OnnxImageEmbedding):
@classmethod
def _list_supported_models(cls) -> list[DenseModelDescription]:
"""
Lists the supported models.
Returns:
list[DenseModelDescription]: A list of DenseModelDescription objects containing the model information.
"""
return supported_normalized_models
@classmethod
def _get_worker_class(cls) -> Type["ImageEmbeddingWorker[NumpyArray]"]:
return NormalizedEmbeddingWorker
def _post_process_onnx_output(
self, output: OnnxOutputContext, **kwargs: Any
) -> Iterable[NumpyArray]:
# The model emits last_hidden_state, which onnx_embed flattens to (batch, tokens * dim).
# Recover the token axis, take the CLS token (index 0) and normalize, matching the reference
# F.normalize(last_hidden_state[:, 0], p=2, dim=1).
dim = self.model_description.dim
assert dim is not None, "Model description is missing the embedding dim"
hidden_states = output.model_output.reshape(output.model_output.shape[0], -1, dim)
return normalize(hidden_states[:, 0])
class NormalizedEmbeddingWorker(ImageEmbeddingWorker[NumpyArray]):
def init_embedding(
self, model_name: str, cache_dir: str, **kwargs: Any
) -> NormalizedEmbedding:
return NormalizedEmbedding(
model_name=model_name,
cache_dir=cache_dir,
threads=1,
**kwargs,
)
+27 -19
View File
@@ -1,8 +1,7 @@
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Iterable, Sequence, Type
import numpy as np
from fastembed.common.types import NumpyArray
from fastembed.common.types import NumpyArray, Device
from fastembed.common import ImageInput, OnnxProvider
from fastembed.common.onnx_model import OnnxOutputContext
from fastembed.common.utils import define_cache_dir, normalize
@@ -64,14 +63,14 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
device_id: Optional[int] = None,
specific_model_path: Optional[str] = None,
device_id: int | None = None,
specific_model_path: str | None = None,
**kwargs: Any,
):
"""
@@ -83,10 +82,11 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.
providers (Optional[Sequence[OnnxProvider]], optional): The list of onnxruntime providers to use.
Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to False.
cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to Device.AUTO.
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, mutually exclusive
with `providers`. Defaults to None.
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
@@ -99,13 +99,14 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
super().__init__(model_name, cache_dir, threads, **kwargs)
self.providers = providers
self.lazy_load = lazy_load
self._extra_session_options = self._select_exposed_session_options(kwargs)
# List of device ids, that can be used for data parallel processing in workers
self.device_ids = device_ids
self.cuda = cuda
# This device_id will be used if we need to load model in current process
self.device_id: Optional[int] = None
self.device_id: int | None = None
if device_id is not None:
self.device_id = device_id
elif self.device_ids is not None:
@@ -113,11 +114,12 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
self.model_description = self._get_model_description(model_name)
self.cache_dir = str(define_cache_dir(cache_dir))
self._specific_model_path = specific_model_path
self._model_dir = self.download_model(
self.model_description,
self.cache_dir,
local_files_only=self._local_files_only,
specific_model_path=specific_model_path,
specific_model_path=self._specific_model_path,
)
if not self.lazy_load:
@@ -134,6 +136,7 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
providers=self.providers,
cuda=self.cuda,
device_id=self.device_id,
extra_session_options=self._extra_session_options,
)
@classmethod
@@ -148,9 +151,9 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
def embed(
self,
images: Union[ImageInput, Iterable[ImageInput]],
images: ImageInput | Iterable[ImageInput],
batch_size: int = 16,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
"""
@@ -178,6 +181,9 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
providers=self.providers,
cuda=self.cuda,
device_ids=self.device_ids,
local_files_only=self._local_files_only,
specific_model_path=self._specific_model_path,
extra_session_options=self._extra_session_options,
**kwargs,
)
@@ -194,8 +200,10 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
return onnx_input
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[NumpyArray]:
return normalize(output.model_output).astype(np.float32)
def _post_process_onnx_output(
self, output: OnnxOutputContext, **kwargs: Any
) -> Iterable[NumpyArray]:
return normalize(output.model_output)
class OnnxImageEmbeddingWorker(ImageEmbeddingWorker[NumpyArray]):
+41 -18
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
@@ -19,16 +19,27 @@ from fastembed.parallel_processor import ParallelWorkerPool
class OnnxImageModel(OnnxModel[T]):
ONNX_OUTPUT_NAMES: list[str] | None = None
@classmethod
def _get_worker_class(cls) -> Type["ImageEmbeddingWorker[T]"]:
raise NotImplementedError("Subclasses must implement this method")
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[T]:
def _post_process_onnx_output(self, output: OnnxOutputContext, **kwargs: Any) -> Iterable[T]:
"""Post-process the ONNX model output to convert it into a usable format.
Args:
output (OnnxOutputContext): The raw output from the ONNX model.
**kwargs: Additional keyword arguments that may be needed by specific implementations.
Returns:
Iterable[T]: Post-processed output as an iterable of type T.
"""
raise NotImplementedError("Subclasses must implement this method")
def __init__(self) -> None:
super().__init__()
self.processor: Optional[Compose] = None
self.processor: Compose | None = None
def _preprocess_onnx_input(
self, onnx_input: dict[str, NumpyArray], **kwargs: Any
@@ -42,10 +53,11 @@ class OnnxImageModel(OnnxModel[T]):
self,
model_dir: Path,
model_file: str,
threads: Optional[int],
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_id: Optional[int] = None,
threads: int | None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_id: int | None = None,
extra_session_options: dict[str, Any] | None = None,
) -> None:
super()._load_onnx_model(
model_dir=model_dir,
@@ -54,6 +66,7 @@ class OnnxImageModel(OnnxModel[T]):
providers=providers,
cuda=cuda,
device_id=device_id,
extra_session_options=extra_session_options,
)
self.processor = load_preprocessor(model_dir=model_dir)
@@ -65,16 +78,18 @@ class OnnxImageModel(OnnxModel[T]):
return {input_name: encoded}
def onnx_embed(self, images: list[ImageInput], **kwargs: Any) -> OnnxOutputContext:
with contextlib.ExitStack():
with contextlib.ExitStack() as stack:
image_files = [
Image.open(image) if not isinstance(image, Image.Image) else image
stack.enter_context(Image.open(image))
if not isinstance(image, Image.Image)
else image
for image in images
]
assert self.processor is not None, "Processor is not initialized"
encoded = np.array(self.processor(image_files))
onnx_input = self._build_onnx_input(encoded)
onnx_input = self._preprocess_onnx_input(onnx_input)
model_output = self.model.run(None, onnx_input) # type: ignore[union-attr]
model_output = self.model.run(self.ONNX_OUTPUT_NAMES, onnx_input) # type: ignore[union-attr]
embeddings = model_output[0].reshape(len(images), -1)
return OnnxOutputContext(model_output=embeddings)
@@ -82,12 +97,15 @@ class OnnxImageModel(OnnxModel[T]):
self,
model_name: str,
cache_dir: str,
images: Union[ImageInput, Iterable[ImageInput]],
images: ImageInput | Iterable[ImageInput],
batch_size: int = 256,
parallel: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
parallel: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
local_files_only: bool = False,
specific_model_path: str | None = None,
extra_session_options: dict[str, Any] | None = None,
**kwargs: Any,
) -> Iterable[T]:
is_small = False
@@ -104,7 +122,7 @@ class OnnxImageModel(OnnxModel[T]):
self.load_onnx_model()
for batch in iter_batch(images, batch_size):
yield from self._post_process_onnx_output(self.onnx_embed(batch))
yield from self._post_process_onnx_output(self.onnx_embed(batch), **kwargs)
else:
if parallel == 0:
parallel = os.cpu_count()
@@ -114,9 +132,14 @@ class OnnxImageModel(OnnxModel[T]):
"model_name": model_name,
"cache_dir": cache_dir,
"providers": providers,
"local_files_only": local_files_only,
"specific_model_path": specific_model_path,
**kwargs,
}
if extra_session_options is not None:
params.update(extra_session_options)
pool = ParallelWorkerPool(
num_workers=parallel or 1,
worker=self._get_worker_class(),
@@ -125,7 +148,7 @@ class OnnxImageModel(OnnxModel[T]):
start_method=start_method,
)
for batch in pool.ordered_map(iter_batch(images, batch_size), **params):
yield from self._post_process_onnx_output(batch) # type: ignore
yield from self._post_process_onnx_output(batch, **kwargs) # type: ignore
class ImageEmbeddingWorker(EmbeddingWorker[T]):
+44
View File
@@ -0,0 +1,44 @@
from typing import Any, Type
from fastembed.common.model_description import DenseModelDescription, ModelSource
from fastembed.image.onnx_embedding import OnnxImageEmbedding, OnnxImageEmbeddingWorker
supported_siglip_models: list[DenseModelDescription] = [
DenseModelDescription(
model="google/siglip2-base-patch16-224",
dim=768,
description="Image embeddings, Multimodal (text&image), 2025 year",
license="apache-2.0",
size_in_GB=0.37,
sources=ModelSource(hf="onnx-community/siglip2-base-patch16-224-ONNX"),
model_file="onnx/vision_model.onnx",
),
]
class SiglipOnnxImageEmbedding(OnnxImageEmbedding):
"""SigLIP vision tower.
The exported graph returns both the per-patch `last_hidden_state` and the pooled
`pooler_output`; only the latter is the image embedding, so it must be selected explicitly.
"""
ONNX_OUTPUT_NAMES = ["pooler_output"]
@classmethod
def _get_worker_class(cls) -> Type["OnnxImageEmbeddingWorker"]:
return SiglipImageEmbeddingWorker
@classmethod
def _list_supported_models(cls) -> list[DenseModelDescription]:
return supported_siglip_models
class SiglipImageEmbeddingWorker(OnnxImageEmbeddingWorker):
def init_embedding(self, model_name: str, cache_dir: str, **kwargs: Any) -> OnnxImageEmbedding:
return SiglipOnnxImageEmbedding(
model_name=model_name,
cache_dir=cache_dir,
threads=1,
**kwargs,
)
+106 -21
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,43 +62,56 @@ def center_crop(
def normalize(
image: NumpyArray,
mean: Union[float, list[float]],
std: Union[float, list[float]],
mean: float | list[float],
std: float | list[float],
) -> NumpyArray:
num_channels = image.shape[1] if len(image.shape) == 4 else image.shape[0]
if image.ndim < 3:
raise ValueError(f"image must be (C, H, W) or (N, C, H, W), got shape {image.shape}")
# Channels sit on the third axis from the end, which covers (C, H, W) and
# (N, C, H, W) alike. Transposing instead reversed every axis, which put the
# batch dimension where the channels were meant to be.
num_channels = image.shape[-3]
if not np.issubdtype(image.dtype, np.floating):
image = image.astype(np.float32)
mean = mean if isinstance(mean, list) else [mean] * num_channels
mean_list = mean if isinstance(mean, list) else [mean] * num_channels
if len(mean) != num_channels:
if len(mean_list) != num_channels:
raise ValueError(
f"mean must have the same number of channels as the image, image has {num_channels} channels, got "
f"{len(mean)}"
f"{len(mean_list)}"
)
mean_arr = np.array(mean, dtype=np.float32)
# (C, 1, 1) lines the channels up with the trailing (C, H, W) axes under numpy
# broadcasting, whatever batch dimensions lead them.
mean_arr = np.array(mean_list, dtype=np.float32).reshape(-1, 1, 1)
std = std if isinstance(std, list) else [std] * num_channels
if len(std) != num_channels:
std_list = std if isinstance(std, list) else [std] * num_channels
if len(std_list) != num_channels:
raise ValueError(
f"std must have the same number of channels as the image, image has {num_channels} channels, got {len(std)}"
f"std must have the same number of channels as the image, image has {num_channels} channels, got {len(std_list)}"
)
std_arr = np.array(std, dtype=np.float32)
std_arr = np.array(std_list, dtype=np.float32).reshape(-1, 1, 1)
image = ((image.T - mean_arr) / std_arr).T
return image
image_upd = (image - mean_arr) / std_arr
return image_upd
def resize(
image: Image.Image,
size: Union[int, tuple[int, int]],
resample: Union[int, Image.Resampling] = Image.Resampling.BILINEAR,
size: int | tuple[int, int],
resample: int | Image.Resampling = Image.Resampling.BILINEAR,
) -> Image.Image:
if isinstance(size, tuple):
return image.resize(size, resample)
# fastembed keeps sizes as (height, width) — `Compose.from_config` builds the
# tuple as (size["height"], size["width"]) — while Pillow's resize takes
# (width, height). The two agree for square sizes, so this only shows up on a
# non-square image processor configuration.
height, width = size
return image.resize((width, height), resample)
height, width = image.height, image.width
short, long = (width, height) if width <= height else (height, width)
@@ -117,7 +128,7 @@ def rescale(image: NumpyArray, scale: float, dtype: type = np.float32) -> NumpyA
return (image * scale).astype(dtype)
def pil2ndarray(image: Union[Image.Image, NumpyArray]) -> NumpyArray:
def pil2ndarray(image: Image.Image | NumpyArray) -> NumpyArray:
if isinstance(image, Image.Image):
return np.asarray(image).transpose((2, 0, 1))
return image
@@ -126,7 +137,7 @@ def pil2ndarray(image: Union[Image.Image, NumpyArray]) -> NumpyArray:
def pad2square(
image: Image.Image,
size: int,
fill_color: Union[str, int, tuple[int, ...]] = 0,
fill_color: str | int | tuple[int, ...] = 0,
) -> Image.Image:
height, width = image.height, image.width
@@ -147,3 +158,77 @@ def pad2square(
new_image = Image.new(mode="RGB", size=(size, size), color=fill_color)
new_image.paste(image.crop((left, top, right, bottom)) if crop_required else image)
return new_image
def resize_longest_edge(
image: Image.Image,
max_size: int,
resample: int | Image.Resampling = Image.Resampling.LANCZOS,
) -> Image.Image:
height, width = image.height, image.width
aspect_ratio = width / height
if width >= height:
# Width is longer
new_width = max_size
new_height = int(new_width / aspect_ratio)
else:
# Height is longer
new_height = max_size
new_width = int(new_height * aspect_ratio)
# Ensure even dimensions
if new_height % 2 != 0:
new_height += 1
if new_width % 2 != 0:
new_width += 1
return image.resize((new_width, new_height), resample)
def crop_ndarray(
image: NumpyArray,
x1: int,
y1: int,
x2: int,
y2: int,
channel_first: bool = True,
) -> NumpyArray:
if channel_first:
# (C, H, W) format
return image[:, y1:y2, x1:x2]
else:
# (H, W, C) format
return image[y1:y2, x1:x2, :]
def resize_ndarray(
image: NumpyArray,
size: tuple[int, int],
resample: int | Image.Resampling = Image.Resampling.LANCZOS,
channel_first: bool = True,
) -> NumpyArray:
# Convert to PIL-friendly format (H, W, C)
if channel_first:
img_hwc = image.transpose((1, 2, 0))
else:
img_hwc = image
# Handle different dtypes
if img_hwc.dtype == np.float32 or img_hwc.dtype == np.float64:
# Assume normalized, scale to 0-255 for PIL
img_hwc_scaled = (img_hwc * 255).astype(np.uint8)
pil_img = Image.fromarray(img_hwc_scaled, mode="RGB")
resized = pil_img.resize(size, resample)
result = np.array(resized).astype(np.float32) / 255.0
else:
# uint8 or similar
pil_img = Image.fromarray(img_hwc.astype(np.uint8), mode="RGB")
resized = pil_img.resize(size, resample)
result = np.array(resized)
# Convert back to original format
if channel_first:
result = result.transpose((2, 0, 1))
return result
+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,
+93 -55
View File
@@ -1,13 +1,14 @@
import string
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Iterable, Sequence, Type
import numpy as np
from tokenizers import Encoding
from tokenizers import Encoding, Tokenizer
from fastembed.common.types import NumpyArray
from fastembed.common.preprocessor_utils import load_tokenizer
from fastembed.common.types import NumpyArray, Device
from fastembed.common import OnnxProvider
from fastembed.common.onnx_model import OnnxOutputContext
from fastembed.common.utils import define_cache_dir
from fastembed.common.utils import define_cache_dir, iter_batch
from fastembed.late_interaction.late_interaction_embedding_base import (
LateInteractionTextEmbeddingBase,
)
@@ -18,7 +19,7 @@ supported_colbert_models: list[DenseModelDescription] = [
DenseModelDescription(
model="colbert-ir/colbertv2.0",
dim=128,
description="Late interaction model",
description="Text embeddings, Unimodal (text), English, 512 input tokens truncation, 2023 year",
license="mit",
size_in_GB=0.44,
sources=ModelSource(hf="colbert-ir/colbertv2.0"),
@@ -27,7 +28,7 @@ supported_colbert_models: list[DenseModelDescription] = [
DenseModelDescription(
model="answerdotai/answerai-colbert-small-v1",
dim=96,
description="Text embeddings, Unimodal (text), Multilingual (~100 languages), 512 input tokens truncation, 2024 year",
description="Text embeddings, Unimodal (text), English, 512 input tokens truncation, 2024 year",
license="apache-2.0",
size_in_GB=0.13,
sources=ModelSource(hf="answerdotai/answerai-colbert-small-v1"),
@@ -43,26 +44,29 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
MASK_TOKEN = "[MASK]"
def _post_process_onnx_output(
self, output: OnnxOutputContext, is_doc: bool = True
self, output: OnnxOutputContext, is_doc: bool = True, **kwargs: Any
) -> Iterable[NumpyArray]:
if not is_doc:
return output.model_output.astype(np.float32)
for embedding in output.model_output:
yield embedding
else:
if output.input_ids is None or output.attention_mask is None:
raise ValueError(
"input_ids and attention_mask must be provided for document post-processing"
)
if output.input_ids is None or output.attention_mask is None:
raise ValueError(
"input_ids and attention_mask must be provided for document post-processing"
)
for i, token_sequence in enumerate(output.input_ids):
for j, token_id in enumerate(token_sequence): # type: ignore
if token_id in self.skip_list or token_id == self.pad_token_id:
output.attention_mask[i, j] = 0
for i, token_sequence in enumerate(output.input_ids):
for j, token_id in enumerate(token_sequence): # type: ignore
if token_id in self.skip_list or token_id == self.pad_token_id:
output.attention_mask[i, j] = 0
output.model_output *= np.expand_dims(output.attention_mask, 2)
norm = np.linalg.norm(output.model_output, ord=2, axis=2, keepdims=True)
norm_clamped = np.maximum(norm, 1e-12)
output.model_output /= norm_clamped
output.model_output *= np.expand_dims(output.attention_mask, 2).astype(np.float32)
norm = np.linalg.norm(output.model_output, ord=2, axis=2, keepdims=True)
norm_clamped = np.maximum(norm, 1e-12)
output.model_output /= norm_clamped
return output.model_output.astype(np.float32)
for embedding, attention_mask in zip(output.model_output, output.attention_mask):
yield embedding[attention_mask == 1]
def _preprocess_onnx_input(
self, onnx_input: dict[str, NumpyArray], is_doc: bool = True, **kwargs: Any
@@ -84,29 +88,46 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
)
def _tokenize_query(self, query: str) -> list[Encoding]:
assert self.tokenizer is not None
encoded = self.tokenizer.encode_batch([query])
# colbert authors recommend to pad queries with [MASK] tokens for query augmentation to improve performance
if len(encoded[0].ids) < self.MIN_QUERY_LENGTH:
prev_padding = None
if self.tokenizer.padding:
prev_padding = self.tokenizer.padding
self.tokenizer.enable_padding(
pad_token=self.MASK_TOKEN,
pad_id=self.mask_token_id,
length=self.MIN_QUERY_LENGTH,
)
encoded = self.tokenizer.encode_batch([query])
if prev_padding is None:
self.tokenizer.no_padding()
else:
self.tokenizer.enable_padding(**prev_padding)
assert self.query_tokenizer is not None
encoded = self.query_tokenizer.encode_batch([query])
return encoded
def _tokenize_documents(self, documents: list[str]) -> list[Encoding]:
encoded = self.tokenizer.encode_batch(documents) # type: ignore[union-attr]
return encoded
def token_count(
self,
texts: str | Iterable[str],
batch_size: int = 1024,
is_doc: bool = True,
include_extension: bool = False,
**kwargs: Any,
) -> int:
if not hasattr(self, "model") or self.model is None:
self.load_onnx_model() # loads the tokenizer as well
token_num = 0
texts = [texts] if isinstance(texts, str) else texts
tokenizer = self.tokenizer if is_doc else self.query_tokenizer
assert tokenizer is not None
for batch in iter_batch(texts, batch_size):
for tokens in tokenizer.encode_batch(batch):
if is_doc:
token_num += sum(tokens.attention_mask)
else:
attend_count = sum(tokens.attention_mask)
if include_extension:
token_num += max(attend_count, self.MIN_QUERY_LENGTH)
else:
token_num += attend_count
if include_extension:
token_num += len(
batch
) # add 1 for each cls.DOC_MARKER_TOKEN_ID or cls.QUERY_MARKER_TOKEN_ID
return token_num
@classmethod
def _list_supported_models(cls) -> list[DenseModelDescription]:
"""Lists the supported models.
@@ -119,14 +140,14 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
device_id: Optional[int] = None,
specific_model_path: Optional[str] = None,
device_id: int | None = None,
specific_model_path: str | None = None,
**kwargs: Any,
):
"""
@@ -138,10 +159,11 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.
providers (Optional[Sequence[OnnxProvider]], optional): The list of onnxruntime providers to use.
Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to False.
cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to Device.AUTO.
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, mutually exclusive
with `providers`. Defaults to None.
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
@@ -154,13 +176,14 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
super().__init__(model_name, cache_dir, threads, **kwargs)
self.providers = providers
self.lazy_load = lazy_load
self._extra_session_options = self._select_exposed_session_options(kwargs)
# List of device ids, that can be used for data parallel processing in workers
self.device_ids = device_ids
self.cuda = cuda
# This device_id will be used if we need to load model in current process
self.device_id: Optional[int] = None
self.device_id: int | None = None
if device_id is not None:
self.device_id = device_id
elif self.device_ids is not None:
@@ -169,16 +192,19 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
self.model_description = self._get_model_description(model_name)
self.cache_dir = str(define_cache_dir(cache_dir))
self._specific_model_path = specific_model_path
self._model_dir = self.download_model(
self.model_description,
self.cache_dir,
local_files_only=self._local_files_only,
specific_model_path=specific_model_path,
specific_model_path=self._specific_model_path,
)
self.mask_token_id: Optional[int] = None
self.pad_token_id: Optional[int] = None
self.mask_token_id: int | None = None
self.pad_token_id: int | None = None
self.skip_list: set[int] = set()
self.query_tokenizer: Tokenizer | None = None
if not self.lazy_load:
self.load_onnx_model()
@@ -190,7 +216,10 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
providers=self.providers,
cuda=self.cuda,
device_id=self.device_id,
extra_session_options=self._extra_session_options,
)
self.query_tokenizer, _ = load_tokenizer(model_dir=self._model_dir)
assert self.tokenizer is not None
self.mask_token_id = self.special_token_to_id[self.MASK_TOKEN]
self.pad_token_id = self.tokenizer.padding["pad_id"]
@@ -201,12 +230,18 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
current_max_length = self.tokenizer.truncation["max_length"]
# ensure not to overflow after adding document-marker
self.tokenizer.enable_truncation(max_length=current_max_length - 1)
self.query_tokenizer.enable_truncation(max_length=current_max_length - 1)
self.query_tokenizer.enable_padding(
pad_token=self.MASK_TOKEN,
pad_id=self.mask_token_id,
length=self.MIN_QUERY_LENGTH,
)
def embed(
self,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
"""
@@ -233,10 +268,13 @@ class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):
providers=self.providers,
cuda=self.cuda,
device_ids=self.device_ids,
local_files_only=self._local_files_only,
specific_model_path=self._specific_model_path,
extra_session_options=self._extra_session_options,
**kwargs,
)
def query_embed(self, query: Union[str, Iterable[str]], **kwargs: Any) -> Iterable[NumpyArray]:
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:
if isinstance(query, str):
query = [query]
@@ -1,4 +1,4 @@
from typing import Iterable, Optional, Union, Any
from typing import Iterable, Any
from fastembed.common.model_description import DenseModelDescription
from fastembed.common.types import NumpyArray
@@ -9,20 +9,21 @@ class LateInteractionTextEmbeddingBase(ModelManagement[DenseModelDescription]):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
cache_dir: str | None = None,
threads: int | None = None,
**kwargs: Any,
):
self.model_name = model_name
self.cache_dir = cache_dir
self.threads = threads
self._local_files_only = kwargs.pop("local_files_only", False)
self._embedding_size: int | None = None
def embed(
self,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
raise NotImplementedError()
@@ -42,7 +43,7 @@ class LateInteractionTextEmbeddingBase(ModelManagement[DenseModelDescription]):
# This is model-specific, so that different models can have specialized implementations
yield from self.embed(texts, **kwargs)
def query_embed(self, query: Union[str, Iterable[str]], **kwargs: Any) -> Iterable[NumpyArray]:
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:
"""
Embeds queries
@@ -58,3 +59,22 @@ class LateInteractionTextEmbeddingBase(ModelManagement[DenseModelDescription]):
yield from self.embed([query], **kwargs)
else:
yield from self.embed(query, **kwargs)
@classmethod
def get_embedding_size(cls, model_name: str) -> int:
"""Returns embedding size of the chosen model."""
raise NotImplementedError("Subclasses must implement this method")
@property
def embedding_size(self) -> int:
"""Returns embedding size for the current model"""
raise NotImplementedError("Subclasses must implement this method")
def token_count(
self,
texts: str | Iterable[str],
batch_size: int = 1024,
**kwargs: Any,
) -> int:
"""Returns the number of tokens in the texts."""
raise NotImplementedError("Subclasses must implement this method")
@@ -1,8 +1,8 @@
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Iterable, Sequence, Type
from dataclasses import asdict
from fastembed.common.model_description import DenseModelDescription
from fastembed.common.types import NumpyArray
from fastembed.common.types import NumpyArray, Device
from fastembed.common import OnnxProvider
from fastembed.late_interaction.colbert import Colbert
from fastembed.late_interaction.jina_colbert import JinaColbert
@@ -51,11 +51,11 @@ class LateInteractionTextEmbedding(LateInteractionTextEmbeddingBase):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
**kwargs: Any,
):
@@ -80,11 +80,45 @@ class LateInteractionTextEmbedding(LateInteractionTextEmbeddingBase):
"Please check the supported models using `LateInteractionTextEmbedding.list_supported_models()`"
)
@property
def embedding_size(self) -> int:
"""Get the embedding size of the current model"""
if self._embedding_size is None:
self._embedding_size = self.get_embedding_size(self.model_name)
return self._embedding_size
@classmethod
def get_embedding_size(cls, model_name: str) -> int:
"""Get the embedding size of the passed model
Args:
model_name (str): The name of the model to get embedding size for.
Returns:
int: The size of the embedding.
Raises:
ValueError: If the model name is not found in the supported models.
"""
descriptions = cls._list_supported_models()
embedding_size: int | None = None
for description in descriptions:
if description.model.lower() == model_name.lower():
embedding_size = description.dim
break
if embedding_size is None:
model_names = [description.model for description in descriptions]
raise ValueError(
f"Embedding size for model {model_name} was None. "
f"Available model names: {model_names}"
)
return embedding_size
def embed(
self,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
"""
@@ -104,7 +138,7 @@ class LateInteractionTextEmbedding(LateInteractionTextEmbeddingBase):
"""
yield from self.model.embed(documents, batch_size, parallel, **kwargs)
def query_embed(self, query: Union[str, Iterable[str]], **kwargs: Any) -> Iterable[NumpyArray]:
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:
"""
Embeds queries
@@ -117,3 +151,30 @@ class LateInteractionTextEmbedding(LateInteractionTextEmbeddingBase):
# This is model-specific, so that different models can have specialized implementations
yield from self.model.query_embed(query, **kwargs)
def token_count(
self,
texts: str | Iterable[str],
batch_size: int = 1024,
is_doc: bool = True,
include_extension: bool = False,
**kwargs: Any,
) -> int:
"""Returns the number of tokens in the texts.
Args:
texts (str | Iterable[str]): The list of texts to embed.
batch_size (int): Batch size for encoding
is_doc (bool): Whether the texts are documents (disable embedding a query with include_mask=True).
include_extension (bool): Turn on to count DOC / QUERY marker tokens, and [MASK] token in query mode.
Returns:
int: Sum of number of tokens in the texts.
"""
return self.model.token_count(
texts,
batch_size=batch_size,
is_doc=is_doc,
include_extension=include_extension,
**kwargs,
)
@@ -0,0 +1,83 @@
from dataclasses import asdict
from typing import Iterable, Any, Type
from fastembed.common.model_description import DenseModelDescription, ModelSource
from fastembed.common.onnx_model import OnnxOutputContext
from fastembed.common.types import NumpyArray
from fastembed.late_interaction.late_interaction_embedding_base import (
LateInteractionTextEmbeddingBase,
)
from fastembed.text.onnx_embedding import OnnxTextEmbedding
from fastembed.text.onnx_text_model import TextEmbeddingWorker
supported_token_embeddings_models = [
DenseModelDescription(
model="jinaai/jina-embeddings-v2-small-en-tokens",
dim=512,
description="Text embeddings, Unimodal (text), English, 8192 input tokens truncation,"
" Prefixes for queries/documents: not necessary, 2023 year.",
license="apache-2.0",
size_in_GB=0.12,
sources=ModelSource(hf="xenova/jina-embeddings-v2-small-en"),
model_file="onnx/model.onnx",
),
]
class TokenEmbeddingsModel(OnnxTextEmbedding, LateInteractionTextEmbeddingBase):
@classmethod
def _list_supported_models(cls) -> list[DenseModelDescription]:
"""Lists the supported models.
Returns:
list[DenseModelDescription]: A list of DenseModelDescription objects containing the model information.
"""
return supported_token_embeddings_models
@classmethod
def list_supported_models(cls) -> list[dict[str, Any]]:
"""Lists the supported models.
Returns:
list[dict[str, Any]]: A list of dictionaries containing the model information.
"""
return [asdict(model) for model in cls._list_supported_models()]
@classmethod
def _get_worker_class(cls) -> Type[TextEmbeddingWorker[NumpyArray]]:
return TokensEmbeddingWorker
def _post_process_onnx_output(
self, output: OnnxOutputContext, **kwargs: Any
) -> Iterable[NumpyArray]:
# Size: (batch_size, sequence_length, hidden_size)
embeddings = output.model_output
# Size: (batch_size, sequence_length)
assert output.attention_mask is not None
masks = output.attention_mask
# For each document we only select those embeddings that are not masked out
for i in range(embeddings.shape[0]):
yield embeddings[i, masks[i] == 1]
def embed(
self,
documents: str | Iterable[str],
batch_size: int = 256,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
yield from super().embed(documents, batch_size=batch_size, parallel=parallel, **kwargs)
class TokensEmbeddingWorker(TextEmbeddingWorker[NumpyArray]):
def init_embedding(
self, model_name: str, cache_dir: str, **kwargs: Any
) -> TokenEmbeddingsModel:
return TokenEmbeddingsModel(
model_name=model_name,
cache_dir=cache_dir,
threads=1,
**kwargs,
)
@@ -0,0 +1,532 @@
import contextlib
from typing import Any, Iterable, Type, Optional, Sequence
import json
import numpy as np
from tokenizers import Encoding
from PIL import Image
from fastembed.common import ImageInput
from fastembed.common.model_description import DenseModelDescription, ModelSource
from fastembed.common.onnx_model import OnnxOutputContext
from fastembed.common.types import NumpyArray, OnnxProvider
from fastembed.common.utils import define_cache_dir, iter_batch
from fastembed.late_interaction_multimodal.late_interaction_multimodal_embedding_base import (
LateInteractionMultimodalEmbeddingBase,
)
from fastembed.late_interaction_multimodal.onnx_multimodal_model import (
OnnxMultimodalModel,
TextEmbeddingWorker,
ImageEmbeddingWorker,
)
supported_colmodernvbert_models: list[DenseModelDescription] = [
DenseModelDescription(
model="Qdrant/colmodernvbert",
dim=128,
description="The late-interaction version of ModernVBERT, CPU friendly, English, 2025.",
license="mit",
size_in_GB=1.0,
sources=ModelSource(hf="Qdrant/colmodernvbert"),
additional_files=["processor_config.json"],
model_file="model.onnx",
),
]
class ColModernVBERT(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyArray]):
"""
The ModernVBERT/colmodernvbert model implementation. This model uses
bidirectional attention, which proves to work better for retrieval.
See: https://huggingface.co/ModernVBERT/colmodernvbert
"""
VISUAL_PROMPT_PREFIX = (
"<|begin_of_text|>User:<image>Describe the image.<end_of_utterance>\nAssistant:"
)
QUERY_AUGMENTATION_TOKEN = "<end_of_utterance>"
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
lazy_load: bool = False,
device_id: Optional[int] = None,
specific_model_path: Optional[str] = None,
**kwargs: Any,
):
"""
Args:
model_name (str): The name of the model to use.
cache_dir (str, optional): The path to the cache directory.
Can be set using the `FASTEMBED_CACHE_PATH` env variable.
Defaults to `fastembed_cache` in the system's temp directory.
threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.
providers (Optional[Sequence[OnnxProvider]], optional): The list of onnxruntime providers to use.
Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to False.
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
Raises:
ValueError: If the model_name is not in the format <org>/<model> e.g. BAAI/bge-base-en.
"""
super().__init__(model_name, cache_dir, threads, **kwargs)
self.providers = providers
self.lazy_load = lazy_load
self._extra_session_options = self._select_exposed_session_options(kwargs)
# List of device ids, that can be used for data parallel processing in workers
self.device_ids = device_ids
self.cuda = cuda
# This device_id will be used if we need to load model in current process
self.device_id: Optional[int] = None
if device_id is not None:
self.device_id = device_id
elif self.device_ids is not None:
self.device_id = self.device_ids[0]
self.model_description = self._get_model_description(model_name)
self.cache_dir = str(define_cache_dir(cache_dir))
self._specific_model_path = specific_model_path
self._model_dir = self.download_model(
self.model_description,
self.cache_dir,
local_files_only=self._local_files_only,
specific_model_path=self._specific_model_path,
)
self.mask_token_id = None
self.pad_token_id = None
self.image_seq_len: Optional[int] = None
self.max_image_size: Optional[int] = None
self.image_size: Optional[int] = None
if not self.lazy_load:
self.load_onnx_model()
@classmethod
def _list_supported_models(cls) -> list[DenseModelDescription]:
"""Lists the supported models.
Returns:
list[DenseModelDescription]: A list of DenseModelDescription objects containing the model information.
"""
return supported_colmodernvbert_models
def load_onnx_model(self) -> None:
self._load_onnx_model(
model_dir=self._model_dir,
model_file=self.model_description.model_file,
threads=self.threads,
providers=self.providers,
cuda=self.cuda,
device_id=self.device_id,
extra_session_options=self._extra_session_options,
)
# Load image processing configuration
processor_config_path = self._model_dir / "processor_config.json"
with open(processor_config_path) as f:
processor_config = json.load(f)
self.image_seq_len = processor_config.get("image_seq_len", 64)
preprocessor_config_path = self._model_dir / "preprocessor_config.json"
with open(preprocessor_config_path) as f:
preprocessor_config = json.load(f)
self.max_image_size = preprocessor_config.get("max_image_size", {}).get(
"longest_edge", 512
)
# Load model configuration
config_path = self._model_dir / "config.json"
with open(config_path) as f:
model_config = json.load(f)
vision_config = model_config.get("vision_config", {})
self.image_size = vision_config.get("image_size", 512)
def _preprocess_onnx_text_input(
self, onnx_input: dict[str, NumpyArray], **kwargs: Any
) -> dict[str, NumpyArray]:
"""
Post-process the ONNX model output to convert it into a usable format.
Args:
output (OnnxOutputContext): The raw output from the ONNX model.
Returns:
Iterable[NumpyArray]: Post-processed output as NumPy arrays.
"""
batch_size, seq_length = onnx_input["input_ids"].shape
empty_image_placeholder: NumpyArray = np.zeros(
(batch_size, seq_length, 3, self.image_size, self.image_size),
dtype=np.float32, # type: ignore[type-var,arg-type,assignment]
)
onnx_input["pixel_values"] = empty_image_placeholder
return onnx_input
def _post_process_onnx_text_output(
self,
output: OnnxOutputContext,
) -> Iterable[NumpyArray]:
"""
Post-process the ONNX model output to convert it into a usable format.
Args:
output (OnnxOutputContext): The raw output from the ONNX model.
Returns:
Iterable[NumpyArray]: Post-processed output as NumPy arrays.
"""
return output.model_output
def tokenize(self, documents: list[str], **kwargs: Any) -> list[Encoding]:
# Add query augmentation tokens (matching process_queries logic from colpali-engine)
augmented_queries = [doc + self.QUERY_AUGMENTATION_TOKEN * 10 for doc in documents]
encoded = self.tokenizer.encode_batch(augmented_queries) # type: ignore[union-attr]
return encoded
def token_count(
self,
texts: str | Iterable[str],
batch_size: int = 1024,
include_extension: bool = False,
**kwargs: Any,
) -> int:
if not hasattr(self, "model") or self.model is None:
self.load_onnx_model() # loads the tokenizer as well
token_num = 0
texts = [texts] if isinstance(texts, str) else texts
assert self.tokenizer is not None
tokenize_func = self.tokenize if include_extension else self.tokenizer.encode_batch
for batch in iter_batch(texts, batch_size):
token_num += sum([sum(encoding.attention_mask) for encoding in tokenize_func(batch)])
return token_num
def onnx_embed_image(self, images: list[ImageInput], **kwargs: Any) -> OnnxOutputContext:
with contextlib.ExitStack() as stack:
image_files = [
stack.enter_context(Image.open(image))
if not isinstance(image, Image.Image)
else image
for image in images
]
assert self.processor is not None, "Processor is not initialized"
processed = self.processor(image_files)
encoded, attention_mask, metadata = self._process_nested_patches(processed) # type: ignore[arg-type]
onnx_input = {"pixel_values": encoded, "attention_mask": attention_mask}
onnx_input = self._preprocess_onnx_image_input(onnx_input, **kwargs)
model_output = self.model.run(None, onnx_input) # type: ignore[union-attr]
return OnnxOutputContext(
model_output=model_output[0],
attention_mask=attention_mask, # type: ignore[arg-type]
metadata=metadata,
)
@staticmethod
def _process_nested_patches(
processed: list[list[NumpyArray]],
) -> tuple[NumpyArray, NumpyArray, dict[str, Any]]:
"""
Process nested image patches (from ImageSplitter).
Args:
processed: List of patch lists, one per image [[img1_patches], [img2_patches], ...]
Returns:
tuple: (encoded array, attention_mask, metadata)
- encoded: (batch_size, max_patches, C, H, W)
- attention_mask: (batch_size, max_patches) with 1 for real patches, 0 for padding
- metadata: Dict with 'patch_counts' key
"""
patch_counts = [len(patches) for patches in processed]
max_patches = max(patch_counts)
# Get dimensions from first patch
channels, height, width = processed[0][0].shape
batch_size = len(processed)
# Create padded array
encoded = np.zeros(
(batch_size, max_patches, channels, height, width), dtype=processed[0][0].dtype
)
# Create attention mask (1 for real patches, 0 for padding)
attention_mask = np.zeros((batch_size, max_patches), dtype=np.int64)
# Fill in patches and attention mask
for i, patches in enumerate(processed):
for j, patch in enumerate(patches):
encoded[i, j] = patch
attention_mask[i, j] = 1
metadata = {"patch_counts": patch_counts}
return encoded, attention_mask, metadata # type: ignore[return-value]
def _preprocess_onnx_image_input(
self, onnx_input: dict[str, np.ndarray], **kwargs: Any
) -> dict[str, NumpyArray]:
"""
Add text input placeholders for image data, following Idefics3 processing logic.
Constructs input_ids dynamically based on the actual number of image patches,
using the same token expansion logic as Idefics3Processor.
Args:
onnx_input: Dict with 'pixel_values' (batch, num_patches, C, H, W)
and 'attention_mask' (batch, num_patches) indicating real patches
**kwargs: Additional arguments
Returns:
Updated onnx_input with 'input_ids' and updated 'attention_mask' for token sequence
"""
# The attention_mask in onnx_input has a shape of (batch_size, num_patches),
# and should be used to create an attention mask matching the input_ids shape.
patch_attention_mask = onnx_input["attention_mask"]
pixel_values = onnx_input["pixel_values"]
batch_size = pixel_values.shape[0]
batch_input_ids = []
# Build input_ids for each image based on its actual patch count
for i in range(batch_size):
# Count real patches (non-padded) from attention mask
patch_count = int(np.sum(patch_attention_mask[i]))
# Compute rows/cols from patch count
rows, cols = self._compute_rows_cols_from_patches(patch_count)
# Build input_ids for this image
input_ids = self._build_input_ids_for_image(rows, cols)
batch_input_ids.append(input_ids)
# Pad sequences to max length in batch
max_len = max(len(ids) for ids in batch_input_ids)
# Get padding config from tokenizer
padding_direction = self.tokenizer.padding["direction"] # type: ignore[index,union-attr]
pad_token_id = self.tokenizer.padding["pad_id"] # type: ignore[index,union-attr]
# Initialize with pad token
padded_input_ids = np.full((batch_size, max_len), pad_token_id, dtype=np.int64)
attention_mask = np.zeros((batch_size, max_len), dtype=np.int64)
for i, input_ids in enumerate(batch_input_ids):
seq_len = len(input_ids)
if padding_direction == "left":
# Left padding: place tokens at the END of the array
start_idx = max_len - seq_len
padded_input_ids[i, start_idx:] = input_ids
attention_mask[i, start_idx:] = 1
else:
# Right padding: place tokens at the START of the array
padded_input_ids[i, :seq_len] = input_ids
attention_mask[i, :seq_len] = 1
onnx_input["input_ids"] = padded_input_ids
# Update attention_mask with token-level data
onnx_input["attention_mask"] = attention_mask
return onnx_input
@staticmethod
def _compute_rows_cols_from_patches(patch_count: int) -> tuple[int, int]:
if patch_count <= 1:
return 0, 0
# Subtract 1 for the global image
grid_patches = patch_count - 1
# Find rows and cols (assume square or near-square grid)
rows = int(grid_patches**0.5)
cols = grid_patches // rows
# Verify the calculation
if rows * cols + 1 != patch_count:
# Handle non-square grids
for r in range(1, grid_patches + 1):
if grid_patches % r == 0:
c = grid_patches // r
if r * c + 1 == patch_count:
return r, c
# Fallback: treat as unsplit
return 0, 0
return rows, cols
def _create_single_image_prompt_string(self) -> str:
return (
"<fake_token_around_image>"
+ "<global-img>"
+ "<image>" * self.image_seq_len # type: ignore[operator]
+ "<fake_token_around_image>"
)
def _create_split_image_prompt_string(self, rows: int, cols: int) -> str:
text_split_images = ""
# Add tokens for each patch in the grid
for n_h in range(rows):
for n_w in range(cols):
text_split_images += (
"<fake_token_around_image>"
+ f"<row_{n_h + 1}_col_{n_w + 1}>"
+ "<image>" * self.image_seq_len # type: ignore[operator]
)
text_split_images += "\n"
# Add global image at the end
text_split_images += (
"\n<fake_token_around_image>"
+ "<global-img>"
+ "<image>" * self.image_seq_len # type: ignore[operator]
+ "<fake_token_around_image>"
)
return text_split_images
def _build_input_ids_for_image(self, rows: int, cols: int) -> np.ndarray:
# Create the appropriate image prompt string
if rows == 0 and cols == 0:
image_prompt_tokens = self._create_single_image_prompt_string()
else:
image_prompt_tokens = self._create_split_image_prompt_string(rows, cols)
# Replace <image> in visual prompt with expanded tokens
# The visual prompt is: "<|begin_of_text|>User:<image>Describe the image.<end_of_utterance>\nAssistant:"
expanded_prompt = self.VISUAL_PROMPT_PREFIX.replace("<image>", image_prompt_tokens)
# Tokenize the complete prompt
encoded = self.tokenizer.encode(expanded_prompt) # type: ignore[union-attr]
# Convert to numpy array
return np.array(encoded.ids, dtype=np.int64)
def _post_process_onnx_image_output(
self,
output: OnnxOutputContext,
) -> Iterable[NumpyArray]:
"""
Post-process the ONNX model output to convert it into a usable format.
Args:
output (OnnxOutputContext): The raw output from the ONNX model.
Returns:
Iterable[NumpyArray]: Post-processed output as NumPy arrays.
"""
assert self.model_description.dim is not None, "Model dim is not defined"
return output.model_output.reshape(
output.model_output.shape[0], -1, self.model_description.dim
)
def embed_text(
self,
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
"""
Encode a list of documents into list of embeddings.
Args:
documents: Iterator of documents or single document to embed
batch_size: Batch size for encoding -- higher values will use more memory, but be faster
parallel:
If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.
If 0, use all available cores.
If None, don't use data-parallel processing, use default onnxruntime threading instead.
Returns:
List of embeddings, one per document
"""
yield from self._embed_documents(
model_name=self.model_name,
cache_dir=str(self.cache_dir),
documents=documents,
batch_size=batch_size,
parallel=parallel,
providers=self.providers,
cuda=self.cuda,
device_ids=self.device_ids,
local_files_only=self._local_files_only,
specific_model_path=self._specific_model_path,
extra_session_options=self._extra_session_options,
**kwargs,
)
def embed_image(
self,
images: ImageInput | Iterable[ImageInput],
batch_size: int = 16,
parallel: Optional[int] = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
"""
Encode a list of images into list of embeddings.
Args:
images: Iterator of image paths or single image path to embed
batch_size: Batch size for encoding -- higher values will use more memory, but be faster
parallel:
If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.
If 0, use all available cores.
If None, don't use data-parallel processing, use default onnxruntime threading instead.
Returns:
List of embeddings, one per document
"""
yield from self._embed_images(
model_name=self.model_name,
cache_dir=str(self.cache_dir),
images=images,
batch_size=batch_size,
parallel=parallel,
providers=self.providers,
cuda=self.cuda,
device_ids=self.device_ids,
local_files_only=self._local_files_only,
specific_model_path=self._specific_model_path,
extra_session_options=self._extra_session_options,
**kwargs,
)
@classmethod
def _get_text_worker_class(cls) -> Type[TextEmbeddingWorker[NumpyArray]]:
return ColModernVBERTTextEmbeddingWorker
@classmethod
def _get_image_worker_class(cls) -> Type[ImageEmbeddingWorker[NumpyArray]]:
return ColModernVBERTImageEmbeddingWorker
class ColModernVBERTTextEmbeddingWorker(TextEmbeddingWorker[NumpyArray]):
def init_embedding(self, model_name: str, cache_dir: str, **kwargs: Any) -> ColModernVBERT:
return ColModernVBERT(
model_name=model_name,
cache_dir=cache_dir,
threads=1,
**kwargs,
)
class ColModernVBERTImageEmbeddingWorker(ImageEmbeddingWorker[NumpyArray]):
def init_embedding(self, model_name: str, cache_dir: str, **kwargs: Any) -> ColModernVBERT:
return ColModernVBERT(
model_name=model_name,
cache_dir=cache_dir,
threads=1,
**kwargs,
)
@@ -1,12 +1,12 @@
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Iterable, Sequence, Type
import numpy as np
from tokenizers import Encoding
from fastembed.common import OnnxProvider, ImageInput
from fastembed.common.onnx_model import OnnxOutputContext
from fastembed.common.types import NumpyArray
from fastembed.common.utils import define_cache_dir
from fastembed.common.types import NumpyArray, Device
from fastembed.common.utils import define_cache_dir, iter_batch
from fastembed.late_interaction_multimodal.late_interaction_multimodal_embedding_base import (
LateInteractionMultimodalEmbeddingBase,
)
@@ -46,14 +46,14 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
device_id: Optional[int] = None,
specific_model_path: Optional[str] = None,
device_id: int | None = None,
specific_model_path: str | None = None,
**kwargs: Any,
):
"""
@@ -65,10 +65,11 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.
providers (Optional[Sequence[OnnxProvider]], optional): The list of onnxruntime providers to use.
Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to False.
cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to Device.AUTO.
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, mutually exclusive
with `providers`. Defaults to None.
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
@@ -80,13 +81,14 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
super().__init__(model_name, cache_dir, threads, **kwargs)
self.providers = providers
self.lazy_load = lazy_load
self._extra_session_options = self._select_exposed_session_options(kwargs)
# List of device ids, that can be used for data parallel processing in workers
self.device_ids = device_ids
self.cuda = cuda
# This device_id will be used if we need to load model in current process
self.device_id: Optional[int] = None
self.device_id: int | None = None
if device_id is not None:
self.device_id = device_id
elif self.device_ids is not None:
@@ -95,11 +97,12 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
self.model_description = self._get_model_description(model_name)
self.cache_dir = str(define_cache_dir(cache_dir))
self._specific_model_path = specific_model_path
self._model_dir = self.download_model(
self.model_description,
self.cache_dir,
local_files_only=self._local_files_only,
specific_model_path=specific_model_path,
specific_model_path=self._specific_model_path,
)
self.mask_token_id = None
self.pad_token_id = None
@@ -124,6 +127,7 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
providers=self.providers,
cuda=self.cuda,
device_id=self.device_id,
extra_session_options=self._extra_session_options,
)
def _post_process_onnx_image_output(
@@ -142,7 +146,7 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
assert self.model_description.dim is not None, "Model dim is not defined"
return output.model_output.reshape(
output.model_output.shape[0], -1, self.model_description.dim
).astype(np.float32)
)
def _post_process_onnx_text_output(
self,
@@ -157,7 +161,7 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
Returns:
Iterable[NumpyArray]: Post-processed output as NumPy arrays.
"""
return output.model_output.astype(np.float32)
return output.model_output
def tokenize(self, documents: list[str], **kwargs: Any) -> list[Encoding]:
texts_query: list[str] = []
@@ -169,12 +173,29 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
encoded = self.tokenizer.encode_batch(texts_query) # type: ignore[union-attr]
return encoded
def token_count(
self,
texts: str | Iterable[str],
batch_size: int = 1024,
include_extension: bool = False,
**kwargs: Any,
) -> int:
if not hasattr(self, "model") or self.model is None:
self.load_onnx_model() # loads the tokenizer as well
token_num = 0
texts = [texts] if isinstance(texts, str) else texts
assert self.tokenizer is not None
tokenize_func = self.tokenize if include_extension else self.tokenizer.encode_batch
for batch in iter_batch(texts, batch_size):
token_num += sum([sum(encoding.attention_mask) for encoding in tokenize_func(batch)])
return token_num
def _preprocess_onnx_text_input(
self, onnx_input: dict[str, NumpyArray], **kwargs: Any
) -> dict[str, NumpyArray]:
onnx_input["input_ids"] = np.array(
[
self.QUERY_MARKER_TOKEN_ID + input_ids[2:].tolist()
self.QUERY_MARKER_TOKEN_ID + input_ids[2:].tolist() # type: ignore[index]
for input_ids in onnx_input["input_ids"]
]
)
@@ -207,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]:
"""
@@ -235,14 +256,17 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
providers=self.providers,
cuda=self.cuda,
device_ids=self.device_ids,
local_files_only=self._local_files_only,
specific_model_path=self._specific_model_path,
extra_session_options=self._extra_session_options,
**kwargs,
)
def embed_image(
self,
images: Union[ImageInput, Iterable[ImageInput]],
images: ImageInput | Iterable[ImageInput],
batch_size: int = 16,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
"""
@@ -268,6 +292,9 @@ class ColPali(LateInteractionMultimodalEmbeddingBase, OnnxMultimodalModel[NumpyA
providers=self.providers,
cuda=self.cuda,
device_ids=self.device_ids,
local_files_only=self._local_files_only,
specific_model_path=self._specific_model_path,
extra_session_options=self._extra_session_options,
**kwargs,
)
@@ -1,9 +1,10 @@
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Iterable, Sequence, Type
from dataclasses import asdict
from fastembed.common import OnnxProvider, ImageInput
from fastembed.common.types import NumpyArray
from fastembed.common.types import NumpyArray, Device
from fastembed.late_interaction_multimodal.colpali import ColPali
from fastembed.late_interaction_multimodal.colmodernvbert import ColModernVBERT
from fastembed.late_interaction_multimodal.late_interaction_multimodal_embedding_base import (
LateInteractionMultimodalEmbeddingBase,
@@ -12,7 +13,10 @@ from fastembed.common.model_description import DenseModelDescription
class LateInteractionMultimodalEmbedding(LateInteractionMultimodalEmbeddingBase):
EMBEDDINGS_REGISTRY: list[Type[LateInteractionMultimodalEmbeddingBase]] = [ColPali]
EMBEDDINGS_REGISTRY: list[Type[LateInteractionMultimodalEmbeddingBase]] = [
ColPali,
ColModernVBERT,
]
@classmethod
def list_supported_models(cls) -> list[dict[str, Any]]:
@@ -54,11 +58,11 @@ class LateInteractionMultimodalEmbedding(LateInteractionMultimodalEmbeddingBase)
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
**kwargs: Any,
):
@@ -83,11 +87,45 @@ class LateInteractionMultimodalEmbedding(LateInteractionMultimodalEmbeddingBase)
"Please check the supported models using `LateInteractionMultimodalEmbedding.list_supported_models()`"
)
@property
def embedding_size(self) -> int:
"""Get the embedding size of the current model"""
if self._embedding_size is None:
self._embedding_size = self.get_embedding_size(self.model_name)
return self._embedding_size
@classmethod
def get_embedding_size(cls, model_name: str) -> int:
"""Get the embedding size of the passed model
Args:
model_name (str): The name of the model to get embedding size for.
Returns:
int: The size of the embedding.
Raises:
ValueError: If the model name is not found in the supported models.
"""
descriptions = cls._list_supported_models()
embedding_size: int | None = None
for description in descriptions:
if description.model.lower() == model_name.lower():
embedding_size = description.dim
break
if embedding_size is None:
model_names = [description.model for description in descriptions]
raise ValueError(
f"Embedding size for model {model_name} was None. "
f"Available model names: {model_names}"
)
return embedding_size
def embed_text(
self,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
"""
@@ -108,9 +146,9 @@ class LateInteractionMultimodalEmbedding(LateInteractionMultimodalEmbeddingBase)
def embed_image(
self,
images: Union[ImageInput, Iterable[ImageInput]],
images: ImageInput | Iterable[ImageInput],
batch_size: int = 16,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
"""
@@ -128,3 +166,24 @@ class LateInteractionMultimodalEmbedding(LateInteractionMultimodalEmbeddingBase)
List of embeddings, one per image
"""
yield from self.model.embed_image(images, batch_size, parallel, **kwargs)
def token_count(
self,
texts: str | Iterable[str],
batch_size: int = 1024,
include_extension: bool = False,
**kwargs: Any,
) -> int:
"""Returns the number of tokens in the texts.
Args:
texts (str | Iterable[str]): The list of texts to embed.
batch_size (int): Batch size for encoding
include_extension (bool): Whether to include tokens added by preprocessing
Returns:
int: Sum of number of tokens in the texts.
"""
return self.model.token_count(
texts, batch_size=batch_size, include_extension=include_extension, **kwargs
)
@@ -1,4 +1,4 @@
from typing import Iterable, Optional, Union, Any
from typing import Iterable, Any
from fastembed.common import ImageInput
@@ -11,20 +11,21 @@ class LateInteractionMultimodalEmbeddingBase(ModelManagement[DenseModelDescripti
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
cache_dir: str | None = None,
threads: int | None = None,
**kwargs: Any,
):
self.model_name = model_name
self.cache_dir = cache_dir
self.threads = threads
self._local_files_only = kwargs.pop("local_files_only", False)
self._embedding_size: int | None = None
def embed_text(
self,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
"""
@@ -46,9 +47,9 @@ class LateInteractionMultimodalEmbeddingBase(ModelManagement[DenseModelDescripti
def embed_image(
self,
images: Union[ImageInput, Iterable[ImageInput]],
images: ImageInput | Iterable[ImageInput],
batch_size: int = 16,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
"""
@@ -65,3 +66,21 @@ class LateInteractionMultimodalEmbeddingBase(ModelManagement[DenseModelDescripti
List of embeddings, one per image
"""
raise NotImplementedError()
@classmethod
def get_embedding_size(cls, model_name: str) -> int:
"""Returns embedding size of the chosen model."""
raise NotImplementedError("Subclasses must implement this method")
@property
def embedding_size(self) -> int:
"""Returns embedding size for the current model"""
raise NotImplementedError("Subclasses must implement this method")
def token_count(
self,
texts: str | Iterable[str],
**kwargs: Any,
) -> int:
"""Returns the number of tokens in the texts."""
raise NotImplementedError("Subclasses must implement this method")
@@ -2,7 +2,7 @@ import contextlib
import os
from multiprocessing import get_all_start_methods
from pathlib import Path
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Iterable, Sequence, Type
import numpy as np
from PIL import Image
@@ -11,19 +11,19 @@ from tokenizers import Encoding, Tokenizer
from fastembed.common import OnnxProvider, ImageInput
from fastembed.common.onnx_model import EmbeddingWorker, OnnxModel, OnnxOutputContext, T
from fastembed.common.preprocessor_utils import load_tokenizer, load_preprocessor
from fastembed.common.types import NumpyArray
from fastembed.common.types import NumpyArray, Device
from fastembed.common.utils import iter_batch
from fastembed.image.transform.operators import Compose
from fastembed.parallel_processor import ParallelWorkerPool
class OnnxMultimodalModel(OnnxModel[T]):
ONNX_OUTPUT_NAMES: Optional[list[str]] = None
ONNX_OUTPUT_NAMES: list[str] | None = None
def __init__(self) -> None:
super().__init__()
self.tokenizer: Optional[Tokenizer] = None
self.processor: Optional[Compose] = None
self.tokenizer: Tokenizer | None = None
self.processor: Compose | None = None
self.special_token_to_id: dict[str, int] = {}
def _preprocess_onnx_text_input(
@@ -60,10 +60,11 @@ class OnnxMultimodalModel(OnnxModel[T]):
self,
model_dir: Path,
model_file: str,
threads: Optional[int],
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_id: Optional[int] = None,
threads: int | None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_id: int | None = None,
extra_session_options: dict[str, Any] | None = None,
) -> None:
super()._load_onnx_model(
model_dir=model_dir,
@@ -72,6 +73,7 @@ class OnnxMultimodalModel(OnnxModel[T]):
providers=providers,
cuda=cuda,
device_id=device_id,
extra_session_options=extra_session_options,
)
self.tokenizer, self.special_token_to_id = load_tokenizer(model_dir=model_dir)
assert self.tokenizer is not None
@@ -114,12 +116,15 @@ class OnnxMultimodalModel(OnnxModel[T]):
self,
model_name: str,
cache_dir: str,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
parallel: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
local_files_only: bool = False,
specific_model_path: str | None = None,
extra_session_options: dict[str, Any] | None = None,
**kwargs: Any,
) -> Iterable[T]:
is_small = False
@@ -146,9 +151,14 @@ class OnnxMultimodalModel(OnnxModel[T]):
"model_name": model_name,
"cache_dir": cache_dir,
"providers": providers,
"local_files_only": local_files_only,
"specific_model_path": specific_model_path,
**kwargs,
}
if extra_session_options is not None:
params.update(extra_session_options)
pool = ParallelWorkerPool(
num_workers=parallel or 1,
worker=self._get_text_worker_class(),
@@ -160,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"
@@ -177,12 +189,15 @@ class OnnxMultimodalModel(OnnxModel[T]):
self,
model_name: str,
cache_dir: str,
images: Union[Iterable[ImageInput], ImageInput],
images: Iterable[ImageInput] | ImageInput,
batch_size: int = 256,
parallel: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
parallel: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
local_files_only: bool = False,
specific_model_path: str | None = None,
extra_session_options: dict[str, Any] | None = None,
**kwargs: Any,
) -> Iterable[T]:
is_small = False
@@ -209,9 +224,14 @@ class OnnxMultimodalModel(OnnxModel[T]):
"model_name": model_name,
"cache_dir": cache_dir,
"providers": providers,
"local_files_only": local_files_only,
"specific_model_path": specific_model_path,
**kwargs,
}
if extra_session_options is not None:
params.update(extra_session_options)
pool = ParallelWorkerPool(
num_workers=parallel or 1,
worker=self._get_image_worker_class(),
+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:
+3
View File
@@ -0,0 +1,3 @@
from fastembed.postprocess.muvera import Muvera
__all__ = ["Muvera"]
+362
View File
@@ -0,0 +1,362 @@
import numpy as np
from fastembed.common.types import NumpyArray
from fastembed.late_interaction.late_interaction_embedding_base import (
LateInteractionTextEmbeddingBase,
)
from fastembed.late_interaction_multimodal.late_interaction_multimodal_embedding_base import (
LateInteractionMultimodalEmbeddingBase,
)
MultiVectorModel = LateInteractionTextEmbeddingBase | LateInteractionMultimodalEmbeddingBase
MAX_HAMMING_DISTANCE = 65 # 64 bits + 1
POPCOUNT_LUT = np.array([bin(x).count("1") for x in range(256)], dtype=np.uint8)
def hamming_distance_matrix(ids: np.ndarray) -> np.ndarray:
"""Compute full Hamming distance matrix
Args:
ids: shape (n,) - array of ids, only size of the array matters
Return:
np.ndarray (n, n) - hamming distance matrix
"""
n = len(ids)
xor_vals = np.bitwise_xor(ids[:, None], ids[None, :]) # (n, n) uint64
bytes_view = xor_vals.view(np.uint8).reshape(n, n, 8) # (n, n, 8)
return POPCOUNT_LUT[bytes_view].sum(axis=2)
class SimHashProjection:
"""
SimHash projection component for MUVERA clustering.
This class implements locality-sensitive hashing using random hyperplanes
to partition the vector space into 2^k_sim clusters. Each vector is assigned
to a cluster based on which side of k_sim random hyperplanes it falls on.
Attributes:
k_sim (int): Number of SimHash functions (hyperplanes)
dim (int): Dimensionality of input vectors
simhash_vectors (np.ndarray): Random hyperplane normal vectors of shape (dim, k_sim)
"""
def __init__(self, k_sim: int, dim: int, random_generator: np.random.Generator):
"""
Initialize SimHash projection with random hyperplanes.
Args:
k_sim (int): Number of SimHash functions, determines 2^k_sim clusters
dim (int): Dimensionality of input vectors
random_generator (np.random.Generator): Random number generator for reproducibility
"""
self.k_sim = k_sim
self.dim = dim
# Generate k_sim random hyperplanes (normal vectors) from standard normal distribution
self.simhash_vectors = random_generator.normal(size=(dim, k_sim))
def get_cluster_ids(self, vectors: np.ndarray) -> np.ndarray:
"""
Compute the cluster IDs for a given vector using SimHash.
The cluster ID is determined by computing the dot product of the vector
with each hyperplane normal vector, taking the sign, and interpreting
the resulting binary string as an integer.
Args:
vectors (np.ndarray): Input vectors of shape (n, dim,)
Returns:
np.ndarray: Cluster IDs in range [0, 2^k_sim - 1]
Raises:
AssertionError: If a vector shape doesn't match expected dimensionality
"""
dot_product = (
vectors @ self.simhash_vectors
) # (token_num, dim) x (dim, k_sim) -> (token_num, k_sim)
cluster_ids = (dot_product > 0) @ (1 << np.arange(self.k_sim))
return cluster_ids
class Muvera:
"""
MUVERA (Multi-Vector Retrieval Architecture) algorithm implementation.
This class creates Fixed Dimensional Encodings (FDEs) from variable-length
sequences of vectors by using SimHash clustering and random projections.
The process involves:
1. Clustering vectors using multiple SimHash projections
2. Computing cluster centers (with different strategies for docs vs queries)
3. Applying random projections for dimensionality reduction
4. Concatenating results from all projections
Attributes:
k_sim (int): Number of SimHash functions per projection
dim (int): Input vector dimensionality
dim_proj (int): Output dimensionality after random projection
r_reps (int): Number of random projection repetitions
random_seed (int): Random seed for consistent random matrix generation
simhash_projections (List[SimHashProjection]): SimHash instances for clustering
dim_reduction_projections (np.ndarray): Random projection matrices of shape (R_reps, d, d_proj)
"""
def __init__(
self,
dim: int,
k_sim: int = 5,
dim_proj: int = 16,
r_reps: int = 20,
random_seed: int = 42,
):
"""
Initialize MUVERA algorithm with specified parameters.
Args:
dim (int): Dimensionality of individual input vectors
k_sim (int, optional): Number of SimHash functions (creates 2^k_sim clusters).
Defaults to 5.
dim_proj (int, optional): Dimensionality after random projection (must be <= dim).
Defaults to 16.
r_reps (int, optional): Number of random projection repetitions for robustness.
Defaults to 20.
random_seed (int, optional): Seed for random number generator to ensure
reproducible results. Defaults to 42.
Raises:
ValueError: If dim_proj > dim (cannot project to higher dimensionality)
"""
if dim_proj > dim:
raise ValueError(
f"Cannot project to a higher dimensionality (dim_proj={dim_proj} > dim={dim})"
)
self.k_sim = k_sim
self.dim = dim
self.dim_proj = dim_proj
self.r_reps = r_reps
# Create r_reps independent SimHash projections for robustness
generator = np.random.default_rng(random_seed)
self.simhash_projections = [
SimHashProjection(k_sim=self.k_sim, dim=self.dim, random_generator=generator)
for _ in range(r_reps)
]
# Random projection matrices with entries from {-1, +1} for each repetition
self.dim_reduction_projections = generator.choice([-1, 1], size=(r_reps, dim, dim_proj))
@classmethod
def from_multivector_model(
cls,
model: MultiVectorModel,
k_sim: int = 5,
dim_proj: int = 16,
r_reps: int = 20, # noqa[naming]
random_seed: int = 42,
) -> "Muvera":
"""
Create a Muvera instance from a multi-vector embedding model.
This class method provides a convenient way to initialize a MUVERA
that is compatible with a given multi-vector model by automatically extracting
the embedding dimensionality from the model.
Args:
model (MultiVectorModel): A late interaction text or multimodal embedding model
that provides multi-vector embeddings. Must have an
`embedding_size` attribute specifying the dimensionality
of individual vectors.
k_sim (int, optional): Number of SimHash functions (creates 2^k_sim clusters).
Defaults to 5.
dim_proj (int, optional): Dimensionality after random projection (must be <= model's
embedding_size). Defaults to 16.
r_reps (int, optional): Number of random projection repetitions for robustness.
Defaults to 20.
random_seed (int, optional): Seed for random number generator to ensure
reproducible results. Defaults to 42.
Returns:
Muvera: A configured MUVERA instance ready to process embeddings from the given model.
Raises:
ValueError: If dim_proj > model.embedding_size (cannot project to higher dimensionality)
Example:
>>> from fastembed import LateInteractionTextEmbedding
>>> model = LateInteractionTextEmbedding(model_name="colbert-ir/colbertv2.0")
>>> muvera = Muvera.from_multivector_model(
... model=model,
... k_sim=6,
... dim_proj=32
... )
>>> # Now use postprocessor with embeddings from the model
>>> embeddings = np.array(list(model.embed(["sample text"])))
>>> fde = muvera.process_document(embeddings[0])
"""
return cls(
dim=model.embedding_size,
k_sim=k_sim,
dim_proj=dim_proj,
r_reps=r_reps,
random_seed=random_seed,
)
def _get_output_dimension(self) -> int:
"""
Get the output dimension of the MUVERA algorithm.
Returns:
int: Output dimension (r_reps * num_partitions * dim_proj) where b = 2^k_sim
"""
num_partitions = 2**self.k_sim
return self.r_reps * num_partitions * self.dim_proj
@property
def embedding_size(self) -> int:
return self._get_output_dimension()
def process_document(self, vectors: NumpyArray) -> NumpyArray:
"""
Encode a document's vectors into a Fixed Dimensional Encoding (FDE).
Uses document-specific settings: normalizes cluster centers by vector count
and fills empty clusters using Hamming distance-based selection.
Args:
vectors (NumpyArray): Document vectors of shape (n_tokens, dim)
Returns:
NumpyArray: Fixed dimensional encodings of shape (r_reps * b * dim_proj,)
"""
return self.process(vectors, fill_empty_clusters=True, normalize_by_count=True)
def process_query(self, vectors: NumpyArray) -> NumpyArray:
"""
Encode a query's vectors into a Fixed Dimensional Encoding (FDE).
Uses query-specific settings: no normalization by count and no empty
cluster filling to preserve query vector magnitudes.
Args:
vectors (NumpyArray]): Query vectors of shape (n_tokens, dim)
Returns:
NumpyArray: Fixed dimensional encoding of shape (r_reps * b * dim_proj,)
"""
return self.process(vectors, fill_empty_clusters=False, normalize_by_count=False)
def process(
self,
vectors: NumpyArray,
fill_empty_clusters: bool = True,
normalize_by_count: bool = True,
) -> NumpyArray:
"""
Core encoding method that transforms variable-length vector sequences into FDEs.
The encoding process:
1. For each of r_reps random projections:
a. Assign vectors to clusters using SimHash
b. Compute cluster centers (sum of vectors in each cluster)
c. Optionally normalize by cluster size
d. Fill empty clusters using Hamming distance if requested
e. Apply random projection for dimensionality reduction
f. Flatten cluster centers into a vector
2. Concatenate all projection results
Args:
vectors (np.ndarray): Input vectors of shape (n_vectors, dim)
fill_empty_clusters (bool): Whether to fill empty clusters using nearest
vectors based on Hamming distance of cluster IDs
normalize_by_count (bool): Whether to normalize cluster centers by the
number of vectors assigned to each cluster
Returns:
np.ndarray: Fixed dimensional encoding of shape (r_reps * b * dim_proj)
where B = 2^k_sim is the number of clusters
Raises:
AssertionError: If input vectors don't have expected dimensionality
"""
assert (
vectors.shape[1] == self.dim
), f"Expected vectors of shape (n, {self.dim}), got {vectors.shape}"
# Store results from each random projection
output_vectors = []
# num of space partitions in SimHash
num_partitions = 2**self.k_sim
cluster_center_ids = np.arange(num_partitions)
precomputed_hamming_matrix = (
hamming_distance_matrix(cluster_center_ids) if fill_empty_clusters else None
)
for projection_index, simhash in enumerate(self.simhash_projections):
# Initialize cluster centers and count vectors assigned to each cluster
cluster_centers = np.zeros((num_partitions, self.dim))
cluster_center_id_to_vectors: dict[int, list[int]] = {
cluster_center_id: [] for cluster_center_id in cluster_center_ids
}
cluster_vector_counts = None
empty_mask = None
# Assign each vector to its cluster and accumulate cluster centers
vector_cluster_ids = simhash.get_cluster_ids(vectors)
for cluster_id, (vec_idx, vec) in zip(vector_cluster_ids, enumerate(vectors)):
cluster_centers[cluster_id] += vec
cluster_center_id_to_vectors[cluster_id].append(vec_idx)
if normalize_by_count or fill_empty_clusters:
cluster_vector_counts = np.bincount(vector_cluster_ids, minlength=num_partitions)
empty_mask = cluster_vector_counts == 0
if normalize_by_count:
assert empty_mask is not None
assert cluster_vector_counts is not None
non_empty_mask = ~empty_mask
cluster_centers[non_empty_mask] /= cluster_vector_counts[non_empty_mask][:, None]
# Fill empty clusters using vectors with minimum Hamming distance
if fill_empty_clusters:
assert empty_mask is not None
assert precomputed_hamming_matrix is not None
masked_hamming = np.where(
empty_mask[None, :], MAX_HAMMING_DISTANCE, precomputed_hamming_matrix
)
nearest_non_empty = np.argmin(masked_hamming, axis=1)
fill_vectors = np.array(
[
vectors[cluster_center_id_to_vectors[cluster_id][0]]
for cluster_id in nearest_non_empty[empty_mask]
]
).reshape(-1, self.dim)
cluster_centers[empty_mask] = fill_vectors
# Apply random projection for dimensionality reduction if needed
if self.dim_proj < self.dim:
dim_reduction_projection = self.dim_reduction_projections[
projection_index
] # Get projection matrix for this repetition
projected_centers = (1 / np.sqrt(self.dim_proj)) * (
cluster_centers @ dim_reduction_projection
)
# Flatten cluster centers into a single vector and add to output
output_vectors.append(projected_centers.flatten())
continue
# If no projection needed (dim_proj == dim), use original cluster centers
output_vectors.append(cluster_centers.flatten())
# Concatenate results from all R_reps projections into final FDE
return np.concatenate(output_vectors)
if __name__ == "__main__":
v_arrs = np.random.randn(10, 100, 128)
muvera = Muvera(128, 4, 8, 20, 42)
for v_arr in v_arrs:
muvera.process(v_arr) # type: ignore
@@ -0,0 +1,78 @@
from typing import Sequence, Any, Type
from fastembed.common import OnnxProvider
from fastembed.common.model_description import BaseModelDescription
from fastembed.common.types import Device
from fastembed.rerank.cross_encoder.onnx_text_cross_encoder import OnnxTextCrossEncoder
from fastembed.rerank.cross_encoder.onnx_text_model import TextRerankerWorker
class CustomTextCrossEncoder(OnnxTextCrossEncoder):
SUPPORTED_MODELS: list[BaseModelDescription] = []
def __init__(
self,
model_name: str,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
device_id: int | None = None,
specific_model_path: str | None = None,
**kwargs: Any,
):
super().__init__(
model_name=model_name,
cache_dir=cache_dir,
threads=threads,
providers=providers,
cuda=cuda,
device_ids=device_ids,
lazy_load=lazy_load,
device_id=device_id,
specific_model_path=specific_model_path,
**kwargs,
)
@classmethod
def _list_supported_models(cls) -> list[BaseModelDescription]:
return cls.SUPPORTED_MODELS
@classmethod
def _get_worker_class(cls) -> Type[TextRerankerWorker]:
return CustomTextCrossEncoderWorker
def _get_worker_init_kwargs(self) -> dict[str, Any]:
return {"model_description": self.model_description}
@classmethod
def add_model(
cls,
model_description: BaseModelDescription,
) -> None:
cls.SUPPORTED_MODELS.append(model_description)
class CustomTextCrossEncoderWorker(TextRerankerWorker):
def init_embedding(
self,
model_name: str,
cache_dir: str,
model_description: BaseModelDescription | None = None,
**kwargs: Any,
) -> CustomTextCrossEncoder:
if model_description is None:
raise ValueError(
"`model_description` is required to initialize a custom model in a worker "
"process, it is provided by `CustomTextCrossEncoder._get_worker_init_kwargs`"
)
# custom models live in a class-level registry, which spawned workers don't inherit
CustomTextCrossEncoder.add_model(model_description)
return CustomTextCrossEncoder(
model_name=model_name,
cache_dir=cache_dir,
threads=1,
**kwargs,
)
@@ -1,9 +1,10 @@
from typing import Any, Iterable, Optional, Sequence, Type
from typing import Any, Iterable, Sequence, Type
from loguru import logger
from fastembed.common import OnnxProvider
from fastembed.common.onnx_model import OnnxOutputContext
from fastembed.common.types import Device
from fastembed.common.utils import define_cache_dir
from fastembed.rerank.cross_encoder.onnx_text_model import (
OnnxCrossEncoderModel,
@@ -77,14 +78,14 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
device_id: Optional[int] = None,
specific_model_path: Optional[str] = None,
device_id: int | None = None,
specific_model_path: str | None = None,
**kwargs: Any,
):
"""
@@ -96,10 +97,11 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.
providers (Optional[Sequence[OnnxProvider]], optional): The list of onnxruntime providers to use.
Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to False.
cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to Device.AUTO.
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, mutually exclusive
with `providers`. Defaults to None.
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
@@ -111,6 +113,7 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
super().__init__(model_name, cache_dir, threads, **kwargs)
self.providers = providers
self.lazy_load = lazy_load
self._extra_session_options = self._select_exposed_session_options(kwargs)
# List of device ids, that can be used for data parallel processing in workers
self.device_ids = device_ids
@@ -123,7 +126,7 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
)
# This device_id will be used if we need to load model in current process
self.device_id: Optional[int] = None
self.device_id: int | None = None
if device_id is not None:
self.device_id = device_id
elif self.device_ids is not None:
@@ -131,11 +134,12 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
self.model_description = self._get_model_description(model_name)
self.cache_dir = str(define_cache_dir(cache_dir))
self._specific_model_path = specific_model_path
self._model_dir = self.download_model(
self.model_description,
self.cache_dir,
local_files_only=self._local_files_only,
specific_model_path=specific_model_path,
specific_model_path=self._specific_model_path,
)
if not self.lazy_load:
@@ -149,6 +153,7 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
providers=self.providers,
cuda=self.cuda,
device_id=self.device_id,
extra_session_options=self._extra_session_options,
)
def rerank(
@@ -177,7 +182,7 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
self,
pairs: Iterable[tuple[str, str]],
batch_size: int = 64,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[float]:
yield from self._rerank_pairs(
@@ -189,6 +194,9 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
providers=self.providers,
cuda=self.cuda,
device_ids=self.device_ids,
local_files_only=self._local_files_only,
specific_model_path=self._specific_model_path,
extra_session_options=self._extra_session_options,
**kwargs,
)
@@ -196,9 +204,25 @@ class OnnxTextCrossEncoder(TextCrossEncoderBase, OnnxCrossEncoderModel):
def _get_worker_class(cls) -> Type[TextRerankerWorker]:
return TextCrossEncoderWorker
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[float]:
def _post_process_onnx_output(
self, output: OnnxOutputContext, **kwargs: Any
) -> Iterable[float]:
return (float(elem) for elem in output.model_output)
def token_count(
self, pairs: Iterable[tuple[str, str]], batch_size: int = 1024, **kwargs: Any
) -> int:
"""Returns the number of tokens in the pairs.
Args:
pairs: Iterable of tuples, where each tuple contains a query and a document to be tokenized
batch_size: Batch size for tokenizing
Returns:
token count: overall number of tokens in the pairs
"""
return self._token_count(pairs, batch_size=batch_size, **kwargs)
class TextCrossEncoderWorker(TextRerankerWorker):
def init_embedding(
@@ -1,7 +1,7 @@
import os
from multiprocessing import get_all_start_methods
from pathlib import Path
from typing import Any, Iterable, Optional, Sequence, Type
from typing import Any, Iterable, Sequence, Type
import numpy as np
from tokenizers import Encoding
@@ -12,14 +12,14 @@ from fastembed.common.onnx_model import (
OnnxOutputContext,
OnnxProvider,
)
from fastembed.common.types import NumpyArray
from fastembed.common.types import NumpyArray, Device
from fastembed.common.preprocessor_utils import load_tokenizer
from fastembed.common.utils import iter_batch
from fastembed.parallel_processor import ParallelWorkerPool
class OnnxCrossEncoderModel(OnnxModel[float]):
ONNX_OUTPUT_NAMES: Optional[list[str]] = None
ONNX_OUTPUT_NAMES: list[str] | None = None
@classmethod
def _get_worker_class(cls) -> Type["TextRerankerWorker"]:
@@ -29,10 +29,11 @@ class OnnxCrossEncoderModel(OnnxModel[float]):
self,
model_dir: Path,
model_file: str,
threads: Optional[int],
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_id: Optional[int] = None,
threads: int | None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_id: int | None = None,
extra_session_options: dict[str, Any] | None = None,
) -> None:
super()._load_onnx_model(
model_dir=model_dir,
@@ -41,6 +42,7 @@ class OnnxCrossEncoderModel(OnnxModel[float]):
providers=providers,
cuda=cuda,
device_id=device_id,
extra_session_options=extra_session_options,
)
self.tokenizer, _ = load_tokenizer(model_dir=model_dir)
assert self.tokenizer is not None
@@ -90,10 +92,13 @@ class OnnxCrossEncoderModel(OnnxModel[float]):
cache_dir: str,
pairs: Iterable[tuple[str, str]],
batch_size: int,
parallel: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
parallel: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
local_files_only: bool = False,
specific_model_path: str | None = None,
extra_session_options: dict[str, Any] | None = None,
**kwargs: Any,
) -> Iterable[float]:
is_small = False
@@ -120,9 +125,15 @@ class OnnxCrossEncoderModel(OnnxModel[float]):
"model_name": model_name,
"cache_dir": cache_dir,
"providers": providers,
"local_files_only": local_files_only,
"specific_model_path": specific_model_path,
**kwargs,
**self._get_worker_init_kwargs(),
}
if extra_session_options is not None:
params.update(extra_session_options)
pool = ParallelWorkerPool(
num_workers=parallel or 1,
worker=self._get_worker_class(),
@@ -133,7 +144,18 @@ class OnnxCrossEncoderModel(OnnxModel[float]):
for batch in pool.ordered_map(iter_batch(pairs, batch_size), **params):
yield from self._post_process_onnx_output(batch) # type: ignore
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[float]:
def _post_process_onnx_output(
self, output: OnnxOutputContext, **kwargs: Any
) -> Iterable[float]:
"""Post-process the ONNX model output to convert it into a usable format.
Args:
output (OnnxOutputContext): The raw output from the ONNX model.
**kwargs: Additional keyword arguments that may be needed by specific implementations.
Returns:
Iterable[float]: Post-processed output as an iterable of float values.
"""
raise NotImplementedError("Subclasses must implement this method")
def _preprocess_onnx_input(
@@ -144,6 +166,20 @@ class OnnxCrossEncoderModel(OnnxModel[float]):
"""
return onnx_input
def _token_count(
self, pairs: Iterable[tuple[str, str]], batch_size: int = 1024, **_: Any
) -> int:
if not hasattr(self, "model") or self.model is None:
self.load_onnx_model() # loads the tokenizer as well
token_num = 0
assert self.tokenizer is not None
for batch in iter_batch(pairs, batch_size):
for tokens in self.tokenizer.encode_batch(batch):
token_num += sum(tokens.attention_mask)
return token_num
class TextRerankerWorker(EmbeddingWorker[float]):
def __init__(
@@ -1,15 +1,22 @@
from typing import Any, Iterable, Optional, Sequence, Type
from typing import Any, Iterable, Sequence, Type
from dataclasses import asdict
from fastembed.common import OnnxProvider
from fastembed.common.types import Device
from fastembed.rerank.cross_encoder.onnx_text_cross_encoder import OnnxTextCrossEncoder
from fastembed.rerank.cross_encoder.custom_text_cross_encoder import CustomTextCrossEncoder
from fastembed.rerank.cross_encoder.text_cross_encoder_base import TextCrossEncoderBase
from fastembed.common.model_description import BaseModelDescription
from fastembed.common.model_description import (
ModelSource,
BaseModelDescription,
)
class TextCrossEncoder(TextCrossEncoderBase):
CROSS_ENCODER_REGISTRY: list[Type[TextCrossEncoderBase]] = [
OnnxTextCrossEncoder,
CustomTextCrossEncoder,
]
@classmethod
@@ -47,11 +54,11 @@ class TextCrossEncoder(TextCrossEncoderBase):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
**kwargs: Any,
):
@@ -96,7 +103,7 @@ class TextCrossEncoder(TextCrossEncoderBase):
self,
pairs: Iterable[tuple[str, str]],
batch_size: int = 64,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[float]:
"""
@@ -124,3 +131,48 @@ class TextCrossEncoder(TextCrossEncoderBase):
yield from self.model.rerank_pairs(
pairs, batch_size=batch_size, parallel=parallel, **kwargs
)
@classmethod
def add_custom_model(
cls,
model: str,
sources: ModelSource,
model_file: str = "onnx/model.onnx",
description: str = "",
license: str = "",
size_in_gb: float = 0.0,
additional_files: list[str] | None = None,
) -> None:
registered_models = cls._list_supported_models()
for registered_model in registered_models:
if model == registered_model.model:
raise ValueError(
f"Model {model} is already registered in CrossEncoderModel, if you still want to add this model, "
f"please use another model name"
)
CustomTextCrossEncoder.add_model(
BaseModelDescription(
model=model,
sources=sources,
model_file=model_file,
description=description,
license=license,
size_in_GB=size_in_gb,
additional_files=additional_files or [],
)
)
def token_count(
self, pairs: Iterable[tuple[str, str]], batch_size: int = 1024, **kwargs: Any
) -> int:
"""Returns the number of tokens in the pairs.
Args:
pairs: Iterable of tuples, where each tuple contains a query and a document to be tokenized
batch_size: Batch size for tokenizing
Returns:
token count: overall number of tokens in the pairs
"""
return self.model.token_count(pairs, batch_size=batch_size, **kwargs)
@@ -1,4 +1,4 @@
from typing import Any, Iterable, Optional
from typing import Any, Iterable
from fastembed.common.model_description import BaseModelDescription
from fastembed.common.model_management import ModelManagement
@@ -8,8 +8,8 @@ class TextCrossEncoderBase(ModelManagement[BaseModelDescription]):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
cache_dir: str | None = None,
threads: int | None = None,
**kwargs: Any,
):
self.model_name = model_name
@@ -41,7 +41,7 @@ class TextCrossEncoderBase(ModelManagement[BaseModelDescription]):
self,
pairs: Iterable[tuple[str, str]],
batch_size: int = 64,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[float]:
"""Rerank query-document pairs.
@@ -57,3 +57,7 @@ class TextCrossEncoderBase(ModelManagement[BaseModelDescription]):
Iterable[float]: Scores for each individual pair
"""
raise NotImplementedError("This method should be overridden by subclasses")
def token_count(self, pairs: Iterable[tuple[str, str]], **kwargs: Any) -> int:
"""Returns the number of tokens in the pairs."""
raise NotImplementedError("This method should be overridden by subclasses")
+27 -23
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
@@ -21,13 +21,9 @@ from fastembed.sparse.sparse_embedding_base import (
from fastembed.sparse.utils.tokenizer import SimpleTokenizer
from fastembed.common.model_description import SparseModelDescription, ModelSource
supported_languages = [
"arabic",
"azerbaijani",
"basque",
"bengali",
"catalan",
"chinese",
"danish",
"dutch",
"english",
@@ -35,21 +31,15 @@ supported_languages = [
"french",
"german",
"greek",
"hebrew",
"hinglish",
"hungarian",
"indonesian",
"italian",
"kazakh",
"nepali",
"norwegian",
"portuguese",
"romanian",
"russian",
"slovene",
"spanish",
"swedish",
"tajik",
"tamil",
"turkish",
]
@@ -101,14 +91,14 @@ class Bm25(SparseTextEmbeddingBase):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
cache_dir: str | None = None,
k: float = 1.2,
b: float = 0.75,
avg_len: float = 256.0,
language: str = "english",
token_max_length: int = 40,
disable_stemmer: bool = False,
specific_model_path: Optional[str] = None,
specific_model_path: str | None = None,
**kwargs: Any,
):
super().__init__(model_name, cache_dir, **kwargs)
@@ -125,11 +115,12 @@ class Bm25(SparseTextEmbeddingBase):
model_description = self._get_model_description(model_name)
self.cache_dir = str(define_cache_dir(cache_dir))
self._specific_model_path = specific_model_path
self._model_dir = self.download_model(
model_description,
self.cache_dir,
local_files_only=self._local_files_only,
specific_model_path=specific_model_path,
specific_model_path=self._specific_model_path,
)
self.token_max_length = token_max_length
@@ -167,9 +158,11 @@ class Bm25(SparseTextEmbeddingBase):
self,
model_name: str,
cache_dir: str,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
parallel: int | None = None,
local_files_only: bool = False,
specific_model_path: str | None = None,
) -> Iterable[SparseEmbedding]:
is_small = False
@@ -198,6 +191,8 @@ class Bm25(SparseTextEmbeddingBase):
"language": self.language,
"token_max_length": self.token_max_length,
"disable_stemmer": self.disable_stemmer,
"local_files_only": local_files_only,
"specific_model_path": specific_model_path,
}
pool = ParallelWorkerPool(
num_workers=parallel or 1,
@@ -210,9 +205,9 @@ class Bm25(SparseTextEmbeddingBase):
def embed(
self,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[SparseEmbedding]:
"""
@@ -236,6 +231,8 @@ class Bm25(SparseTextEmbeddingBase):
documents=documents,
batch_size=batch_size,
parallel=parallel,
local_files_only=self._local_files_only,
specific_model_path=self._specific_model_path,
)
def _stem(self, tokens: list[str]) -> list[str]:
@@ -271,6 +268,15 @@ class Bm25(SparseTextEmbeddingBase):
embeddings.append(SparseEmbedding.from_dict(token_id2value))
return embeddings
def token_count(self, texts: str | Iterable[str], **kwargs: Any) -> int:
token_num = 0
texts = [texts] if isinstance(texts, str) else texts
for text in texts:
document = remove_non_alphanumeric(text)
tokens = self.tokenizer.tokenize(document)
token_num += len(tokens)
return token_num
def _term_frequency(self, tokens: list[str]) -> dict[int, float]:
"""Calculate the term frequency part of the BM25 formula.
@@ -305,9 +311,7 @@ class Bm25(SparseTextEmbeddingBase):
def compute_token_id(cls, token: str) -> int:
return abs(mmh3.hash(token))
def query_embed(
self, query: Union[str, Iterable[str]], **kwargs: Any
) -> Iterable[SparseEmbedding]:
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[SparseEmbedding]:
"""To emulate BM25 behaviour, we don't need to use weights in the query, and
it's enough to just hash the tokens and assign a weight of 1.0 to them.
"""
+44 -21
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,
@@ -31,9 +32,17 @@ supported_bm42_models: list[SparseModelDescription] = [
),
]
MODEL_TO_LANGUAGE = {
_MODEL_TO_LANGUAGE = {
"Qdrant/bm42-all-minilm-l6-v2-attentions": "english",
}
MODEL_TO_LANGUAGE = {
model_name.lower(): language for model_name, language in _MODEL_TO_LANGUAGE.items()
}
def get_language_by_model_name(model_name: str) -> str:
return MODEL_TO_LANGUAGE[model_name.lower()]
class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
@@ -57,15 +66,15 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
alpha: float = 0.5,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
device_id: Optional[int] = None,
specific_model_path: Optional[str] = None,
device_id: int | None = None,
specific_model_path: str | None = None,
**kwargs: Any,
):
"""
@@ -79,10 +88,11 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
alpha (float, optional): Parameter, that defines the importance of the token weight in the document
versus the importance of the token frequency in the corpus. Defaults to 0.5, based on empirical testing.
It is recommended to only change this parameter based on training data for a specific dataset.
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to False.
cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to Device.AUTO.
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, mutually exclusive
with `providers`. Defaults to None.
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
@@ -95,13 +105,14 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
super().__init__(model_name, cache_dir, threads, **kwargs)
self.providers = providers
self.lazy_load = lazy_load
self._extra_session_options = self._select_exposed_session_options(kwargs)
# List of device ids, that can be used for data parallel processing in workers
self.device_ids = device_ids
self.cuda = cuda
# This device_id will be used if we need to load model in current process
self.device_id: Optional[int] = None
self.device_id: int | None = None
if device_id is not None:
self.device_id = device_id
elif self.device_ids is not None:
@@ -110,11 +121,12 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
self.model_description = self._get_model_description(model_name)
self.cache_dir = str(define_cache_dir(cache_dir))
self._specific_model_path = specific_model_path
self._model_dir = self.download_model(
self.model_description,
self.cache_dir,
local_files_only=self._local_files_only,
specific_model_path=specific_model_path,
specific_model_path=self._specific_model_path,
)
self.invert_vocab: dict[int, str] = {}
@@ -123,7 +135,7 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
self.special_tokens_ids: set[int] = set()
self.punctuation = set(string.punctuation)
self.stopwords = set(self._load_stopwords(self._model_dir))
self.stemmer = SnowballStemmer(MODEL_TO_LANGUAGE[model_name])
self.stemmer = SnowballStemmer(get_language_by_model_name(self.model_name))
self.alpha = alpha
if not self.lazy_load:
@@ -137,6 +149,7 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
providers=self.providers,
cuda=self.cuda,
device_id=self.device_id,
extra_session_options=self._extra_session_options,
)
for token, idx in self.tokenizer.get_vocab().items(): # type: ignore[union-attr]
@@ -217,7 +230,9 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
return new_vector
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[SparseEmbedding]:
def _post_process_onnx_output(
self, output: OnnxOutputContext, **kwargs: Any
) -> Iterable[SparseEmbedding]:
if output.input_ids is None:
raise ValueError("input_ids must be provided for document post-processing")
@@ -269,9 +284,9 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
def embed(
self,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[SparseEmbedding]:
"""
@@ -299,6 +314,9 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
cuda=self.cuda,
device_ids=self.device_ids,
alpha=self.alpha,
local_files_only=self._local_files_only,
specific_model_path=self._specific_model_path,
extra_session_options=self._extra_session_options,
)
@classmethod
@@ -309,9 +327,7 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
result[token_id] = 1.0
return result
def query_embed(
self, query: Union[str, Iterable[str]], **kwargs: Any
) -> Iterable[SparseEmbedding]:
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[SparseEmbedding]:
"""
To emulate BM25 behaviour, we don't need to use smart weights in the query, and
it's enough to just hash the tokens and assign a weight of 1.0 to them.
@@ -336,6 +352,13 @@ class Bm42(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
def _get_worker_class(cls) -> Type[TextEmbeddingWorker[SparseEmbedding]]:
return Bm42TextEmbeddingWorker
def token_count(
self, texts: str | Iterable[str], batch_size: int = 1024, **kwargs: Any
) -> int:
if not hasattr(self, "model") or self.model is None:
self.load_onnx_model() # loads the tokenizer as well
return self._token_count(texts, batch_size=batch_size, **kwargs)
class Bm42TextEmbeddingWorker(TextEmbeddingWorker[SparseEmbedding]):
def init_embedding(self, model_name: str, cache_dir: str, **kwargs: Any) -> Bm42:
+246
View File
@@ -0,0 +1,246 @@
import json
from typing import Any, Iterable, Sequence, Type
import numpy as np
from fastembed.common import OnnxProvider
from fastembed.common.model_description import ModelSource, SparseModelDescription
from fastembed.common.onnx_model import OnnxOutputContext
from fastembed.common.preprocessor_utils import load_tokenizer
from fastembed.common.types import Device
from fastembed.common.utils import define_cache_dir, iter_batch
from fastembed.sparse.sparse_embedding_base import (
SparseEmbedding,
SparseTextEmbeddingBase,
)
from fastembed.text.onnx_text_model import OnnxTextModel, TextEmbeddingWorker
IDF_FILE = "idf.json"
supported_if_splade_models: list[SparseModelDescription] = [
SparseModelDescription(
model="opensearch-project/opensearch-neural-sparse-encoding-doc-v3-gte",
vocab_size=30522,
description="Inference-free SPLADE model. Documents are expanded with an ONNX encoder at index "
"time, queries are encoded with a tokenizer and an IDF lookup table only, "
"without any model inference.",
license="apache-2.0",
size_in_GB=0.55,
sources=ModelSource(hf="Qdrant/opensearch-neural-sparse-encoding-doc-v3-gte"),
model_file="model.onnx",
additional_files=[IDF_FILE],
requires_idf=None,
),
]
class IfSplade(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
"""Inference-free (asymmetric) SPLADE model.
Documents are encoded with a neural encoder which expands them into a sparse vocabulary-sized
vector, while queries are encoded by tokenizing the text and looking up a precomputed IDF
weight per token — no neural inference happens at query time.
Query and document embeddings are compared with a dot product.
Special tokens are excluded from both document and query embeddings.
"""
def _post_process_onnx_output(
self, output: OnnxOutputContext, **kwargs: Any
) -> Iterable[SparseEmbedding]:
if output.attention_mask is None:
raise ValueError("attention_mask must be provided for document post-processing")
# Max-pool token logits over the sequence, masking out the padding
pooled = np.max(
output.model_output * np.expand_dims(output.attention_mask, axis=-1), axis=1
)
# v3 models of the opensearch-neural-sparse family use a double log activation,
# log(1 + log(1 + relu(x))), to increase sparsity of document embeddings
scores = np.log1p(np.log1p(np.maximum(pooled, 0.0)))
if self.special_tokens_ids:
scores[:, list(self.special_tokens_ids)] = 0.0
for row_scores in scores:
indices = row_scores.nonzero()[0]
yield SparseEmbedding(values=row_scores[indices], indices=indices)
def token_count(
self, texts: str | Iterable[str], batch_size: int = 1024, **kwargs: Any
) -> int:
# unlike `OnnxTextModel._token_count`, does not require the onnx model to be loaded
token_num = 0
texts = [texts] if isinstance(texts, str) else texts
for batch in iter_batch(texts, batch_size):
for tokens in self.tokenizer.encode_batch(batch): # type: ignore[union-attr]
token_num += sum(tokens.attention_mask)
return token_num
@classmethod
def _list_supported_models(cls) -> list[SparseModelDescription]:
"""Lists the supported models.
Returns:
list[SparseModelDescription]: A list of SparseModelDescription objects containing the model information.
"""
return supported_if_splade_models
def __init__(
self,
model_name: str,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
device_id: int | None = None,
specific_model_path: str | None = None,
**kwargs: Any,
):
"""
Args:
model_name (str): The name of the model to use.
cache_dir (str, optional): The path to the cache directory.
Can be set using the `FASTEMBED_CACHE_PATH` env variable.
Defaults to `fastembed_cache` in the system's temp directory.
threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.
providers (Optional[Sequence[OnnxProvider]], optional): The list of onnxruntime providers to use.
Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to Device.
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, mutually exclusive
with `providers`. Defaults to None.
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
specific_model_path (Optional[str], optional): The specific path to the onnx model dir if it should be imported from somewhere else
Raises:
ValueError: If the model_name is not in the format <org>/<model> e.g. BAAI/bge-base-en.
"""
super().__init__(model_name, cache_dir, threads, **kwargs)
self.providers = providers
self.lazy_load = lazy_load
self._extra_session_options = self._select_exposed_session_options(kwargs)
# List of device ids, that can be used for data parallel processing in workers
self.device_ids = device_ids
self.cuda = cuda
# This device_id will be used if we need to load model in current process
self.device_id: int | None = None
if device_id is not None:
self.device_id = device_id
elif self.device_ids is not None:
self.device_id = self.device_ids[0]
self.model_description = self._get_model_description(model_name)
self.cache_dir = str(define_cache_dir(cache_dir))
self._specific_model_path = specific_model_path
self._model_dir = self.download_model(
self.model_description,
self.cache_dir,
local_files_only=self._local_files_only,
specific_model_path=self._specific_model_path,
)
# The tokenizer and the idf table are lightweight and are required for query embedding,
# which does not involve any model inference, so they are loaded eagerly, while
# `lazy_load` only defers the initialization of the onnx model
self.tokenizer, self.special_token_to_id = load_tokenizer(model_dir=self._model_dir)
self.special_tokens_ids: set[int] = set(self.special_token_to_id.values())
self._token_id_to_idf = self._load_idf()
if not self.lazy_load:
self.load_onnx_model()
def load_onnx_model(self) -> None:
self._load_onnx_model(
model_dir=self._model_dir,
model_file=self.model_description.model_file,
threads=self.threads,
providers=self.providers,
cuda=self.cuda,
device_id=self.device_id,
extra_session_options=self._extra_session_options,
)
def _load_idf(self) -> dict[int, float]:
with open(self._model_dir / IDF_FILE) as f:
token_to_idf: dict[str, float] = json.load(f)
vocab: dict[str, int] = self.tokenizer.get_vocab() # type: ignore[union-attr]
return {vocab[token]: idf for token, idf in token_to_idf.items() if token in vocab}
def embed(
self,
documents: str | Iterable[str],
batch_size: int = 256,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[SparseEmbedding]:
"""
Encode a list of documents into list of embeddings.
Args:
documents: Iterator of documents or single document to embed
batch_size: Batch size for encoding -- higher values will use more memory, but be faster
parallel:
If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.
If 0, use all available cores.
If None, don't use data-parallel processing, use default onnxruntime threading instead.
Returns:
List of embeddings, one per document
"""
yield from self._embed_documents(
model_name=self.model_name,
cache_dir=str(self.cache_dir),
documents=documents,
batch_size=batch_size,
parallel=parallel,
providers=self.providers,
cuda=self.cuda,
device_ids=self.device_ids,
local_files_only=self._local_files_only,
specific_model_path=self._specific_model_path,
extra_session_options=self._extra_session_options,
**kwargs,
)
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[SparseEmbedding]:
"""
Encode a list of queries into list of sparse embeddings without any model inference.
A query is tokenized, and each unique token is assigned its IDF weight from
a precomputed lookup table shipped with the model. Special tokens are ignored.
"""
if isinstance(query, str):
query = [query]
for text in query:
token_ids = set(self.tokenizer.encode(text).ids) - self.special_tokens_ids # type: ignore[union-attr]
embedding = {
token_id: self._token_id_to_idf[token_id]
for token_id in sorted(token_ids)
if token_id in self._token_id_to_idf
}
yield SparseEmbedding.from_dict(embedding)
@classmethod
def _get_worker_class(cls) -> Type[TextEmbeddingWorker[SparseEmbedding]]:
return IfSpladeEmbeddingWorker
class IfSpladeEmbeddingWorker(TextEmbeddingWorker[SparseEmbedding]):
def init_embedding(self, model_name: str, cache_dir: str, **kwargs: Any) -> IfSplade:
return IfSplade(
model_name=model_name,
cache_dir=cache_dir,
threads=1,
**kwargs,
)
+372
View File
@@ -0,0 +1,372 @@
from pathlib import Path
from typing import Any, Sequence, Iterable, Type
import numpy as np
from numpy.typing import NDArray
from py_rust_stemmers import SnowballStemmer
from tokenizers import Tokenizer
from fastembed.common.model_description import SparseModelDescription, ModelSource
from fastembed.common.onnx_model import OnnxOutputContext
from fastembed.common import OnnxProvider
from fastembed.common.types import Device
from fastembed.common.utils import define_cache_dir
from fastembed.sparse.sparse_embedding_base import (
SparseEmbedding,
SparseTextEmbeddingBase,
)
from fastembed.sparse.utils.minicoil_encoder import Encoder
from fastembed.sparse.utils.sparse_vectors_converter import SparseVectorConverter, WordEmbedding
from fastembed.sparse.utils.vocab_resolver import VocabResolver, VocabTokenizer
from fastembed.text.onnx_text_model import OnnxTextModel, TextEmbeddingWorker
MINICOIL_MODEL_FILE = "minicoil.triplet.model.npy"
MINICOIL_VOCAB_FILE = "minicoil.triplet.model.vocab"
STOPWORDS_FILE = "stopwords.txt"
supported_minicoil_models: list[SparseModelDescription] = [
SparseModelDescription(
model="Qdrant/minicoil-v1",
vocab_size=19125,
description="Sparse embedding model, that resolves semantic meaning of the words, "
"while keeping exact keyword match behavior. "
"Based on jinaai/jina-embeddings-v2-small-en-tokens",
license="apache-2.0",
size_in_GB=0.09,
sources=ModelSource(hf="Qdrant/minicoil-v1"),
model_file="onnx/model.onnx",
additional_files=[
STOPWORDS_FILE,
MINICOIL_MODEL_FILE,
MINICOIL_VOCAB_FILE,
],
requires_idf=True,
),
]
_MODEL_TO_LANGUAGE = {
"Qdrant/minicoil-v1": "english",
}
MODEL_TO_LANGUAGE = {
model_name.lower(): language for model_name, language in _MODEL_TO_LANGUAGE.items()
}
def get_language_by_model_name(model_name: str) -> str:
return MODEL_TO_LANGUAGE[model_name.lower()]
class MiniCOIL(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
"""
MiniCOIL is a sparse embedding model, that resolves semantic meaning of the words,
while keeping exact keyword match behavior.
Each vocabulary token is converted into 4d component of a sparse vector, which is then weighted by the token frequency in the corpus.
If the token is not found in the corpus, it is treated exactly like in BM25.
`
The model is based on `jinaai/jina-embeddings-v2-small-en-tokens`
"""
def __init__(
self,
model_name: str,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
k: float = 1.2,
b: float = 0.75,
avg_len: float = 150.0,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
device_id: int | None = None,
specific_model_path: str | None = None,
**kwargs: Any,
):
"""
Args:
model_name (str): The name of the model to use.
cache_dir (str, optional): The path to the cache directory.
Can be set using the `FASTEMBED_CACHE_PATH` env variable.
Defaults to `fastembed_cache` in the system's temp directory.
threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.
providers (Optional[Sequence[OnnxProvider]], optional): The providers to use for onnxruntime.
k (float, optional): The k parameter in the BM25 formula. Defines the saturation of the term frequency.
I.e. defines how fast the moment when additional terms stop to increase the score. Defaults to 1.2.
b (float, optional): The b parameter in the BM25 formula. Defines the importance of the document length.
Defaults to 0.75.
avg_len (float, optional): The average length of the documents in the corpus. Defaults to 150.0.
cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to Device.AUTO.
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, mutually exclusive
with `providers`. Defaults to None.
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
specific_model_path (Optional[str], optional): The specific path to the onnx model dir if it should be imported from somewhere else
Raises:
ValueError: If the model_name is not in the format <org>/<model> e.g. BAAI/bge-base-en.
"""
super().__init__(model_name, cache_dir, threads, **kwargs)
self.providers = providers
self.lazy_load = lazy_load
self.device_ids = device_ids
self.cuda = cuda
self.device_id = device_id
self._extra_session_options = self._select_exposed_session_options(kwargs)
self.k = k
self.b = b
self.avg_len = avg_len
# Initialize class attributes
self.tokenizer: Tokenizer | None = None
self.invert_vocab: dict[int, str] = {}
self.special_tokens: set[str] = set()
self.special_tokens_ids: set[int] = set()
self.stopwords: set[str] = set()
self.vocab_resolver: VocabResolver | None = None
self.encoder: Encoder | None = None
self.output_dim: int | None = None
self.sparse_vector_converter: SparseVectorConverter | None = None
self.model_description = self._get_model_description(model_name)
self.cache_dir = str(define_cache_dir(cache_dir))
self._specific_model_path = specific_model_path
self._model_dir = self.download_model(
self.model_description,
self.cache_dir,
local_files_only=self._local_files_only,
specific_model_path=self._specific_model_path,
)
if not self.lazy_load:
self.load_onnx_model()
def load_onnx_model(self) -> None:
self._load_onnx_model(
model_dir=self._model_dir,
model_file=self.model_description.model_file,
threads=self.threads,
providers=self.providers,
cuda=self.cuda,
device_id=self.device_id,
extra_session_options=self._extra_session_options,
)
assert self.tokenizer is not None
for token, idx in self.tokenizer.get_vocab().items(): # type: ignore[union-attr]
self.invert_vocab[idx] = token
self.special_tokens = set(self.special_token_to_id.keys())
self.special_tokens_ids = set(self.special_token_to_id.values())
self.stopwords = set(self._load_stopwords(self._model_dir))
stemmer = SnowballStemmer(get_language_by_model_name(self.model_name))
self.vocab_resolver = VocabResolver(
tokenizer=VocabTokenizer(self.tokenizer),
stopwords=self.stopwords,
stemmer=stemmer,
)
self.vocab_resolver.load_json_vocab(str(self._model_dir / MINICOIL_VOCAB_FILE))
weights = np.load(str(self._model_dir / MINICOIL_MODEL_FILE), mmap_mode="r")
self.encoder = Encoder(weights)
self.output_dim = self.encoder.output_dim
self.sparse_vector_converter = SparseVectorConverter(
stopwords=self.stopwords,
stemmer=stemmer,
k=self.k,
b=self.b,
avg_len=self.avg_len,
)
def token_count(
self, texts: str | Iterable[str], batch_size: int = 1024, **kwargs: Any
) -> int:
return self._token_count(texts, batch_size=batch_size, **kwargs)
def embed(
self,
documents: str | Iterable[str],
batch_size: int = 256,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[SparseEmbedding]:
"""
Encode a list of documents into list of embeddings.
We use mean pooling with attention so that the model can handle variable-length inputs.
Args:
documents: Iterator of documents or single document to embed
batch_size: Batch size for encoding -- higher values will use more memory, but be faster
parallel:
If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.
If 0, use all available cores.
If None, don't use data-parallel processing, use default onnxruntime threading instead.
Returns:
List of embeddings, one per document
"""
yield from self._embed_documents(
model_name=self.model_name,
cache_dir=str(self.cache_dir),
documents=documents,
batch_size=batch_size,
parallel=parallel,
providers=self.providers,
cuda=self.cuda,
device_ids=self.device_ids,
k=self.k,
b=self.b,
avg_len=self.avg_len,
is_query=False,
local_files_only=self._local_files_only,
specific_model_path=self._specific_model_path,
extra_session_options=self._extra_session_options,
**kwargs,
)
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[SparseEmbedding]:
"""
Encode a list of queries into list of embeddings.
"""
yield from self._embed_documents(
model_name=self.model_name,
cache_dir=str(self.cache_dir),
documents=query,
providers=self.providers,
cuda=self.cuda,
device_ids=self.device_ids,
k=self.k,
b=self.b,
avg_len=self.avg_len,
is_query=True,
local_files_only=self._local_files_only,
specific_model_path=self._specific_model_path,
**kwargs,
)
@classmethod
def _load_stopwords(cls, model_dir: Path) -> list[str]:
stopwords_path = model_dir / STOPWORDS_FILE
if not stopwords_path.exists():
return []
with open(stopwords_path, "r") as f:
return f.read().splitlines()
@classmethod
def _list_supported_models(cls) -> list[SparseModelDescription]:
"""Lists the supported models.
Returns:
list[SparseModelDescription]: A list of SparseModelDescription objects containing the model information.
"""
return supported_minicoil_models
def _post_process_onnx_output(
self, output: OnnxOutputContext, is_query: bool = False, **kwargs: Any
) -> Iterable[SparseEmbedding]:
if output.input_ids is None:
raise ValueError("input_ids must be provided for document post-processing")
assert self.vocab_resolver is not None
assert self.encoder is not None
assert self.sparse_vector_converter is not None
# Size: (batch_size, sequence_length, hidden_size)
embeddings = output.model_output
# Size: (batch_size, sequence_length)
assert output.attention_mask is not None
masks = output.attention_mask
vocab_size = self.vocab_resolver.vocab_size()
embedding_size = self.encoder.output_dim
# For each document we only select those embeddings that are not masked out
for i in range(embeddings.shape[0]):
# Size: (sequence_length, hidden_size)
token_embeddings = embeddings[i, masks[i] == 1]
# Size: (sequence_length)
token_ids: NDArray[np.int64] = output.input_ids[i, masks[i] == 1]
word_ids_array, counts, oov, forms = self.vocab_resolver.resolve_tokens(token_ids)
# Size: (1, words)
word_ids_array_expanded: NDArray[np.int64] = np.expand_dims(word_ids_array, axis=0)
# Size: (1, words, embedding_size)
token_embeddings_array: NDArray[np.float32] = np.expand_dims(token_embeddings, axis=0)
assert word_ids_array_expanded.shape[1] == token_embeddings_array.shape[1]
# Size of word_ids_mapping: (unique_words, 2) - [vocab_id, batch_id]
# Size of embeddings: (unique_words, embedding_size)
ids_mapping, minicoil_embeddings = self.encoder.forward(
word_ids_array_expanded, token_embeddings_array
)
# Size of counts: (unique_words)
words_ids: list[int] = ids_mapping[:, 0].tolist() # type: ignore[assignment]
sentence_result: dict[str, WordEmbedding] = {}
words = [self.vocab_resolver.lookup_word(word_id) for word_id in words_ids]
for word, word_id, emb in zip(words, words_ids, minicoil_embeddings.tolist()): # type: ignore[arg-type]
if word_id == 0:
continue
sentence_result[word] = WordEmbedding(
word=word,
forms=forms[word],
count=int(counts[word_id]),
word_id=int(word_id),
embedding=emb, # type: ignore[arg-type]
)
for oov_word, count in oov.items():
# {
# "word": oov_word,
# "forms": [oov_word],
# "count": int(count),
# "word_id": -1,
# "embedding": [1]
# }
sentence_result[oov_word] = WordEmbedding(
word=oov_word, forms=[oov_word], count=int(count), word_id=-1, embedding=[1]
)
if not is_query:
yield self.sparse_vector_converter.embedding_to_vector(
sentence_result, vocab_size=vocab_size, embedding_size=embedding_size
)
else:
yield self.sparse_vector_converter.embedding_to_vector_query(
sentence_result, vocab_size=vocab_size, embedding_size=embedding_size
)
@classmethod
def _get_worker_class(cls) -> Type["MiniCoilTextEmbeddingWorker"]:
return MiniCoilTextEmbeddingWorker
class MiniCoilTextEmbeddingWorker(TextEmbeddingWorker[SparseEmbedding]):
def init_embedding(self, model_name: str, cache_dir: str, **kwargs: Any) -> MiniCOIL:
return MiniCOIL(
model_name=model_name,
cache_dir=cache_dir,
threads=1,
**kwargs,
)
+11 -9
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
@@ -86,3 +84,7 @@ class SparseTextEmbeddingBase(ModelManagement[SparseModelDescription]):
yield from self.embed([query], **kwargs)
else:
yield from self.embed(query, **kwargs)
def token_count(self, texts: str | Iterable[str], **kwargs: Any) -> int:
"""Returns the number of tokens in the texts."""
raise NotImplementedError("Subclasses must implement this method")
+34 -13
View File
@@ -1,9 +1,12 @@
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Iterable, Sequence, Type
from dataclasses import asdict
from fastembed.common import OnnxProvider
from fastembed.common.types import Device
from fastembed.sparse.bm25 import Bm25
from fastembed.sparse.bm42 import Bm42
from fastembed.sparse.if_splade import IfSplade
from fastembed.sparse.minicoil import MiniCOIL
from fastembed.sparse.sparse_embedding_base import (
SparseEmbedding,
SparseTextEmbeddingBase,
@@ -14,7 +17,13 @@ from fastembed.common.model_description import SparseModelDescription
class SparseTextEmbedding(SparseTextEmbeddingBase):
EMBEDDINGS_REGISTRY: list[Type[SparseTextEmbeddingBase]] = [SpladePP, Bm42, Bm25]
EMBEDDINGS_REGISTRY: list[Type[SparseTextEmbeddingBase]] = [
SpladePP,
Bm42,
Bm25,
MiniCOIL,
IfSplade,
]
@classmethod
def list_supported_models(cls) -> list[dict[str, Any]]:
@@ -52,16 +61,16 @@ class SparseTextEmbedding(SparseTextEmbeddingBase):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
**kwargs: Any,
):
super().__init__(model_name, cache_dir, threads, **kwargs)
if model_name == "prithvida/Splade_PP_en_v1":
if model_name.lower() == "prithvida/Splade_PP_en_v1".lower():
warnings.warn(
"The right spelling is prithivida/Splade_PP_en_v1. "
"Support of this name will be removed soon, please fix the model_name",
@@ -92,9 +101,9 @@ class SparseTextEmbedding(SparseTextEmbeddingBase):
def embed(
self,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[SparseEmbedding]:
"""
@@ -114,9 +123,7 @@ class SparseTextEmbedding(SparseTextEmbeddingBase):
"""
yield from self.model.embed(documents, batch_size, parallel, **kwargs)
def query_embed(
self, query: Union[str, Iterable[str]], **kwargs: Any
) -> Iterable[SparseEmbedding]:
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[SparseEmbedding]:
"""
Embeds queries
@@ -127,3 +134,17 @@ class SparseTextEmbedding(SparseTextEmbeddingBase):
Iterable[SparseEmbedding]: The sparse embeddings.
"""
yield from self.model.query_embed(query, **kwargs)
def token_count(
self, texts: str | Iterable[str], batch_size: int = 1024, **kwargs: Any
) -> int:
"""Returns the number of tokens in the texts.
Args:
texts (str | Iterable[str]): The list of texts to embed.
batch_size (int): Batch size for encoding
Returns:
int: Sum of number of tokens in the texts.
"""
return self.model.token_count(texts, batch_size=batch_size, **kwargs)
+33 -18
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,
@@ -18,7 +19,7 @@ supported_splade_models: list[SparseModelDescription] = [
description="Independent Implementation of SPLADE++ Model for English.",
license="apache-2.0",
size_in_GB=0.532,
sources=ModelSource(hf="Qdrant/SPLADE_PP_en_v1"),
sources=ModelSource(hf="Qdrant/Splade_PP_en_v1"),
model_file="model.onnx",
),
SparseModelDescription(
@@ -27,14 +28,16 @@ supported_splade_models: list[SparseModelDescription] = [
description="Independent Implementation of SPLADE++ Model for English.",
license="apache-2.0",
size_in_GB=0.532,
sources=ModelSource(hf="Qdrant/SPLADE_PP_en_v1"),
sources=ModelSource(hf="Qdrant/Splade_PP_en_v1"),
model_file="model.onnx",
),
]
class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[SparseEmbedding]:
def _post_process_onnx_output(
self, output: OnnxOutputContext, **kwargs: Any
) -> Iterable[SparseEmbedding]:
if output.attention_mask is None:
raise ValueError("attention_mask must be provided for document post-processing")
@@ -51,6 +54,11 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
scores = row_scores[indices]
yield SparseEmbedding(values=scores, indices=indices)
def token_count(
self, texts: str | Iterable[str], batch_size: int = 1024, **kwargs: Any
) -> int:
return self._token_count(texts, batch_size=batch_size, **kwargs)
@classmethod
def _list_supported_models(cls) -> list[SparseModelDescription]:
"""Lists the supported models.
@@ -63,14 +71,14 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
device_id: Optional[int] = None,
specific_model_path: Optional[str] = None,
device_id: int | None = None,
specific_model_path: str | None = None,
**kwargs: Any,
):
"""
@@ -82,10 +90,11 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.
providers (Optional[Sequence[OnnxProvider]], optional): The list of onnxruntime providers to use.
Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to False.
cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to Device.
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, mutually exclusive
with `providers`. Defaults to None.
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
@@ -97,13 +106,14 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
super().__init__(model_name, cache_dir, threads, **kwargs)
self.providers = providers
self.lazy_load = lazy_load
self._extra_session_options = self._select_exposed_session_options(kwargs)
# List of device ids, that can be used for data parallel processing in workers
self.device_ids = device_ids
self.cuda = cuda
# This device_id will be used if we need to load model in current process
self.device_id: Optional[int] = None
self.device_id: int | None = None
if device_id is not None:
self.device_id = device_id
elif self.device_ids is not None:
@@ -112,11 +122,12 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
self.model_description = self._get_model_description(model_name)
self.cache_dir = str(define_cache_dir(cache_dir))
self._specific_model_path = specific_model_path
self._model_dir = self.download_model(
self.model_description,
self.cache_dir,
local_files_only=self._local_files_only,
specific_model_path=specific_model_path,
specific_model_path=self._specific_model_path,
)
if not self.lazy_load:
@@ -130,13 +141,14 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
providers=self.providers,
cuda=self.cuda,
device_id=self.device_id,
extra_session_options=self._extra_session_options,
)
def embed(
self,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[SparseEmbedding]:
"""
@@ -163,6 +175,9 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
providers=self.providers,
cuda=self.cuda,
device_ids=self.device_ids,
local_files_only=self._local_files_only,
specific_model_path=self._specific_model_path,
extra_session_options=self._extra_session_options,
**kwargs,
)
+146
View File
@@ -0,0 +1,146 @@
"""
Pure numpy implementation of encoder model for a single word.
This model is not trainable, and should only be used for inference.
"""
import numpy as np
from fastembed.common.types import NumpyArray
class Encoder:
"""
Encoder(768, 4, 10000)
Will look like this:
Per-word
Encoder Matrix
Token Embedding(768) (10k, 768, 4)
Tanh
Final linear transformation is accompanied by a non-linear activation function: Tanh.
Tanh is used to ensure that the output is in the range [-1, 1].
It would be easier to visually interpret the output of the model, assuming that each dimension
would need to encode a type of semantic cluster.
"""
def __init__(
self,
weights: NumpyArray,
):
self.weights = weights
self.vocab_size, self.input_dim, self.output_dim = weights.shape
self.encoder_weights: NumpyArray = weights
# Activation function
self.activation = np.tanh
@staticmethod
def convert_vocab_ids(vocab_ids: NumpyArray) -> NumpyArray:
"""
Convert vocab_ids of shape (batch_size, seq_len) into (batch_size, seq_len, 2)
by appending batch_id alongside each vocab_id.
"""
batch_size, seq_len = vocab_ids.shape
batch_ids = np.arange(batch_size, dtype=vocab_ids.dtype).reshape(batch_size, 1)
batch_ids = np.repeat(batch_ids, seq_len, axis=1)
# Stack vocab_ids and batch_ids along the last dimension
combined: NumpyArray = np.stack((vocab_ids, batch_ids), axis=2).astype(np.int32)
return combined
@classmethod
def avg_by_vocab_ids(
cls, vocab_ids: NumpyArray, embeddings: NumpyArray
) -> tuple[NumpyArray, NumpyArray]:
"""
Takes:
vocab_ids: (batch_size, seq_len) int array
embeddings: (batch_size, seq_len, input_dim) float array
Returns:
unique_flattened_vocab_ids: (total_unique, 2) array of [vocab_id, batch_id]
unique_flattened_embeddings: (total_unique, input_dim) averaged embeddings
"""
input_dim = embeddings.shape[2]
# Flatten vocab_ids and embeddings
# flattened_vocab_ids: (batch_size*seq_len, 2)
flattened_vocab_ids = cls.convert_vocab_ids(vocab_ids).reshape(-1, 2)
# flattened_embeddings: (batch_size*seq_len, input_dim)
flattened_embeddings = embeddings.reshape(-1, input_dim)
# Find unique (vocab_id, batch_id) pairs
unique_flattened_vocab_ids, inverse_indices = np.unique(
flattened_vocab_ids, axis=0, return_inverse=True
)
# Prepare arrays to accumulate sums
unique_count = unique_flattened_vocab_ids.shape[0]
unique_flattened_embeddings = np.zeros((unique_count, input_dim), dtype=np.float32)
unique_flattened_count = np.zeros(unique_count, dtype=np.int32)
# Use np.add.at to accumulate sums based on inverse indices
np.add.at(unique_flattened_embeddings, inverse_indices, flattened_embeddings)
np.add.at(unique_flattened_count, inverse_indices, 1)
# Compute averages
unique_flattened_embeddings /= unique_flattened_count[:, None]
return unique_flattened_vocab_ids.astype(np.int32), unique_flattened_embeddings.astype(
np.float32
)
def forward(
self, vocab_ids: NumpyArray, embeddings: NumpyArray
) -> tuple[NumpyArray, NumpyArray]:
"""
Args:
vocab_ids: (batch_size, seq_len) int array
embeddings: (batch_size, seq_len, input_dim) float array
Returns:
unique_flattened_vocab_ids_and_batch_ids: (total_unique, 2)
unique_flattened_encoded: (total_unique, output_dim)
"""
# Average embeddings for duplicate vocab_ids
unique_flattened_vocab_ids_and_batch_ids, unique_flattened_embeddings = (
self.avg_by_vocab_ids(vocab_ids, embeddings)
)
# Select the encoder weights for each unique vocab_id
unique_flattened_vocab_ids = unique_flattened_vocab_ids_and_batch_ids[:, 0].astype(
np.int32
)
# unique_encoder_weights: (total_unique, input_dim, output_dim)
unique_encoder_weights = self.encoder_weights[unique_flattened_vocab_ids]
# Compute linear transform: (total_unique, output_dim)
# Using Einstein summation for matrix multiplication:
# 'bi,bio->bo' means: for each "b" (batch element), multiply embeddings (b,i) by weights (b,i,o) -> (b,o)
unique_flattened_encoded = np.einsum(
"bi,bio->bo", unique_flattened_embeddings, unique_encoder_weights
)
# Apply Tanh activation and ensure float32 type
unique_flattened_encoded = self.activation(unique_flattened_encoded).astype(np.float32)
return unique_flattened_vocab_ids_and_batch_ids.astype(np.int32), unique_flattened_encoded
@@ -0,0 +1,244 @@
import copy
from dataclasses import dataclass
import mmh3
import numpy as np
from py_rust_stemmers import SnowballStemmer
from fastembed.common.utils import get_all_punctuation, remove_non_alphanumeric
from fastembed.sparse.sparse_embedding_base import SparseEmbedding
GAP = 32000
INT32_MAX = 2**31 - 1
@dataclass
class WordEmbedding:
word: str
forms: list[str]
count: int
word_id: int
embedding: list[float]
class SparseVectorConverter:
def __init__(
self,
stopwords: set[str],
stemmer: SnowballStemmer,
k: float = 1.2,
b: float = 0.75,
avg_len: float = 150.0,
):
punctuation = set(get_all_punctuation())
special_tokens = {"[CLS]", "[SEP]", "[PAD]", "[UNK]", "[MASK]"}
self.stemmer = stemmer
self.unwanted_tokens = punctuation | special_tokens | stopwords
self.k = k
self.b = b
self.avg_len = avg_len
@classmethod
def unkn_word_token_id(
cls, word: str, shift: int
) -> int: # 2-3 words can collide in 1 index with this mapping, not considering mm3 collisions
token_hash = abs(mmh3.hash(word))
range_size = INT32_MAX - shift
remapped_hash = shift + (token_hash % range_size)
return remapped_hash
def bm25_tf(self, num_occurrences: int, sentence_len: int) -> float:
res = num_occurrences * (self.k + 1)
res /= num_occurrences + self.k * (1 - self.b + self.b * sentence_len / self.avg_len)
return res
@classmethod
def normalize_vector(cls, vector: list[float]) -> list[float]:
norm = sum([x**2 for x in vector]) ** 0.5
if norm < 1e-8:
return vector
return [x / norm for x in vector]
def clean_words(
self, sentence_embedding: dict[str, WordEmbedding], token_max_length: int = 40
) -> dict[str, WordEmbedding]:
"""
Clean miniCOIL-produced sentence_embedding, as unknown to the miniCOIL's stemmer tokens should fully resemble
our BM25 token representation.
sentence_embedding = {"": {"word": "", "word_id": -1, "count": 2, "embedding": [1], "forms": [""]},
"9": {"word": "9", "word_id": -1, "count": 2, "embedding": [1], "forms": ["9"]},
"bat": {"word": "bat", "word_id": 2, "count": 3, "embedding": [0.2, 0.1, -0.2, -0.2], "forms": ["bats", "bat"]},
"9°9": {"word": "9°9", "word_id": -1, "count": 1, "embedding": [1], "forms": ["9°9"]},
"screech": {"word": "screech", "word_id": -1, "count": 1, "embedding": [1], "forms": ["screech"]},
"screeched": {"word": "screeched", "word_id": -1, "count": 1, "embedding": [1], "forms": ["screeched"]}
}
cleaned_embedding_ground_truth = {
"9": {"word": "9", "word_id": -1, "count": 6, "embedding": [1], "forms": ["", "9", "9°9", "9°9"]},
"bat": {"word": "bat", "word_id": 2, "count": 3, "embedding": [0.2, 0.1, -0.2, -0.2], "forms": ["bats", "bat"]},
"screech": {"word": "screech", "word_id": -1, "count": 2, "embedding": [1], "forms": ["screech", "screeched"]}
}
"""
new_sentence_embedding: dict[str, WordEmbedding] = {}
for word, embedding in sentence_embedding.items():
# embedding = {
# "word": "vector",
# "forms": ["vector", "vectors"],
# "count": 2,
# "word_id": 1231,
# "embedding": [0.1, 0.2, 0.3, 0.4]
# }
if embedding.word_id > 0:
# Known word, no need to clean
new_sentence_embedding[word] = embedding
else:
# Unknown word
if word in self.unwanted_tokens:
continue
# Example complex word split:
# word = `word^vec`
word_cleaned = remove_non_alphanumeric(word).strip()
# word_cleaned = `word vec`
if len(word_cleaned) > 0:
# Subwords: ['word', 'vec']
for subword in word_cleaned.split():
stemmed_subword: str = self.stemmer.stem_word(subword)
if (
len(stemmed_subword) <= token_max_length
and stemmed_subword not in self.unwanted_tokens
):
if stemmed_subword not in new_sentence_embedding:
new_sentence_embedding[stemmed_subword] = copy.deepcopy(embedding)
new_sentence_embedding[stemmed_subword].word = stemmed_subword
else:
new_sentence_embedding[stemmed_subword].count += embedding.count
new_sentence_embedding[stemmed_subword].forms += embedding.forms
return new_sentence_embedding
def embedding_to_vector(
self,
sentence_embedding: dict[str, WordEmbedding],
embedding_size: int,
vocab_size: int,
) -> SparseEmbedding:
"""
Convert miniCOIL sentence embedding to Qdrant sparse vector
Example input:
```
{
"vector": WordEmbedding({ // Vocabulary word, encoded with miniCOIL normally
"word": "vector",
"forms": ["vector", "vectors"],
"count": 2,
"word_id": 1231,
"embedding": [0.1, 0.2, 0.3, 0.4]
}),
"axiotic": WordEmbedding({ // Out-of-vocabulary word, fallback to BM25
"word": "axiotic",
"forms": ["axiotics"],
"count": 1,
"word_id": -1,
})
}
```
"""
indices: list[int] = []
values: list[float] = []
# Example:
# vocab_size = 10000
# embedding_size = 4
# GAP = 32000
#
# We want to start random words section from the bucket, that is guaranteed to not
# include any vocab words.
# We need (vocab_size * embedding_size) slots for vocab words.
# Therefore we need (vocab_size * embedding_size) // GAP + 1 buckets for vocab words.
# Therefore, we can start random words from bucket (vocab_size * embedding_size) // GAP + 1 + 1
# ID at which the scope of OOV words starts
unknown_words_shift = ((vocab_size * embedding_size) // GAP + 2) * GAP
sentence_embedding_cleaned = self.clean_words(sentence_embedding)
# Calculate sentence length after cleaning
sentence_len = 0
for embedding in sentence_embedding_cleaned.values():
sentence_len += embedding.count
for embedding in sentence_embedding_cleaned.values():
word_id = embedding.word_id
num_occurrences = embedding.count
tf = self.bm25_tf(num_occurrences, sentence_len)
if (
word_id > 0
): # miniCOIL starts with ID 1, we generally won't have word_id == 0 (UNK), as we don't add
# these words to sentence_embedding
embedding_values = embedding.embedding
normalized_embedding = self.normalize_vector(embedding_values)
for val_id, value in enumerate(normalized_embedding):
indices.append(
word_id * embedding_size + val_id
) # since miniCOIL IDs start with 1
values.append(value * tf)
else:
indices.append(self.unkn_word_token_id(embedding.word, unknown_words_shift))
values.append(tf)
return SparseEmbedding(
indices=np.array(indices, dtype=np.int32),
values=np.array(values, dtype=np.float32),
)
def embedding_to_vector_query(
self,
sentence_embedding: dict[str, WordEmbedding],
embedding_size: int,
vocab_size: int,
) -> SparseEmbedding:
"""
Same as `embedding_to_vector`, but no TF
"""
indices: list[int] = []
values: list[float] = []
# ID at which the scope of OOV words starts
unknown_words_shift = ((vocab_size * embedding_size) // GAP + 2) * GAP
sentence_embedding_cleaned = self.clean_words(sentence_embedding)
for embedding in sentence_embedding_cleaned.values():
word_id = embedding.word_id
tf = 1.0
if word_id >= 0: # miniCOIL starts with ID 1
embedding_values = embedding.embedding
normalized_embedding = self.normalize_vector(embedding_values)
for val_id, value in enumerate(normalized_embedding):
indices.append(
word_id * embedding_size + val_id
) # since miniCOIL IDs start with 1
values.append(value * tf)
else:
indices.append(self.unkn_word_token_id(embedding.word, unknown_words_shift))
values.append(tf)
return SparseEmbedding(
indices=np.array(indices, dtype=np.int32),
values=np.array(values, dtype=np.float32),
)
+202
View File
@@ -0,0 +1,202 @@
from collections import defaultdict
from typing import Iterable
from py_rust_stemmers import SnowballStemmer
import numpy as np
from tokenizers import Tokenizer
from numpy.typing import NDArray
from fastembed.common.types import NumpyArray
class VocabTokenizerBase:
def tokenize(self, sentence: str) -> NumpyArray:
raise NotImplementedError()
def convert_ids_to_tokens(self, token_ids: NumpyArray) -> list[str]:
raise NotImplementedError()
class VocabTokenizer(VocabTokenizerBase):
def __init__(self, tokenizer: Tokenizer):
self.tokenizer = tokenizer
def tokenize(self, sentence: str) -> NumpyArray:
return np.array(self.tokenizer.encode(sentence).ids)
def convert_ids_to_tokens(self, token_ids: NumpyArray) -> list[str]:
return [self.tokenizer.id_to_token(token_id) for token_id in token_ids]
class VocabResolver:
def __init__(self, tokenizer: VocabTokenizerBase, stopwords: set[str], stemmer: SnowballStemmer):
# Word to id mapping
self.vocab: dict[str, int] = {}
# Id to word mapping
self.words: list[str] = []
# Lemma to word mapping
self.stem_mapping: dict[str, str] = {}
self.tokenizer: VocabTokenizerBase = tokenizer
self.stemmer = stemmer
self.stopwords: set[str] = stopwords
def tokenize(self, sentence: str) -> NumpyArray:
return self.tokenizer.tokenize(sentence)
def lookup_word(self, word_id: int) -> str:
if word_id == 0:
return "UNK"
return self.words[word_id - 1]
def convert_ids_to_tokens(self, token_ids: NumpyArray) -> list[str]:
return self.tokenizer.convert_ids_to_tokens(token_ids)
def vocab_size(self) -> int:
# We need +1 for UNK token
return len(self.vocab) + 1
def save_vocab(self, path: str) -> None:
with open(path, "w") as f:
for word in self.words:
f.write(word + "\n")
def save_json_vocab(self, path: str) -> None:
import json
with open(path, "w") as f:
json.dump({"vocab": self.words, "stem_mapping": self.stem_mapping}, f, indent=2)
def load_json_vocab(self, path: str) -> None:
import json
with open(path, "r") as f:
data = json.load(f)
self.words = data["vocab"]
self.vocab = {word: idx + 1 for idx, word in enumerate(self.words)}
self.stem_mapping = data["stem_mapping"]
def add_word(self, word: str) -> None:
if word not in self.vocab:
self.vocab[word] = len(self.vocab) + 1
self.words.append(word)
stem = self.stemmer.stem_word(word)
if stem not in self.stem_mapping:
self.stem_mapping[stem] = word
else:
existing_word = self.stem_mapping[stem]
if len(existing_word) > len(word):
# Prefer shorter words for the same stem
# Example: "swim" is preferred over "swimming"
self.stem_mapping[stem] = word
def load_vocab(self, path: str) -> None:
with open(path, "r") as f:
for line in f:
self.add_word(line.strip())
@classmethod
def _reconstruct_bpe(
cls, bpe_tokens: Iterable[tuple[int, str]]
) -> list[tuple[str, list[int]]]:
result: list[tuple[str, list[int]]] = []
acc: str = ""
acc_idx: list[int] = []
continuing_subword_prefix = "##"
continuing_subword_prefix_len = len(continuing_subword_prefix)
for idx, token in bpe_tokens:
if token.startswith(continuing_subword_prefix):
acc += token[continuing_subword_prefix_len:]
acc_idx.append(idx)
else:
if acc:
result.append((acc, acc_idx))
acc_idx = []
acc = token
acc_idx.append(idx)
if acc:
result.append((acc, acc_idx))
return result
def resolve_tokens(
self, token_ids: NDArray[np.int64]
) -> tuple[NDArray[np.int64], dict[int, int], dict[str, int], dict[str, list[str]]]:
"""
Mark known tokens (including composed tokens) with vocab ids.
Args:
token_ids: (seq_len) - list of ids of tokens
Example:
[
101, 3897, 19332, 12718, 23348,
1010, 1996, 7151, 2296, 4845,
2359, 2005, 4234, 1010, 4332,
2871, 3191, 2062, 102
]
returns:
- token_ids with vocab ids
[
0, 151, 151, 0, 0,
912, 0, 0, 0, 332,
332, 332, 0, 7121, 191,
0, 0, 332, 0
]
- counts of each token
{
151: 1,
332: 3,
7121: 1,
191: 1,
912: 1
}
- oov counts of each token
{
"the": 1,
"a": 1,
"[CLS]": 1,
"[SEP]": 1,
...
}
- forms of each token
{
"hello": ["hello"],
"world": ["worlds", "world", "worlding"],
}
"""
tokens = self.convert_ids_to_tokens(token_ids)
tokens_mapping = self._reconstruct_bpe(enumerate(tokens))
counts: dict[int, int] = defaultdict(int)
oov_count: dict[str, int] = defaultdict(int)
forms: dict[str, list[str]] = defaultdict(list)
for token, mapped_token_ids in tokens_mapping:
vocab_id = 0
if token in self.stopwords:
vocab_id = 0
elif token in self.vocab:
vocab_id = self.vocab[token]
forms[token].append(token)
elif token in self.stem_mapping:
vocab_id = self.vocab[self.stem_mapping[token]]
forms[self.stem_mapping[token]].append(token)
else:
stem = self.stemmer.stem_word(token)
if stem in self.stem_mapping:
vocab_id = self.vocab[self.stem_mapping[stem]]
forms[self.stem_mapping[stem]].append(token)
for token_id in mapped_token_ids:
token_ids[token_id] = vocab_id
if vocab_id == 0:
oov_count[token] += 1
else:
counts[vocab_id] += 1
return token_ids, counts, oov_count, forms
@@ -0,0 +1,69 @@
from typing import Any, Iterable, Type
from fastembed.common.types import NumpyArray
from fastembed.common.onnx_model import OnnxOutputContext
from fastembed.text.onnx_embedding import OnnxTextEmbedding, OnnxTextEmbeddingWorker
from fastembed.common.model_description import DenseModelDescription, ModelSource
supported_builtin_sentence_embedding_models: list[DenseModelDescription] = [
DenseModelDescription(
model="google/embeddinggemma-300m",
dim=768,
description=(
"Text embeddings, Unimodal (text), multilingual, 2048 input tokens truncation, "
"Prefixes for queries/documents: `task: search result | query: {content}` for query, "
"`title: {title | 'none'} | text: {content}` for documents, 2025 year."
),
license="gemma",
size_in_GB=1.24,
sources=ModelSource(
hf="onnx-community/embeddinggemma-300m-ONNX",
),
model_file="onnx/model.onnx",
additional_files=["onnx/model.onnx_data"],
),
]
class BuiltinSentenceEmbedding(OnnxTextEmbedding):
"""Builtin Sentence Embedding uses built-in pooling and normalization of underlying onnx models"""
@classmethod
def _get_worker_class(cls) -> Type[OnnxTextEmbeddingWorker]:
return BuiltinSentenceEmbeddingWorker
@classmethod
def _list_supported_models(cls) -> list[DenseModelDescription]:
"""Lists the supported models.
Returns:
list[DenseModelDescription]: A list of DenseModelDescription objects containing the model information.
"""
return supported_builtin_sentence_embedding_models
def _post_process_onnx_output(
self, output: OnnxOutputContext, **kwargs: Any
) -> Iterable[NumpyArray]:
return output.model_output
def _run_model(
self, onnx_input: dict[str, Any], onnx_output_names: list[str] | None = None
) -> NumpyArray:
return self.model.run(onnx_output_names, onnx_input)[1] # type: ignore[union-attr]
class BuiltinSentenceEmbeddingWorker(OnnxTextEmbeddingWorker):
def init_embedding(
self,
model_name: str,
cache_dir: str,
**kwargs: Any,
) -> OnnxTextEmbedding:
return BuiltinSentenceEmbedding(
model_name=model_name,
cache_dir=cache_dir,
threads=1,
**kwargs,
)
+3 -1
View File
@@ -35,7 +35,9 @@ class CLIPOnnxEmbedding(OnnxTextEmbedding):
"""
return supported_clip_models
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[NumpyArray]:
def _post_process_onnx_output(
self, output: OnnxOutputContext, **kwargs: Any
) -> Iterable[NumpyArray]:
return output.model_output
+68 -15
View File
@@ -1,5 +1,4 @@
from typing import Optional, Sequence, Any, Iterable
from typing import Sequence, Any, Iterable, Type
from dataclasses import dataclass
import numpy as np
@@ -11,9 +10,10 @@ from fastembed.common.model_description import (
DenseModelDescription,
)
from fastembed.common.onnx_model import OnnxOutputContext
from fastembed.common.types import NumpyArray
from fastembed.common.utils import normalize, mean_pooling
from fastembed.common.types import NumpyArray, Device
from fastembed.common.utils import normalize, mean_pooling, last_token_pooling
from fastembed.text.onnx_embedding import OnnxTextEmbedding
from fastembed.text.onnx_text_model import TextEmbeddingWorker
@dataclass(frozen=True)
@@ -29,14 +29,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__(
@@ -51,18 +51,31 @@ class CustomTextEmbedding(OnnxTextEmbedding):
specific_model_path=specific_model_path,
**kwargs,
)
self._pooling = self.POSTPROCESSING_MAPPING[model_name].pooling
self._normalization = self.POSTPROCESSING_MAPPING[model_name].normalization
postprocessing_config = self.POSTPROCESSING_MAPPING[self.model_description.model]
self._pooling = postprocessing_config.pooling
self._normalization = postprocessing_config.normalization
@classmethod
def _list_supported_models(cls) -> list[DenseModelDescription]:
return cls.SUPPORTED_MODELS
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[NumpyArray]:
@classmethod
def _get_worker_class(cls) -> Type["TextEmbeddingWorker[NumpyArray]"]:
return CustomTextEmbeddingWorker
def _get_worker_init_kwargs(self) -> dict[str, Any]:
return {
"model_description": self.model_description,
"postprocessing_config": self.POSTPROCESSING_MAPPING[self.model_description.model],
}
def _post_process_onnx_output(
self, output: OnnxOutputContext, **kwargs: Any
) -> Iterable[NumpyArray]:
return self._normalize(self._pool(output.model_output, output.attention_mask))
def _pool(
self, embeddings: NumpyArray, attention_mask: Optional[NDArray[np.int64]] = None
self, embeddings: NumpyArray, attention_mask: NDArray[np.int64] | None = None
) -> NumpyArray:
if self._pooling == PoolingType.CLS:
return embeddings[:, 0]
@@ -72,9 +85,20 @@ class CustomTextEmbedding(OnnxTextEmbedding):
raise ValueError("attention_mask must be provided for mean pooling")
return mean_pooling(embeddings, attention_mask)
if self._pooling == PoolingType.LAST_TOKEN:
if attention_mask is None:
raise ValueError("attention_mask must be provided for last token pooling")
return last_token_pooling(embeddings, attention_mask)
if self._pooling == PoolingType.DISABLED:
return embeddings
raise ValueError(
f"Unsupported pooling type {self._pooling}. "
f"Supported types are: {PoolingType.CLS}, {PoolingType.MEAN}, "
f"{PoolingType.LAST_TOKEN}, {PoolingType.DISABLED}."
)
def _normalize(self, embeddings: NumpyArray) -> NumpyArray:
return normalize(embeddings) if self._normalization else embeddings
@@ -89,3 +113,32 @@ class CustomTextEmbedding(OnnxTextEmbedding):
cls.POSTPROCESSING_MAPPING[model_description.model] = PostprocessingConfig(
pooling=pooling, normalization=normalization
)
class CustomTextEmbeddingWorker(TextEmbeddingWorker[NumpyArray]):
def init_embedding(
self,
model_name: str,
cache_dir: str,
model_description: DenseModelDescription | None = None,
postprocessing_config: PostprocessingConfig | None = None,
**kwargs: Any,
) -> CustomTextEmbedding:
if model_description is None or postprocessing_config is None:
raise ValueError(
"`model_description` and `postprocessing_config` are required to initialize a "
"custom model in a worker process, they are provided by "
"`CustomTextEmbedding._get_worker_init_kwargs`"
)
# custom models live in a class-level registry, which spawned workers don't inherit
CustomTextEmbedding.add_model(
model_description,
pooling=postprocessing_config.pooling,
normalization=postprocessing_config.normalization,
)
return CustomTextEmbedding(
model_name=model_name,
cache_dir=cache_dir,
threads=1,
**kwargs,
)
@@ -0,0 +1,94 @@
from typing import Any, Iterable, Type
import onnxruntime as ort
from fastembed.common.types import NumpyArray
from fastembed.common.onnx_model import OnnxOutputContext
from fastembed.common.utils import last_token_pooling, normalize
from fastembed.text.onnx_embedding import OnnxTextEmbedding, OnnxTextEmbeddingWorker
from fastembed.common.model_description import DenseModelDescription, ModelSource
supported_last_token_normalized_models: list[DenseModelDescription] = [
DenseModelDescription(
model="Qwen/Qwen3-Embedding-0.6B",
dim=1024,
description=(
"Text embeddings, Unimodal (text), multilingual, 32768 input tokens truncation, "
"Prefixes for queries/documents: `Instruct: {task_description}\\nQuery:{query}` "
"for queries, none for documents, 2025 year."
),
license="apache-2.0",
size_in_GB=2.38,
sources=ModelSource(hf="Qdrant/Qwen3-Embedding-0.6B-onnx"),
model_file="onnx/model.onnx",
additional_files=["onnx/model.onnx_data"],
),
DenseModelDescription(
model="Qwen/Qwen3-Embedding-0.6B-Q",
dim=1024,
description=(
"Text embeddings, Unimodal (text), multilingual, 32768 input tokens truncation, "
"Prefixes for queries/documents: `Instruct: {task_description}\\nQuery:{query}` "
"for queries, none for documents, int8 weights, requires onnxruntime>=1.23, "
"2025 year."
),
license="apache-2.0",
size_in_GB=1.12,
sources=ModelSource(hf="Qdrant/Qwen3-Embedding-0.6B-onnx"),
model_file="onnx/model_quantized.onnx",
),
]
class LastTokenNormalizedEmbedding(OnnxTextEmbedding):
"""Decoder-based embedding models, which pool the last non-padding token and normalize it"""
@classmethod
def _get_worker_class(cls) -> Type[OnnxTextEmbeddingWorker]:
return LastTokenNormalizedEmbeddingWorker
def load_onnx_model(self) -> None:
try:
super().load_onnx_model()
except Exception as e:
# int8 weights are stored as 8-bit MatMulNBits, which onnxruntime only
# implements since 1.23; older versions fail with "nbits_ == 4 was false"
if "nbits" not in str(e).lower():
raise
raise RuntimeError(
f"Could not load {self.model_name}: its int8 weights require "
f"onnxruntime>=1.23, but onnxruntime {ort.__version__} is installed. "
f"Either upgrade onnxruntime or use a non-quantized model."
) from e
@classmethod
def _list_supported_models(cls) -> list[DenseModelDescription]:
"""Lists the supported models.
Returns:
list[DenseModelDescription]: A list of DenseModelDescription objects containing the model information.
"""
return supported_last_token_normalized_models
def _post_process_onnx_output(
self, output: OnnxOutputContext, **kwargs: Any
) -> Iterable[NumpyArray]:
if output.attention_mask is None:
raise ValueError("attention_mask must be provided for last token pooling")
return normalize(last_token_pooling(output.model_output, output.attention_mask))
class LastTokenNormalizedEmbeddingWorker(OnnxTextEmbeddingWorker):
def init_embedding(
self,
model_name: str,
cache_dir: str,
**kwargs: Any,
) -> OnnxTextEmbedding:
return LastTokenNormalizedEmbedding(
model_name=model_name,
cache_dir=cache_dir,
threads=1,
**kwargs,
)
+28 -19
View File
@@ -1,8 +1,9 @@
from enum import Enum
from typing import Any, Type, Iterable, Union, Optional
from typing import Any, Type, Iterable
import numpy as np
from fastembed.common.onnx_model import OnnxOutputContext
from fastembed.common.types import NumpyArray
from fastembed.text.pooled_normalized_embedding import PooledNormalizedEmbedding
from fastembed.text.onnx_embedding import OnnxTextEmbeddingWorker
@@ -44,9 +45,9 @@ class JinaEmbeddingV3(PooledNormalizedEmbedding):
PASSAGE_TASK = Task.RETRIEVAL_PASSAGE
QUERY_TASK = Task.RETRIEVAL_QUERY
def __init__(self, *args: Any, **kwargs: Any):
def __init__(self, *args: Any, task_id: int | None = None, **kwargs: Any):
super().__init__(*args, **kwargs)
self.current_task_id: Union[Task, int] = self.PASSAGE_TASK
self.default_task_id: Task | int = task_id if task_id is not None else self.PASSAGE_TASK
@classmethod
def _get_worker_class(cls) -> Type[OnnxTextEmbeddingWorker]:
@@ -57,30 +58,34 @@ class JinaEmbeddingV3(PooledNormalizedEmbedding):
return supported_multitask_models
def _preprocess_onnx_input(
self, onnx_input: dict[str, NumpyArray], **kwargs: Any
self,
onnx_input: dict[str, NumpyArray],
task_id: int | Task | None = None,
**kwargs: Any,
) -> dict[str, NumpyArray]:
onnx_input["task_id"] = np.array(self.current_task_id, dtype=np.int64)
if task_id is None:
raise ValueError(f"task_id must be provided for JinaEmbeddingV3, got <{task_id}>")
onnx_input["task_id"] = np.array(task_id, dtype=np.int64)
return onnx_input
def embed(
self,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
task_id: int = PASSAGE_TASK,
parallel: int | None = None,
task_id: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
self.current_task_id = task_id
kwargs["task_id"] = task_id
yield from super().embed(documents, batch_size, parallel, **kwargs)
task_id = (
task_id if task_id is not None else self.default_task_id
) # required for multiprocessing
yield from super().embed(documents, batch_size, parallel, task_id=task_id, **kwargs)
def query_embed(self, query: Union[str, Iterable[str]], **kwargs: Any) -> Iterable[NumpyArray]:
self.current_task_id = self.QUERY_TASK
yield from super().embed(query, **kwargs)
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:
yield from super().embed(query, task_id=self.QUERY_TASK, **kwargs)
def passage_embed(self, texts: Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:
self.current_task_id = self.PASSAGE_TASK
yield from super().embed(texts, **kwargs)
yield from super().embed(texts, task_id=self.PASSAGE_TASK, **kwargs)
class JinaEmbeddingV3Worker(OnnxTextEmbeddingWorker):
@@ -90,11 +95,15 @@ class JinaEmbeddingV3Worker(OnnxTextEmbeddingWorker):
cache_dir: str,
**kwargs: Any,
) -> JinaEmbeddingV3:
model = JinaEmbeddingV3(
return JinaEmbeddingV3(
model_name=model_name,
cache_dir=cache_dir,
threads=1,
**kwargs,
)
model.current_task_id = kwargs["task_id"]
return model
def process(self, items: Iterable[tuple[int, Any]]) -> Iterable[tuple[int, OnnxOutputContext]]:
self.model: JinaEmbeddingV3 # mypy complaints `self.model` does not have `default_task_id`
for idx, batch in items:
onnx_output = self.model.onnx_embed(batch, task_id=self.model.default_task_id)
yield idx, onnx_output
+75 -23
View File
@@ -1,12 +1,11 @@
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Iterable, Sequence, Type
import numpy as np
from fastembed.common.types import NumpyArray, OnnxProvider
from fastembed.common.model_description import DenseModelDescription, ModelSource
from fastembed.common.onnx_model import OnnxOutputContext
from fastembed.common.types import Device, NumpyArray, OnnxProvider
from fastembed.common.utils import define_cache_dir, normalize
from fastembed.text.onnx_text_model import OnnxTextModel, TextEmbeddingWorker
from fastembed.text.text_embedding_base import TextEmbeddingBase
from fastembed.common.model_description import DenseModelDescription, ModelSource
supported_onnx_models: list[DenseModelDescription] = [
DenseModelDescription(
@@ -21,6 +20,7 @@ supported_onnx_models: list[DenseModelDescription] = [
sources=ModelSource(
hf="Qdrant/fast-bge-base-en",
url="https://storage.googleapis.com/qdrant-fastembed/fast-bge-base-en.tar.gz",
_deprecated_tar_struct=True,
),
model_file="model_optimized.onnx",
),
@@ -34,8 +34,9 @@ supported_onnx_models: list[DenseModelDescription] = [
license="mit",
size_in_GB=0.21,
sources=ModelSource(
hf="qdrant/bge-base-en-v1.5-onnx-q",
hf="Qdrant/bge-base-en-v1.5-onnx-Q",
url="https://storage.googleapis.com/qdrant-fastembed/fast-bge-base-en-v1.5.tar.gz",
_deprecated_tar_struct=True,
),
model_file="model_optimized.onnx",
),
@@ -63,6 +64,7 @@ supported_onnx_models: list[DenseModelDescription] = [
sources=ModelSource(
hf="Qdrant/bge-small-en",
url="https://storage.googleapis.com/qdrant-fastembed/BAAI-bge-small-en.tar.gz",
_deprecated_tar_struct=True,
),
model_file="model_optimized.onnx",
),
@@ -75,7 +77,7 @@ supported_onnx_models: list[DenseModelDescription] = [
),
license="mit",
size_in_GB=0.067,
sources=ModelSource(hf="qdrant/bge-small-en-v1.5-onnx-q"),
sources=ModelSource(hf="Qdrant/bge-small-en-v1.5-onnx-Q"),
model_file="model_optimized.onnx",
),
DenseModelDescription(
@@ -90,6 +92,7 @@ supported_onnx_models: list[DenseModelDescription] = [
sources=ModelSource(
hf="Qdrant/bge-small-zh-v1.5",
url="https://storage.googleapis.com/qdrant-fastembed/fast-bge-small-zh-v1.5.tar.gz",
_deprecated_tar_struct=True,
),
model_file="model_optimized.onnx",
),
@@ -177,6 +180,42 @@ supported_onnx_models: list[DenseModelDescription] = [
sources=ModelSource(hf="jinaai/jina-clip-v1"),
model_file="onnx/text_model.onnx",
),
DenseModelDescription(
model="minishlab/potion-base-8M",
dim=256,
description=(
"Text embeddings, Unimodal (text), English, 512 input tokens truncation, "
"Prefixes for queries/documents: not necessary, 2024 year."
),
license="mit",
size_in_GB=0.030,
sources=ModelSource(hf="minishlab/potion-base-8m-onnx"),
model_file="model.onnx",
),
DenseModelDescription(
model="minishlab/potion-retrieval-32M",
dim=512,
description=(
"Text embeddings, Unimodal (text), English, 512 input tokens truncation, "
"Prefixes for queries/documents: not necessary, 2025 year."
),
license="mit",
size_in_GB=0.129,
sources=ModelSource(hf="minishlab/potion-retrieval-32m-onnx"),
model_file="model.onnx",
),
DenseModelDescription(
model="minishlab/potion-multilingual-128M",
dim=256,
description=(
"Text embeddings, Unimodal (text), Multilingual, 512 input tokens truncation, "
"Prefixes for queries/documents: not necessary, 2025 year."
),
license="mit",
size_in_GB=0.512,
sources=ModelSource(hf="minishlab/potion-multilingual-128m-onnx"),
model_file="model.onnx",
),
]
@@ -196,14 +235,14 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[NumpyArray]):
def __init__(
self,
model_name: str = "BAAI/bge-small-en-v1.5",
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
lazy_load: bool = False,
device_id: Optional[int] = None,
specific_model_path: Optional[str] = None,
device_id: int | None = None,
specific_model_path: str | None = None,
**kwargs: Any,
):
"""
@@ -215,10 +254,11 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[NumpyArray]):
threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.
providers (Optional[Sequence[OnnxProvider]], optional): The list of onnxruntime providers to use.
Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to False.
cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`
Defaults to Device.AUTO.
device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, mutually exclusive
with `providers`. Defaults to None.
lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
@@ -230,13 +270,13 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[NumpyArray]):
super().__init__(model_name, cache_dir, threads, **kwargs)
self.providers = providers
self.lazy_load = lazy_load
self._extra_session_options = self._select_exposed_session_options(kwargs)
# List of device ids, that can be used for data parallel processing in workers
self.device_ids = device_ids
self.cuda = cuda
# This device_id will be used if we need to load model in current process
self.device_id: Optional[int] = None
self.device_id: int | None = None
if device_id is not None:
self.device_id = device_id
elif self.device_ids is not None:
@@ -244,11 +284,12 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[NumpyArray]):
self.model_description = self._get_model_description(model_name)
self.cache_dir = str(define_cache_dir(cache_dir))
self._specific_model_path = specific_model_path
self._model_dir = self.download_model(
self.model_description,
self.cache_dir,
local_files_only=self._local_files_only,
specific_model_path=specific_model_path,
specific_model_path=self._specific_model_path,
)
if not self.lazy_load:
@@ -256,9 +297,9 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[NumpyArray]):
def embed(
self,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
"""
@@ -285,6 +326,9 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[NumpyArray]):
providers=self.providers,
cuda=self.cuda,
device_ids=self.device_ids,
local_files_only=self._local_files_only,
specific_model_path=self._specific_model_path,
extra_session_options=self._extra_session_options,
**kwargs,
)
@@ -300,7 +344,9 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[NumpyArray]):
"""
return onnx_input
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[NumpyArray]:
def _post_process_onnx_output(
self, output: OnnxOutputContext, **kwargs: Any
) -> Iterable[NumpyArray]:
embeddings = output.model_output
if embeddings.ndim == 3: # (batch_size, seq_len, embedding_dim)
@@ -309,7 +355,7 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[NumpyArray]):
processed_embeddings = embeddings
else:
raise ValueError(f"Unsupported embedding shape: {embeddings.shape}")
return normalize(processed_embeddings).astype(np.float32)
return normalize(processed_embeddings)
def load_onnx_model(self) -> None:
self._load_onnx_model(
@@ -319,8 +365,14 @@ class OnnxTextEmbedding(TextEmbeddingBase, OnnxTextModel[NumpyArray]):
providers=self.providers,
cuda=self.cuda,
device_id=self.device_id,
extra_session_options=self._extra_session_options,
)
def token_count(
self, texts: str | Iterable[str], batch_size: int = 1024, **kwargs: Any
) -> int:
return self._token_count(texts, batch_size=batch_size, **kwargs)
class OnnxTextEmbeddingWorker(TextEmbeddingWorker[NumpyArray]):
def init_embedding(
+61 -19
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,23 +15,32 @@ from fastembed.parallel_processor import ParallelWorkerPool
class OnnxTextModel(OnnxModel[T]):
ONNX_OUTPUT_NAMES: Optional[list[str]] = None
ONNX_OUTPUT_NAMES: list[str] | None = None
@classmethod
def _get_worker_class(cls) -> Type["TextEmbeddingWorker[T]"]:
raise NotImplementedError("Subclasses must implement this method")
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[T]:
def _post_process_onnx_output(self, output: OnnxOutputContext, **kwargs: Any) -> Iterable[T]:
"""Post-process the ONNX model output to convert it into a usable format.
Args:
output (OnnxOutputContext): The raw output from the ONNX model.
**kwargs: Additional keyword arguments that may be needed by specific implementations.
Returns:
Iterable[T]: Post-processed output as an iterable of type T.
"""
raise NotImplementedError("Subclasses must implement this method")
def __init__(self) -> None:
super().__init__()
self.tokenizer: Optional[Tokenizer] = None
self.tokenizer: Tokenizer | None = None
self.special_token_to_id: dict[str, int] = {}
def _preprocess_onnx_input(
self, onnx_input: dict[str, NumpyArray], **kwargs: Any
) -> dict[str, Union[NumpyArray, NDArray[np.int64]]]:
) -> dict[str, NumpyArray | NDArray[np.int64]]:
"""
Preprocess the onnx input.
"""
@@ -41,10 +50,11 @@ class OnnxTextModel(OnnxModel[T]):
self,
model_dir: Path,
model_file: str,
threads: Optional[int],
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_id: Optional[int] = None,
threads: int | None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_id: int | None = None,
extra_session_options: dict[str, Any] | None = None,
) -> None:
super()._load_onnx_model(
model_dir=model_dir,
@@ -53,6 +63,7 @@ class OnnxTextModel(OnnxModel[T]):
providers=providers,
cuda=cuda,
device_id=device_id,
extra_session_options=extra_session_options,
)
self.tokenizer, self.special_token_to_id = load_tokenizer(model_dir=model_dir)
@@ -81,24 +92,34 @@ class OnnxTextModel(OnnxModel[T]):
[np.zeros(len(e), dtype=np.int64) for e in input_ids], dtype=np.int64
)
onnx_input = self._preprocess_onnx_input(onnx_input, **kwargs)
model_output = self._run_model(
onnx_input=onnx_input, onnx_output_names=self.ONNX_OUTPUT_NAMES
)
model_output = self.model.run(self.ONNX_OUTPUT_NAMES, onnx_input) # type: ignore[union-attr]
return OnnxOutputContext(
model_output=model_output[0],
model_output=model_output,
attention_mask=onnx_input.get("attention_mask", attention_mask),
input_ids=onnx_input.get("input_ids", input_ids),
)
def _run_model(
self, onnx_input: dict[str, Any], onnx_output_names: list[str] | None = None
) -> NumpyArray:
return self.model.run(onnx_output_names, onnx_input)[0] # type: ignore[union-attr]
def _embed_documents(
self,
model_name: str,
cache_dir: str,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
providers: Optional[Sequence[OnnxProvider]] = None,
cuda: bool = False,
device_ids: Optional[list[int]] = None,
parallel: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
cuda: bool | Device = Device.AUTO,
device_ids: list[int] | None = None,
local_files_only: bool = False,
specific_model_path: str | None = None,
extra_session_options: dict[str, Any] | None = None,
**kwargs: Any,
) -> Iterable[T]:
is_small = False
@@ -115,7 +136,9 @@ class OnnxTextModel(OnnxModel[T]):
if not hasattr(self, "model") or self.model is None:
self.load_onnx_model()
for batch in iter_batch(documents, batch_size):
yield from self._post_process_onnx_output(self.onnx_embed(batch))
yield from self._post_process_onnx_output(
self.onnx_embed(batch, **kwargs), **kwargs
)
else:
if parallel == 0:
parallel = os.cpu_count()
@@ -125,9 +148,15 @@ class OnnxTextModel(OnnxModel[T]):
"model_name": model_name,
"cache_dir": cache_dir,
"providers": providers,
"local_files_only": local_files_only,
"specific_model_path": specific_model_path,
**kwargs,
**self._get_worker_init_kwargs(),
}
if extra_session_options is not None:
params.update(extra_session_options)
pool = ParallelWorkerPool(
num_workers=parallel or 1,
worker=self._get_worker_class(),
@@ -136,7 +165,20 @@ class OnnxTextModel(OnnxModel[T]):
start_method=start_method,
)
for batch in pool.ordered_map(iter_batch(documents, batch_size), **params):
yield from self._post_process_onnx_output(batch) # type: ignore
yield from self._post_process_onnx_output(batch, **kwargs) # type: ignore
def _token_count(self, texts: str | Iterable[str], batch_size: int = 1024, **_: Any) -> int:
if not hasattr(self, "model") or self.model is None:
self.load_onnx_model() # loads the tokenizer as well
token_num = 0
assert self.tokenizer is not None
texts = [texts] if isinstance(texts, str) else texts
for batch in iter_batch(texts, batch_size):
for tokens in self.tokenizer.encode_batch(batch):
token_num += sum(tokens.attention_mask)
return token_num
class TextEmbeddingWorker(EmbeddingWorker[T]):
+5 -2
View File
@@ -82,6 +82,7 @@ supported_pooled_models: list[DenseModelDescription] = [
sources=ModelSource(
hf="qdrant/multilingual-e5-large-onnx",
url="https://storage.googleapis.com/qdrant-fastembed/fast-multilingual-e5-large.tar.gz",
_deprecated_tar_struct=True,
),
model_file="model.onnx",
additional_files=["model.onnx_data"],
@@ -109,13 +110,15 @@ class PooledEmbedding(OnnxTextEmbedding):
"""
return supported_pooled_models
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[NumpyArray]:
def _post_process_onnx_output(
self, output: OnnxOutputContext, **kwargs: Any
) -> Iterable[NumpyArray]:
if output.attention_mask is None:
raise ValueError("attention_mask must be provided for document post-processing")
embeddings = output.model_output
attn_mask = output.attention_mask
return self.mean_pooling(embeddings, attn_mask).astype(np.float32)
return self.mean_pooling(embeddings, attn_mask)
class PooledEmbeddingWorker(OnnxTextEmbeddingWorker):
@@ -1,6 +1,5 @@
from typing import Any, Iterable, Type
import numpy as np
from fastembed.common.types import NumpyArray
from fastembed.common.onnx_model import OnnxOutputContext
@@ -22,6 +21,7 @@ supported_pooled_normalized_models: list[DenseModelDescription] = [
sources=ModelSource(
url="https://storage.googleapis.com/qdrant-fastembed/sentence-transformers-all-MiniLM-L6-v2.tar.gz",
hf="qdrant/all-MiniLM-L6-v2-onnx",
_deprecated_tar_struct=True,
),
model_file="model.onnx",
),
@@ -57,9 +57,9 @@ supported_pooled_normalized_models: list[DenseModelDescription] = [
"Prefixes for queries/documents: not necessary, 2024 year."
),
license="apache-2.0",
size_in_GB=0.32,
size_in_GB=0.64,
sources=ModelSource(hf="jinaai/jina-embeddings-v2-base-de"),
model_file="onnx/model_fp16.onnx",
model_file="onnx/model.onnx",
),
DenseModelDescription(
model="jinaai/jina-embeddings-v2-base-code",
@@ -138,13 +138,15 @@ class PooledNormalizedEmbedding(PooledEmbedding):
"""
return supported_pooled_normalized_models
def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[NumpyArray]:
def _post_process_onnx_output(
self, output: OnnxOutputContext, **kwargs: Any
) -> Iterable[NumpyArray]:
if output.attention_mask is None:
raise ValueError("attention_mask must be provided for document post-processing")
embeddings = output.model_output
attn_mask = output.attention_mask
return normalize(self.mean_pooling(embeddings, attn_mask)).astype(np.float32)
return normalize(self.mean_pooling(embeddings, attn_mask))
class PooledNormalizedEmbeddingWorker(OnnxTextEmbeddingWorker):
+71
View File
@@ -0,0 +1,71 @@
from typing import Any, Type
from fastembed.common.model_description import DenseModelDescription, ModelSource
from fastembed.text.onnx_embedding import OnnxTextEmbedding, OnnxTextEmbeddingWorker
supported_siglip_models: list[DenseModelDescription] = [
DenseModelDescription(
model="google/siglip2-base-patch16-224",
dim=768,
description=(
"Text embeddings, Multimodal (text&image), multilingual, 64 input tokens truncation, "
"Prefixes for queries/documents: not necessary, 2025 year"
),
license="apache-2.0",
size_in_GB=1.13,
sources=ModelSource(hf="onnx-community/siglip2-base-patch16-224-ONNX"),
model_file="onnx/text_model.onnx",
),
]
class SiglipOnnxTextEmbedding(OnnxTextEmbedding):
"""SigLIP text tower.
SigLIP always pools the hidden state at the last sequence position, whether or not it is
padding, so every batch must be padded to the model's fixed training length (rather than to
the longest sequence in the batch) or the resulting embeddings become dependent on what else
is in the batch.
The exported graph also returns both `last_hidden_state` and `pooler_output`; only the
latter is the text embedding, so it must be selected explicitly.
"""
ONNX_OUTPUT_NAMES = ["pooler_output"]
@classmethod
def _get_worker_class(cls) -> Type[OnnxTextEmbeddingWorker]:
return SiglipTextEmbeddingWorker
@classmethod
def _list_supported_models(cls) -> list[DenseModelDescription]:
return supported_siglip_models
def load_onnx_model(self) -> None:
super().load_onnx_model()
if self.tokenizer is not None:
truncation = self.tokenizer.truncation
padding = self.tokenizer.padding
if truncation and padding and padding.get("length") is None:
self.tokenizer.enable_padding(
direction=padding["direction"],
pad_id=padding["pad_id"],
pad_type_id=padding["pad_type_id"],
pad_token=padding["pad_token"],
length=truncation["max_length"],
)
class SiglipTextEmbeddingWorker(OnnxTextEmbeddingWorker):
def init_embedding(
self,
model_name: str,
cache_dir: str,
**kwargs: Any,
) -> OnnxTextEmbedding:
return SiglipOnnxTextEmbedding(
model_name=model_name,
cache_dir=cache_dir,
threads=1,
**kwargs,
)
+68 -28
View File
@@ -1,14 +1,17 @@
import warnings
from typing import Any, Iterable, Optional, Sequence, Type, Union
from typing import Any, Iterable, Sequence, Type
from dataclasses import asdict
from fastembed.common.types import NumpyArray, OnnxProvider
from fastembed.common.types import NumpyArray, OnnxProvider, Device
from fastembed.text.clip_embedding import CLIPOnnxEmbedding
from fastembed.text.custom_text_embedding import CustomTextEmbedding
from fastembed.text.pooled_normalized_embedding import PooledNormalizedEmbedding
from fastembed.text.pooled_embedding import PooledEmbedding
from fastembed.text.multitask_embedding import JinaEmbeddingV3
from fastembed.text.builtin_sentence_embedding import BuiltinSentenceEmbedding
from fastembed.text.last_token_normalized_embedding import LastTokenNormalizedEmbedding
from fastembed.text.onnx_embedding import OnnxTextEmbedding
from fastembed.text.siglip_embedding import SiglipOnnxTextEmbedding
from fastembed.text.text_embedding_base import TextEmbeddingBase
from fastembed.common.model_description import DenseModelDescription, ModelSource, PoolingType
@@ -17,9 +20,12 @@ class TextEmbedding(TextEmbeddingBase):
EMBEDDINGS_REGISTRY: list[Type[TextEmbeddingBase]] = [
OnnxTextEmbedding,
CLIPOnnxEmbedding,
SiglipOnnxTextEmbedding,
PooledNormalizedEmbedding,
PooledEmbedding,
JinaEmbeddingV3,
BuiltinSentenceEmbedding,
LastTokenNormalizedEmbedding,
CustomTextEmbedding,
]
@@ -51,11 +57,11 @@ 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:
if model == registered_model.model:
if model.lower() == registered_model.model.lower():
raise ValueError(
f"Model {model} is already registered in TextEmbedding, if you still want to add this model, "
f"please use another model name"
@@ -79,32 +85,18 @@ 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,
):
super().__init__(model_name, cache_dir, threads, **kwargs)
if model_name == "nomic-ai/nomic-embed-text-v1.5-Q":
if model_name.lower() == "jinaai/jina-embeddings-v2-base-de":
warnings.warn(
"The model 'nomic-ai/nomic-embed-text-v1.5-Q' has been updated on HuggingFace. Please review "
"the latest documentation on HF and release notes to ensure compatibility with your workflow. ",
UserWarning,
stacklevel=2,
)
if model_name in {
"sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2",
"thenlper/gte-large",
"intfloat/multilingual-e5-large",
"sentence-transformers/paraphrase-multilingual-mpnet-base-v2",
}:
warnings.warn(
f"The model {model_name} now uses mean pooling instead of CLS embedding. "
f"In order to preserve the previous behaviour, consider either pinning fastembed version to 0.5.1 or "
"using `add_custom_model` functionality.",
"The model 'jinaai/jina-embeddings-v2-base-de' used to run with fp16 model, but due to onnxruntime updates, now it runs with the original fp32 model.",
UserWarning,
stacklevel=2,
)
@@ -128,11 +120,45 @@ class TextEmbedding(TextEmbeddingBase):
"Please check the supported models using `TextEmbedding.list_supported_models()`"
)
@property
def embedding_size(self) -> int:
"""Get the embedding size of the current model"""
if self._embedding_size is None:
self._embedding_size = self.get_embedding_size(self.model_name)
return self._embedding_size
@classmethod
def get_embedding_size(cls, model_name: str) -> int:
"""Get the embedding size of the passed model
Args:
model_name (str): The name of the model to get embedding size for.
Returns:
int: The size of the embedding.
Raises:
ValueError: If the model name is not found in the supported models.
"""
descriptions = cls._list_supported_models()
embedding_size: int | None = None
for description in descriptions:
if description.model.lower() == model_name.lower():
embedding_size = description.dim
break
if embedding_size is None:
model_names = [description.model for description in descriptions]
raise ValueError(
f"Embedding size for model {model_name} was None. "
f"Available model names: {model_names}"
)
return embedding_size
def embed(
self,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
"""
@@ -152,7 +178,7 @@ class TextEmbedding(TextEmbeddingBase):
"""
yield from self.model.embed(documents, batch_size, parallel, **kwargs)
def query_embed(self, query: Union[str, Iterable[str]], **kwargs: Any) -> Iterable[NumpyArray]:
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:
"""
Embeds queries
@@ -178,3 +204,17 @@ class TextEmbedding(TextEmbeddingBase):
"""
# This is model-specific, so that different models can have specialized implementations
yield from self.model.passage_embed(texts, **kwargs)
def token_count(
self, texts: str | Iterable[str], batch_size: int = 1024, **kwargs: Any
) -> int:
"""Returns the number of tokens in the texts.
Args:
texts (str | Iterable[str]): The list of texts to embed.
batch_size (int): Batch size for encoding
Returns:
int: Sum of number of tokens in the texts.
"""
return self.model.token_count(texts, batch_size=batch_size, **kwargs)
+21 -6
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,20 +9,21 @@ class TextEmbeddingBase(ModelManagement[DenseModelDescription]):
def __init__(
self,
model_name: str,
cache_dir: Optional[str] = None,
threads: Optional[int] = None,
cache_dir: str | None = None,
threads: int | None = None,
**kwargs: Any,
):
self.model_name = model_name
self.cache_dir = cache_dir
self.threads = threads
self._local_files_only = kwargs.pop("local_files_only", False)
self._embedding_size: int | None = None
def embed(
self,
documents: Union[str, Iterable[str]],
documents: str | Iterable[str],
batch_size: int = 256,
parallel: Optional[int] = None,
parallel: int | None = None,
**kwargs: Any,
) -> Iterable[NumpyArray]:
raise NotImplementedError()
@@ -42,7 +43,7 @@ class TextEmbeddingBase(ModelManagement[DenseModelDescription]):
# This is model-specific, so that different models can have specialized implementations
yield from self.embed(texts, **kwargs)
def query_embed(self, query: Union[str, Iterable[str]], **kwargs: Any) -> Iterable[NumpyArray]:
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:
"""
Embeds queries
@@ -58,3 +59,17 @@ class TextEmbeddingBase(ModelManagement[DenseModelDescription]):
yield from self.embed([query], **kwargs)
else:
yield from self.embed(query, **kwargs)
@classmethod
def get_embedding_size(cls, model_name: str) -> int:
"""Returns embedding size of the passed model."""
raise NotImplementedError("Subclasses must implement this method")
@property
def embedding_size(self) -> int:
"""Returns embedding size for the current model"""
raise NotImplementedError("Subclasses must implement this method")
def token_count(self, texts: str | Iterable[str], **kwargs: Any) -> int:
"""Returns the number of tokens in the texts."""
raise NotImplementedError("Subclasses must implement this method")
+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
+4510
View File
File diff suppressed because it is too large Load Diff
+28 -18
View File
@@ -1,9 +1,9 @@
[tool.poetry]
name = "fastembed"
version = "0.6.0"
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"
license = "Apache-2.0"
readme = "README.md"
packages = [{include = "fastembed"}]
homepage = "https://github.com/qdrant/fastembed"
@@ -11,46 +11,56 @@ repository = "https://github.com/qdrant/fastembed"
keywords = ["vector", "embedding", "neural", "search", "qdrant", "sentence-transformers"]
[tool.poetry.dependencies]
python = ">=3.9.0"
python = ">=3.10.0"
numpy = [
{ version = ">=1.21", python = ">=3.10,<3.12" },
{ version = ">=1.26", python = ">=3.12,<3.13" },
{ version = ">=2.1.0", python = ">=3.13" },
{ version = ">=1.21,<2.1.0", python = "<3.10" },
{ version = ">=1.21,<2.3.0", python = "3.10" },
{ version = ">=1.21", python = "3.11" },
{ version = ">=1.26", python = "3.12" },
{ version = ">=2.1.0", python = "3.13" },
{ version = ">=2.3.0", python = ">=3.14" },
]
onnxruntime = [
{ version = ">1.20.0", python = ">=3.13" },
{ version = ">=1.17.0,<1.20.0", python = "<3.10" },
{ version = ">=1.17.0,!=1.20.0", python = ">=3.10,<3.13" },
{ version = ">=1.17.0,!=1.20.0,<1.24", python = "3.10" },
{ version = ">=1.17.0,!=1.20.0,!=1.24.0,!=1.24.1", python = ">=3.11,<3.13" },
{ version = ">1.21.0,!=1.24.0,!=1.24.1", python = "3.13" },
{ version = ">=1.24.2", python = ">=3.14" },
]
tqdm = "^4.66"
requests = "^2.31"
tokenizers = ">=0.15,<1.0"
huggingface-hub = ">=0.20,<1.0"
huggingface-hub = ">=0.20,<2.0"
loguru = "^0.7.2"
pillow = ">=10.3.0,<12.0.0"
pillow = [
{ version = ">=10.3.0,<13.0", python = ">=3.10,<3.13" },
{ version = ">=11.0.0,<13.0", python = "3.13" },
{ version = ">=12.0.0,<13.0", python = ">=3.14" },
]
mmh3 = ">=4.1.0,<6.0.0"
py-rust-stemmers = "^0.1.0"
[tool.poetry.group.test.dependencies]
pytest = "^7.4.2"
pytest = ">=7.4.2,<10.0.0"
ruff = ">=0.3.1,<1.0"
[tool.poetry.group.dev.dependencies]
notebook = ">=7.0.2"
pre-commit = "^3.6.2"
onnx = ">=1.15.0"
pre-commit = ">=3.6.2,<5.0.0"
onnx = [
{ version = ">=1.15.0", python = ">=3.10,<3.13" },
{ version = ">=1.18.0", python = "3.13" },
{ version = ">=1.20.0", python = ">=3.14" },
]
[tool.poetry.group.docs.dependencies]
mkdocs-material = "^9.5.10"
mkdocstrings = "^0.24.0"
pillow = ">=10.3.0,<12.0.0"
mkdocstrings = ">=0.24,<1.1"
pillow = ">=10.3.0,<13.0.0"
cairosvg = "^2.7.1"
mknotebooks = "^0.8.0"
[tool.poetry.group.types.dependencies]
pyright = ">=1.1.293"
mypy = "^1.0.0"
mypy = ">=1,<3"
[build-system]
requires = ["poetry-core"]
+116 -103
View File
@@ -1,4 +1,5 @@
import os
from contextlib import contextmanager
import numpy as np
import pytest
@@ -7,98 +8,119 @@ from fastembed import SparseTextEmbedding
from tests.utils import delete_model_cache
@pytest.mark.parametrize("model_name", ["Qdrant/bm42-all-minilm-l6-v2-attentions", "Qdrant/bm25"])
def test_attention_embeddings(model_name: str) -> None:
_MODELS_TO_CACHE = ("Qdrant/bm42-all-minilm-l6-v2-attentions", "Qdrant/bm25")
MODELS_TO_CACHE = tuple([x.lower() for x in _MODELS_TO_CACHE])
@pytest.fixture(scope="module")
def model_cache():
is_ci = os.getenv("CI")
model = SparseTextEmbedding(model_name=model_name)
cache = {}
output = list(
model.query_embed(
[
"I must not fear. Fear is the mind-killer.",
]
)
)
@contextmanager
def get_model(model_name: str):
lowercase_model_name = model_name.lower()
if lowercase_model_name not in cache:
cache[lowercase_model_name] = SparseTextEmbedding(lowercase_model_name)
yield cache[lowercase_model_name]
if lowercase_model_name not in MODELS_TO_CACHE:
print("deleting model")
model_inst = cache.pop(lowercase_model_name)
if is_ci:
delete_model_cache(model_inst.model._model_dir)
del model_inst
assert len(output) == 1
for result in output:
assert len(result.indices) == len(result.values)
assert np.allclose(result.values, np.ones(len(result.values)))
quotes = [
"I must not fear. Fear is the mind-killer.",
"All animals are equal, but some animals are more equal than others.",
"It was a pleasure to burn.",
"The sky above the port was the color of television, tuned to a dead channel.",
"In the beginning, the universe was created."
" This has made a lot of people very angry and been widely regarded as a bad move.",
"It's a truth universally acknowledged that a zombie in possession of brains must be in want of more brains.",
"War is peace. Freedom is slavery. Ignorance is strength.",
"We're not in Infinity; we're in the suburbs.",
"I was a thousand times more evil than thou!",
"History is merely a list of surprises... It can only prepare us to be surprised yet again.",
".", # Empty string
]
output = list(model.embed(quotes))
assert len(output) == len(quotes)
for result in output[:-1]:
assert len(result.indices) == len(result.values)
assert len(result.indices) > 0
assert len(output[-1].indices) == 0
# Test support for unknown languages
output = list(
model.query_embed(
[
"привет мир!",
]
)
)
assert len(output) == 1
for result in output:
assert len(result.indices) == len(result.values)
assert len(result.indices) == 2
yield get_model
if is_ci:
delete_model_cache(model.model._model_dir)
for name, model in cache.items():
delete_model_cache(model.model._model_dir)
cache.clear()
@pytest.mark.parametrize("model_name", ["Qdrant/bm42-all-minilm-l6-v2-attentions", "Qdrant/bm25"])
def test_parallel_processing(model_name: str) -> None:
is_ci = os.getenv("CI")
def test_attention_embeddings(model_cache, model_name: str) -> None:
with model_cache(model_name) as model:
output = list(
model.query_embed(
[
"I must not fear. Fear is the mind-killer.",
]
)
)
model = SparseTextEmbedding(model_name=model_name)
assert len(output) == 1
docs = ["hello world", "attention embedding", "Mangez-vous vraiment des grenouilles?"] * 100
embeddings = list(model.embed(docs, batch_size=10, parallel=2))
for result in output:
assert len(result.indices) == len(result.values)
assert np.allclose(result.values, np.ones(len(result.values)))
embeddings_2 = list(model.embed(docs, batch_size=10, parallel=None))
quotes = [
"I must not fear. Fear is the mind-killer.",
"All animals are equal, but some animals are more equal than others.",
"It was a pleasure to burn.",
"The sky above the port was the color of television, tuned to a dead channel.",
"In the beginning, the universe was created."
" This has made a lot of people very angry and been widely regarded as a bad move.",
"It's a truth universally acknowledged that a zombie in possession of brains must be in want of more brains.",
"War is peace. Freedom is slavery. Ignorance is strength.",
"We're not in Infinity; we're in the suburbs.",
"I was a thousand times more evil than thou!",
"History is merely a list of surprises... It can only prepare us to be surprised yet again.",
".", # Empty string
]
embeddings_3 = list(model.embed(docs, batch_size=10, parallel=0))
output = list(model.embed(quotes))
assert len(embeddings) == len(docs)
assert len(output) == len(quotes)
for emb_1, emb_2, emb_3 in zip(embeddings, embeddings_2, embeddings_3):
assert np.allclose(emb_1.indices, emb_2.indices)
assert np.allclose(emb_1.indices, emb_3.indices)
assert np.allclose(emb_1.values, emb_2.values)
assert np.allclose(emb_1.values, emb_3.values)
for result in output[:-1]:
assert len(result.indices) == len(result.values)
assert len(result.indices) > 0
if is_ci:
delete_model_cache(model.model._model_dir)
assert len(output[-1].indices) == 0
# Test support for unknown languages
output = list(
model.query_embed(
[
"привет мир!",
]
)
)
assert len(output) == 1
for result in output:
assert len(result.indices) == len(result.values)
assert len(result.indices) == 2
@pytest.mark.parametrize("model_name", ["Qdrant/bm42-all-minilm-l6-v2-attentions", "Qdrant/bm25"])
def test_parallel_processing(model_cache, model_name: str) -> None:
with model_cache(model_name) as model:
docs = [
"hello world",
"attention embedding",
"Mangez-vous vraiment des grenouilles?",
] * 100
embeddings = list(model.embed(docs, batch_size=10, parallel=2))
embeddings_2 = list(model.embed(docs, batch_size=10, parallel=None))
embeddings_3 = list(model.embed(docs, batch_size=10, parallel=0))
assert len(embeddings) == len(docs)
for emb_1, emb_2, emb_3 in zip(embeddings, embeddings_2, embeddings_3):
assert np.allclose(emb_1.indices, emb_2.indices)
assert np.allclose(emb_1.indices, emb_3.indices)
assert np.allclose(emb_1.values, emb_2.values)
assert np.allclose(emb_1.values, emb_3.values)
@pytest.mark.parametrize("model_name", ["Qdrant/bm25"])
def test_multilanguage(model_name: str) -> None:
is_ci = os.getenv("CI")
def test_multilanguage(model_cache, model_name: str) -> None:
docs = ["Mangez-vous vraiment des grenouilles?", "Je suis au lit"]
model = SparseTextEmbedding(model_name=model_name, language="french")
@@ -109,39 +131,30 @@ def test_multilanguage(model_name: str) -> None:
assert embeddings[1].values.shape == (1,)
assert embeddings[1].indices.shape == (1,)
model = SparseTextEmbedding(model_name=model_name, language="english")
embeddings = list(model.embed(docs))[:2]
assert embeddings[0].values.shape == (5,)
assert embeddings[0].indices.shape == (5,)
with model_cache(model_name) as model: # language = "english"
embeddings = list(model.embed(docs))[:2]
assert embeddings[0].values.shape == (5,)
assert embeddings[0].indices.shape == (5,)
assert embeddings[1].values.shape == (4,)
assert embeddings[1].indices.shape == (4,)
if is_ci:
delete_model_cache(model.model._model_dir)
assert embeddings[1].values.shape == (4,)
assert embeddings[1].indices.shape == (4,)
@pytest.mark.parametrize("model_name", ["Qdrant/bm25"])
def test_special_characters(model_name: str) -> None:
is_ci = os.getenv("CI")
docs = [
"Über den größten Flüssen Österreichs äußern sich Experten häufig: Öko-Systeme müssen geschützt werden!",
"L'élève français s'écrie : « Où est mon crayon ? J'ai besoin de finir cet exercice avant la récréation!",
"Într-o zi însorită, Ștefan și Ioana au mâncat mămăligă cu brânză și au băut țuică la cabană.",
"Üzgün öğretmen öğrencilere seslendi: Lütfen gürültü yapmayın, sınavınızı bitirmeye çalışıyorum!",
"Ο Ξενοφών είπε: «Ψάχνω για ένα ωραίο δώρο για τη γιαγιά μου. Ίσως ένα φυτό ή ένα βιβλίο;»",
"Hola! ¿Cómo estás? Estoy muy emocionado por el cumpleaños de mi hermano, ¡va a ser increíble! También quiero comprar un pastel de chocolate con fresas y un regalo especial: un libro titulado «Cien años de soledad",
]
model = SparseTextEmbedding(model_name=model_name, language="english")
embeddings = list(model.embed(docs))
for idx, shape in enumerate([14, 18, 15, 10, 15]):
assert embeddings[idx].values.shape == (shape,)
assert embeddings[idx].indices.shape == (shape,)
if is_ci:
delete_model_cache(model.model._model_dir)
def test_special_characters(model_cache, model_name: str) -> None:
with model_cache(model_name) as model:
docs = [
"Über den größten Flüssen Österreichs äußern sich Experten häufig: Öko-Systeme müssen geschützt werden!",
"L'élève français s'écrie : « Où est mon crayon ? J'ai besoin de finir cet exercice avant la récréation!",
"Într-o zi însorită, Ștefan și Ioana au mâncat mămăligă cu brânză și au băut țuică la cabană.",
"Üzgün öğretmen öğrencilere seslendi: Lütfen gürültü yapmayın, sınavınızı bitirmeye çalışıyorum!",
"Ο Ξενοφών είπε: «Ψάχνω για ένα ωραίο δώρο για τη γιαγιά μου. Ίσως ένα φυτό ή ένα βιβλίο;»",
"Hola! ¿Cómo estás? Estoy muy emocionado por el cumpleaños de mi hermano, ¡va a ser increíble! También quiero comprar un pastel de chocolate con fresas y un regalo especial: un libro titulado «Cien años de soledad",
]
embeddings = list(model.embed(docs))
for idx, shape in enumerate([14, 18, 15, 10, 15]):
assert embeddings[idx].values.shape == (shape,)
assert embeddings[idx].indices.shape == (shape,)
@pytest.mark.parametrize("model_name", ["Qdrant/bm42-all-minilm-l6-v2-attentions"])
+31
View File
@@ -1,3 +1,5 @@
import numpy as np
from fastembed import (
TextEmbedding,
SparseTextEmbedding,
@@ -5,6 +7,7 @@ from fastembed import (
LateInteractionMultimodalEmbedding,
LateInteractionTextEmbedding,
)
from fastembed.common.utils import last_token_pooling
def test_text_list_supported_models():
@@ -28,3 +31,31 @@ def test_text_list_supported_models():
assert "model_file" in description and description["model_file"]
assert "sources" in description and description["sources"]
assert "hf" in description["sources"] or "url" in description["sources"]
def test_last_token_pooling():
token_embeddings = np.array(
[
[[1.0, 1.0], [2.0, 2.0], [9.0, 9.0], [9.0, 9.0]], # 2 real tokens, then padding
[[3.0, 3.0], [4.0, 4.0], [5.0, 5.0], [6.0, 6.0]], # no padding
]
)
attention_mask = np.array([[1, 1, 0, 0], [1, 1, 1, 1]], dtype=np.int64)
pooled = last_token_pooling(token_embeddings, attention_mask)
assert np.allclose(pooled, [[2.0, 2.0], [6.0, 6.0]])
def test_last_token_pooling_with_left_padding():
token_embeddings = np.array(
[
[[9.0, 9.0], [9.0, 9.0], [1.0, 1.0], [2.0, 2.0]], # padding, then 2 real tokens
[[3.0, 3.0], [4.0, 4.0], [5.0, 5.0], [6.0, 6.0]], # no padding
]
)
attention_mask = np.array([[0, 0, 1, 1], [1, 1, 1, 1]], dtype=np.int64)
pooled = last_token_pooling(token_embeddings, attention_mask)
assert np.allclose(pooled, [[2.0, 2.0], [6.0, 6.0]])
+152 -9
View File
@@ -3,10 +3,17 @@ import os
import numpy as np
import pytest
from fastembed.common.model_description import PoolingType, ModelSource, DenseModelDescription
from fastembed.common.model_description import (
PoolingType,
ModelSource,
DenseModelDescription,
BaseModelDescription,
)
from fastembed.common.onnx_model import OnnxOutputContext
from fastembed.common.utils import normalize, mean_pooling
from fastembed.common.utils import normalize, mean_pooling, last_token_pooling
from fastembed.text.custom_text_embedding import CustomTextEmbedding, PostprocessingConfig
from fastembed.rerank.cross_encoder.custom_text_cross_encoder import CustomTextCrossEncoder
from fastembed.rerank.cross_encoder import TextCrossEncoder
from fastembed.text.text_embedding import TextEmbedding
from tests.utils import delete_model_cache
@@ -14,8 +21,12 @@ from tests.utils import delete_model_cache
@pytest.fixture(autouse=True)
def restore_custom_models_fixture():
CustomTextEmbedding.SUPPORTED_MODELS = []
CustomTextEmbedding.POSTPROCESSING_MAPPING = {}
CustomTextCrossEncoder.SUPPORTED_MODELS = []
yield
CustomTextEmbedding.SUPPORTED_MODELS = []
CustomTextEmbedding.POSTPROCESSING_MAPPING = {}
CustomTextCrossEncoder.SUPPORTED_MODELS = []
def test_text_custom_model():
@@ -61,6 +72,89 @@ def test_text_custom_model():
assert embeddings.shape == (2, dim)
assert np.allclose(embeddings[0, : canonical_vector.shape[0]], canonical_vector, atol=1e-3)
if is_ci:
delete_model_cache(model.model._model_dir)
def test_text_custom_model_parallel_processing():
is_ci = os.getenv("CI")
custom_model_name = "intfloat/multilingual-e5-small"
dim = 384
TextEmbedding.add_custom_model(
custom_model_name,
pooling=PoolingType.MEAN,
normalization=True,
sources=ModelSource(hf=custom_model_name),
dim=dim,
size_in_gb=0.47,
)
model = TextEmbedding(custom_model_name)
docs = ["hello world", "flag embedding"] * 50
embeddings = np.stack(list(model.embed(docs, batch_size=10, parallel=2)), axis=0)
assert embeddings.shape == (len(docs), dim)
if is_ci:
delete_model_cache(model.model._model_dir)
def test_cross_encoder_custom_model():
is_ci = os.getenv("CI")
custom_model_name = "Xenova/ms-marco-MiniLM-L-4-v2"
size_in_gb = 0.08
source = ModelSource(hf=custom_model_name)
canonical_vector = np.array([-5.7170815, -11.112114], dtype=np.float32)
TextCrossEncoder.add_custom_model(
custom_model_name,
model_file="onnx/model.onnx",
sources=source,
size_in_gb=size_in_gb,
)
assert CustomTextCrossEncoder.SUPPORTED_MODELS[0] == BaseModelDescription(
model=custom_model_name,
sources=source,
model_file="onnx/model.onnx",
description="",
license="",
size_in_GB=size_in_gb,
)
model = TextCrossEncoder(custom_model_name)
pairs = [
("What is AI?", "Artificial intelligence is ..."),
("What is ML?", "Machine learning is ..."),
]
scores = list(model.rerank_pairs(pairs))
embeddings = np.stack(scores, axis=0)
assert embeddings.shape == (2,)
assert np.allclose(embeddings, canonical_vector, atol=1e-3)
if is_ci:
delete_model_cache(model.model._model_dir)
def test_cross_encoder_custom_model_parallel_processing():
is_ci = os.getenv("CI")
custom_model_name = "Xenova/ms-marco-MiniLM-L-4-v2"
TextCrossEncoder.add_custom_model(
custom_model_name,
model_file="onnx/model.onnx",
sources=ModelSource(hf=custom_model_name),
size_in_gb=0.08,
)
model = TextCrossEncoder(custom_model_name)
pairs = [("What is AI?", "Artificial intelligence is ...")] * 50
scores = np.stack(list(model.rerank_pairs(pairs, batch_size=10, parallel=2)), axis=0)
assert scores.shape == (len(pairs),)
if is_ci:
delete_model_cache(model.model._model_dir)
@@ -84,6 +178,8 @@ def test_mock_add_custom_models():
f"{PoolingType.MEAN.lower()}": dummy_token_output,
f"{PoolingType.CLS.lower()}-normalized": dummy_token_output,
f"{PoolingType.CLS.lower()}": dummy_token_output,
f"{PoolingType.LAST_TOKEN.lower()}-normalized": dummy_token_output,
f"{PoolingType.LAST_TOKEN.lower()}": dummy_token_output,
f"{PoolingType.DISABLED.lower()}-normalized": dummy_pooled_output,
f"{PoolingType.DISABLED.lower()}": dummy_pooled_output,
}
@@ -91,20 +187,23 @@ def test_mock_add_custom_models():
expected_output = {
f"{PoolingType.MEAN.lower()}-normalized": normalize(
mean_pooling(dummy_token_embedding, dummy_attention_mask)
).astype(np.float32),
),
f"{PoolingType.MEAN.lower()}": mean_pooling(dummy_token_embedding, dummy_attention_mask),
f"{PoolingType.CLS.lower()}-normalized": normalize(dummy_token_embedding[:, 0]).astype(
np.float32
),
f"{PoolingType.CLS.lower()}-normalized": normalize(dummy_token_embedding[:, 0]),
f"{PoolingType.CLS.lower()}": dummy_token_embedding[:, 0],
f"{PoolingType.DISABLED.lower()}-normalized": normalize(dummy_pooled_embedding).astype(
np.float32
f"{PoolingType.LAST_TOKEN.lower()}-normalized": normalize(
last_token_pooling(dummy_token_embedding, dummy_attention_mask)
),
f"{PoolingType.LAST_TOKEN.lower()}": last_token_pooling(
dummy_token_embedding, dummy_attention_mask
),
f"{PoolingType.DISABLED.lower()}-normalized": normalize(dummy_pooled_embedding),
f"{PoolingType.DISABLED.lower()}": dummy_pooled_embedding,
}
for pooling, normalization in itertools.product(
(PoolingType.MEAN, PoolingType.CLS, PoolingType.DISABLED), (True, False)
(PoolingType.MEAN, PoolingType.CLS, PoolingType.LAST_TOKEN, PoolingType.DISABLED),
(True, False),
):
model_name = f"{pooling.name.lower()}{'-normalized' if normalization else ''}"
TextEmbedding.add_custom_model(
@@ -128,6 +227,25 @@ def test_mock_add_custom_models():
assert np.allclose(post_processed_output, expected_output[model_name], atol=1e-3)
def test_custom_text_model_lookup_is_case_insensitive():
model_name = "Org/Model"
TextEmbedding.add_custom_model(
model_name,
pooling=PoolingType.MEAN,
normalization=True,
sources=ModelSource(hf="artificial"),
dim=5,
size_in_gb=0.1,
)
model = TextEmbedding("org/model", lazy_load=True, specific_model_path="./")
assert isinstance(model.model, CustomTextEmbedding)
assert model.model._pooling == PoolingType.MEAN
assert model.model._normalization is True
def test_do_not_add_existing_model():
existing_base_model = "sentence-transformers/all-MiniLM-L6-v2"
custom_model_name = "intfloat/multilingual-e5-small"
@@ -160,3 +278,28 @@ def test_do_not_add_existing_model():
dim=384,
size_in_gb=0.47,
)
def test_do_not_add_existing_cross_encoder():
existing_base_model = "Xenova/ms-marco-MiniLM-L-6-v2"
custom_model_name = "Xenova/ms-marco-MiniLM-L-4-v2"
with pytest.raises(ValueError, match=f"Model {existing_base_model} is already registered"):
TextCrossEncoder.add_custom_model(
existing_base_model,
sources=ModelSource(hf=existing_base_model),
size_in_gb=0.08,
)
TextCrossEncoder.add_custom_model(
custom_model_name,
sources=ModelSource(hf=existing_base_model),
size_in_gb=0.08,
)
with pytest.raises(ValueError, match=f"Model {custom_model_name} is already registered"):
TextCrossEncoder.add_custom_model(
custom_model_name,
sources=ModelSource(hf=custom_model_name),
size_in_gb=0.08,
)
+126 -60
View File
@@ -1,4 +1,6 @@
import os
import platform
from contextlib import contextmanager
from io import BytesIO
import numpy as np
@@ -8,7 +10,7 @@ from PIL import Image
from fastembed import ImageEmbedding
from tests.config import TEST_MISC_DIR
from tests.utils import delete_model_cache
from tests.utils import delete_model_cache, should_test_model
CANONICAL_VECTOR_VALUES = {
"Qdrant/clip-ViT-B-32-vision": np.array([-0.0098, 0.0128, -0.0274, 0.002, -0.0059]),
@@ -24,88 +26,124 @@ CANONICAL_VECTOR_VALUES = {
"jinaai/jina-clip-v1": np.array(
[-0.029, 0.0216, 0.0396, 0.0283, -0.0023, 0.0151, 0.011, -0.0235, 0.0251, -0.0343]
),
"nomic-ai/nomic-embed-vision-v1.5": np.array(
[0.0048, -0.0254, 0.0067, -0.0296, -0.0435, -0.0123, 0.0024, -0.0361, -0.0703, -0.0186]
),
"nomic-ai/nomic-embed-vision-v1.5-Q": np.array(
[-0.0011, -0.0477, 0.0024, -0.049, -0.0458, -0.0314, 0.017, -0.0383, -0.0537, -0.021]
),
"google/siglip2-base-patch16-224": np.array(
[-0.02095927, -0.0075177, -0.00144479, -0.0080948, 0.05031789]
),
}
_MODELS_TO_CACHE = ("Qdrant/clip-ViT-B-32-vision",)
MODELS_TO_CACHE = tuple([x.lower() for x in _MODELS_TO_CACHE])
def test_embedding() -> None:
@pytest.fixture(scope="module")
def model_cache():
is_ci = os.getenv("CI")
cache = {}
@contextmanager
def get_model(model_name: str):
lowercase_model_name = model_name.lower()
if lowercase_model_name not in cache:
cache[lowercase_model_name] = ImageEmbedding(lowercase_model_name)
yield cache[lowercase_model_name]
if lowercase_model_name not in MODELS_TO_CACHE:
model_inst = cache.pop(lowercase_model_name)
if is_ci:
delete_model_cache(model_inst.model._model_dir)
del model_inst
yield get_model
if is_ci:
for name, model in cache.items():
delete_model_cache(model.model._model_dir)
cache.clear()
@pytest.mark.parametrize("model_name", ["Qdrant/clip-ViT-B-32-vision"])
def test_embedding(model_cache, model_name: str) -> None:
is_ci = os.getenv("CI")
is_mac = platform.system() == "Darwin"
is_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
for model_desc in ImageEmbedding._list_supported_models():
if not is_ci and model_desc.size_in_GB > 1:
# quantized int8 ops diverge on macOS; canonical vector is generated on linux/amd64 (CI)
if is_mac and model_desc.model == "nomic-ai/nomic-embed-vision-v1.5-Q":
continue
if not should_test_model(model_desc, model_name, is_ci, is_manual):
continue
dim = model_desc.dim
model = ImageEmbedding(model_name=model_desc.model)
with model_cache(model_desc.model) as model:
images = [
TEST_MISC_DIR / "image.jpeg",
str(TEST_MISC_DIR / "small_image.jpeg"),
Image.open((TEST_MISC_DIR / "small_image.jpeg")),
Image.open(BytesIO(requests.get("https://qdrant.tech/img/logo.png").content)),
]
embeddings = list(model.embed(images))
embeddings = np.stack(embeddings, axis=0)
assert embeddings.shape == (len(images), dim)
images = [
canonical_vector = CANONICAL_VECTOR_VALUES[model_desc.model]
assert np.allclose(
embeddings[0, : canonical_vector.shape[0]], canonical_vector, atol=1e-3
), model_desc.model
assert np.allclose(embeddings[1], embeddings[2]), model_desc.model
@pytest.mark.parametrize("n_dims,model_name", [(512, "Qdrant/clip-ViT-B-32-vision")])
def test_batch_embedding(model_cache, n_dims: int, model_name: str) -> None:
with model_cache(model_name) as model:
n_images = 32
test_images = [
TEST_MISC_DIR / "image.jpeg",
str(TEST_MISC_DIR / "small_image.jpeg"),
Image.open((TEST_MISC_DIR / "small_image.jpeg")),
Image.open(BytesIO(requests.get("https://qdrant.tech/img/logo.png").content)),
Image.open(TEST_MISC_DIR / "small_image.jpeg"),
]
embeddings = list(model.embed(images))
images = test_images * n_images
embeddings = list(model.embed(images, batch_size=10))
embeddings = np.stack(embeddings, axis=0)
assert embeddings.shape == (len(images), dim)
assert np.allclose(embeddings[1], embeddings[2])
canonical_vector = CANONICAL_VECTOR_VALUES[model_desc.model]
canonical_vector = CANONICAL_VECTOR_VALUES[model_name]
assert np.allclose(
embeddings[0, : canonical_vector.shape[0]], canonical_vector, atol=1e-3
), model_desc.model
assert np.allclose(embeddings[1], embeddings[2]), model_desc.model
if is_ci:
delete_model_cache(model.model._model_dir)
assert embeddings.shape == (len(test_images) * n_images, n_dims)
assert np.allclose(embeddings[0, : canonical_vector.shape[0]], canonical_vector, atol=1e-3)
@pytest.mark.parametrize("n_dims,model_name", [(512, "Qdrant/clip-ViT-B-32-vision")])
def test_batch_embedding(n_dims: int, model_name: str) -> None:
is_ci = os.getenv("CI")
model = ImageEmbedding(model_name=model_name)
n_images = 32
test_images = [
TEST_MISC_DIR / "image.jpeg",
str(TEST_MISC_DIR / "small_image.jpeg"),
Image.open(TEST_MISC_DIR / "small_image.jpeg"),
]
images = test_images * n_images
def test_parallel_processing(model_cache, n_dims: int, model_name: str) -> None:
with model_cache(model_name) as model:
n_images = 32
test_images = [
TEST_MISC_DIR / "image.jpeg",
str(TEST_MISC_DIR / "small_image.jpeg"),
Image.open(TEST_MISC_DIR / "small_image.jpeg"),
]
images = test_images * n_images
embeddings = list(model.embed(images, batch_size=10, parallel=2))
embeddings = np.stack(embeddings, axis=0)
embeddings = list(model.embed(images, batch_size=10))
embeddings = np.stack(embeddings, axis=0)
embeddings_2 = list(model.embed(images, batch_size=10, parallel=None))
embeddings_2 = np.stack(embeddings_2, axis=0)
assert embeddings.shape == (len(test_images) * n_images, n_dims)
if is_ci:
delete_model_cache(model.model._model_dir)
embeddings_3 = list(model.embed(images, batch_size=10, parallel=0))
embeddings_3 = np.stack(embeddings_3, axis=0)
@pytest.mark.parametrize("n_dims,model_name", [(512, "Qdrant/clip-ViT-B-32-vision")])
def test_parallel_processing(n_dims: int, model_name: str) -> None:
is_ci = os.getenv("CI")
model = ImageEmbedding(model_name=model_name)
n_images = 32
test_images = [
TEST_MISC_DIR / "image.jpeg",
str(TEST_MISC_DIR / "small_image.jpeg"),
Image.open(TEST_MISC_DIR / "small_image.jpeg"),
]
images = test_images * n_images
embeddings = list(model.embed(images, batch_size=10, parallel=2))
embeddings = np.stack(embeddings, axis=0)
embeddings_2 = list(model.embed(images, batch_size=10, parallel=None))
embeddings_2 = np.stack(embeddings_2, axis=0)
embeddings_3 = list(model.embed(images, batch_size=10, parallel=0))
embeddings_3 = np.stack(embeddings_3, axis=0)
assert embeddings.shape == (n_images * len(test_images), n_dims)
assert np.allclose(embeddings, embeddings_2, atol=1e-3)
assert np.allclose(embeddings, embeddings_3, atol=1e-3)
if is_ci:
delete_model_cache(model.model._model_dir)
assert embeddings.shape == (n_images * len(test_images), n_dims)
assert np.allclose(embeddings, embeddings_2, atol=1e-3)
assert np.allclose(embeddings, embeddings_3, atol=1e-3)
@pytest.mark.parametrize("model_name", ["Qdrant/clip-ViT-B-32-vision"])
@@ -121,3 +159,31 @@ def test_lazy_load(model_name: str) -> None:
assert hasattr(model.model, "model")
if is_ci:
delete_model_cache(model.model._model_dir)
def test_get_embedding_size() -> None:
assert ImageEmbedding.get_embedding_size(model_name="Qdrant/clip-ViT-B-32-vision") == 512
assert ImageEmbedding.get_embedding_size(model_name="Qdrant/clip-vit-b-32-vision") == 512
def test_embedding_size() -> None:
is_ci = os.getenv("CI")
model_name = "Qdrant/clip-ViT-B-32-vision"
model = ImageEmbedding(model_name=model_name, lazy_load=True)
assert model.embedding_size == 512
model_name = "Qdrant/clip-vit-b-32-vision"
model = ImageEmbedding(model_name=model_name, lazy_load=True)
assert model.embedding_size == 512
if is_ci:
delete_model_cache(model.model._model_dir)
@pytest.mark.parametrize("model_name", ["Qdrant/clip-ViT-B-32-vision"])
def test_session_options(model_cache, model_name) -> None:
with model_cache(model_name) as default_model:
default_session_options = default_model.model.model.get_session_options()
assert default_session_options.enable_cpu_mem_arena is True
model = ImageEmbedding(model_name=model_name, enable_cpu_mem_arena=False)
session_options = model.model.model.get_session_options()
assert session_options.enable_cpu_mem_arena is False
+90
View File
@@ -0,0 +1,90 @@
import numpy as np
import pytest
from PIL import Image
from fastembed.image.transform.functional import normalize, resize
@pytest.mark.parametrize(
("size", "expected"),
[
((100, 200), (200, 100)), # the bug: a non-square size came back transposed
((224, 224), (224, 224)), # the square path every shipped model takes
],
)
def test_resize_tuple_is_height_width(size: tuple[int, int], expected: tuple[int, int]) -> None:
"""A ``(height, width)`` size must reach Pillow as ``(width, height)``."""
resized = resize(Image.new("RGB", (300, 300)), size=size)
assert resized.size == expected # PIL reports (width, height)
def test_resize_int_keeps_shortest_edge_behaviour() -> None:
"""The int branch already emitted Pillow order; it must not be disturbed."""
landscape = Image.new("RGB", (400, 200))
portrait = Image.new("RGB", (200, 400))
# size sets the shortest edge, and the aspect ratio is preserved.
assert resize(landscape, size=100).size == (200, 100)
assert resize(portrait, size=100).size == (100, 200)
@pytest.mark.parametrize(
("mean", "std"),
[
([0.1, 0.2, 0.3], [0.5, 0.6, 0.7]), # per-channel, as every model config gives it
(0.5, 0.25), # scalar, expanded to one value per channel
],
)
def test_normalize_chw_is_channel_wise(
mean: list[float] | float, std: list[float] | float
) -> None:
"""Each channel must be normalized by its own mean/std, not by any other axis."""
rng = np.random.default_rng(0)
image = rng.random((3, 5, 7)).astype(np.float32)
means = mean if isinstance(mean, list) else [mean] * 3
stds = std if isinstance(std, list) else [std] * 3
result = normalize(image, mean=mean, std=std)
for c in range(3):
assert np.allclose(result[c], (image[c] - means[c]) / stds[c], atol=1e-6)
@pytest.mark.parametrize("batch_size", [2, 3])
def test_normalize_batched_matches_per_image(batch_size: int) -> None:
"""A batch must give exactly what the (C, H, W) path gives image by image.
batch_size 2 used to raise, since transposing reversed every axis; batch_size 3
matched the channel count and silently normalized along the batch axis instead.
"""
rng = np.random.default_rng(2)
batch = rng.random((batch_size, 3, 4, 4)).astype(np.float32)
mean, std = [0.1, 0.2, 0.3], [0.5, 0.6, 0.7]
result = normalize(batch, mean=mean, std=std)
per_image = np.stack([normalize(image, mean=mean, std=std) for image in batch])
assert result.shape == batch.shape
assert np.array_equal(result, per_image)
def test_normalize_rejects_input_without_a_channel_axis() -> None:
"""Every pipeline runs ConvertToRGB first, so normalize only ever sees (C, H, W)."""
with pytest.raises(ValueError, match=r"must be \(C, H, W\)"):
normalize(np.zeros((4, 6), dtype=np.float32), mean=0.5, std=0.25)
@pytest.mark.parametrize(
("mean", "std", "expected"),
[
([0.1, 0.2], [1.0, 1.0, 1.0], "mean must"),
([0.1, 0.2, 0.3], [1.0, 1.0], "std must"),
],
)
def test_normalize_channel_count_mismatch_raises(
mean: list[float], std: list[float], expected: str
) -> None:
image = np.zeros((3, 4, 4), dtype=np.float32)
with pytest.raises(ValueError, match=expected):
normalize(image, mean=mean, std=std)
+144 -46
View File
@@ -1,4 +1,5 @@
import os
from contextlib import contextmanager
import pytest
import numpy as np
@@ -6,7 +7,7 @@ import numpy as np
from fastembed.late_interaction.late_interaction_text_embedding import (
LateInteractionTextEmbedding,
)
from tests.utils import delete_model_cache
from tests.utils import delete_model_cache, should_test_model
# vectors are abridged and rounded for brevity
CANONICAL_COLUMN_VALUES = {
@@ -150,82 +151,124 @@ CANONICAL_QUERY_VALUES = {
),
}
_MODELS_TO_CACHE = ("answerdotai/answerai-colbert-small-v1",)
MODELS_TO_CACHE = tuple([x.lower() for x in _MODELS_TO_CACHE])
@pytest.fixture(scope="module")
def model_cache():
is_ci = os.getenv("CI")
cache = {}
@contextmanager
def get_model(model_name: str):
lowercase_model_name = model_name.lower()
if lowercase_model_name not in cache:
cache[lowercase_model_name] = LateInteractionTextEmbedding(lowercase_model_name)
yield cache[lowercase_model_name]
if lowercase_model_name not in MODELS_TO_CACHE:
model_inst = cache.pop(lowercase_model_name)
if is_ci:
delete_model_cache(model_inst.model._model_dir)
del model_inst
yield get_model
if is_ci:
for name, model in cache.items():
delete_model_cache(model.model._model_dir)
cache.clear()
docs = ["Hello World"]
def test_batch_embedding():
is_ci = os.getenv("CI")
@pytest.mark.parametrize("model_name", ["answerdotai/answerai-colbert-small-v1"])
def test_batch_embedding(model_cache, model_name: str):
docs_to_embed = docs * 10
for model_name, expected_result in CANONICAL_COLUMN_VALUES.items():
print("evaluating", model_name)
model = LateInteractionTextEmbedding(model_name=model_name)
with model_cache(model_name) as model:
result = list(model.embed(docs_to_embed, batch_size=6))
expected_result = CANONICAL_COLUMN_VALUES[model_name]
for value in result:
token_num, abridged_dim = expected_result.shape
assert np.allclose(value[:, :abridged_dim], expected_result, atol=2e-3)
if is_ci:
delete_model_cache(model.model._model_dir)
@pytest.mark.parametrize("model_name", ["answerdotai/answerai-colbert-small-v1"])
def test_batch_inference_size_same_as_single_inference(model_cache, model_name: str):
with model_cache(model_name) as model:
docs_to_embed = [
"short document",
"A bit longer document, which should not affect the size",
]
result = list(model.embed(docs_to_embed, batch_size=1))
result_2 = list(model.embed(docs_to_embed, batch_size=2))
assert len(result[0]) == len(result_2[0])
def test_single_embedding():
@pytest.mark.parametrize("model_name", ["answerdotai/answerai-colbert-small-v1"])
def test_single_embedding(model_cache, model_name: str):
is_ci = os.getenv("CI")
is_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
docs_to_embed = docs
for model_name, expected_result in CANONICAL_COLUMN_VALUES.items():
for model_desc in LateInteractionTextEmbedding._list_supported_models():
if not should_test_model(model_desc, model_name, is_ci, is_manual):
continue
print("evaluating", model_name)
model = LateInteractionTextEmbedding(model_name=model_name)
result = next(iter(model.embed(docs_to_embed, batch_size=6)))
token_num, abridged_dim = expected_result.shape
assert np.allclose(result[:, :abridged_dim], expected_result, atol=2e-3)
if is_ci:
delete_model_cache(model.model._model_dir)
with model_cache(model_desc.model) as model:
whole_result = list(model.embed(docs_to_embed, batch_size=6))
assert len(whole_result) == 1
result = whole_result[0]
expected_result = CANONICAL_COLUMN_VALUES[model_desc.model]
token_num, abridged_dim = expected_result.shape
assert np.allclose(result[:, :abridged_dim], expected_result, atol=2e-3)
def test_single_embedding_query():
@pytest.mark.parametrize("model_name", ["answerdotai/answerai-colbert-small-v1"])
def test_single_embedding_query(model_cache, model_name: str):
is_ci = os.getenv("CI")
is_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
queries_to_embed = docs
for model_name, expected_result in CANONICAL_QUERY_VALUES.items():
print("evaluating", model_name)
model = LateInteractionTextEmbedding(model_name=model_name)
result = next(iter(model.query_embed(queries_to_embed)))
token_num, abridged_dim = expected_result.shape
assert np.allclose(result[:, :abridged_dim], expected_result, atol=2e-3)
for model_desc in LateInteractionTextEmbedding._list_supported_models():
if not should_test_model(model_desc, model_name, is_ci, is_manual):
continue
if is_ci:
delete_model_cache(model.model._model_dir)
print("evaluating", model_desc.model)
with model_cache(model_desc.model) as model:
whole_result = list(model.query_embed(queries_to_embed))
assert len(whole_result) == 1
result = whole_result[0]
expected_result = CANONICAL_QUERY_VALUES[model_desc.model]
token_num, abridged_dim = expected_result.shape
assert np.allclose(result[:, :abridged_dim], expected_result, atol=2e-3)
def test_parallel_processing():
is_ci = os.getenv("CI")
model = LateInteractionTextEmbedding(model_name="colbert-ir/colbertv2.0")
token_dim = 128
docs = ["hello world", "flag embedding"] * 100
embeddings = list(model.embed(docs, batch_size=10, parallel=2))
embeddings = np.stack(embeddings, axis=0)
@pytest.mark.parametrize("token_dim,model_name", [(96, "answerdotai/answerai-colbert-small-v1")])
def test_parallel_processing(model_cache, token_dim: int, model_name: str):
with model_cache(model_name) as model:
docs = ["hello world", "flag embedding"] * 100
embeddings = list(model.embed(docs, batch_size=10, parallel=2))
embeddings_2 = list(model.embed(docs, batch_size=10, parallel=None))
embeddings_2 = np.stack(embeddings_2, axis=0)
embeddings_2 = list(model.embed(docs, batch_size=10, parallel=None))
embeddings_3 = list(model.embed(docs, batch_size=10, parallel=0))
embeddings_3 = np.stack(embeddings_3, axis=0)
# embeddings_3 = list(model.embed(docs, batch_size=10, parallel=0)) # inherits OnnxTextModel which
# # is tested in TextEmbedding, disabling it here to reduce number of requests to hf
# # multiprocessing is enough to test with `parallel=2`, and `parallel=None` is okay to tests since it reuses
# # model from cache
assert embeddings.shape[0] == len(docs) and embeddings.shape[-1] == token_dim
assert np.allclose(embeddings, embeddings_2, atol=1e-3)
assert np.allclose(embeddings, embeddings_3, atol=1e-3)
assert len(embeddings) == len(docs) and embeddings[0].shape[-1] == token_dim
if is_ci:
delete_model_cache(model.model._model_dir)
for i in range(len(embeddings)):
assert np.allclose(embeddings[i], embeddings_2[i], atol=1e-3)
# assert np.allclose(embeddings[i], embeddings_3[i], atol=1e-3)
@pytest.mark.parametrize(
"model_name",
["colbert-ir/colbertv2.0"],
)
@pytest.mark.parametrize("model_name", ["answerdotai/answerai-colbert-small-v1"])
def test_lazy_load(model_name: str):
is_ci = os.getenv("CI")
@@ -244,3 +287,58 @@ def test_lazy_load(model_name: str):
if is_ci:
delete_model_cache(model.model._model_dir)
def test_get_embedding_size():
model_name = "answerdotai/answerai-colbert-small-v1"
assert LateInteractionTextEmbedding.get_embedding_size(model_name) == 96
model_name = "answerdotai/answerai-ColBERT-small-v1"
assert LateInteractionTextEmbedding.get_embedding_size(model_name) == 96
def test_embedding_size():
is_ci = os.getenv("CI")
model_name = "answerdotai/answerai-colbert-small-v1"
model = LateInteractionTextEmbedding(model_name=model_name, lazy_load=True)
assert model.embedding_size == 96
model_name = "answerdotai/answerai-ColBERT-small-v1"
model = LateInteractionTextEmbedding(model_name=model_name, lazy_load=True)
assert model.embedding_size == 96
if is_ci:
delete_model_cache(model.model._model_dir)
@pytest.mark.parametrize("model_name", ["answerdotai/answerai-ColBERT-small-v1"])
def test_session_options(model_cache, model_name) -> None:
with model_cache(model_name) as default_model:
default_session_options = default_model.model.model.get_session_options()
assert default_session_options.enable_cpu_mem_arena is True
model = LateInteractionTextEmbedding(model_name=model_name, enable_cpu_mem_arena=False)
session_options = model.model.model.get_session_options()
assert session_options.enable_cpu_mem_arena is False
@pytest.mark.parametrize("model_name", ["answerdotai/answerai-colbert-small-v1"])
def test_token_count(model_cache, model_name) -> None:
with model_cache(model_name) as model:
documents = ["short doc", "it is a long document to check attention mask for paddings"]
short_doc_token_count = model.token_count(documents[0])
long_doc_token_count = model.token_count(documents[1])
documents_token_count = model.token_count(documents)
assert short_doc_token_count + long_doc_token_count == documents_token_count
# 2 is 2*DOC_MARKER_TOKEN_ID for each document
assert short_doc_token_count + long_doc_token_count + 2 == model.token_count(
documents, include_extension=True
)
assert short_doc_token_count + long_doc_token_count == model.token_count(
documents, batch_size=1
)
assert short_doc_token_count + long_doc_token_count == model.token_count(
documents, is_doc=False
)
# query min length is 32
assert model.token_count(documents, is_doc=False, include_extension=True) == 64
very_long_query = "It's a very long query which definitely contains more than 32 tokens and we're using it to check whether the method can handle large query properly without cutting it to 32 tokens"
assert model.token_count(very_long_query, is_doc=False, include_extension=True) > 32
+104 -21
View File
@@ -1,11 +1,13 @@
import os
from contextlib import contextmanager
import pytest
from PIL import Image
import numpy as np
from fastembed import LateInteractionMultimodalEmbedding
from tests.config import TEST_MISC_DIR
from tests.utils import delete_model_cache
# vectors are abridged and rounded for brevity
CANONICAL_IMAGE_VALUES = {
@@ -20,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 = {
@@ -34,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"]
@@ -43,14 +67,42 @@ images = [
Image.open((TEST_MISC_DIR / "image.jpeg")),
]
_MODELS_TO_CACHE = ("Qdrant/colmodernvbert",)
MODELS_TO_CACHE = tuple(model_name.lower() for model_name in _MODELS_TO_CACHE)
def test_batch_embedding():
@pytest.fixture(scope="module")
def model_cache():
is_ci = os.getenv("CI")
cache = {}
if not is_ci:
for model_name, expected_result in CANONICAL_IMAGE_VALUES.items():
print("evaluating", model_name)
model = LateInteractionMultimodalEmbedding(model_name=model_name)
@contextmanager
def get_model(model_name: str):
lowercase_model_name = model_name.lower()
if lowercase_model_name not in cache:
cache[lowercase_model_name] = LateInteractionMultimodalEmbedding(lowercase_model_name)
yield cache[lowercase_model_name]
if lowercase_model_name not in MODELS_TO_CACHE:
model_inst = cache.pop(lowercase_model_name)
if is_ci:
delete_model_cache(model_inst.model._model_dir)
del model_inst
yield get_model
if is_ci:
for _, model in cache.items():
delete_model_cache(model.model._model_dir)
cache.clear()
def test_batch_embedding(model_cache):
for model_name, expected_result in CANONICAL_IMAGE_VALUES.items():
if model_name.lower() == "Qdrant/colpali-v1.3-fp16".lower() and os.getenv("CI"):
continue # colpali is too large for ci
print("evaluating", model_name)
with model_cache(model_name) as model:
result = list(model.embed_image(images, batch_size=2))
for value in result:
@@ -58,25 +110,56 @@ def test_batch_embedding():
assert np.allclose(value[:token_num, :abridged_dim], expected_result, atol=2e-3)
def test_single_embedding():
is_ci = os.getenv("CI")
if not is_ci:
for model_name, expected_result in CANONICAL_IMAGE_VALUES.items():
print("evaluating", model_name)
model = LateInteractionMultimodalEmbedding(model_name=model_name)
def test_single_embedding(model_cache):
for model_name, expected_result in CANONICAL_IMAGE_VALUES.items():
if model_name.lower() == "Qdrant/colpali-v1.3-fp16".lower() and os.getenv("CI"):
continue # colpali is too large for ci
print("evaluating", model_name)
with model_cache(model_name) as model:
result = next(iter(model.embed_image(images, batch_size=6)))
token_num, abridged_dim = expected_result.shape
assert np.allclose(result[:token_num, :abridged_dim], expected_result, atol=2e-3)
def test_single_embedding_query():
is_ci = os.getenv("CI")
if not is_ci:
queries_to_embed = queries
for model_name, expected_result in CANONICAL_QUERY_VALUES.items():
print("evaluating", model_name)
model = LateInteractionMultimodalEmbedding(model_name=model_name)
result = next(iter(model.embed_text(queries_to_embed)))
def test_single_embedding_query(model_cache):
for model_name, expected_result in CANONICAL_QUERY_VALUES.items():
if model_name.lower() == "Qdrant/colpali-v1.3-fp16".lower() and os.getenv("CI"):
continue # colpali is too large for ci
print("evaluating", model_name)
with model_cache(model_name) as model:
result = next(iter(model.embed_text(queries)))
token_num, abridged_dim = expected_result.shape
assert np.allclose(result[:token_num, :abridged_dim], expected_result, atol=2e-3)
def test_get_embedding_size():
model_name = "Qdrant/colpali-v1.3-fp16"
assert LateInteractionMultimodalEmbedding.get_embedding_size(model_name) == 128
model_name = "Qdrant/ColPali-v1.3-fp16"
assert LateInteractionMultimodalEmbedding.get_embedding_size(model_name) == 128
model_name = "Qdrant/colmodernvbert"
assert LateInteractionMultimodalEmbedding.get_embedding_size(model_name) == 128
def test_embedding_size():
model_name = "Qdrant/colmodernvbert"
model = LateInteractionMultimodalEmbedding(model_name=model_name, lazy_load=True)
assert model.embedding_size == 128
def test_token_count(model_cache) -> None:
model_name = "Qdrant/colmodernvbert"
with model_cache(model_name) as model:
documents = ["short doc", "it is a long document to check attention mask for paddings"]
short_doc_token_count = model.token_count(documents[0])
long_doc_token_count = model.token_count(documents[1])
documents_token_count = model.token_count(documents)
assert short_doc_token_count + long_doc_token_count == documents_token_count
assert short_doc_token_count + long_doc_token_count == model.token_count(
documents, batch_size=1
)
assert short_doc_token_count + long_doc_token_count < model.token_count(
documents, include_extension=True
)
+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
+38
View File
@@ -0,0 +1,38 @@
import numpy as np
from fastembed import LateInteractionTextEmbedding
from fastembed.postprocess import Muvera
CANONICAL_VALUES = [-2.61810007e-04, 1.89005750e00, -2.32070747e00]
CANONICAL_QUERY_VALUES = [
-0.85783903,
1.1077204,
-0.09522747,
] # part of the values are zeros, should be compared with the result of nonzero mask
DIM = 128
K_SIM = 5
DIM_PROJ = 16
R_REPS = 20
def test_single_input():
model = LateInteractionTextEmbedding("colbert-ir/colbertv2.0", lazy_load=True)
random_generator = np.random.default_rng(42)
multivector = random_generator.random((10, 128))
for muvera in (
Muvera(dim=DIM, k_sim=K_SIM, dim_proj=DIM_PROJ, r_reps=R_REPS, random_seed=42),
Muvera.from_multivector_model(model, k_sim=K_SIM, dim_proj=DIM_PROJ, r_reps=R_REPS),
):
fde = muvera.process(multivector)
assert fde.shape[0] == muvera.embedding_size
assert np.allclose(fde[:3], CANONICAL_VALUES)
fde_doc = muvera.process_document(multivector)
assert fde_doc.shape[0] == muvera.embedding_size
assert np.allclose(fde, fde_doc)
fde_query = muvera.process_query(multivector)
assert fde_query.shape[0] == muvera.embedding_size
assert np.allclose(fde_query[np.nonzero(fde_query)][:3], CANONICAL_QUERY_VALUES)
+333
View File
@@ -0,0 +1,333 @@
import itertools
import json
import os
import shutil
from pathlib import Path
from typing import Any
import numpy as np
import pytest
from tokenizers import Tokenizer
from fastembed.common.preprocessor_utils import load_tokenizer
from fastembed.text.text_embedding import TextEmbedding
from tests.utils import delete_model_cache
# transformers writes its VERY_LARGE_INTEGER in place of `model_max_length` when the real
# value is unknown, which is more than `enable_truncation` can accept
HF_SENTINEL = int(1e30)
# a lightweight model whose config files serve as a realistic starting point for the cases below
BASE_MODEL = "BAAI/bge-small-en-v1.5"
TOKENIZER_FILES = (
"config.json",
"tokenizer.json",
"tokenizer_config.json",
"special_tokens_map.json",
)
def _patch_json(path: Path, overrides: dict[str, Any], drop: tuple[str, ...] = ()) -> None:
with open(path) as source:
content = json.load(source)
content.update(overrides)
for key in drop:
content.pop(key, None)
with open(path, "w") as target:
json.dump(content, target)
def _set_serialized_padding(path: Path, padding: dict[str, Any] | None) -> None:
"""Rewrite tokenizer.json through the tokenizers library, so the format stays authoritative."""
tokenizer = Tokenizer.from_file(str(path))
if padding is None:
tokenizer.no_padding()
else:
tokenizer.enable_padding(**padding)
tokenizer.save(str(path))
@pytest.fixture(scope="module")
def make_model_dir(tmp_path_factory):
"""Build model directories from a real model's config files, with targeted overrides.
`load_tokenizer` reads only the four files in `TOKENIZER_FILES`, so the onnx weights are
never copied.
"""
is_ci = os.getenv("CI")
base_model = TextEmbedding(BASE_MODEL)
source_dir = Path(base_model.model._model_dir)
counter = itertools.count()
def factory(
tokenizer_config: dict[str, Any] | None = None,
config: dict[str, Any] | None = None,
padding: dict[str, Any] | None = None,
drop_from_tokenizer_config: tuple[str, ...] = (),
drop_from_config: tuple[str, ...] = (),
drop_files: tuple[str, ...] = (),
) -> Path:
model_dir = tmp_path_factory.mktemp(f"model_dir_{next(counter)}")
for file_name in TOKENIZER_FILES:
if file_name in drop_files:
continue
shutil.copy(source_dir / file_name, model_dir / file_name)
if "tokenizer_config.json" not in drop_files:
_patch_json(
model_dir / "tokenizer_config.json",
tokenizer_config or {},
drop_from_tokenizer_config,
)
if "config.json" not in drop_files:
_patch_json(model_dir / "config.json", config or {}, drop_from_config)
if padding is not None:
_set_serialized_padding(model_dir / "tokenizer.json", padding)
return model_dir
yield factory
if is_ci:
delete_model_cache(base_model.model._model_dir)
def test_fixed_padding_is_relaxed_to_batch_longest(make_model_dir) -> None:
"""Fixed padding shorter than the truncation limit leaves longer encodings ragged."""
model_dir = make_model_dir(
tokenizer_config={"model_max_length": 512, "max_length": None},
padding={"length": 128, "pad_id": 0, "pad_token": "[PAD]", "direction": "right"},
)
tokenizer, _ = load_tokenizer(model_dir)
assert tokenizer.padding["length"] is None
assert tokenizer.truncation["max_length"] == 512
encoded = tokenizer.encode_batch(["hello world", "retrieval " * 200])
# ragged encodings make this raise, the same way onnx_embed does
input_ids = np.array([encoding.ids for encoding in encoded])
assert input_ids.shape == (2, len(encoded[0].ids))
assert input_ids.shape[1] > 128
def test_batch_longest_padding_does_not_pad_to_the_truncation_limit(make_model_dir) -> None:
model_dir = make_model_dir(
tokenizer_config={"model_max_length": 512, "max_length": None},
padding={"length": 128, "pad_id": 0, "pad_token": "[PAD]", "direction": "right"},
)
tokenizer, _ = load_tokenizer(model_dir)
encoded = tokenizer.encode_batch(["hello world", "hello"])
assert len(encoded[0].ids) == len(encoded[1].ids) < 128
def test_serialized_left_padding_is_preserved(make_model_dir) -> None:
"""ColModernVBERT pads on the left, normalizing the length must not reset the direction."""
model_dir = make_model_dir(
padding={"length": None, "pad_id": 0, "pad_token": "[PAD]", "direction": "left"},
)
tokenizer, _ = load_tokenizer(model_dir)
assert tokenizer.padding["direction"] == "left"
assert tokenizer.padding["length"] is None
encoded = tokenizer.encode_batch(["hello world and then some", "hello"])
assert encoded[1].ids[0] == 0
assert encoded[1].attention_mask[0] == 0
def test_serialized_pad_to_multiple_of_is_preserved(make_model_dir) -> None:
"""Everything the tokenizer declared is kept; only the fixed length is overridden."""
model_dir = make_model_dir(
padding={"length": 128, "pad_id": 0, "pad_token": "[PAD]", "pad_to_multiple_of": 8},
)
tokenizer, _ = load_tokenizer(model_dir)
assert tokenizer.padding["length"] is None
assert tokenizer.padding["pad_to_multiple_of"] == 8
encoded = tokenizer.encode_batch(["hello world", "hello"])
assert len(encoded[0].ids) % 8 == 0
def test_serialized_pad_id_takes_precedence_over_config(make_model_dir) -> None:
mask_pad_id = 103 # [MASK] in the bert-base vocab, any id other than the config's works
model_dir = make_model_dir(
config={"pad_token_id": 7},
padding={"length": None, "pad_id": mask_pad_id, "pad_token": "[MASK]"},
)
tokenizer, _ = load_tokenizer(model_dir)
assert tokenizer.padding["pad_id"] == mask_pad_id
assert tokenizer.padding["pad_token"] == "[MASK]"
def test_pad_token_falls_back_to_tokenizer_config(make_model_dir) -> None:
model_dir = make_model_dir(config={"pad_token_id": 3}, padding=None)
tokenizer, _ = load_tokenizer(model_dir)
assert tokenizer.padding["pad_token"] == "[PAD]"
assert tokenizer.padding["pad_id"] == 3
def test_missing_pad_token_raises(make_model_dir) -> None:
model_dir = make_model_dir(drop_from_tokenizer_config=("pad_token",))
with pytest.raises(ValueError, match="Could not find a pad token"):
load_tokenizer(model_dir)
@pytest.mark.parametrize(
"model_max_length,max_length,expected",
[
(512, 128, 128), # both usable, the stricter one wins
(128, 512, 128),
(512, None, 512),
(None, 256, 256),
(HF_SENTINEL, 128, 128), # qdrant/gte-large-onnx
(0, 256, 256),
(512, 0, 512),
],
)
def test_max_context_resolution(make_model_dir, model_max_length, max_length, expected) -> None:
model_dir = make_model_dir(
tokenizer_config={"model_max_length": model_max_length, "max_length": max_length},
)
tokenizer, _ = load_tokenizer(model_dir)
assert tokenizer.truncation["max_length"] == expected
@pytest.mark.parametrize(
"model_max_length,max_length",
[
(HF_SENTINEL, None), # transformers' placeholder is not a limit
(0, None), # a zero would truncate everything away
(None, 0),
(None, None),
("512", None), # not an integer
],
)
def test_unusable_max_context_raises(make_model_dir, model_max_length, max_length) -> None:
model_dir = make_model_dir(
tokenizer_config={"model_max_length": model_max_length, "max_length": max_length},
)
with pytest.raises(ValueError, match="Could not determine the maximum context length"):
load_tokenizer(model_dir)
def test_absent_max_context_keys_raise(make_model_dir) -> None:
model_dir = make_model_dir(
drop_from_tokenizer_config=("model_max_length", "max_length"),
)
with pytest.raises(ValueError, match="Could not determine the maximum context length"):
load_tokenizer(model_dir)
@pytest.fixture(scope="module")
def token_id(make_model_dir):
"""Resolve vocabulary ids by name, so the cases below carry no magic numbers."""
return Tokenizer.from_file(str(make_model_dir() / "tokenizer.json")).token_to_id
@pytest.mark.parametrize(
"dropped",
[("config.json",), ("special_tokens_map.json",), ("config.json", "special_tokens_map.json")],
ids=["no-config", "no-special-tokens-map", "neither"],
)
def test_optional_files_do_not_change_what_is_loaded(make_model_dir, dropped) -> None:
"""Both files are redundant: everything they carry is already in the tokenizer."""
baseline, baseline_specials = load_tokenizer(make_model_dir())
tokenizer, specials = load_tokenizer(make_model_dir(drop_files=dropped))
assert specials == baseline_specials
assert tokenizer.padding == baseline.padding
assert tokenizer.encode("hello world").ids == baseline.encode("hello world").ids
@pytest.mark.parametrize("missing", ("tokenizer.json", "tokenizer_config.json"))
def test_the_remaining_files_are_still_required(make_model_dir, missing) -> None:
"""Relaxing the optional two must not relax the two that carry irreplaceable data."""
model_dir = make_model_dir(drop_files=(missing,))
with pytest.raises(ValueError, match=f"Could not find {missing}"):
load_tokenizer(model_dir)
@pytest.mark.parametrize(
"model_files",
[
pytest.param({"drop_from_config": ("pad_token_id",)}, id="config-omits-pad-token-id"),
pytest.param({"drop_files": ("config.json",)}, id="config-is-absent"),
],
)
def test_pad_id_falls_back_to_the_vocabulary(make_model_dir, token_id, model_files) -> None:
"""Last link of the chain; a hardcoded 0 would silently disagree with `pad_token`."""
expected = token_id("[SEP]")
assert expected != 0, "a pad token whose id is 0 would pass even without a lookup"
model_dir = make_model_dir(tokenizer_config={"pad_token": "[SEP]"}, **model_files)
tokenizer, _ = load_tokenizer(model_dir)
assert tokenizer.padding["pad_id"] == expected
def test_pad_token_that_resolves_nowhere_raises(make_model_dir) -> None:
"""Without config.json a pad token outside the vocabulary has no id left to fall back on."""
model_dir = make_model_dir(
tokenizer_config={"pad_token": "[NOT_IN_VOCAB]"},
drop_files=("config.json",),
)
with pytest.raises(ValueError, match="Could not resolve an id for the pad token"):
load_tokenizer(model_dir)
def test_pad_token_named_only_in_the_map_resolves(make_model_dir) -> None:
"""The map is read first, so it can name a pad token tokenizer.json does not carry."""
model_dir = make_model_dir(
tokenizer_config={"pad_token": "<|mypad|>"},
drop_files=("config.json",),
)
_patch_json(model_dir / "special_tokens_map.json", {"pad_token": "<|mypad|>"})
tokenizer, specials = load_tokenizer(model_dir)
assert tokenizer.padding["pad_token"] == "<|mypad|>"
assert tokenizer.padding["pad_id"] == specials["<|mypad|>"]
@pytest.mark.parametrize(
"additional",
[
pytest.param(["<|list_str|>"], id="list-of-strings"),
pytest.param([{"content": "<|list_str|>"}], id="list-of-added-token-dicts"),
],
)
def test_list_valued_map_entries_are_registered(make_model_dir, additional) -> None:
"""`additional_special_tokens` holds a list, which the str/dict dispatch alone drops.
Real repos ship both spellings, and their tokens are in tokenizer.json already, so
only a token living nowhere else shows the drop.
"""
model_dir = make_model_dir()
_patch_json(model_dir / "special_tokens_map.json", {"additional_special_tokens": additional})
_, specials = load_tokenizer(model_dir)
assert "<|list_str|>" in specials
+271 -80
View File
@@ -1,14 +1,16 @@
import os
from contextlib import contextmanager
import pytest
import numpy as np
from fastembed.sparse.bm25 import Bm25
from fastembed.sparse.sparse_text_embedding import SparseTextEmbedding
from tests.utils import delete_model_cache
from tests.utils import delete_model_cache, should_test_model
CANONICAL_COLUMN_VALUES = {
"prithvida/Splade_PP_en_v1": {
"prithivida/Splade_PP_en_v1": {
"indices": [
2040,
2047,
@@ -43,112 +45,249 @@ CANONICAL_COLUMN_VALUES = {
2.1904349327087402,
1.0531445741653442,
],
}
},
"Qdrant/minicoil-v1": {
"indices": [80, 81, 82, 83, 6664, 6665, 6666, 6667],
"values": [
0.52634597,
0.8711344,
1.2264385,
0.52123857,
0.974713,
-0.97803956,
-0.94312465,
-0.12508166,
],
},
# first 15 non-zero dimensions of the embedding
"opensearch-project/opensearch-neural-sparse-encoding-doc-v3-gte": {
"indices": [
999,
1010,
1011,
1024,
1028,
1029,
1045,
1074,
1993,
2017,
2033,
2054,
2073,
2080,
2088,
],
"values": [
0.16544909,
0.00529129,
0.0392109,
0.12337475,
0.09640586,
0.05325737,
0.09611791,
0.03159865,
0.01349991,
0.09392473,
0.01928805,
0.05238346,
0.05515401,
0.03156782,
0.98263124,
],
},
}
CANONICAL_QUERY_VALUES = {
"opensearch-project/opensearch-neural-sparse-encoding-doc-v3-gte": {
"indices": [2088, 7592],
"values": [3.42086864, 6.93775654],
},
"Qdrant/minicoil-v1": {
"indices": [80, 81, 82, 83, 6664, 6665, 6666, 6667],
"values": [
0.31389374,
0.5195128,
0.7314033,
0.3108479,
0.5812834,
-0.5832673,
-0.5624452,
-0.0745942,
],
},
}
_MODELS_TO_CACHE = (
"prithivida/Splade_PP_en_v1",
"Qdrant/minicoil-v1",
"Qdrant/bm25",
"Qdrant/bm42-all-minilm-l6-v2-attentions",
"opensearch-project/opensearch-neural-sparse-encoding-doc-v3-gte",
)
MODELS_TO_CACHE = tuple([x.lower() for x in _MODELS_TO_CACHE])
@pytest.fixture(scope="module")
def model_cache():
is_ci = os.getenv("CI")
cache = {}
@contextmanager
def get_model(model_name: str):
lowercase_model_name = model_name.lower()
if lowercase_model_name not in cache:
cache[lowercase_model_name] = SparseTextEmbedding(lowercase_model_name)
yield cache[lowercase_model_name]
if lowercase_model_name not in MODELS_TO_CACHE:
model_inst = cache.pop(lowercase_model_name)
if is_ci:
delete_model_cache(model_inst.model._model_dir)
del model_inst
yield get_model
if is_ci:
for name, model in cache.items():
delete_model_cache(model.model._model_dir)
cache.clear()
docs = ["Hello World"]
def test_batch_embedding() -> None:
is_ci = os.getenv("CI")
@pytest.mark.parametrize(
"model_name",
["prithivida/Splade_PP_en_v1", "Qdrant/minicoil-v1"],
)
def test_batch_embedding(model_cache, model_name: str) -> None:
docs_to_embed = docs * 10
for model_name, expected_result in CANONICAL_COLUMN_VALUES.items():
model = SparseTextEmbedding(model_name=model_name)
with model_cache(model_name) as model:
result = next(iter(model.embed(docs_to_embed, batch_size=6)))
expected_result = CANONICAL_COLUMN_VALUES[model_name]
assert result.indices.tolist() == expected_result["indices"]
for i, value in enumerate(result.values):
assert pytest.approx(value, abs=0.001) == expected_result["values"][i]
if is_ci:
delete_model_cache(model.model._model_dir)
def test_single_embedding() -> None:
def test_single_embedding(model_cache) -> None:
is_ci = os.getenv("CI")
for model_name, expected_result in CANONICAL_COLUMN_VALUES.items():
model = SparseTextEmbedding(model_name=model_name)
is_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
passage_result = next(iter(model.embed(docs, batch_size=6)))
query_result = next(iter(model.query_embed(docs)))
for result in [passage_result, query_result]:
assert result.indices.tolist() == expected_result["indices"]
for model_desc in SparseTextEmbedding._list_supported_models():
if (
model_desc.model not in CANONICAL_COLUMN_VALUES
): # attention models and bm25 are also parts of
# SparseTextEmbedding, however, they have their own tests
continue
if not should_test_model(model_desc, model_desc.model, is_ci, is_manual):
continue
with model_cache(model_desc.model) as model:
passage_result = next(iter(model.embed(docs, batch_size=6)))
query_result = next(iter(model.query_embed(docs)))
expected_result = CANONICAL_COLUMN_VALUES[model_desc.model]
expected_query_result = CANONICAL_QUERY_VALUES.get(model_desc.model, expected_result)
for i, value in enumerate(result.values):
# canonical values might contain only a prefix of the non-zero dimensions
num_dims = len(expected_result["indices"])
assert passage_result.indices.tolist()[:num_dims] == expected_result["indices"]
for i, value in enumerate(passage_result.values[:num_dims]):
assert pytest.approx(value, abs=0.001) == expected_result["values"][i]
if is_ci:
delete_model_cache(model.model._model_dir)
num_query_dims = len(expected_query_result["indices"])
assert (
query_result.indices.tolist()[:num_query_dims] == expected_query_result["indices"]
)
for i, value in enumerate(query_result.values[:num_query_dims]):
assert pytest.approx(value, abs=0.001) == expected_query_result["values"][i]
def test_parallel_processing() -> None:
is_ci = os.getenv("CI")
model = SparseTextEmbedding(model_name="prithivida/Splade_PP_en_v1")
docs = ["hello world", "flag embedding"] * 30
sparse_embeddings_duo = list(model.embed(docs, batch_size=10, parallel=2))
sparse_embeddings_all = list(model.embed(docs, batch_size=10, parallel=0))
sparse_embeddings = list(model.embed(docs, batch_size=10, parallel=None))
@pytest.mark.parametrize(
"model_name",
["prithivida/Splade_PP_en_v1", "Qdrant/minicoil-v1"],
)
def test_parallel_processing(model_cache, model_name: str) -> None:
with model_cache(model_name) as model:
docs = ["hello world", "flag embedding"] * 30
sparse_embeddings_duo = list(model.embed(docs, batch_size=10, parallel=2))
# sparse_embeddings_all = list(model.embed(docs, batch_size=10, parallel=0)) # inherits OnnxTextModel which
# is tested in TextEmbedding, disabling it here to reduce number of requests to hf
# multiprocessing is enough to test with `parallel=2`, and `parallel=None` is okay to tests since it reuses
# model from cache
sparse_embeddings = list(model.embed(docs, batch_size=10, parallel=None))
assert (
len(sparse_embeddings)
== len(sparse_embeddings_duo)
== len(sparse_embeddings_all)
== len(docs)
)
for sparse_embedding, sparse_embedding_duo, sparse_embedding_all in zip(
sparse_embeddings, sparse_embeddings_duo, sparse_embeddings_all
):
assert (
sparse_embedding.indices.tolist()
== sparse_embedding_duo.indices.tolist()
== sparse_embedding_all.indices.tolist()
len(sparse_embeddings)
== len(sparse_embeddings_duo)
# == len(sparse_embeddings_all)
== len(docs)
)
assert np.allclose(sparse_embedding.values, sparse_embedding_duo.values, atol=1e-3)
assert np.allclose(sparse_embedding.values, sparse_embedding_all.values, atol=1e-3)
if is_ci:
delete_model_cache(model.model._model_dir)
for (
sparse_embedding,
sparse_embedding_duo,
# sparse_embedding_all
) in zip(
sparse_embeddings,
sparse_embeddings_duo,
# sparse_embeddings_all
):
assert (
sparse_embedding.indices.tolist() == sparse_embedding_duo.indices.tolist()
# == sparse_embedding_all.indices.tolist()
)
assert np.allclose(sparse_embedding.values, sparse_embedding_duo.values, atol=1e-3)
# assert np.allclose(sparse_embedding.values, sparse_embedding_all.values, atol=1e-3)
@pytest.fixture
def bm25_instance() -> None:
ci = os.getenv("CI", True)
model = Bm25("Qdrant/bm25", language="english")
yield model
if ci:
delete_model_cache(model._model_dir)
def test_stem_with_stopwords_and_punctuation(model_cache) -> None:
with model_cache("Qdrant/bm25") as model:
bm25_instance = model.model
# Setup
original_stopwords = bm25_instance.stopwords.copy()
original_punctuation = bm25_instance.punctuation.copy()
bm25_instance.stopwords = {"the", "is", "a"}
bm25_instance.punctuation = {".", ",", "!"}
# Test data
tokens = ["The", "quick", "brown", "fox", "is", "a", "test", "sentence", ".", "!"]
# Execute
result = bm25_instance._stem(tokens)
# Assert
expected = ["quick", "brown", "fox", "test", "sentenc"]
assert result == expected, f"Expected {expected}, but got {result}"
bm25_instance.stopwords = original_stopwords
bm25_instance.punctuation = original_punctuation
def test_stem_with_stopwords_and_punctuation(bm25_instance: Bm25) -> None:
# Setup
bm25_instance.stopwords = {"the", "is", "a"}
bm25_instance.punctuation = {".", ",", "!"}
def test_stem_case_insensitive_stopwords(model_cache) -> None:
with model_cache("Qdrant/bm25") as model:
bm25_instance = model.model
original_stopwords = bm25_instance.stopwords.copy()
original_punctuation = bm25_instance.punctuation.copy()
# Test data
tokens = ["The", "quick", "brown", "fox", "is", "a", "test", "sentence", ".", "!"]
# Setup
bm25_instance.stopwords = {"the", "is", "a"}
bm25_instance.punctuation = {".", ",", "!"}
# Execute
result = bm25_instance._stem(tokens)
# Test data
tokens = ["THE", "Quick", "Brown", "Fox", "IS", "A", "Test", "Sentence", ".", "!"]
# Assert
expected = ["quick", "brown", "fox", "test", "sentenc"]
assert result == expected, f"Expected {expected}, but got {result}"
# Execute
result = bm25_instance._stem(tokens)
def test_stem_case_insensitive_stopwords(bm25_instance: Bm25) -> None:
# Setup
bm25_instance.stopwords = {"the", "is", "a"}
bm25_instance.punctuation = {".", ",", "!"}
# Test data
tokens = ["THE", "Quick", "Brown", "Fox", "IS", "A", "Test", "Sentence", ".", "!"]
# Execute
result = bm25_instance._stem(tokens)
# Assert
expected = ["quick", "brown", "fox", "test", "sentenc"]
assert result == expected, f"Expected {expected}, but got {result}"
# Assert
expected = ["quick", "brown", "fox", "test", "sentenc"]
assert result == expected, f"Expected {expected}, but got {result}"
bm25_instance.stopwords = original_stopwords
bm25_instance.punctuation = original_punctuation
@pytest.mark.parametrize("disable_stemmer", [True, False])
@@ -172,10 +311,23 @@ def test_disable_stemmer_behavior(disable_stemmer: bool) -> None:
assert result == expected, f"Expected {expected}, but got {result}"
@pytest.mark.parametrize(
"model_name",
["prithivida/Splade_PP_en_v1"],
)
def test_if_splade_query_embed_is_inference_free() -> None:
is_ci = os.getenv("CI")
model = SparseTextEmbedding(
model_name="opensearch-project/opensearch-neural-sparse-encoding-doc-v3-gte",
lazy_load=True,
)
embeddings = list(model.query_embed(["hello world", "flag embedding"]))
# queries are embedded with a tokenizer and an idf lookup table only,
# the onnx model must stay unloaded
assert not hasattr(model.model, "model")
assert all(len(embedding.indices) > 0 for embedding in embeddings)
if is_ci:
delete_model_cache(model.model._model_dir)
@pytest.mark.parametrize("model_name", ["prithivida/Splade_PP_en_v1"])
def test_lazy_load(model_name: str) -> None:
is_ci = os.getenv("CI")
model = SparseTextEmbedding(model_name=model_name, lazy_load=True)
@@ -193,3 +345,42 @@ def test_lazy_load(model_name: str) -> None:
if is_ci:
delete_model_cache(model.model._model_dir)
@pytest.mark.parametrize(
"model_name",
[
"prithivida/Splade_PP_en_v1",
"Qdrant/minicoil-v1",
"Qdrant/bm42-all-minilm-l6-v2-attentions",
],
)
def test_session_options(model_cache, model_name) -> None:
with model_cache(model_name) as default_model:
default_session_options = default_model.model.model.get_session_options()
assert default_session_options.enable_cpu_mem_arena is True
model = SparseTextEmbedding(model_name=model_name, enable_cpu_mem_arena=False)
session_options = model.model.model.get_session_options()
assert session_options.enable_cpu_mem_arena is False
@pytest.mark.parametrize(
"model_name",
[
"prithivida/Splade_PP_en_v1",
"Qdrant/minicoil-v1",
"Qdrant/bm42-all-minilm-l6-v2-attentions",
"Qdrant/bm25",
],
)
def test_token_count(model_cache, model_name) -> None:
with model_cache(model_name) as model:
documents = [
"Name me a couple of cities were the capitals of Germany?",
"Berlin is the current capital of Germany, Bonn is a former capital of Germany.",
]
first_doc_token_count = model.token_count(documents[0])
second_doc_token_count = model.token_count(documents[1])
doc_token_count = model.token_count(documents)
assert first_doc_token_count + second_doc_token_count == doc_token_count
assert doc_token_count == model.token_count(documents, batch_size=1)
+105 -73
View File
@@ -1,10 +1,11 @@
import os
from contextlib import contextmanager
import numpy as np
import pytest
from fastembed.rerank.cross_encoder import TextCrossEncoder
from tests.utils import delete_model_cache
from tests.utils import delete_model_cache, should_test_model
CANONICAL_SCORE_VALUES = {
"Xenova/ms-marco-MiniLM-L-6-v2": np.array([8.500708, -2.541011]),
@@ -15,73 +16,84 @@ CANONICAL_SCORE_VALUES = {
"jinaai/jina-reranker-v2-base-multilingual": np.array([1.6533, -1.6455]),
}
SELECTED_MODELS = {
"Xenova": "Xenova/ms-marco-MiniLM-L-6-v2",
"BAAI": "BAAI/bge-reranker-base",
"jinaai": "jinaai/jina-reranker-v1-tiny-en",
}
_MODELS_TO_CACHE = ("Xenova/ms-marco-MiniLM-L-6-v2",)
MODELS_TO_CACHE = tuple([x.lower() for x in _MODELS_TO_CACHE])
@pytest.mark.parametrize(
"model_name",
[model_name for model_name in CANONICAL_SCORE_VALUES],
)
def test_rerank(model_name: str) -> None:
@pytest.fixture(scope="module")
def model_cache():
is_ci = os.getenv("CI")
cache = {}
model = TextCrossEncoder(model_name=model_name)
@contextmanager
def get_model(model_name: str):
lowercase_model_name = model_name.lower()
if lowercase_model_name not in cache:
cache[lowercase_model_name] = TextCrossEncoder(lowercase_model_name)
yield cache[lowercase_model_name]
if lowercase_model_name not in MODELS_TO_CACHE:
model_inst = cache.pop(lowercase_model_name)
if is_ci:
delete_model_cache(model_inst.model._model_dir)
del model_inst
query = "What is the capital of France?"
documents = ["Paris is the capital of France.", "Berlin is the capital of Germany."]
scores = np.array(list(model.rerank(query, documents)))
yield get_model
pairs = [(query, doc) for doc in documents]
scores2 = np.array(list(model.rerank_pairs(pairs)))
assert np.allclose(
scores, scores2, atol=1e-5
), f"Model: {model_name}, Scores: {scores}, Scores2: {scores2}"
canonical_scores = CANONICAL_SCORE_VALUES[model_name]
assert np.allclose(
scores, canonical_scores, atol=1e-3
), f"Model: {model_name}, Scores: {scores}, Expected: {canonical_scores}"
if is_ci:
delete_model_cache(model.model._model_dir)
for name, model in cache.items():
delete_model_cache(model.model._model_dir)
cache.clear()
@pytest.mark.parametrize(
"model_name",
[model_name for model_name in SELECTED_MODELS.values()],
)
def test_batch_rerank(model_name: str) -> None:
@pytest.mark.parametrize("model_name", ["Xenova/ms-marco-MiniLM-L-6-v2"])
def test_rerank(model_cache, model_name: str) -> None:
is_ci = os.getenv("CI")
is_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
model = TextCrossEncoder(model_name=model_name)
for model_desc in TextCrossEncoder._list_supported_models():
if not should_test_model(model_desc, model_name, is_ci, is_manual):
continue
query = "What is the capital of France?"
documents = ["Paris is the capital of France.", "Berlin is the capital of Germany."] * 50
scores = np.array(list(model.rerank(query, documents, batch_size=10)))
with model_cache(model_desc.model) as model:
query = "What is the capital of France?"
documents = ["Paris is the capital of France.", "Berlin is the capital of Germany."]
scores = np.array(list(model.rerank(query, documents)))
pairs = [(query, doc) for doc in documents]
scores2 = np.array(list(model.rerank_pairs(pairs)))
assert np.allclose(
scores, scores2, atol=1e-5
), f"Model: {model_name}, Scores: {scores}, Scores2: {scores2}"
pairs = [(query, doc) for doc in documents]
scores2 = np.array(list(model.rerank_pairs(pairs)))
assert np.allclose(
scores, scores2, atol=1e-5
), f"Model: {model_desc.model}, Scores: {scores}, Scores2: {scores2}"
canonical_scores = np.tile(CANONICAL_SCORE_VALUES[model_name], 50)
assert scores.shape == canonical_scores.shape, f"Unexpected shape for model {model_name}"
assert np.allclose(
scores, canonical_scores, atol=1e-3
), f"Model: {model_name}, Scores: {scores}, Expected: {canonical_scores}"
if is_ci:
delete_model_cache(model.model._model_dir)
canonical_scores = CANONICAL_SCORE_VALUES[model_desc.model]
assert np.allclose(
scores, canonical_scores, atol=1e-3
), f"Model: {model_desc.model}, Scores: {scores}, Expected: {canonical_scores}"
@pytest.mark.parametrize(
"model_name",
["Xenova/ms-marco-MiniLM-L-6-v2"],
)
@pytest.mark.parametrize("model_name", ["Xenova/ms-marco-MiniLM-L-6-v2"])
def test_batch_rerank(model_cache, model_name: str) -> None:
with model_cache(model_name) as model:
query = "What is the capital of France?"
documents = ["Paris is the capital of France.", "Berlin is the capital of Germany."] * 50
scores = np.array(list(model.rerank(query, documents, batch_size=10)))
pairs = [(query, doc) for doc in documents]
scores2 = np.array(list(model.rerank_pairs(pairs)))
assert np.allclose(
scores, scores2, atol=1e-5
), f"Model: {model_name}, Scores: {scores}, Scores2: {scores2}"
canonical_scores = np.tile(CANONICAL_SCORE_VALUES[model_name], 50)
assert scores.shape == canonical_scores.shape, f"Unexpected shape for model {model_name}"
assert np.allclose(
scores, canonical_scores, atol=1e-3
), f"Model: {model_name}, Scores: {scores}, Expected: {canonical_scores}"
@pytest.mark.parametrize("model_name", ["Xenova/ms-marco-MiniLM-L-6-v2"])
def test_lazy_load(model_name: str) -> None:
is_ci = os.getenv("CI")
model = TextCrossEncoder(model_name=model_name, lazy_load=True)
@@ -95,25 +107,45 @@ def test_lazy_load(model_name: str) -> None:
delete_model_cache(model.model._model_dir)
@pytest.mark.parametrize(
"model_name",
[model_name for model_name in SELECTED_MODELS.values()],
)
def test_rerank_pairs_parallel(model_name: str) -> None:
is_ci = os.getenv("CI")
@pytest.mark.parametrize("model_name", ["Xenova/ms-marco-MiniLM-L-6-v2"])
def test_rerank_pairs_parallel(model_cache, model_name: str) -> None:
with model_cache(model_name) as model:
query = "What is the capital of France?"
documents = ["Paris is the capital of France.", "Berlin is the capital of Germany."] * 10
pairs = [(query, doc) for doc in documents]
scores_parallel = np.array(list(model.rerank_pairs(pairs, parallel=2, batch_size=10)))
scores_sequential = np.array(list(model.rerank_pairs(pairs, batch_size=10)))
assert np.allclose(
scores_parallel, scores_sequential, atol=1e-5
), f"Model: {model_name}, Scores (Parallel): {scores_parallel}, Scores (Sequential): {scores_sequential}"
canonical_scores = CANONICAL_SCORE_VALUES[model_name]
assert np.allclose(
scores_parallel[: len(canonical_scores)], canonical_scores, atol=1e-3
), f"Model: {model_name}, Scores (Parallel): {scores_parallel}, Expected: {canonical_scores}"
model = TextCrossEncoder(model_name=model_name)
query = "What is the capital of France?"
documents = ["Paris is the capital of France.", "Berlin is the capital of Germany."] * 10
pairs = [(query, doc) for doc in documents]
scores_parallel = np.array(list(model.rerank_pairs(pairs, parallel=2, batch_size=10)))
scores_sequential = np.array(list(model.rerank_pairs(pairs, batch_size=10)))
assert np.allclose(
scores_parallel, scores_sequential, atol=1e-5
), f"Model: {model_name}, Scores (Parallel): {scores_parallel}, Scores (Sequential): {scores_sequential}"
canonical_scores = CANONICAL_SCORE_VALUES[model_name]
assert np.allclose(
scores_parallel[: len(canonical_scores)], canonical_scores, atol=1e-3
), f"Model: {model_name}, Scores (Parallel): {scores_parallel}, Expected: {canonical_scores}"
if is_ci:
delete_model_cache(model.model._model_dir)
@pytest.mark.parametrize("model_name", ["Xenova/ms-marco-MiniLM-L-6-v2"])
def test_token_count(model_cache, model_name: str) -> None:
with model_cache(model_name) as model:
pairs = [
("What is the capital of France?", "Paris is the capital of France."),
(
"Name me a couple of cities were the capitals of Germany?",
"Berlin is the current capital of Germany, Bonn is a former capital of Germany.",
),
]
first_pair_token_count = model.token_count([pairs[0]])
second_pair_token_count = model.token_count([pairs[1]])
pairs_token_count = model.token_count(pairs)
assert first_pair_token_count + second_pair_token_count == pairs_token_count
assert pairs_token_count == model.token_count(pairs, batch_size=1)
@pytest.mark.parametrize("model_name", ["Xenova/ms-marco-MiniLM-L-6-v2"])
def test_session_options(model_cache, model_name) -> None:
with model_cache(model_name) as default_model:
default_session_options = default_model.model.model.get_session_options()
assert default_session_options.enable_cpu_mem_arena is True
model = TextCrossEncoder(model_name=model_name, enable_cpu_mem_arena=False)
session_options = model.model.model.get_session_options()
assert session_options.enable_cpu_mem_arena is False
+64 -76
View File
@@ -4,7 +4,7 @@ import numpy as np
import pytest
from fastembed import TextEmbedding
from fastembed.text.multitask_embedding import Task
from fastembed.text.multitask_embedding import JinaEmbeddingV3, Task
from tests.utils import delete_model_cache
@@ -60,52 +60,43 @@ CANONICAL_VECTOR_VALUES = {
docs = ["Hello World", "Follow the white rabbit."]
def test_batch_embedding():
@pytest.mark.parametrize("dim,model_name", [(1024, "jinaai/jina-embeddings-v3")])
def test_batch_embedding(dim: int, model_name: str):
is_ci = os.getenv("CI")
is_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
if is_ci and not is_manual:
pytest.skip("Skipping multitask models in CI non-manual mode")
docs_to_embed = docs * 10
default_task = Task.RETRIEVAL_PASSAGE
for model_desc in TextEmbedding._list_supported_models():
if not is_ci and model_desc.size_in_GB > 1:
continue
model = TextEmbedding(model_name=model_name)
model_name = model_desc.model
dim = model_desc.dim
embeddings = list(model.embed(documents=docs_to_embed, batch_size=6))
embeddings = np.stack(embeddings, axis=0)
if model_name not in CANONICAL_VECTOR_VALUES.keys():
continue
assert embeddings.shape == (len(docs_to_embed), dim)
model = TextEmbedding(model_name=model_name)
canonical_vector = CANONICAL_VECTOR_VALUES[model_name][default_task]["vectors"]
assert np.allclose(
embeddings[: len(docs), : canonical_vector.shape[1]], canonical_vector, atol=1e-4
), model_name
print(f"evaluating {model_name} default task")
embeddings = list(model.embed(documents=docs_to_embed, batch_size=6))
embeddings = np.stack(embeddings, axis=0)
assert embeddings.shape == (len(docs_to_embed), dim)
canonical_vector = CANONICAL_VECTOR_VALUES[model_name][default_task]["vectors"]
assert np.allclose(
embeddings[: len(docs), : canonical_vector.shape[1]], canonical_vector, atol=1e-4
), model_desc.model
if is_ci:
delete_model_cache(model.model._model_dir)
if is_ci:
delete_model_cache(model.model._model_dir)
def test_single_embedding():
is_ci = os.getenv("CI")
is_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
if is_ci and not is_manual:
pytest.skip("Skipping multitask models in CI non-manual mode")
for model_desc in TextEmbedding._list_supported_models():
if not is_ci and model_desc.size_in_GB > 1:
continue
for model_desc in JinaEmbeddingV3._list_supported_models():
# todo: once we add more models, we should not test models >1GB size locally
model_name = model_desc.model
dim = model_desc.dim
if model_name not in CANONICAL_VECTOR_VALUES.keys():
continue
model = TextEmbedding(model_name=model_name)
for task in CANONICAL_VECTOR_VALUES[model_name]:
@@ -118,27 +109,42 @@ def test_single_embedding():
canonical_vector = task["vectors"]
assert np.allclose(
embeddings[: len(docs), : canonical_vector.shape[1]], canonical_vector, atol=1e-4
embeddings[:, : canonical_vector.shape[1]], canonical_vector, atol=1e-4
), model_desc.model
classification_embeddings = list(model.embed(documents=docs, task_id=Task.CLASSIFICATION))
classification_embeddings = np.stack(classification_embeddings, axis=0)
assert classification_embeddings.shape == (len(docs), dim)
model = TextEmbedding(model_name=model_name, task_id=Task.CLASSIFICATION)
default_embeddings = list(model.embed(documents=docs))
default_embeddings = np.stack(default_embeddings, axis=0)
assert default_embeddings.shape == (len(docs), dim)
assert np.allclose(
classification_embeddings,
default_embeddings,
atol=1e-4,
), model_desc.model
if is_ci:
delete_model_cache(model.model._model_dir)
def test_single_embedding_query():
is_ci = os.getenv("CI")
is_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
if is_ci and not is_manual:
pytest.skip("Skipping multitask models in CI non-manual mode")
task_id = Task.RETRIEVAL_QUERY
for model_desc in TextEmbedding._list_supported_models():
if not is_ci and model_desc.size_in_GB > 1:
continue
for model_desc in JinaEmbeddingV3._list_supported_models():
# todo: once we add more models, we should not test models >1GB size locally
model_name = model_desc.model
dim = model_desc.dim
if model_name not in CANONICAL_VECTOR_VALUES.keys():
continue
model = TextEmbedding(model_name=model_name)
print(f"evaluating {model_name} query_embed task_id: {task_id}")
@@ -150,7 +156,7 @@ def test_single_embedding_query():
canonical_vector = CANONICAL_VECTOR_VALUES[model_name][task_id]["vectors"]
assert np.allclose(
embeddings[: len(docs), : canonical_vector.shape[1]], canonical_vector, atol=1e-4
embeddings[:, : canonical_vector.shape[1]], canonical_vector, atol=1e-4
), model_desc.model
if is_ci:
@@ -159,18 +165,18 @@ def test_single_embedding_query():
def test_single_embedding_passage():
is_ci = os.getenv("CI")
is_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
if is_ci and not is_manual:
pytest.skip("Skipping multitask models in CI non-manual mode")
task_id = Task.RETRIEVAL_PASSAGE
for model_desc in TextEmbedding._list_supported_models():
if not is_ci and model_desc.size_in_GB > 1:
continue
for model_desc in JinaEmbeddingV3._list_supported_models():
# todo: once we add more models, we should not test models >1GB size locally
model_name = model_desc.model
dim = model_desc.dim
if model_name not in CANONICAL_VECTOR_VALUES.keys():
continue
model = TextEmbedding(model_name=model_name)
print(f"evaluating {model_name} passage_embed task_id: {task_id}")
@@ -182,21 +188,22 @@ def test_single_embedding_passage():
canonical_vector = CANONICAL_VECTOR_VALUES[model_name][task_id]["vectors"]
assert np.allclose(
embeddings[: len(docs), : canonical_vector.shape[1]], canonical_vector, atol=1e-4
embeddings[:, : canonical_vector.shape[1]], canonical_vector, atol=1e-4
), model_desc.model
if is_ci:
delete_model_cache(model.model._model_dir)
def test_parallel_processing():
@pytest.mark.parametrize("dim,model_name", [(1024, "jinaai/jina-embeddings-v3")])
def test_parallel_processing(dim: int, model_name: str):
is_ci = os.getenv("CI")
is_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
if is_ci and not is_manual:
pytest.skip("Skipping in CI non-manual mode")
docs = ["Hello World", "Follow the white rabbit."] * 10
model_name = "jinaai/jina-embeddings-v3"
dim = 1024
model = TextEmbedding(model_name=model_name)
task_id = Task.SEPARATION
@@ -216,33 +223,14 @@ def test_parallel_processing():
delete_model_cache(model.model._model_dir)
def test_task_assignment():
is_ci = os.getenv("CI")
for model_desc in TextEmbedding._list_supported_models():
if not is_ci and model_desc.size_in_GB > 1:
continue
model_name = model_desc.model
if model_name not in CANONICAL_VECTOR_VALUES.keys():
continue
model = TextEmbedding(model_name=model_name)
for i, task_id in enumerate(Task):
_ = list(model.embed(documents=docs, batch_size=1, task_id=i))
assert model.model.current_task_id == task_id
if is_ci:
delete_model_cache(model.model._model_dir)
@pytest.mark.parametrize(
"model_name",
["jinaai/jina-embeddings-v3"],
)
@pytest.mark.parametrize("model_name", ["jinaai/jina-embeddings-v3"])
def test_lazy_load(model_name: str):
is_ci = os.getenv("CI")
is_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
if is_ci and not is_manual:
pytest.skip("Skipping in CI non-manual mode")
model = TextEmbedding(model_name=model_name, lazy_load=True)
assert not hasattr(model.model, "model")
+227 -59
View File
@@ -1,11 +1,14 @@
import os
import platform
from contextlib import contextmanager
import numpy as np
import pytest
from fastembed.text.last_token_normalized_embedding import LastTokenNormalizedEmbedding
from fastembed.text.onnx_embedding import OnnxTextEmbedding
from fastembed.text.text_embedding import TextEmbedding
from tests.utils import delete_model_cache
from tests.utils import delete_model_cache, should_test_model
CANONICAL_VECTOR_VALUES = {
"BAAI/bge-small-en": np.array([-0.0232, -0.0255, 0.0174, -0.0639, -0.0006]),
@@ -67,86 +70,194 @@ CANONICAL_VECTOR_VALUES = {
"Qdrant/clip-ViT-B-32-text": np.array([0.0083, 0.0103, -0.0138, 0.0199, -0.0069]),
"thenlper/gte-base": np.array([0.0038, 0.0355, 0.0181, 0.0092, 0.0654]),
"jinaai/jina-clip-v1": np.array([-0.0862, -0.0101, -0.0056, 0.0375, -0.0472]),
"google/embeddinggemma-300m": np.array(
[-0.08181356, 0.0214127, 0.05120273, -0.03690156, -0.0254504]
),
"Qwen/Qwen3-Embedding-0.6B": np.array(
[-0.01476084, 0.01723184, -0.01195498, -0.07275258, 0.00281229]
),
"Qwen/Qwen3-Embedding-0.6B-Q": np.array(
[-0.01599521, 0.01676456, -0.01195119, -0.07132675, 0.00346729]
),
"google/siglip2-base-patch16-224": np.array(
[-0.01181389, 0.00737596, 0.01118064, 0.0103095, 0.3451049]
),
"minishlab/potion-base-8M": np.array(
[-0.03432461, -0.08020256, -0.14396408, 0.08480079, 0.01958815]
),
"minishlab/potion-retrieval-32M": np.array(
[0.019733, -0.01530093, -0.08678473, 0.0229059, 0.04700558]
),
"minishlab/potion-multilingual-128M": np.array(
[0.02366836, 0.02973341, 0.05140258, -0.00745248, -0.06740689]
),
}
QWEN3_INSTRUCT_PREFIX = (
"Instruct: Given a web search query, retrieve relevant passages that answer the query\nQuery:"
)
DOC_PREFIXES = {
"google/embeddinggemma-300m": "title: none | text: ",
}
QUERY_PREFIXES = {
"google/embeddinggemma-300m": "task: search result | query: ",
"Qwen/Qwen3-Embedding-0.6B": QWEN3_INSTRUCT_PREFIX,
"Qwen/Qwen3-Embedding-0.6B-Q": QWEN3_INSTRUCT_PREFIX,
}
CANONICAL_QUERY_VECTOR_VALUES = {
"google/embeddinggemma-300m": np.array(
[-0.22990295, 0.03311195, 0.04290345, -0.03558498, -0.01399477]
),
"Qwen/Qwen3-Embedding-0.6B": np.array(
[-0.01908712, 0.01635596, -0.00356586, -0.03947155, -0.01387356]
),
"Qwen/Qwen3-Embedding-0.6B-Q": np.array(
[-0.02221339, 0.01932909, -0.00361797, -0.03888897, -0.01362813]
),
}
MULTI_TASK_MODELS = ["jinaai/jina-embeddings-v3"]
_MODELS_TO_CACHE = ("BAAI/bge-small-en-v1.5",)
MODELS_TO_CACHE = tuple([x.lower() for x in _MODELS_TO_CACHE])
def test_embedding() -> None:
@pytest.fixture(scope="module")
def model_cache():
is_ci = os.getenv("CI")
cache = {}
@contextmanager
def get_model(model_name: str):
lowercase_model_name = model_name.lower()
if lowercase_model_name not in cache:
cache[lowercase_model_name] = TextEmbedding(lowercase_model_name)
yield cache[lowercase_model_name]
if lowercase_model_name not in MODELS_TO_CACHE:
model_inst = cache.pop(lowercase_model_name)
if is_ci:
delete_model_cache(model_inst.model._model_dir)
del model_inst
yield get_model
if is_ci:
for name, model in cache.items():
delete_model_cache(model.model._model_dir)
cache.clear()
@pytest.mark.parametrize("model_name", ["BAAI/bge-small-en-v1.5"])
def test_embedding(model_cache, model_name: str) -> None:
is_ci = os.getenv("CI")
is_mac = platform.system() == "Darwin"
is_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
for model_desc in TextEmbedding._list_supported_models():
if (
(not is_ci and model_desc.size_in_GB > 1)
or model_desc.model in MULTI_TASK_MODELS
or (is_mac and model_desc.model == "nomic-ai/nomic-embed-text-v1.5-Q")
if model_desc.model in MULTI_TASK_MODELS or (
is_mac and model_desc.model == "nomic-ai/nomic-embed-text-v1.5-Q"
):
continue
if not should_test_model(model_desc, model_name, is_ci, is_manual):
continue
dim = model_desc.dim
model = TextEmbedding(model_name=model_desc.model)
docs = ["hello world", "flag embedding"]
embeddings = list(model.embed(docs))
with model_cache(model_desc.model) as model:
docs = ["hello world", "flag embedding"]
if model_desc.model in DOC_PREFIXES:
docs = [DOC_PREFIXES[model_desc.model] + doc for doc in docs]
embeddings = list(model.embed(docs))
embeddings = np.stack(embeddings, axis=0)
assert embeddings.shape == (2, dim)
canonical_vector = CANONICAL_VECTOR_VALUES[model_desc.model]
assert np.allclose(
embeddings[0, : canonical_vector.shape[0]], canonical_vector, atol=1e-3
), model_desc.model
def test_query_embedding(model_cache) -> None:
is_ci = os.getenv("CI")
is_mac = platform.system() == "Darwin"
is_manual = os.getenv("GITHUB_EVENT_NAME") == "workflow_dispatch"
for model_desc in TextEmbedding._list_supported_models():
if model_desc.model in MULTI_TASK_MODELS or (
is_mac and model_desc.model == "nomic-ai/nomic-embed-text-v1.5-Q"
):
continue
if model_desc.model not in CANONICAL_QUERY_VECTOR_VALUES:
continue
if not should_test_model(model_desc, "", is_ci, is_manual):
continue
dim = model_desc.dim
with model_cache(model_desc.model) as model:
queries = ["hello world", "flag embedding"]
if model_desc.model in QUERY_PREFIXES:
queries = [QUERY_PREFIXES[model_desc.model] + query for query in queries]
embeddings = list(model.query_embed(queries))
embeddings = np.stack(embeddings, axis=0)
assert embeddings.shape == (2, dim)
canonical_vector = CANONICAL_QUERY_VECTOR_VALUES[model_desc.model]
assert np.allclose(
embeddings[0, : canonical_vector.shape[0]], canonical_vector, atol=1e-3
), model_desc.model
def test_quantized_model_reports_onnxruntime_requirement(monkeypatch) -> None:
"""Old onnxruntime only implements 4-bit MatMulNBits, the error should say so."""
monkeypatch.setattr(
OnnxTextEmbedding,
"load_onnx_model",
lambda self: (_ for _ in ()).throw(RuntimeError("nbits_ == 4 was false")),
)
model = LastTokenNormalizedEmbedding(
"Qwen/Qwen3-Embedding-0.6B-Q",
lazy_load=True,
specific_model_path="./", # disable model downloading and loading
)
with pytest.raises(RuntimeError, match="onnxruntime>=1.23"):
model.load_onnx_model()
@pytest.mark.parametrize("n_dims,model_name", [(384, "BAAI/bge-small-en-v1.5")])
def test_batch_embedding(model_cache, n_dims: int, model_name: str) -> None:
with model_cache(model_name) as model:
docs = ["hello world", "flag embedding"] * 100
embeddings = list(model.embed(docs, batch_size=10))
embeddings = np.stack(embeddings, axis=0)
assert embeddings.shape == (2, dim)
canonical_vector = CANONICAL_VECTOR_VALUES[model_desc.model]
assert np.allclose(
embeddings[0, : canonical_vector.shape[0]], canonical_vector, atol=1e-3
), model_desc["model"]
if is_ci:
delete_model_cache(model.model._model_dir)
assert embeddings.shape == (len(docs), n_dims)
@pytest.mark.parametrize(
"n_dims,model_name",
[(384, "BAAI/bge-small-en-v1.5"), (768, "jinaai/jina-embeddings-v2-base-en")],
)
def test_batch_embedding(n_dims: int, model_name: str) -> None:
is_ci = os.getenv("CI")
model = TextEmbedding(model_name=model_name)
@pytest.mark.parametrize("n_dims,model_name", [(384, "BAAI/bge-small-en-v1.5")])
def test_parallel_processing(model_cache, n_dims: int, model_name: str) -> None:
with model_cache(model_name) as model:
docs = ["hello world", "flag embedding"] * 100
embeddings = list(model.embed(docs, batch_size=10, parallel=2))
embeddings = np.stack(embeddings, axis=0)
docs = ["hello world", "flag embedding"] * 100
embeddings = list(model.embed(docs, batch_size=10))
embeddings = np.stack(embeddings, axis=0)
embeddings_2 = list(model.embed(docs, batch_size=10, parallel=None))
embeddings_2 = np.stack(embeddings_2, axis=0)
assert embeddings.shape == (200, n_dims)
if is_ci:
delete_model_cache(model.model._model_dir)
embeddings_3 = list(model.embed(docs, batch_size=10, parallel=0))
embeddings_3 = np.stack(embeddings_3, axis=0)
assert embeddings.shape == (len(docs), n_dims)
assert np.allclose(embeddings, embeddings_2, atol=1e-3)
assert np.allclose(embeddings, embeddings_3, atol=1e-3)
@pytest.mark.parametrize(
"n_dims,model_name",
[(384, "BAAI/bge-small-en-v1.5"), (768, "jinaai/jina-embeddings-v2-base-en")],
)
def test_parallel_processing(n_dims: int, model_name: str) -> None:
is_ci = os.getenv("CI")
model = TextEmbedding(model_name=model_name)
docs = ["hello world", "flag embedding"] * 100
embeddings = list(model.embed(docs, batch_size=10, parallel=2))
embeddings = np.stack(embeddings, axis=0)
embeddings_2 = list(model.embed(docs, batch_size=10, parallel=None))
embeddings_2 = np.stack(embeddings_2, axis=0)
embeddings_3 = list(model.embed(docs, batch_size=10, parallel=0))
embeddings_3 = np.stack(embeddings_3, axis=0)
assert embeddings.shape == (200, n_dims)
assert np.allclose(embeddings, embeddings_2, atol=1e-3)
assert np.allclose(embeddings, embeddings_3, atol=1e-3)
if is_ci:
delete_model_cache(model.model._model_dir)
@pytest.mark.parametrize(
"model_name",
["BAAI/bge-small-en-v1.5"],
)
@pytest.mark.parametrize("model_name", ["BAAI/bge-small-en-v1.5"])
def test_lazy_load(model_name: str) -> None:
is_ci = os.getenv("CI")
model = TextEmbedding(model_name=model_name, lazy_load=True)
@@ -163,3 +274,60 @@ def test_lazy_load(model_name: str) -> None:
if is_ci:
delete_model_cache(model.model._model_dir)
def test_get_embedding_size() -> None:
assert TextEmbedding.get_embedding_size("sentence-transformers/all-MiniLM-L6-v2") == 384
assert TextEmbedding.get_embedding_size("sentence-transformers/all-minilm-l6-v2") == 384
def test_embedding_size() -> None:
is_ci = os.getenv("CI")
model_name = "sentence-transformers/all-MiniLM-L6-v2"
model = TextEmbedding(model_name=model_name, lazy_load=True)
assert model.embedding_size == 384
model_name = "sentence-transformers/all-minilm-l6-v2"
model = TextEmbedding(model_name=model_name, lazy_load=True)
assert model.embedding_size == 384
if is_ci:
delete_model_cache(model.model._model_dir)
@pytest.mark.parametrize("model_name", ["sentence-transformers/all-MiniLM-L6-v2"])
def test_session_options(model_cache, model_name) -> None:
with model_cache(model_name) as default_model:
default_session_options = default_model.model.model.get_session_options()
assert default_session_options.enable_cpu_mem_arena is True
model = TextEmbedding(model_name=model_name, enable_cpu_mem_arena=False)
session_options = model.model.model.get_session_options()
assert session_options.enable_cpu_mem_arena is False
@pytest.mark.parametrize("model_name", ["sentence-transformers/all-MiniLM-L6-v2"])
def test_token_count(model_cache, model_name) -> None:
with model_cache(model_name) as model:
documents = [
"Name me a couple of cities were the capitals of Germany?",
"Berlin is the current capital of Germany, Bonn is a former capital of Germany.",
]
first_doc_token_count = model.token_count(documents[0])
second_doc_token_count = model.token_count(documents[1])
doc_token_count = model.token_count(documents)
assert first_doc_token_count + second_doc_token_count == doc_token_count
assert doc_token_count == model.token_count(documents, batch_size=1)
@pytest.mark.parametrize(
"model_name,dim",
[("sentence-transformers/all-MiniLM-L6-v2", 384), ("thenlper/gte-base", 768)],
)
def test_mixed_length_batch_with_fixed_padding(model_cache, model_name: str, dim: int) -> None:
# both models serialize a fixed padding length of 128 in tokenizer.json; gte-base truncates
# at 512, so a document longer than 128 makes the batch ragged unless the padding is relaxed
with model_cache(model_name) as model:
assert model.model.tokenizer.padding["length"] is None
embeddings = np.stack(list(model.embed(["hello world", "retrieval " * 200])), axis=0)
assert embeddings.shape == (2, dim)
+32 -2
View File
@@ -3,10 +3,12 @@ import traceback
from pathlib import Path
from types import TracebackType
from typing import Union, Callable, Any, Type
from typing import Callable, Any, Type
from fastembed.common.model_description import BaseModelDescription
def delete_model_cache(model_dir: Union[str, Path]) -> None:
def delete_model_cache(model_dir: str | Path) -> None:
"""Delete the model cache directory.
If a model was downloaded from the HuggingFace model hub, then _model_dir is the dir to snapshots, removing
@@ -35,3 +37,31 @@ def delete_model_cache(model_dir: Union[str, Path]) -> None:
if model_dir.exists():
# todo: PermissionDenied is raised on blobs removal in Windows, with blobs > 2GB
shutil.rmtree(model_dir, onerror=on_error)
def should_test_model(
model_desc: BaseModelDescription,
autotest_model_name: str,
is_ci: str | None,
is_manual: bool,
):
"""Determine if a model should be tested based on environment
Tests can be run either in ci or locally.
Testing all models each time in ci is too long.
The testing scheme in ci and on a local machine are different, therefore, there are 3 possible scenarios.
1) Run lightweight tests in ci:
- test only one model that has been manually chosen as a representative for a certain class family
2) Run heavyweight (manual) tests in ci:
- test all models
Running tests in ci each time is too expensive, however, it's fine to run it one time with a manual dispatch
3) Run tests locally:
- test all models, which are not too heavy, since network speed might be a bottleneck
"""
if not is_ci:
if model_desc.size_in_GB > 1:
return False
elif not is_manual and model_desc.model != autotest_model_name:
return False
return True