| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800 |
- """
- RAG 全链路优化系统
- 优化点:
- 1. 语义切块(SemanticChunker)替代固定切片
- 2. 多路召回(BM25 + 向量检索 + 混合检索)
- 3. 查询侧优化(查询重写、查询分解、HyDE)
- 4. Reranker 重排序精排
- 5. Milvus 向量数据库替代 Chroma
- """
- import os
- import re
- import json
- from typing import List, Optional
- from concurrent.futures import ThreadPoolExecutor
- 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 Milvus
- from langchain_community.retrievers import BM25Retriever
- from langchain_classic.retrievers import EnsembleRetriever
- from langchain_classic.retrievers.multi_query import MultiQueryRetriever
- from langchain_core.prompts import ChatPromptTemplate, PromptTemplate
- from langchain_openai import ChatOpenAI
- from langchain_core.output_parsers import StrOutputParser
- from langchain_core.documents import Document
- from langchain_experimental.text_splitter import SemanticChunker
- # 加载环境变量
- load_dotenv()
- # ====================================================================
- # 日志工具
- # ====================================================================
- def log_info(message: str):
- """打印日志信息"""
- print(f"[INFO] {message}")
- def log_error(message: str):
- """打印错误日志"""
- print(f"[ERROR] {message}")
- def log_warn(message: str):
- """打印警告日志"""
- print(f"[WARN] {message}")
- # ====================================================================
- # 查询侧优化模块
- # ====================================================================
- class QueryOptimizer:
- """查询侧优化器:查询重写、查询分解、HyDE"""
- def __init__(self, llm: ChatOpenAI):
- self.llm = llm
- self._init_chains()
- def _init_chains(self):
- """初始化各优化链"""
- # 查询重写 Prompt
- self.rewrite_prompt = PromptTemplate(
- input_variables=["query"],
- template="""你是一个查询优化助手。请将用户的口语化问题改写为更适合信息检索的精确查询。
- 改写要求:
- 1. 补充隐含的上下文信息
- 2. 将口语化表达转为专业表述
- 3. 消除歧义,明确查询意图
- 4. 保持原意不变,不要添加原问题未提及的内容
- 5. 直接输出改写后的查询,不要解释
- 用户问题: {query}
- 改写后的查询:"""
- )
- self.rewrite_chain = self.rewrite_prompt | self.llm | StrOutputParser()
- # 查询分解 Prompt
- self.decompose_prompt = PromptTemplate(
- input_variables=["question"],
- template="""你是一个问题分解助手。请将用户的复杂问题分解为 2-4 个独立的子问题,
- 每个子问题应该能独立检索和回答。
- 要求:
- 1. 子问题之间互不依赖,可以并行检索
- 2. 子问题覆盖原始问题的所有方面
- 3. 每个子问题简洁明确
- 4. 以 JSON 数组格式输出
- 用户问题: {question}
- 输出格式: ["子问题1", "子问题2", "子问题3"]
- 子问题列表:"""
- )
- self.decompose_chain = self.decompose_prompt | self.llm | StrOutputParser()
- # HyDE Prompt(假设性文档嵌入)
- self.hyde_prompt = PromptTemplate(
- input_variables=["question"],
- template="""请根据以下问题,写一段可能出现在相关文档中的内容(约100-200字)。
- 不需要完全准确,只需要包含相关的关键词和表述方式。
- 问题: {question}
- 假设性文档内容:"""
- )
- self.hyde_chain = self.hyde_prompt | self.llm | StrOutputParser()
- def rewrite(self, query: str) -> str:
- """查询重写:将口语化问题转为精确检索查询"""
- try:
- rewritten = self.rewrite_chain.invoke({"query": query})
- log_info(f"查询重写: '{query}' → '{rewritten}'")
- return rewritten
- except Exception as e:
- log_warn(f"查询重写失败: {e},使用原始查询")
- return query
- def decompose(self, question: str) -> List[str]:
- """查询分解:将复杂问题拆分为子问题"""
- try:
- result = self.decompose_chain.invoke({"question": question})
- sub_queries = json.loads(result.strip())
- log_info(f"查询分解: '{question}' → {len(sub_queries)} 个子问题")
- return sub_queries
- except (json.JSONDecodeError, Exception) as e:
- log_warn(f"查询分解失败: {e},使用原始问题")
- return [question]
- def hyde(self, question: str) -> str:
- """HyDE:生成假设性文档用于检索"""
- try:
- hypothetical_doc = self.hyde_chain.invoke({"question": question})
- log_info(f"HyDE 生成假设文档: {hypothetical_doc[:80]}...")
- return hypothetical_doc
- except Exception as e:
- log_warn(f"HyDE 生成失败: {e},使用原始问题")
- return question
- # ====================================================================
- # Reranker 重排序模块
- # ====================================================================
- class LLMReranker:
- """基于 LLM 的重排序器(无需额外模型,用大模型做相关性打分)"""
- def __init__(self, llm: ChatOpenAI):
- self.llm = llm
- self.score_prompt = PromptTemplate(
- input_variables=["query", "document"],
- template="""请评估以下文档与用户问题的相关性,给出 0-10 的分数。
- 只输出数字分数,不要任何解释。
- 用户问题: {query}
- 文档内容: {document}
- 相关性分数(0-10):"""
- )
- self.score_chain = self.score_prompt | self.llm | StrOutputParser()
- def rerank(self, query: str, documents: List[Document], top_k: int = 3) -> List[Document]:
- """对检索结果重排序,返回最相关的 top_k 个文档"""
- if not documents:
- return []
- if len(documents) <= top_k:
- return documents
- log_info(f"Reranker: 对 {len(documents)} 个文档重排序,保留 Top-{top_k}")
- scored_docs = []
- for doc in documents:
- try:
- score_str = self.score_chain.invoke({
- "query": query,
- "document": doc.page_content[:500]
- })
- # 提取数字分数
- score = float(re.search(r'[\d.]+', score_str).group())
- scored_docs.append((doc, score))
- except Exception:
- # 打分失败时给默认分数
- scored_docs.append((doc, 5.0))
- # 按分数降序排序
- scored_docs.sort(key=lambda x: x[1], reverse=True)
- # 返回 top_k 个
- result = [doc for doc, score in scored_docs[:top_k]]
- log_info(f"Reranker 完成,分数范围: {scored_docs[0][1]:.1f} - {scored_docs[-1][1]:.1f}")
- return result
- # ====================================================================
- # 核心 RAG 系统(全链路优化版)
- # ====================================================================
- class OptimizedRAGSystem:
- """
- 全链路优化 RAG 系统
- 优化链路:
- 文档加载 → 数据清洗 → 语义切块 → Milvus存储 → 多路召回 → Reranker精排 → 生成回答
- 查询优化:
- 用户问题 → 查询重写/分解/HyDE → 多路召回 → 重排序 → 生成
- """
- def __init__(
- self,
- pdf_path: str,
- collection_name: str = "rag_optimized",
- milvus_uri: str = "http://localhost:19530",
- use_semantic_chunking: bool = True,
- ):
- """
- 初始化优化版 RAG 系统
- Args:
- pdf_path: PDF 文件路径
- collection_name: Milvus 集合名称
- milvus_uri: Milvus 连接地址
- use_semantic_chunking: 是否使用语义切块(需要较多 API 调用)
- """
- self.pdf_path = pdf_path
- self.collection_name = collection_name
- self.milvus_uri = milvus_uri
- self.use_semantic_chunking = use_semantic_chunking
- self.pages = None
- self.split_docs = None
- self.vectorstore = None
- # 初始化 Embedding 模型
- self.embedding_model = DashScopeEmbeddings(
- model=os.getenv("EMBEDDING_MODEL", "text-embedding-v3"),
- dashscope_api_key=os.getenv("DASHSCOPE_API_KEY", "")
- )
- # 初始化大模型
- 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
- )
- # 初始化查询优化器和重排序器
- self.query_optimizer = QueryOptimizer(self.llm)
- self.reranker = LLMReranker(self.llm)
- log_info("全链路优化 RAG 系统初始化完成")
- log_info(f"模型: {os.getenv('MODEL_NAME', 'deepseek-chat')}")
- log_info(f"Milvus: {milvus_uri}")
- log_info(f"切块策略: {'语义切块' if use_semantic_chunking else '递归字符切块'}")
- # ================================================================
- # 文档处理阶段
- # ================================================================
- def load_pdf(self) -> List[Document]:
- """加载 PDF 文档"""
- log_info(f"加载 PDF: {self.pdf_path}")
- if not os.path.exists(self.pdf_path):
- raise FileNotFoundError(f"PDF 文件不存在: {self.pdf_path}")
- loader = PyMuPDFLoader(self.pdf_path)
- self.pages = loader.load()
- total_content = sum(len(p.page_content) for p in self.pages)
- if total_content == 0:
- raise ValueError("PDF 文件无可提取文本内容,请使用文本型 PDF 或进行 OCR 处理")
- log_info(f"加载完成: {len(self.pages)} 页, {total_content} 字符")
- return self.pages
- def clean_text(self, text: str) -> str:
- """清洗文本:去除噪声、规范化"""
- if not text or not text.strip():
- return ""
- # 删除多余换行(保留最多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)
- return text.strip()
- def clean_all_pages(self) -> List[Document]:
- """清洗所有页面"""
- if not self.pages:
- raise ValueError("请先调用 load_pdf()")
- log_info("清洗文档...")
- for page in self.pages:
- page.page_content = self.clean_text(page.page_content)
- # 过滤空页
- self.pages = [p for p in self.pages if p.page_content.strip()]
- log_info(f"清洗完成: {len(self.pages)} 页有效内容")
- return self.pages
- def split_documents(self) -> List[Document]:
- """
- 文档切块(支持语义切块和递归字符切块两种策略)
- 语义切块:基于相邻句子的语义相似度,在话题转变处切分
- 递归字符切块:按字符数切分,带重叠防止上下文断裂
- """
- if not self.pages:
- raise ValueError("请先调用 load_pdf()")
- if self.use_semantic_chunking:
- log_info("使用语义切块策略(SemanticChunker)...")
- try:
- text_splitter = SemanticChunker(
- embeddings=self.embedding_model,
- breakpoint_threshold_type="percentile",
- # 90 表示只有相似度排名后 10% 的位置才切分
- # 值越大切得越少,块越大
- breakpoint_threshold_amount=90,
- )
- self.split_docs = text_splitter.split_documents(self.pages)
- except Exception as e:
- log_warn(f"语义切块失败: {e},回退到递归字符切块")
- self.split_docs = self._fallback_split()
- else:
- self.split_docs = self._fallback_split()
- # 过滤空块
- self.split_docs = [
- doc for doc in self.split_docs
- if doc.page_content and doc.page_content.strip()
- ]
- log_info(f"切块完成: {len(self.split_docs)} 个文档块")
- if self.split_docs:
- avg_len = sum(len(d.page_content) for d in self.split_docs) / len(self.split_docs)
- log_info(f"平均块长度: {avg_len:.0f} 字符")
- return self.split_docs
- def _fallback_split(self) -> List[Document]:
- """回退策略:递归字符切块"""
- log_info("使用递归字符切块策略...")
- text_splitter = RecursiveCharacterTextSplitter(
- separators=["\n\n", "\n", "。", "!", "?", ";", " ", ""],
- chunk_size=300,
- chunk_overlap=50,
- length_function=len,
- )
- return text_splitter.split_documents(self.pages)
- # ================================================================
- # 向量存储与多路召回
- # ================================================================
- def create_vectorstore(self) -> Milvus:
- """将文档块存入 Milvus 向量数据库"""
- if not self.split_docs:
- raise ValueError("请先调用 split_documents()")
- log_info(f"创建 Milvus 向量库: {self.collection_name}")
- self.vectorstore = Milvus.from_documents(
- documents=self.split_docs,
- embedding=self.embedding_model,
- connection_args={"uri": self.milvus_uri},
- collection_name=self.collection_name,
- drop_old=True, # 开发阶段覆盖旧数据
- )
- log_info(f"Milvus 集合创建完成: {len(self.split_docs)} 个文档块已入库")
- return self.vectorstore
- def load_vectorstore(self) -> Milvus:
- """加载已有的 Milvus 向量数据库"""
- log_info(f"加载 Milvus 集合: {self.collection_name}")
- self.vectorstore = Milvus(
- embedding_function=self.embedding_model,
- connection_args={"uri": self.milvus_uri},
- collection_name=self.collection_name,
- )
- log_info("Milvus 向量库加载完成")
- return self.vectorstore
- def build_ensemble_retriever(self, k: int = 10) -> EnsembleRetriever:
- """
- 构建混合检索器(BM25 + 向量检索)
- BM25: 关键词精确匹配,擅长处理精确术语
- 向量检索: 语义相似度匹配,擅长处理同义词和语义理解
- 混合: 两者加权融合,取长补短
- """
- if not self.vectorstore or not self.split_docs:
- raise ValueError("请先创建向量库")
- # BM25 检索器(关键词匹配)
- bm25_retriever = BM25Retriever.from_documents(self.split_docs)
- bm25_retriever.k = k
- # 向量检索器(语义匹配)
- vector_retriever = self.vectorstore.as_retriever(
- search_kwargs={"k": k}
- )
- # 混合检索器:BM25 权重 0.4,向量检索权重 0.6
- ensemble_retriever = EnsembleRetriever(
- retrievers=[bm25_retriever, vector_retriever],
- weights=[0.4, 0.6],
- )
- log_info(f"混合检索器构建完成 (BM25:0.4 + Vector:0.6, Top-{k})")
- return ensemble_retriever
- def build_multi_query_retriever(
- self, base_retriever: EnsembleRetriever
- ) -> MultiQueryRetriever:
- """
- 构建多查询检索器
- 原理:用 LLM 将一个问题改写为多个不同角度的问题,
- 分别检索后合并去重,增加召回率
- """
- multi_query_prompt = PromptTemplate(
- input_variables=["question"],
- template="""你是一个问题改写助手。请为下面的问题生成 3 个不同的改写版本。
- 每个版本应从不同角度表达相同意图,使用不同的关键词。
- 每个问题单独一行,不要编号。
- 原始问题: {question}
- 改写后的问题:"""
- )
- retriever = MultiQueryRetriever.from_llm(
- retriever=base_retriever,
- llm=self.llm,
- prompt=multi_query_prompt,
- )
- log_info("多查询检索器构建完成")
- return retriever
- # ================================================================
- # 检索与生成
- # ================================================================
- def retrieve_with_optimization(
- self,
- query: str,
- retriever,
- mode: str = "rewrite",
- rerank_top_k: int = 3,
- ) -> List[Document]:
- """
- 带查询优化 + Reranker 的检索流程
- Args:
- query: 用户原始问题
- retriever: 检索器实例
- mode: 查询优化模式
- - "none": 不做优化,直接检索
- - "rewrite": 查询重写后检索
- - "decompose": 查询分解后并行检索
- - "hyde": HyDE 假设文档检索
- - "multi_query": 多查询检索(需要传入 MultiQueryRetriever)
- rerank_top_k: Reranker 精排后保留的文档数
- Returns:
- 精排后的文档列表
- """
- log_info(f"检索模式: {mode} | 问题: {query}")
- # 第一步:查询优化
- if mode == "rewrite":
- optimized_query = self.query_optimizer.rewrite(query)
- raw_docs = retriever.invoke(optimized_query)
- elif mode == "decompose":
- sub_queries = self.query_optimizer.decompose(query)
- raw_docs = self._parallel_retrieve(sub_queries, retriever)
- elif mode == "hyde":
- hypothetical_doc = self.query_optimizer.hyde(query)
- raw_docs = retriever.invoke(hypothetical_doc)
- elif mode == "multi_query":
- # MultiQueryRetriever 内部自带多查询逻辑
- raw_docs = retriever.invoke(query)
- else:
- # mode == "none"
- raw_docs = retriever.invoke(query)
- log_info(f"初步召回: {len(raw_docs)} 个文档")
- # 第二步:Reranker 精排
- if len(raw_docs) > rerank_top_k:
- reranked_docs = self.reranker.rerank(query, raw_docs, top_k=rerank_top_k)
- else:
- reranked_docs = raw_docs
- log_info(f"精排后: {len(reranked_docs)} 个文档")
- return reranked_docs
- def _parallel_retrieve(
- self, queries: List[str], retriever
- ) -> List[Document]:
- """并行检索多个查询并合并去重"""
- all_docs = []
- seen_contents = set()
- def retrieve_one(q):
- return retriever.invoke(q)
- with ThreadPoolExecutor(max_workers=min(len(queries), 4)) as executor:
- futures = [executor.submit(retrieve_one, q) for q in queries]
- for future in futures:
- try:
- docs = future.result()
- for doc in docs:
- content_hash = hash(doc.page_content)
- if content_hash not in seen_contents:
- seen_contents.add(content_hash)
- all_docs.append(doc)
- except Exception as e:
- log_warn(f"并行检索子任务失败: {e}")
- return all_docs
- def generate_answer(self, query: str, relevant_docs: List[Document]) -> str:
- """基于检索结果生成回答"""
- if not relevant_docs:
- return "抱歉,根据现有资料,我找不到与您问题相关的信息。"
- log_info("生成回答...")
- context = "\n\n---\n\n".join([doc.page_content for doc in relevant_docs])
- prompt = ChatPromptTemplate.from_template("""你是一个专业的知识库助手。请根据以下检索到的上下文回答用户问题。
- **规则:**
- - 只基于提供的上下文回答,不要编造信息
- - 如果上下文中没有相关信息,请明确说明
- - 回答要简洁直接,条理清晰
- - 引用原文时用引号标注
- **检索到的上下文:**
- {context}
- **用户问题:**
- {question}
- **回答:**""")
- chain = prompt | self.llm | StrOutputParser()
- answer = chain.invoke({"context": context, "question": query})
- log_info("回答生成完成")
- return answer
- def query(
- self,
- question: str,
- mode: str = "rewrite",
- rerank_top_k: int = 3,
- ) -> dict:
- """
- 完整查询流程:查询优化 → 多路召回 → 重排序 → 生成回答
- Args:
- question: 用户问题
- mode: 查询优化模式 ("none"/"rewrite"/"decompose"/"hyde"/"multi_query")
- rerank_top_k: 精排后保留文档数
- Returns:
- 包含回答和中间结果的字典
- """
- # 构建检索器
- ensemble_retriever = self.build_ensemble_retriever(k=10)
- # 如果是 multi_query 模式,额外包装一层
- if mode == "multi_query":
- retriever = self.build_multi_query_retriever(ensemble_retriever)
- else:
- retriever = ensemble_retriever
- # 检索 + 重排序
- relevant_docs = self.retrieve_with_optimization(
- query=question,
- retriever=retriever,
- mode=mode,
- rerank_top_k=rerank_top_k,
- )
- # 生成回答
- answer = self.generate_answer(question, relevant_docs)
- return {
- "question": question,
- "mode": mode,
- "doc_count": len(relevant_docs),
- "answer": answer,
- "sources": [doc.page_content[:100] for doc in relevant_docs],
- }
- # ================================================================
- # 构建与对比
- # ================================================================
- def build_knowledge_base(self):
- """构建完整知识库:加载 → 清洗 → 切块 → 向量化"""
- log_info("=" * 60)
- log_info("开始构建优化版知识库...")
- log_info("=" * 60)
- self.load_pdf()
- self.clean_all_pages()
- self.split_documents()
- self.create_vectorstore()
- log_info("=" * 60)
- log_info("知识库构建完成!")
- log_info("=" * 60)
- def compare_retrieval_modes(self, question: str):
- """对比不同检索模式的效果"""
- log_info("=" * 60)
- log_info(f"对比测试: {question}")
- log_info("=" * 60)
- modes = ["none", "rewrite", "decompose", "hyde"]
- results = {}
- for mode in modes:
- log_info(f"\n--- 模式: {mode} ---")
- try:
- result = self.query(question, mode=mode)
- results[mode] = result
- print(f"\n[{mode}] 召回 {result['doc_count']} 个文档")
- print(f"回答: {result['answer'][:200]}...")
- except Exception as e:
- log_error(f"模式 {mode} 执行失败: {e}")
- results[mode] = {"error": str(e)}
- return results
- # ====================================================================
- # 交互式问答
- # ====================================================================
- def interactive_query(rag_system: OptimizedRAGSystem):
- """交互式问答模式"""
- print("\n" + "=" * 60)
- print("RAG 全链路优化系统 - 交互式问答")
- print("=" * 60)
- print("检索模式: none / rewrite / decompose / hyde / multi_query")
- print("命令: 'mode <模式名>' 切换模式, 'compare' 对比模式, 'quit' 退出")
- print()
- current_mode = "rewrite"
- while True:
- try:
- user_input = input(f"[{current_mode}] 你的问题: ").strip()
- if user_input.lower() in ['quit', 'exit', 'q']:
- print("\n再见!")
- break
- if not user_input:
- continue
- # 切换模式命令
- if user_input.startswith("mode "):
- new_mode = user_input[5:].strip()
- valid_modes = ["none", "rewrite", "decompose", "hyde", "multi_query"]
- if new_mode in valid_modes:
- current_mode = new_mode
- print(f"已切换到模式: {current_mode}")
- else:
- print(f"无效模式。可选: {valid_modes}")
- continue
- # 对比命令
- if user_input == "compare":
- question = input("输入对比问题: ").strip()
- if question:
- rag_system.compare_retrieval_modes(question)
- continue
- # 正常查询
- print("\n检索并生成中...\n")
- result = rag_system.query(user_input, mode=current_mode)
- print("-" * 60)
- print(f"回答:\n{result['answer']}")
- print(f"\n[召回文档数: {result['doc_count']}]")
- print("-" * 60)
- print()
- except KeyboardInterrupt:
- print("\n\n再见!")
- break
- except Exception as e:
- log_error(f"查询出错: {e}")
- # ====================================================================
- # 主函数
- # ====================================================================
- def main():
- """主函数"""
- # PDF 文件路径
- pdf_path = r"./car_info.pdf"
- # Milvus 配置
- milvus_uri = "http://localhost:19530"
- collection_name = "car_info_optimized"
- # 检查 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-api-key")
- return
- if not dashscope_api_key:
- print("错误: 未设置 DASHSCOPE_API_KEY 环境变量")
- print("请在 .env 文件中配置: DASHSCOPE_API_KEY=your-api-key")
- return
- # 检查 PDF 文件
- if not os.path.exists(pdf_path):
- print(f"错误: PDF 文件不存在: {pdf_path}")
- return
- try:
- # 创建优化版 RAG 系统
- rag = OptimizedRAGSystem(
- pdf_path=pdf_path,
- collection_name=collection_name,
- milvus_uri=milvus_uri,
- # 语义切块需要较多 Embedding API 调用
- # 如果 API 调用有限制,可以设为 False 使用递归字符切块
- use_semantic_chunking=True,
- )
- # 构建知识库
- print("\n是否重新构建知识库?(y/n,默认 n 加载已有数据): ", end="")
- choice = input().strip().lower()
- if choice == 'y':
- rag.build_knowledge_base()
- else:
- try:
- rag.load_vectorstore()
- # 加载文档用于 BM25 检索器
- rag.load_pdf()
- rag.clean_all_pages()
- rag.split_documents()
- log_info("已加载现有向量库和文档")
- except Exception as e:
- log_warn(f"加载失败: {e},重新构建...")
- rag.build_knowledge_base()
- # 进入交互式问答
- interactive_query(rag)
- except Exception as e:
- log_error(f"系统运行出错: {e}")
- import traceback
- traceback.print_exc()
- if __name__ == "__main__":
- main()
|