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:
143
core/engine.py
143
core/engine.py
@@ -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]:
|
||||
# 兼容两种 BM25Index:core.bm25_index 返回 dict,knowledge.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
|
||||
|
||||
# ---------------- 流式生成 ----------------
|
||||
|
||||
Reference in New Issue
Block a user