Day10:从零开始写一个Agent项目(2)
开头
嘿,这里是惬鹤频道!
昨晚上学习了新的一课,现在我带来了项目的最新进展!
这次学习的是向量数据库相关的代码,内容挺多的,并且和上次的五个工具方法关系密切,
结构很是复杂,不过我在代码部分写了一些注释,对理解代码应该会有帮助。
那么接下来就进入代码展示环节,如果有任何问题欢迎在评论区提出。
总会有的。。。项目结构图

这次我们完成的就是右上角的部分和与它连接的factory
工厂文件:factory.py
在真正的大餐端上来前,先来点开胃小菜。
注意!这个文件内部有一座工厂!
"""
文件名:factory.py
描述:工厂文件,用来创建两个工厂:嵌入式模型工厂,和聊天模型工厂。
对比硬编码创建虽然复杂一点,但硬编码一旦需要替换模型,就必须在所有使用这些变量的文件中修改代码(如果分散在多处会很麻烦)。
而工厂模式将模型的选择和构造集中在一个地方,维护成本更低。
"""
# 文件依赖导入
from utils.config_handler import rag_conf
# 依赖导入
from abc import ABC, abstractmethod
from typing import Optional
from langchain_core.embeddings import Embeddings
from langchain_core.language_models import BaseChatModel
from langchain_community.embeddings import DashScopeEmbeddings
from langchain_community.chat_models.tongyi import ChatTongyi
# 工厂化的模型定义
class BaseModelFactory(ABC):
# 使用 @abstractmethod 装饰器声明一个抽象方法 generator
@abstractmethod
# 表示该方法可能返回两种类型之一(嵌入模型或聊天模型),也可能返回 None
def generator(self) -> Optional[Embeddings | BaseChatModel]:
pass
# ChatTongyi 是 BaseChatModel 的子类,所以返回类型符合基类声明的 BaseChatModel
class ChatModelFactory(BaseModelFactory):
def generator(self) -> Optional[Embeddings | BaseChatModel]:
return ChatTongyi(model=rag_conf["chat_model_name"])
# DashScopeEmbeddings 是 Embeddings 的子类,返回类型符合基类声明的 Embeddings
class EmbeddingsFactory(BaseModelFactory):
def generator(self) -> Optional[Embeddings | BaseChatModel]:
return DashScopeEmbeddings(model=rag_conf["embedding_model_name"])
# 创建两个用工厂方法创建的模型,这样就可以在其他文件中直接使用了。
chat_model = ChatModelFactory().generator()
embed_model = EmbeddingsFactory().generator()
这个工厂用于生产各类模型,这里是嵌入模型和聊天模型。它的结构如下:
BaseModelFactory (抽象基类)
├── ChatModelFactory (具体工厂)
│ └── generator() → ChatTongyi 实例
└── EmbeddingsFactory (具体工厂)
└── generator() → DashScopeEmbeddings 实例
有人可能不明白为什么要搞这么复杂,“直接用ChatTongyi硬编码不好吗?”
其实,如果是在学习阶段,或者使用量很少,那随便用用就算了。
但是,如果你定义的模型多起来,那么每次你想要修改模型内的参数时,所有用到这个模型变量的地方都要同步修改,
工作量一下子大起来了,而且不利于维护。
当然,在这样一个小型项目中使用工厂方法看上去有点小题大做,但是在中大型项目中,硬编码模型变量基本会被淘汰,
所以还是早点熟悉这种方法比较好。
(而且,这个工厂看上去比硬编码酷多了不是吗?)
代码文件:vector_store.py
这是本次分享的重点,从开头的依赖导入你就能知道结构有多复杂。
这个文件实现了一个向量数据库服务,其核心功能是:把本地文件夹里的文本或 PDF 文件,自动切分成小段,转换成向量,然后存到 Chroma 向量数据库里,方便以后做语义搜索。
同时,它还会用 MD5 记录哪些文件已经处理过,避免重复加载。
"""
文件名:vector_store.py
描述:用于创建向量数据库的基类
"""
# 文件依赖导入
from utils.config_handler import chroma_conf
from utils.path_tool import get_abs_path
from utils.file_handler import pdf_loader, txt_loader, listdir_with_allowed_type, get_file_md5_hex
from utils.logger_handler import logger
from model.factory import embed_model
# langchain依赖导入
from langchain_chroma import Chroma
from langchain_core.documents import Document
from langchain_text_splitters import RecursiveCharacterTextSplitter
import os
# 向量库服务基础类
class VectorStoreService:
def __init__(self):
# 定义chroma数据库
self.vector_store = Chroma(
collection_name=chroma_conf["collection_name"],
embedding_function=embed_model,
persist_directory=chroma_conf["persist_directory"],
)
# 定义递归文本分割器
self.spliter = RecursiveCharacterTextSplitter(
chunk_size=chroma_conf["chunk_size"],
chunk_overlap=chroma_conf["chunk_overlap"],
separators=chroma_conf["separators"],
length_function=len,
)
# 返回检索器(k是检索的最大数量)
def get_retriever(self):
return self.vector_store.as_retriever(search_kwargs={"k": chroma_conf["k"]})
# 从数据文件夹内读取数据文件,转为向量存入数据库(要计算MD5值做去重)
def load_document(self):
# 方法:检查输入的文档是否已经处理过(md5_for_chunk是需要传入并检查的字符串)
def check_md5_hex(md5_for_chunk: str):
if not os.path.exists(get_abs_path(chroma_conf["md5_hex_store"])):
# 如果找不到存储MD5的文件,则创建这个文件(方法是执行打开操作,默认如果不存在,那么打开操作相当于创建。)
open(get_abs_path(chroma_conf["md5_hex_store"]), "w", encoding="utf-8").close()
return False
# 通过检验后,进行读取,验证
with open(get_abs_path(chroma_conf["md5_hex_store"]), "r", encoding="utf-8") as f:
for line in f.readlines():
line = line.strip()
if line == md5_for_chunk:
return True
return False
# 方法:检查通过后,写入md5文档,下一次就可以检测到了
def save_md5_hex(md5_for_chunk: str):
with open(get_abs_path(chroma_conf["md5_hex_store"]), "a", encoding="utf-8") as f:
f.write(md5_for_chunk + "\n")
# 方法:得到所有的documents
def get_file_documents(read_path: str):
# 以txt结尾
if read_path.endswith("txt"):
return txt_loader(read_path)
if read_path.endswith("pdf"):
return pdf_loader(read_path, None)
# 啥也不是,返回空列表
return []
# 判断路径中有哪些符合类型的文件
allowed_files_path: list[str] = listdir_with_allowed_type(
get_abs_path(chroma_conf["data_path"]),
tuple(chroma_conf["allow_knowledge_file_type"]),
)
# 处理文件列表
for path in allowed_files_path:
# 获取文件的MD5
md5_hex = get_file_md5_hex(path)
# 如果检测到文件已经被加载过
if check_md5_hex(md5_hex):
logger.info(f"[加载知识库]路径{path}对应的内容已经存在库中,已跳过。")
continue
try:
# 尝试读取路径中的document内容
documents: list[Document] = get_file_documents(path)
# 警告一:路径内没有documents
if not documents:
logger.warning(f"[加载知识库]路径{path}内没有有效内容,已跳过")
continue
# 开始尝试分片
spliter_document: list[Document] = self.spliter.split_documents(documents)
# 警告二:分片后没有有效内容
if not spliter_document:
logger.warning(f"[加载知识库]路径{path}分片后没有有效内容,已跳过")
continue
# 把内容存进向量库
self.vector_store.add_documents(spliter_document)
# 记录已经处理好的文件的MD5,避免下次重复加载
save_md5_hex(md5_hex)
logger.info(f"[加载知识库]当前路径{path}内容加载成功!")
# exc_info=True时,会记录详细的堆栈报错
except Exception as e:
logger.error(f"[知识库加载]路径{path}加载失败,原因:{str(e)}", exc_info=True)
continue
# 测试
if __name__ == '__main__':
vs = VectorStoreService()
vs.load_document()
retriever = vs.get_retriever()
res = retriever.invoke("迷路")
for r in res:
print(r.page_content)
print("="*20)
我在代码中添加了很多注释,方便大家理解。这个文件主要的部分就是load_document方法,其内部还有一些子方法,
不过都是为它自己服务的。
这个文件负责把本地文件夹中的 TXT 和 PDF 文件自动转换成向量,并存入 Chroma 数据库。
它通过 MD5 记录哪些文件已处理过,支持增量更新,避免重复加载。
基本结构:
class VectorStoreService:
def __init__(self):
# 创建 Chroma 向量库
# 创建文本分割器
def get_retriever(self):
# 返回检索器(用于查询)
def load_document(self):
# 内部定义三个辅助函数:
# - check_md5_hex():检查文件是否处理过
# - save_md5_hex():记录已处理文件的MD5
# - get_file_documents():加载txt或pdf为Document
# 主流程:
# 1. 扫描数据文件夹,得到所有待处理文件
# 2. 对每个文件计算MD5,若未处理则加载、分割、存入向量库
# 3. 记录MD5,避免重复
大家可以看出来,代码内部需要使用的变量,相关函数,基本都是从外部导入的(从配置文件,工具文件等),这就是结构化项目的魅力之处!
结尾
这次的代码分享就到这里了,这次学习到的东西真的挺多的,代码比较复杂,需要多理解。
如果对你有帮助,可以点个赞,点点关注,有任何问题的话,可以在评论区发表留言。
下次见!
AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。
更多推荐



所有评论(0)