From 5e2a9280133815dcd14e8285cfa2d8aa782edb82 Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Tue, 14 Jul 2026 10:13:45 +0300 Subject: [PATCH] Allow transformers 5.x in the mxbai extra --- CHANGELOG.md | 4 ++++ haiku_rag_slim/haiku/rag/reranking/mxbai.py | 17 +++++++++++++++++ haiku_rag_slim/pyproject.toml | 2 +- tests/test_reranker.py | 11 +++++++++++ uv.lock | 2 +- 5 files changed, 34 insertions(+), 2 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 61c216b0..6615dcb0 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,6 +1,10 @@ # Changelog ## [Unreleased] +### Changed + +- `mxbai` extra allows `transformers` 5.x (`<6.0.0`); `MxBAIReranker` supplies the `tokenizer.prepare_for_model` variant `mxbai-rerank` needs when the tokenizer lacks it. + ### Fixed - Context expansion no longer drops a retrieved result from a merged group when clipping to `search.max_context_chars`; groups whose clip would evict a constituent's evidence are returned as separate expanded results. diff --git a/haiku_rag_slim/haiku/rag/reranking/mxbai.py b/haiku_rag_slim/haiku/rag/reranking/mxbai.py index c53713a8..03e99078 100644 --- a/haiku_rag_slim/haiku/rag/reranking/mxbai.py +++ b/haiku_rag_slim/haiku/rag/reranking/mxbai.py @@ -16,6 +16,21 @@ from haiku.rag.store.models.chunk import Chunk tqdm.tqdm.set_lock(threading.RLock()) +def _prepare_for_model( + ids: list[int], + pair_ids: list[int] | None = None, + max_length: int | None = None, + **_, +) -> dict[str, list[int]]: + # transformers 5.x removed tokenizer.prepare_for_model, which mxbai-rerank + # calls with add_special_tokens=False and truncation="only_second"; for that + # call pattern it reduces to truncating the pair and concatenating. + pair_ids = pair_ids or [] + if max_length is not None: + pair_ids = pair_ids[: max(0, max_length - len(ids))] + return {"input_ids": ids + pair_ids} + + class MxBAIReranker(RerankerBase): def __init__(self): model_name = ( @@ -24,6 +39,8 @@ class MxBAIReranker(RerankerBase): else "mixedbread-ai/mxbai-rerank-base-v2" ) self._client = MxbaiRerankV2(model_name, disable_transformers_warnings=True) + if not hasattr(self._client.tokenizer, "prepare_for_model"): + self._client.tokenizer.prepare_for_model = _prepare_for_model async def _rerank( self, query: str, chunks: list[Chunk], top_n: int = 10 diff --git a/haiku_rag_slim/pyproject.toml b/haiku_rag_slim/pyproject.toml index 75c085d6..e5ce96d5 100644 --- a/haiku_rag_slim/pyproject.toml +++ b/haiku_rag_slim/pyproject.toml @@ -51,7 +51,7 @@ s3 = ["obstore>=0.9,<0.10"] # Embedding providers voyageai = ["pydantic-ai-slim[voyageai]"] # Rerankers -mxbai = ["mxbai-rerank>=0.1.6", "transformers>=4.49.0,<5.0.0"] +mxbai = ["mxbai-rerank>=0.1.6", "transformers>=4.49.0,<6.0.0"] cohere = ["cohere>=5.21.1"] zeroentropy = ["zeroentropy>=0.1.0a11"] jina = ["transformers>=4.40.0", "torch>=2.0.0"] diff --git a/tests/test_reranker.py b/tests/test_reranker.py index 30dc6171..2fcfc9ac 100644 --- a/tests/test_reranker.py +++ b/tests/test_reranker.py @@ -76,6 +76,17 @@ async def test_mxbai_reranker(): pytest.skip("MxBAI package not installed") +def test_mxbai_prepare_for_model_shim(): + pytest.importorskip("mxbai_rerank") + from haiku.rag.reranking.mxbai import _prepare_for_model + + assert _prepare_for_model([1, 2], [3, 4]) == {"input_ids": [1, 2, 3, 4]} + assert _prepare_for_model([1, 2], [3, 4, 5], max_length=4) == { + "input_ids": [1, 2, 3, 4] + } + assert _prepare_for_model([1, 2], [3, 4], max_length=2) == {"input_ids": [1, 2]} + + @pytest.mark.asyncio @pytest.mark.vcr() async def test_cohere_reranker(): diff --git a/uv.lock b/uv.lock index b3b43f2e..22c3a97d 100644 --- a/uv.lock +++ b/uv.lock @@ -1793,7 +1793,7 @@ requires-dist = [ { name = "textual-image", specifier = ">=0.8.5" }, { name = "torch", marker = "extra == 'jina'", specifier = ">=2.0.0" }, { name = "transformers", marker = "extra == 'jina'", specifier = ">=4.40.0" }, - { name = "transformers", marker = "extra == 'mxbai'", specifier = ">=4.49.0,<5.0.0" }, + { name = "transformers", marker = "extra == 'mxbai'", specifier = ">=4.49.0,<6.0.0" }, { name = "tree-sitter", marker = "extra == 'tui'", specifier = ">=0.25.2" }, { name = "tree-sitter-json", marker = "extra == 'tui'", specifier = ">=0.24.8" }, { name = "typer", specifier = ">=0.21.0,<0.22.0" },