Spaces:
Runtime error
Runtime error
| """ | |
| Reranking module for improving search result relevance. | |
| Uses cross-encoder models to rerank retrieved chunks based on query relevance. | |
| """ | |
| from typing import List, Dict, Tuple | |
| from sentence_transformers import CrossEncoder | |
| class Reranker: | |
| """Reranks search results using cross-encoder models.""" | |
| def __init__(self, model_name: str = 'cross-encoder/ms-marco-MiniLM-L-6-v2'): | |
| """ | |
| Initialize reranker with cross-encoder model. | |
| Args: | |
| model_name: HuggingFace model name for cross-encoder | |
| """ | |
| self.model_name = model_name | |
| self.model = CrossEncoder(model_name) | |
| print(f"[Reranker] Loaded model: {model_name}") | |
| def rerank( | |
| self, | |
| query: str, | |
| chunks: List[Dict], | |
| top_k: int = 10 | |
| ) -> List[Dict]: | |
| """ | |
| Rerank chunks based on relevance to query. | |
| Args: | |
| query: Search query | |
| chunks: List of chunk dictionaries with 'content' key | |
| top_k: Number of top results to return after reranking | |
| Returns: | |
| Reranked list of chunks (top_k most relevant) | |
| """ | |
| if not chunks: | |
| return [] | |
| # Prepare query-document pairs | |
| pairs = [(query, chunk['content']) for chunk in chunks] | |
| # Get relevance scores | |
| scores = self.model.predict(pairs) | |
| # Combine chunks with scores and sort | |
| chunks_with_scores = [ | |
| {**chunk, 'rerank_score': float(score)} | |
| for chunk, score in zip(chunks, scores) | |
| ] | |
| # Sort by rerank score (highest first) | |
| reranked = sorted( | |
| chunks_with_scores, | |
| key=lambda x: x['rerank_score'], | |
| reverse=True | |
| ) | |
| # Return top_k | |
| return reranked[:top_k] | |
| def rerank_with_scores( | |
| self, | |
| query: str, | |
| chunks: List[Dict] | |
| ) -> List[Tuple[Dict, float]]: | |
| """ | |
| Rerank and return chunks with their relevance scores. | |
| Args: | |
| query: Search query | |
| chunks: List of chunk dictionaries | |
| Returns: | |
| List of (chunk, score) tuples sorted by relevance | |
| """ | |
| if not chunks: | |
| return [] | |
| pairs = [(query, chunk['content']) for chunk in chunks] | |
| scores = self.model.predict(pairs) | |
| results = list(zip(chunks, scores)) | |
| results.sort(key=lambda x: x[1], reverse=True) | |
| return results | |
| # Global reranker instance | |
| reranker = Reranker() | |