-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmain.py
More file actions
321 lines (275 loc) · 11.8 KB
/
Copy pathmain.py
File metadata and controls
321 lines (275 loc) · 11.8 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
import streamlit as st
from PyPDF2 import PdfReader # 用于读取PDF文件内容
from langchain.text_splitter import RecursiveCharacterTextSplitter # 文本分割工具
from langchain_core.prompts import ChatPromptTemplate # 聊天提示模板
from langchain_community.vectorstores import FAISS # Facebook开源的向量数据库
from langchain.tools.retriever import create_retriever_tool # 创建检索工具
from langchain.agents import AgentExecutor, create_tool_calling_agent # 代理执行器和创建代理
# from langchain_community.embeddings import DashScopeEmbeddings # 阿里云提供的嵌入模型
# from langchain.chat_models import init_chat_model # 初始化聊天模型
from langchain_openai import OpenAIEmbeddings
from langchain_openai import ChatOpenAI
import os # 操作系统接口
from dotenv import load_dotenv # 从.env文件加载环境变量
from tqdm import tqdm # 用于进度条
import tiktoken # 添加在文件顶部
load_dotenv(override=True) # 覆盖现有环境变量
# 从环境变量获取API密钥
# DeepSeek_API_KEY = os.getenv("DEEPSEEK_API_KEY")
# dashscope_api_key = os.getenv("dashscope_api_key")
SILICONFLOW_API_KEY = os.getenv("SILICONFLOW_API_KEY")
# 解决某些环境下的库冲突问题
os.environ["KMP_DUPLICATE_LIB_OK"] = "TRUE"
# # 初始化DashScope嵌入模型(阿里云)
# embeddings = DashScopeEmbeddings(
# model="text-embedding-v1", # 嵌入模型名称
# dashscope_api_key=dashscope_api_key # API密钥
# )
# 使用硅基流动的免费嵌入模型
embeddings = OpenAIEmbeddings(
model="BAAI/bge-m3",
api_key=SILICONFLOW_API_KEY,
base_url="https://api.siliconflow.cn/v1",
)
def pdf_read(pdf_doc):
"""读取PDF文件并提取所有文本内容"""
text = ""
# 遍历所有上传的PDF文件
for pdf in pdf_doc:
pdf_reader = PdfReader(pdf)
# 遍历PDF的每一页
for page in pdf_reader.pages:
# 提取页面文本并追加到总文本
text += page.extract_text()
return text
# def get_chunks(text):
# """将长文本分割成小块"""
# # 创建递归字符文本分割器
# text_splitter = RecursiveCharacterTextSplitter(
# chunk_size=500, # 每个块的大小1000
# chunk_overlap=80 # 块之间的重叠部分200
# )
# # 分割文本并返回块列表
# chunks = text_splitter.split_text(text)
# return chunks
def get_chunks(text):
"""将长文本分割成小块(基于token计数)"""
# 使用硅基流动模型对应的分词器
tokenizer = tiktoken.get_encoding("cl100k_base")
def token_length(text):
"""计算文本的token数量"""
# 使用更精确的token计数方法
tokens = tokenizer.encode(text, disallowed_special=())
return len(tokens)
# 创建基于token的文本分割器
text_splitter = RecursiveCharacterTextSplitter(
chunk_size=1000, # 更小的目标token数
chunk_overlap=200, # 块间重叠token数
length_function=token_length, # 使用token计数函数
separators=["\n\n", "\n", "。", "!", "?", ";", " ", ""] # 中文友好的分割点
)
# 分割文本并返回块列表
chunks = text_splitter.split_text(text)
return chunks
# def vector_store(text_chunks):
# """将文本块转换为向量并存储在FAISS数据库中"""
# # 使用嵌入模型将文本转换为向量
# vector_store = FAISS.from_texts(text_chunks, embedding=embeddings)
# # 将向量数据库保存到本地
# vector_store.save_local("faiss_db")
def vector_store(text_chunks):
"""手动控制分批处理"""
# 添加详细日志记录
tokenizer = tiktoken.get_encoding("cl100k_base")
# 记录每个块的token数量
token_counts = []
for i, chunk in enumerate(text_chunks):
tokens = tokenizer.encode(chunk, disallowed_special=())
token_counts.append(len(tokens))
if len(tokens) > 1000:
st.warning(f"⚠️ 警告: 块 {i} 包含 {len(tokens)} tokens (接近上限)")
st.info(f"文本块统计: 平均 {sum(token_counts) / len(token_counts):.1f} tokens, 最大 {max(token_counts)} tokens")
max_batch_size = 32
vector_store = None
for i in tqdm(range(0, len(text_chunks), max_batch_size)):
batch = text_chunks[i:i + max_batch_size]
# 添加延迟避免速率限制
import time
time.sleep(0.5)
batch_store = FAISS.from_texts(batch, embedding=embeddings)
if vector_store is None:
vector_store = batch_store
else:
vector_store.merge_from(batch_store)
vector_store.save_local("faiss_db")
st.success(f"✅ 成功处理 {len(text_chunks)} 个文本块")
return vector_store
def get_conversational_chain(tools, ques):
"""创建对话链并执行问题查询"""
# 初始化DeepSeek聊天模型
# llm = init_chat_model("deepseek-chat", model_provider="deepseek")
# 使用硅基流动的API - 兼容OpenAI格式
llm = ChatOpenAI(
model="Qwen/Qwen3-8B", # 通义千问模型(硅基流动托管)
temperature=0.7,
max_tokens=1024,
api_key=SILICONFLOW_API_KEY, # 硅基流动API密钥
base_url="https://api.siliconflow.cn/v1", # 硅基流动API地址
)
# 定义对话提示模板
prompt = ChatPromptTemplate.from_messages([
(
"system", # 系统指令
"""你是AI助手,请根据提供的上下文回答问题,确保提供所有细节,
如果答案不在上下文中,请说"答案不在上下文中",不要提供错误的答案"""
),
("placeholder", "{chat_history}"), # 聊天历史占位符
("human", "{input}"), # 用户输入占位符
("placeholder", "{agent_scratchpad}"), # 代理工作区占位符
])
# 将工具包装成列表
tool = [tools]
# 创建能调用工具的代理
agent = create_tool_calling_agent(llm, tool, prompt)
# 创建代理执行器
agent_executor = AgentExecutor(agent=agent, tools=tool, verbose=True)
# 使用代理执行问题查询
response = agent_executor.invoke({"input": ques})
print(response) # 打印响应(调试用)
# 在Streamlit界面显示回答
st.write("🤖 回答: ", response['output'])
def check_database_exists():
"""检查FAISS向量数据库是否存在"""
# 检查数据库文件和索引文件是否存在
return os.path.exists("faiss_db") and os.path.exists("faiss_db/index.faiss")
def user_input(user_question):
"""处理用户输入的问题"""
# 检查数据库是否存在
if not check_database_exists():
st.error("❌ 请先上传PDF文件并点击'Submit & Process'按钮来处理文档!")
st.info("💡 步骤:1️⃣ 上传PDF → 2️⃣ 点击处理 → 3️⃣ 开始提问")
return
try:
# 加载本地的FAISS向量数据库
new_db = FAISS.load_local("faiss_db", embeddings, allow_dangerous_deserialization=True)
# 创建检索器
retriever = new_db.as_retriever()
# 创建检索工具
retrieval_chain = create_retriever_tool(
retriever,
"pdf_extractor",
"This tool is to give answer to queries from the pdf"
)
# 使用对话链处理问题
get_conversational_chain(retrieval_chain, user_question)
except Exception as e:
# 处理加载数据库时的错误
st.error(f"❌ 加载数据库时出错: {str(e)}")
import traceback
st.text("详细错误信息:")
st.code(traceback.format_exc(), language="text") # 显示完整的错误堆栈
st.info("请重新处理PDF文件")
def main():
"""主函数,构建Streamlit应用界面"""
# 设置页面配置
st.set_page_config("🤖 TeliangWang的RAG实战")
st.header("🤖 TeliangWang的RAG实战")
# 创建两列布局
col1, col2 = st.columns([3, 1])
# 左侧列:显示数据库状态
with col1:
if check_database_exists():
pass # 数据库存在时不显示警告
else:
st.warning("⚠️ 请先上传并处理PDF文件")
# 右侧列:清除数据库按钮
with col2:
if st.button("🗑️ 清除数据库"):
try:
import shutil
if os.path.exists("faiss_db"):
# 递归删除数据库目录
shutil.rmtree("faiss_db")
st.success("数据库已清除")
st.rerun() # 刷新页面
except Exception as e:
st.error(f"清除失败: {e}")
# 用户问题输入框
user_question = st.text_input(
"💬 请输入问题",
placeholder="例如:这个文档的主要内容是什么?",
disabled=not check_database_exists() # 无数据库时禁用输入
)
# 当用户输入问题后
if user_question:
if check_database_exists():
# 显示加载动画
with st.spinner("🤔 AI正在分析文档..."):
user_input(user_question)
else:
st.error("❌ 请先上传并处理PDF文件!")
# 侧边栏
with st.sidebar:
st.title("📁 文档管理")
# 显示当前数据库状态
if check_database_exists():
st.success("✅ 数据库状态:已就绪")
else:
st.info("📝 状态:等待上传PDF")
st.markdown("---") # 分隔线
# PDF文件上传器
pdf_doc = st.file_uploader(
"📎 上传PDF文件",
accept_multiple_files=True, # 允许多文件上传
type=['pdf'], # 限制文件类型
help="支持上传多个PDF文件" # 帮助文本
)
# 显示已上传的文件信息
if pdf_doc:
st.info(f"📄 已选择 {len(pdf_doc)} 个文件")
for i, pdf in enumerate(pdf_doc, 1):
st.write(f"{i}. {pdf.name}")
# 处理按钮
process_button = st.button(
"🚀 提交并处理",
disabled=not pdf_doc, # 无文件时禁用
use_container_width=True # 使用全宽度
)
# 当点击处理按钮时
if process_button:
if pdf_doc:
with st.spinner("📊 正在处理PDF文件..."):
try:
# 读取PDF文本
raw_text = pdf_read(pdf_doc)
# 检查是否成功提取文本
if not raw_text.strip():
st.error("❌ 无法从PDF中提取文本,请检查文件是否有效")
return
# 分割文本成块
text_chunks = get_chunks(raw_text)
st.info(f"📝 文本已分割为 {len(text_chunks)} 个片段")
# 创建并保存向量数据库
vector_store(text_chunks)
st.success("✅ PDF处理完成!现在可以开始提问了")
st.balloons() # 显示庆祝动画
st.rerun() # 刷新页面
except Exception as e:
st.error(f"❌ 处理PDF时出错: {str(e)}")
else:
st.warning("⚠️ 请先选择PDF文件")
# 使用说明折叠面板
with st.expander("💡 使用说明"):
st.markdown("""
**步骤:**
1. 📎 上传一个或多个PDF文件
2. 🚀 点击"Submit & Process"处理文档
3. 💬 在主页面输入您的问题
4. 🤖 AI将基于PDF内容回答问题
**提示:**
- 支持多个PDF文件同时上传
- 处理大文件可能需要一些时间
- 可以随时清除数据库重新开始
""")
if __name__ == "__main__":
main()