问答大模型技术方案算法实现-RAPTOR树构建算法与BEG集成使用

算法核心特点

递归聚类:逐层对文档块进行聚类,构建层次结构

摘要生成:为每个聚类生成摘要,捕获语义信息

展开树检索:将所有节点展开,直接进行向量检索

层次化检索:从根节点开始,逐层深入检索

灵活聚类:支持层次聚类和基于阈值的聚类

自适应:根据节点数量动态调整树结构

可扩展:支持自定义摘要生成器和检索策略

python 复制代码
import numpy as np
from typing import List, Dict, Any, Optional, Tuple
from sklearn.cluster import AgglomerativeClustering
from sklearn.metrics.pairwise import cosine_similarity
import networkx as nx
import json
import logging
from collections import defaultdict
import torch
from dataclasses import dataclass, field
from abc import ABC, abstractmethod

# 配置日志
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)


@dataclass
class RAPTORNode:
    """
    RAPTOR树节点类
    """
    node_id: str
    content: str
    embedding: np.ndarray
    children: List['RAPTORNode'] = field(default_factory=list)
    parent: Optional['RAPTORNode'] = None
    level: int = 0
    cluster_id: int = -1
    is_leaf: bool = True
    metadata: Dict[str, Any] = field(default_factory=dict)
    
    def add_child(self, child: 'RAPTORNode'):
        """添加子节点"""
        child.parent = self
        child.level = self.level + 1
        self.children.append(child)
        self.is_leaf = False
    
    def get_all_children(self) -> List['RAPTORNode']:
        """获取所有子节点(递归)"""
        children = []
        for child in self.children:
            children.append(child)
            children.extend(child.get_all_children())
        return children
    
    def get_leaf_nodes(self) -> List['RAPTORNode']:
        """获取所有叶子节点"""
        if self.is_leaf:
            return [self]
        leaves = []
        for child in self.children:
            leaves.extend(child.get_leaf_nodes())
        return leaves
    
    def to_dict(self) -> Dict[str, Any]:
        """转换为字典"""
        return {
            'node_id': self.node_id,
            'content': self.content[:100] + '...' if len(self.content) > 100 else self.content,
            'level': self.level,
            'cluster_id': self.cluster_id,
            'is_leaf': self.is_leaf,
            'children_count': len(self.children),
            'metadata': self.metadata
        }


class BaseSummarizer(ABC):
    """摘要生成器基类"""
    
    @abstractmethod
    def generate_summary(self, texts: List[str]) -> str:
        """生成摘要"""
        pass


class SimpleSummarizer(BaseSummarizer):
    """简单摘要生成器(基于提取式)"""
    
    def __init__(self, max_length: int = 100):
        self.max_length = max_length
    
    def generate_summary(self, texts: List[str]) -> str:
        """简单的提取式摘要"""
        if not texts:
            return ""
        
        # 简单的实现:取每个句子的前几个词
        combined_text = " ".join(texts)
        words = combined_text.split()
        
        if len(words) <= self.max_length:
            return combined_text
        
        # 提取最重要的句子(简化版:取中间部分)
        return " ".join(words[:self.max_length]) + "..."


class LLMSummarizer(BaseSummarizer):
    """基于LLM的摘要生成器"""
    
    def __init__(self, model_name: str = "Qwen/Qwen2-7B-Instruct"):
        self.model_name = model_name
        # 这里可以集成实际的LLM
        logger.info(f"初始化LLM摘要生成器: {model_name}")
    
    def generate_summary(self, texts: List[str]) -> str:
        """使用LLM生成摘要"""
        if not texts:
            return ""
        
        # 构建提示词
        prompt = f"""
        请为以下文本生成一个简洁的摘要,保留核心语义信息:
        
        {' '.join(texts)}
        
        摘要(50-100字):
        """
        
        # 这里调用实际的LLM API
        # 简化版实现
        combined = " ".join(texts)
        return combined[:150] + "..." if len(combined) > 150 else combined


class RAPTORTreeBuilder:
    """
    RAPTOR树构建器
    实现递归抽象处理树状检索算法
    """
    
    def __init__(
        self,
        embedding_model: Any,
        summarizer: Optional[BaseSummarizer] = None,
        max_cluster_size: int = 20,
        min_cluster_size: int = 3,
        similarity_threshold: float = 0.5,
        cluster_method: str = 'agglomerative',
        max_levels: int = 5,
        verbose: bool = True
    ):
        """
        初始化RAPTOR树构建器
        
        Args:
            embedding_model: 向量化模型
            summarizer: 摘要生成器
            max_cluster_size: 最大聚类大小
            min_cluster_size: 最小聚类大小
            similarity_threshold: 相似度阈值
            cluster_method: 聚类方法 ('agglomerative', 'kmeans')
            max_levels: 最大树层数
            verbose: 是否输出详细信息
        """
        self.embedding_model = embedding_model
        self.summarizer = summarizer or SimpleSummarizer()
        self.max_cluster_size = max_cluster_size
        self.min_cluster_size = min_cluster_size
        self.similarity_threshold = similarity_threshold
        self.cluster_method = cluster_method
        self.max_levels = max_levels
        self.verbose = verbose
        
        self.root = None
        self.node_counter = 0
        self.all_nodes = []
        self.level_nodes = defaultdict(list)
    
    def _create_node(
        self,
        content: str,
        embedding: Optional[np.ndarray] = None,
        metadata: Optional[Dict[str, Any]] = None
    ) -> RAPTORNode:
        """创建节点"""
        self.node_counter += 1
        node_id = f"node_{self.node_counter:04d}"
        
        if embedding is None:
            embedding = self.embedding_model.encode([content])[0]
        
        return RAPTORNode(
            node_id=node_id,
            content=content,
            embedding=embedding,
            metadata=metadata or {}
        )
    
    def _cluster_nodes(
        self,
        nodes: List[RAPTORNode],
        embeddings: np.ndarray
    ) -> List[List[RAPTORNode]]:
        """
        对节点进行聚类
        
        Args:
            nodes: 节点列表
            embeddings: 节点向量
            
        Returns:
            聚类结果
        """
        n_nodes = len(nodes)
        if n_nodes <= self.min_cluster_size:
            # 节点数量少于最小聚类大小,直接返回一个组
            return [nodes]
        
        # 计算相似度矩阵
        similarity_matrix = cosine_similarity(embeddings)
        distance_matrix = 1 - similarity_matrix
        
        # 使用层次聚类
        if self.cluster_method == 'agglomerative':
            clustering = AgglomerativeClustering(
                n_clusters=None,
                metric='precomputed',
                linkage='average',
                distance_threshold=1 - self.similarity_threshold
            )
            labels = clustering.fit_predict(distance_matrix)
        else:
            # 简单的基于阈值的聚类
            labels = self._threshold_clustering(similarity_matrix)
        
        # 组织聚类结果
        clusters = defaultdict(list)
        for node, label in zip(nodes, labels):
            clusters[label].append(node)
        
        # 过滤太小的聚类
        result_clusters = []
        for cluster_nodes in clusters.values():
            if len(cluster_nodes) >= self.min_cluster_size:
                result_clusters.append(cluster_nodes)
            else:
                # 将小聚类合并到最近的大聚类
                self._merge_small_cluster(cluster_nodes, result_clusters, embeddings)
        
        return result_clusters
    
    def _threshold_clustering(self, similarity_matrix: np.ndarray) -> np.ndarray:
        """基于阈值的简单聚类"""
        n = similarity_matrix.shape[0]
        labels = np.zeros(n, dtype=int)
        cluster_id = 0
        visited = set()
        
        for i in range(n):
            if i not in visited:
                # 开始新的聚类
                cluster = [i]
                visited.add(i)
                
                # 添加相似度高于阈值的节点
                for j in range(n):
                    if j not in visited and similarity_matrix[i][j] > self.similarity_threshold:
                        cluster.append(j)
                        visited.add(j)
                
                # 如果聚类太小,分配随机标签
                if len(cluster) < self.min_cluster_size:
                    for idx in cluster:
                        labels[idx] = -1
                else:
                    for idx in cluster:
                        labels[idx] = cluster_id
                    cluster_id += 1
        
        # 处理未分配的节点
        for i in range(n):
            if labels[i] == -1:
                # 分配到最近的聚类
                max_sim = -1
                best_cluster = 0
                for j in range(n):
                    if labels[j] != -1 and similarity_matrix[i][j] > max_sim:
                        max_sim = similarity_matrix[i][j]
                        best_cluster = labels[j]
                labels[i] = best_cluster if max_sim > 0 else cluster_id
        
        return labels
    
    def _merge_small_cluster(
        self,
        small_cluster: List[RAPTORNode],
        clusters: List[List[RAPTORNode]],
        embeddings: np.ndarray
    ):
        """将小聚类合并到最近的聚类"""
        if not clusters:
            clusters.append(small_cluster)
            return
        
        # 计算小聚类中心
        small_embeddings = np.array([node.embedding for node in small_cluster])
        small_center = np.mean(small_embeddings, axis=0)
        
        # 找到最近的聚类
        best_cluster_idx = 0
        best_similarity = -1
        
        for idx, cluster in enumerate(clusters):
            cluster_embeddings = np.array([node.embedding for node in cluster])
            cluster_center = np.mean(cluster_embeddings, axis=0)
            similarity = cosine_similarity([small_center], [cluster_center])[0][0]
            
            if similarity > best_similarity:
                best_similarity = similarity
                best_cluster_idx = idx
        
        # 合并到最近的聚类
        clusters[best_cluster_idx].extend(small_cluster)
    
    def _generate_summary(
        self,
        nodes: List[RAPTORNode],
        level: int,
        cluster_id: int
    ) -> str:
        """为聚类生成摘要"""
        texts = [node.content for node in nodes]
        summary = self.summarizer.generate_summary(texts)
        
        # 添加层级信息
        summary = f"[Level {level}, Cluster {cluster_id}] {summary}"
        
        return summary
    
    def _build_level(
        self,
        nodes: List[RAPTORNode],
        level: int
    ) -> List[RAPTORNode]:
        """
        构建一层RAPTOR树
        
        Args:
            nodes: 当前层级的节点列表
            level: 当前层级
            
        Returns:
            上层节点列表
        """
        if not nodes:
            return []
        
        if len(nodes) <= self.max_cluster_size:
            # 节点数量小于最大聚类大小,停止递归
            return nodes
        
        # 提取embeddings
        embeddings = np.array([node.embedding for node in nodes])
        
        # 聚类
        clusters = self._cluster_nodes(nodes, embeddings)
        
        if not clusters or len(clusters) <= 1:
            # 无法继续聚类或只有一个聚类
            return nodes
        
        logger.info(f"Level {level}: {len(nodes)} 个节点聚类为 {len(clusters)} 个簇")
        
        # 为每个聚类创建摘要节点
        summary_nodes = []
        for cluster_id, cluster_nodes in enumerate(clusters):
            # 生成摘要
            summary_text = self._generate_summary(cluster_nodes, level, cluster_id)
            summary_embedding = self.embedding_model.encode([summary_text])[0]
            
            # 创建摘要节点
            summary_node = self._create_node(
                content=summary_text,
                embedding=summary_embedding,
                metadata={
                    'level': level,
                    'cluster_id': cluster_id,
                    'children_count': len(cluster_nodes)
                }
            )
            
            # 添加子节点
            for child in cluster_nodes:
                summary_node.add_child(child)
            
            # 存储节点信息
            self.all_nodes.append(summary_node)
            self.level_nodes[level].append(summary_node)
            
            summary_nodes.append(summary_node)
        
        return summary_nodes
    
    def build_tree(self, chunks: List[Dict[str, Any]]) -> RAPTORNode:
        """
        构建RAPTOR树
        
        Args:
            chunks: 文档块列表,每个块包含 'text' 和 'metadata'
            
        Returns:
            根节点
        """
        logger.info(f"开始构建RAPTOR树,文档块数量: {len(chunks)}")
        
        # 1. 创建叶子节点
        leaf_nodes = []
        for chunk in chunks:
            text = chunk.get('text', '')
            if not text:
                continue
            
            embedding = chunk.get('embedding')
            if embedding is None:
                embedding = self.embedding_model.encode([text])[0]
            elif isinstance(embedding, list):
                embedding = np.array(embedding)
            
            metadata = chunk.get('metadata', {})
            metadata['chunk_id'] = chunk.get('chunk_id', '')
            
            leaf_node = self._create_node(text, embedding, metadata)
            leaf_nodes.append(leaf_node)
            self.all_nodes.append(leaf_node)
            self.level_nodes[0].append(leaf_node)
        
        if not leaf_nodes:
            raise ValueError("没有有效的文档块")
        
        logger.info(f"创建了 {len(leaf_nodes)} 个叶子节点")
        
        # 2. 递归构建上层
        current_level = 0
        current_nodes = leaf_nodes
        
        while len(current_nodes) > self.max_cluster_size and current_level < self.max_levels:
            current_level += 1
            current_nodes = self._build_level(current_nodes, current_level)
        
        # 3. 创建根节点
        if len(current_nodes) > 1:
            # 创建根节点
            root_summary = self._generate_summary(
                current_nodes,
                current_level + 1,
                0
            )
            root_embedding = self.embedding_model.encode([root_summary])[0]
            
            self.root = self._create_node(
                content=root_summary,
                embedding=root_embedding,
                metadata={
                    'level': current_level + 1,
                    'is_root': True,
                    'children_count': len(current_nodes)
                }
            )
            
            # 添加子节点
            for node in current_nodes:
                self.root.add_child(node)
            
            self.all_nodes.append(self.root)
            self.level_nodes[current_level + 1].append(self.root)
        else:
            self.root = current_nodes[0] if current_nodes else None
        
        logger.info(f"RAPTOR树构建完成,总共 {len(self.all_nodes)} 个节点,{current_level + 1} 层")
        self._print_tree_statistics()
        
        return self.root
    
    def _print_tree_statistics(self):
        """打印树统计信息"""
        if not self.root:
            return
        
        leaf_count = len(self.root.get_leaf_nodes())
        total_nodes = len(self.all_nodes)
        levels = len(self.level_nodes)
        
        logger.info(f"树统计: 总节点={total_nodes}, 叶子节点={leaf_count}, 层级={levels}")
        
        for level, nodes in sorted(self.level_nodes.items()):
            logger.info(f"  Level {level}: {len(nodes)} 个节点")
    
    def collapsed_tree_retrieval(
        self,
        query: str,
        query_embedding: Optional[np.ndarray] = None,
        top_k: int = 10
    ) -> List[Dict[str, Any]]:
        """
        展开树检索(Collapsed Tree Retrieval)
        将树展开为单层,对所有节点进行向量检索
        
        Args:
            query: 查询文本
            query_embedding: 查询向量(可选)
            top_k: 返回结果数量
            
        Returns:
            检索结果列表
        """
        if not self.all_nodes:
            logger.warning("RAPTOR树为空")
            return []
        
        # 获取查询向量
        if query_embedding is None:
            query_embedding = self.embedding_model.encode_queries([query])[0]
        
        # 获取所有节点的embeddings
        node_embeddings = np.array([node.embedding for node in self.all_nodes])
        node_contents = [node.content for node in self.all_nodes]
        
        # 计算相似度
        similarities = cosine_similarity([query_embedding], node_embeddings)[0]
        
        # 获取top-k索引
        top_indices = np.argsort(similarities)[::-1][:top_k]
        
        # 构建结果
        results = []
        for idx in top_indices:
            node = self.all_nodes[idx]
            results.append({
                'node_id': node.node_id,
                'content': node.content,
                'score': float(similarities[idx]),
                'level': node.level,
                'is_leaf': node.is_leaf,
                'children_count': len(node.children),
                'metadata': node.metadata
            })
        
        return results
    
    def hierarchical_retrieval(
        self,
        query: str,
        query_embedding: Optional[np.ndarray] = None,
        top_k: int = 5,
        expand_children: bool = True
    ) -> List[Dict[str, Any]]:
        """
        层次化检索
        从根节点开始逐层检索
        
        Args:
            query: 查询文本
            query_embedding: 查询向量
            top_k: 每层返回结果数量
            expand_children: 是否展开子节点
            
        Returns:
            检索结果列表
        """
        if not self.root:
            return []
        
        if query_embedding is None:
            query_embedding = self.embedding_model.encode_queries([query])[0]
        
        results = []
        current_nodes = [self.root]
        visited = set()
        
        while current_nodes:
            # 计算当前层节点的相似度
            embeddings = np.array([node.embedding for node in current_nodes])
            similarities = cosine_similarity([query_embedding], embeddings)[0]
            
            # 获取top-k
            top_indices = np.argsort(similarities)[::-1][:top_k]
            
            next_nodes = []
            for idx in top_indices:
                node = current_nodes[idx]
                if node.node_id in visited:
                    continue
                
                visited.add(node.node_id)
                
                # 添加结果
                results.append({
                    'node_id': node.node_id,
                    'content': node.content,
                    'score': float(similarities[idx]),
                    'level': node.level,
                    'is_leaf': node.is_leaf,
                    'children_count': len(node.children),
                    'metadata': node.metadata
                })
                
                # 如果展开子节点,将子节点添加到下一层
                if expand_children and node.children:
                    next_nodes.extend(node.children)
            
            current_nodes = next_nodes
        
        # 按分数排序
        results.sort(key=lambda x: x['score'], reverse=True)
        return results
    
    def visualize_tree(self, max_depth: int = 3) -> str:
        """
        可视化树结构(简化版)
        
        Args:
            max_depth: 最大显示深度
            
        Returns:
            树的文本表示
        """
        if not self.root:
            return "树为空"
        
        def _visualize_node(node: RAPTORNode, depth: int = 0, max_depth: int = 3) -> str:
            if depth > max_depth:
                return "  " * depth + "...\n"
            
            indent = "  " * depth
            node_info = f"{node.node_id} (L{node.level}, {len(node.children)} children)"
            content_preview = node.content[:50] + "..." if len(node.content) > 50 else node.content
            
            result = f"{indent}├─ {node_info}: {content_preview}\n"
            
            if node.children:
                for child in node.children[:3]:  # 限制显示数量
                    result += _visualize_node(child, depth + 1, max_depth)
                if len(node.children) > 3:
                    result += f"{indent}  └─ ... 还有 {len(node.children) - 3} 个子节点\n"
            
            return result
        
        return _visualize_node(self.root, 0, max_depth)


class RAPTORVectorStore:
    """
    基于RAPTOR的向量存储系统
    """
    
    def __init__(self, raptor_builder: RAPTORTreeBuilder):
        self.raptor_builder = raptor_builder
        self.chunks = []
        self.tree = None
    
    def add_documents(self, chunks: List[Dict[str, Any]]):
        """
        添加文档块并构建RAPTOR树
        """
        self.chunks.extend(chunks)
        self.tree = self.raptor_builder.build_tree(self.chunks)
        logger.info(f"RAPTOR向量存储已更新,包含 {len(self.chunks)} 个文档块")
    
    def search(
        self,
        query: str,
        top_k: int = 10,
        method: str = 'collapsed'
    ) -> List[Dict[str, Any]]:
        """
        搜索文档
        
        Args:
            query: 查询文本
            top_k: 返回结果数量
            method: 检索方法 ('collapsed' 或 'hierarchical')
            
        Returns:
            搜索结果
        """
        if method == 'hierarchical':
            return self.raptor_builder.hierarchical_retrieval(query, top_k=top_k)
        else:
            return self.raptor_builder.collapsed_tree_retrieval(query, top_k=top_k)
    
    def get_chunk_by_node_id(self, node_id: str) -> Optional[Dict[str, Any]]:
        """通过节点ID获取原始chunk"""
        for chunk in self.chunks:
            if chunk.get('chunk_id') == node_id:
                return chunk
        return None


def test_raptor_builder():
    """
    测试RAPTOR树构建器
    """
    # 模拟BGE模型
    class MockEmbeddingModel:
        def encode(self, texts):
            if isinstance(texts, str):
                texts = [texts]
            # 随机生成向量
            return np.random.randn(len(texts), 768)
        
        def encode_queries(self, texts):
            return self.encode(texts)
    
    # 创建测试数据
    test_chunks = [
        {"text": f"装备维修文档第{i}段:这是关于装备维修的详细描述内容。", "chunk_id": f"chunk_{i}"}
        for i in range(50)
    ]
    
    # 初始化RAPTOR构建器
    embedding_model = MockEmbeddingModel()
    summarizer = SimpleSummarizer(max_length=100)
    
    raptor_builder = RAPTORTreeBuilder(
        embedding_model=embedding_model,
        summarizer=summarizer,
        max_cluster_size=10,
        min_cluster_size=3,
        similarity_threshold=0.3,
        max_levels=3
    )
    
    # 构建树
    root = raptor_builder.build_tree(test_chunks)
    
    print("RAPTOR树构建完成!")
    print(f"根节点: {root.node_id}")
    print(f"总节点数: {len(raptor_builder.all_nodes)}")
    print(f"叶子节点数: {len(root.get_leaf_nodes())}")
    print(f"层级: {len(raptor_builder.level_nodes)}")
    
    # 测试检索
    query = "装备维修方法"
    results = raptor_builder.collapsed_tree_retrieval(query, top_k=5)
    
    print(f"\n检索结果 (查询: '{query}'):")
    for i, result in enumerate(results):
        print(f"{i+1}. Score: {result['score']:.4f}, Level: {result['level']}, Leaf: {result['is_leaf']}")
        print(f"   Content: {result['content'][:80]}...")
    
    return raptor_builder


if __name__ == "__main__":
    print("="*60)
    print("RAPTOR树构建算法测试")
    print("="*60)
    
    test_raptor_builder()
相关推荐
zlinear数据采集卡2 小时前
数据采集卡从入门到精通(10):采样率与分辨率的核心关系——反比律、架构分布与过采样
arm开发·嵌入式硬件·算法·fpga开发·架构·开源
从小就看凹凸曼^o^3 小时前
Python基础5 - 字典与集合:(1)字典
开发语言·python
GeekZHR3 小时前
C语言指针进阶补充6:动态内存管理、mem系列内存函数、复杂指针声明,一次补齐指针的“三大盲区“
java·c语言·算法·指针
小赵AI手记3 小时前
技术拆解(十七)具身智能:机器人动作生成为何走向Diffusion Policy?
人工智能·笔记·python·机器人
.道阻且长.3 小时前
8.LeetCode算法习题讲解--滑动窗口--长度最小的子数组
算法·leetcode·职场和发展
本地化文档3 小时前
poethepoet-docs-l10n
python·github·gitcode·sphinx
(❁´◡`❁)Jimmy(❁´◡`❁)3 小时前
P1156 [USACO01OPEN] 垃圾陷阱
算法·动态规划
baopixiaoz3 小时前
AI量化策略师|Web3 量化交易研究员
大数据·人工智能·python·区块链
疯狂打码的少年4 小时前
【数据结构】二叉排序树(BST)的定义与操作
数据结构·笔记·算法