import os import chromadb from dotenv import load_dotenv from langchain_community.document_loaders import PyMuPDFLoader from langchain_community.embeddings import DashScopeEmbeddings from langchain_community.retrievers import BM25Retriever from schema import RAGWithQueryRewriting from schema import RAGWithDecomposition from langchain_community.vectorstores import Chroma from langchain_core.output_parsers import StrOutputParser from langchain_core.prompts import ChatPromptTemplate from langchain_openai import ChatOpenAI from langchain_community.retrievers import BM25Retriever from langchain_community.embeddings import DashScopeEmbeddings from langchain_classic.retrievers import EnsembleRetriever from langchain_classic.chains.combine_documents import create_stuff_documents_chain from langchain_text_splitters import RecursiveCharacterTextSplitter from langchain_core.prompts import PromptTemplate from langchain_community.chat_models import ChatTongyi from langchain_classic.retrievers.multi_query import MultiQueryRetriever from langchain_experimental.text_splitter import SemanticChunker # ── 配置 ────────────────────────────────────────────────────────── BASE_DIR = os.path.dirname(os.path.abspath(__file__)) PDF_PATH = os.path.join(BASE_DIR, "docs", "浦发上海浦东发展银行西安分行个金客户经理考核办法.pdf") CHROMA_DB_PATH = os.path.join(BASE_DIR, "chroma") COLLECTION_NAME = "shanghai_bank_policy" PROMPT_TEMPLATE = """ 你是一个专业的知识库助手。请根据以下上下文回答问题。 **规则:** - 只基于提供的上下文回答,不要编造 - 如果上下文中没有相关信息,直接说「根据现有资料,我找不到这个问题的答案」 - 回答要简洁直接,引用原文时用引号 **上下文:** {context} **问题:** {question} """ def load_config() -> dict: """加载 .env 中的配置项""" load_dotenv() config = { "api_key": os.getenv("ALIYUN_API_KEY"), "base_url": os.getenv("ALIYUN_BASE_URL"), "chat_model": os.getenv("ALIYUN_CHAT_MODEL"), "embedding_model": os.getenv("ALIYUN_EMBEDDING_MODEL", "text-embedding-v3"), } missing = [k for k, v in config.items() if not v and k != "embedding_model"] if missing: raise EnvironmentError(f"缺少环境变量: {missing},请检查 .env 文件") return config def init_llm(config: dict) -> ChatOpenAI: return ChatOpenAI( base_url=config["base_url"], api_key=config["api_key"], model=config["chat_model"], temperature=0, timeout=60, max_retries=1, ) def init_embeddings(config: dict) -> DashScopeEmbeddings: return DashScopeEmbeddings( model=config["embedding_model"], dashscope_api_key=config["api_key"], ) def clean_pdf_text(text: str) -> str: """清洗 PDF 解析出的文本,去除常见噪声""" import re # 删除非中文字符之间的换行符 text = re.sub(r'[^一](\n)[^一]', lambda m: m.group(0).replace('\n', ''), text) # 删除项目符号和多余空格 text = text.replace('•', '').replace(' ', ' ') # 删除连续的换行符(保留一个) text = re.sub(r'\n{2,}', '\n', text) return text.strip() def build_or_load_vectorstore(embeddings: DashScopeEmbeddings) -> Chroma: """如果 collection 已存在则直接加载,否则解析 PDF 并建库""" chroma_client = chromadb.PersistentClient(path=CHROMA_DB_PATH) existing_collections = [col.name for col in chroma_client.list_collections()] if COLLECTION_NAME in existing_collections: print(f"✅ 检测到已有 collection「{COLLECTION_NAME}」,直接加载") return Chroma( embedding_function=embeddings, collection_name=COLLECTION_NAME, client=chroma_client, ) print(f"ℹ️ 未检测到 collection「{COLLECTION_NAME}」,开始解析文档并建库") loader = PyMuPDFLoader(PDF_PATH) pages = loader.load() # ========== 第二步:清洗数据(可选,根据文档质量决定)========== pages = clean_documents(pages) # 创建语义分块器 # breakpoint_threshold_type="percentile" 表示用百分位数法确定切分阈值 # breakpoint_threshold_amount=95 表示只有相似度排名后 5% 的位置才会被切开 text_splitter = SemanticChunker( embeddings=embeddings, breakpoint_threshold_type="percentile", breakpoint_threshold_amount=95 # 值越大,切出来的块越少(越粗) ) docs = text_splitter.split_documents(pages) vectorstore = Chroma.from_documents( embedding=embeddings, collection_name=COLLECTION_NAME, client=chroma_client, documents=docs, ) print(f"✅ 建库完成,共 {len(docs)} 个分块") return vectorstore def build_ensemble_retriever(embeddings: DashScopeEmbeddings) -> EnsembleRetriever: """如果 collection 已存在则直接加载,否则解析 PDF 并建库""" chroma_client = chromadb.PersistentClient(path=CHROMA_DB_PATH) existing_collections = [col.name for col in chroma_client.list_collections()] loader = PyMuPDFLoader(PDF_PATH) pages = loader.load() # ========== 第二步:清洗数据(可选,根据文档质量决定)========== pages = clean_documents(pages) # 创建语义分块器 # breakpoint_threshold_type="percentile" 表示用百分位数法确定切分阈值 # breakpoint_threshold_amount=95 表示只有相似度排名后 5% 的位置才会被切开 text_splitter = SemanticChunker( embeddings=embeddings, breakpoint_threshold_type="percentile", breakpoint_threshold_amount=95 # 值越大,切出来的块越少(越粗) ) docs = text_splitter.split_documents(pages) vectorstore = None if COLLECTION_NAME in existing_collections: print(f"✅ 检测到已有 collection「{COLLECTION_NAME}」,直接加载") vectorstore = Chroma( embedding_function=embeddings, collection_name=COLLECTION_NAME, client=chroma_client, ) else: print(f"ℹ️ 未检测到 collection「{COLLECTION_NAME}」,开始解析文档并建库") vectorstore = Chroma.from_documents( embedding=embeddings, collection_name=COLLECTION_NAME, client=chroma_client, documents=docs, ) # ---- 创建 BM25 检索器 ---- # BM25 不需要向量,直接基于文本的关键词匹配 bm25_retriever = BM25Retriever.from_documents(docs) bm25_retriever.k = 10 vector_retriever = vectorstore.as_retriever(search_kwargs={"k": 10}) ensemble_retriever = EnsembleRetriever( retrievers=[bm25_retriever, vector_retriever], weights=[0.4, 0.6], normalize_scores=True # 将不同检索器的分数归一化到 [0,1],避免偏差 ) return ensemble_retriever def build_multi_query_retriever(embeddings: DashScopeEmbeddings, llm) -> MultiQueryRetriever: # 自定义改写提示词(可选,不写则用默认的) CUSTOM_PROMPT = PromptTemplate( input_variables=["question"], template="""你是一个专业的问题改写助手。请为下面的问题生成 4 个不同的改写版本, 每个版本应该: - 从不同角度表达相同的意图 - 使用不同的关键词和表达方式 - 保持问题的核心含义 每个问题单独一行,不要编号。 原始问题: {question} 改写后的问题:""" ) ensemble_retriever = build_ensemble_retriever(embeddings=embeddings) # 创建多查询检索器 # 底层检索器用的是上面的混合检索器,这样每个改写问题都会走混合检索 multi_query_retriever = MultiQueryRetriever.from_llm( retriever=ensemble_retriever, llm=llm, prompt=CUSTOM_PROMPT ) return multi_query_retriever def ask(query: str, retriever, llm) -> str: """检索 + 生成回答""" relevant_docs = retriever.invoke(query) context = "\n\n---\n\n".join(d.page_content for d in relevant_docs) prompt = ChatPromptTemplate.from_template(PROMPT_TEMPLATE) chain = prompt | llm | StrOutputParser() return chain.invoke({"context": context, "question": query}) def build_hyde_chain(llm) -> dict: # HyDE Prompt:生成假设性文档 hyde_prompt = PromptTemplate( input_variables=["question"], template="""请根据以下问题,写一段可能包含答案的文档片段。 要求: 1. 像真实文档一样专业、详细 2. 包含具体的数据、步骤或事实 3. 长度在 100-200 字之间 4. 即使你不确定答案,也要根据问题合理推测,写出一段"看起来像真的"文档 问题: {question} 假设性文档:""" ) hyde_chain = hyde_prompt | llm | StrOutputParser() return hyde_chain def build_rewrite_chain(llm) -> dict: # 查询重写 Prompt rewrite_prompt = PromptTemplate( input_variables=["query"], template="""你是一个查询优化助手。请将用户的口语化问题改写为更适合信息检索的精确查询。 改写要求: 1. 补充隐含的上下文信息 2. 将口语化表达转为专业表述 3. 消除歧义,明确查询意图 4. 保持原意不变,不要添加原问题未提及的内容 5. 直接输出改写后的查询,不要解释 用户问题: {query} 改写后的查询:""" ) # 构建重写链 rewrite_chain = rewrite_prompt | llm | StrOutputParser() return rewrite_chain def clean_pdf_text(text: str) -> str: """清洗 PDF 解析出的文本,去除常见噪声""" import re # 删除非中文字符之间的换行符 text = re.sub(r'[^一](\n)[^一]', lambda m: m.group(0).replace('\n', ''), text) # 删除项目符号和多余空格 text = text.replace('•', '').replace(' ', ' ') # 删除连续的换行符(保留一个) text = re.sub(r'\n{2,}', '\n', text) return text.strip() def clean_documents(pages: list) -> list: """对 Document 列表逐个清洗 page_content,返回清洗后的新 Document 列表""" for page in pages: page.page_content = clean_pdf_text(page.page_content) return pages def main(): config = load_config() llm = init_llm(config) embeddings = init_embeddings(config) print("✅ 模型初始化完成") rewrite_chain = build_rewrite_chain(llm) multi_query_retriever = build_multi_query_retriever(embeddings, llm) rag = RAGWithDecomposition(retriever = multi_query_retriever,llm = llm) query = "什么是工作质量考核标准,还有聘任考核程序是什么?" result = rag.invoke(query) print("\n" + "=" * 60) print(f"问题:{query}") print("=" * 60) print(result["answer"]) if __name__ == "__main__": main()