沐沐 1 miesiąc temu
rodzic
commit
806663507d
1 zmienionych plików z 800 dodań i 0 usunięć
  1. 800 0
      02_RAG_study/02-2.RAG_optimize_task.py

+ 800 - 0
02_RAG_study/02-2.RAG_optimize_task.py

@@ -0,0 +1,800 @@
+"""
+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()