Compare commits

...
50 Commits
Author SHA1 Message Date
George 9b04753f58 new: gpu package (#224) 2026-09-23 01:25:04 +07:00
George Panchuk a107d5c994 add eofl 2026-09-23 01:23:37 +07:00
George Panchuk 20863c9ba9 sync publih with main 2026-09-23 01:23:37 +07:00
George Panchuk 833525a5fa fix: workflow dispatch can only be triggered from the default branch 2026-09-23 01:23:37 +07:00
George Panchuk 05abc3744f refactoring: alter workflow names 2026-09-23 01:23:37 +07:00
George Panchuk 8d590e3fb0 fix: do not run windows and mac os tests on gpu branch 2026-09-23 01:23:37 +07:00
George Panchuk f4db4c5244 new: gpu package publish workflow 2026-09-23 01:23:37 +07:00
George Panchuk a2bef9821d bump version to v0.8.1 2026-09-23 01:23:18 +07:00
George fb68e86c23 fix: stage GCS downloads instead of deleting the caller's cache dir (#718)
* fix: stage GCS downloads instead of deleting the caller's cache dir

* fix: verify archive integrity and give each download its own staging dir
2026-09-23 01:15:54 +07:00
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
40 changed files with 3818 additions and 2007 deletions
+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
+3 -3
View File
@@ -21,9 +21,9 @@ 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.10.x'
- name: Install dependencies
@@ -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 }}
+4 -19
View File
@@ -1,10 +1,11 @@
name: Tests
run-name: Tests (gpu)
on:
pull_request:
branches: [ master, main, gpu ]
workflow_dispatch:
env:
CARGO_TERM_COLOR: always
@@ -21,31 +22,15 @@ jobs:
- '3.13.x'
os:
- ubuntu-latest
- macos-latest
- windows-latest
exclude:
# Exclude 3.103.12 for macOS and Windows
- os: macos-latest
python-version: '3.10.x'
- os: macos-latest
python-version: '3.11.x'
- os: macos-latest
python-version: '3.12.x'
- os: windows-latest
python-version: '3.10.x'
- os: windows-latest
python-version: '3.11.x'
- os: windows-latest
python-version: '3.12.x'
runs-on: ${{ matrix.os }}
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
+2 -2
View File
@@ -14,10 +14,10 @@ jobs:
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 }}
+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
+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()
+1
View File
@@ -49,4 +49,5 @@ class SparseModelDescription(BaseModelDescription):
class PoolingType(str, Enum):
CLS = "CLS"
MEAN = "MEAN"
LAST_TOKEN = "LAST_TOKEN"
DISABLED = "DISABLED"
+106 -41
View File
@@ -1,10 +1,13 @@
import os
import time
import gzip
import json
import shutil
import tarfile
import tempfile
import contextlib
from copy import deepcopy
from pathlib import Path
from pathlib import Path, PureWindowsPath
from typing import Any, TypeVar, Generic
import requests
@@ -21,6 +24,8 @@ from fastembed.common.model_description import BaseModelDescription
T = TypeVar("T", bound=BaseModelDescription)
_DOWNLOAD_CHUNK_SIZE = 256 * 1024
class ModelManagement(Generic[T]):
METADATA_FILE = "files_metadata.json"
@@ -97,9 +102,7 @@ class ModelManagement(Generic[T]):
str: The path to the downloaded file.
"""
if os.path.exists(output_path):
return output_path
response = requests.get(url, stream=True)
response = requests.get(url, stream=True, timeout=(10, 120))
# Handle HTTP errors
if response.status_code == 403:
@@ -107,6 +110,8 @@ class ModelManagement(Generic[T]):
"Authentication Error: You do not have permission to access this resource. "
"Please check your credentials."
)
# Otherwise an error page gets written out as though it were the archive.
response.raise_for_status()
# Get the total size of the file
total_size_in_bytes = int(response.headers.get("content-length", 0))
@@ -124,7 +129,7 @@ class ModelManagement(Generic[T]):
disable=not show_progress,
) as progress_bar:
with open(output_path, "wb") as file:
for chunk in response.iter_content(chunk_size=1024):
for chunk in response.iter_content(chunk_size=_DOWNLOAD_CHUNK_SIZE):
if chunk: # Filter out keep-alive new chunks
progress_bar.update(len(chunk))
file.write(chunk)
@@ -286,12 +291,18 @@ class ModelManagement(Generic[T]):
"""
Decompresses a .tar.gz file to a cache directory.
Nothing is deleted on failure, since `cache_dir` may hold more than this archive.
Cleaning up a partial extraction is the caller's job.
Args:
targz_path (str): Path to the .tar.gz file.
cache_dir (str): Path to the cache directory.
Returns:
cache_dir (str): Path to the cache directory.
Raises:
ValueError: If the archive is missing, corrupt, or holds an unsafe member.
"""
# Check if targz_path exists and is a file
if not os.path.isfile(targz_path):
@@ -304,20 +315,50 @@ class ModelManagement(Generic[T]):
try:
# Open the tar.gz file
with tarfile.open(targz_path, "r:gz") as tar:
# Extract all files into the cache directory
tar.extractall(
path=cache_dir,
)
except tarfile.TarError as e:
# If any error occurs while opening or extracting the tar.gz file,
# delete the cache directory (if it was created in this function)
# and raise the error again
if "tmp" in cache_dir:
shutil.rmtree(cache_dir)
raise ValueError(f"An error occurred while decompressing {targz_path}: {e}")
if hasattr(tarfile, "data_filter"):
tar.extractall(path=cache_dir, filter="data")
else:
# No PEP 706 filter before 3.10.12, so vet the members by hand.
members = tar.getmembers()
for member in members:
cls._validate_tar_member(member)
tar.extractall(path=cache_dir, members=members)
# tarfile stops at the end-of-archive marker, short of the gzip trailer, so
# the CRC is only checked if the rest of the stream is read.
while tar.fileobj.read(1 << 20):
pass
except (tarfile.TarError, ValueError, EOFError, gzip.BadGzipFile) as e:
# gzip raises EOFError for a truncated stream and BadGzipFile for a corrupted one.
raise ValueError(f"An error occurred while decompressing {targz_path}: {e}") from e
return cache_dir
@staticmethod
def _is_unsafe_tar_path(path: str) -> bool:
"""Checks whether a tar member name or link target may escape the extraction dir.
Lexical on purpose: resolving against the extraction directory is unsound before
extraction, since `link/../escape` only escapes once an earlier member has been
written as a symlink. Any `..` component is therefore rejected outright.
"""
# PureWindowsPath splits on both separators, so `root` covers POSIX "/evil" as
# well as "\\evil", which escapes on Windows without being absolute.
windows_path = PureWindowsPath(path)
return bool(windows_path.drive or windows_path.root) or ".." in windows_path.parts
@classmethod
def _validate_tar_member(cls, member: tarfile.TarInfo) -> None:
"""Raises ValueError if a member could write outside the extraction directory."""
if cls._is_unsafe_tar_path(member.name):
raise ValueError(f"Unsafe tar member path: {member.name}")
if member.issym() or member.islnk():
if cls._is_unsafe_tar_path(member.linkname):
raise ValueError(f"Unsafe tar link target: {member.name} -> {member.linkname}")
elif not (member.isfile() or member.isdir()):
# Devices, fifos and the like have no place in a model archive.
raise ValueError(f"Unsupported tar member type: {member.name}")
@classmethod
def retrieve_model_gcs(
cls,
@@ -329,36 +370,13 @@ class ModelManagement(Generic[T]):
) -> Path:
fast_model_name = f"{'fast-' if deprecated_tar_struct else ''}{model_name.split('/')[-1]}"
cache_tmp_dir = Path(cache_dir) / "tmp"
model_tmp_dir = cache_tmp_dir / fast_model_name
model_dir = Path(cache_dir) / fast_model_name
# check if the model_dir and the model files are both present for macOS
if model_dir.exists() and len(list(model_dir.glob("*"))) > 0:
return model_dir
if model_tmp_dir.exists():
shutil.rmtree(model_tmp_dir)
cache_tmp_dir.mkdir(parents=True, exist_ok=True)
model_tar_gz = Path(cache_dir) / f"{fast_model_name}.tar.gz"
if model_tar_gz.exists():
model_tar_gz.unlink()
if not local_files_only:
cls.download_file_from_gcs(
source_url,
output_path=str(model_tar_gz),
)
cls.decompress_to_cache(targz_path=str(model_tar_gz), cache_dir=str(cache_tmp_dir))
assert model_tmp_dir.exists(), f"Could not find {model_tmp_dir} in {cache_tmp_dir}"
model_tar_gz.unlink()
# Rename from tmp to final name is atomic
model_tmp_dir.rename(model_dir)
else:
if local_files_only:
logger.error(
f"Could not find the model tar.gz file at {model_dir} and local_files_only=True."
)
@@ -366,6 +384,45 @@ class ModelManagement(Generic[T]):
f"Could not find the model tar.gz file at {model_dir} and local_files_only=True."
)
if cache_tmp_dir.is_symlink():
raise ValueError(
f"{cache_tmp_dir} is a symlink, refusing to stage downloads through it"
)
cache_tmp_dir.mkdir(parents=True, exist_ok=True)
# The archive and everything extracted from it go in a directory of this attempt's own,
# so removing it undoes the attempt without touching any other download of the model.
staging_dir = Path(tempfile.mkdtemp(dir=cache_tmp_dir, prefix=f"{fast_model_name}-"))
try:
model_tar_gz = staging_dir / f"{fast_model_name}.tar.gz"
cls.download_file_from_gcs(
source_url,
output_path=str(model_tar_gz),
)
cls.decompress_to_cache(targz_path=str(model_tar_gz), cache_dir=str(staging_dir))
model_tmp_dir = staging_dir / fast_model_name
if not model_tmp_dir.is_dir() or model_tmp_dir.is_symlink():
raise ValueError(
f"The archive from {source_url} has no {fast_model_name} directory"
)
# Replace a stale empty model_dir, which Windows will not rename onto. rmdir leaves
# anything else alone, including one another download has just filled.
with contextlib.suppress(OSError):
model_dir.rmdir()
try:
# Rename from the staging dir to the final name is atomic
model_tmp_dir.rename(model_dir)
except OSError:
# Another download of the same model finished first, so keep its copy.
if not (model_dir.is_dir() and any(model_dir.iterdir())):
raise
finally:
shutil.rmtree(staging_dir, ignore_errors=True)
return model_dir
@classmethod
@@ -395,6 +452,10 @@ class ModelManagement(Generic[T]):
Path: The path to the downloaded model directory.
"""
local_files_only = kwargs.get("local_files_only", False)
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)
@@ -409,7 +470,7 @@ class ModelManagement(Generic[T]):
try:
cache_kwargs = deepcopy(kwargs)
cache_kwargs["local_files_only"] = True
return Path(
resolved_path = Path(
cls.download_files_from_huggingface(
hf_source,
cache_dir=cache_dir,
@@ -417,6 +478,10 @@ class ModelManagement(Generic[T]):
**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:
+12
View File
@@ -31,6 +31,18 @@ class OnnxModel(Generic[T]):
def _get_worker_class(cls) -> Type["EmbeddingWorker[T]"]:
raise NotImplementedError("Subclasses must implement this method")
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.
+96 -29
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,44 +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)
if not tokenizer.padding:
tokenizer.enable_padding(
pad_id=config.get("pad_token_id", 0), pad_token=tokenizer_config["pad_token"]
)
for token in tokens_map.values():
# 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
+10
View File
@@ -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))
+7 -1
View File
@@ -5,11 +5,17 @@ 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]]:
+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,
)
-1
View File
@@ -63,7 +63,6 @@ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[NumpyArray]):
def __init__(
self,
model_name: str,
cache_dir: str | None = None,
threads: int | None = None,
providers: Sequence[OnnxProvider] | None = None,
+3 -1
View File
@@ -19,6 +19,8 @@ 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")
@@ -87,7 +89,7 @@ class OnnxImageModel(OnnxModel[T]):
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)
+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,
)
+18 -5
View File
@@ -65,7 +65,13 @@ def normalize(
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)
@@ -78,7 +84,9 @@ def normalize(
f"{len(mean_list)}"
)
mean_arr = np.array(mean_list, 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_list = std if isinstance(std, list) else [std] * num_channels
if len(std_list) != num_channels:
@@ -86,9 +94,9 @@ def normalize(
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_list, dtype=np.float32)
std_arr = np.array(std_list, dtype=np.float32).reshape(-1, 1, 1)
image_upd = ((image.T - mean_arr) / std_arr).T
image_upd = (image - mean_arr) / std_arr
return image_upd
@@ -98,7 +106,12 @@ def resize(
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)
@@ -1,9 +1,10 @@
from typing import Sequence, Any
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):
@@ -39,9 +40,39 @@ class CustomTextCrossEncoder(OnnxTextCrossEncoder):
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,
)
@@ -128,6 +128,7 @@ class OnnxCrossEncoderModel(OnnxModel[float]):
"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:
+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,
)
+8 -1
View File
@@ -5,6 +5,7 @@ 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,
@@ -16,7 +17,13 @@ from fastembed.common.model_description import SparseModelDescription
class SparseTextEmbedding(SparseTextEmbeddingBase):
EMBEDDINGS_REGISTRY: list[Type[SparseTextEmbeddingBase]] = [SpladePP, Bm42, Bm25, MiniCOIL]
EMBEDDINGS_REGISTRY: list[Type[SparseTextEmbeddingBase]] = [
SpladePP,
Bm42,
Bm25,
MiniCOIL,
IfSplade,
]
@classmethod
def list_supported_models(cls) -> list[dict[str, Any]]:
@@ -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,
)
+52 -5
View File
@@ -1,4 +1,4 @@
from typing import Sequence, Any, Iterable
from typing import Sequence, Any, Iterable, Type
from dataclasses import dataclass
import numpy as np
@@ -11,8 +11,9 @@ from fastembed.common.model_description import (
)
from fastembed.common.onnx_model import OnnxOutputContext
from fastembed.common.types import NumpyArray, Device
from fastembed.common.utils import normalize, mean_pooling
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)
@@ -50,13 +51,24 @@ 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
@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]:
@@ -73,12 +85,18 @@ 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}, {PoolingType.DISABLED}."
f"Supported types are: {PoolingType.CLS}, {PoolingType.MEAN}, "
f"{PoolingType.LAST_TOKEN}, {PoolingType.DISABLED}."
)
def _normalize(self, embeddings: NumpyArray) -> NumpyArray:
@@ -95,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,
)
+40 -4
View File
@@ -1,11 +1,11 @@
from typing import Any, Iterable, Sequence, Type
from fastembed.common.types import NumpyArray, OnnxProvider, Device
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(
@@ -34,7 +34,7 @@ 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,
),
@@ -77,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(
@@ -180,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",
),
]
+10 -2
View File
@@ -92,14 +92,21 @@ 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,
@@ -144,6 +151,7 @@ class OnnxTextModel(OnnxModel[T]):
"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:
@@ -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",
+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,
)
+8 -16
View File
@@ -8,7 +8,10 @@ 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,
]
@@ -88,23 +94,9 @@ class TextEmbedding(TextEmbeddingBase):
**kwargs: Any,
):
super().__init__(model_name, cache_dir, threads, **kwargs)
if model_name.lower() == "nomic-ai/nomic-embed-text-v1.5-Q".lower():
if model_name.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.lower() in {
"sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2".lower(),
"thenlper/gte-large".lower(),
"intfloat/multilingual-e5-large".lower(),
"sentence-transformers/paraphrase-multilingual-mpnet-base-v2".lower(),
}:
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,
)
Generated
+1883 -1834
View File
File diff suppressed because it is too large Load Diff
+22 -19
View File
@@ -1,9 +1,9 @@
[tool.poetry]
name = "fastembed"
version = "0.7.4"
name = "fastembed-gpu"
version = "0.8.1"
description = "Fast, light, accurate library built for retrieval embedding generation"
authors = ["Qdrant Team <info@qdrant.tech>", "NirantK <nirant.bits@gmail.com>"]
license = "Apache License"
license = "Apache-2.0"
readme = "README.md"
packages = [{include = "fastembed"}]
homepage = "https://github.com/qdrant/fastembed"
@@ -13,15 +13,17 @@ keywords = ["vector", "embedding", "neural", "search", "qdrant", "sentence-trans
[tool.poetry.dependencies]
python = ">=3.10.0"
numpy = [
{ version = ">=1.21,<2.3.0", python = ">=3.10,<3.11" },
{ version = ">=1.21", python = ">=3.11,<3.12" },
{ version = ">=1.26", python = ">=3.12,<3.13" },
{ version = ">=2.1.0", python = ">=3.13,<3.14" },
{ version = ">=1.21,<2.3.0", python = "3.10" },
{ version = ">=1.21", python = "3.11" },
{ version = ">=1.26", python = "3.12" },
{ version = ">=2.1.0", python = "3.13" },
{ version = ">=2.3.0", python = ">=3.14" },
]
onnxruntime = [
{ version = ">1.20.0", python = ">=3.13" },
{ version = ">=1.17.0,!=1.20.0", python = ">=3.10,<3.13" },
onnxruntime-gpu = [
{ version = ">=1.17.0,!=1.20.0,<1.24", python = "3.10" },
{ version = ">=1.17.0,!=1.20.0,!=1.24.0,!=1.24.1", python = ">=3.11,<3.13" },
{ version = ">1.21.0,!=1.24.0,!=1.24.1", python = "3.13" },
{ version = ">=1.24.2", python = ">=3.14" },
]
tqdm = "^4.66"
requests = "^2.31"
@@ -29,35 +31,36 @@ tokenizers = ">=0.15,<1.0"
huggingface-hub = ">=0.20,<2.0"
loguru = "^0.7.2"
pillow = [
{ version = ">=10.3.0,<11.0", python = "<3.10" },
{ version = ">=10.3.0,<12.0", python = ">=3.10,<3.13" },
{ version = ">=11.0.0,<12.0", python = ">=3.13" },
{ version = ">=10.3.0,<13.0", python = ">=3.10,<3.13" },
{ version = ">=11.0.0,<13.0", python = "3.13" },
{ version = ">=12.0.0,<13.0", python = ">=3.14" },
]
mmh3 = ">=4.1.0,<6.0.0"
py-rust-stemmers = "^0.1.0"
[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"
pre-commit = ">=3.6.2,<5.0.0"
onnx = [
{ version = ">=1.15.0", python = "<3.13" },
{ version = ">=1.18.0", python = ">=3.13" },
{ version = ">=1.15.0", python = ">=3.10,<3.13" },
{ version = ">=1.18.0", python = "3.13" },
{ version = ">=1.20.0", python = ">=3.14" },
]
[tool.poetry.group.docs.dependencies]
mkdocs-material = "^9.5.10"
mkdocstrings = "^0.24.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"]
+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]])
+74 -12
View File
@@ -10,7 +10,7 @@ from fastembed.common.model_description import (
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
@@ -21,9 +21,11 @@ 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 = []
@@ -74,8 +76,29 @@ def test_text_custom_model():
if is_ci:
delete_model_cache(model.model._model_dir)
CustomTextEmbedding.SUPPORTED_MODELS.clear()
CustomTextEmbedding.POSTPROCESSING_MAPPING.clear()
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():
@@ -114,7 +137,26 @@ def test_cross_encoder_custom_model():
if is_ci:
delete_model_cache(model.model._model_dir)
CustomTextCrossEncoder.SUPPORTED_MODELS.clear()
def test_cross_encoder_custom_model_parallel_processing():
is_ci = os.getenv("CI")
custom_model_name = "Xenova/ms-marco-MiniLM-L-4-v2"
TextCrossEncoder.add_custom_model(
custom_model_name,
model_file="onnx/model.onnx",
sources=ModelSource(hf=custom_model_name),
size_in_gb=0.08,
)
model = TextCrossEncoder(custom_model_name)
pairs = [("What is AI?", "Artificial intelligence is ...")] * 50
scores = np.stack(list(model.rerank_pairs(pairs, batch_size=10, parallel=2)), axis=0)
assert scores.shape == (len(pairs),)
if is_ci:
delete_model_cache(model.model._model_dir)
def test_mock_add_custom_models():
@@ -136,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,
}
@@ -147,12 +191,19 @@ def test_mock_add_custom_models():
f"{PoolingType.MEAN.lower()}": mean_pooling(dummy_token_embedding, dummy_attention_mask),
f"{PoolingType.CLS.lower()}-normalized": normalize(dummy_token_embedding[:, 0]),
f"{PoolingType.CLS.lower()}": dummy_token_embedding[:, 0],
f"{PoolingType.LAST_TOKEN.lower()}-normalized": normalize(
last_token_pooling(dummy_token_embedding, dummy_attention_mask)
),
f"{PoolingType.LAST_TOKEN.lower()}": last_token_pooling(
dummy_token_embedding, dummy_attention_mask
),
f"{PoolingType.DISABLED.lower()}-normalized": normalize(dummy_pooled_embedding),
f"{PoolingType.DISABLED.lower()}": dummy_pooled_embedding,
}
for pooling, normalization in itertools.product(
(PoolingType.MEAN, PoolingType.CLS, PoolingType.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(
@@ -175,8 +226,24 @@ def test_mock_add_custom_models():
)
assert np.allclose(post_processed_output, expected_output[model_name], atol=1e-3)
CustomTextEmbedding.SUPPORTED_MODELS.clear()
CustomTextEmbedding.POSTPROCESSING_MAPPING.clear()
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():
@@ -212,9 +279,6 @@ def test_do_not_add_existing_model():
size_in_gb=0.47,
)
CustomTextEmbedding.SUPPORTED_MODELS.clear()
CustomTextEmbedding.POSTPROCESSING_MAPPING.clear()
def test_do_not_add_existing_cross_encoder():
existing_base_model = "Xenova/ms-marco-MiniLM-L-6-v2"
@@ -239,5 +303,3 @@ def test_do_not_add_existing_cross_encoder():
sources=ModelSource(hf=custom_model_name),
size_in_gb=0.08,
)
CustomTextCrossEncoder.SUPPORTED_MODELS.clear()
+14
View File
@@ -1,4 +1,5 @@
import os
import platform
from contextlib import contextmanager
from io import BytesIO
@@ -25,6 +26,15 @@ 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",)
@@ -59,9 +69,13 @@ def model_cache():
@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():
# 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
+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)
+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
+69 -5
View File
@@ -8,6 +8,7 @@ from fastembed.sparse.bm25 import Bm25
from fastembed.sparse.sparse_text_embedding import SparseTextEmbedding
from tests.utils import delete_model_cache, should_test_model
CANONICAL_COLUMN_VALUES = {
"prithivida/Splade_PP_en_v1": {
"indices": [
@@ -58,9 +59,50 @@ CANONICAL_COLUMN_VALUES = {
-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": [
@@ -82,6 +124,7 @@ _MODELS_TO_CACHE = (
"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])
@@ -142,18 +185,23 @@ def test_single_embedding(model_cache) -> None:
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)
assert passage_result.indices.tolist() == expected_result["indices"]
for i, value in enumerate(passage_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]
assert query_result.indices.tolist() == expected_query_result["indices"]
for i, value in enumerate(query_result.values):
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]
@@ -263,6 +311,22 @@ def test_disable_stemmer_behavior(disable_stemmer: bool) -> None:
assert result == expected, f"Expected {expected}, but got {result}"
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")
+114
View File
@@ -5,6 +5,8 @@ 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, should_test_model
@@ -68,6 +70,51 @@ 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"]
@@ -119,6 +166,9 @@ def test_embedding(model_cache, model_name: str) -> None:
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)
@@ -129,6 +179,56 @@ def test_embedding(model_cache, model_name: str) -> None:
), 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:
@@ -217,3 +317,17 @@ def test_token_count(model_cache, model_name) -> None:
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)