瀏覽代碼

第三次作业

ZhouYI 1 月之前
父節點
當前提交
926c8c5657

+ 15 - 0
03_RAG优化/agent/chain.py

@@ -0,0 +1,15 @@
+from langchain.agents import create_agent
+from .llm import create_llm
+from .config import AppConfig
+
+
+def create_chat_agent(config:AppConfig):
+    llm = create_llm(config)
+    agent = create_agent(
+        llm=llm,
+        tools=[],
+        system_prompt="""
+        你是一个耐心、清晰的 Python 和 Agent 开发学习助手。
+        """
+    )
+    return agent

+ 22 - 0
03_RAG优化/agent/cli.py

@@ -0,0 +1,22 @@
+from .chain import create_chat_agent
+from .config import load_config
+
+
+def run_chat() -> None:
+    config = load_config()
+    agent = create_chat_agent(config)
+
+
+    while True:
+        question = input("请输入: ").strip()
+        if not question:
+            continue
+
+        if question.lower() in ["exit", "quit", "q"]:
+            break
+
+        result = agent.invoke(
+            {"messages": [{"role": "user", "content": question}]}
+        )
+        answer = result["messages"][-1].content
+        print(f"AI: {answer}")

+ 29 - 0
03_RAG优化/agent/config.py

@@ -0,0 +1,29 @@
+import os
+from dataclasses import dataclass
+
+
+@dataclass(frozen=True)
+class AppConfig:
+    api_key: str
+    api_url: str
+    model: str
+    database_url: str
+    history_table_name: str
+    default_session_id: str
+
+def load_config() -> AppConfig:
+    api_key = os.getenv("API_KEY")
+    if not api_key:
+        raise ValueError("API_KEY is not set in environment variables.")
+
+    return AppConfig(
+        api_key=api_key,
+        api_url=os.getenv(
+            "API_URL",
+            "https://dashscope.aliyuncs.com/compatible-mode/v1",
+        ),
+        model=os.getenv("MODEL", "qwen3.6-plus"),
+        database_url=os.getenv("DATABASE_URL", "mysql+pymysql://root:1234@localhost:3306/mydb"),
+        history_table_name=os.getenv("HISTORY_TABLE_NAME", "chat_history"),
+        default_session_id=os.getenv("SESSION_ID", "user_001"),
+    )

+ 12 - 0
03_RAG优化/agent/llm.py

@@ -0,0 +1,12 @@
+from langchain_openai import ChatOpenAI
+
+from .config import AppConfig
+
+
+def create_llm(config: AppConfig) -> ChatOpenAI:
+    return ChatOpenAI(
+        model_name=config.model,
+        api_key=config.api_key,
+        base_url=config.api_url,
+        temperature=0.7,
+    )

+ 18 - 0
03_RAG优化/agent/rag/document.py

@@ -0,0 +1,18 @@
+"""文档加载和文本清洗。"""
+
+import re
+
+from langchain_community.document_loaders import PyMuPDFLoader
+from langchain_core.documents import Document
+
+
+def load_pdf_document(pdf_path: str) -> list[Document]:
+    """按页加载 PDF 文档。"""
+    return PyMuPDFLoader(pdf_path).load()
+
+
+def clean_pdf_text(text: str) -> str:
+    """规整 PDF 提取文本中的空白字符。"""
+    text = re.sub(r"[\t\r\f\v ]+", " ", text)
+    text = re.sub(r"\n{3,}", "\n\n", text)
+    return text.strip()

+ 14 - 0
03_RAG优化/agent/rag/embedding.py

@@ -0,0 +1,14 @@
+"""Embedding 模型工厂。"""
+
+from langchain_community.embeddings import DashScopeEmbeddings
+
+from agent.config import load_config
+
+
+def get_embedding_model() -> DashScopeEmbeddings:
+    """创建与 Milvus 1024 维集合配套的 DashScope Embedding 模型。"""
+    config = load_config()
+    return DashScopeEmbeddings(
+        model="text-embedding-v4",
+        dashscope_api_key=config.api_key,
+    )

+ 65 - 0
03_RAG优化/agent/rag/multi_retriever.py

@@ -0,0 +1,65 @@
+"""BM25、Milvus 向量和混合检索。"""
+
+from langchain_classic.retrievers import EnsembleRetriever
+from langchain_community.retrievers import BM25Retriever
+from langchain_core.documents import Document
+from langchain_milvus import Milvus
+
+
+class MultiRetrieverSystem:
+    """多路召回检索系统,默认使用混合检索。"""
+
+    def __init__(self, documents, embedding_model, llm):
+        self.documents = documents
+        self.embedding_model = embedding_model
+        self.llm = llm
+        self.setup_retrievers()
+
+    def setup_retrievers(self):
+        """初始化 BM25、向量和混合检索器。"""
+        self.bm25 = BM25Retriever.from_documents(self.documents)
+        self.bm25.k = 10
+
+        self.vectorstore = Milvus(
+            embedding_function=self.embedding_model,
+            connection_args={"uri": "http://localhost:19530"},
+            collection_name="car_info_collection",
+            primary_field="id",
+            text_field="content",
+            vector_field="embedding",
+            auto_id=True,
+            search_params={"metric_type": "COSINE", "params": {}},
+        )
+        self._index_documents_if_empty()
+        self.vector = self.vectorstore.as_retriever(search_kwargs={"k": 10})
+
+        # LangChain 当前版本使用 RRF 融合排序,不需要 normalize_scores 参数。
+        self.ensemble = EnsembleRetriever(
+            retrievers=[self.bm25, self.vector],
+            weights=[0.4, 0.6],
+        )
+
+    def _index_documents_if_empty(self):
+        """Milvus 集合为空时才写入分块,避免重复入库。"""
+        rows = self.vectorstore.client.query(
+            collection_name="car_info_collection",
+            filter="id >= 0",
+            output_fields=["id"],
+            limit=1,
+            consistency_level="Strong",
+        )
+        if not rows:
+            self.vectorstore.add_documents(self.documents)
+
+    def search(self, query, mode="ensemble") -> list[Document]:
+        """执行检索,mode 可选 bm25、vector、ensemble。"""
+        retriever_map = {
+            "bm25": self.bm25,
+            "vector": self.vector,
+            "ensemble": self.ensemble,
+        }
+        if mode not in retriever_map:
+            choices = ", ".join(retriever_map)
+            raise ValueError(f"不支持的检索模式: {mode},可选: {choices}")
+
+        return retriever_map[mode].invoke(query)[:10]

+ 68 - 0
03_RAG优化/agent/rag/query_optimization.py

@@ -0,0 +1,68 @@
+from langchain_core.prompts import ChatPromptTemplate, PromptTemplate
+
+def get_prompt_template(prompt_type:str) -> ChatPromptTemplate:
+    """返回 RAG 问答链的提示模板。"""
+
+    rewrite_prompt = PromptTemplate(
+input_variables=["query"],
+template="""你是⼀个查询优化助⼿。请将⽤户的⼝语化问题改写为更适合信息检索的
+精确查询。
+改写要求:
+1. 补充隐含的上下⽂信息
+2. 将⼝语化表达转为专业表述
+在 RAG 链路中集成
+将查询重写作为检索前的预处理步骤:
+3. 消除歧义,明确查询意图
+4. 保持原意不变,不要添加原问题未提及的内容
+5. 直接输出改写后的查询,不要解释
+⽤户问题: {query}
+改写后的查询:"""
+)
+
+    decompose_prompt = PromptTemplate(
+input_variables=["question"],
+template="""你是⼀个问题分解助⼿。请将⽤户的复杂问题分解为 2-4 个独⽴的⼦问
+题,
+每个⼦问题应该能独⽴检索和回答。
+要求:
+1. ⼦问题之间互不依赖,可以并⾏检索
+2. ⼦问题覆盖原始问题的所有⽅⾯
+3. 每个⼦问题简洁明确
+4. 以 JSON 数组格式输出
+⽤户问题: {question}
+输出格式示例: ["⼦问题1", "⼦问题2", "⼦问题3"]
+⼦问题列表:"""
+)
+
+    clarify_prompt = PromptTemplate(
+input_variables=["query"],
+template="""你是⼀个智能客服助⼿。请判断⽤户的提问是否包含⾜够的信息来进⾏准
+确检索和回答。
+判断标准:
+1. 问题是否有明确的主语(指代的对象是否清晰)
+2. 问题是否包含必要的上下⽂(时间、地点、产品型号等)
+3. 问题是否存在歧义(可能有多种理解⽅式)
+请以 JSON 格式返回:
+- 如果问题清晰:{{"need_clarify": false, "reason": "问题清晰的原因"}}
+- 如果需要澄清:{{"need_clarify": true, "clarification": "向⽤户提出的澄清
+问题", "assumption": "如果必须回答时的合理假设"}}
+⽤户问题: {query}
+结果:"""
+)
+    if prompt_type == "rewrite":
+        return rewrite_prompt
+    elif prompt_type == "decompose":
+        return decompose_prompt
+    elif prompt_type == "clarify":
+        return clarify_prompt
+    else:
+        return (ChatPromptTemplate.from_messages(
+                [
+                    (
+                        "system",
+                        "你是专业的汽车知识库助手。只能依据给出的上下文回答;"
+                        "上下文没有答案时,请直接说明“抱歉,我无法从知识库中回答这个问题”。",
+                    ),
+                    ("user", "上下文:\n{context}\n\n用户问题:{question}"),
+                ]
+            )) 

+ 64 - 0
03_RAG优化/agent/rag/rag_cli.py

@@ -0,0 +1,64 @@
+"""RAG 命令行入口。"""
+
+from pathlib import Path
+
+from langchain_core.prompts import ChatPromptTemplate
+
+from agent.config import load_config
+from agent.llm import create_llm
+from agent.rag.document import load_pdf_document
+from agent.rag.embedding import get_embedding_model
+from agent.rag.multi_retriever import MultiRetrieverSystem
+from agent.rag.splitter import recursive_split
+
+from query_optimization import get_prompt_template
+DOCUMENT_PATH = Path("car_info.pdf")
+
+
+def create_rag_system() -> tuple[MultiRetrieverSystem, object, ChatPromptTemplate]:
+    """加载文档并创建检索器和问答链所需组件。"""
+    pages = load_pdf_document(str(DOCUMENT_PATH))
+    documents = recursive_split(pages)
+
+    config = load_config()
+    llm = create_llm(config)
+    retriever_system = MultiRetrieverSystem(
+        documents=documents,
+        embedding_model=get_embedding_model(),
+        llm=llm,
+    )
+    
+    return retriever_system, llm
+
+
+def run_cli() -> None:
+    """启动交互式 RAG 命令行。"""
+    retriever_system, llm = create_rag_system()
+    prompt =get_prompt_template("rewrite")
+    mode = "ensemble"
+    valid_modes = {"bm25", "vector", "ensemble"}
+    print("检索模式:ensemble(可输入 /mode bm25、/mode vector 或 /mode ensemble 切换)")
+
+    while True:
+        question = input("请输入问题:").strip()
+        if not question:
+            continue
+        if question.lower() in {"exit", "quit", "q"}:
+            break
+        if question.startswith("/mode "):
+            new_mode = question.removeprefix("/mode ").strip().lower()
+            if new_mode in valid_modes:
+                mode = new_mode
+                print(f"已切换为 {mode} 检索。")
+            else:
+                print("不支持的检索模式,可选:bm25、vector、ensemble")
+            continue
+
+        relevant_docs = retriever_system.search(question, mode)
+        context = "\n\n".join(document.page_content for document in relevant_docs)
+        answer = (prompt | llm).invoke({"context": context, "question": question})
+        print(f"AI 回答:{answer.content}")
+
+
+if __name__ == "__main__":
+    run_cli()

+ 38 - 0
03_RAG优化/agent/rag/splitter.py

@@ -0,0 +1,38 @@
+"""文本分块策略。"""
+
+from langchain_community.embeddings import DashScopeEmbeddings
+from langchain_core.documents import Document
+from langchain_experimental.text_splitter import SemanticChunker
+from langchain_text_splitters import RecursiveCharacterTextSplitter
+
+from agent.config import load_config
+
+
+def recursive_split(
+    documents: list[Document],
+    chunk_size: int = 500,
+    chunk_overlap: int = 100,
+) -> list[Document]:
+    """使用适合中文文本的递归字符分块。"""
+    splitter = RecursiveCharacterTextSplitter(
+        separators=["\n\n", "\n", "。", "!", "?", ".", " ", ""],
+        chunk_size=chunk_size,
+        chunk_overlap=chunk_overlap,
+        length_function=len,
+    )
+    return splitter.split_documents(documents)
+
+
+def semantic_split(text: str) -> list[str]:
+    """使用 Embedding 的语义断点进行分块。"""
+    config = load_config()
+    embedding_model = DashScopeEmbeddings(
+        model="text-embedding-v4",
+        dashscope_api_key=config.api_key,
+    )
+    splitter = SemanticChunker(
+        embeddings=embedding_model,
+        breakpoint_threshold_type="percentile",
+        breakpoint_threshold_amount=85,
+    )
+    return splitter.split_text(text)

+ 49 - 0
03_RAG优化/main.py

@@ -0,0 +1,49 @@
+from agent.config import load_config
+
+from langchain_text_splitters import RecursiveCharacterTextSplitter  
+from langchain_experimental.text_splitter import SemanticChunker
+from langchain_community.embeddings import DashScopeEmbeddings
+from agent.config import load_config
+
+from langchain_community.document_loaders import PyMuPDFLoader
+
+
+from agent.config import load_config
+from langchain_community.vectorstores import Chroma
+import os
+def main():
+    loader=PyMuPDFLoader("car_info.pdf")
+
+    embedding_model=DashScopeEmbeddings(
+        model="text-embedding-v4",
+        dashscope_api_key=load_config().api_key,
+    )
+    text_splitter = RecursiveCharacterTextSplitter(
+        separators=["\n\n", "\n", "。", "!", "?", ".", " ",""],
+        chunk_size=500,
+        chunk_overlap=100,
+        length_function=len,
+    )
+    chuncks=text_splitter.split_documents(loader.load())
+
+    vectorstore=Chroma.from_documents(
+        documents=chuncks,
+        embedding=embedding_model,
+        collection_metadata={"hnsw:space": "cosine"},
+        persist_directory="./chroma_db"
+    )
+
+    retrever=vectorstore.as_retriever(search_kwargs={"k": 3})
+
+    results=retrever.invoke("汽车的主要组成部分有哪些?")
+    
+    print(results.page_content)
+
+
+
+
+
+if __name__ == "__main__":
+    main()
+
+