前言

嘿,这里是惬鹤频道!
这几天有点忙,其实RAG项目的代码我已经完成了,但是一直没时间来给大家分享,这里和大家道个歉。
这次分享的是两个代码文件:rag和vector_stores,分别是实现RAG的关键代码和返回向量检索器的方法。剩余的代码我会在之后再写一篇文章发出来。
所以,不多说了,我们开始吧!
哦对了,关于上期的两个代码文件,在项目完成后我检查了一下,发现我在原本的代码中又加上了一些东西,并不是最终版本,后续我会一起发出来的!

正文

代码:rag.py

# 导入文件
import config_md5 as config
from vector_stores import VectorStore_service
from file_history_store import get_history
# 导入依赖
from langchain_community.embeddings import DashScopeEmbeddings
from langchain_community.chat_models.tongyi import ChatTongyi
from langchain_core.prompts import ChatPromptTemplate, MessagesPlaceholder
from langchain_core.runnables import RunnablePassthrough, RunnableWithMessageHistory, RunnableLambda
from langchain_core.documents import Document
from langchain_core.output_parsers import StrOutputParser

# 创建类
class RagService(object):
    def __init__(self):
        # 嵌入模型服务
        self.vector_service = VectorStore_service(DashScopeEmbeddings(model=config.embedding_model_name))

        # 提示词模板基础定义
        self.prompt_template = ChatPromptTemplate.from_messages([
            ("system","请主要根据你检索的信息回答问题。检索信息:{messages}"),
            ("system","并且这里提供了用户对话的历史记录"),
            MessagesPlaceholder("history"),
            ("human","用户问题:{input}"),
        ])

        # 聊天模型
        self.chat_model = ChatTongyi(model=config.chat_model_name,streaming=True)

        # 将由方法组建的链调用为自己的对象
        self.chain = self.__get_chain()

    def __get_chain(self):
        retrievers = self.vector_service.get_retrievers()

        # 方法:将retrievers返回的VectorStoreRetriever(本质为Document类型的列表)对象拼凑为字符串
        def format_document(document: list[Document]):
            if not document:
                return "没有可参考的资料"
            docs = ""
            for doc in document:
                docs = docs + f"文档片段:{doc.page_content}\n,源数据:{doc.metadata}\n\n"
            return docs

        # 方法:打印提示词
        def print_prompt(input_prompt):
            print("=" * 20)
            print(input_prompt.to_string())
            print("=" * 20)

            return input_prompt

        # 为了解决链的传输过程中格式不匹配的问题,这里有两个方法用于转换格式。
        def format_for_dic(input):
            return input["input"]

        def rebuild_dic(input):
            new_value = {}
            new_value["input"] = input["input"]["input"]
            new_value["history"] = input["input"]["history"]
            new_value["messages"] = input["messages"]
            return new_value

        # 构建链,其中给模板的是一个字典格式,包含一个子链
        chain = ({
            "input": RunnablePassthrough(),
            "messages":RunnableLambda(format_for_dic)| retrievers | format_document} | RunnableLambda(rebuild_dic) | self.prompt_template | print_prompt | self.chat_model | StrOutputParser()
                 )

        # 构建一个增强的链
        conversion_chain = RunnableWithMessageHistory(
            chain,
            get_history,
            input_messages_key="input",
            history_messages_key="history",
        )

        return conversion_chain

这个代码最重要的是定义了一个类RagService,其中定义了所有重要的部分:提示词,模型,增强链,并返回了一个组装好的链。代码中的一些关键配置(如模型名)写在了配置文件中方便管理。

代码vector_stores.py

"""
文件名:retriever
作用:创建一个类,定义一个文本分割器,并在类中定义一个返回retrievers的方法
"""
# 依赖导入
from langchain_chroma import Chroma
from langchain_community.embeddings import DashScopeEmbeddings
import config_md5

# 定义类
class VectorStore_service():
    def __init__(self,embeddings):
        self.embeddings = embeddings

        # 定义chroma对象
        self.chroma = Chroma(
            collection_name=config_md5.collection_name,
            embedding_function=embeddings,
            persist_directory=config_md5.persist_directory,
        )

    # 返回一个向量检索器,后续加入chain
    def get_retrievers(self):
        return self.chroma.as_retriever(search_kwargs={"k":config_md5.k_num})

这个文件就是在rag文件中使用的返回检索器的代码。代码中的一些关键配置(如K值)写在了配置文件中方便管理。

结尾

还有一些代码没有放出来,这几天忙完了我会抓紧写的,非常感谢大家的关注!如果你觉得有用,可以点个赞谢谢!

Logo

AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。

更多推荐