""" =========================================== 节点检索模块 - NodeRetrieval =========================================== 功能:根据路由结果,使用混合检索获取候选入口节点 使用HybridRetriever混合检索(向量+全文) =========================================== """ import re import asyncio import logging import time from typing import List, Dict, Any, Optional logger = logging.getLogger(__name__) # ========== ========== 导入依赖 ========== ========== try: from neo4j_graphrag.retrievers import HybridRetriever import jieba # 过滤 jieba 缓存写入权限错误的日志(不影响功能,只是无法写入缓存) jieba_logger = logging.getLogger('jieba') jieba_logger.setLevel(logging.CRITICAL) # 只显示 CRITICAL 级别,过滤所有 ERROR/WARNING/INFO/DEBUG NODE_RETRIEVAL_AVAILABLE = True except ImportError: NODE_RETRIEVAL_AVAILABLE = False class NodeRetrieval: """ 节点检索类 功能:根据路由结果,检索候选入口节点 """ def __init__(self, driver, embeddings): """ 初始化节点检索器 Args: driver: Neo4j driver实例 embeddings: 嵌入模型实例(需要有embed_documents方法) """ self.driver = driver self.embeddings = embeddings async def retrieve_nodes( self, route_res: List[Dict[str, str]], # [{"label": "Drug", "entity": "xxx"}] top_k: int = 10 ) -> Dict[str, List[Dict[str, Any]]]: """ 节点检索:根据标签和实体,检索入口节点 功能: 使用混合检索(向量+全文)进行检索 Args: route_res: 路由结果,包含标签和实体信息 top_k: 检索返回的节点数量上限 Returns: Dict[str, List[Dict]]: 检索到的节点结果,以标签为键 """ if not NODE_RETRIEVAL_AVAILABLE: logger.error("节点检索依赖未安装") return {} pairs = [] # 用于存储需要检索的标签-实体对 retrieved_nodes = {} # 用于存储检索到的节点结果,以标签为键 for item in route_res: label = item.get("label", "") entity = item.get("entity", "") if not entity: # 如果实体为空则跳过 continue # 将标签和实体作为一个元组添加到pairs列表中 pairs.append((label, entity)) if not pairs: # 如果没有需要检索的标签-实体对,直接返回 return retrieved_nodes # 将标签-实体对分离成两个独立的列表 labels, entities = zip(*pairs) labels, entities = list(labels), list(entities) # ========== 对每个实体进行中文分词处理,生成全文检索查询文本 ========== query_texts = [] for entity in entities: try: words = jieba.lcut(entity) filtered_words = [ word.strip() for word in words if word.strip() and re.fullmatch(r"[a-zA-Z0-9\u4e00-\u9fa5]+", word.strip()) ] # 如果过滤后没有词,使用原始实体 if not filtered_words: filtered_words = [entity] if entity.strip() else [] query_text = " OR ".join(filtered_words) if filtered_words else entity query_texts.append(query_text) except Exception as e: logger.warning(f"实体 '{entity}' 分词失败: {e},使用原始实体作为查询文本") query_texts.append(entity) # ========== 对实体进行向量化处理 ========== try: query_vectors = self.embeddings.embed_documents(entities) except Exception as e: logger.error(f"向量化失败: {e}", exc_info=True) query_vectors = [None] * len(entities) # ========== 为每个标签创建混合检索任务 ========== tasks = [] for label, entity, query_text, query_vector in zip(labels, entities, query_texts, query_vectors): try: # vector_index_name = label.lower() + "_vector" # fulltext_index_name = label.lower() + "_fulltext" vector_index_name = "global_searchable_embedding" fulltext_index_name = "global_searchable_fulltext_search" retriever = HybridRetriever( self.driver, vector_index_name=vector_index_name, fulltext_index_name=fulltext_index_name ) tasks.append( asyncio.to_thread( retriever.get_search_results, query_text, query_vector, top_k, effective_search_ratio=2, ) ) except Exception as e: logger.error(f"创建HybridRetriever失败 ({label}): {e}", exc_info=True) tasks.append(None) # ========== 并发执行所有检索任务 ========== valid_tasks = [task for task in tasks if task is not None] if valid_tasks: results = await asyncio.gather(*valid_tasks, return_exceptions=True) else: logger.warning("没有有效的检索任务") results = [] # ========== 处理检索结果 ========== valid_results = [] exception_results = [] for i, result in enumerate(results): if isinstance(result, Exception): label = labels[i] if i < len(labels) else "unknown" logger.error(f"检索任务 {i} (标签: {label}) 执行异常: {result}", exc_info=result) exception_results.append((i, result)) else: valid_results.append((i, result)) # ========== 按标签分组,准备进行RRF融合 ========== # 结构: {label: [(entity, result_index, nodes_list), ...]} label_results_map = {} for i, result in valid_results: label = labels[i] if i < len(labels) else "unknown" entity = entities[i] if i < len(entities) else "unknown" if result is None: logger.warning(f"标签 '{label}' (实体: '{entity}') 的检索结果为 None") continue if not hasattr(result, 'records'): logger.warning(f"标签 '{label}' (实体: '{entity}') 的结果没有 'records' 属性,类型: {type(result)}") continue # 提取节点信息(优先使用"名称"属性,兼容"name") nodes_list = [ { "name": record["node"].get("名称") or record["node"].get("name") or str(record["node"]), "score": record.get("score", 0.0), "node": dict(record["node"]) } for record in result.records ] if nodes_list: if label not in label_results_map: label_results_map[label] = [] label_results_map[label].append((entity, i, nodes_list)) # ========== 对每个标签的多个实体结果进行RRF融合 ========== # 策略:只对同一标签下相同实体名的结果进行RRF融合 rrf_fusion_start = time.time() for label, entity_results in label_results_map.items(): # ========== 按实体名分组 ========== entity_name_groups = {} # {entity_name: [(entity, result_index, nodes_list), ...]} for entity, result_index, nodes_list in entity_results: if entity not in entity_name_groups: entity_name_groups[entity] = [] entity_name_groups[entity].append((entity, result_index, nodes_list)) # ========== 对每个实体名的结果进行RRF融合,然后合并 ========== all_fused_nodes = [] for entity_name, group_results in entity_name_groups.items(): if len(group_results) == 1: # 只有一个结果,不需要RRF融合,直接使用 _, _, nodes_list = group_results[0] all_fused_nodes.extend(nodes_list) logger.debug(f"标签 '{label}' 实体 '{entity_name}' 只有一个结果,跳过RRF融合") else: # 多个结果(相同实体名),使用RRF融合 fused_nodes = self._rrf_fusion(group_results, top_k=top_k) all_fused_nodes.extend(fused_nodes) # ========== 对同一标签内的所有节点按score排序 ========== all_fused_nodes.sort(key=lambda x: x.get('score', 0), reverse=True) # ========== 限制每个标签的数量 ========== retrieved_nodes[label] = all_fused_nodes[:top_k] # 统计去重后的结果 dedup_stats = {k: len(v) for k, v in retrieved_nodes.items()} logger.info(f"检索结果统计(RRF融合后): {list(dedup_stats.items())}") return retrieved_nodes def _rrf_fusion( self, entity_results: List[tuple], top_k: int = 10, k: int = 60 ) -> List[Dict[str, Any]]: """ 使用RRF (Reciprocal Rank Fusion) 融合多个实体的检索结果 RRF公式: RRF(d) = Σ 1/(k + rank_i(d)) 其中: - d 是文档/节点 - rank_i(d) 是文档在第i个排序结果中的排名(从1开始) - k 是常数,通常为60 Args: entity_results: [(entity, result_index, nodes_list), ...] top_k: 返回的节点数量上限 k: RRF常数,默认60 Returns: List[Dict]: 融合后的节点列表,按RRF分数降序排序 """ if not entity_results: return [] # ========== 1. 为每个节点计算RRF分数 ========== node_rrf_scores = {} # {node_name: {"rrf_score": float, "node_data": dict, "ranks": [rank1, rank2, ...]}} for entity, result_index, nodes_list in entity_results: for rank, node in enumerate(nodes_list, start=1): # rank从1开始 node_name = node.get("name") if not node_name: continue if node_name not in node_rrf_scores: node_rrf_scores[node_name] = { "rrf_score": 0.0, "node_data": node, # 保留第一个出现的节点数据 "ranks": [] } # 累加RRF分数: 1/(k + rank) rrf_contribution = 1.0 / (k + rank) node_rrf_scores[node_name]["rrf_score"] += rrf_contribution node_rrf_scores[node_name]["ranks"].append(rank) # ========== 2. 转换为列表并按RRF分数排序 ========== fused_nodes = [] for node_name, data in node_rrf_scores.items(): node_data = data["node_data"].copy() node_data["score"] = data["rrf_score"] # 使用RRF分数替换原始score node_data["_rrf_ranks"] = data["ranks"] # 保留排名信息用于调试 fused_nodes.append(node_data) # 按RRF分数降序排序 fused_nodes.sort(key=lambda x: x["score"], reverse=True) # ========== 3. 限制返回数量 ========== fused_nodes = fused_nodes[:top_k] logger.debug( f"RRF融合: {len(entity_results)} 个实体, " f"融合前总节点数: {sum(len(nodes) for _, _, nodes in entity_results)}, " f"融合后: {len(fused_nodes)} 个节点" ) return fused_nodes