| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205 |
- """
- 一. 构建知识库
- 1. 收集资料(无优化空间,但为质量的前提)
- 2. 解析并清洗数据
- ①清洗噪音
- ②清洗多余无意义字符
- 3. 分块(决定7成左右RAG质量)
- ①按固定字符切分
- ②按递归符号切分(递归符号:需人工分析,少量可官网AI总结)
- ③按文档块标识符切分(一般固定文档格式类型,如markdown)
- ④按语义模块切分(费钱,使用embadding模型)
- ⑤按LLM切分(烧钱,不考虑钱直接冲)
- 4. 向量化(智普embadding模型, 余弦相似度)
- 5. 向量数据入库(测试chroma,主流milvus)
- 二. 用户基于知识库提问
- 1. 用户提问
- 2. 用户问题向量化,到库中检索,top_k返回前n条完整结果
- 3. 将结果封装进提示词
- 4. 请求LLM返回最终结果
- """
- from dotenv import load_dotenv
- from langchain_openai import ChatOpenAI
- import os, re
- from fastapi import FastAPI, UploadFile, File, HTTPException
- from langchain_community.document_loaders import PyPDFLoader, DirectoryLoader
- import tempfile
- from langchain_text_splitters import RecursiveCharacterTextSplitter
- from langchain_community.vectorstores import Chroma
- from langchain_community.embeddings import ZhipuAIEmbeddings
- from langchain_core.prompts import PromptTemplate
- from zai import ZhipuAiClient
- load_dotenv(override=True)
- deepseek_base_url = os.getenv("DEEPSEEK_BASE_URL")
- deepseek_base_key = os.getenv("DEEPSEEK_BASE_KEY")
- deepseek_base_name = os.getenv("DEEPSEEK_BASE_NAME")
- # ==============================一. 构建知识库==================================================
- # ==============================1. 收集资料==================================================
- class KnowledgeBaseBuilder:
- def __init__(self, pdf_dir="./resources", persist_dir="./chroma_db"):
- self.pdf_dir = pdf_dir
- self.persist_dir = persist_dir
- self.embeddings = ZhipuAIEmbeddings(
- model="embedding-3", # 或者 "embedding-3"
- api_key=os.getenv("ZHIPU_API_KEY")
- )
- self.documents = []
-
- def load_pdfs(self):
- """加载PDF文件"""
- loader = DirectoryLoader(
- self.pdf_dir,
- glob="**/*.pdf",
- loader_cls=PyPDFLoader,
- show_progress=True
- )
- self.documents = loader.load()
- print(f"加载了 {len(self.documents)} 个文档")
- return self.documents
- # ==============================2. 解析并清洗数据==================================================
- def clean_documents(self, docs):
- """清洗文档数据"""
- cleaned_docs = []
- empty_count = 0
-
- for i, doc in enumerate(docs):
- text = doc.page_content
-
- # 检查文本是否为空
- if not text or len(text.strip()) == 0:
- empty_count += 1
- print(f"⚠️ 文档 {i+1} 内容为空,跳过")
- continue
-
- # ① 清洗噪音:移除页眉页脚、水印等
- text = re.sub(r'第\s*\d+\s*页\s*/\s*共\s*\d+\s*页', '', text)
- text = re.sub(r'Copyright.*?\n', '', text, flags=re.IGNORECASE)
- text = re.sub(r'机密|内部资料|仅供内部使用', '', text)
-
- # ② 去除无意义字符
- text = re.sub(r'\n\s*\n+', '\n\n', text) # 合并多个空行
- text = re.sub(r'[ \t]+', ' ', text) # 合并多个空格
- # 保留中英文、数字和常用标点
- text = re.sub(r'[^\u4e00-\u9fa5a-zA-Z0-9\.\,\,\。\!\?\:\;\(\)\n]', ' ', text)
- text = re.sub(r'\s+', ' ', text).strip()
-
- # 清洗后再次检查
- if not text or len(text) < 10:
- empty_count += 1
- print(f"⚠️ 文档 {i+1} 清洗后内容过短,跳过")
- continue
-
- doc.page_content = text
- cleaned_docs.append(doc)
-
- print(f"✅ 清洗完成,有效文档: {len(cleaned_docs)}, 跳过: {empty_count}")
-
- if not cleaned_docs:
- raise ValueError("清洗后没有有效的文档内容")
-
- return cleaned_docs
- # ==============================3. 分块==================================================
- def chunk_documents(self, docs, strategy="recursive"):
- """
- 多种分块策略
- strategy: fixed, recursive, semantic, markdown
- 当前实现递归分块策略
- """
- if strategy == "recursive":
- splitter = RecursiveCharacterTextSplitter(
- chunk_size=500,
- chunk_overlap=50,
- separators=[
- "\n\n", # 段落
- "\n", # 行
- "。", # 中文句号
- "!", # 感叹号
- "?", # 问号
- ";", # 分号
- ",", # 逗号
- ".", # 英文句号
- "!",
- "?",
- " "
- ],
- length_function=len,
- keep_separator=True
- )
- chunked_docs = splitter.split_documents(docs)
- return chunked_docs
- # ==============================4. 向量化==================================================
- # ==============================5. 向量数据入库==================================================
- def build_vectorstore(self, chunks, collection_name="knowledge_base"):
- """向量化并存入Chroma数据库"""
- # 使用智谱Embedding模型
- vectorstore = Chroma.from_documents(
- documents=chunks,
- embedding=self.embeddings,
- persist_directory=self.persist_dir,
- collection_name=collection_name
- )
- vectorstore.persist()
- print(f"向量数据库构建完成,存储于: {self.persist_dir}")
- return vectorstore
- # ==============================二. 用户基于知识库提问==================================================
- # ==============================1. 用户提问==================================================
- # ==============================2. 用户问题向量化,到库中检索,top_k返回前n条完整结果==================================================
- # ==============================3. 将结果封装进提示词==================================================
- # ==============================4. 请求LLM返回最终结果==================================================
- llm = ChatOpenAI(
- base_url = deepseek_base_url,
- api_key = deepseek_base_key,
- model= deepseek_base_name
- )
- if __name__ == "__main__":
- knowlege = KnowledgeBaseBuilder()
- docs = knowlege.load_pdfs()
- clean_docs = knowlege.clean_documents(docs)
- chunks = knowlege.chunk_documents(clean_docs)
- vectorstore = knowlege.build_vectorstore(chunks)
- # vectorstore = Chroma(
- # persist_directory="./chroma_db", # 持久化目录
- # embedding_function=knowlege.embeddings,
- # collection_name="knowledge_base" # 集合名称(与构建时一致)
- # )
- user_question = '遥控器如何操作?'
- retrieved_docs = vectorstore.similarity_search(user_question, k=3)
- context = "\n\n".join([doc.page_content for doc in retrieved_docs])
- prompt_template = """
- 你是一个专业的AI助手,请基于以下参考资料回答用户的问题。
-
- 【参考资料】:
- {context}
-
- 【用户问题】:
- {question}
-
- 要求:
- 1. 严格基于参考资料回答,不要编造信息
- 2. 如果参考资料中没有相关信息,请明确告知
- 3. 回答要简洁、准确、有条理
-
- 回答:
- """
- prompt = PromptTemplate(
- template=prompt_template,
- input_variables=["context", "question"]
- )
-
- # 4. 请求LLM
- formatted_prompt = prompt.format(context=context, question=user_question)
- response = llm.invoke(formatted_prompt)
- print(response)
|