diff --git a/docs/tuning.md b/docs/tuning.md index 507aeb9b..182990c8 100644 --- a/docs/tuning.md +++ b/docs/tuning.md @@ -164,7 +164,7 @@ Reranking retrieves more candidates than needed, then uses a cross-encoder to re reranking: model: provider: mxbai # or cohere, zeroentropy, vllm - name: mxbai-rerank-base-v1 + name: mixedbread-ai/mxbai-rerank-base-v2 ``` **When to use reranking:** @@ -291,7 +291,7 @@ search: reranking: model: provider: mxbai - name: mxbai-rerank-base-v1 + name: mixedbread-ai/mxbai-rerank-base-v2 ``` ### Long-Form Content (Articles, Reports) diff --git a/haiku_rag_slim/haiku/rag/reranking/mxbai.py b/haiku_rag_slim/haiku/rag/reranking/mxbai.py index e36f4f1a..6e42f1a4 100644 --- a/haiku_rag_slim/haiku/rag/reranking/mxbai.py +++ b/haiku_rag_slim/haiku/rag/reranking/mxbai.py @@ -10,7 +10,7 @@ class MxBAIReranker(RerankerBase): model_name = ( Config.reranking.model.name if Config.reranking.model - else "mxbai-rerank-base-v2" + else "mixedbread-ai/mxbai-rerank-base-v2" ) self._client = MxbaiRerankV2(model_name, disable_transformers_warnings=True) diff --git a/tests/test_reranker.py b/tests/test_reranker.py index e7a4e8fe..80a73b2d 100644 --- a/tests/test_reranker.py +++ b/tests/test_reranker.py @@ -31,7 +31,8 @@ async def test_reranker_base(): from haiku.rag.config import Config reranker = RerankerBase() - assert reranker._model == Config.reranking.model + expected_model = Config.reranking.model.name if Config.reranking.model else None + assert reranker._model == expected_model with pytest.raises(NotImplementedError): await reranker.rerank("query", [])