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":
|
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,
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -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]]
|
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue