bge-reranker/workers/bge_reranker.py

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
}