Cognitive-rag / 04_vectoredb / reranker.py
reranker.py
Raw
"""
Cross-Encoder Reranker
Uses ms-marco-MiniLM-L6-v2 to rerank retrieved passages by scoring
each (query, passage) pair with a cross-encoder.
"""

from __future__ import annotations

import logging
from typing import Dict, List, Optional

logger = logging.getLogger(__name__)
_cross_encoder = None

def _get_model(model_name: str = None):
    """Load the cross-encoder model"""
    global _cross_encoder

    if _cross_encoder is not None:
        return _cross_encoder

    try:
        from . import config as _cfg
    except ImportError:
        import config as _cfg

    model_name = model_name or _cfg.RERANKER_MODEL

    from sentence_transformers import CrossEncoder

    logger.info(f"Loading cross-encoder model: {model_name}")
    _cross_encoder = CrossEncoder(model_name)
    logger.info("Cross-encoder model loaded.")
    return _cross_encoder


def rerank(
    query: str,
    chunks: List[Dict],
    top_n: Optional[int] = None,
    model_name: Optional[str] = None,
) -> List[Dict]:
    """
    rerank chunks using the cross-encoder.

    args:
        query: The user query string.
        chunks: List of chunk dicts, each must have a 'text' field.
        top_n: Number of top passages to return. If None, uses config.RERANKER_TOP_N.
        model_name: Override model name (uses config default if None).
    returns:
        top N chunks sorted by cross-encoder score, each with
        'cross_encoder_score' added.
    """
    if not chunks:
        return []

    try:
        from . import config as _cfg
    except ImportError:
        import config as _cfg

    top_n = top_n if top_n is not None else _cfg.RERANKER_TOP_N

    model = _get_model(model_name)

    # build (query, passage) pairs
    pairs = [(query, chunk.get("text", "")) for chunk in chunks]

    # score all pairs
    scores = model.predict(pairs)

    # attach scores and sort descending
    scored_chunks = []
    for chunk, score in zip(chunks, scores):
        enriched = chunk.copy()
        enriched["cross_encoder_score"] = round(float(score), 6)
        scored_chunks.append(enriched)

    scored_chunks.sort(key=lambda x: x["cross_encoder_score"], reverse=True)

    return scored_chunks[:top_n]