| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425 |
- 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=120,
- 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 hybrid_hyde_search(llm, embeddings: DashScopeEmbeddings, hyde_chain, question) -> dict:
- """如果 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
- hypothetical_doc = hyde_chain.invoke({"question": question})
- # 方案一:用假设性文档检索(偏语义)
- hyde_docs = vectorstore.similarity_search(hypothetical_doc, k=10)
- bm25_docs = bm25_retriever.invoke(hypothetical_doc, k=10) # question must be a str here
- # 方案二:用原始问题检索(偏精确)
- original_docs = vectorstore.similarity_search(question, k=10)
- # 方案二:用原始问题检索(偏精确)
- # 合并去重
- seen = set()
- merged_docs = []
- for doc in hyde_docs + bm25_docs:
- if hash(doc.page_content) not in seen:
- seen.add(hash(doc.page_content))
- merged_docs.append(doc)
- all_docs = merged_docs[:10]
- # 第三步:用原始问题 + 检索结果生成答案
- context = "\n\n".join([doc.page_content for doc in all_docs])
- answer_prompt = f"""基于以下上下文回答用户问题。如果上下文中没有相关信息,请说明。
- 上下文:{context}
- 用户问题:{question}
- 答案:"""
- answer = llm.invoke(answer_prompt).content
- return {
- "original_query": question,
- "hypothetical_doc": hypothetical_doc,
- "retrieved_docs": docs,
- "answer": answer
- }
- 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("✅ 模型初始化完成")
- query = "什么是工作质量考核标准,还有聘任考核程序是什么?"
- hyde_chain = build_hyde_chain(llm)
- result = hybrid_hyde_search(llm, embeddings, hyde_chain, query)
- print("\n" + "=" * 60)
- print(f"问题:{query}")
- print("=" * 60)
- print(result["answer"])
- if __name__ == "__main__":
- main()
|