| 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 |
|
|
| |
|
|
| |
| 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}) |
|
|
| |
| print(f"正在加载大语言模型: {MODEL_PATH}") |
| llm = LlamaCpp( |
| model_path=MODEL_PATH, |
| n_ctx=2048, |
| n_threads=8, |
| n_gpu_layers=0, |
| verbose=False, |
| temperature=0.1, |
| stop=["</s>"] |
| ) |
|
|
| |
| qa_chain = RetrievalQA.from_chain_type( |
| llm=llm, |
| chain_type="stuff", |
| retriever=retriever, |
| return_source_documents=True, |
| verbose=False |
| ) |
|
|
| |
|
|
| def answer_query(query): |
| """ |
| 处理用户查询并返回答案。 |
| """ |
| if not query: |
| return "请输入您的问题。" |
|
|
| print(f"收到查询: {query}") |
| |
| try: |
| |
| 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 "抱歉,我在处理您的请求时遇到了问题。请稍后再试。" |
|
|
| |
|
|
| 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]) |
|
|
| |
|
|
| if __name__ == "__main__": |
| print("启动Gradio应用...") |
| demo.launch(server_name="0.0.0.0", server_port=7860) |