Allow transformers 5.x in the mxbai extra
This commit is contained in:
parent
8f085101fc
commit
5e2a928013
5 changed files with 34 additions and 2 deletions
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
2
uv.lock
2
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" },
|
||||
|
|
|
|||
Loading…
Reference in a new issue