#!/usr/bin/env python3 """ rag_reranker.py — Reranking con cross-encoder (bge-reranker-v2-m3). Pipeline: 1. Hybrid search → top 20 candidatos 2. Cross-encoder reranker → top 5 finales Uso: python3 rag_reranker.py "cómo configurar certificados SSL" """ import argparse import os import sys import sqlite3 import math import numpy as np from numpy.linalg import norm DB_PATH = os.path.expanduser("~/.cerebro/rag_conocimiento.db") EMBEDDING_BYTES = 1024 * 4 MODELO_RERANKER = "BAAI/bge-reranker-v2-m3" def obtener_embedding(texto: str) -> np.ndarray | None: import requests try: resp = requests.post( "http://localhost:11434/api/embeddings", json={"model": "bge-m3", "prompt": texto}, timeout=15, ) resp.raise_for_status() return np.array(resp.json()["embedding"], dtype=np.float32) except Exception as e: print(f"[aviso] Error: {e}", file=sys.stderr) return None def hybrid_search(query: str, alpha: float = 0.4, top_k: int = 20) -> list[dict]: if not os.path.exists(DB_PATH): print(f"Error: BD no encontrada", file=sys.stderr) return [] query_emb = obtener_embedding(query) if query_emb is None: return [] query_emb = query_emb / (norm(query_emb) + 1e-10) conn = sqlite3.connect(DB_PATH) conn.row_factory = sqlite3.Row cursor = conn.execute( """SELECT c.id, c.content, e.vector, c.doc_id, bm25(chunks_fts) as fts_score FROM chunks_fts JOIN chunks c ON chunks_fts.rowid = c.id JOIN embeddings e ON e.chunk_id = c.id WHERE chunks_fts MATCH ? ORDER BY fts_score DESC LIMIT ?""", (query, top_k * 2), ) rows = cursor.fetchall() conn.close() results = [] for row in rows: emb_vec = np.frombuffer(row["vector"], dtype=np.float32).copy() emb_vec = emb_vec / (norm(emb_vec) + 1e-10) cos_sim = float(np.dot(query_emb, emb_vec)) fts_norm = 1.0 / (1.0 + math.exp(-row["fts_score"] / 10.0)) score_total = alpha * fts_norm + (1.0 - alpha) * cos_sim results.append({ "id": row["id"], "contenido": row["content"], "doc_id": row["doc_id"], "score": score_total, }) results.sort(key=lambda x: x["score"], reverse=True) return results[:top_k] class Reranker: def __init__(self, model_name: str = MODELO_RERANKER, use_fp16: bool = True): from sentence_transformers import CrossEncoder self.model = CrossEncoder(model_name, max_length=512, device="cpu") def rerank(self, query: str, candidates: list[dict], top_k: int = 5) -> list[dict]: if not candidates: return [] pairs = [(query, c["contenido"]) for c in candidates] scores = self.model.predict(pairs) for i, score in enumerate(scores): candidates[i]["rerank_score"] = float(score) candidates.sort(key=lambda x: x["rerank_score"], reverse=True) return candidates[:top_k] def search_with_rerank(query: str, alpha: float = 0.4, top_k_hybrid: int = 20, top_k_final: int = 5): candidates = hybrid_search(query, alpha=alpha, top_k=top_k_hybrid) if not candidates: return [] reranker = Reranker() return reranker.rerank(query, candidates, top_k=top_k_final) def main(): parser = argparse.ArgumentParser(description="Reranking con cross-encoder") parser.add_argument("consulta", help="Texto de la consulta") parser.add_argument("-k", "--top-k", type=int, default=5, help="Resultados finales") parser.add_argument("-c", "--candidatos", type=int, default=20, help="Candidatos para reranking") args = parser.parse_args() results = search_with_rerank(query=args.consulta, top_k_hybrid=args.candidatos, top_k_final=args.top_k) if not results: print("Sin resultados.") return print(f"\nResultados rerankeados (top {len(results)}):") for i, r in enumerate(results, 1): preview = r["contenido"].replace("\n", " ")[:150] print(f" [{i:2d}] Rerank: {r.get(rerank_score, 0):.4f}") print(f" {preview}\n") if __name__ == "__main__": main()