| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899 |
- """
- BGE-Reranker-v2-m3 重排序
- Cross-Encoder 精排,将 Top-20 精准压缩到 Top-5
- 在当前未加载 Cross-Encoder 模型的情况下,提供基于信号融合的精排:
- - 最低相似度阈值过滤
- - 关键词重叠度加权
- - 内容去重
- """
- import re
- from typing import Optional
- from app.core.config import get_settings
- settings = get_settings()
- # 最低相似度阈值:低于此分数的结果很可能不相关,直接丢弃
- MIN_SIMILARITY_THRESHOLD = 0.3
- class Reranker:
- def __init__(self, model_name: Optional[str] = None):
- self.model_name = model_name or settings.reranker_model
- self._model = None
- def _load_model(self):
- """Phase 3 实现:加载 BGE-Reranker-v2-m3 Cross-Encoder"""
- pass
- def rerank(
- self,
- query: str,
- documents: list[dict],
- top_k: int = 5,
- ) -> list[dict]:
- """重排序:融合向量相似度 + 关键词重叠度,过滤低质量结果并去重。"""
- if not documents:
- return []
- # 1. 过滤:低于阈值的结果直接丢弃(避免答非所问)
- filtered = [d for d in documents if d.get("score", 0) >= MIN_SIMILARITY_THRESHOLD]
- # 2. 关键词加权:查询词在文档中出现越多,得分越高
- query_terms = self._tokenize_query(query)
- for doc in filtered:
- content = doc.get("content", "")
- keyword_bonus = self._keyword_overlap_score(query_terms, content)
- # 原始相似度 + 关键词奖励(关键词匹配最高加 0.3)
- doc["score"] = doc.get("score", 0) + keyword_bonus * 0.3
- # 3. 按融合分数排序
- sorted_docs = sorted(filtered, key=lambda d: d.get("score", 0), reverse=True)
- # 4. 去重:移除内容高度重叠的 chunk(Jaccard 相似度 > 0.8)
- deduped = []
- seen_texts = []
- for doc in sorted_docs:
- content = doc.get("content", "")
- if self._is_duplicate(content, seen_texts, threshold=0.8):
- continue
- deduped.append(doc)
- seen_texts.append(content)
- return deduped[:top_k]
- def _tokenize_query(self, query: str) -> set[str]:
- """提取查询中的关键词(中文按 1-4 字切分,英文按空格切分)。"""
- tokens = set()
- # 中文:提取 2-4 字短语
- for n in range(2, 5):
- for i in range(len(query) - n + 1):
- seg = query[i:i + n]
- if all('一' <= c <= '鿿' for c in seg):
- tokens.add(seg)
- # 英文/数字词
- for word in re.findall(r'[a-zA-Z0-9]+', query):
- tokens.add(word.lower())
- return tokens
- def _keyword_overlap_score(self, query_terms: set[str], content: str) -> float:
- """计算查询词在文档中的覆盖率(0.0 ~ 1.0)。"""
- if not query_terms:
- return 0.0
- matched = sum(1 for t in query_terms if t in content)
- return matched / len(query_terms)
- def _is_duplicate(self, content: str, seen_texts: list[str], threshold: float = 0.8) -> bool:
- """检查当前内容是否与已选中的内容高度重复(基于 Jaccard 字符集相似度)。"""
- if not content or not seen_texts:
- return False
- # 采样前 200 字符做快速比较
- content_sample = set(content[:200])
- if not content_sample:
- return False
- for seen in seen_texts[-5:]: # 只比较最近 5 个
- seen_sample = set(seen[:200])
- intersection = len(content_sample & seen_sample)
- union = len(content_sample | seen_sample)
- if union > 0 and intersection / union > threshold:
- return True
- return False
|