Add jina, mxbai and sentence-transformers to download-models when used

This commit is contained in:
Yiorgis Gozadinos 2026-01-21 10:40:22 +02:00
parent 7ff014684a
commit 36a675e1bb
No known key found for this signature in database
3 changed files with 53 additions and 3 deletions

View file

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

View file

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

View file

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