SAP-ERP-AI-Agent / src /rag /reranker.py
daisysooyeon's picture
perf: lighter reranker (ms-marco-MiniLM-L-6-v2) via config; bake it in image
38a0085
Raw
History Blame Contribute Delete
7.29 kB
"""
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