AI_Menu_Search / core /reranker.py
Juhaha
HF Spaces 데모 배포 (Streamlit + Qdrant 임베디드, 색인 빌드타임 생성)
fbd1091
Raw
History Blame Contribute Delete
3.57 kB
"""
Cross-Encoder 기반 리랭킹 모듈 (로컬, API 키 불필요)
사용 방법:
from core.reranker import CrossEncoderReranker
reranker = CrossEncoderReranker.get_instance()
reranked = reranker.rerank(query, candidates, top_n=5)
동작 원리:
- RRF 하이브리드 검색 결과 Top-N 후보를 입력으로 받음
- (query, document) 쌍을 Cross-Encoder에 입력 → relevance 점수 계산
- sigmoid 변환으로 [0,1] 정규화 → similarity 필드 갱신
- 재정렬 후 Top-N 반환
모델: BAAI/bge-reranker-v2-m3
- bge-m3 임베더와 동일 BAAI 패밀리 (한국어 최적화)
- 100개 언어 지원, BEIR 다국어 벤치마크 SOTA
- 첫 실행 시 자동 다운로드 (~300MB)
- sentence-transformers CrossEncoder 인터페이스 호환
"""
import numpy as np
from sentence_transformers import CrossEncoder
class CrossEncoderReranker:
"""
Cross-Encoder 리랭킹 래퍼 (로컬, API 키 불필요).
싱글톤 패턴으로 모델 로드 1회만 수행.
"""
_instance = None
MODEL_NAME = "BAAI/bge-reranker-v2-m3"
@classmethod
def get_instance(cls) -> "CrossEncoderReranker":
if cls._instance is None:
cls._instance = cls()
return cls._instance
def __init__(self):
print(f"[리랭킹] Cross-Encoder 모델 로드 중: {self.MODEL_NAME}")
self.model = CrossEncoder(self.MODEL_NAME)
print("[리랭킹] 모델 로드 완료")
def rerank(self, query: str, candidates: list, top_n: int) -> list:
"""
RRF 결과를 Cross-Encoder로 재정렬.
Args:
query: 정규화된 검색 쿼리 (norm_query)
candidates: search_engine.search() 반환 dict 리스트
각 dict에 "document" 또는 "menu_name" 필드 필요
top_n: 반환할 최대 결과 수
Returns:
재정렬된 dict 리스트 (similarity, similarity_pct 필드 갱신됨)
원본 RRF 점수는 "_rrf_similarity" 필드로 보존
"""
if not candidates:
return candidates
# (query, document) 쌍 구성
# 【식별문맥】 prefix 제거 후 원본 embedding_text 사용
# (Contextual Retrieval로 추가된 긴 컨텍스트가 reranker 혼동 유발)
CONTEXT_MARKER = "【식별문맥】"
pairs = []
for c in candidates:
doc = c.get("document") or c.get("menu_name", "")
if CONTEXT_MARKER in doc:
# 마커 이후 두 줄바꿈 뒤의 원본 텍스트만 사용
parts = doc.split("\n\n", 1)
doc = parts[1].strip() if len(parts) > 1 else doc.split(CONTEXT_MARKER)[-1].strip()
pairs.append((query, doc))
# Cross-Encoder 점수 계산 (raw logit 값, 음수 가능)
raw_scores = self.model.predict(pairs)
# Sigmoid 변환 → [0,1] 정규화
norm_scores = 1.0 / (1.0 + np.exp(-raw_scores))
# 점수 기준 내림차순 정렬 후 top_n 반환
scored = list(zip(norm_scores, candidates))
scored.sort(key=lambda x: -x[0])
result = []
for score, cand in scored[:top_n]:
updated = dict(cand)
# 기존 RRF 점수 보존 (디버깅용)
updated["_rrf_similarity"] = cand.get("similarity", 0.0)
updated["similarity"] = round(float(score), 4)
updated["similarity_pct"] = f"{float(score) * 100:.1f}%"
result.append(updated)
return result