Allow transformers 5.x in the mxbai extra

This commit is contained in:
Yiorgis Gozadinos 2026-07-14 10:13:45 +03:00
parent 8f085101fc
commit 5e2a928013
No known key found for this signature in database
5 changed files with 34 additions and 2 deletions

View file

@ -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.

View file

@ -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

View file

@ -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"]

View file

@ -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():

View file

@ -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" },