| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159 |
- 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.vectorstores import Chroma
- from langchain_core.output_parsers import StrOutputParser
- from langchain_core.prompts import ChatPromptTemplate
- from langchain_openai import ChatOpenAI
- from langchain_classic.chains.combine_documents import create_stuff_documents_chain
- from langchain_text_splitters import RecursiveCharacterTextSplitter
- # ── 配置 ──────────────────────────────────────────────────────────
- 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=10,
- 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()
- # ========== 第二步:清洗数据(可选,根据文档质量决定)==========
- clean_pages = [clean_pdf_text(page) for page in pages]
- text_splitter = RecursiveCharacterTextSplitter(
- chunk_size=500,
- chunk_overlap=50,
- separators=["\n\n", "\n", "。", ";", ",", " ", ""],
- )
- 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 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 main():
- config = load_config()
- llm = init_llm(config)
- embeddings = init_embeddings(config)
- print("✅ 模型初始化完成")
- vectorstore = build_or_load_vectorstore(embeddings)
- retriever = vectorstore.as_retriever(search_kwargs={"k": 3})
- query = "什么是工作质量考核标准?"
- answer = ask(query, retriever, llm)
- print("\n" + "=" * 60)
- print(f"问题:{query}")
- print("=" * 60)
- print(answer)
- if __name__ == "__main__":
- main()
|