GraphRAG图检索增强生成深度解析:从知识图谱到结构化推理的工程实现

举报
柠檬🍋 发表于 2026/08/30 14:37:40 2026/08/30
【摘要】 GraphRAG图检索增强生成深度解析:从知识图谱到结构化推理的工程实现 一、引言:传统RAG的结构化局限传统RAG基于向量检索,将文档切分为chunk后按语义相似度召回。这种方式在处理实体间复杂关系(如"A公司的CEO同时也是B公司的董事")时存在局限——关系信息分散在不同chunk中,向量检索难以完整捕获。GraphRAG将知识图谱引入RAG流程,通过实体-关系-属性的结构化表示,支持...

GraphRAG图检索增强生成深度解析:从知识图谱到结构化推理的工程实现

一、引言:传统RAG的结构化局限

传统RAG基于向量检索,将文档切分为chunk后按语义相似度召回。这种方式在处理实体间复杂关系(如"A公司的CEO同时也是B公司的董事")时存在局限——关系信息分散在不同chunk中,向量检索难以完整捕获。GraphRAG将知识图谱引入RAG流程,通过实体-关系-属性的结构化表示,支持多跳推理和关系查询。微软提出的GraphRAG方法结合了图谱构建、社区检测和层次摘要,在复杂关系推理任务上显著优于传统向量RAG。本文将深入解析GraphRAG的架构设计和工程实现。

二、GraphRAG架构设计

import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
import re
import json
from typing import List, Dict, Any, Optional, Tuple, Set
from dataclasses import dataclass, field
from collections import defaultdict, deque
import math

@dataclass
class Entity:
    """实体"""
    id: str
    name: str
    entity_type: str
    description: str = ""
    properties: Dict[str, Any] = field(default_factory=dict)

@dataclass
class Relationship:
    """关系"""
    id: str
    source_id: str
    target_id: str
    relation_type: str
    description: str = ""
    properties: Dict[str, Any] = field(default_factory=dict)
    weight: float = 1.0

@dataclass
class Community:
    """社区"""
    id: str
    entity_ids: List[str]
    summary: str = ""
    level: int = 0

class KnowledgeGraph:
    """知识图谱"""
    
    def __init__(self):
        self.entities: Dict[str, Entity] = {}
        self.relationships: Dict[str, Relationship] = {}
        self.adjacency: Dict[str, List[str]] = defaultdict(list)  # entity_id -> [related_entity_ids]
        self.communities: Dict[str, Community] = {}
    
    def add_entity(self, entity: Entity):
        self.entities[entity.id] = entity
    
    def add_relationship(self, rel: Relationship):
        self.relationships[rel.id] = rel
        self.adjacency[rel.source_id].append(rel.target_id)
        self.adjacency[rel.target_id].append(rel.source_id)  # 无向图
    
    def get_entity(self, entity_id: str) -> Optional[Entity]:
        return self.entities.get(entity_id)
    
    def get_relationships(self, entity_id: str) -> List[Relationship]:
        """获取实体的所有关系"""
        rels = []
        for rel in self.relationships.values():
            if rel.source_id == entity_id or rel.target_id == entity_id:
                rels.append(rel)
        return rels
    
    def get_neighbors(self, entity_id: str, depth: int = 1) -> Set[str]:
        """获取N跳邻居"""
        visited = set()
        queue = deque([(entity_id, 0)])
        
        while queue:
            current, d = queue.popleft()
            if current in visited or d > depth:
                continue
            visited.add(current)
            for neighbor in self.adjacency.get(current, []):
                if neighbor not in visited:
                    queue.append((neighbor, d + 1))
        
        visited.discard(entity_id)
        return visited
    
    def find_path(self, source: str, target: str) -> List[str]:
        """查找两个实体间的路径(BFS)"""
        if source == target:
            return [source]
        
        visited = {source}
        queue = deque([(source, [source])])
        
        while queue:
            current, path = queue.popleft()
            
            for neighbor in self.adjacency.get(current, []):
                if neighbor == target:
                    return path + [neighbor]
                
                if neighbor not in visited:
                    visited.add(neighbor)
                    queue.append((neighbor, path + [neighbor]))
        
        return []
    
    def detect_communities(self, algorithm: str = "louvain") -> Dict[str, Community]:
        """社区检测"""
        if algorithm == "louvain":
            return self._louvain_communities()
        else:
            return self._connected_components()
    
    def _connected_components(self) -> Dict[str, Community]:
        """连通分量作为社区"""
        visited = set()
        communities = {}
        comm_id = 0
        
        for entity_id in self.entities:
            if entity_id in visited:
                continue
            
            # BFS找连通分量
            component = set()
            queue = deque([entity_id])
            
            while queue:
                current = queue.popleft()
                if current in component:
                    continue
                component.add(current)
                visited.add(current)
                
                for neighbor in self.adjacency.get(current, []):
                    if neighbor not in component:
                        queue.append(neighbor)
            
            communities[f"comm_{comm_id}"] = Community(
                id=f"comm_{comm_id}",
                entity_ids=list(component),
                level=0
            )
            comm_id += 1
        
        self.communities = communities
        return communities
    
    def _louvain_communities(self) -> Dict[str, Community]:
        """简化的Louvain社区检测"""
        # 实际Louvain算法更复杂,这里用简化版
        return self._connected_components()
    
    def to_text(self) -> str:
        """将图谱转为文本描述"""
        lines = []
        for entity in self.entities.values():
            lines.append(f"Entity: {entity.name} ({entity.entity_type}) - {entity.description}")
        
        for rel in self.relationships.values():
            source = self.entities.get(rel.source_id)
            target = self.entities.get(rel.target_id)
            if source and target:
                lines.append(f"Relation: {source.name} --[{rel.relation_type}]--> {target.name}")
        
        return '\n'.join(lines)

class GraphExtractor:
    """图谱抽取器:从文本中抽取实体和关系"""
    
    def __init__(self, llm_generate: callable = None):
        self.llm = llm_generate
        self.entity_patterns = {
            'PERSON': r'[A-Z][a-z]+ [A-Z][a-z]+',
            'ORG': r'[A-Z][a-zA-Z]+ (?:Inc|Corp|Ltd|Company|Corporation)',
            'LOC': r'(?:[A-Z][a-z]+ )?(?:City|Country|State|Province)',
            'DATE': r'\d{4}|\d{1,2}/\d{1,2}/\d{4}',
        }
        self.relation_patterns = [
            (r'(\w+)\s+(?:is|was)\s+(?:the\s+)?(?:CEO|CTO|CFO|founder|director)\s+of\s+(\w+)', 'CEO_OF'),
            (r'(\w+)\s+(?:is|was)\s+(?:born|located|based)\s+in\s+(\w+)', 'LOCATED_IN'),
            (r'(\w+)\s+(?:acquired|bought|purchased)\s+(\w+)', 'ACQUIRED'),
            (r'(\w+)\s+(?:is|are)\s+(?:part of|subsidiary of)\s+(\w+)', 'SUBSIDIARY_OF'),
            (r'(\w+)\s+(?:founded|created|established)\s+(\w+)', 'FOUNDED'),
            (r'(\w+)\s+(?:works? at|employed by)\s+(\w+)', 'WORKS_AT'),
        ]
    
    def extract(self, text: str) -> Tuple[List[Entity], List[Relationship]]:
        """从文本抽取实体和关系"""
        entities = {}
        relationships = []
        
        # 抽取实体
        for etype, pattern in self.entity_patterns.items():
            for match in re.finditer(pattern, text):
                name = match.group()
                eid = f"e_{name.lower().replace(' ', '_')}"
                if eid not in entities:
                    entities[eid] = Entity(
                        id=eid, name=name, entity_type=etype,
                        description=f"{name} found in text"
                    )
        
        # 抽取关系
        for pattern, rel_type in self.relation_patterns:
            for match in re.finditer(pattern, text):
                source_name = match.group(1)
                target_name = match.group(2)
                
                source_id = f"e_{source_name.lower().replace(' ', '_')}"
                target_id = f"e_{target_name.lower().replace(' ', '_')}"
                
                # 确保实体存在
                if source_id not in entities:
                    entities[source_id] = Entity(
                        id=source_id, name=source_name, entity_type="UNKNOWN"
                    )
                if target_id not in entities:
                    entities[target_id] = Entity(
                        id=target_id, name=target_name, entity_type="UNKNOWN"
                    )
                
                rel_id = f"r_{source_id}_{rel_type}_{target_id}"
                relationships.append(Relationship(
                    id=rel_id,
                    source_id=source_id,
                    target_id=target_id,
                    relation_type=rel_type,
                    description=f"{source_name} {rel_type} {target_name}"
                ))
        
        return list(entities.values()), relationships

class GraphRAGRetriever:
    """GraphRAG检索器"""
    
    def __init__(self, kg: KnowledgeGraph, embed_func: callable = None):
        self.kg = kg
        self.embed = embed_func or (lambda x: np.random.randn(384))
        self.entity_embeddings = {}
        self._build_embeddings()
    
    def _build_embeddings(self):
        """为实体构建embedding"""
        for eid, entity in self.kg.entities.items():
            text = f"{entity.name} {entity.description}"
            self.entity_embeddings[eid] = self.embed(text)
    
    def retrieve(self, query: str, top_k: int = 5, 
                n_hops: int = 2) -> Dict[str, Any]:
        """检索相关图谱信息"""
        # 1. 向量检索找到初始实体
        query_vec = self.embed(query)
        
        entity_scores = []
        for eid, emb in self.entity_embeddings.items():
            sim = np.dot(query_vec, emb) / (
                np.linalg.norm(query_vec) * np.linalg.norm(emb) + 1e-8
            )
            entity_scores.append((eid, sim))
        
        entity_scores.sort(key=lambda x: x[1], reverse=True)
        initial_entities = [eid for eid, _ in entity_scores[:top_k]]
        
        # 2. 图扩展:获取N跳邻居
        expanded_entities = set()
        for eid in initial_entities:
            expanded_entities.add(eid)
            neighbors = self.kg.get_neighbors(eid, depth=n_hops)
            expanded_entities.update(neighbors)
        
        # 3. 收集相关关系
        relevant_rels = []
        for eid in expanded_entities:
            rels = self.kg.get_relationships(eid)
            for rel in rels:
                if rel.source_id in expanded_entities and rel.target_id in expanded_entities:
                    relevant_rels.append(rel)
        
        # 4. 查找路径
        paths = []
        if len(initial_entities) >= 2:
            path = self.kg.find_path(initial_entities[0], initial_entities[1])
            if path:
                paths.append(path)
        
        # 5. 查找社区
        communities = self.kg.detect_communities()
        relevant_communities = []
        for comm in communities.values():
            if any(eid in comm.entity_ids for eid in initial_entities):
                relevant_communities.append(comm)
        
        # 6. 构建上下文文本
        context = self._build_context(expanded_entities, relevant_rels, paths)
        
        return {
            'query': query,
            'initial_entities': initial_entities,
            'expanded_entities': list(expanded_entities),
            'relationships': relevant_rels,
            'paths': paths,
            'communities': relevant_communities,
            'context': context
        }
    
    def _build_context(self, entity_ids: Set[str], 
                      relationships: List[Relationship],
                      paths: List[List[str]]) -> str:
        """构建图谱上下文文本"""
        parts = []
        
        # 实体信息
        parts.append("=== 相关实体 ===")
        for eid in entity_ids:
            entity = self.kg.get_entity(eid)
            if entity:
                parts.append(f"- {entity.name} ({entity.entity_type}): {entity.description}")
        
        # 关系信息
        parts.append("\n=== 相关关系 ===")
        seen_rels = set()
        for rel in relationships:
            if rel.id not in seen_rels:
                source = self.kg.get_entity(rel.source_id)
                target = self.kg.get_entity(rel.target_id)
                if source and target:
                    parts.append(f"- {source.name} --[{rel.relation_type}]--> {target.name}")
                    seen_rels.add(rel.id)
        
        # 路径信息
        if paths:
            parts.append("\n=== 推理路径 ===")
            for path in paths:
                path_names = []
                for eid in path:
                    entity = self.kg.get_entity(eid)
                    if entity:
                        path_names.append(entity.name)
                parts.append(" -> ".join(path_names))
        
        return '\n'.join(parts)

class GraphRAGSystem:
    """完整GraphRAG系统"""
    
    def __init__(self):
        self.kg = KnowledgeGraph()
        self.extractor = GraphExtractor()
        self.retriever = None
    
    def ingest_documents(self, documents: List[str]):
        """导入文档,抽取图谱"""
        for doc in documents:
            entities, relationships = self.extractor.extract(doc)
            for entity in entities:
                self.kg.add_entity(entity)
            for rel in relationships:
                self.kg.add_relationship(rel)
        
        # 社区检测
        self.kg.detect_communities()
        
        # 构建检索器
        self.retriever = GraphRAGRetriever(self.kg)
        
        print(f"导入完成: {len(self.kg.entities)} 实体, "
              f"{len(self.kg.relationships)} 关系, "
              f"{len(self.kg.communities)} 社区")
    
    def query(self, question: str, llm_generate: callable = None) -> Dict[str, Any]:
        """查询"""
        # 图检索
        retrieval = self.retriever.retrieve(question, top_k=5, n_hops=2)
        
        # 构建prompt
        prompt = f"""Based on the following knowledge graph context, answer the question.

Context:
{retrieval['context']}

Question: {question}

Answer:"""
        
        if llm_generate:
            answer = llm_generate(prompt)
        else:
            answer = f"Based on the graph with {len(retrieval['expanded_entities'])} entities..."
        
        return {
            'question': question,
            'answer': answer,
            'retrieval': retrieval
        }

def test_graphrag():
    """测试GraphRAG"""
    system = GraphRAGSystem()
    
    documents = [
        "John Smith is the CEO of TechCorp Inc. TechCorp is located in Silicon Valley.",
        "TechCorp acquired DataSoft Ltd in 2023. DataSoft was founded by Alice Brown.",
        "Alice Brown works at TechCorp after the acquisition. She is based in New York City.",
        "Bob Johnson is the CTO of TechCorp. He founded CloudService Company.",
    ]
    
    system.ingest_documents(documents)
    
    # 图谱信息
    print("\n=== 知识图谱 ===")
    print(system.kg.to_text()[:500])
    
    # 查询
    print("\n=== GraphRAG查询 ===")
    questions = [
        "Who is the CEO of TechCorp?",
        "What is the relationship between John Smith and Alice Brown?",
        "What companies are related to TechCorp?",
    ]
    
    for q in questions:
        result = system.query(q)
        print(f"\nQ: {q}")
        print(f"初始实体: {result['retrieval']['initial_entities']}")
        print(f"扩展实体数: {len(result['retrieval']['expanded_entities'])}")
        print(f"关系数: {len(result['retrieval']['relationships'])}")
        print(f"路径: {result['retrieval']['paths']}")
    
    # 对比传统RAG
    print("\n=== GraphRAG vs 传统RAG ===")
    comparisons = [
        ("实体关系查询", "支持多跳推理", "依赖chunk共现"),
        ("结构化信息", "实体-关系-属性", "非结构化文本"),
        ("多跳推理", "图遍历自然支持", "需要多个chunk拼接"),
        ("更新效率", "增量更新图谱", "重新embed文档"),
        ("可解释性", "路径可视化", "相似度排序"),
        ("复杂查询", "Cypher类查询", "向量相似度"),
    ]
    for aspect, graph_rag, vector_rag in comparisons:
        print(f"  {aspect}: GraphRAG={graph_rag} | VectorRAG={vector_rag}")

if __name__ == "__main__":
    test_graphrag()

三、社区摘要与层次检索

微软GraphRAG的一个关键创新是层次社区摘要:将图谱通过社区检测划分为社区,为每个社区生成摘要,形成层次结构。查询时先匹配社区摘要,再深入社区内部检索。

def community_summarization():
    print("GraphRAG层次社区摘要:")
    print("1. 图谱构建: 从文档抽取实体和关系")
    print("2. 社区检测: Louvain算法划分社区")
    print("3. 层次聚类: 构建多层社区层次")
    print("4. 社区摘要: LLM为每个社区生成摘要")
    print("5. 层次检索: 全局问题->社区摘要, 局部问题->实体检索")
    
    print("\n检索策略选择:")
    strategies = [
        ("全局摘要检索", "概述性问题", "'总结AI行业趋势'"),
        ("局部实体检索", "具体事实问题", "'谁是X公司CEO'"),
        ("混合检索", "复合问题", "'X公司CEO与Y公司的关系'"),
        ("多跳路径", "推理问题", "'A如何影响B'"),
    ]
    for name, use_case, example in strategies:
        print(f"  - {name}: {use_case} (例: {example})")

if __name__ == "__main__":
    community_summarization()

四、总结

GraphRAG通过将知识图谱引入RAG流程,克服了传统向量RAG在实体关系推理上的局限。图谱的结构化表示(实体-关系-属性)支持多跳推理和路径查询,社区检测和层次摘要实现了从全局到局部的层次化检索。在处理"A公司的CEO与B公司的什么人有关系"这类需要跨文档实体关联的复杂查询时,GraphRAG显著优于传统向量RAG。随着知识图谱构建自动化和图谱规模的增长,GraphRAG将在企业知识管理、法律分析和科学研究等需要复杂关系推理的场景中发挥越来越重要的作用。

【声明】本内容来自华为云开发者社区博主,不代表华为云及华为云开发者社区的观点和立场。转载时必须标注文章的来源(华为云社区)、文章链接、文章作者等基本信息,否则作者和本社区有权追究责任。如果您发现本社区中有涉嫌抄袭的内容,欢迎发送邮件进行举报,并提供相关证据,一经查实,本社区将立刻删除涉嫌侵权内容,举报邮箱: cloudbbs@huaweicloud.com
  • 点赞
  • 收藏
  • 关注作者

评论(0

0/1000
抱歉,系统识别当前为高风险访问,暂不支持该操作

全部回复

上滑加载中

设置昵称

在此一键设置昵称,即可参与社区互动!

*长度不超过10个汉字或20个英文字符,设置后3个月内不可修改。

*长度不超过10个汉字或20个英文字符,设置后3个月内不可修改。