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=[""] # 模型停止生成的标志 ) # 创建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( """

成都二中初59级二班同学信息查询机器人

您可以查询同学的联系方式、当前状态(有联系/无法联系/已故)等信息。

例如:“阎国蜀的电话是多少?”、“谁去世了?”、“崔厚佳能联系上吗?”

""" ) 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)