61 lines
1.5 KiB
Python
61 lines
1.5 KiB
Python
# -*- coding:utf-8 -*-
|
|
"""BGE-Reranker-M3 lazy-loading wrapper."""
|
|
import os
|
|
import json
|
|
import torch
|
|
from FlagEmbedding import FlagReranker
|
|
|
|
_model = None
|
|
_lock = False
|
|
|
|
MODEL_PATH = "/data/ymq/models/BAAI/bge-reranker-v2-m3"
|
|
DEVICE = "cuda:0"
|
|
|
|
def get_model():
|
|
global _model, _lock
|
|
if _model is None and not _lock:
|
|
_lock = True
|
|
try:
|
|
_model = FlagReranker(MODEL_PATH, use_fp16=True, device=DEVICE)
|
|
except Exception as e:
|
|
_lock = False
|
|
raise RuntimeError(f"Failed to load BGE-Reranker: {e}")
|
|
return _model
|
|
|
|
|
|
def rerank(query: str, documents: list, top_k: int = 5) -> dict:
|
|
"""Rerank documents by relevance to query."""
|
|
model = get_model()
|
|
|
|
# Create pairs: [[query, doc1], [query, doc2], ...]
|
|
pairs = [[query, doc] for doc in documents]
|
|
|
|
# Compute scores
|
|
scores = model.compute_score(pairs, normalize=True)
|
|
|
|
if isinstance(scores, float):
|
|
scores = [scores]
|
|
|
|
# Pair with documents and sort
|
|
ranked = list(zip(scores, documents))
|
|
ranked.sort(key=lambda x: -x[0])
|
|
|
|
# Return top_k
|
|
results = [
|
|
{"doc": doc, "score": round(score, 6)}
|
|
for score, doc in ranked[:top_k]
|
|
]
|
|
|
|
return {"ranked_docs": results, "total": len(documents), "returned": len(results)}
|
|
|
|
|
|
def health_check():
|
|
"""Check model status."""
|
|
model = get_model()
|
|
return {
|
|
"model": "BAAI/bge-reranker-v2-m3",
|
|
"device": DEVICE,
|
|
"loaded": model is not None,
|
|
"model_path": MODEL_PATH
|
|
}
|