salim0986/graph-bug-ai
0
1from sentence_transformers import CrossEncoder2from typing import List, Dict, Any3from .logger import setup_logger4 5logger = setup_logger(__name__)6 7class Reranker:8 """9 Reranks semantic search results using a cross-encoder model (e.g. BAAI/bge-reranker-base).10 Cross-encoders score pairs of (query, document) directly and are more accurate 11 than bi-encoders, but slower.12 """13 def __init__(self, model_name: str = "BAAI/bge-reranker-base"):14 logger.info(f"Loading reranker model: {model_name}")15 self.model = CrossEncoder(model_name)16 17 def rerank(self, query: str, documents: List[Dict[str, Any]], top_k: int = 5) -> List[Dict[str, Any]]:18 """19 Rerank a list of documents based on a query.20 Assumes documents have a 'code' or 'raw_code' field to evaluate.21 """22 if not documents:23 return []24 25 # Extract the text to score against the query26 # Usually it's 'raw_code' or 'code' from the VectorBuilder output27 pairs = []28 for doc in documents:29 # Handle both dictionary keys and Qdrant payload objects30 if isinstance(doc, dict):31 text = doc.get("raw_code") or doc.get("code") or doc.get("text", "")32 else:33 # If it's a Qdrant point object34 text = getattr(doc, "payload", {}).get("raw_code", "")35 36 pairs.append((query, str(text)))37 38 # Score pairs39 try:40 scores = self.model.predict(pairs)41 42 # Attach scores to documents43 scored_docs = []44 for doc, score in zip(documents, scores):45 doc_copy = doc.copy() if isinstance(doc, dict) else doc46 if isinstance(doc_copy, dict):47 doc_copy["rerank_score"] = float(score)48 scored_docs.append((score, doc_copy))49 50 # Sort by score descending51 scored_docs.sort(key=lambda x: x[0], reverse=True)52 53 # Return top_k54 return [doc for _, doc in scored_docs[:top_k]]55 56 except Exception as e:57 logger.error(f"Reranking failed: {e}")58 # Fallback to original order if reranking fails59 return documents[:top_k]60 