Add jina, mxbai and sentence-transformers to download-models when used
This commit is contained in:
parent
7ff014684a
commit
36a675e1bb
3 changed files with 53 additions and 3 deletions
|
|
@ -1,6 +1,8 @@
|
||||||
# Changelog
|
# Changelog
|
||||||
## [Unreleased]
|
## [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
|
- **Reranker Factory**: Removed unreliable `id(config)`-based caching from `get_reranker()`; factory now always instantiates fresh
|
||||||
|
|
||||||
## [0.26.7] - 2026-01-20
|
## [0.26.7] - 2026-01-20
|
||||||
|
|
|
||||||
|
|
@ -1656,9 +1656,11 @@ class HaikuRAG:
|
||||||
"""Download required models, yielding progress events.
|
"""Download required models, yielding progress events.
|
||||||
|
|
||||||
Yields DownloadProgress events for:
|
Yields DownloadProgress events for:
|
||||||
- Docling models (status="docling_start", "docling_done")
|
- Docling models
|
||||||
- HuggingFace tokenizer (status="tokenizer_start", "tokenizer_done")
|
- HuggingFace tokenizer
|
||||||
- Ollama models (status="pulling", "downloading", "done", or other Ollama statuses)
|
- Sentence-transformers embedder (if configured)
|
||||||
|
- HuggingFace reranker models (mxbai, jina-local)
|
||||||
|
- Ollama models
|
||||||
"""
|
"""
|
||||||
# Docling models
|
# Docling models
|
||||||
try:
|
try:
|
||||||
|
|
@ -1678,6 +1680,51 @@ class HaikuRAG:
|
||||||
await asyncio.to_thread(AutoTokenizer.from_pretrained, tokenizer_name)
|
await asyncio.to_thread(AutoTokenizer.from_pretrained, tokenizer_name)
|
||||||
yield DownloadProgress(model=tokenizer_name, status="done")
|
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
|
# Collect Ollama models from config
|
||||||
required_models: set[str] = set()
|
required_models: set[str] = set()
|
||||||
if self._config.embeddings.model.provider == "ollama":
|
if self._config.embeddings.model.provider == "ollama":
|
||||||
|
|
|
||||||
|
|
@ -259,6 +259,7 @@ async def test_jina_reranker(monkeypatch):
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
async def test_jina_local_reranker():
|
async def test_jina_local_reranker():
|
||||||
try:
|
try:
|
||||||
from haiku.rag.reranking.jina_local import JinaLocalReranker
|
from haiku.rag.reranking.jina_local import JinaLocalReranker
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue