fzd9 / app.py
fdbw's picture
Rename app1.py to app.py
7b172e6 verified
Raw
History Blame Contribute Delete
4.5 kB
import gradio as gr
from langchain_community.document_loaders import TextLoader
from langchain_community.vectorstores import Chroma
from langchain_community.embeddings import HuggingFaceEmbeddings
from langchain_community.llms import LlamaCpp
from langchain.chains import RetrievalQA
import os
# --- 1. 初始化组件 ---
# 定义模型和向量存储的路径
PERSIST_DIRECTORY = './db'
MODEL_PATH = './models/mistral-7b-instruct-v0.2.Q4_K_M.gguf'
# 加载或创建向量数据库
print("正在加载或创建向量数据库...")
embeddings = HuggingFaceEmbeddings(model_name="all-MiniLM-L6-v2")
if not os.path.exists(PERSIST_DIRECTORY):
print("向量数据库不存在,正在从知识库文件创建...")
# 加载知识库文件
loader = TextLoader("knowledge_base.txt", encoding='utf-8')
documents = loader.load()
# 创建并持久化向量数据库
db = Chroma.from_documents(documents, embeddings, persist_directory=PERSIST_DIRECTORY)
db.persist()
else:
print("正在从本地加载向量数据库...")
# 从已保存的目录加载向量数据库
db = Chroma(persist_directory=PERSIST_DIRECTORY, embedding_function=embeddings)
# 创建检索器
retriever = db.as_retriever(search_kwargs={"k": 5}) # 每次检索最多返回5个最相关的结果
# 加载大语言模型 (LLM)
print(f"正在加载大语言模型: {MODEL_PATH}")
llm = LlamaCpp(
model_path=MODEL_PATH,
n_ctx=2048, # 上下文窗口大小
n_threads=8, # 使用的CPU线程数
n_gpu_layers=0, # 如果有GPU,可以设置为正整数来加速
verbose=False, # 是否输出详细日志
temperature=0.1, # 随机性,越低越确定
stop=["</s>"] # 模型停止生成的标志
)
# 创建RAG链
qa_chain = RetrievalQA.from_chain_type(
llm=llm,
chain_type="stuff", # 将检索到的所有信息一次性"stuff"给LLM
retriever=retriever,
return_source_documents=True, # 是否返回原始的检索结果
verbose=False # 是否输出链的详细运行日志
)
# --- 2. 定义问答函数 ---
def answer_query(query):
"""
处理用户查询并返回答案。
"""
if not query:
return "请输入您的问题。"
print(f"收到查询: {query}")
try:
# 运行RAG链
result = qa_chain({"query": query})
# 提取答案和来源
answer = result["result"]
source_documents = result["source_documents"]
# 格式化来源信息
sources = ""
seen_sources = set() # 用于去重
for doc in source_documents:
content = doc.page_content
if content not in seen_sources:
seen_sources.add(content)
sources += f"\n• {content}"
# 构建最终回答
final_answer = f"{answer}\n\n---\n信息来源:{sources}"
return final_answer
except Exception as e:
error_message = f"处理查询时发生错误: {e}"
print(error_message)
return "抱歉,我在处理您的请求时遇到了问题。请稍后再试。"
# --- 3. 创建Gradio界面 ---
with gr.Blocks(title="成都二中初59级二班同学信息查询机器人") as demo:
gr.Markdown(
"""
<h1 style="text-align: center;">成都二中初59级二班同学信息查询机器人</h1>
<p style="text-align: center;">您可以查询同学的联系方式、当前状态(有联系/无法联系/已故)等信息。</p>
<p style="text-align: center;">例如:“阎国蜀的电话是多少?”、“谁去世了?”、“崔厚佳能联系上吗?”</p>
"""
)
chatbot = gr.Chatbot(label="聊天记录")
history = gr.State([])
with gr.Row():
user_input = gr.Textbox(
show_label=False,
placeholder="请输入您的问题...",
lines=3,
container=False
)
submit_btn = gr.Button("发送", variant="primary")
# 绑定事件
def chat(user_message, history):
history = history or []
bot_message = answer_query(user_message)
history.append((user_message, bot_message))
return history, history
submit_btn.click(chat, [user_input, history], [chatbot, history])
user_input.submit(chat, [user_input, history], [chatbot, history])
# --- 4. 运行应用 ---
if __name__ == "__main__":
print("启动Gradio应用...")
demo.launch(server_name="0.0.0.0", server_port=7860)