File size: 2,637 Bytes
f3997d4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
"""
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()