项目介绍

RAG回顾

RAG即检索、增强和生成,其主要分为2条线:

  1. 离线处理: 向私有知识库(向量存储)源源不断添加私有知识文档。
  • 向知识库添加来自未来的知识文档(基于模型训练完成时间)
  • 向模型添加私有知识文档
  • 给出模型参考资料,规避模型幻觉(一本正经的胡说八道)
  1. 在线处理: 用户提问会先基于私有知识库做检索,获取参考资料,同步组装新提示词询问大模型获取结果。

项目需求和思路

本次项目以“某东商品衣服"为例,以衣服属性构建本地知识。使用者可以自由更新本地知识,用户问题的答案也是基于本地知识生成的。

项目主要会实现如下代码

离线模块

离线流程

  1. app_file_uploader.py是文件服务
  • uploader_file.get_value(): 拿到用户上传文件的内容
  • KnowledgetBaseService: 创建知识库服务实例,把文件内容存入向量数据库
  1. knowledge_base.py是知识库服务
  • check_md5: 读取记录文件,对比md5值
  • save_md5: 把md5值写入记录文件
  • get_string_md5: 把文件内容转md5值
  • chroma: 向量数据库实例
  • spliter: 字符串分割器实例
  • upload_by_str: 数据入库
  1. config_data.py是全局配置文件

完成app_file_uploader.py文件的编写

"""
基于Streamlit完成Web页面上传服务
pip install streamlit
"""
import streamlit as st

# 添加网页标题
st.title("知识库更新服务")

# 文件上传
uploaded_file = st.file_uploader(
    "请上传TXT文件",  # 标题
    type=["txt"],  # 允许上传的文件类型
    accept_multiple_files=False,  # 是否允许上传多个文件
)

if uploaded_file is not None:
    # 提取文件信息
    file_name = uploaded_file.name
    file_type = uploaded_file.type
    file_size = uploaded_file.size / 1024 # KB
    # 展示文件信息
    st.subheader(f"文件名称:{file_name}")
    st.write(f"类型:{file_type} | 大小:{file_size:.2f} KB")
    # 提取文件内容
    # getvalue()默认返回bytes, 使用decode()转码为str
    text = uploaded_file.getvalue().decode("utf-8")
    st.write(text)
  1. 通过命令行启动服务

  1. 到此,完成了一部分工作

完成knowledge_base.py文件的编写

"""
知识库
"""
import hashlib
import os

from datetime import datetime
from langchain_community.embeddings import DashScopeEmbeddings
from langchain_community.vectorstores import Chroma
from langchain_text_splitters import RecursiveCharacterTextSplitter

import config_data as config


def check_md5(md5_str: str):
    """
    检查md5字符串是否存在
    """
    if not os.path.exists(config.md5_path):
        # 文件不存在
        open(config.md5_path, "w", encoding="utf-8").close()
        return False
    else:
        for line in open(config.md5_path, "r", encoding="utf-8").readlines():
            line = line.strip()  # 去除字符串首尾的空格/换行符
            if line == md5_str:
                return True
        return False


def save_md5(md5_str: str):
    """
    保存md5值
    """
    with open(config.md5_path, "a", encoding="utf-8") as f:
        f.write(md5_str + "\n")


def get_string_md5(input_str: str, encoding="utf-8"):
    """
    计算md5
    """
    # 将字符串转成bytes字节数组
    str_bytes = input_str.encode(encoding)
    # 创建MD5对象
    md5_obj = hashlib.md5()
    # 更新转换内容
    md5_obj.update(str_bytes)
    # 得到md5的十六进制字符串
    md5_str = md5_obj.hexdigest()
    return md5_str


class KnowLedgeBaseService(object):
    """
    知识库服务
    """

    def __init__(self):
        # 如果文件夹不存在则创建
        os.makedirs(config.persist_directory, exist_ok=True)
        # Chroma向量库实例
        self.chroma = Chroma(
            collection_name=config.collection_name,  # 数据库表名
            embedding_function=DashScopeEmbeddings(  # 创建嵌入模型实例
                model="text-embedding-v4",
                dashscope_api_key="sk-0af497cdd462495aa2f08dddaead1ded",
            ),
            persist_directory=config.persist_directory,  # 持久化存储目录
        )
        # 文本分割器实例
        self.spliter = RecursiveCharacterTextSplitter(
            chunk_size=config.chunk_size,  # 分割后的文本段最大长度
            chunk_overlap=config.chunk_overlap,  # 相邻段落间之间的字符重叠数
            separators=config.separators,  # 分割的分隔符
            length_function=len,  # 使用py自带的len函数做字符长度统计
        )

    def upload_by_str(self, data, filename):
        """将字符串向量化,存入数据库"""
        md5_hex = get_string_md5(data)

        if check_md5(md5_hex):
            return "[跳过]内容已存在知识库中"

        if len(data) > config.max_split_char_number:
            knowledge_chunks: list[str] = self.spliter.split_text(data)
        else:
            knowledge_chunks = [data]

        metadate = {
            "source": filename,
            "create_time": datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
            "operator": "王",
        }

        # 内容加载到向量库
        self.chroma.add_texts(
            # 要求传入Iterable[str], Iterable即可迭代对象,list数组/tuple字典都是可迭代对象,然后容器中是str
            texts=knowledge_chunks, # 待加载的文本
            metadate=[metadate for _ in knowledge_chunks], # 数据的元数据(额外信息)
        )
        # 记录md5值
        save_md5(md5_hex)
        return "[完成]内容已经加载到向量库中"


if __name__ == '__main__':
    # 单元测试
    # print(get_string_md5("hello world"))
    # print(check_md5("5eb63bbbe01eeed093cb22bb8f5acdc3"))
    # print(save_md5("5eb63bbbe01eeed093cb22bb8f5acdc3"))
    service = KnowLedgeBaseService()
    r = service.upload_by_str("周杰伦", "testfile")
    print(r)
  1. 配置文件
md5_path = "./md5.text"

# Chroma
collection_name = "rag"
persist_directory = "./chroma_db"

# spliter
chunk_size = 1000
chunk_overlap = 100
separators = ["\n\n", "\n", " ", "", ".", "!", "?", "。","!","?"]
max_split_char_number = 1000 # 文本分隔的阈值 (过短的文本没必要分割)
  1. 执行单元测试

  1. 至此,完成一部分工作

完成离线流程的开发

"""
基于Streamlit完成Web页面上传服务
  pip install streamlit

Streamlit的重要特点
1. 当web网页元素发生变化,则代码重新执行一遍,
2. 会造成页面状态的丢失
3. 如果有数据需要保存,要借助session_state实现
4. session_state就是一个持久化存储的字典
"""
import time

import streamlit as st
from knowledge_base import KnowLedgeBaseService

# 添加网页标题
st.title("知识库更新服务")

# 文件上传
uploaded_file = st.file_uploader(
    "请上传TXT文件",  # 标题
    type=["txt"],  # 允许上传的文件类型
    accept_multiple_files=False,  # 是否允许上传多个文件
)


if "service" not in st.session_state:
    st.session_state["service"] = KnowLedgeBaseService()

if uploaded_file is not None:
    # 提取文件信息
    file_name = uploaded_file.name
    file_type = uploaded_file.type
    file_size = uploaded_file.size / 1024 # KB
    # 展示文件信息
    st.subheader(f"文件名称:{file_name}")
    st.write(f"类型:{file_type} | 大小:{file_size:.2f} KB")
    # 提取文件内容
    # getvalue()默认返回bytes, 使用decode()转码为str
    text = uploaded_file.getvalue().decode("utf-8")
    # st.write(text)

    # 存储文件 (加一个loding效果)
    with st.spinner("载入知识库中..."):
        time.sleep(1)
        # 调用知识库服务完成文件内容的向量化以及存储
        result = st.session_state["service"].upload_by_str(text, file_name)
        st.write(result)
  1. 调试

  1. 至此,完成的工作,整个离线流程全部打通

在线模块

在线流程

  1. vector_stores: 向量存储服务
  • 通过get_retriever方法获取检索器,用于向量检索
  1. file_history_store.py: 历史信息存储
  2. rag.py:检索增强服务
  • vector_service: 向量检索器的实例
  • prompt_template: 提示词模板
  • chain: 链实例
  • chat_model: 模型对象
  • __get_chain(): 获取链
  • invoke: 向模型提问
  1. app_qa.py 用户交互服务

vector_stores 向量存储服务开发

from langchain_chroma import Chroma
import config_data as config


class VectorStoreService(object):
    def __init__(self, embedding):
        """
        :param embedding: 嵌入模型的传入
        """
        self.embedding = embedding
        self.vector_store = Chroma( # 创建向量存储实例
            collection_name=config.collection_name,
            embedding_function=self.embedding,
            persist_directory=config.persist_directory,
        )

    def get_retriever(self):
        """
        获取向量检索器
        """
        return self.vector_store.as_retriever(search_kwargs={"k": config.simplify_threshold})

if __name__ == '__main__':
    from langchain_community.embeddings import DashScopeEmbeddings
    # retriever就是检索器对象
    retriever = service = VectorStoreService(DashScopeEmbeddings(model="text-embedding-v4", dashscope_api_key="sk-0af497cdd462495aa2f08dddaead1ded")).get_retriever()
    res = retriever.invoke("我的体重180斤,尺码推荐")
    print(res)
  1. 测试:已经把匹配到的 尺码推荐.text 的知识拿到了

rag.py 知识检索开发

from langchain_community.chat_models import ChatTongyi
from langchain_community.embeddings import DashScopeEmbeddings
from langchain_core.documents import Document
from langchain_core.output_parsers import StrOutputParser
from langchain_core.prompts import ChatPromptTemplate
from langchain_core.runnables import RunnablePassthrough

import config_data as config

from vector_stores import VectorStoreService


def print_prompt(prompt):
    print("=" * 20)
    print(prompt.to_string())
    print("=" * 20)
    return prompt


class RagService(object):
    def __init__(self):
        # 构建向量存储对象
        self.vector_store = VectorStoreService(
            embedding=DashScopeEmbeddings(model=config.embedding_model_name,
                                          dashscope_api_key="sk-0af497cdd462495aa2f08dddaead1ded")
        )
        # 构建提示词模板
        self.prompt_template = ChatPromptTemplate.from_messages(
            [
                ("system", "以我提供的已知参考资料为主,简洁专业的回答用户问题。参考资料:{context}"),
                ("user", "请回答用户提问:{input}"),
            ]
        )
        # 构建模型对象
        self.chat_model = ChatTongyi(model=config.chat_model_name, api_key="sk-0af497cdd462495aa2f08dddaead1ded")
        # 构建执行链
        self.chain = self.__get_chain()

    def __get_chain(self):
        """获取最终的执行链"""
        retriever = self.vector_store.get_retriever()

        def format_document(docs: list[Document]):
            if not docs:
                return "无相关参考资料"

            formatted_str = ""
            for doc in docs:
                formatted_str += f"文档片段:{doc.page_content} \n文档元数据:{doc.metadata}\n\n"

            return formatted_str

        chain = (
                {
                    "input": RunnablePassthrough(),
                    "context": retriever | format_document,
                } | self.prompt_template | print_prompt | self.chat_model | StrOutputParser()
        )

        return chain


if __name__ == '__main__':
    res = RagService().chain.invoke("我身高180厘米,尺码推荐")
    print(res)

file_history_store.py 历史信息存储开发

import json
import os
from typing import Sequence

from langchain_core.chat_history import BaseChatMessageHistory
from langchain_core.messages import BaseMessage, messages_from_dict, message_to_dict


# 自定义会话存储的方法
class FileChatMessageHistory(BaseChatMessageHistory):
    def __init__(self, session_id, storage_path):
        self.session_id = session_id # 会话id
        self.storage_path = storage_path # 存储路径
        self.file_path = os.path.join(self.storage_path, self.session_id)  # 拼接完整的文件路径

        # 判断文件夹是否存在 (不存在自动创建)
        os.makedirs(os.path.dirname(self.file_path), exist_ok=True)

    def add_messages(self, messages: Sequence[BaseMessage]):   # Sequence序列 类似list/tuple
        all_messages = list(self.messages) # 已有的消息列表, messages来自父类
        all_messages.extend(messages) # 把新的消息和旧的消息合并为一个list

        # 将数据同步写入到本地文件
        # 类对象不能直接写入文件, 如果直接写入是一堆二进制数据
        # 要将BaseMessage对象转成字典, 再借助json模块以json格式写入文件
        # 完整代码
        # new_message = []
        # for message in all_messages:
        #     new_message.append(message.to_dict())

        # 列表推导式
        # 流程: 遍历all_messages, 赋值给message, 调用message_to_dict()方法, 把返回值放入list
        new_message = [message_to_dict(message) for message in all_messages]

        # 将数据写入文件
        with open(self.file_path, "w", encoding="utf-8") as f:
            json.dump(new_message, f)

    @property # 该注解不能省略, 作用是把该方法变成成员属性, LangChain调用该方法是以属性的方式调用的
    def messages(self) -> list[BaseMessage]:
        try:
            with open(self.file_path, "r", encoding="utf-8") as f:
                messages_data = json.load(f) # 返回的是 list[字典]
                return messages_from_dict(messages_data)
        except FileNotFoundError:
            return []

    def clear(self) -> None:
        with open(self.file_path, "w", encoding="utf-8") as f:
            json.dump([], f)


# 历史会话容器 (由内存存储变成文件存储)
def get_history(session_id):
    return FileChatMessageHistory(session_id, "./chat_history")

from langchain_community.chat_models import ChatTongyi
from langchain_community.embeddings import DashScopeEmbeddings
from langchain_core.documents import Document
from langchain_core.output_parsers import StrOutputParser
from langchain_core.prompts import ChatPromptTemplate, MessagesPlaceholder
from langchain_core.runnables import RunnablePassthrough, RunnableWithMessageHistory, RunnableLambda

import config_data as config
from file_history_store import get_history

from vector_stores import VectorStoreService


def print_prompt(prompt):
    print("=" * 20)
    print(prompt.to_string())
    print("=" * 20)
    return prompt


class RagService(object):
    def __init__(self):
        # 构建向量存储对象
        self.vector_store = VectorStoreService(
            embedding=DashScopeEmbeddings(model=config.embedding_model_name,
                                          dashscope_api_key="sk-0af497cdd462495aa2f08dddaead1ded")
        )
        # 构建提示词模板
        self.prompt_template = ChatPromptTemplate.from_messages(
            [
                ("system", "以我提供的已知参考资料为主,简洁专业的回答用户问题。参考资料:{context}"),
                ("system", "并且我提供用户的对话历史记录,如下:"),
                MessagesPlaceholder("history"),
                ("user", "请回答用户提问:{input}"),
            ]
        )
        # 构建模型对象
        self.chat_model = ChatTongyi(model=config.chat_model_name, api_key="sk-0af497cdd462495aa2f08dddaead1ded")
        # 构建执行链
        self.chain = self.__get_chain()

    def __get_chain(self):
        """获取最终的执行链"""
        retriever = self.vector_store.get_retriever()

        def format_document(docs: list[Document]):
            if not docs:
                return "无相关参考资料"

            formatted_str = ""
            for doc in docs:
                formatted_str += f"文档片段:{doc.page_content} \n文档元数据:{doc.metadata}\n\n"

            return formatted_str

        # 提取字符串
        def format_for_retriever(value):
            return value["input"]

        # 重组参数
        def format_for_prompt_template(value):
            new_value = {}
            new_value["input"] = value["input"]["input"]
            new_value["history"] = value["input"]["history"]
            new_value["context"] = value["context"]
            return new_value

        # 基础链
        chain = (
                {
                    "input": RunnablePassthrough(),
                    "context": RunnableLambda(format_for_retriever) | retriever | format_document,
                } | RunnableLambda(format_for_prompt_template) | self.prompt_template | print_prompt | self.chat_model | StrOutputParser()
        )

        # 增强链
        conversation_chain = RunnableWithMessageHistory(
            chain,
            get_history,
            input_messages_key="input",
            history_messages_key="history",
        )

        return conversation_chain


if __name__ == '__main__':
    # session id配置
    session_config = {
        "configurable": {
            "session_id": "user_001"
        }
    }

    res = RagService().chain.invoke({"input": "我身高180厘米,尺码推荐"}, session_config)
    print(res)

app_qa.py 用户交互服务开发

import time
from rag import RagService
import streamlit as st
import config_data as config

# 标题
st.title("智能客服")
st.divider() # 分割线

if "message" not in st.session_state:
    st.session_state["message"] = [{"role": "assistant", "content": "你好,有什么可以帮你?"}]

for message in st.session_state["message"]:
    st.chat_message(message["role"]).write(message["content"])

if "rag" not in st.session_state:
    st.session_state["rag"] = RagService()

# 用户输入栏
prompt = st.chat_input()


if prompt:
    # 用户输入提问
    st.chat_message("user").write(prompt)
    st.session_state["message"].append({"role": "user", "content": prompt})

    ai_res_list = []
    with st.spinner("AI思考中..."):
        res_stream = st.session_state["rag"].chain.stream({"input": prompt}, config.session_config)

        def capture(generator, cache_list):
            for chunk in generator:
                cache_list.append(chunk)
                yield chunk

        st.chat_message("assistant").write(capture(res_stream, ai_res_list))
        st.session_state["message"].append({"role": "assistant", "content": "".join(ai_res_list)})

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐