File size: 4,504 Bytes
24d02fa
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
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)