From 36a675e1bb802b61c0fe9b6854ad6b755bbade5d Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Wed, 21 Jan 2026 10:40:22 +0200 Subject: [PATCH] Add jina, mxbai and sentence-transformers to download-models when used --- CHANGELOG.md | 2 ++ haiku_rag_slim/haiku/rag/client.py | 53 ++++++++++++++++++++++++++++-- tests/test_reranker.py | 1 + 3 files changed, 53 insertions(+), 3 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 3b2f862b..53441028 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,6 +1,8 @@ # Changelog ## [Unreleased] +- **Jina Reranker v3**: Added support for Jina reranking with API mode (`provider: jina`) and local inference (`provider: jina-local`, requires `[jina]` extra) +- **Model Downloads**: `download-models` now pre-downloads HuggingFace models for `sentence-transformers`, `mxbai`, and `jina-local` - **Reranker Factory**: Removed unreliable `id(config)`-based caching from `get_reranker()`; factory now always instantiates fresh ## [0.26.7] - 2026-01-20 diff --git a/haiku_rag_slim/haiku/rag/client.py b/haiku_rag_slim/haiku/rag/client.py index 9f1460b1..056446fe 100644 --- a/haiku_rag_slim/haiku/rag/client.py +++ b/haiku_rag_slim/haiku/rag/client.py @@ -1656,9 +1656,11 @@ class HaikuRAG: """Download required models, yielding progress events. Yields DownloadProgress events for: - - Docling models (status="docling_start", "docling_done") - - HuggingFace tokenizer (status="tokenizer_start", "tokenizer_done") - - Ollama models (status="pulling", "downloading", "done", or other Ollama statuses) + - Docling models + - HuggingFace tokenizer + - Sentence-transformers embedder (if configured) + - HuggingFace reranker models (mxbai, jina-local) + - Ollama models """ # Docling models try: @@ -1678,6 +1680,51 @@ class HaikuRAG: await asyncio.to_thread(AutoTokenizer.from_pretrained, tokenizer_name) yield DownloadProgress(model=tokenizer_name, status="done") + # Sentence-transformers embedder + if self._config.embeddings.model.provider == "sentence-transformers": + try: + from sentence_transformers import ( # type: ignore[import-not-found] + SentenceTransformer, + ) + + model_name = self._config.embeddings.model.name + yield DownloadProgress(model=model_name, status="start") + await asyncio.to_thread(SentenceTransformer, model_name) + yield DownloadProgress(model=model_name, status="done") + except ImportError: + pass + + # HuggingFace reranker models + if self._config.reranking.model: + provider = self._config.reranking.model.provider + model_name = self._config.reranking.model.name + + if provider == "mxbai": + try: + from mxbai_rerank import MxbaiRerankV2 + + yield DownloadProgress(model=model_name, status="start") + await asyncio.to_thread( + MxbaiRerankV2, model_name, disable_transformers_warnings=True + ) + yield DownloadProgress(model=model_name, status="done") + except ImportError: + pass + + elif provider == "jina-local": + try: + from transformers import AutoModelForSequenceClassification + + yield DownloadProgress(model=model_name, status="start") + await asyncio.to_thread( + AutoModelForSequenceClassification.from_pretrained, + model_name, + trust_remote_code=True, + ) + yield DownloadProgress(model=model_name, status="done") + except ImportError: + pass + # Collect Ollama models from config required_models: set[str] = set() if self._config.embeddings.model.provider == "ollama": diff --git a/tests/test_reranker.py b/tests/test_reranker.py index d79cfcc4..74f19b6f 100644 --- a/tests/test_reranker.py +++ b/tests/test_reranker.py @@ -259,6 +259,7 @@ async def test_jina_reranker(monkeypatch): @pytest.mark.asyncio +@pytest.mark.integration async def test_jina_local_reranker(): try: from haiku.rag.reranking.jina_local import JinaLocalReranker