299 lines
12 KiB
Python
299 lines
12 KiB
Python
"""
|
||
===========================================
|
||
节点检索模块 - 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"
|
||
|
||
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
|
||
|