Jina uses AutoModel not AutoModelForSequenceClassification
This commit is contained in:
parent
ce95dc47a5
commit
a1abfd9666
2 changed files with 7 additions and 16 deletions
|
|
@ -1715,11 +1715,11 @@ class HaikuRAG:
|
|||
|
||||
elif provider == "jina-local":
|
||||
try:
|
||||
from transformers import AutoModelForSequenceClassification
|
||||
from transformers import AutoModel
|
||||
|
||||
yield DownloadProgress(model=model_name, status="start")
|
||||
await asyncio.to_thread(
|
||||
AutoModelForSequenceClassification.from_pretrained,
|
||||
AutoModel.from_pretrained,
|
||||
model_name,
|
||||
trust_remote_code=True,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
try:
|
||||
from transformers import (
|
||||
AutoModelForSequenceClassification, # pyright: ignore[reportMissingImports]
|
||||
AutoModel, # pyright: ignore[reportMissingImports]
|
||||
)
|
||||
except ImportError as e:
|
||||
raise ImportError(
|
||||
|
|
@ -21,9 +21,8 @@ class JinaLocalReranker(RerankerBase): # pragma: no cover
|
|||
|
||||
def __init__(self, model: str = "jinaai/jina-reranker-v3"):
|
||||
self._model = model
|
||||
self._reranker = AutoModelForSequenceClassification.from_pretrained(
|
||||
model, trust_remote_code=True
|
||||
)
|
||||
self._reranker = AutoModel.from_pretrained(model, trust_remote_code=True)
|
||||
self._reranker.eval()
|
||||
|
||||
async def rerank(
|
||||
self, query: str, chunks: list[Chunk], top_n: int = 10
|
||||
|
|
@ -32,15 +31,7 @@ class JinaLocalReranker(RerankerBase): # pragma: no cover
|
|||
return []
|
||||
|
||||
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
|
||||
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]]
|
||||
return [(chunks[r["index"]], float(r["relevance_score"])) for r in results]
|
||||
|
|
|
|||
Loading…
Reference in a new issue