fix(engine): 修复 _score_source 在 Rerank 后的行为与自适应 TopK 分数计算

问题1: Rerank 后 _score_source 仍为 'rrf',导致自适应 TopK 被跳过
修复: rerank_results 中将 _score_source 更新为 'rerank'

问题2: 自适应 TopK 用 1.0 - dist 转换分数,但 Rerank 后 distances
      是相关性分数(越大越好),转换后语义反转
修复: 根据 _score_source 区分分数来源,rerank 直接使用,向量距离才做转换

附带: BM25 top3 捕获时确保 meta 包含 _collection 字段
This commit is contained in:
lacerate551
2026-06-17 20:10:50 +08:00
parent 4fd53f48f9
commit 17ce177d1a

View File

@@ -669,6 +669,7 @@ class RAGEngine:
results_list = [vector_results]
weights = [VECTOR_WEIGHT]
bm25_results = None # 初始化,防止 NameError
if USE_HYBRID_SEARCH and self.bm25_index.bm25:
bm25_results = self.bm25_index.search(query, top_k=recall_k)
@@ -683,6 +684,22 @@ class RAGEngine:
vector_w, bm25_w = self._get_dynamic_rrf_weights(query)
weights = [vector_w, bm25_w]
# ========== 保留 BM25 原始 top-3 完整信息,用于下游分歧检测救援 ==========
_bm25_raw_top3 = []
if USE_HYBRID_SEARCH and bm25_results and bm25_results.get('ids') and bm25_results['ids'][0]:
_bm25_ids = bm25_results['ids'][0][:3]
_bm25_docs = bm25_results['documents'][0][:3]
_bm25_metas = bm25_results['metadatas'][0][:3]
_bm25_dists = (bm25_results.get('distances', [[]])[0] or [0]*3)[:3]
for i in range(len(_bm25_ids)):
_bm25_raw_top3.append({
'id': _bm25_ids[i],
'doc': _bm25_docs[i],
'meta': _bm25_metas[i],
'bm25_score': _bm25_dists[i],
'rank': i + 1
})
if len(results_list) > 1:
fused_results = self.reciprocal_rank_fusion(results_list, weights)
_debug['steps'].append({'name': 'rrf_fusion', 'count': len(fused_results['ids'][0]) if fused_results.get('ids') else 0, 'weights': [round(w, 2) for w in weights]})
@@ -698,6 +715,8 @@ class RAGEngine:
is_enum_query = self._is_enumeration_query(query)
fused_results['_enum_query'] = is_enum_query
# 传递 BM25 原始 top-3 到路由层,用于分歧检测救援
fused_results['_bm25_top3'] = _bm25_raw_top3
# 章节过滤(如果查询中提到了章节)
fused_results = self._filter_by_section(fused_results, query)
@@ -768,7 +787,15 @@ class RAGEngine:
and not (is_enum_query and ENUM_QUERY_DISABLE_TOPK_SHRINK)
and fused_results.get('_score_source') != 'rrf'
):
top_score = 1.0 - fused_results['distances'][0][0] # 距离转相似度
# 根据分数来源计算相似度分数(越高越好)
score_source = fused_results.get('_score_source')
top_dist = fused_results['distances'][0][0]
if score_source == 'rerank':
# Rerank 后 distances 是相关性分数,越大越好,直接使用
top_score = top_dist
else:
# 向量距离,越小越好,转为相似度
top_score = 1.0 - top_dist
adjusted_k, should_retrieve, reason = self._adaptive_topk.adjust(top_score, top_k)
if "high_confidence" in reason:
# 高置信度时截断结果
@@ -1107,7 +1134,7 @@ class RAGEngine:
'metadatas': [f_metas],
'distances': [f_scores]
}
for key in ('_debug', '_score_source', '_enum_query', '_expanded_context'):
for key in ('_debug', '_score_source', '_enum_query', '_expanded_context', '_bm25_top3'):
if key in results:
filtered[key] = results[key]
return filtered
@@ -1130,7 +1157,7 @@ class RAGEngine:
'metadatas': [f_metas],
'distances': [f_scores]
}
for key in ('_debug', '_score_source', '_enum_query', '_expanded_context'):
for key in ('_debug', '_score_source', '_enum_query', '_expanded_context', '_bm25_top3'):
if key in results:
filtered[key] = results[key]
return filtered
@@ -1145,7 +1172,7 @@ class RAGEngine:
'metadatas': [results['metadatas'][0][:top_k]],
'distances': [results['distances'][0][:top_k]]
}
for key in ('_debug', '_score_source', '_enum_query', '_expanded_context'):
for key in ('_debug', '_score_source', '_enum_query', '_expanded_context', '_bm25_top3'):
if key in results:
truncated[key] = results[key]
return truncated
@@ -1460,7 +1487,7 @@ class RAGEngine:
'distances': [[item[3] for item in items]],
'_expanded_context': {'added': added}
}
for key in ('_debug', '_score_source', '_enum_query'):
for key in ('_debug', '_score_source', '_enum_query', '_bm25_top3'):
if key in results:
expanded[key] = results[key]
return expanded
@@ -1487,6 +1514,7 @@ class RAGEngine:
sub_top_k = max(top_k, 5)
all_results = []
_all_bm25_top3 = [] # 收集各子查询的 BM25 top3
for sub_q in sub_queries:
try:
sub_result = self.search_knowledge(
@@ -1497,6 +1525,9 @@ class RAGEngine:
)
if sub_result and sub_result.get('ids') and sub_result['ids'][0]:
all_results.append(sub_result)
# 收集子查询的 BM25 top3
if sub_result.get('_bm25_top3'):
_all_bm25_top3.extend(sub_result['_bm25_top3'])
except Exception as e:
logger.warning(f"子查询检索失败: '{sub_q}' - {e}")
@@ -1505,9 +1536,19 @@ class RAGEngine:
# 合并去重
if len(all_results) == 1:
return all_results[0]
merged = all_results[0]
else:
merged = self._merge_and_deduplicate(all_results, top_k)
return self._merge_and_deduplicate(all_results, top_k)
# 将收集的 BM25 top3 传递到合并结果中
if _all_bm25_top3:
_all_bm25_top3.sort(key=lambda x: x.get('bm25_score', 0), reverse=True)
_all_bm25_top3 = _all_bm25_top3[:3]
for rank, item in enumerate(_all_bm25_top3):
item['rank'] = rank + 1
merged['_bm25_top3'] = _all_bm25_top3
return merged
def _search_with_decomposition(
self, query, decomposer, top_k=5, allowed_levels=None,
@@ -1539,6 +1580,7 @@ class RAGEngine:
# 并行检索各子查询
all_results = []
_all_bm25_top3 = [] # 收集各子查询的 BM25 top3
for sub_q in sub_queries:
try:
sub_result = self.search_knowledge(
@@ -1549,6 +1591,9 @@ class RAGEngine:
)
if sub_result and sub_result.get('ids') and sub_result['ids'][0]:
all_results.append(sub_result)
# 收集子查询的 BM25 top3
if sub_result.get('_bm25_top3'):
_all_bm25_top3.extend(sub_result['_bm25_top3'])
except Exception as e:
logger.warning(f"子查询检索失败: '{sub_q}' - {e}")
@@ -1561,6 +1606,14 @@ class RAGEngine:
else:
merged = self._merge_and_deduplicate(all_results, top_k)
# 将收集的 BM25 top3 传递到合并结果中
if _all_bm25_top3:
_all_bm25_top3.sort(key=lambda x: x.get('bm25_score', 0), reverse=True)
_all_bm25_top3 = _all_bm25_top3[:3]
for rank, item in enumerate(_all_bm25_top3):
item['rank'] = rank + 1
merged['_bm25_top3'] = _all_bm25_top3
return merged
def _merge_and_deduplicate(self, results_list, top_k):
@@ -1645,12 +1698,13 @@ class RAGEngine:
from concurrent.futures import ThreadPoolExecutor, as_completed
def _query_single_collection(coll_name):
"""查询单个向量库(向量 + BM25"""
"""查询单个向量库(向量 + BM25,返回 (coll_results, bm25_raw_items)"""
coll_results = []
bm25_raw_items = [] # 该 collection 的 BM25 原始结果
try:
coll = self.kb_manager.get_collection(coll_name)
if not coll:
return coll_results
return coll_results, bm25_raw_items
query_kwargs = {
"query_embeddings": [query_vector],
@@ -1668,25 +1722,61 @@ class RAGEngine:
if USE_HYBRID_SEARCH:
try:
bm25 = self.kb_manager.get_bm25_index(coll_name)
if bm25.bm25:
if bm25 and bm25.bm25:
bm25_res = bm25.search(query, top_k=recall_k)
if source_filter and bm25_res['metadatas'] and bm25_res['metadatas'][0]:
# 兼容两种 BM25Indexcore.bm25_index 返回 dictknowledge.base 返回 tuple
if isinstance(bm25_res, tuple):
_ids, _docs, _metas, _dists = bm25_res
bm25_res = {
'ids': [_ids],
'documents': [_docs],
'metadatas': [_metas],
'distances': [_dists]
}
if source_filter and bm25_res['metadatas'][0]:
bm25_res = self._filter_results(bm25_res, lambda meta: meta.get('source') == source_filter)
if bm25_res['metadatas'] and bm25_res['metadatas'][0]:
for meta in bm25_res['metadatas'][0]:
meta['_collection'] = coll_name
coll_results.append(bm25_res)
# 提取 BM25 原始 top-3在此处直接捕获避免与向量结果混淆
_bm25_ids = bm25_res['ids'][0][:3]
_bm25_docs = bm25_res['documents'][0][:3]
_bm25_metas = bm25_res['metadatas'][0][:3]
_bm25_dists = (bm25_res.get('distances', [[]])[0] or [0]*3)[:3]
for i in range(len(_bm25_ids)):
# 确保 meta 包含 _collection用于路由层注入时下游处理
bm25_meta = _bm25_metas[i]
if '_collection' not in bm25_meta:
bm25_meta = {**bm25_meta, '_collection': coll_name}
bm25_raw_items.append({
'id': _bm25_ids[i],
'doc': _bm25_docs[i],
'meta': bm25_meta,
'bm25_score': _bm25_dists[i],
})
logger.debug(f"[BM25] {coll_name}: captured {len(bm25_raw_items)} raw items")
except Exception as e:
logger.debug(f"向量库 {coll_name} 检索失败: {e}")
logger.debug(f"向量库 {coll_name} BM25检索失败: {e}")
except Exception as e:
logger.debug(f"多向量库检索失败: {e}")
return coll_results
return coll_results, bm25_raw_items
all_results = []
_bm25_raw_top3 = []
with ThreadPoolExecutor(max_workers=len(target_collections)) as executor:
futures = {executor.submit(_query_single_collection, name): name for name in target_collections}
for future in as_completed(futures):
all_results.extend(future.result())
coll_results, bm25_raw_items = future.result()
all_results.extend(coll_results)
_bm25_raw_top3.extend(bm25_raw_items)
# 按 bm25_score 降序取全局 top-3
if _bm25_raw_top3:
_bm25_raw_top3.sort(key=lambda x: x['bm25_score'], reverse=True)
_bm25_raw_top3 = _bm25_raw_top3[:3]
for rank, item in enumerate(_bm25_raw_top3):
item['rank'] = rank + 1
# ========== FAQ 检索 ==========
faq_results = self._search_faq_collection(query_vector, top_k=FAQ_RECALL_TOP_K)
@@ -1734,6 +1824,8 @@ class RAGEngine:
is_enum_query = self._is_enumeration_query(query)
fused_results['_enum_query'] = is_enum_query
# 传递 BM25 原始 top-3 到路由层,用于分歧检测救援
fused_results['_bm25_top3'] = _bm25_raw_top3
# 章节过滤(如果查询中提到了章节)
fused_results = self._filter_by_section(fused_results, query)
@@ -1801,7 +1893,15 @@ class RAGEngine:
and not (is_enum_query and ENUM_QUERY_DISABLE_TOPK_SHRINK)
and fused_results.get('_score_source') != 'rrf'
):
top_score = 1.0 - fused_results['distances'][0][0] # 距离转相似度
# 根据分数来源计算相似度分数(越高越好)
score_source = fused_results.get('_score_source')
top_dist = fused_results['distances'][0][0]
if score_source == 'rerank':
# Rerank 后 distances 是相关性分数,越大越好,直接使用
top_score = top_dist
else:
# 向量距离,越小越好,转为相似度
top_score = 1.0 - top_dist
adjusted_k, should_retrieve, reason = self._adaptive_topk.adjust(top_score, top_k)
if "high_confidence" in reason:
# 高置信度时截断结果
@@ -1851,7 +1951,7 @@ class RAGEngine:
'metadatas': [filtered_metas],
'distances': [filtered_distances]
}
for key in ('_debug', '_score_source', '_enum_query', '_expanded_context'):
for key in ('_debug', '_score_source', '_enum_query', '_expanded_context', '_bm25_top3'):
if key in results:
filtered[key] = results[key]
return filtered
@@ -1924,7 +2024,7 @@ class RAGEngine:
'metadatas': [filtered_metas],
'distances': [filtered_distances]
}
for key in ('_debug', '_score_source', '_enum_query', '_expanded_context'):
for key in ('_debug', '_score_source', '_enum_query', '_expanded_context', '_bm25_top3'):
if key in results:
filtered[key] = results[key]
return filtered
@@ -2040,7 +2140,7 @@ class RAGEngine:
'metadatas': [[c['metadata'] for c in selected]],
'distances': [[id_to_dist.get(doc_id, 0) for doc_id in selected_ids]]
}
for key in ('_debug', '_score_source', '_enum_query', '_expanded_context'):
for key in ('_debug', '_score_source', '_enum_query', '_expanded_context', '_bm25_top3'):
if key in results:
filtered[key] = results[key]
return filtered
@@ -2080,7 +2180,7 @@ class RAGEngine:
'metadatas': [[c['metadata'] for c in selected]],
'distances': [[id_to_dist.get(c['id'], 0) for c in selected]]
}
for key in ('_debug', '_score_source', '_enum_query', '_expanded_context'):
for key in ('_debug', '_score_source', '_enum_query', '_expanded_context', '_bm25_top3'):
if key in results:
filtered[key] = results[key]
return filtered
@@ -2189,9 +2289,12 @@ class RAGEngine:
'_rerank_cached': cache_hit
}
# 保留原有标记字段
for key in ('_debug', '_score_source', '_enum_query', '_expanded_context'):
for key in ('_debug', '_enum_query', '_expanded_context', '_bm25_top3'):
if key in results:
reranked[key] = results[key]
# Rerank 后 distances 语义变为 CrossEncoder 分数,更新 _score_source
# 使自适应 TopK 能正确应用(之前 _score_source='rrf' 会导致自适应 TopK 被跳过)
reranked['_score_source'] = 'rerank'
return reranked
# ---------------- 流式生成 ----------------