Spaces:
Sleeping
Sleeping
| """ | |
| src/rag/reranker.py | |
| bge-reranker-v2-m3λ₯Ό μ¬μ©ν λ¬Έμ μ¬μμν | |
| - sentence-transformers CrossEncoder μ¬μ© (transformers 5.x νΈν) | |
| - CPU/GPU μλ κ°μ§, μ±κΈν΄ μΊμ±μΌλ‘ λͺ¨λΈ 1νλ§ λ‘λ© | |
| - λͺ¨λΈ λ‘λ μ€ν¨ μ hybrid κ²μ μμ κ·Έλλ‘ graceful fallback | |
| - top_nμ configs.yaml rag.top_k_rerankμμ μ½μ | |
| """ | |
| import logging | |
| import re | |
| from langchain_core.documents import Document | |
| from src.config import get_config | |
| logger = logging.getLogger(__name__) | |
| # λ©ν μΌμΉ κ³μ° μ 쿼리 μͺ½μμ 무μν νν λΆμ©μ΄ (λΆλͺ¨ ν¬μ λ°©μ§) | |
| _META_STOP = frozenset( | |
| "the a an of to in on at for and or is are how do does i we you can with " | |
| "what when where which that this it its my our your be by from as".split() | |
| ) | |
| # μ±κΈν΄: bge-reranker 1νλ§ λ‘λ© | |
| _reranker = None # CrossEncoder μΈμ€ν΄μ€ λλ "fallback" | |
| _FALLBACK_SENTINEL = "fallback" | |
| _MODEL_NAME = "BAAI/bge-reranker-v2-m3" | |
| _RRF_K = 60 # κ²μ reciprocal-rank μμ (retriever._rrf_merge μ λμΌ κ΄λ‘) | |
| def _minmax(vals: list[float]) -> list[float]: | |
| """리μ€νΈλ₯Ό [0,1]λ‘ min-max μ κ·ν. λͺ¨λ κ°μ κ°μ΄λ©΄ 0.5λ‘ μ±μ.""" | |
| lo, hi = min(vals), max(vals) | |
| if hi - lo < 1e-9: | |
| return [0.5] * len(vals) | |
| return [(v - lo) / (hi - lo) for v in vals] | |
| def _meta_match(query_text: str, doc: Document) -> float: | |
| """쿼리 ν€μλμ μ²ν¬ λ©ν(unit/lesson/section)μ ν ν° κ²ΉμΉ¨ λΉμ¨ [0,1]. | |
| νλ νν°κ° μλλΌ μννΈ λΆμ€νΈμ© μ νΈ β ν€λκ° μλ² λ©μ μ΄λ―Έ μΌλΆ λ°μνμ§λ§, | |
| μ¬κΈ°μ λͺ μμ Β·νλλΈνκ² λ³΄κ°νλ€. λΆμ©μ΄λ λΆλͺ¨μμ μ μΈν΄ ν¬μμ μ€μΈλ€.""" | |
| q = {t for t in re.findall(r"\w+", query_text.lower()) if t not in _META_STOP} | |
| if not q: | |
| return 0.0 | |
| meta = " ".join(str(doc.metadata.get(k, "")) for k in ("unit", "lesson", "section")) | |
| m = set(re.findall(r"\w+", meta.lower())) | |
| return len(q & m) / len(q) | |
| def _get_reranker(): | |
| global _reranker | |
| if _reranker is not None: | |
| return _reranker | |
| # sentence-transformersμ CrossEncoderλ‘ λ‘λ©νλ€. | |
| # (ꡬ FlagEmbeddingμ transformers 5.xμμ is_torch_fx_available μ κ±°λ‘ import μ€ν¨) | |
| model_name = get_config().rag.reranker_model or _MODEL_NAME | |
| try: | |
| from sentence_transformers import CrossEncoder | |
| logger.info("[reranker] Loading %s via CrossEncoder β¦", model_name) | |
| _reranker = CrossEncoder(model_name, max_length=512) | |
| logger.info("[reranker] %s loaded.", model_name) | |
| except Exception as e: | |
| logger.warning( | |
| "[reranker] Failed to load CrossEncoder: %s. " | |
| "Falling back to score-aggregation ordering.", | |
| e, | |
| ) | |
| _reranker = _FALLBACK_SENTINEL | |
| return _reranker | |
| def rerank( | |
| query: str, | |
| docs: list[Document], | |
| top_n: int | None = None, | |
| queries: list[str] | None = None, | |
| ) -> list[Document]: | |
| """ | |
| bge-reranker-v2-m3λ‘ λ¬Έμ μ¬μμν ν top_n λ°ν. | |
| queriesκ° μ£Όμ΄μ§λ©΄ Multi-Query Reranking: | |
| κ° λ¬Έμλ₯Ό λͺ¨λ 쿼리μ λν΄ μ μ μ°μ β λ¬Έμλ³ max μ μλ‘ μμ κ²°μ . | |
| queriesκ° μμΌλ©΄ query λ¨μΌ κΈ°μ€μΌλ‘ μμ κ²°μ . | |
| FlagEmbedding λ‘λ μ€ν¨ μ EnsembleRetriever μ μ ν©μ° μμ κ·Έλλ‘ λ°ν. | |
| Args: | |
| query: λν κ²μ 쿼리 (λ¨μΌ λͺ¨λ λλ multi-queryμ primary) | |
| docs: Hybrid Retrieverκ° λ°νν ν보 λ¬Έμ λͺ©λ‘ | |
| top_n: μ΅μ’ λ°ν λ¬Έμ μ (Noneμ΄λ©΄ configs.rag.top_k_rerank μ¬μ©) | |
| queries: 쿼리 λ³ν λͺ©λ‘ (Query Expansion κ²°κ³Ό). μμΌλ©΄ multi-query λͺ¨λ. | |
| """ | |
| if top_n is None: | |
| top_n = get_config().rag.top_k_rerank | |
| if not docs: | |
| return [] | |
| reranker = _get_reranker() | |
| # ββ fallback: EnsembleRetriever μ μ ν©μ° μμ μ μ§ ββββββββββββββββββββββ | |
| if reranker is _FALLBACK_SENTINEL: | |
| logger.debug("[reranker] Using fallback ordering (top %d of %d)", top_n, len(docs)) | |
| return docs[:top_n] | |
| # ββ 쿼리 λͺ©λ‘ κ²°μ ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| all_queries = queries if queries else [query] | |
| # ββ bge-reranker-v2-m3: λ¬Έμλ³ max μ μ μ°μ βββββββββββββββββββββββββββββ | |
| # CrossEncoder.predictλ μ λͺ©λ‘μ λν relevance μ μ(λμμλ‘ κ΄λ ¨)λ₯Ό λ°ννλ€. | |
| doc_scores = [float("-inf")] * len(docs) | |
| for q in all_queries: | |
| pairs = [(q, doc.page_content) for doc in docs] | |
| try: | |
| scores = reranker.predict(pairs) | |
| if isinstance(scores, float): | |
| scores = [scores] | |
| for i, s in enumerate(scores): | |
| s = float(s) | |
| if s > doc_scores[i]: | |
| doc_scores[i] = s | |
| except Exception as e: | |
| logger.warning("[reranker] predict failed for query %r: %s", q[:40], e) | |
| # ββ μ μ λΈλ λ©: cross-encoder μ μ + κ²μ(RRF) μμ κ²°ν© ββββββββββββββββ | |
| # μ λ ₯ docs λ νμ΄λΈλ¦¬λ retriever κ° RRF μμΌλ‘ λκΈ΄ ν보λΌ, μ λ ₯ μΈλ±μ€ i κ° κ³§ | |
| # κ²μ μμλ€. cross-encoder κ° μ νκ°ν΄λ dense/BM25 κ° κ°νκ² λ―Ό μμ νλ³΄κ° | |
| # top_n λ°μΌλ‘ λ°λ €λμ§ μλλ‘, λ μ νΈλ₯Ό [0,1] μ κ·ν ν κ°μ€ν©νλ€. | |
| # final = alpha * cross_encoder + (1 - alpha) * κ²μ_reciprocal_rank | |
| # alpha=1.0 μ΄λ©΄ μμ 리λ컀(κΈ°μ‘΄ λμ). (νΉν ννΈνλ OCR μ²ν¬κ° dense μμμΈλ° | |
| # 리λμ»€κ° μ νκ°νλ κ²½μ°λ₯Ό ꡬμ ) | |
| cfg_rag = get_config().rag | |
| alpha = cfg_rag.rerank_blend_alpha | |
| beta = cfg_rag.meta_boost_beta | |
| # λΈλ λ©(alpha<1) λλ λ©ν λΆμ€νΈ(beta>0)κ° μΌμ§λ©΄ μ κ·ν κ²°ν© μ μλ₯Ό μ΄λ€. | |
| # (alpha=1.0 & beta>0 μ΄λ©΄ final = μ κ·ν cross_encoder + beta*meta β μμ 리λ컀 + λ©ν) | |
| if (alpha < 1.0 or beta > 0.0) and len(docs) > 1: | |
| finite = [s for s in doc_scores if s != float("-inf")] | |
| floor = min(finite) if finite else 0.0 | |
| ce = [s if s != float("-inf") else floor for s in doc_scores] | |
| rr = [1.0 / (_RRF_K + i + 1) for i in range(len(docs))] | |
| ce_n, rr_n = _minmax(ce), _minmax(rr) | |
| if beta > 0.0: | |
| meta_q = " ".join(all_queries) # λͺ¨λ 쿼리 λ³νμ ν ν°μ λ©νμ λμ‘° | |
| mm_n = _minmax([_meta_match(meta_q, d) for d in docs]) | |
| else: | |
| mm_n = [0.0] * len(docs) | |
| final = [alpha * c + (1.0 - alpha) * r + beta * m | |
| for c, r, m in zip(ce_n, rr_n, mm_n)] | |
| else: | |
| final = doc_scores | |
| ranked = sorted(zip(final, docs), key=lambda x: x[0], reverse=True) | |
| result = [doc for _, doc in ranked[:top_n]] | |
| logger.debug( | |
| "[reranker] Reranked %d docs β top %d (queries=%d, blend_alpha=%.2f) | scores: %s", | |
| len(docs), top_n, len(all_queries), alpha, | |
| [f"{s:.3f}" for s, _ in ranked[:top_n]], | |
| ) | |
| return result | |