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