kgrag/graph_search/node_retrieval.py

300 lines
12 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""
===========================================
节点检索模块 - 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