diff --git a/core/engine.py b/core/engine.py index 4e3b878..f4283fc 100644 --- a/core/engine.py +++ b/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 # ---------------- 流式生成 ----------------