reranker.py 3.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899
  1. """
  2. BGE-Reranker-v2-m3 重排序
  3. Cross-Encoder 精排,将 Top-20 精准压缩到 Top-5
  4. 在当前未加载 Cross-Encoder 模型的情况下,提供基于信号融合的精排:
  5. - 最低相似度阈值过滤
  6. - 关键词重叠度加权
  7. - 内容去重
  8. """
  9. import re
  10. from typing import Optional
  11. from app.core.config import get_settings
  12. settings = get_settings()
  13. # 最低相似度阈值:低于此分数的结果很可能不相关,直接丢弃
  14. MIN_SIMILARITY_THRESHOLD = 0.3
  15. class Reranker:
  16. def __init__(self, model_name: Optional[str] = None):
  17. self.model_name = model_name or settings.reranker_model
  18. self._model = None
  19. def _load_model(self):
  20. """Phase 3 实现:加载 BGE-Reranker-v2-m3 Cross-Encoder"""
  21. pass
  22. def rerank(
  23. self,
  24. query: str,
  25. documents: list[dict],
  26. top_k: int = 5,
  27. ) -> list[dict]:
  28. """重排序:融合向量相似度 + 关键词重叠度,过滤低质量结果并去重。"""
  29. if not documents:
  30. return []
  31. # 1. 过滤:低于阈值的结果直接丢弃(避免答非所问)
  32. filtered = [d for d in documents if d.get("score", 0) >= MIN_SIMILARITY_THRESHOLD]
  33. # 2. 关键词加权:查询词在文档中出现越多,得分越高
  34. query_terms = self._tokenize_query(query)
  35. for doc in filtered:
  36. content = doc.get("content", "")
  37. keyword_bonus = self._keyword_overlap_score(query_terms, content)
  38. # 原始相似度 + 关键词奖励(关键词匹配最高加 0.3)
  39. doc["score"] = doc.get("score", 0) + keyword_bonus * 0.3
  40. # 3. 按融合分数排序
  41. sorted_docs = sorted(filtered, key=lambda d: d.get("score", 0), reverse=True)
  42. # 4. 去重:移除内容高度重叠的 chunk(Jaccard 相似度 > 0.8)
  43. deduped = []
  44. seen_texts = []
  45. for doc in sorted_docs:
  46. content = doc.get("content", "")
  47. if self._is_duplicate(content, seen_texts, threshold=0.8):
  48. continue
  49. deduped.append(doc)
  50. seen_texts.append(content)
  51. return deduped[:top_k]
  52. def _tokenize_query(self, query: str) -> set[str]:
  53. """提取查询中的关键词(中文按 1-4 字切分,英文按空格切分)。"""
  54. tokens = set()
  55. # 中文:提取 2-4 字短语
  56. for n in range(2, 5):
  57. for i in range(len(query) - n + 1):
  58. seg = query[i:i + n]
  59. if all('一' <= c <= '鿿' for c in seg):
  60. tokens.add(seg)
  61. # 英文/数字词
  62. for word in re.findall(r'[a-zA-Z0-9]+', query):
  63. tokens.add(word.lower())
  64. return tokens
  65. def _keyword_overlap_score(self, query_terms: set[str], content: str) -> float:
  66. """计算查询词在文档中的覆盖率(0.0 ~ 1.0)。"""
  67. if not query_terms:
  68. return 0.0
  69. matched = sum(1 for t in query_terms if t in content)
  70. return matched / len(query_terms)
  71. def _is_duplicate(self, content: str, seen_texts: list[str], threshold: float = 0.8) -> bool:
  72. """检查当前内容是否与已选中的内容高度重复(基于 Jaccard 字符集相似度)。"""
  73. if not content or not seen_texts:
  74. return False
  75. # 采样前 200 字符做快速比较
  76. content_sample = set(content[:200])
  77. if not content_sample:
  78. return False
  79. for seen in seen_texts[-5:]: # 只比较最近 5 个
  80. seen_sample = set(seen[:200])
  81. intersection = len(content_sample & seen_sample)
  82. union = len(content_sample | seen_sample)
  83. if union > 0 and intersection / union > threshold:
  84. return True
  85. return False