""" 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()