Spaces:
Sleeping
Sleeping
| """ | |
| 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" | |
| 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 | |