Jina uses AutoModel not AutoModelForSequenceClassification

This commit is contained in:
Yiorgis Gozadinos 2026-01-21 11:26:23 +02:00
parent ce95dc47a5
commit a1abfd9666
No known key found for this signature in database
2 changed files with 7 additions and 16 deletions

View file

@ -1715,11 +1715,11 @@ class HaikuRAG:
elif provider == "jina-local": elif provider == "jina-local":
try: try:
from transformers import AutoModelForSequenceClassification from transformers import AutoModel
yield DownloadProgress(model=model_name, status="start") yield DownloadProgress(model=model_name, status="start")
await asyncio.to_thread( await asyncio.to_thread(
AutoModelForSequenceClassification.from_pretrained, AutoModel.from_pretrained,
model_name, model_name,
trust_remote_code=True, trust_remote_code=True,
) )

View file

@ -1,6 +1,6 @@
try: try:
from transformers import ( from transformers import (
AutoModelForSequenceClassification, # pyright: ignore[reportMissingImports] AutoModel, # pyright: ignore[reportMissingImports]
) )
except ImportError as e: except ImportError as e:
raise ImportError( raise ImportError(
@ -21,9 +21,8 @@ class JinaLocalReranker(RerankerBase): # pragma: no cover
def __init__(self, model: str = "jinaai/jina-reranker-v3"): def __init__(self, model: str = "jinaai/jina-reranker-v3"):
self._model = model self._model = model
self._reranker = AutoModelForSequenceClassification.from_pretrained( self._reranker = AutoModel.from_pretrained(model, trust_remote_code=True)
model, trust_remote_code=True self._reranker.eval()
)
async def rerank( async def rerank(
self, query: str, chunks: list[Chunk], top_n: int = 10 self, query: str, chunks: list[Chunk], top_n: int = 10
@ -32,15 +31,7 @@ class JinaLocalReranker(RerankerBase): # pragma: no cover
return [] return []
documents = [chunk.content for chunk in chunks] documents = [chunk.content for chunk in chunks]
sentence_pairs = [[query, doc] for doc in documents]
scores = self._reranker.compute_score(sentence_pairs) results = self._reranker.rerank(query, documents, top_n=top_n)
# Handle both single score and list of scores return [(chunks[r["index"]], float(r["relevance_score"])) for r in results]
if isinstance(scores, (int, float)):
scores = [scores]
scored_chunks = list(zip(chunks, scores, strict=False))
scored_chunks.sort(key=lambda x: x[1], reverse=True)
return [(chunk, float(score)) for chunk, score in scored_chunks[:top_n]]