-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathreranker.py
More file actions
151 lines (119 loc) · 4.83 KB
/
Copy pathreranker.py
File metadata and controls
151 lines (119 loc) · 4.83 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
"""
Rerank 精排模块: 用 Cross-Encoder 对粗排候选做精排
粗排(混合检索)从几百个 chunk 里快速筛出 top-k 候选,
精排(Cross-Encoder)对这 k 个候选仔细评估, 重新排名.
【面试重点】Bi-Encoder vs Cross-Encoder:
- Bi-Encoder: query 和 chunk 分别编码成向量, 比距离. 快但精度中等.
- Cross-Encoder: query 和 chunk 拼在一起输入模型, 内部交叉 Attention.
精度更高, 但每对都要过模型, 所以只对少量候选做.
Cross-Encoder 输入: "[CLS] query [SEP] chunk [SEP]"
Cross-Encoder 输出: 一个浮点数(相关度分数)
"""
from sentence_transformers import CrossEncoder
from models import SearchResult
# ============ 模型配置 ============
# BAAI/bge-reranker-base: 北京智源的 Cross-Encoder 精排模型
# 跟我们的 Embedding 模型(bge-large-zh-v1.5)是同一家族, 中英文都支持
RERANK_MODEL_NAME = "BAAI/bge-reranker-base"
# 单例模式: 模型全局只加载一次
_rerank_model = None
def get_rerank_model() -> CrossEncoder:
"""
获取 Cross-Encoder 模型实例. 第一次调用时加载, 之后复用.
CrossEncoder 跟 SentenceTransformer 不同:
- SentenceTransformer: 输入一段文本 -> 输出一个向量
- CrossEncoder: 输入一对文本(query, chunk) -> 输出一个分数
"""
global _rerank_model
if _rerank_model is None:
print(f"正在加载 Rerank 模型: {RERANK_MODEL_NAME} ...")
_rerank_model = CrossEncoder(RERANK_MODEL_NAME)
print("Rerank 模型加载完成!")
return _rerank_model
def rerank(
query: str,
candidates: list[SearchResult],
top_k: int = 5,
) -> list[SearchResult]:
"""
用 Cross-Encoder 对候选结果做精排.
流程:
1. 把 query 和每个候选的原文组成配对 [(query, chunk1), (query, chunk2), ...]
2. Cross-Encoder 对每对打分(内部: 拼接 -> BERT 前向传播 -> 输出分数)
3. 按分数从高到低排序, 取 top_k
参数:
query: 用户的搜索文本
candidates: 粗排返回的候选 SearchResult 列表
top_k: 精排后返回几条结果
返回:
精排后的 top_k SearchResult 列表, score 更新为 Cross-Encoder 分数
"""
if not candidates:
return []
model = get_rerank_model()
# 第一步: 组成 query-chunk 配对
# Cross-Encoder 的 predict() 接受 [(text_a, text_b), ...] 格式
pairs = [(query, result.content) for result in candidates]
# 第二步: Cross-Encoder 打分
# 返回一个数组, scores[i] = 第 i 个配对的相关度分数
# 分数越高越相关(不像余弦相似度有 0-1 范围, CE 分数可以是任意实数)
scores = model.predict(pairs)
# 第三步: 按分数排序
# 把 (分数, 原始索引) 配对, 按分数降序
scored_indices = sorted(
range(len(scores)),
key=lambda i: scores[i],
reverse=True,
)
# 取 top_k, 重新编排名
reranked = []
for rank, idx in enumerate(scored_indices[:top_k], start=1):
original = candidates[idx]
reranked.append(SearchResult(
chunk_id=original.chunk_id,
content=original.content,
metadata=original.metadata,
score=float(scores[idx]), # 更新为 Cross-Encoder 的分数
rank=rank,
))
return reranked
# ============ 测试代码 ============
if __name__ == "__main__":
# 模拟粗排结果, 测试 Rerank 是否能正确重排
print("=== 测试 Cross-Encoder Rerank ===\n")
query = "MCP协议是什么"
# 构造几个假的候选, 故意把不太相关的排前面
candidates = [
SearchResult(
chunk_id="test_chunk_0",
content="今天天气真好, 适合出去散步",
metadata={"filename": "weather.md"},
score=0.9,
rank=1,
),
SearchResult(
chunk_id="test_chunk_1",
content="MCP 是 Model Context Protocol, Anthropic 在 2024 年发布的协议标准, "
"让 AI 应用通过统一接口调用外部工具和数据源",
metadata={"filename": "mcp.md"},
score=0.5,
rank=2,
),
SearchResult(
chunk_id="test_chunk_2",
content="Python 的 list 和 C++ 的 vector 类似, 都是动态数组",
metadata={"filename": "python.md"},
score=0.7,
rank=3,
),
]
print(f"Query: {query}")
print(f"\n--- 粗排顺序(故意不准) ---")
for r in candidates:
print(f" #{r.rank} [{r.score:.2f}] {r.content[:50]}...")
# Rerank
reranked = rerank(query, candidates, top_k=3)
print(f"\n--- 精排顺序(Cross-Encoder) ---")
for r in reranked:
print(f" #{r.rank} [{r.score:.4f}] {r.content[:50]}...")
print("\n期望: MCP 相关的 chunk 应该排第一")