Bladeren bron

初步搭建RAG系统

沐沐 1 maand geleden
bovenliggende
commit
826d65b14b
1 gewijzigde bestanden met toevoegingen van 543 en 0 verwijderingen
  1. 543 0
      02_RAG_study/02_RAG_task.py

+ 543 - 0
02_RAG_study/02_RAG_task.py

@@ -0,0 +1,543 @@
+"""
+RAG系统实战 - 投资方法论知识库
+功能:加载PDF文档、数据清洗、分块、向量化、存储到向量库、检索生成回答
+"""
+
+import os
+import re
+from typing import List
+from dotenv import load_dotenv
+
+# LangChain 核心组件
+from langchain_community.document_loaders import PyMuPDFLoader
+from langchain_text_splitters import RecursiveCharacterTextSplitter
+from langchain_community.embeddings import DashScopeEmbeddings
+from langchain_community.vectorstores import Chroma
+from langchain_core.prompts import ChatPromptTemplate
+from langchain_openai import ChatOpenAI
+from langchain_core.output_parsers import StrOutputParser
+from langchain_core.documents import Document
+
+# 加载环境变量
+load_dotenv()
+
+
+class RAGSystem:
+    """RAG系统类,封装完整的检索增强生成流程"""
+    
+    def __init__(self, pdf_path: str, persist_directory: str = "./knowledge_db"):
+        """
+        初始化RAG系统
+        
+        Args:
+            pdf_path: PDF文件路径
+            persist_directory: 向量数据库持久化目录
+        """
+        self.pdf_path = pdf_path
+        self.persist_directory = persist_directory
+        self.pages = None
+        self.split_docs = None
+        self.vectorstore = None
+        self.retriever = None
+        
+        # 初始化Embedding模型(使用阿里通义)
+        self.embedding_model = DashScopeEmbeddings(
+            model=os.getenv("EMBEDDING_MODEL", "text-embedding-v3"),
+            dashscope_api_key=os.getenv("DASHSCOPE_API_KEY", "")
+        )
+        
+        # 初始化大模型(使用DeepSeek)
+        self.llm = ChatOpenAI(
+            model=os.getenv("MODEL_NAME", "deepseek-chat"),
+            openai_api_key=os.getenv("OPENAI_API_KEY", ""),
+            openai_api_base=os.getenv("OPENAI_API_BASE", "https://api.deepseek.com/v1"),
+            temperature=0.7
+        )
+        
+        log_info("RAG系统初始化完成")
+        log_info(f"使用模型: {os.getenv('MODEL_NAME', 'deepseek-chat')}")
+        log_info(f"API地址: {os.getenv('OPENAI_API_BASE', 'https://api.deepseek.com/v1')}")
+    
+    def load_pdf(self) -> List[Document]:
+        """
+        步骤1:加载PDF文档
+        
+        Returns:
+            文档页面列表
+        """
+        log_info(f"正在加载PDF文档: {self.pdf_path}")
+        
+        # 检查文件是否存在
+        if not os.path.exists(self.pdf_path):
+            raise FileNotFoundError(f"PDF文件不存在: {self.pdf_path}")
+        
+        # 使用PyMuPDF加载PDF
+        loader = PyMuPDFLoader(self.pdf_path)
+        self.pages = loader.load()
+        
+        log_info(f"PDF加载完成,共 {len(self.pages)} 页")
+        
+        # 检查是否有实际内容
+        total_content = sum(len(page.page_content) for page in self.pages)
+        
+        if total_content == 0:
+            log_error("PDF文件没有可提取的文本内容!")
+            log_error("可能原因:")
+            log_error("  1. PDF是扫描版或图片型PDF")
+            log_error("  2. PDF有加密保护")
+            log_error("  3. PDF格式特殊")
+            log_error("\n解决方案:")
+            log_error("  1. 使用OCR工具将PDF转换为文本")
+            log_error("  2. 尝试其他PDF文件")
+            raise ValueError("PDF文件无可提取文本内容,请使用文本型PDF或进行OCR处理")
+        
+        # 打印统计信息
+        non_empty_pages = sum(1 for page in self.pages if len(page.page_content) > 0)
+        log_info(f"有内容的页数: {non_empty_pages}/{len(self.pages)}")
+        log_info(f"总字符数: {total_content}")
+        
+        # 打印第一页预览
+        if self.pages and len(self.pages[0].page_content) > 0:
+            preview = self.pages[0].page_content[:500]
+            log_info(f"第一页内容预览:\n{preview}...")
+        
+        return self.pages
+    
+    def clean_text(self, text: str) -> str:
+        """
+        步骤2:清洗PDF解析出的文本,去除常见噪声
+        
+        Args:
+            text: 原始文本
+            
+        Returns:
+            清洗后的文本
+        """
+        if not text or not text.strip():
+            return ""
+        
+        # 保存原始长度用于对比
+        original_len = len(text)
+        
+        # 删除连续的换行符(保留最多2个)
+        text = re.sub(r'\n{3,}', '\n\n', text)
+        
+        # 删除项目符号
+        text = text.replace('•', '').replace('·', '')
+        
+        # 合并多余空格(但不删除换行符周围的空格)
+        text = re.sub(r'[^\S\n]+', ' ', text)
+        
+        # 删除特殊控制字符
+        text = re.sub(r'[\x00-\x08\x0b\x0c\x0e-\x1f\x7f-\x9f]', '', text)
+        
+        # 如果清洗后文本为空或太短,返回原始文本
+        cleaned_text = text.strip()
+        if len(cleaned_text) < 10:  # 如果清洗后少于10个字符,返回原始文本
+            log_info(f"警告: 清洗后文本过短({len(cleaned_text)}字符),保留原始文本({original_len}字符)")
+            return text.strip()
+        
+        return cleaned_text
+    
+    def clean_all_pages(self) -> List[Document]:
+        """
+        清洗所有页面
+        
+        Returns:
+            清洗后的文档列表
+        """
+        if not self.pages:
+            raise ValueError("请先调用 load_pdf() 加载文档")
+        
+        log_info("开始清洗文档...")
+        
+        # 清洗每一页的内容
+        for page in self.pages:
+            page.page_content = self.clean_text(page.page_content)
+        
+        log_info(f"文档清洗完成,共处理 {len(self.pages)} 页")
+        
+        return self.pages
+    
+    def split_documents(
+        self, 
+        chunk_size: int = 100, 
+        chunk_overlap: int = 20
+    ) -> List[Document]:
+        """
+        步骤3:文档分块(使用递归字符分割器)
+        
+        Args:
+            chunk_size: 每个块的最大字符数
+            chunk_overlap: 相邻块之间的重叠字符数
+            
+        Returns:
+            分块后的文档列表
+        """
+        if not self.pages:
+            raise ValueError("请先调用 load_pdf() 加载文档")
+        
+        log_info(f"开始分块,chunk_size={chunk_size}, chunk_overlap={chunk_overlap}")
+        
+        # 创建递归字符分割器
+        text_splitter = RecursiveCharacterTextSplitter(
+            # 分隔符优先级:段落 → 换行 → 句号 → 空格 → 硬切
+            separators=["\n\n", "\n", "。", "!", "?", ";", " ", ""],
+            chunk_size=chunk_size,
+            chunk_overlap=chunk_overlap,
+            length_function=len,
+            is_separator_regex=False
+        )
+        
+        # 执行分块
+        self.split_docs = text_splitter.split_documents(self.pages)
+        
+        # 过滤空块
+        self.split_docs = [
+            doc for doc in self.split_docs 
+            if doc.page_content and doc.page_content.strip()
+        ]
+        
+        # 统计信息
+        total_chars = sum(len(doc.page_content) for doc in self.split_docs)
+        log_info(f"分块完成,共 {len(self.split_docs)} 个块,总字符数: {total_chars}")
+        
+        # 打印前3个块预览
+        for i, doc in enumerate(self.split_docs[:3]):
+            log_info(f"块{i+1} 预览: {doc.page_content[:100]}...")
+        
+        return self.split_docs
+    
+    def create_vectorstore(self) -> Chroma:
+        """
+        步骤4:创建向量数据库并存储文档
+        
+        Returns:
+            向量数据库实例
+        """
+        if not self.split_docs:
+            raise ValueError("请先调用 split_documents() 分块文档")
+        
+        log_info("开始创建向量数据库...")
+        
+        # 创建向量数据库并持久化
+        self.vectorstore = Chroma.from_documents(
+            documents=self.split_docs,
+            embedding=self.embedding_model,
+            collection_metadata={"hnsw:space": "cosine"},  # 使用余弦相似度
+            persist_directory=self.persist_directory
+        )
+        
+        log_info(f"向量数据库创建完成,保存至: {self.persist_directory}")
+        
+        return self.vectorstore
+    
+    def load_vectorstore(self) -> Chroma:
+        """
+        加载已存在的向量数据库
+        
+        Returns:
+            向量数据库实例
+        """
+        if not os.path.exists(self.persist_directory):
+            raise FileNotFoundError(f"向量数据库不存在: {self.persist_directory}")
+        
+        log_info(f"加载向量数据库: {self.persist_directory}")
+        
+        self.vectorstore = Chroma(
+            persist_directory=self.persist_directory,
+            embedding_function=self.embedding_model
+        )
+        
+        log_info("向量数据库加载完成")
+        
+        return self.vectorstore
+    
+    def create_retriever(self, k: int = 3):
+        """
+        步骤5:创建检索器
+        
+        Args:
+            k: 返回的最相关文档数量
+            
+        Returns:
+            检索器实例
+        """
+        if not self.vectorstore:
+            raise ValueError("请先创建或加载向量数据库")
+        
+        log_info(f"创建检索器,返回Top-{k}文档")
+        
+        self.retriever = self.vectorstore.as_retriever(
+            search_kwargs={"k": k}
+        )
+        
+        return self.retriever
+    
+    def retrieve(self, query: str, k: int = 3) -> List[Document]:
+        """
+        步骤6:检索相关文档
+        
+        Args:
+            query: 用户问题
+            k: 返回的文档数量
+            
+        Returns:
+            相关文档列表
+        """
+        if not self.retriever:
+            raise ValueError("请先创建检索器")
+        
+        log_info(f"检索问题: {query}")
+        
+        # 执行检索
+        relevant_docs = self.retriever.invoke(query)
+        
+        log_info(f"检索完成,找到 {len(relevant_docs)} 个相关文档")
+        
+        # 打印检索结果
+        for i, doc in enumerate(relevant_docs):
+            log_info(f"文档{i+1}: {doc.page_content[:150]}...")
+        
+        return relevant_docs
+    
+    def generate_answer(self, query: str, relevant_docs: List[Document]) -> str:
+        """
+        步骤7:基于检索结果生成回答
+        
+        Args:
+            query: 用户问题
+            relevant_docs: 相关文档列表
+            
+        Returns:
+            生成的回答
+        """
+        log_info("开始生成回答...")
+        
+        # 拼接上下文
+        context = "\n\n---\n\n".join([doc.page_content for doc in relevant_docs])
+        
+        # 构建Prompt模板
+        prompt = ChatPromptTemplate.from_template("""
+你是一个专业的投资知识库助手。请根据以下检索到的上下文回答用户问题。
+
+**规则:**
+- 只基于提供的上下文回答,不要编造
+- 如果上下文中没有相关信息,直接说「根据现有资料,我找不到这个问题的答案」
+- 回答要简洁直接,引用原文时用引号
+- 回答时请标注信息来源的页码
+
+**检索到的上下文:**
+{context}
+
+**用户问题:**
+{question}
+""")
+        
+        # 构建Chain
+        chain = prompt | self.llm | StrOutputParser()
+        
+        # 生成回答
+        answer = chain.invoke({
+            "context": context,
+            "question": query
+        })
+        
+        log_info("回答生成完成")
+        
+        return answer
+    
+    def query(self, question: str, k: int = 3) -> str:
+        """
+        完整查询流程:检索 + 生成
+        
+        Args:
+            question: 用户问题
+            k: 检索文档数量
+            
+        Returns:
+            生成的回答
+        """
+        # 检索相关文档
+        relevant_docs = self.retrieve(question, k)
+        
+        # 生成回答
+        answer = self.generate_answer(question, relevant_docs)
+        
+        return answer
+    
+    def build_knowledge_base(self):
+        """
+        构建完整的知识库:加载 → 清洗 → 分块 → 向量化 → 存储
+        """
+        log_info("=" * 50)
+        log_info("开始构建知识库...")
+        log_info("=" * 50)
+        
+        # 步骤1:加载PDF
+        self.load_pdf()
+        
+        # 步骤2:清洗数据
+        self.clean_all_pages()
+        
+        # 步骤3:分块
+        self.split_documents()
+        
+        # 步骤4:向量化并存储
+        self.create_vectorstore()
+        
+        # 步骤5:创建检索器
+        self.create_retriever()
+        
+        log_info("=" * 50)
+        log_info("知识库构建完成!")
+        log_info("=" * 50)
+
+
+def log_info(message: str):
+    """打印日志信息"""
+    print(f"[INFO] {message}")
+
+
+def interactive_query(rag_system: RAGSystem):
+    """
+    交互式问答模式
+    
+    Args:
+        rag_system: RAG系统实例
+    """
+    print("\n" + "=" * 60)
+    print("RAG知识库问答系统 - 投资方法论")
+    print("=" * 60)
+    print("输入问题开始查询,输入 'quit' 或 'exit' 退出\n")
+    
+    while True:
+        try:
+            # 获取用户输入
+            question = input("你的问题: ").strip()
+            
+            # 检查退出命令
+            if question.lower() in ['quit', 'exit', 'q']:
+                print("\n感谢使用,再见!")
+                break
+            
+            # 跳过空问题
+            if not question:
+                print("请输入有效问题\n")
+                continue
+            
+            # 执行查询
+            print("\n正在检索并生成答案...\n")
+            answer = rag_system.query(question)
+            
+            # 显示结果
+            print("-" * 60)
+            print(f"回答:\n{answer}")
+            print("-" * 60)
+            print()
+            
+        except KeyboardInterrupt:
+            print("\n\n感谢使用,再见!")
+            break
+        except Exception as e:
+            log_error(f"查询出错: {str(e)}")
+
+
+def log_error(message: str):
+    """打印错误日志"""
+    print(f"[ERROR] {message}")
+
+
+def main():
+    """主函数"""
+    # PDF文件路径
+    pdf_path = r"D:\investment\data\疯狂的里海_投资体系框架.pdf"
+    
+    # 备用PDF文件路径(用于测试)
+    backup_pdf_path = r"./car_info.pdf"
+    
+    # 向量数据库存储路径
+    persist_directory = "D:\agentlearning\lqq-agent-study\investment_db"
+    
+    # 检查API Key配置
+    openai_api_key = os.getenv("OPENAI_API_KEY")
+    dashscope_api_key = os.getenv("DASHSCOPE_API_KEY")
+    
+    if not openai_api_key:
+        print("警告: 未设置 OPENAI_API_KEY 环境变量")
+        print("请在 .env 文件中添加: OPENAI_API_KEY=your-deepseek-api-key")
+        print("或者访问 https://platform.deepseek.com/ 获取API Key")
+        return
+    
+    if not dashscope_api_key:
+        print("警告: 未设置 DASHSCOPE_API_KEY 环境变量")
+        print("请在 .env 文件中添加: DASHSCOPE_API_KEY=your-dashscope-api-key")
+        print("或者访问 https://dashscope.console.aliyun.com/ 获取API Key")
+        return
+    
+    try:
+        # 检查主PDF文件是否存在
+        if not os.path.exists(pdf_path):
+            print(f"警告: 指定的PDF文件不存在: {pdf_path}")
+            
+            # 尝试使用备用PDF文件
+            if os.path.exists(backup_pdf_path):
+                print(f"将使用备用PDF文件进行测试: {backup_pdf_path}")
+                pdf_path = backup_pdf_path
+                persist_directory = "./car_info_knowledge_db"
+            else:
+                print("备用PDF文件也不存在,请检查文件路径")
+                return
+        
+        # 创建RAG系统实例
+        rag = RAGSystem(
+            pdf_path=pdf_path,
+            persist_directory=persist_directory
+        )
+        
+        # 检查向量数据库是否已存在
+        if os.path.exists(persist_directory):
+            print(f"检测到已有向量数据库: {persist_directory}")
+            choice = input("是否重新构建知识库?(y/n): ").strip().lower()
+            
+            if choice == 'y':
+                # 重新构建知识库
+                rag.build_knowledge_base()
+            else:
+                # 加载已有向量数据库
+                rag.load_vectorstore()
+                rag.create_retriever()
+        else:
+            # 构建新知识库
+            rag.build_knowledge_base()
+        
+        # 进入交互式问答
+        interactive_query(rag)
+        
+    except ValueError as e:
+        if "PDF文件无可提取文本内容" in str(e):
+            print("\n" + "=" * 60)
+            print("PDF文件处理失败")
+            print("=" * 60)
+            print("\n您的PDF文件是扫描版或图片型PDF,无法直接提取文本。")
+            print("\n解决方案:")
+            print("1. 使用OCR工具将PDF转换为文本")
+            print("   推荐工具:Adobe Acrobat、福昕PDF、ABBYY FineReader")
+            print("   在线工具:https://www.pdf2go.com/zh/ocr-pdf")
+            print("\n2. 使用项目中的示例PDF文件进行测试:")
+            print(f"   文件路径: {backup_pdf_path}")
+            print("\n3. 寻找其他文本型PDF文件")
+            print("\n提示:您可以修改代码中的pdf_path变量指向可用的PDF文件")
+        else:
+            log_error(str(e))
+    except FileNotFoundError as e:
+        log_error(str(e))
+        log_error("请检查PDF文件路径是否正确")
+    except Exception as e:
+        log_error(f"系统运行出错: {str(e)}")
+        import traceback
+        traceback.print_exc()
+
+
+if __name__ == "__main__":
+    main()