开头

嘿,这里是惬鹤频道!

昨晚上学习了新的一课,现在我带来了项目的最新进展!

这次学习的是向量数据库相关的代码,内容挺多的,并且和上次的五个工具方法关系密切,
结构很是复杂,不过我在代码部分写了一些注释,对理解代码应该会有帮助。

那么接下来就进入代码展示环节,如果有任何问题欢迎在评论区提出。

总会有的。。。项目结构图

在这里插入图片描述
这次我们完成的就是右上角的部分和与它连接的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,避免重复

大家可以看出来,代码内部需要使用的变量,相关函数,基本都是从外部导入的(从配置文件,工具文件等),这就是结构化项目的魅力之处!

结尾

这次的代码分享就到这里了,这次学习到的东西真的挺多的,代码比较复杂,需要多理解。
如果对你有帮助,可以点个赞,点点关注,有任何问题的话,可以在评论区发表留言。
下次见!

Logo

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

更多推荐