Day7:RAG项目实战(2)
·
前言
嘿,这里是惬鹤频道!
这几天有点忙,其实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值)写在了配置文件中方便管理。
结尾
还有一些代码没有放出来,这几天忙完了我会抓紧写的,非常感谢大家的关注!如果你觉得有用,可以点个赞谢谢!
AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。
更多推荐



所有评论(0)