RAG04-基于Milvus的RAG系统

基于Milvus的RAG系统

1 整体架构与工程流程

1.1 学习目标 ¶

  • 1.理解RAG系统的基本原理及其在教育领域的应用场景。

  • 2.掌握EduRAG系统的模块化设计和各模块的核心功能。

  • 3.熟悉RAG系统从查询到生成回答的完整工作流程。

1.2 RAG系统整体架构介绍 ¶

1 系统背景 ¶

EduRAG智慧问答系统是一个基于 RAG(Retrieval-Augmented Generation,检索增强生成) 技术的智能问答平台,专为IT教育培训设计。EduRAG系统采用双层架构:FQA系统 + RAG系统。FQA系统负责高频问答对匹配,RAG系统负责 检索知识库 和 LLM生成最终答案。

flowchart LR;
%% =========================
    %% 主数据流
    %% =========================
    Query([用户问题query]) --> FQA系统
    FQA系统 -->|不超过阈值| RAG系统
    RAG系统 --> FinalAnswer([最终答案])
    FQA系统 -->|超过阈值| FinalAnswer([最终答案])

2 模块化架构 ¶

系统的代码组织分为以下几个核心模块:

  • base/ :基础支持模块,负责配置、日志处理。

  • core/ :核心逻辑模块,实现RAG的关键功能。

  • main.py :系统运行入口,支持数据处理和交互查询。

3 代码目录结构 ¶

integrated_qa_system/
├── config.ini                 # 配置文件,包含所有模块的配置
├── base/
│   ├── config.py              # 配置管理,加载 config.ini
│   ├── logger.py              # 日志设置
├── rag_qa/
│   ├── core/
│   │   ├── document_processor.py # 文档处理模块
│   │   ├── prompts.py         # RAG 提示模板
│   │   ├── query_classifier.py # 查询分类器
│   │   ├── strategy_selector.py # 检索策略选择器
│   │   ├── vector_store.py    # 向量存储与检索
│   │   ├── rag_system.py      # RAG 系统核心逻辑
│   ├── edu_document_loaders/  # 文档加载模块
│   ├── edu_text_spliter/      # 文本分割模块
│   ├── models/                # 模型模块
│   ├── rag_assesment/         # RAGAS评估模块
│   ├── main.py                # RAG 系统独立入口,支持存储和查询
├── requirements.txt           # 依赖文件
└── logs/
    └── app.log                # 日志文件

1.3 RAG系统基本工作流程 ¶

1 步骤 ¶

RAG系统的工作流程分为5个步骤:

  • 0.文档向量存储

    • 文档加载 -> 文档分块 -> 文档向量化 -> 保存到向量数据库中
  • 1.查询分类 :

    • 调用本地BERT判断查询类型(如“通用知识”或“专业咨询”)。

    • 通用知识直接由大语言模型回答,专业咨询进入检索流程。

  • 2.检索策略选择:

    • 直接检索 :适用于明确查询。

    • 假设问题检索 :适用于抽象问题,生成假设答案后检索。

    • 子查询检索 :分解复杂查询。

    • 回溯检索 :简化复杂问题后检索。

  • 3.检索向量库 :

    • 使用 vector_store.py 从向量数据库中检索相关文档。

    • 支持稠密向量和稀疏向量的混合检索,结果经过重排序优化。

  • 4.生成回答 :

    • 将检索到的文档作为上下文,结合用户查询输入大语言模型。

    • 生成自然语言回答,若无答案则引导人工支持。

2 流程图 ¶

flowchart LR;
%%{init:{"themeVariables":{"fontSize":"18px","nodeSpacing":15,"rankSpacing":15}}}%%

    %% =========================
    %% 在线阶段
    %% =========================
    subgraph 用户提问阶段-在线

        direction TB

        Start([用户查询]) --> Classify{查询分类<br/>bert意图识别模型}

        Classify -->|通用知识| DirectLLM["直接调用LLM<br/>生成回答"]

        Classify -->|专业咨询| Strategy{选择检索策略<br/>LLM}

        Strategy -->|明确查询| DirectRetrieval["直接检索"]

        Strategy -->|抽象问题| HyDERetrieval["假设问题检索<br/>生成假设答案"]

        Strategy -->|复杂查询| SubQuery["子查询检索"]

        Strategy -->|啰嗦问题| BacktrackRetrieval["回溯问题检索"]

        DirectRetrieval --> OptimizeQuery["LLM优化query"]
        HyDERetrieval --> OptimizeQuery
        SubQuery --> OptimizeQuery
        BacktrackRetrieval --> OptimizeQuery

        OptimizeQuery --> Search["检索向量库<br/>得到上下文context"]

        Search --> BuildPrompt["拼接Prompt<br/>优化query+context+history"]

        BuildPrompt --> Generate["LLM生成回答"]

        DirectLLM --> Output([返回回答])

        Generate --> HasAnswer["是否有答案<br/>无答案则引导人工支持"]

        HasAnswer --> Output

    end

    style 用户提问阶段-在线 fill:none,stroke:#000000

    %% =========================
    %% 离线阶段
    %% =========================
    subgraph 文档向量存储阶段-离线

        direction TB

            StartDoc([本地知识文档])

            StartDoc --> PDFLoader["PDF加载器"]
            StartDoc --> DOCLoader["DOC加载器"]
            StartDoc --> TXTLoader["TXT加载器"]
            StartDoc --> PPTLoader["PPT加载器"]

            PDFLoader --> Document["文档对象"]
            DOCLoader --> Document
            TXTLoader --> Document
            PPTLoader --> Document

            Document --> ParentChunk["父块分块"]

            ParentChunk --> SubChunk["子块分块"]

            SubChunk --> SubChunkList["子块集合"]

            SubChunkList --> Embed["文档向量化<br/>(bge-m3本地模型)"]

            Embed --> Milvus[(Milvus)]

            classDef milvusLarge font-size:32px,stroke-width:5px;
            class Milvus milvusLarge;

    end

    style 文档向量存储阶段-离线 fill:none,stroke:#000000

    %% =========================
    %% 在线检索连接离线向量库
    %% =========================
    Milvus --> Search

1.4 本章小结 ¶

EduRAG系统的RAG系统,实现了从查询分类到生成回答的完整流程。架构包含:

离线存储部分: 文档加载 -> 文档分块 -> 文档向量化 -> 保存到向量数据库中

在线问答部分: 用户问题 -> 查询分类 -> 选择检索策略 -> 根据检索策略优化query -> 根据优化后query,检索向量库相关文档得到检索上下文context -> 构造prompt(query+context+history)-> 调用LLM生成回答

2 基础模块

2.1 学习目标 ¶

  • 1.理解并掌握如何通过Config类集中管理系统的配置参数。

  • 2.学会配置和使用日志记录器,实现对系统运行状态的监控。

  • 3.认识base模块在系统架构中的作用,为学习后续核心逻辑奠定基础。

base 模块是EduRAG智慧问答系统的基础,负责提供系统运行所需的核心功能,包括配置管理、日志记录。这些功能为系统的其他模块提供了稳定的支持,确保系统能够灵活配置、监控运行状态。

2.2 配置管理 ¶

1 功能 ¶

config.py 文件定义了 Config 类,用于集中管理系统中的所有配置参数。这些参数包括数据库连接信息、模型选择、分块策略、API设置等。通过集中管理配置,系统可以方便地调整参数、适配不同环境,并支持通过环境变量进行灵活配置。

2 代码实现 ¶

# base/config.py
# 导入配置解析库
import configparser
# 导入路径操作库
import os

# print(__file__)
# print(os.path.dirname(__file__))

class Config:
    # 初始化配置,加载 config.ini 文件
    def __init__(self, config_file=None):
        # 创建配置解析器,启用插值功能
        self.config = configparser.ConfigParser(interpolation=configparser.ExtendedInterpolation())
        # 如果没有提供配置文件路径,则使用默认路径
        # __file__: 当前文件的路径
        # os.path.dirname: 当前文件的上一级目录

        # self.PROJECT_ROOT = os.path.dirname(os.path.dirname(__file__))
        self.PROJECT_ROOT = os.getcwd()

        self.LOG_DIR = os.path.join(self.PROJECT_ROOT, 'logs')
        # self.DATA_DIR = os.path.join(self.PROJECT_ROOT, 'rag_qa/data')
        self.MODELS_DIR = os.path.join(self.PROJECT_ROOT, 'rag_qa/models')
        self.EDU_DOCUMENT_LOADERS_DIR = os.path.join(self.PROJECT_ROOT, 'rag_qa/edu_document_loaders')

        if config_file is None:
            config_file = os.path.join(self.PROJECT_ROOT, 'config.ini')
        # 读取配置文件
        self.config.read(config_file)

        # MySQL 配置
        # MySQL 主机地址
        self.MYSQL_HOST = os.getenv('MYSQL_HOST', self.config.get('mysql', 'host', fallback='localhost'))
        # MySQL 用户名
        self.MYSQL_USER = os.getenv('MYSQL_USER', self.config.get('mysql', 'user', fallback='edu_rag'))
        # MySQL 密码
        self.MYSQL_PASSWORD = os.getenv('MYSQL_PASSWORD', self.config.get('mysql', 'password', fallback='123456'))
        # MySQL 数据库名
        self.MYSQL_DATABASE = os.getenv('MYSQL_DATABASE', self.config.get('mysql', 'database', fallback='subjects_kg'))

        # Redis 配置
        # Redis 主机地址
        self.REDIS_HOST = os.getenv('REDIS_HOST', self.config.get('redis', 'host', fallback='localhost'))
        # Redis 端口
        self.REDIS_PORT = int(os.getenv('REDIS_PORT', self.config.get('redis', 'port', fallback=6379)))
        # Redis 密码
        self.REDIS_PASSWORD = os.getenv('REDIS_PASSWORD', self.config.get('redis', 'password', fallback='1234'))
        # Redis 数据库编号
        self.REDIS_DB = int(os.getenv('REDIS_DB', self.config.get('redis', 'db', fallback=0)))

        # Milvus 配置
        # Milvus 主机地址
        self.MILVUS_HOST = os.getenv('MILVUS_HOST', self.config.get('milvus', 'host', fallback='localhost'))
        # Milvus 端口
        self.MILVUS_PORT = os.getenv('MILVUS_PORT', self.config.get('milvus', 'port', fallback='19530'))
        # Milvus 数据库名
        self.MILVUS_DATABASE_NAME = os.getenv('MILVUS_DATABASE_NAME',
                                              self.config.get('milvus', 'database_name', fallback='itcast'))
        # Milvus 集合名
        self.MILVUS_COLLECTION_NAME = os.getenv('MILVUS_COLLECTION_NAME',
                                                self.config.get('milvus', 'collection_name', fallback='edurag_xian_1'))

        # LLM 配置
        # LLM 模型名
        self.LLM_MODEL = self.config.get('llm', 'model', fallback='qwen-plus')
        # DashScope API 密钥
        self.DASHSCOPE_API_KEY = os.getenv('DASHSCOPE_API_KEY', self.config.get('llm', 'dashscope_api_key',
                                                                                fallback='sk-b7314bb9c71a444293456ce5dcedf57e'))
        # DashScope API 地址
        self.DASHSCOPE_BASE_URL = self.config.get('llm', 'dashscope_base_url',
                                                  fallback='https://dashscope.aliyuncs.com/compatible-mode/v1')

        # 检索参数
        # 父块大小
        self.PARENT_CHUNK_SIZE = self.config.getint('retrieval', 'parent_chunk_size', fallback=1200)
        # 子块大小
        self.CHILD_CHUNK_SIZE = self.config.getint('retrieval', 'child_chunk_size', fallback=300)
        # 块重叠大小
        self.CHUNK_OVERLAP = self.config.getint('retrieval', 'chunk_overlap', fallback=50)
        # 检索返回数量
        self.RETRIEVAL_K = self.config.getint('retrieval', 'retrieval_k', fallback=5)
        # 最终候选数量
        self.CANDIDATE_M = self.config.getint('retrieval', 'candidate_m', fallback=2)

        # 应用配置
        # 有效来源列表
        self.VALID_SOURCES = eval(
            self.config.get('app', 'valid_sources', fallback='["ai", "java", "test", "ops", "bigdata"]'))
        # 客服电话
        self.CUSTOMER_SERVICE_PHONE = self.config.get('app', 'customer_service_phone', fallback='12345678')

        # 日志文件路径
        self.LOG_FILE = os.path.join(self.LOG_DIR, 'app.log')


config = Config()

if __name__ == '__main__':
    conf = Config()
    print(conf.MYSQL_USER)

3 说明 ¶

  • 环境变量支持 :使用 dotenv 加载 .env 文件中的环境变量,避免敏感信息硬编码。

  • 默认值 :每个参数设有默认值,确保未配置环境变量时系统仍可运行。

  • 参数分类 :按功能分类(如数据库、模型、分块等),便于管理和维护。

2.3 日志记录(logger.py) ¶

1 功能 ¶

logger.py 文件定义了 setup_logging 函数,用于配置系统的日志记录器。日志记录器将运行信息、警告和错误输出到文件和控制台,便于开发、调试和运维人员监控系统状态。

2 代码实现 ¶

# base/logger.py
# 导入日志库
import logging
# 导入路径操作库
import os
# 导入配置类(从base.config中已定义的config对象)
# from base.config import config

def setup_logger(logger_name='EduRAG', logger_file=config.LOG_FILE):
    # 1. 确保日志目录存在,不存在创建
    # logger_file 是日志的文件名,所以我们先取到它所在的目录 os.path.dirname
    dirname = os.path.dirname(logger_file)
    if not os.path.exists(dirname):
        os.makedirs(dirname)
    # 2. 创建日志记录器: Logger
    # 2.1 获取Logger对象
    logger = logging.getLogger(logger_name)
    # 2.2 设置日志级别为所有控制器最低的(设置全局的日志级别)
    logger.setLevel(logging.DEBUG)
    if not logger.handlers:
        # 3. 创建控制台控制器:StreamHandler
        # 3.1 创建控制台处理器对象
        stream_handler = logging.StreamHandler()
        # 3.2 设置日志级别为INFO
        stream_handler.setLevel(logging.INFO)
        # 4. 创建文件处理器:FileHandler,并指定目录
        # 4.1 创建文件处理对象
        file_handler = logging.FileHandler(logger_file, mode='a', encoding='utf-8')
        # 4.2 设置日志级别为DEBUG
        file_handler.setLevel(logging.DEBUG)
        # 5. 定义并设置日志格式:
        # 5.1 定义日志格式:logging.Formatter('%(asctime)s - %(levelname)s - %(name)s - %(message)s')
        formatter = logging.Formatter('%(asctime)s - %(levelname)s - %(pathname)s - %(funcName)s - %(module)s - %(lineno)d - %(message)s')
        # 5.2 设置处理器日志格式
        stream_handler.setFormatter(formatter)
        file_handler.setFormatter(formatter)
        # 6. 把处理器添加到logger中
        logger.addHandler(stream_handler)
        logger.addHandler(file_handler)
    return logger

logger = setup_logger('EduRAG')

3 说明 ¶

  • 日志级别 :默认设为 INFO ,记录关键运行信息。

  • 双重输出 :同时输出到文件和控制台,便于实时监控和后续分析。

  • 格式化 :日志包含时间戳、名称、级别和内容,便于问题定位。

2.4 本章小结 ¶

base 模块为EduRAG系统提供了以下核心支持:

  • 配置管理 :通过 Config 类实现灵活的参数配置。

  • 日志记录 :通过 logger 实现运行状态的实时监控和记录。

本章内容为学习者理解EduRAG系统的基础功能奠定了基础,为后续深入学习核心逻辑提供了支持。

3 文档处理模块

3.1 学习目标 ¶

  • 1.了解不通类型文档处理的基本逻辑。

  • 2.掌握文档加载和分块的基本原理。

3.2 文档解析 ¶

1 基本介绍 ¶

document_processor.py 是EduRAG系统的核心模块之一,用于文档解析。
主要负责加载多种格式的文档(如 .txt 、 .pdf 等),并对其进行分层切分,生成父块和子块,为后续的向量存储和检索做好准备。

2 代码实现 ¶

# core/document_processor.py
import os
from langchain_community.document_loaders import TextLoader
from langchain_community.document_loaders.markdown import UnstructuredMarkdownLoader
from langchain_text_splitters import MarkdownTextSplitter
from datetime import datetime
from rag_qa.edu_text_spliter import ChineseRecursiveTextSplitter
from rag_qa.edu_text_spliter import AliTextSplitter
from rag_qa.edu_document_loaders import OCRPDFLoader, OCRDOCLoader, OCRPPTLoader, OCRIMGLoader
from base.config import config
from base.logger import logger

# 使用全局config对象
conf = config

# 定义支持的文件类型及其对应的加载器字典
document_loaders = {
    # 文本文件使用 TextLoader
    ".txt": TextLoader,
    # PDF 文件使用 OCRPDFLoader
    ".pdf": OCRPDFLoader,
    # Word 文件使用 OCRDOCLoader
    ".docx": OCRDOCLoader,
    # PPT 文件使用 OCRPPTLoader
    ".ppt": OCRPPTLoader,
    # PPTX 文件使用 OCRPPTLoader
    ".pptx": OCRPPTLoader,
    # JPG 文件使用 OCRIMGLoader
    ".jpg": OCRIMGLoader,
    # PNG 文件使用 OCRIMGLoader
    ".png": OCRIMGLoader,
    # Markdown 文件使用 UnstructuredMarkdownLoader
    ".md": UnstructuredMarkdownLoader
}

# 定义函数,从指定文件夹加载多种类型文件并添加元数据
def load_documents_from_directory(directory_path):
    # 初始化空列表,用于存储加载的文档
    documents = []
    # 获取支持的文件扩展名集合
    supported_extensions = document_loaders.keys()
    # 从目录名提取学科类别(如 "ai_data" -> "ai")
    source = os.path.basename(directory_path).replace("_data", "")
    # 遍历指定目录及其子目录
    for root, _, files in os.walk(directory_path):
        # 遍历当前目录下的所有文件
        for file in files:
            # 构造文件的完整路径
            file_path = os.path.join(root, file)
            # 获取文件扩展名并转换为小写
            file_extension = os.path.splitext(file_path)[1].lower()
            # 检查文件类型是否在支持的扩展名列表中
            if file_extension in supported_extensions:
                # 使用 try-except 捕获加载过程中的异常
                try:
                    # 根据文件扩展名获取对应的加载器类
                    loader_class = document_loaders[file_extension]
                    # 实例化加载器对象,传入文件路径
                    if file_extension == ".txt":
                        loader = loader_class(file_path, encoding="utf-8")
                    else:
                        loader = loader_class(file_path)
                    # 调用加载器加载文档内容,返回文档列表
                    loaded_docs = loader.load()
                    # 遍历加载的每个文档
                    for doc in loaded_docs:
                        # 为文档添加学科类别元数据
                        doc.metadata["source"] = source
                        # 为文档添加文件路径元数据
                        doc.metadata["file_path"] = file_path
                        # 为文档添加当前时间戳元数据
                        doc.metadata["timestamp"] = datetime.now().isoformat()
                    # 将加载的文档添加到总列表中
                    documents.extend(loaded_docs)
                    # 记录成功加载文件的日志
                    logger.info(f"成功加载文件: {file_path}")
                # 捕获加载过程中可能出现的异常
                except Exception as e:
                    # 记录加载失败的日志,包含错误信息
                    logger.error(f"加载文件 {file_path} 失败: {str(e)}")
            # 如果文件类型不在支持列表中
            else:
                # 记录警告日志,提示不支持的文件类型
                logger.warning(f"不支持的文件类型: {file_path}")
    # 返回加载的所有文档列表
    return documents

# 定义函数,处理文档并进行分层切分,返回子块结果
def process_documents(directory_path, parent_chunk_size=conf.PARENT_CHUNK_SIZE,
                     child_chunk_size=conf.CHILD_CHUNK_SIZE,
                     chunk_overlap=conf.CHUNK_OVERLAP):
    # 从指定目录加载所有文档
    documents = load_documents_from_directory(directory_path)
    # 记录加载的文档总数日志
    logger.info(f"加载的文档数量: {len(documents)}")
    # 初始化父块和子块分词器(通用)
    parent_splitter = ChineseRecursiveTextSplitter(chunk_size=parent_chunk_size, chunk_overlap=chunk_overlap)
    child_splitter = ChineseRecursiveTextSplitter(chunk_size=child_chunk_size, chunk_overlap=chunk_overlap)
    # 初始化 Markdown 专用分词器
    markdown_parent_splitter = MarkdownTextSplitter(chunk_size=parent_chunk_size, chunk_overlap=chunk_overlap)
    markdown_child_splitter = MarkdownTextSplitter(chunk_size=child_chunk_size, chunk_overlap=chunk_overlap)
    # 初始化空列表,用于存储所有子块
    child_chunks = []
    # 遍历每个原始文档,带上索引 i
    for i, doc in enumerate(documents):
        # print(doc)
        # 获取文件扩展名
        file_extension = os.path.splitext(doc.metadata.get("file_path", ""))[1].lower()
        # 选择切分器
        is_markdown = (file_extension == ".md")
        parent_splitter_to_use = markdown_parent_splitter if is_markdown else parent_splitter
        # print(f'parent_splitter_to_use-->{parent_splitter_to_use}')
        child_splitter_to_use = markdown_child_splitter if is_markdown else child_splitter
        logger.info(f"处理文档: {doc.metadata['file_path']}, 使用切分器: {'Markdown' if is_markdown else 'ChineseRecursive'}")
        # 使用父块分词器将文档切分为父块
        parent_docs = parent_splitter_to_use.split_documents([doc])
        # 遍历每个父块,带上索引 j
        for j, parent_doc in enumerate(parent_docs):
            # 为父块生成唯一 ID,格式为 "doc_i_parent_j"
            parent_id = f"doc_{i}_parent_{j}"
            # 将父块 ID 添加到元数据
            parent_doc.metadata["parent_id"] = parent_id
            # 将父块内容存储到元数据
            parent_doc.metadata["parent_content"] = parent_doc.page_content
            # 使用子块分词器将父块切分为子块
            sub_chunks = child_splitter_to_use.split_documents([parent_doc])
            # 遍历每个子块,带上索引 k
            for k, sub_chunk in enumerate(sub_chunks):
                # 为子块添加父块 ID 到元数据
                sub_chunk.metadata["parent_id"] = parent_id
                # 为子块添加父块内容到元数据
                sub_chunk.metadata["parent_content"] = parent_doc.page_content
                # 为子块生成唯一 ID,格式为 "parent_id_child_k"
                sub_chunk.metadata["id"] = f"{parent_id}_child_{k}"
                # 将子块添加到子块列表中
                child_chunks.append(sub_chunk)
    # 记录子块总数日志
    logger.info(f"子块数量: {len(child_chunks)}")
    # 返回所有子块列表
    return child_chunks

if __name__ == '__main__':
    chunks = process_documents(
        '/Users/ligang/PycharmProjects/LLM/ITCAST_EduRAG/data/ai_data',
        conf.PARENT_CHUNK_SIZE,
        conf.CHILD_CHUNK_SIZE,
        conf.CHUNK_OVERLAP,
    )
    print(chunks)

3.3 说明 ¶

  • 文档加载 :支持多种格式(如 .txt 、 .pdf ),使用专用加载器处理复杂文档。

  • 分层切分 :采用 ChineseRecursiveTextSplitter 生成父块和子块,优化中文文本处理。

  • 元数据管理 :为每个块添加唯一ID、来源和时间戳,便于检索和溯源。

3.4 本章小结 ¶

document_processor 模块为EduRAG系统提供了以下核心支持:

  • 文档处理 :实现多格式文档的高效加载和分块。

  • 图片或表格 :采用paddleOCR实现高效识别。

4 向量存储模块

1 学习目标 ¶

  • 1.理解向量存储在RAG系统中的功能和重要性。

  • 2.学会创建和管理向量数据库。

  • 3.掌握如何将文本转化为向量并存入数据库。

  • 4.理解混合检索与重排序的实现原理。

2 模块功能概述 ¶

vector_store.py 是EduRAG系统的核心模块之一,封装了与Milvus向量数据库的交互逻辑。它负责将文档转化为向量并存储到数据库中,并提供高效的混合检索功能。结合BGE-M3嵌入模型和重排序机制,确保系统能够快速检索到与用户查询最相关的文档。

VectorStore 类提供了以下主要功能:

  • 初始化与集合管理 :创建或加载Milvus向量数据库集合。

  • 文档向量化与存储 :将分块后的文档转换为向量并存储。

  • 混合检索与重排序 :结合稠密和稀疏向量进行检索,并通过重排序优化结果。

以下将逐一讲解每个方法的实现细节。
导入必备工具包:

"""
实现VectorStore
作用:
    1.初始化 Milvus
    2.创建或加载 database 和 collection
    3.文档向量化 并写入 Milvus
    4.根据用户问题 进行混合检索(稠密向量+稀疏向量) + 重排/精排
注意:
    1.文档向量化:嵌入模型为bge-m3
    2.query向量化: 嵌入模型为bge-m3
    3.两层排序策略:先粗排(混合检索) + 再精排(CrossEncoder,bge-reranker-large),效果与效率的均衡
细节:
    路径中不要有中文

常见的面试题:
Q1:什么是向量数据库?为什么用 Milvus?
A:向量数据库用来存储向量,并按相似度快速检索。Milvus 是专门做大规模向量检索的工具,速度快,索引多,适合 RAG 场景。

Q2:为什么嵌入模型用 bge-m3?
A:因为 bge-m3 可以同时生成 dense 和 sparse 两种向量:
- dense:稠密向量,每一维都有值,负责语义相似;
- sparse:稀疏向量,只有少量非0值,负责关键词匹配;
一个模型同时支持两种检索,方便又实用。

Q3:嵌入模型和精排模型有什么区别?可以用同一个吗?
A:
- 嵌入模型:把文本变成向量,用来召回候选结果;
- 精排模型:对候选结果重新打分,选出最相关的内容。
一般不建议用同一个模型,因为两者任务不同。

Q4:dense 向量和 sparse 向量有什么区别?
A:
- dense(稠密向量):每一维都有值,适合找“意思相近”的内容;
- sparse(稀疏向量):只有少量维度非零,适合找“关键词相同”的内容。
比如“怎么登录失败”和“无法进入系统”,dense 可能都能找到;“Navicat 乱码”,sparse 更容易直接命中。

Q5:为什么要做混合检索?
A:因为用户问题有时看重语义,有时看重关键词。混合检索把 dense 和 sparse 结合起来,召回更全面,效果通常更稳。

Q6:IVF_FLAT 是什么?nlist 和 nprobe 是什么?
A:IVF_FLAT 是一种常见向量索引。它先把向量分成多个簇,再只在部分簇里搜索。
- nlist:簇的数量,越大越细;
- nprobe:查询时找多少个簇,越大越准但越慢。

Q7:为什么 sparse_index 常用 SPARSE_INVERTED_INDEX + IP?
A:因为 sparse 向量本质上是“词项-权重”,倒排索引适合做关键词检索,IP(内积)适合做权重匹配,所以这是主流做法。

Q8:WeightedRanker 和 RRF 有什么区别?怎么选?
A:
- WeightedRanker:按权重加权融合相似度结果,简单直接,但要人为设置权重;
- RRF:按相似度排名融合,不太受相似度数值大小影响,更稳定。
如果想设置权重,可以用 WeightedRanker;如果追求稳定,RRF 更常见。

Q9:为什么要先混合检索,再用 CrossEncoder 精排?直接精排不行吗?
A:不行,通常效率太低。先混合检索快速召回候选,再用 CrossEncoder 精排,能兼顾速度和准确率。

Q10:CrossEncoder 是什么?精排逻辑是什么?
A:CrossEncoder 会把 query 和 doc 一起输入模型,直接输出相关性分数。分数越高,说明这个文档越适合回答问题。它常用于“精排”。

Q11:hashlib 是什么?这里有什么用?
A:hashlib 是 Python 的哈希工具。这里用 MD5 给文本生成稳定 ID,方便 upsert。相同文本会得到相同 ID。

Q12:为什么要先连 default,再创建业务库?
A:如果业务库不存在,直接连接可能失败。先连 default 更稳定,再创建并切换到业务库,流程更安全。

Q13:这个检索流程一句话怎么说?
A:先把文档转成 dense+sparse 向量,并写入 Milvus;查询时做混合检索,再用 CrossEncoder 精排,最后返回最相关的父文档。

"""

# core/vector_store.py
# 导入 bge-m3 嵌入函数,用于生成文档和查询的向量表示
from milvus_model.hybrid import BGEM3EmbeddingFunction
# 导入 Milvus 相关类,用于操作向量数据库
from pymilvus import MilvusClient, DataType, AnnSearchRequest, WeightedRanker, RRFRanker
# 导入 Document 类,用于创建文档对象
from langchain_core.documents import Document
# 导入 CrossEncoder(交叉编码器),用于重排序(精排)
# 输入是一对文本 [query, doc],模型让两段文本“同时参与注意力计算”,
# 直接输出一个相关性分数(分数越高,query 与 doc 越相关)。
from sentence_transformers import CrossEncoder
# 导入 hashlib 模块-哈希工具包,用于生成唯一 ID 的哈希值,任意长度的文本,都会生成一段固定长度的哈希字符串
import hashlib
# 导入 time 模块,用于生成时间戳
import time
from base.config import config
from base.logger import logger
import sys
import os
import torch

# 选择推理设备:cuda > cpu
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")

# 1.定义类,实现向量存储与检索
class VectorStore:
    # 1.初始化 Milvus
    def __init__(
        self, collection_name=config.MILVUS_COLLECTION_NAME, host=config.MILVUS_HOST,
        port=config.MILVUS_PORT, database = config.MILVUS_DATABASE_NAME,
    ):
        """
        初始化 Milvus
        1.构造 milvus连接参数
        2.创建 milvus客户端
        3.加载 embedding模型
        4.加载 rerank重排模型
        5.创建或加载 milvus集合

        :param collection_name: milvus集合名称
        :param host: milvus服务器地址
        :param port: milvus端口
        :param database: milvus数据库名称
        """
        # 1.构造 milvus连接参数
        self.collection_name = collection_name
        self.host = host
        self.port = port
        self.database = database
        self.logger = logger

        # 2.创建 milvus客户端
        # 先连接 milvus 默认数据库default, 避免目标数据库不存在而初始化失败
        self.client = MilvusClient(
            uri = f"http://{self.host}:{self.port}", db_name="default",
        )

        # 3.加载 embedding模型
        # bge-m3嵌入模型,同时生成 稠密向量 和 稀疏向量
        bge_m3_model_path = os.path.join(config.MODELS_DIR, "bge-m3")
        self.embedding_function = BGEM3EmbeddingFunction(
            model_name_or_path=bge_m3_model_path, # 传入模型名,自动下载;传入路径,直接加载
            use_fp16=False, # 是否使用FP16精度;True表示使用FP16精度(精度低,速度快),False表示使用FP32精度(精度高,速度慢)
            device=DEVICE, # 推理设备
        )
        # 获取稠密向量维度,由嵌入模型决定
        self.dense_dim = self.embedding_function.dim['dense']

        # 4.加载 rerank重排模型
        # bge-reranker-large, 用于重排(精排)
        rerank_model_path = os.path.join(config.MODELS_DIR, "bge-reranker-large")
        self.reranker = CrossEncoder(
            model_name=rerank_model_path,
            device=DEVICE,
        )

        # 5.创建或加载 milvus集合
        self._create_or_load_collection()

    # 2.创建或加载 database 和 collection
    def _create_or_load_collection(self):
        """
        创建或加载 database 和 collection
        1.创建并切换到指定database
        2.创建集合:如果集合不存在,创建字段与索引
        3.加载集合:如果集合已经存在
        :return: None
        """
        # 1.创建并切换到指定database
        if self.database not in self.client.list_databases():
            self.client.create_database(self.database)
            self.logger.info(f"创建数据库 {self.database} 成功")
        # 切换到指定database
        self.client.use_database(self.database)
        self.logger.info(f"切换到数据库 {self.database}")

        # 2.创建集合:
        # 判断集合是否已存在且有数据
        collection_exists = self.client.has_collection(self.collection_name)
        if collection_exists:
            # 检查已有集合中是否有数据
            stats = self.client.get_collection_stats(self.collection_name)
            row_count = int(stats.get("row_count", 0))
            if row_count > 0:
                # 集合已存在且有数据,直接复用,不删除
                self.logger.info(f"集合 {self.collection_name} 已存在且有 {row_count} 条数据,直接加载")
            else:
                # 集合存在但为空,删除后重建
                self.client.drop_collection(self.collection_name)
                self.logger.info(f"已删除空集合: {self.collection_name}")
                collection_exists = False

        # 如果集合不存在(或刚被删除),创建字段与索引
        if not collection_exists:
            # 1.创建schema
            schema = self.client.create_schema(
                auto_id=False, # 是否自动生成id,False表示主键不自增
                enable_dynamic_field=True, # 是否支持动态字段
            )
            # 2.设置field
            # 子块ID
            schema.add_field(field_name="id", datatype=DataType.VARCHAR, is_primary=True, max_length=100)
            # 子块文本内容
            schema.add_field(field_name="text", datatype=DataType.VARCHAR, max_length=65535)
            # 稠密向量,用于语义相似检索
            schema.add_field(field_name="dense_vector", datatype=DataType.FLOAT_VECTOR, dim=self.dense_dim)
            # 稀疏向量,用于关键词匹配
            schema.add_field(field_name="sparse_vector", datatype=DataType.SPARSE_FLOAT_VECTOR)
            # 父块ID
            schema.add_field(field_name="parent_id", datatype=DataType.VARCHAR, max_length=100)
            # 父块文本内容
            schema.add_field(field_name="parent_content", datatype=DataType.VARCHAR, max_length=65535)
            # 数据来源(学科名称),用于检索过滤
            schema.add_field(field_name="source", datatype=DataType.VARCHAR, max_length=50)
            # 时间戳(字符串形式)
            schema.add_field(field_name="timestamp", datatype=DataType.VARCHAR, max_length=50)

            # 3.创建索引
            # 创建索引参数
            index_params = self.client.prepare_index_params()
            # 为稠密向量添加 IVF_FLAT索引,相似度度量方式为IP
            # nlist: 聚类中心的数量。nlist 越大,索引构建越慢,但检索精度通常越高;nlist 越小,检索速度越快,但可能牺牲精度。
            # nprobe: 查询时搜索的聚类中心数量。nprobe 越大,召回率越高,但检索耗时增加;nprobe 越小,检索速度越快,但可能漏掉相关结果。
            index_params.add_index(
                field_name="dense_vector",
                index_name="dense_index",
                index_type="IVF_FLAT", # IVF_FLAT: 先聚类,再查询
                metric_type="IP",
                params={"nlist": 16}
            )

            # 为稀疏向量添加 SPARSE_INVERTED_INDEX索引,相似度度量方式为IP
            index_params.add_index(
                field_name="sparse_vector",
                index_name="sparse_index",
                index_type="SPARSE_INVERTED_INDEX",
                metric_type="IP",
                params={"drop_ratio_build":0.2} # drop_ratio_build: 稀疏向量索引构建时,对低贡献项的裁剪比例。drop_ratio_build 越小,索引构建越慢,但检索精度越高;drop_ratio_build 越大,索引构建越快,但可能丢失部分向量。
            )

            # 4.创建集合
            self.client.create_collection(
                collection_name=self.collection_name,
                schema=schema,
                index_params=index_params
            )
            self.logger.info(f"创建集合 {self.collection_name} 成功")

        else:
            logger.info(f"集合 {self.collection_name} 已经存在")

        # 3.加载集合:如果集合已经存在
        self.client.load_collection(self.collection_name)
        self.logger.info(f"加载集合 {self.collection_name} 成功")

    # 3.文档向量化 并写入 Milvus
    def add_documents(self, documents, batch_size=1000):
        """
        对文档进行向量化 并写入 Milvus
        1.分批处理文档
        2.批量生成向量(稠密/稀疏)
        3.组装数据
        4.upsert 写入 Milvus(同 ID 会覆盖)
        :param documents: 子块文档列表 list[Document]
        :param batch_size: 每批次处理文档的数量,避免一次性占用过多内存
        :return: None
        """
        # 1.分批处理文档
        total_docs = len(documents)
        logger.info(f"开始处理文档,文档总数:{total_docs},批次大小:{batch_size}")
        # 遍历批次,通过获取 start_idx 和 end_idx 来获取当前批次的文档
        for start_idx in range(0, total_docs, batch_size):
            # 获取当前批次的结束索引
            end_idx = min(start_idx + batch_size, total_docs)
            # 获取当前批次的文档列表
            batch_docs = documents[start_idx:end_idx]
            logger.info(f"正在处理文档范围: {start_idx} - {end_idx-1},批次大小:{end_idx - start_idx}")

            # 提取当前批次的文本page_content
            texts = [doc.page_content for doc in batch_docs]

            # 2.批量生成向量(稠密/稀疏)
            try:
                # 使用嵌入模型获取文档向量, bge-m3模型 会同时获取稠密向量 和 稀疏向量
                embeddings = self.embedding_function(texts)

                # 3.组装数据
                # dict列表,list[dict{id:,text:,dense_vector:,sparse_vector:,parent_id:,parent_content:,source:,timestamp:}]
                # 初始化数据列表
                data = []
                # 遍历当前批次的文档列表
                for i, child_doc in enumerate(batch_docs):
                    # 使用MD5生成子块文档ID
                    # MD5: 哈希算法,将任意长度的数据转换为固定长度字符串,32位十六进制字符串
                    text_hash = hashlib.md5(child_doc.page_content.encode("utf-8")).hexdigest()

                    # 获取稠密向量
                    dense_vector = embeddings["dense"][i].tolist()

                    # 获取稀疏向量,并构造 索引-权重 格式的dict
                    # {2:0.333,9:0.333,50:0.333}
                    sparse_row = embeddings["sparse"][i]
                    # 初始化稀疏向量字典
                    sparse_vector = {}
                    # 获取稀疏向量的非零值索引
                    indices = sparse_row.col
                    # 获取稀疏向量的非零值
                    values = sparse_row.data
                    # 将 索引 和 值 撇对,构造稀疏向量字典
                    for idx, weight in zip(indices, values):
                        sparse_vector[int(idx)] = float(weight)

                    # 组装数据并添加到数据列表中
                    data.append({
                        "id": text_hash,
                        "text": child_doc.page_content,
                        "dense_vector": dense_vector,
                        "sparse_vector": sparse_vector,
                        "parent_id": child_doc.metadata["parent_id"],
                        "parent_content": child_doc.metadata["parent_content"],
                        "source": child_doc.metadata.get("source", "unknown"),
                        "timestamp": child_doc.metadata.get("timestamp", "unknown")
                    })

                # 4.upsert 写入 Milvus(同 ID 会覆盖)
                # 检查是否有数据,如果有则写入 milvus
                if data:
                    # 使用upsert 写入 Milvus,覆盖同 ID 的数据
                    self.client.upsert(
                        collection_name=self.collection_name,
                        data=data
                    )
                    self.logger.info(f"当前批次文档 (索引:{start_idx}-{end_idx-1}) 写入 Milvus 成功,文档数量:{len(data)}")
                else:
                    self.logger.info(f"当前批次文档 (索引:{start_idx}-{end_idx-1}) 无数据,跳过写入 Milvus")
            except Exception as e:
                self.logger.error(f"处理文档 (索引:{start_idx}-{end_idx-1}) 时出错:{e}")
                continue

        self.logger.info(f"处理文档完成,总文档数:{total_docs}")

        # 刷新、释放、重新加载集合:
        # 1. flush:确保 upsert 的数据全部落盘
        # 2. release_collection:释放内存中旧的索引(该索引是在空集合上训练的)
        # 3. load_collection:重新加载集合,强制 IVF_FLAT 索引在有数据的情况下重新训练
        # 如果不做这一步,IVF_FLAT 索引是在空集合上初始化的,新插入的数据不在任何簇中,搜索会返回空结果
        self.client.flush(self.collection_name)
        self.client.release_collection(self.collection_name)
        self.client.load_collection(self.collection_name)
        self.logger.info(f"upsert 后 flush + release + reload 集合 {self.collection_name} 成功")

    # 4.根据用户问题 进行混合检索(稠密向量+稀疏向量) + 重排/精排
    def hybrid_search_with_rerank(self, query, top_k=config.RETRIEVAL_K, source_filter=None):
        """
        根据用户问题 进行混合检索 + 重排。
        1.查询向量化:对 query 生成稠密/稀疏向量
        2.分别执行 dense/sparse 两路检索
        3.用 WeightedRanker/RRFRanker 融合结果(粗排)
        4.子块查父块,对父块合并去重
        5.用 reranker 做精排,返回最终检索的父块结果
        :param query: 用户问题
        :param top_k: 粗排检索返回的候选数量
        :param source_filter: 可选来源过滤(比如指定学科)
        :return: 最终检索的父块结果
        """
        # 1.查询向量化:对 query 生成稠密/稀疏向量
        # 使用 写入Milvus 时的相同的嵌入模型bge-m3, 保证一样的文本的向量数值完全一致
        # bge-m3可以同时生成 稠密向量 和 稀疏向量
        query_embeddings = self.embedding_function([str(query)])

        # 获取稠密向量
        dense_query_vector = query_embeddings["dense"][0]

        # 获取稀疏向量
        sparse_row = query_embeddings["sparse"][0]
        # 初始化稀疏向量字典
        sparse_query_vector = {}
        # 获取稀疏向量的非零值索引
        indices = sparse_row.col
        # 获取稀疏向量的非零值
        values = sparse_row.data
        # 将 索引 和 值 撇对,构造稀疏向量字典
        for idx, weight in zip(indices, values):
            sparse_query_vector[int(idx)] = float(weight)

        # 可选来源过滤, 例如 source == 'ai'
        filter_expr = f"source == '{source_filter}'" if source_filter else None

        # 2.分别执行 dense/sparse 两路检索
        # 构建稠密检索请求:用于语义相似匹配
        dense_request = AnnSearchRequest(
            data=[dense_query_vector],
            anns_field="dense_vector",
            param={"metric_type": "IP", "params": {"nprobe": 4}}, # IP: 内积; nprobe: 查询最近的几个簇
            limit=top_k,
            expr=filter_expr
        )

        # 构建稀疏检索请求:用于关键词匹配
        sparse_request = AnnSearchRequest(
            data=[sparse_query_vector],
            anns_field="sparse_vector",
            param={"metric_type": "IP", "params": {}},  # IP: 内积
            limit=top_k,
            expr=filter_expr
        )

        # 3.用 WeightedRanker/RRFRanker 融合结果
        # 构建混合检索器
        # 加权排名:稠密向量0.7,稀疏向量1.0.这里认为稀疏向量更加重要
        # ranker = WeightedRanker(0.7,1.0)
        # 倒数融合排序RRFRanker: 基于倒数排序,不考虑相似度数值的绝对值,更加通用
        ranker = RRFRanker()

        # 查询结果为二维列表(n=1,limit=topK)
        # 1表示1个query,limit表示返回的topK个结果
        results = self.client.hybrid_search(
            collection_name=self.collection_name,
            reqs=[dense_request, sparse_request],
            ranker=ranker,
            limit=top_k,
            output_fields=["id", "text", "parent_id", "parent_content", "source", "timestamp"]
        )[0]

        # 把子块结果转换为 Document
        res_child_chunks = [self._doc_from_hit(hit['entity']) for hit in results]

        # 4.子块查父块,对父块合并去重
        res_parent_docs = self._get_unique_parent_docs(res_child_chunks)

        """
        以上完成了粗排,下面开始 重排/精排
        """

        # 5.用 reranker 做精排,返回config.CANDIDATE_M
        if res_parent_docs:
            # 如果只有一条,则直接返回
            if len(res_parent_docs) < 2:
                return res_parent_docs
            # 这里的 res_parent_docs 其实就是context
            # 如果父块超过一个,需要进行重排序:基于query 和context的匹配程度做重排序
            # 构造 (query, context) 对,一起送入reranker模型做 重排/精排:计算query和context的相关性,计算精度更高
            # 要求传入reranker的数据格式为 [[query, context1],[query, context2],[query, context3],...]
            # 形状(n,2),n=参与检索的父块数量
            pairs = [[query, doc.page_content] for doc in res_parent_docs]
            # 送入reranker模型做 重排/精排
            scores = self.reranker.predict(pairs) # 相似度分数,形状(n,)

            # 排序,从大到小排序scores
            ranked_parent_docs = [doc for score, doc in sorted(zip(scores, res_parent_docs), reverse=True)]

        else:
            # 直接返回空列表
            self.logger.info("无父块结果,返回空列表")
            return []
        # 返回CANDIDATE_M条父块结果
        return ranked_parent_docs[:config.CANDIDATE_M]

    def _doc_from_hit(self, hit):
        """
        把 Milvus 命中结果(dict)转换为 Document 对象。
        :param hit: Milvus 命中结果
        :return:
        """
        return Document(
            page_content=hit['text'],
            metadata={
                # 'id': hit['id'],
                'source': hit['source'],
                'timestamp': hit['timestamp'],
                'parent_id': hit['parent_id'],
                'parent_content': hit['parent_content']
            }
        )
    
    def _get_unique_parent_docs(self, child_chunks):
        """
        从子块列表中提取父块,并按父块内容去重。
        目的:避免同一父块因多个子块命中而重复返回。
        :param child_chunks: 子块列表
        :return: 去重后父块列表
        """
        parent_docs = set()
        # 返回值
        unique_parent_docs = []

        for chunk in child_chunks:
            # 优先取 parent_content,缺失时退化为子块文本
            parent_content = chunk.metadata.get('parent_content', chunk.page_content)
            # 如果父块内容非空,且不重复
            if parent_content and parent_content not in parent_docs:
                # 构建一个父块对象,放到返回值集合中
                unique_parent_docs.append(
                    Document(
                        # 返回的 page_content 存放父块文本
                        page_content=parent_content,
                        metadata=chunk.metadata
                    )
                )
                parent_docs.add(parent_content)

        return unique_parent_docs


# 主程序
if __name__ == "__main__":
    import document_processor
    documents = document_processor.process_documents(
        os.path.join(config.PROJECT_ROOT, "rag_qa/ai_data")
    )
    vector_store = VectorStore()
    vector_store.add_documents(documents=documents)
    result = vector_store.hybrid_search_with_rerank("大模型学什么")
    print(result)

3 本章小结 ¶

本章节全面讲解了 vector_store.py 模块的每个方法:

  • 1.初始化 Milvus

  • 2.创建或加载 database 和 collection

  • 3.文档向量化 并写入 Milvus

  • 4.根据用户问题 进行混合检索(稠密向量+稀疏向量) + 重排/精排

  • 5.根据检索到的子块匹配父块,从子块列表中提取父块,并按父块内容去重。

学习者通过本章节掌握了向量存储的完整流程,为RAG系统的检索功能奠定了基础。

5 查询分类模块

5.1 学习目标 ¶

  • 1.学会查询分类的基本原理,了解如何通过分类优化输入处理流程。

查询分类模块 query_classifier.py 负责 区分用户查询query为 通用知识 还是 专业咨询。通用知识 则直接使用模型回答,专业咨询 则走RAG流程。查询分类模块采用 BERT微调方案,基于预训练bert-base-chinese在 二分类数据集(model_generic.json)上进行继续训练, 实现二分类预测

1 功能概述 ¶

QueryClassifier 提供以下功能:

  • 1.数据预处理

    加载JSON数据

    将查询文本和预测标签转化为模型输入

  • 2.构建数据集

    自定义DataSet

  • 3.加载预训练BERT模型

    预训练 bert-base-chinese

  • 4.模型微调

    设置配置参数,在训练集上训练

  • 5.模型评估

    在测试集上评估,生成分类报告和混淆矩阵

  • 6.模型预测

    加载训练好的模型,进行二分类预测

2 代码实现 ¶

"""
查询分类模块:
    使用预训练 bert-base-chinese, 在二分类任务上微调,实现 通用知识 和 专业咨询 二分类
工作流:
    1.数据预处理
        加载JSON数据
        将查询文本和预测标签转化为模型输入
    2.构建数据集
        自定义DataSet
    3.加载预训练BERT模型
        预训练 bert-base-chinese
    4.模型微调
        设置配置参数,在训练集上训练
    5.模型评估
        在测试集上评估,生成分类报告和混淆矩阵
    6.模型预测
        加载训练好的模型,进行二分类预测

"""

# 导入标准库
import json
import os
import torch
# 导入日志
import sys
from base.logger import logger
from base.config import config
# 导入numpy
import numpy as np
# 导入 Transformers 库
from transformers import BertTokenizer, BertForSequenceClassification
# 模型训练和预测使用的工具
from transformers import Trainer, TrainingArguments

from sklearn.model_selection import train_test_split
from sklearn.metrics import classification_report, confusion_matrix

# 0.全局配置
# 设置设备
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "mps" if torch.mps.is_available() else "cpu" )
# 设置路径
CURRENT_DIR = os.path.dirname(os.path.abspath(__file__))
RAG_QA_PATH = os.path.abspath(os.path.dirname(os.path.abspath(CURRENT_DIR)))
PROJECT_ROOT = os.path.abspath(os.path.dirname(os.path.abspath(RAG_QA_PATH)))
# 超参数
BATCH_SIZE = 32
EPOCHS = 3
LR = 1e-5
WEIGHT_DECAY = 1e-3

class QueryClassifier(object):
    """
    0.初始化方法
    工作流:
        1.加载预训练BERT的tokenizer
        2.初始化微调后模型
        3.创建标签映射字典
        4.加载模型
    """
    def __init__(self, model_path='models/bert_query_classifier'):
        # 1.加载预训练BERT的tokenizer
        # 使用预训练模型,必须使用预训练的分词器,保证模型和分词器是匹配的
        self.pre_trained_model_path = f'{RAG_QA_PATH}/models/bert-base-chinese'
        self.tokenizer = BertTokenizer.from_pretrained(self.pre_trained_model_path)

        # 2.初始化模型和设备
        self.model_path = model_path
        # 模型对象
        self.model = None
        # 设备
        self.device = DEVICE
        logger.info(f"使用设备: {self.device}")
        # 3.创建标签映射字典
        self.label_map = {"通用知识": 0, "专业咨询": 1}
        # 4.加载模型
        self.load_model()

    # 加载微调好的模型 或 预训练BERT
    def load_model(self):
        # 1.优先加载训练好的模型
        if os.path.exists(self.model_path):
            self.model = BertForSequenceClassification.from_pretrained(self.model_path)
            logger.info(f"模型加载成功:{self.model_path}")
        # 2.如果不存在训练好的模型,则加载预训练BERT
        else:
            # num_labels=2:实现二分类任务,对应输出线性层的输出维度为2
            self.model = BertForSequenceClassification.from_pretrained(
                self.pre_trained_model_path, num_labels=2
            )
            logger.info("加载预训练BERT模型")
        # 迁移模型到设备
        self.model.to(self.device)
        # print(self.model)

    # 保存 模型 和 词表
    def save_model(self):
        # 1. 保存模型
        os.makedirs(self.model_path, exist_ok=True)
        self.model.save_pretrained(self.model_path, safe_serialization=False)
        # 2. 保存词表。 token->id映射关系
        self.tokenizer.save_pretrained(self.model_path)
        logger.info(f"保存模型成功:{self.model_path}")

    """
    1.数据预处理,将查询文本和标签转化为模型输入
    工作流:
        1.对输入文本进行分词器编码
        2.返回编码结果和标签列表
    """
    def preprocess_data(self, texts, labels):
        # texts: [问题1,问题2,...问题n]
        # labels:[标签1,标签2,...标签n]
        # pt: pytorch的缩写。 tf: tensorflow
        encodings = self.tokenizer(
            texts, # 输入的文本
            truncation=True,  # 是否截断文本到 max_length
            padding=True, # 是否填充文本,按照当前批次最大长度填充
            max_length=128, # 最大长度
            return_tensors="pt", # 返回张量类型, pt: pytorch
        )

        # encodings 字典 -> input_ids , attention_mask , token_type_ids
        # label_map 字典 {'通用知识': 0}

        return encodings, [self.label_map[label] for label in labels]

    """
    2.构建数据集,用于模型训练
    工作流:
        1. 自定义DataSet类
            1 定义初始化方法
            2 定义__getitem__,根据索引获取对应的数据 {input_ids,attention_mask,token_type_ids,labels}
            3 定义__len__,获取数据集长度
        2. 返回Dataset对象
    """
    def create_dataset(self, encodings, labels):
        class MyDataset(torch.utils.data.Dataset):
            def __init__(self, encodings, labels):
                self.encodings = encodings
                self.labels = labels

            # 根据索引,获取对应的值
            def __getitem__(self, idx):
                # encodings: { 'input_ids': input_ids, 'attention_mask':attention_mask, 'token_type_ids'  }
                # input_ids: (batch_size, seq_len)
                # val[idx]: 第idx条数据的编码以后的id
                item = {key: val[idx] for key, val in self.encodings.items()}
                # 标签 转为 张量
                item["labels"] = torch.tensor(self.labels[idx])
                # item {"labels":0或1, "attention_mask":tensor, "input_ids":tensor, "token_type_ids":tensor}
                return item

            def __len__(self):
                return len(self.labels)

        return MyDataset(encodings, labels)

    """
    3.模型微调,在训练集上训练模型
    工作流:
        1. 数据预处理
            1.1 加载数据集
            1.2 把数据集划分成8:2的训练集和验证集 
            1.3 把数据进行数值化
            1.4 构建Dataset
        2. 设置训练参数
        3. 初始化Trainer,传入参数、数据集、模型对象等
        4. 开始训练
        5. 保存模型
        6. 评估模型
    """
    def train_model(self, data_file='raining_dataset_hybrid_5000.json'):
        # 1.数据预处理
        # 保护性代码,确保训练的数据是存在的
        if not os.path.exists(data_file):
            logger.error(f"数据集文件 {data_file} 不存在")
            raise FileNotFoundError(f"数据集文件 {data_file} 不存在")

        # 打开文件作为f变量,最后退出的时候,自动调用close
        with open(data_file, "r", encoding="utf-8") as f:
            data = [json.loads(value) for value in f.readlines()]
        # 去重
        print("数据集长度:", len(data))
        # 按 (query, label) 去重,保留首次出现样本(顺序稳定)
        seen = set()
        dedup_data = []
        for item in data:
            key = (item.get("query", "").strip(), item.get("label", "").strip())
            if key in seen:
                continue
            seen.add(key)
            dedup_data.append(item)

        data = dedup_data
        print("去重后数据集长度:", len(data))

        texts = [item["query"] for item in data]
        labels = [item["label"] for item in data]

        train_texts, val_texts, train_labels, val_labels = train_test_split(
            texts, labels, test_size=0.2, random_state=42, stratify=labels)


        # preprocess_data :传入文本,返回张量
        train_encodings, train_labels = self.preprocess_data(train_texts, train_labels)
        val_encodings, val_labels = self.preprocess_data(val_texts, val_labels)

        train_dataset = self.create_dataset(train_encodings, train_labels)
        val_dataset = self.create_dataset(val_encodings, val_labels)

        # 2. 设置训练参数
        training_args = TrainingArguments(
            output_dir="bert_results", # 模型和检查点保存的目录路径
            save_total_limit=1, # 最多保存1个检查点文件,超出时自动删除旧的
            num_train_epochs=EPOCHS, # 训练轮数
            per_device_train_batch_size=BATCH_SIZE, # 批次大小
            per_device_eval_batch_size=BATCH_SIZE,
            warmup_steps=500, # 学习率预热步数为500步,训练初期学习率从0逐渐增加到设定值
            weight_decay=WEIGHT_DECAY, # 权重衰减系数
            logging_dir="./bert_logs", # 日志文件保存的目录路径
            logging_steps=20, # 多少个训练步骤记录一次日志
            evaluation_strategy="epoch", # 评估策略为每个epoch结束后进行评估
            save_strategy="epoch", # 模型保存策略:开发阶段为每个epoch结束后保存,生产阶段为不保存模型
            load_best_model_at_end=True, # 训练结束后加载最佳模型而非最后一个模型
            metric_for_best_model="eval_loss", # 判断最佳模型的指标为评估损失
            fp16=False, # 禁用FP16混合精度训练,使用FP32精度
        )

        # 3.初始化 Trainer
        trainer = Trainer(
            # 传入要训练的模型实例
            model=self.model,
            # 传入上面定义的训练参数配置
            args=training_args,
            # 传入训练数据集
            train_dataset=train_dataset,
            # 传入验证数据集,用于训练过程中评估模型性能
            eval_dataset=val_dataset,
            # 传入计算评估指标的函数,用于在验证集上计算准确率等指标
            compute_metrics=self.compute_metrics
        )

        # 训练模型
        logger.info("开始训练 BERT 模型...")
        trainer.train()
        self.save_model()

        # 评估模型
        self.evaluate_model(val_texts, val_labels)

    def compute_metrics(self, eval_pred):
        """计算评估指标acc"""
        # logits:预测权重值 (batch_size,num_classes): [[-1.5, 2.0]]
        # labels: 真实值 (batch_size,): [[0]]
        logits, labels = eval_pred
        # softmax不会影响数据前后的单调性(logits里面最大值,转成softmax归一化以后得结果,还是最大值)
        # argmax需要的是索引, argmax得到的结果就是:1
        # prediction = 1 , label = 0
        # predictions: [batch_size] -> labels:[batch_size]
        predictions = np.argmax(logits, axis=-1)
        accuracy = (predictions == labels).mean()
        return {"accuracy": accuracy}

    """
    4.模型评估,输出分类报告和混淆矩阵
    工作流:
       1. 数据预处理
           1.1 对输入文本进行分词编码(截断/填充至128长度)
           1.2 创建包含编码和标签的tensor数据集
       2. 初始化预测工具
           2.1 创建Trainer实例加载当前模型
       3. 执行预测
           3.1 使用predict方法获取原始预测结果
           3.2 通过argmax解析预测标签,得到概率最大的预测值的标签id(0 ~ 1)
       4. 生成评估报告
           4.1 输出分类报告(含精确率/召回率/F1值)
           4.2 输出混淆矩阵
    """

    def evaluate_model(self, texts, labels):
        """评估模型性能"""
        # 对 texts 进行分词器编码,获得 encodings: {input_ids, attention_mask, token_type_ids}
        encodings = self.tokenizer(
            texts,
            truncation=True,
            padding=True,
            max_length=128,
            return_tensors="pt"
        )
        dataset = self.create_dataset(encodings, labels)

        trainer = Trainer(model=self.model)
        predictions = trainer.predict(dataset)
        # predictions.predictions : (batch_size, 2)
        # np.softmax(predictions.predictions) ->  [-3.1, 2.7 ] /[ 0.3, 0.7 ]
        # argmax操作1维向量,所以我需要给它一个维度,它在哪个维度上去计算最大值 , array[-1]
        # predictions (batch,seq, num_classes=2)
        pred_labels = np.argmax(predictions.predictions, axis=-1)
        true_labels = labels  # 直接使用数字标签

        logger.info("分类报告:")
        logger.info(classification_report(
            true_labels,
            pred_labels,
            target_names=["通用知识", "专业咨询"]
        ))
        logger.info("混淆矩阵:")
        logger.info(confusion_matrix(true_labels, pred_labels))

    """
    5.模型预测
      比如: "Java学费一年多少钱"  -> "专业咨询"
    工作流:
        1. 加载模型, 并检查模型的状态
          1.1 验证模型是否已加载,未加载则记录错误并返回默认类别: 0->通用知识,让大模型处理query
        2. 输入数据处理
          2.1 对查询语句进行分词和编码(截断/填充至128长度)
          2.2 将编码数据移动到模型所在的设备
        3. 执行预测
          3.1 在无梯度模式下进行推理
          3.2 获取模型输出并解析预测结果(取logits最大值对应的类别)
        4. 结果映射
          4.1 将数字标签转换为对应的类别名称(0->通用知识,1->专业咨询)
    """
    def predict_category(self, query):
        # 检查模型是否加载
        if self.model is None:
            # 模型未加载,记录错误
            logger.error("模型未训练或加载")
            # 默认返回通用知识
            return "通用知识"
        # 对查询进行编码
        encoding = self.tokenizer(
            query,
            truncation=True,
            padding=True,
            max_length=128,
            return_tensors="pt"
        )
        # 将编码移到指定设备
        encoding = {k: v.to(self.device) for k, v in encoding.items()}
        # 不计算梯度,进行预测
        with torch.no_grad():
            # 获取模型输出
            # {"attention_mask":attention_mask, "input_ids":input_ids,"token_type_ids": token_type_ids}
            outputs = self.model(**encoding)
            # 获取预测结果
            prediction = torch.argmax(outputs.logits, dim=1).item()
        # 根据预测结果返回类别
        return "专业咨询" if prediction == 1 else "通用知识"


if __name__ == '__main__':

    model = QueryClassifier(model_path=os.path.join(config.MODELS_DIR,"bert_query_classifier"))
    path = os.path.join(config.PROJECT_ROOT, r'data/model_generic.json')

    train_model = model.train_model(path)

    test_queries = [
        "AI学科的课程大纲是什么",
        "JAVA课程费用多少?",
        "5*9等于多少?",
        "AI培训有哪些老师?",
        "你是人吗",
        "蒙特卡罗树怎么用在风险投资的",
        "transformers有哪些常用的API",
        "大模型学费多少?",
        "大模型学什么?",
        "python大模型学科和智能应用开发有什么区别?"
    ]
    for query in test_queries:
        category = model.predict_category(query)
        print(f"查询: {query} -> 分类: {category}")

3 实现细节 ¶

  • __init__ 方法

    • 作用:初始化分词器与分类模型(bert-base-chinese,二分类)。

    • 设备策略:优先 CUDA,不可用则使用 CPU(不启用 MPS,兼容性更稳)。

    • 标签映射label_map = {"通用知识": 0, "专业咨询": 1},统一训练标签格式。

  • preprocess_data 方法

    • 作用:将文本转为 BERT 输入(input_idsattention_mask),并将标签转为数字。

    • 关键参数max_length=128,兼顾性能与信息保留。

  • create_dataset 方法

    • 作用:构建 PyTorch Dataset,适配 Trainer 训练/评估流程。

    • 要点:确保 labels 为数值类型,避免训练报错。

  • train_model 方法

    • 作用:加载数据并微调模型(80% 训练 / 20% 验证)。

    • 核心参数

      • num_train_epochs=3

      • per_device_train_batch_size=8

      • fp16=False(提升 CPU/PyTorch 兼容性)

    • 流程

      1. 加载 training_dataset_hybrid_5000.json

      2. 预处理文本与标签

      3. 使用 Trainer 训练并保存最佳模型

  • evaluate_model 方法(重点优化)

    • 作用:在验证集输出分类报告与混淆矩阵。

    • 已修复问题:避免将数字标签(0/1)再次映射,修复 KeyError 风险。

    • 当前逻辑:仅对 texts 分词,true_labels = labels 直接参与评估。

  • predict_category 方法

    • 作用:对单条查询进行分类。

    • 输出:返回可读标签——“通用知识”“专业咨询”

5.2 章节小结 ¶

  • 1.查询分类模块 query_classifier.py 负责 区分用户查询query为 通用知识 还是 专业咨询。

  • 2.查询分类模块采用 BERT微调方案,基于预训练bert-base-chinese在 二分类数据集上进行继续训练, 实现二分类预测

6 prompts模块

6.1 学习目标 ¶

  • 1.掌握如何设计和使用Prompt模板来调用大语言模型优化query。

6.2 功能概述 ¶

prompts.py 定义了 RAGPrompts 类,管理系统中使用的所有Prompt模板。用于根据检索策略 调用大语言模型优化query。检索策略包括 直接检索、假设问题检索、子查询检索、回溯问题检索。

6.3 代码实现 ¶

# core/prompts.py
# 导入 PromptTemplate 类,用于创建 Prompt 模板
from langchain_core.prompts import PromptTemplate

# 定义 RAGPrompts 类,用于管理所有 Prompt 模板
class RAGPrompts:
    """
    RAGPrompts 类用于管理所有 Prompt 模板。
    包含两类 Prompt:
    1. system_prompt:LLM API 调用时的 system role 消息,定义 LLM 的角色和行为准则
    2. 用户消息模板(PromptTemplate):拼接 context/history/question 等变量后作为 user role 消息
    """

    # ======================== System Prompts ========================

    @staticmethod
    def rag_system_prompt():
        """RAG 检索模式的 system prompt:定义角色和回答风格,具体约束由 rag_prompt 控制"""
        return (
            "你是一名专业的IT教育领域智能助手,面向学生和教育从业者。"
            "回答应准确、清晰、简洁,易于理解。"
        )

    @staticmethod
    def general_system_prompt():
        """通用知识模式的 system prompt:允许使用自身知识直接回答"""
        return (
            "你是一名专业的IT教育领域智能助手。"
            "请根据你的知识准确、清晰、简洁地回答用户问题。"
            "如果你不确定答案,请如实说明。"
        )

    @staticmethod
    def strategy_system_prompt():
        """策略选择的 system prompt:严格执行 prompt 指令"""
        return "你是一个有用的助手,能够根据用户输入的Prompt严格执行并返回可靠的结果。"

    # ======================== User Message Templates ========================

    # 1. 定义 走RAG检索流程 的拼接Prompt的 Prompt 模板, 用于回答 专业咨询 的query
    @staticmethod
    def rag_prompt():
        return PromptTemplate(
            template="""
        你是一个智能助手,负责帮助用户回答问题。请按照以下步骤处理:

        1. **分析问题和上下文**:
           - 仅基于提供的上下文回答,禁止使用自身知识。如果上下文中没有答案,直接回复“信息不足,无法回答,请联系人工客服,电话:{phone}。”。
           - 如果答案来源于检索到的文档,请在回答中明确说明,例如:“根据提供的文档,……”。

        2. **评估对话历史**:
           - 检查对话历史是否与当前问题相关(例如,是否涉及相同的话题、实体或问题背景)。
           - 如果对话历史与问题相关,请结合历史信息生成更准确的回答。
           - 如果对话历史无关(例如,仅包含问候或不相关的内容),忽略历史,仅基于上下文和问题回答。

        3. **生成回答**:
           - 提供清晰、准确的回答,避免无关信息。
           - 如果上下文和历史消息均不足以回答问题,请回复:“信息不足,无法回答,请联系人工客服,电话:{phone}。”

        **上下文**: {context}
        **对话历史**:
        {history}
        **问题**: {question}

        **回答**:
        """,
            input_variables=["context", "history", "question", "phone"],
        )

    # 2. 定义直接调用LLM的 拼接Prompt的 Prompt 模板,用于回答 通用知识 的query
    @staticmethod
    def general_prompt():
        return PromptTemplate(
            template="""
        你是一个智能助手,负责帮助用户回答问题。

        请根据你的知识直接回答用户的问题,回答应准确、清晰、简洁。
        如果你不确定答案,请如实说明。

        **对话历史**:
        {history}
        **问题**: {question}

        **回答**:
        """,
            input_variables=["history", "question"],
        )

    # 3. 定义 假设问题检索策略(HyDE) 的 Prompt 模板
    @staticmethod
    def hyde_prompt():
        #   创建并返回 PromptTemplate 对象
        return PromptTemplate(
            template="""  
            用户想了解以下问题,请生成一个简短的假设答案:  
            问题: {query}  
            假设答案:  
            """,
            #   定义输入变量
            input_variables=["query"],
        )

    # 4. 定义 子查询检索策略 的 Prompt 模板
    @staticmethod
    def subquery_prompt():
        #   创建并返回 PromptTemplate 对象
        return PromptTemplate(
            template="""  
            将以下复杂查询分解为多个简单子查询,每行一个子查询:  
            查询: {query}  
            子查询:  
            """,
            #   定义输入变量
            input_variables=["query"],
        )

    # 5. 定义 回溯问题检索策略 的 Prompt 模板
    @staticmethod
    def backtracking_prompt():
        #   创建并返回 PromptTemplate 对象
        return PromptTemplate(
            template="""  
            将以下复杂查询简化为一个更简单的问题:  
            查询: {query}  
            简化问题:  
            """,
            #   定义输入变量
            input_variables=["query"],
        )

# 主程序
if __name__ == "__main__":
    rag_prompt = RAGPrompts.rag_prompt()
    result = rag_prompt.format(
        context="月入过万,就来黑马程序员",
        question="如何月入过万?",
        history="",
        phone="010-12345678",
    )
    print(result)

6.4 实现细节

  1. rag_prompt

    1. 作用:核心回答模板,结合检索到的上下文生成最终答案。

    2. 输入变量context(检索文档内容)、question(用户查询)、phone(客服电话)。

    3. 设计逻辑:支持“有上下文/无上下文”两种回答路径,并提供兜底回复,保障用户体验。

  2. hyde_prompt

    1. 作用:生成假设答案,用于 HyDE(Hypothetical Document Embeddings)策略,优化抽象查询的检索效果。

    2. 输入变量query(用户查询)。

    3. 设计逻辑:通过生成假设答案,增强查询与文档之间的语义匹配能力。

  3. subquery_prompt

    1. 作用:将复杂查询拆分为多个子查询,适用于多维度问题。

    2. 输入变量query(用户查询)。

    3. 设计逻辑:通过问题分解提高检索覆盖率与召回效果。

  4. backtracking_prompt

    1. 作用:将复杂查询简化为更基础的问题,便于检索。

    2. 输入变量query(用户查询)。

    3. 设计逻辑:通过降低查询复杂度,减少检索难度并提升命中率。

7 检索策略选择 ¶

7.1 学习目标 ¶

  • 掌握如何根据查询query 选择检索策略。

7.2 整体概述 ¶

strategy_selector.py 定义了 StrategySelector 类,根据用户问题query, 通过提示词工程,调用LLM 选择检索策略:直接检索、假设问题检索、子查询检索、回溯问题检索。然后prompts模块根据检索策略来优化query。

7.3 代码示例 ¶

"""
通过提示词工程,调用LLM实现检索策略分类:直接检索、假设问题检索、子查询检索、回溯问题检索。然后prompts模块根据检索策略来优化query

另一种低成本方案:微调本地BERT模型实现这个文本分类任务
"""

# core/strategy_selector.py 源码
# 导入 LangChain 提示模板
from langchain_core.prompts import PromptTemplate
# 导入日志和配置
from base.config import config
from base.logger import logger
# 导入 OpenAI
from openai import OpenAI

class StrategySelector:
    def __init__(self):
        # 初始化 OpenAI 客户端
        self.client = OpenAI(api_key=config.DASHSCOPE_API_KEY,
                             base_url=config.DASHSCOPE_BASE_URL)
        # 获取策略选择提示模板
        self.strategy_prompt_template = self._get_strategy_prompt()

    def _get_strategy_prompt(self):
        #   定义私有方法,获取策略选择 Prompt 模板
        return PromptTemplate(
            template="""
               你是一个智能助手,负责分析用户查询 {query},并从以下四种检索增强策略中选择一个最适合的策略,直接返回策略名称,不需要解释过程。

               以下是几种检索增强策略及其适用场景:

               1.  **直接检索**
                   * 描述:对用户查询直接进行检索,不进行任何增强处理。
                   * 适用场景:适用于查询意图明确,需要从知识库中检索**特定信息**的问题,例如:
                       * 示例:
                           * 查询:AI 学科学费是多少?
                           * 策略:直接检索
                       * 查询:JAVA的课程大纲是什么?
                           * 策略:直接检索
               2.  **假设问题检索**
                   * 描述:使用 LLM 生成一个假设的答案,然后基于假设答案进行检索。
                   * 适用场景:适用于查询较为抽象,直接检索效果不佳的问题,例如:
                       * 示例:
                           * 查询:人工智能在教育领域的应用有哪些?
                           * 策略:假设问题检索
               3.  **子查询检索**
                   * 描述:将复杂的用户查询拆分为多个简单的子查询,分别检索并合并结果。
                   * 适用场景:适用于查询涉及多个实体或方面,需要分别检索不同信息的问题,例如:
                       * 示例:
                           * 查询:比较 Milvus 和 Faiss 的优缺点。
                           * 策略:子查询检索
               4.  **回溯问题检索**
                   * 描述:将复杂的用户查询转化为更基础、更易于检索的问题,然后进行检索。
                   * 适用场景:适用于查询较为复杂,需要简化后才能有效检索的问题,例如:
                       * 示例:
                           * 查询:我有一个包含 100 亿条记录的数据集,想把它存储到 Milvus 中进行查询。可以吗?
                           * 策略:回溯问题检索

               根据用户查询 {query},直接返回最适合的策略名称,例如 "直接检索"。不要输出任何分析过程或其他内容。
               """
            ,
            input_variables=["query"],
        )

    def call_dashscope(self, prompt):
        # 调用 DashScope API
        try:
            # 创建聊天完成请求
            completion = self.client.chat.completions.create(
                model=config.LLM_MODEL,
                messages=[
                    {"role": "system",
                     "content": "你是一个有用的助手,能够根据用户输入的Prompt严格执行并返回可靠的结果"},
                    {"role": "user", "content": prompt},
                ],
                temperature=0.1
            )
            # 返回完成结果
            return completion.choices[0].message.content if completion.choices else "直接检索"
        except Exception as e:
            # 记录 API 调用失败
            logger.error(f"DashScope API 调用失败: {e}")
            # 默认返回直接检索
            return "直接检索"

    #   定义方法,选择检索策略
    def select_strategy(self, query):
        #   调用 LLM 获取检索策略
        prompt = self.strategy_prompt_template.format(query=query)
        print(f"选择检索策略的prompt: {prompt}")
        strategy = self.call_dashscope(prompt).strip()
        logger.info(f"为查询 '{query}' 选择的检索策略:{strategy}")
        return strategy

if __name__ == '__main__':
    selector = StrategySelector()
    strategy = selector.select_strategy(query="比较mysql和postgresql的优缺点")
    # print(strategy)

7.4 实现细节 ¶

__init__

  • 作用:初始化 DashScope 客户端与策略选择 Prompt 模板。

  • 核心逻辑:建立大模型 API 连接,并提前加载模板,减少后续调用开销。

call_dashscope

  • 作用:封装 DashScope API 调用,统一异常处理与结果返回。

  • 核心逻辑:发送请求获取策略结果;若调用失败则记录日志并返回默认策略,保证流程稳定。

_get_strategy_prompt

  • 作用:定义策略选择用的 Prompt 模板。

  • 核心逻辑:清晰描述四类检索策略及适用场景,约束模型仅返回策略名称,避免冗余输出。

select_strategy

  • 作用:根据用户查询选择并返回最合适的检索策略。

  • 核心逻辑:将查询填充到模板后调用模型,记录最终策略,便于调试与追踪。

1 调用流程

select_strategy(query)format promptcall_dashscope(prompt) → 返回策略名称(如“直接检索”)

8 RAG系统设计

8.1 学习目标 ¶

  • 理解RAG系统如何整合查询分类、检索和生成阶段。

8.2 功能概述 ¶

rag_system.py 定义了 RAGSystem 类,整合系统的各个模块,完成从用户问题query 到生成答案的完整流程。用户问题query -> 查询分类 -> 选择检索策略 -> 根据检索策略优化query -> 检索向量库得到检索上下文context -> 拼接提示词 -> LLM生成最终答案。

8.3 代码示例 ¶

"""
定义 RAGSystem 类,封装 RAG 系统的核心逻辑
1.初始化方法,设置 RAG 系统的基本参数
2.假设问题检索策略(HyDE)
    1. 获取 问题检索策略 对应提示词模板
    2. 调用大模型生成假设答案
    3. 基于假设答案查询父块作为上下文
3.子查询检索策略
    1.获取子查询检索策略的 Prompt 模板
    2.调用大模型生成子查询列表
    3.使用子查询进行混合检索
    4.对所有检索结果进行去重
4.回溯问题检索策略
    1.获取回溯问题检索策略的 Prompt 模板
    2.调用大模型生成回溯问题
    3.使用回溯问题进行检索
5.动态选择检索策略并整合结果
    1.未指定策略时通过策略选择器选择策略
    2.根据检索策略进行文档检索
    3.截取上下文文档数量(CANDIDATE_M)
6.端到端处理用户查询并生成答案
    1. 使用意图识别模型判断问题类型(通用/专业)
    2. 通用知识,直接调用 LLM 生成答案
    3. 专业咨询:
      3.1 选择最佳检索策略
      3.2 检索合并相关文档
      3.3 构建上下文
      3.4 组合提示模板调用 LLM

"""
import os.path

# core/rag_system.py 源码
# RAGPrompts包含: 1. augment提示词,用于结合query和上下文生成答案;2. 假设问题检索策略、子查询检索策略、回溯问题检索策略 对应的提示词模板
from rag_qa.core.prompts import RAGPrompts
# 导入 time 模块,用于计算时间
import time
from base.config import config
from base.logger import logger

# 区分 专业咨询 和 通用知识
from rag_qa.core.query_classifier import QueryClassifier  # 导入查询分类器
# 将专业咨询进一步分类,做策略选择
from rag_qa.core.strategy_selector import StrategySelector  # 导入策略选择器

# 定义RAGSystem类,实现RAG系统核心逻辑
class RAGSystem:
    # 1.初始化方法,设置 RAG 系统的基本参数
    def __init__(self, vector_store, llm):
        # 1.设置向量数据库对象
        self.vector_store = vector_store

        # 2.设置大模型调用函数
        self.llm = llm

        # 3.获取RAG提示词模板
        self.rag_prompt = RAGPrompts.rag_prompt()

        # 4.初始化查询分类器
        self.query_classifier = QueryClassifier(model_path=os.path.join(config.MODELS_DIR,"bert_query_classifier"))

        # 5.初始化策略选择器
        self.strategy_selector = StrategySelector()

    # 2.假设问题检索策略(HyDE)
    # 获取 假设问题检索策略 的文档检索得到的 上下文
    def _retrieve_with_hyde(self, query):
        logger.info(f"使用 假设问题检索策略, query: {query}")
        # 1. 获取 假设问题检索策略 对应提示词模板
        prompt_template = RAGPrompts.hyde_prompt()
        try:
            # 2. 调用大模型生成假设答案
            # 注意:self.llm 是生成器函数(yield),需要用 ''.join() 消耗生成器拿到完整字符串
            hypo_answer = ''.join(self.llm(prompt_template.format(query=query))).strip()
            # 3. 基于假设答案 查询父块作为上下文
            # 进行文档检索:混合检索 + 重排
            return self.vector_store.hybrid_search_with_rerank(
                query=hypo_answer, # 这里输入的是假设答案hypo_answer,而不是原始的query,因为假设问题检索策略基于假设答案进行检索
                top_k=config.RETRIEVAL_K,
            )
        except Exception as e:
            logger.error(f"假设问题检索策略 执行错误: {e}")
            return []

    # 3.子查询检索策略
    def _retrieve_with_subqueries(self, query):
        logger.info(f"使用 子查询检索策略, query: {query}")
        # 1.获取子查询检索策略的 Prompt 模板
        prompt_template = RAGPrompts.subquery_prompt()

        try:
            # 2.调用大模型生成子查询列表
            # 注意:self.llm 是生成器函数(yield),需要用 ''.join() 消耗生成器拿到完整字符串
            subqueries_text = ''.join(self.llm(prompt_template.format(query=query))).strip()
            subqueries = [q.strip() for q in subqueries_text.split("\n") if q.strip()]

            # 3.使用子查询进行混合检索
            # 初始化空列表,存储检索结果
            all_docs = []
            # 遍历每个子查询
            for subquery in subqueries:
                # 1.对每个子查询执行hybrid_search_with_rerank(混合检索 + 重排)
                docs = self.vector_store.hybrid_search_with_rerank(
                    query=subquery,
                    top_k=config.RETRIEVAL_K,
                )

                # 2.添加结果到列表中
                all_docs.extend(docs)
                logger.info(f"子查询: {subquery}, 检索到文档数量: {len(docs)}")

            # 4.对所有检索结果进行去重
            # 基于文档内容 或 ID 进行去重
            unique_docs_dict = {doc.page_content: doc for doc in all_docs}
            unique_docs = list(unique_docs_dict.values())

            logger.info(f"子查询检索策略, 检索到文档的去重后数量: {len(unique_docs)}")
            # 返回去重后的唯一文档
            return unique_docs
        except Exception as e:
            logger.error(f"子查询检索策略 执行错误: {e}")
            return []

    # 4.回溯问题检索策略
    # 返回 回溯问题检索策略 的文档检索得到的 上下文
    def _retrieve_with_backtracking(self, query):
        logger.info(f"使用 回溯问题检索策略, query: {query}")
        # 1.获取回溯问题检索策略的 Prompt 模板
        prompt_template = RAGPrompts.backtracking_prompt()
        try:
            # 2.调用大模型生成回溯问题
            # 注意:self.llm 是生成器函数(yield),需要用 ''.join() 消耗生成器拿到完整字符串
            backtracking_question = ''.join(self.llm(prompt_template.format(query=query))).strip()
            logger.info(f"生成的回溯问题: {backtracking_question}")

            # 3.使用回溯问题进行检索
            return self.vector_store.hybrid_search_with_rerank(
                query=backtracking_question, # 这里输入的是回溯问题backtracking_question,而不是原始的query,因为回溯问题检索策略基于回溯问题进行检索
                top_k=config.RETRIEVAL_K,
            )
        except Exception as e:
            logger.error(f"回溯问题检索策略 执行错误: {e}")
            return []

    # 5.动态选择检索策略并整合结果
    # 返回 整合后的文档检索结果
    def retrieve_and_merge(self, query, source_filter=None, strategy=None):
        """
        动态选择检索策略并整合结果:
        未指定strategy时,根据query选择检索策略,然后执行对应策略的文档检索,返回检索结果
        :param query: 查询
        :param source_filter: 学科过滤
        :param strategy: 检索策略
        :return: 文档检索结果
        """

        # 1.未指定策略时通过策略选择器选择策略
        if not strategy:
            strategy = self.strategy_selector.select_strategy(query)

        # 2.根据检索策略进行文档检索
        # 初始化检索到的文档列表
        ranked_chunks = []
        if strategy == "假设问题检索":
            ranked_chunks = self._retrieve_with_hyde(query)
        elif strategy == "子查询检索":
            ranked_chunks = self._retrieve_with_subqueries(query)
        elif strategy == "回溯问题检索":
            ranked_chunks = self._retrieve_with_backtracking(query)
        else: # 默认 直接检索
            logger.info(f"使用 直接检索策略, query: {query}")
            ranked_chunks = self.vector_store.hybrid_search_with_rerank(
                query=query,
                top_k=config.RETRIEVAL_K,
                source_filter=source_filter
            )
        logger.info(f"检索策略: {strategy}, 检索到文档数量: {len(ranked_chunks)}")

        # 3.截取上下文文档数量(CANDIDATE_M)
        # 子查询检索策略 的文档数量可以增大
        num_docs = config.CANDIDATE_M if strategy != "子查询检索" else config.CANDIDATE_M*2
        final_context_docs = ranked_chunks[:num_docs]
        logger.info(f"最终的文档数量: {len(final_context_docs)}")
        return final_context_docs

    # 6.端到端处理用户查询并生成答案
    def generate_answer(self, query, history=None, source_filter=None):
        """
        根据用于查询query,调用RAG系统,生成最终答案answer
        :param query: 用户查询
        :param history: 历史对话
        :param source_filter: 学科过滤
        :return: 最终答案
        """
        # 记录开始时间
        start_time = time.time()
        logger.info(f"用户查询: {query}, 学科过滤: {source_filter}")

        # 1. 使用意图识别模型判断问题类型(通用知识 / 专业咨询)
        query_category = self.query_classifier.predict_category(query)
        logger.info(f"问题类型: {query_category}")

        # 2. 通用知识,直接调用 LLM 生成答案
        if query_category == "通用知识":
            logger.info("query为通用知识,直接调用 LLM 生成答案")
            # 直接拼接提示词,不进行文档检索
            prompt_input = self.rag_prompt.format(
                context="", history="", question=query, phone=config.CUSTOMER_SERVICE_PHONE
            )
            try:
                # self.llm 是生成器函数,用 yield from 逐 token 转发给调用方
                yield from self.llm(prompt_input)
            except Exception as e:
                logger.error(f"直接调用 LLM 执行错误: {e}")
                yield f"抱歉,我无法回答您的问题。请联系人工客服: {config.CUSTOMER_SERVICE_PHONE}"
            process_time = time.time() - start_time
            logger.info(f"通用知识查询完成, 耗时: {process_time}s, 查询:{query}")
            return  # 生成器函数中,return 表示停止生成,不再继续执行后面的代码

        # 3. 专业咨询:
        logger.info("query为专业咨询,执行 RAG 流程")
        # 3.1 选择最佳检索策略
        strategy = self.strategy_selector.select_strategy(query)

        # 3.2 检索合并相关文档
        # list[Doc]
        context_docs = self.retrieve_and_merge(
            query, source_filter=source_filter, strategy=strategy)

        # 3.3 构建上下文
        if context_docs:
            # 兼容 Document / dict / str 三种结构,避免类型不一致导致崩溃
            context_parts = []
            for doc in context_docs:
                if hasattr(doc, "page_content"):
                    context_parts.append(doc.page_content)
                elif isinstance(doc, dict):
                    context_parts.append(doc.get("page_content") or doc.get("text") or doc.get("parent_content") or "")
                else:
                    context_parts.append(str(doc))
            context = "\n\n".join([p for p in context_parts if p])
            logger.info(f"构建上下文完成, 文档数量: {len(context_docs)}")
        else:
            context = ""
            logger.info("没有检索到相关文档, 上下文为空")

        # 3.4 组合提示模板调用 LLM
        # 准备历史对话
        # 验证历史格式:[{}]
        if history and not isinstance(history, list):
            logger.warning(f"无效的历史格式: {type(history)},忽略历史")
            history = []
            history_str = ""
        elif history:
            history_str = "\n\n".join(f"human:{row['question']}; ai:{row['answer']}" for row in history)
        else:
            history_str = ""
        # 构造 prompt
        prompt_input = self.rag_prompt.format(
            context=context,
            history=history_str,
            question=query,
            phone=config.CUSTOMER_SERVICE_PHONE
        )
        # logger.info(f"最终组合的提示词: {prompt_input}")
        # 调用 LLM
        try:
            # self.llm 是生成器函数,用 yield from 逐 token 转发给调用方
            yield from self.llm(prompt_input)
        except Exception as e:
            logger.error(f"RAG流程调用 LLM 执行错误: {e}")
            yield f"抱歉,我无法回答您的问题。请联系人工客服: {config.CUSTOMER_SERVICE_PHONE}"
        # 记录查询日志
        process_time = time.time() - start_time
        logger.info(f"专业咨询查询完成, 耗时: {process_time}s, 查询:{query}")

1 实现细节 ¶

  • __init__

    • 作用:初始化 RAG 系统,整合向量存储、大语言模型和其他核心组件。

    • 依赖VectorStoreRAGPromptsQueryClassifierStrategySelector

  • _retrieve_with_hyde

    • 作用:生成假设答案后再执行混合检索,适用于语义较抽象的查询。

    • 逻辑

      1. 使用 hyde_prompt 生成假设答案。

      2. 将假设答案传递给 hybrid_search_with_rerank

      3. 返回检索结果。

  • _retrieve_with_subqueries

    • 作用:将复杂查询拆分为多个子查询,分别检索后再去重。

    • 逻辑

      1. 使用 subquery_prompt 生成子查询列表。

      2. 对每个子查询分别执行检索。

      3. 合并结果并按内容去重。

      4. 返回去重后的文档列表。

  • _retrieve_with_backtracking

    • 作用:先简化复杂问题,再执行检索,降低检索难度。

    • 逻辑

      1. 使用 backtracking_prompt 生成更基础的问题。

      2. 调用 hybrid_search_with_rerank 执行检索。

      3. 返回检索结果。

  • retrieve_and_merge

    • 作用:根据策略选择不同检索方式,并返回最终候选文档。

    • 优化点

      • 移除冗余合并逻辑。

      • 直接使用 hybrid_search_with_rerank 的结果,即已经去重后的父文档。

    • 支持策略

      • 直接检索

      • 假设问题检索

      • 子查询检索

      • 回溯问题检索

  • generate_answer

    • 作用:整合查询分类、检索向量库与答案生成,输出最终结果。

    • 流程

      1. 使用 QueryClassifier 判断查询类型。

      2. 若为 “通用知识”,直接调用大语言模型生成答案。

      3. 若为 “专业咨询”,先选择检索策略并召回相关文档。

      4. 将检索到的上下文填入 rag_prompt,生成最终回答。

8.4 从查询到回答的流程 ¶

  • 输入处理 :

  • QueryClassifier 分类查询,决定是否需要检索。

  • 策略选择 :

  • StrategySelector 根据查询选择最佳检索策略。

  • 检索向量库 :

  • 根据策略调用 VectorStore 的混合检索,获取相关文档。

  • 答案生成 :

  • 使用 RAGPrompts 的模板,结合上下文调用大语言模型生成答案。

  • 输出 :

  • 返回最终答案,并记录日志。

9 RAG系统运行

9.1 学习目标: ¶

  • 1.掌握如何通过命令行参数和增强的错误处理运行EduRAG系统。

  • 2.了解整合核心的RAG逻辑。

9.2 系统运行入口 ¶

main.py 是EduRAG系统的运行模块,涵盖错误处理、参数灵活性和日志记录等功能。 main.py 作为命令行入口,支持数据处理和交互式查询,适合开发和调试;该模块整合了前几章的核心逻辑,为用户提供了健壮的交互方式。

1 功能概述 ¶

main.py是EduRAG系统的运行入口,提供两种运行模式:

  • 数据处理模式 :加载并向量化文档,构建向量数据库,支持多学科目录处理。

  • 查询模式 :通过命令行交互式回答用户查询,支持学科过滤。

import os
import sys
from base.config import config
from base.logger import logger
from rag_qa.core.document_processor import process_documents  # 导入处理文档的函数
from rag_qa.core.vector_store import VectorStore
from rag_qa.core.rag_system import RAGSystem
from rag_qa.core.prompts import RAGPrompts
from openai import OpenAI  # 使用 OpenAI 接口


def main(query_mode=True, directory_path="./ai_data"):

    #   初始化 DashScope API 客户端 (通过 OpenAI 接口)
    #   确保环境变量 DASHSCOPE_API_KEY 和 DASHSCOPE_BASE_URL 已设置
    try:

        client = OpenAI(api_key=config.DASHSCOPE_API_KEY,
                        base_url=config.DASHSCOPE_BASE_URL)

    except Exception as e:
        logger.error(f"初始化 OpenAI 客户端失败 (请检查 API Key 和 Base URL): {e}")
        # 如果客户端初始化失败,可能无法继续,取决于模式
        if query_mode: # 查询模式下必须要有 LLM
             print("错误:无法初始化语言模型客户端,无法进入查询模式。")
             return
        # 数据处理模式可能不需要 LLM,可以继续,但最好记录错误
        client = None # 标记客户端不可用


    # 定义 LLM 调用函数 (仅在需要时定义和使用)
    def call_dashscope(prompt, system_prompt=None):
        if not client: # 检查客户端是否可用
            logger.error("LLM 客户端未初始化,无法调用 call_dashscope")
            return f"错误: LLM客户端不可用"
        try:
            if system_prompt is None:
                system_prompt = RAGPrompts.rag_system_prompt()
            completion = client.chat.completions.create(
                model=config.LLM_MODEL,
                messages=[
                    {"role": "system", "content": system_prompt},
                    {"role": "user", "content": prompt},
                ]
                # 可以添加 temperature 等参数
            )
            if completion.choices and completion.choices[0].message:
                 return completion.choices[0].message.content
            else:
                 logger.error("LLM API 调用返回无效响应或空消息")
                 return "错误: LLM返回无效响应"
        except Exception as e:
            logger.error(f"LLM API (call_dashscope) 调用失败: {e}")
            return f"错误: 调用LLM失败 - {e}"

    # 初始化 VectorStore
    try:
        vector_store = VectorStore(
            collection_name=config.MILVUS_COLLECTION_NAME,
            host=config.MILVUS_HOST,
            port=config.MILVUS_PORT,
            database=config.MILVUS_DATABASE_NAME,
        )
    except Exception as e:
        logger.error(f"初始化 VectorStore 失败 (请检查 Milvus 连接配置): {e}")
        print("错误:无法连接到向量数据库,程序无法继续。")
        return

    # 根据模式执行不同操作
    if not query_mode:
        # --- 数据处理模式 ---
        logger.info("进入数据处理模式...")
        total_chunks_added = 0
        for source_dir in config.VALID_SOURCES:
            dir_path = os.path.join(directory_path, f"{source_dir}_data")
            if os.path.exists(dir_path):
                logger.info(f"开始处理目录: {dir_path}")
                try:
                    chunks = process_documents(
                        dir_path,
                        config.PARENT_CHUNK_SIZE,
                        config.CHILD_CHUNK_SIZE,
                        config.CHUNK_OVERLAP,
                    )
                    if chunks:
                        vector_store.add_documents(chunks)
                        total_chunks_added += len(chunks)
                        logger.info(f"成功处理目录 {dir_path},添加了 {len(chunks)} 个文档块")
                    else:
                        logger.info(f"目录 {dir_path} 未发现有效文档或处理结果为空")
                except Exception as e:
                    logger.error(f"处理目录 {dir_path} 时出错: {e}")
            else:
                logger.warning(f"目录 {dir_path} 不存在,跳过处理")
        logger.info(f"数据处理完成,共添加了 {total_chunks_added} 个文档块到向量存储")
    else:
        # --- 交互式查询模式 ---
        if not client: # 再次检查 LLM 客户端是否必须且可用
            print("错误:查询模式需要语言模型客户端,但初始化失败。")
            return

        logger.info("进入交互式查询模式...")
        try:
            rag_system = RAGSystem(vector_store, call_dashscope)
        except Exception as e:
             logger.error(f"初始化 RAGSystem 失败: {e}")
             print("错误:无法初始化 RAG 系统,无法进入查询模式。")
             return

        valid_sources = config.VALID_SOURCES
        print("\n欢迎使用 EduRAG 交互式查询系统!")
        print(f"支持的学科类别:{valid_sources}")
        print("输入您的问题,或输入 'exit' 退出。")

        while True:
            query = input("\n请输入您的问题:")
            if query.lower() == "exit":
                logger.info("用户退出查询模式")
                print("再见!")
                break
            if not query.strip():
                print("用户问题为空,无法回答")
                continue

            source_filter_input = input(f"请输入学科类别 ({'/'.join(valid_sources)}) (直接回车默认不过滤):").strip()
            source_filter = None # 默认不过滤
            if source_filter_input:
                if source_filter_input in valid_sources:
                    source_filter = source_filter_input
                    logger.info(f"用户选择了学科过滤: {source_filter}")
                else:
                    logger.warning(
                        f"无效的学科类别 '{source_filter_input}',将不过滤"
                    )
                    print(f"提示:输入的学科 '{source_filter_input}' 无效,将不过滤。")


            try:
                print("正在生成答案,请稍候...")
                answer = "".join(rag_system.generate_answer(query, source_filter=source_filter))
                print("-" * 30)
                print(f"问题: {query}")
                print(f"回答: {answer}")
                print("-" * 30)
            except Exception as e:
                logger.error(f"处理查询 '{query}' 时失败: {str(e)}")
                print(f"抱歉,处理您的问题时遇到了错误,请稍后重试或联系管理员。\n")


if __name__ == "__main__":
    # 默认进入查询模式
    # 若要执行数据处理,可以修改调用方式,例如:
    # main(query_mode=False)
    # 使用argparse实现命令行参数控制
    # 命令示例:
    # 数据处理模式:python main.py --data-processing --data-dir ./ai_data
    # 查询模式:python main.py --data-dir ./ai_data
    import argparse
    parser = argparse.ArgumentParser(description="EduRAG System Main Entry Point")
    parser.add_argument('--data-processing', action='store_true', help='Run in data processing mode instead of query mode.')
    parser.add_argument('--data-dir', type=str, default='./ai_data', help='Path to the data directory.')
    args = parser.parse_args()
    main(query_mode=(not args.data_processing), directory_path=args.data_dir)

2 实现细节 ¶

  • 环境变量加载:通过 Config 文件统一读取 API 密钥等配置。

  • LLM 客户端初始化:增强错误处理——查询模式下初始化失败则退出,数据处理模式可继续;call_dashscope 函数封装 API 调用并含异常处理。

  • VectorStore 初始化:显式传入配置参数,失败时终止程序。

  • 数据处理模式:遍历 VALID_SOURCES,按学科目录处理文档并记录分块数量,支持自定义分块参数。

  • 查询模式:展示支持的学科类别,校验 source_filter 有效性,格式化输出答案。

  • 命令行参数:通过 argparse 支持 --data-processing--data-dir,提升运行灵活性。

9.3 RAG系统运行 ¶

1 数据处理 ¶

python main.py --data-processing --data-dir ./data

2 查询模式 ¶

python main.py
欢迎使用 EduRAG 交互式查询系统!
支持的学科类别:['ai', 'java', 'test', 'ops', 'bigdata']
输入您的问题,或输入 'exit' 退出。
请输入您的问题:AI学科学费是多少?
请输入学科类别 (ai/java/test/ops/bigdata) (直接回车默认不过滤):ai
正在生成答案,请稍候...
------------------------------
问题: AI学科学费是多少?
回答: ...
------------------------------

9.4 章节总结 ¶

  • main.py :命令行入口,支持数据处理和交互式查询。

10 文档解析工具(扩展资料)

10.1 学习目标 ¶

  • 掌握系统中用于解析不同文档格式(PDF, DOCX, PPTX, Images)的核心工具。

  • 理解光学字符识别(OCR)工具 RapidOCR 如何集成并应用于文档解析流程。

  • 了解系统中提供的两种文本切分工具:基于规则的递归切分器和基于模型的语义切分器。

  • 熟悉 edu_document_loaders 和 edu_text_spliter 目录下各脚本的具体实现和功能。

10.2 文档解析工具 ( edu_document_loaders/ ) ¶

为了从各种常见的 IT 教育文档格式中提取信息,系统实现了一系列专门的加载器(Loaders)。这些加载器不仅能提取文档中的原生文本,还能利用 OCR 技术识别并提取图片中嵌入的文字。它们都继承自 Langchain 的 BaseLoader ,并实现了 lazy_load 方法来按需生成 Document 对象。

1 OCR 引擎核心 (edu_ocr.py) ¶

该脚本提供了一个标准化的函数 get_ocr() 来初始化和获取 OCR 识别引擎实例。这是所有需要图片文字识别功能的加载器的基础。

  • 功能 : 初始化 RapidOCR 实例。

  • 特点 :

  • 引擎选择 : 优先尝试 rapidocr_paddle (利用 PaddlePaddle 推理,推荐 GPU 环境),若失败则回退到 rapidocr_onnxruntime (利用 ONNX Runtime 推理,适合 CPU 环境或需要跨平台部署的场景)。

  • 参数控制 : 允许通过 use_cuda 参数控制是否启用 GPU 加速(如果使用 PaddlePaddle 引擎)。

# edu_document_loaders/edu_ocr.py 源码
from typing import TYPE_CHECKING
'''
paddleocr:解析图片中的文字,也可以进行表格识别
rapidocr_paddle 和 rapidocr_onnxruntime 两种导入方式
主要区别在于它们所使用的推理引擎和硬件支持
选择哪种方式最合适取决于你的硬件环境和性能需求。
当你有 GPU 且追求速度时:使用 rapidocr_paddle。PaddlePaddle 原生支持在 GPU 上推理 PaddleOCR 模型,速度更快。
当只有 CPU 且需要高效推理时:使用 rapidocr_onnxruntime。它在 CPU 上进行了优化,资源占用较低.
'''

def get_ocr(use_cuda: bool = True) -> "RapidOCR":
    try:
        from rapidocr_paddle import RapidOCR
        '''
        det_use_cuda=True:启用检测模型的GPU加速。cls_use_cuda=True:启用分类模型的GPU加速。rec_use_cuda=True:启用识别模型的GPU加速。
        '''
        ocr = RapidOCR(det_use_cuda=use_cuda, cls_use_cuda=use_cuda, rec_use_cuda=use_cuda)
    except ImportError:
        #
        from rapidocr_onnxruntime import RapidOCR
        ocr = RapidOCR()
    return ocr

2 PDF 文档加载器 (edu_pdfloader.py) ¶

OCRPDFLoader 类专门用于处理 PDF 文件。

  • 功能 : 解析 PDF,提取文本和图片中的文字。

  • 依赖 : PyMuPDF (fitz), Pillow , numpy , opencv-python , tqdm 以及 edu_ocr.py 。

  • 核心逻辑 :

  • 使用 fitz.open() 打开 PDF。

  • 逐页 ( page ) 处理。

  • 使用 page.get_text() 提取原生文本。

  • 使用 page.get_image_info(xrefs=True) 获取页面上的图片信息。

  • OCR 应用 : 对获取到的图片,检查其尺寸是否超过预设阈值 PDF_OCR_THRESHOLD (默认为页面宽高的 60%)。仅对大于阈值的图片执行 OCR。

  • 处理页面旋转 ( page.rotation ),确保 OCR 时图像方向正确。

  • 调用 get_ocr() 获取的 OCR 实例识别图片文字。

  • 合并原生文本和 OCR 结果。

  • 使用 tqdm 显示处理进度。

# edu_document_loaders/edu_pdfloader.py 源码
import cv2
import fitz  # pyMuPDF里面的fitz包,不要与pip install fitz混淆
import numpy as np
from PIL import Image
from tqdm import tqdm
from typing import Iterator
from edu_ocr import get_ocr
from langchain_core.documents import Document
from langchain_core.document_loaders import BaseLoader
from langchain.text_splitter import CharacterTextSplitter
# PDF OCR 控制:只对宽高超过页面一定比例(图片宽/页面宽,图片高/页面高)的图片进行 OCR。
# 这样可以避免 PDF 中一些小图片的干扰,提高非扫描版 PDF 处理速度
PDF_OCR_THRESHOLD = (0.6, 0.6)


class OCRPDFLoader(BaseLoader):
    """An example document loader that reads a file line by line."""

    def __init__(self, file_path: str) -> None:
        """Initialize the loader with a file path.

        Args:
            file_path: The path to the file to load.
        """
        self.file_path = file_path

    def lazy_load(self) -> Iterator[Document]:
        # <-- Does not take any arguments
        """A lazy loader that reads a file line by line.

        When you're implementing lazy load methods, you should use a generator
        to yield documents one by one.
        """

        line = self.pdf2text()
        yield Document(page_content=line, metadata={"source": self.file_path})



    def pdf2text(self):
        ocr = get_ocr()
        # 打开pdf文件
        doc = fitz.open(self.file_path)
        ## 获取页数
        # print(f'len(doc)-->{len(doc)}')
        resp = ""
        b_unit = tqdm(total=doc.page_count, desc="OCRPDFLoader context page index: 0")
        for i, page in enumerate(doc):
            b_unit.set_description("OCRPDFLoader context page index: {}".format(i))
            b_unit.refresh()
            # 提取文本:默认使用 "text" 模式提取文本。
            text = page.get_text("")
            resp += text + "\n"
            # print(f'resp-->{resp}')
            # 获取图片:获得所有显示的图像的元信息列表。
            # 它适用于所有文档类型,不仅限于 PDF。
            img_list = page.get_image_info(xrefs=True)
            # print(f'img_list--》{img_list}')
            # print(f'img_list--》{len(img_list)}')
            for img in img_list:
                # xref一种编号,指向该图像对象在PDF文件中的位置,程序可以通过这个编号快速定位和提取图像数据。
                if xref := img.get("xref"):
                    # 图像在页面上的位置和尺寸。
                    bbox = img["bbox"]
                    # 检查图片尺寸是否超过设定的阈值
                    # if ((bbox[2] - bbox[0]) / (page.rect.width) < PDF_OCR_THRESHOLD[0]
                    #         or (bbox[3] - bbox[1]) / (page.rect.height) < PDF_OCR_THRESHOLD[1]):
                    #     continue
                    pix = fitz.Pixmap(doc, xref)
                    # print(f'page.rotation-->{page.rotation}')
                    if int(page.rotation) != 0:  # 如果Page有旋转角度,则旋转图片
                        img_array = np.frombuffer(pix.samples, dtype=np.uint8).reshape(pix.height, pix.width, -1)
                        tmp_img = Image.fromarray(img_array)
                        ori_img = cv2.cvtColor(np.array(tmp_img), cv2.COLOR_RGB2BGR)
                        rot_img = self.rotate_img(img=ori_img, angle=360 - page.rotation)
                        img_array = cv2.cvtColor(rot_img, cv2.COLOR_RGB2BGR)
                    else:
                        img_array = np.frombuffer(pix.samples, dtype=np.uint8).reshape(pix.height, pix.width, -1)

                    # result:包含了图像中检测到的所有文本框的位置、文本内容和置信度信息。
                    # _:它是一个包含了时间数据的列表,可以用于优化模型运行速度。
                    result, _ = ocr(img_array)
                    if result:
                        ocr_result = [line[1] for line in result]
                        resp += "\n".join(ocr_result)
            # 更新进度
            b_unit.update(1)
        return resp



if __name__ == '__main__':
    pdf_loader = OCRPDFLoader(file_path="./data/Python机器学习基础教程.pdf")
    doc = pdf_loader.load()

    print(type(doc))
    print(doc)
    # text_spliter = CharacterTextSplitter(chunk_size=300, chunk_overlap=20)
    # result = text_spliter.split_documents(doc)
    # print(len(result))
    # print(result[0])

(注意:上述代码中 self.rotate_img 方法未在提供的代码段中定义,实际使用时需要确保该方法存在或移除相关调用)

3 Word 文档加载器 (edu_docloader.py) ¶

OCRDOCLoader 类用于处理 .docx 文件。

  • 功能 : 解析 DOCX 文件,提取段落、表格文本,并对嵌入的图片进行 OCR。

功能 : 解析 DOCX 文件,提取段落、表格文本,并对嵌入的图片进行 OCR。

  • 依赖 : python-docx , Pillow , numpy , tqdm 以及 edu_ocr.py 。

依赖 : python-docx , Pillow , numpy , tqdm 以及 edu_ocr.py 。

  • 核心逻辑 :

核心逻辑 :

  • 使用 docx.Document() 打开 DOCX 文件。

  • 定义 iter_block_items 辅助函数,用于统一遍历文档中的段落 ( Paragraph ) 和表格 ( Table ) 块。

  • 遍历所有块: 如果是段落,提取 block.text 。同时,使用 XPath ( .//pic:pic , .//a:blip/@r:embed ) 查找并提取段落内嵌入的图片。对提取的图片执行 OCR。 如果是表格,遍历所有单元格 ( cell ),提取单元格内段落的文本。

  • 如果是段落,提取 block.text 。同时,使用 XPath ( .//pic:pic , .//a:blip/@r:embed ) 查找并提取段落内嵌入的图片。对提取的图片执行 OCR。

  • 如果是表格,遍历所有单元格 ( cell ),提取单元格内段落的文本。

  • 合并所有提取的文本和 OCR 结果。

  • 使用 tqdm 显示处理进度。

# edu_document_loaders/edu_docloader.py 源码
from typing import Iterator
from .edu_ocr import get_ocr
# 导入必要的模块
from tqdm import tqdm
from docx.table import _Cell, Table  # 用于处理表格
from docx.oxml.table import CT_Tbl  # 用于处理表格XML结构
from docx.oxml.text.paragraph import CT_P  # 用于处理段落XML结构
from docx.text.paragraph import Paragraph  # 用于处理段落内容
from docx import Document as Docu1
from docx.document import Document as Docu2
from docx import ImagePart  # 用于处理Word文档和图片
from PIL import Image  # 用于处理图片
from io import BytesIO  # 用于将字节流转换为图片
import numpy as np  # 用于处理数组
from langchain_core.documents import Document
from langchain_core.document_loaders import BaseLoader
class OCRDOCLoader(BaseLoader):
    """An example document loader that reads a file line by line."""
    def __init__(self, filepath: str) -> None:
        """Initialize the loader with a file path.
        Args:
            filepath_path: The path to the filepath to load.
        """
        self.filepath = filepath
    def lazy_load(self) -> Iterator[Document]:
        # <-- Does not take any arguments
        """A lazy loader that reads a file line by line.
        When you're implementing lazy load methods, you should use a generator
        to yield documents one by one.
        """
        line = self.doc2text(self.filepath)
        yield Document(page_content=line, metadata={"source": self.filepath})
    def doc2text(self, filepath):
        # 创建OCR识别对象
        ocr = get_ocr()
        # print(f'ocr--》{ocr}')  # 输出OCR对象信息
        # 读取Word文档
        doc = Docu1(filepath)
        # print(f'doc-->{doc}')  # 输出读取到的文档信息
        # 定义一个空字符串用于存储最终的文本内容
        resp = ""
        # 定义一个迭代器,用于遍历文档中的块(段落、表格等)
        def iter_block_items(parent):
            # 判断parent对象类型,如果是Document类型,则获取其元素
            if isinstance(parent, Docu2):
                parent_elm = parent.element.body
            # 如果是表格单元格类型,获取单元格的XML元素
            elif isinstance(parent, _Cell):
                parent_elm = parent._tc
            else:
                raise ValueError("OCRDOCLoader parse fail")  # 如果都不是,则抛出错误
            # print(f'parent_elm--》{parent_elm}')
            # print('*'*80)
            # 遍历parent_elm中的所有子元素
            for child in parent_elm.iterchildren():
                # print(f'child--》{child}')
                if isinstance(child, CT_P):  # 如果是段落类型
                    yield Paragraph(child, parent)  # 返回段落
                elif isinstance(child, CT_Tbl):  # 如果是表格类型
                    yield Table(child, parent)  # 返回表格
        # print(f'doc.paragraphs-->{doc.paragraphs}')
        # print(f'doc.tables-->{doc.tables}')
        # 创建进度条,表示文档处理的进度
        b_unit = tqdm(total=len(doc.paragraphs) + len(doc.tables),
                      desc="OCRDOCLoader block index: 0")
        # 遍历文档中的所有块(段落和表格)
        for i, block in enumerate(iter_block_items(doc)):
            # 更新进度条描述
            b_unit.set_description("OCRDOCLoader  block index: {}".format(i))
            b_unit.refresh()  # 刷新进度条
            # 如果块是段落类型
            if isinstance(block, Paragraph):
                resp += block.text.strip() + "\n"  # 将段落文本加入到返回字符串中
                # 获取段落中的所有图片
                images = block._element.xpath('.//pic:pic')
                for image in images:
                    # 遍历图片,获取图片ID
                    for img_id in image.xpath('.//a:blip/@r:embed'):
                        part = doc.part.related_parts[img_id]  # 根据图片ID获取图片对象
                        if isinstance(part, ImagePart):  # 如果该部分是图片
                            # BytesIO 是 Python 内置的 io 模块中的一个类,用于在内存中读写二进制数据
                            # part._blob 通常表示从某个文档(如 DOCX 文件)中提取的二进制内容。
                            image = Image.open(BytesIO(part._blob))  # 打开图片
                            result, _ = ocr(np.array(image))  # 使用OCR识别图片中的文字
                            if result:  # 如果识别结果不为空
                                ocr_result = [line[1] for line in result]  # 提取识别出的文字
                                resp += "\n".join(ocr_result)  # 将识别结果加入返回文本中
            # 如果块是表格类型
            elif isinstance(block, Table):
                # 遍历表格中的所有行和单元格
                for row in block.rows:
                    for cell in row.cells:
                        for paragraph in cell.paragraphs:
                            resp += paragraph.text.strip() + "\n"  # 将单元格内的段落文本加入返回文本中
            # 更新进度条
            b_unit.update(1)
        # 返回提取的文本内容
        return resp
if __name__ == '__main__':
    docx_loader = OCRDOCLoader(filepath='./data/b.docx')
    doc = docx_loader.load()
    print(doc)

4 PowerPoint 文档加载器 (edu_pptloader.py) ¶

OCRPPTLoader 类用于处理 .ppt 和 .pptx 文件。

  • 功能 : 解析 PPT/PPTX 文件,提取形状(文本框、表格)、图片中的文本。

功能 : 解析 PPT/PPTX 文件,提取形状(文本框、表格)、图片中的文本。

  • 依赖 : python-pptx , Pillow , numpy , tqdm 以及 edu_ocr.py 。

依赖 : python-pptx , Pillow , numpy , tqdm 以及 edu_ocr.py 。

  • 核心逻辑 :

核心逻辑 :

  • 使用 pptx.Presentation() 打开演示文稿。

使用 pptx.Presentation() 打开演示文稿。

  • 逐张幻灯片 ( slide ) 处理。

逐张幻灯片 ( slide ) 处理。

  • 顺序处理 : 将幻灯片上的形状 ( shape ) 按视觉顺序( top , left 坐标)排序。

顺序处理 : 将幻灯片上的形状 ( shape ) 按视觉顺序( top , left 坐标)排序。

  • 定义 extract_text 递归函数处理单个形状: 提取文本框 ( shape.has_text_frame ) 的文本。 提取表格 ( shape.has_table ) 内所有单元格的文本。 如果形状是图片 ( shape.shape_type 13 ),提取图片数据 ( shape.image.blob ),执行 OCR。 如果形状是组合 ( shape.shape_type 6 ),递归调用 extract_text 处理其包含的子形状。

定义 extract_text 递归函数处理单个形状:

  • 提取文本框 ( shape.has_text_frame ) 的文本。

  • 提取表格 ( shape.has_table ) 内所有单元格的文本。

  • 如果形状是图片 ( shape.shape_type == 13 ),提取图片数据 ( shape.image.blob ),执行 OCR。

  • 如果形状是组合 ( shape.shape_type == 6 ),递归调用 extract_text 处理其包含的子形状。

  • 遍历排序后的形状,调用 extract_text 。

遍历排序后的形状,调用 extract_text 。

  • 合并所有提取的文本和 OCR 结果。

合并所有提取的文本和 OCR 结果。

  • 使用 tqdm 显示处理进度。

使用 tqdm 显示处理进度。

# edu_document_loaders/edu_pptloader.py 源码
from typing import Iterator
from edu_ocr import get_ocr
from langchain_core.documents import Document
from langchain_core.document_loaders import BaseLoader
from pptx import Presentation
from PIL import Image
import numpy as np
from io import BytesIO
from tqdm import tqdm
class OCRPPTLoader(BaseLoader):
    """An example document loader that reads a file line by line."""
    def __init__(self, filepath: str) -> None:
        """Initialize the loader with a file path.
        Args:
            filepath: The path to the ppt to load.
        """
        self.filepath = filepath
    def lazy_load(self) -> Iterator[Document]:
        # <-- Does not take any arguments
        """A lazy loader that reads a file line by line.
        When you're implementing lazy load methods, you should use a generator
        to yield documents one by one.
        """
        line = self.ppt2text(self.filepath)
        yield Document(page_content=line, metadata={"source": self.filepath})
    def ppt2text(self, filepath):
        # 打开指定路径的 PowerPoint 文件
        prs = Presentation(filepath)
        print(f'prs-->{prs}')
        # 获取 OCR 功能的实例
        ocr = get_ocr()
        # 初始化一个空字符串,用于存储提取的文本内容
        resp = ""
        def extract_text(shape):
            # nonlocal指明resp非全局非局部,而是外部嵌套函数中的变量,
            # 允许内部函数访问和修改外部函数中定义的变量resp
            nonlocal resp
            # 检查形状是否有文本框
            if shape.has_text_frame:
                # 将文本框中的文本添加到resp中,并去掉前后空格
                resp += shape.text.strip() + "\n"
            # 检查形状是否为表格
            if shape.has_table:
                # 遍历表格的每一行
                for row in shape.table.rows:
                    # 遍历每一行中的每个单元格
                    for cell in row.cells:
                        # 遍历单元格中的每个段落
                        for paragraph in cell.text_frame.paragraphs:
                            # 将单元格中的文本添加到resp中,并去掉前后空格
                            resp += paragraph.text.strip() + "\n"
            # 检查形状是否为图片(shape_type == 13)
            if shape.shape_type == 13:  # 13 表示图片
                # 使用 BytesIO 打开图片数据并转换为图像对象
                image = Image.open(BytesIO(shape.image.blob))
                # 使用 OCR 处理图像并获取结果
                result, _ = ocr(np.array(image))
                if result:  # 如果 OCR 有结果
                    # 提取 OCR 结果中的文本行
                    ocr_result = [line[1] for line in result]
                    # 将 OCR 提取的文本添加到resp中,以换行分隔
                    resp += "\n".join(ocr_result)
            # 检查形状是否为组合形状(shape_type == 6)
            elif shape.shape_type == 6:  # 6 表示组合
                # 遍历组合形状中的每个子形状,递归调用extract_text函数
                for child_shape in shape.shapes:
                    extract_text(child_shape)
        # 创建一个进度条,用于显示幻灯片处理进度,初始总数为幻灯片数量
        b_unit = tqdm(total=len(prs.slides), desc="OCRPPTLoader slide index: 1")
        # 遍历所有幻灯片
        for slide_number, slide in enumerate(prs.slides, start=1):
            # 更新进度条描述,显示当前处理的幻灯片索引
            b_unit.set_description("OCRPPTLoader slide index: {}".format(slide_number))
            b_unit.refresh()  # 刷新进度条显示
            # 按照从上到下、从左到右的顺序对形状进行排序遍历
            sorted_shapes = sorted(slide.shapes, key=lambda x: (x.top, x.left))
            for shape in sorted_shapes:
                extract_text(shape)  # 调用extract_text函数提取当前形状的文本内容
            b_unit.update(1)  # 更新进度条,表示处理了一张幻灯片
        return resp  # 返回提取到的所有文本内容
if __name__ == '__main__':
    img_loader = OCRPPTLoader(filepath='./data/01.pptx')
    doc = img_loader.load()
    print(doc)

5 图像文件加载器 (edu_imgloader.py) ¶

OCRIMGLoader 类用于直接处理图像文件(如 .png , .jpg )。

  • 功能 : 对单个图像文件执行 OCR。

功能 : 对单个图像文件执行 OCR。

  • 依赖 : Pillow , numpy 以及 edu_ocr.py 。

依赖 : Pillow , numpy 以及 edu_ocr.py 。

  • 核心逻辑 :

核心逻辑 :

  • 接收图像文件路径 img_path 。

  • 调用 get_ocr() 获取 OCR 实例。

  • 直接对图像文件执行 OCR。

  • 将 OCR 结果(所有识别出的文本行)合并成一个字符串。

# edu_document_loaders/edu_imgloader.py 源码
from typing import Iterator
from edu_ocr import get_ocr
from langchain_core.documents import Document
from langchain_core.document_loaders import BaseLoader
class OCRIMGLoader(BaseLoader):
    """An example document loader that reads a file line by line."""
    def __init__(self, img_path: str) -> None:
        """Initialize the loader with a file path.
        Args:
            img_path: The path to the img to load.
        """
        self.img_path = img_path
    def lazy_load(self) -> Iterator[Document]:
        # <-- Does not take any arguments
        """A lazy loader that reads a file line by line.
        When you're implementing lazy load methods, you should use a generator
        to yield documents one by one.
        """
        line = self.img2text()
        yield Document(page_content=line, metadata={"source": self.img_path})
    def img2text(self):
        resp = ""
        ocr = get_ocr()
        result, _ = ocr(self.img_path)
        if result:
            ocr_result = [line[1] for line in result]
            resp += "\n".join(ocr_result)
        return resp
if __name__ == '__main__':
    img_loader = OCRIMGLoader(img_path='./data/test_img.png')
    doc = img_loader.load()
    print(doc)

10.3 文本切分工具 ( edu_text_spliter/ ) ¶

将解析得到的长文本切分成适合向量化和检索的小块是 RAG 流程中的关键一步。本系统提供了两种文本切分工具。

1 中文递归文本切分器 (edu_chinese_recursive_text_splitter.py) ¶

ChineseRecursiveTextSplitter 类是针对中文文本特点定制的切分器。

  • 功能 : 将长文本按照预设的中文分隔符递归地切分成指定大小的块。

功能 : 将长文本按照预设的中文分隔符递归地切分成指定大小的块。

  • 继承 : langchain.text_splitter.RecursiveCharacterTextSplitter 。

继承 : langchain.text_splitter.RecursiveCharacterTextSplitter 。

  • 核心定制 :

核心定制 :

  • _separators : 定义了用于切分的、按优先级排列的分隔符列表,包括常见的中文标点和换行符,如 [“\n\n”, “\n”, “。|!|?”, “.\s|!\s|?\s”, “;|;\s”, “,|,\s”] 。这有助于在切分时尽量保持句子的完整性。

  • 支持通过正则表达式定义分隔符 ( is_separator_regex=True )。

  • 通过 chunk_size 和 chunk_overlap 控制切分块的大小和重叠。

# edu_text_spliter/edu_chinese_recursive_text_splitter.py 源码
import re
from typing import List, Optional, Any
from langchain.text_splitter import RecursiveCharacterTextSplitter
import logging
logger = logging.getLogger(__name__)
def _split_text_with_regex_from_end(
        text: str, separator: str, keep_separator: bool
) -> List[str]:
    # Now that we have the separator, split the text
    if separator:
        if keep_separator:
            # The parentheses in the pattern keep the delimiters in the result.
            _splits = re.split(f"({separator})", text)
            splits = ["".join(i) for i in zip(_splits[0::2], _splits[1::2])]
            if len(_splits) % 2 == 1:
                splits += _splits[-1:]
            # splits = [_splits[0]] + splits
        else:
            splits = re.split(separator, text)
    else:
        splits = list(text)
    return [s for s in splits if s != ""]
class ChineseRecursiveTextSplitter(RecursiveCharacterTextSplitter):
    def __init__(
            self,
            separators: Optional[List[str]] = None,
            keep_separator: bool = True,
            is_separator_regex: bool = True,
            **kwargs: Any,
    ) -> None:
        """Create a new TextSplitter."""
        super().__init__(keep_separator=keep_separator, **kwargs)
        self._separators = separators or [
            "\n\n",
            "\n",
            "。|!|?",
            "\.\s|\!\s|\?\s",
            ";|;\s",
            ",|,\s"
        ]
        self._is_separator_regex = is_separator_regex
    def _split_text(self, text: str, separators: List[str]) -> List[str]:
        """Split incoming text and return chunks."""
        final_chunks = []
        # Get appropriate separator to use
        separator = separators[-1]
        new_separators = []
        for i, _s in enumerate(separators):
            _separator = _s if self._is_separator_regex else re.escape(_s)
            if _s == "":
                separator = _s
                break
            if re.search(_separator, text):
                separator = _s
                new_separators = separators[i + 1:]
                break
        _separator = separator if self._is_separator_regex else re.escape(separator)
        splits = _split_text_with_regex_from_end(text, _separator, self._keep_separator)
        # Now go merging things, recursively splitting longer texts.
        _good_splits = []
        _separator = "" if self._keep_separator else separator
        for s in splits:
            if self._length_function(s) < self._chunk_size:
                _good_splits.append(s)
            else:
                if _good_splits:
                    merged_text = self._merge_splits(_good_splits, _separator)
                    final_chunks.extend(merged_text)
                    _good_splits = []
                if not new_separators:
                    final_chunks.append(s)
                else:
                    other_info = self._split_text(s, new_separators)
                    final_chunks.extend(other_info)
        if _good_splits:
            merged_text = self._merge_splits(_good_splits, _separator)
            final_chunks.extend(merged_text)
        return [re.sub(r"\n{2,}", "\n", chunk.strip()) for chunk in final_chunks if chunk.strip()!=""]
if __name__ == "__main__":
    text_splitter = ChineseRecursiveTextSplitter(
        keep_separator=True,
        is_separator_regex=True,
        chunk_size=150,
        chunk_overlap=10
    )
    ls = [
        """中国对外贸易形势报告(75页)。前 10 个月,一般贸易进出口 19.5 万亿元,增长 25.1%, 比整体进出口增速高出 2.9 个百分点,占进出口总额的 61.7%,较去年同期提升 1.6 个百分点。其中,一般贸易出口 10.6 万亿元,增长 25.3%,占出口总额的 60.9%,提升 1.5 个百分点;进口8.9万亿元,增长24.9%,占进口总额的62.7%, 提升 1.8 个百分点。加工贸易进出口 6.8 万亿元,增长 11.8%, 占进出口总额的 21.5%,减少 2.0 个百分点。其中,出口增 长 10.4%,占出口总额的 24.3%,减少 2.6 个百分点;进口增 长 14.2%,占进口总额的 18.0%,减少 1.2 个百分点。此外, 以保税物流方式进出口 3.96 万亿元,增长 27.9%。其中,出 口 1.47 万亿元,增长 38.9%;进口 2.49 万亿元,增长 22.2%。前三季度,中国服务贸易继续保持快速增长态势。服务 进出口总额 37834.3 亿元,增长 11.6%;其中服务出口 17820.9 亿元,增长 27.3%;进口 20013.4 亿元,增长 0.5%,进口增 速实现了疫情以来的首次转正。服务出口增幅大于进口 26.8 个百分点,带动服务贸易逆差下降 62.9%至 2192.5 亿元。服 务贸易结构持续优化,知识密集型服务进出口 16917.7 亿元, 增长 13.3%,占服务进出口总额的比重达到 44.7%,提升 0.7 个百分点。 二、中国对外贸易发展环境分析和展望 全球疫情起伏反复,经济复苏分化加剧,大宗商品价格 上涨、能源紧缺、运力紧张及发达经济体政策调整外溢等风 险交织叠加。同时也要看到,我国经济长期向好的趋势没有 改变,外贸企业韧性和活力不断增强,新业态新模式加快发 展,创新转型步伐提速。产业链供应链面临挑战。美欧等加快出台制造业回迁计 划,加速产业链供应链本土布局,跨国公司调整产业链供应 链,全球双链面临新一轮重构,区域化、近岸化、本土化、 短链化趋势凸显。疫苗供应不足,制造业“缺芯”、物流受限、 运价高企,全球产业链供应链面临压力。 全球通胀持续高位运行。能源价格上涨加大主要经济体 的通胀压力,增加全球经济复苏的不确定性。世界银行今年 10 月发布《大宗商品市场展望》指出,能源价格在 2021 年 大涨逾 80%,并且仍将在 2022 年小幅上涨。IMF 指出,全 球通胀上行风险加剧,通胀前景存在巨大不确定性。""",
        ]
    # text = """"""
    for inum, text in enumerate(ls):
        print(inum)
        chunks = text_splitter.split_text(text)
        for chunk in chunks:
            print(chunk)

2 基于模型的语义切分器 (edu_model_text_spliter.py) ¶

AliTextSplitter 类提供了另一种基于 AI 模型的文本切分方法。

  • 功能 : 利用预训练的文档语义分割模型对文本进行切分。

  • 继承 : langchain.text_splitter.CharacterTextSplitter 。

  • 核心逻辑 :

  • 初始化时指定是否处理 PDF 文本(包含特定的换行符和空格处理逻辑)。

  • 调用 modelscope.pipeline 加载指定的文档分割模型(代码中为 MODEL_PATH[‘segment_model’][‘ali_model’] ,需要配置 configs.py 或直接指定模型路径/名称,如 ‘damo/nlp_bert_document-segmentation_chinese-base’)。模型运行在 CPU 上。

  • 将输入文本传递给模型 pipeline 进行处理。

  • 模型返回按语义分割好的文本段落,脚本将其整理成列表返回。

  • 优势 : 理论上能更好地根据内容的语义关联性进行切分,而不是仅仅依赖标点符号。

  • 劣势 : 需要额外加载一个模型,增加了计算开销和依赖。

# edu_text_spliter/edu_model_text_spliter.py 源码
from langchain.text_splitter import CharacterTextSplitter
import re
from typing import List
from modelscope.pipelines import pipeline
# from configs import MODEL_PATH # Assume MODEL_PATH is defined elsewhere or replaced
# Placeholder for MODEL_PATH if configs.py is not available
MODEL_PATH = {
    'segment_model': {
        # Replace with the actual model name or path from ModelScope
        'ali_model': 'damo/nlp_bert_document-segmentation_chinese-base'
    }
}
class AliTextSplitter(CharacterTextSplitter):
    def __init__(self, pdf: bool = False, **kwargs):
        super().__init__(**kwargs)
        self.pdf = pdf
        # Initialize the pipeline here or ensure it's initialized before split_text is called
        # Consider adding error handling for model loading
        try:
            self.pipeline = pipeline(
                task="document-segmentation",
                model=MODEL_PATH['segment_model']['ali_model'],
                device="cpu" # Specify CPU device
            )
        except Exception as e:
            print(f"Error initializing ModelScope pipeline: {e}")
            self.pipeline = None
    def split_text(self, text: str) -> List[str]:
        if not self.pipeline:
            print("ModelScope pipeline not initialized. Returning empty list.")
            return []
        # Preprocessing specific to PDF text if needed
        if self.pdf:
            text = re.sub(r"\n{3,}", r"\n", text)
            # Replace multiple spaces with a single space
            text = re.sub('\s+', " ", text)
            # Consider removing single newlines carefully, might merge unrelated lines
            # text = text.replace("\n", " ") # This might be too aggressive
            text = re.sub("\n\n", "\n", text) # Keep paragraph breaks
        try:
            result = self.pipeline(documents=text)
            # The default output format might be a single string with "\n\t" separators
            sent_list = [segment.strip() for segment in result["text"].split("\n\t") if segment.strip()]
            return sent_list
        except Exception as e:
            print(f"Error during ModelScope document segmentation: {e}")
            # Fallback behavior: maybe split by paragraph or return the original text in a list
            return text.split('\n\n') # Simple fallback
if __name__ == '__main__':
    # Example usage requires modelscope and relevant model downloaded
    # pip install "modelscope[nlp]" tensorflow torch -f https://modelscope.oss-cn-beijing.aliyuncs.com/releases/repo.html
    sample_text = """移动端语音唤醒模型,检测关键词为“小云小云”。
模型主体为4层FSMN结构,使用CTC训练准则,参数量750K,适用于移动端设备运行。
模型输入为Fbank特征,输出为基于char建模的中文全集token预测,测试工具根据每一帧的预测数据进行后处理得到输入音频的实时检测结果。
模型训练采用“basetrain + finetune”的模式,basetrain过程使用大量内部移动端数据,在此基础上,使用1万条设备端录制安静场景“小云小云”数据进行微调,得到最终面向业务的模型。
后续用户可在basetrain模型基础上,使用其他关键词数据进行微调,得到新的语音唤醒模型,但暂时未开放模型finetune功能。"""
    # Assuming MODEL_PATH is correctly configured and modelscope is installed
    model_split = AliTextSplitter()
    result = model_split.split_text(text=sample_text)
    print(result)
    # Expected output (example, actual output depends on the model):
    # ['移动端语音唤醒模型,检测关键词为“小云小云”。', '模型主体为4层FSMN结构,使用CTC训练准则,参数量750K,适用于移动端设备运行。', '模型输入为Fbank特征,输出为基于char建模的中文全集token预测,测试工具根据每一帧的预测数据进行后处理得到输入音频的实时检测结果。', '模型训练采用“basetrain + finetune”的模式,basetrain过程使用大量内部移动端数据,在此基础上,使用1万条设备端录制安静场景“小云小云”数据进行微调,得到最终面向业务的模型。', '后续用户可在basetrain模型基础上,使用其他关键词数据进行微调,得到新的语音唤醒模型,但暂时未开放模型finetune功能。']

10.4 本章小结 ¶

本章我们详细介绍了 EduRAG 系统中用于处理原始文档和切分文本的核心工具。在文档解析方面,我们学习了 edu_document_loaders 目录下的各个加载器如何利用 PyMuPDF , python-docx , python-pptx 等库结合 RapidOCR (通过 edu_ocr.py 提供)来处理 PDF、DOCX、PPTX 及图像文件,有效提取文本和图片中的文字。在文本切分方面,我们探讨了 edu_text_spliter 目录提供的两种工具: ChineseRecursiveTextSplitter (针对中文优化的、基于规则的递归切分器)和 AliTextSplitter (利用 AI 模型进行语义切分)。理解这些底层工具的功能和实现是掌握 RAG 系统数据处理流程的关键。

11 RAG中的Query改写(扩展资料)

11.1 学习目标 ¶

  • 理解query改写的意义

  • 掌握qeury改写的实现方法

11.2 前言 ¶

在RAG(检索增强生成)流程中,第一步通常是对用户的提问(query)进行改写。这是因为用户提问的方式与他们期望的答案之间可能存在差距。由于每个用户的提问方式可能千差万别,因此对问题进行改写可以帮助系统更好地理解问题并返回更相关的答案,从而提升RAG系统的鲁棒性和扩展性。
用户提出的问题通常存在两类问题:

  • 信息不完整:用户的提问没有表达清楚所有的关键信息

  • 噪声问题:提问中可能包含了与答案无关的内容。

11.3 信息不完整 ¶

1 历史会话改写 ¶

在对话中,前后文是相互关联的。如果仅凭当前的query进行检索,可能会导致召回精度大幅下降,因为query中往往缺少重要的上下文信息。以下是一个具体的例子:

用户
:
华为meta70手机的性能怎么样?
系统
:
华为meta70手机搭载了强大的处理器和先进的摄像系统,性能表现非常优秀。
用户
:
与上一代相比,它有哪些改进?
系统
:
华为meta70相较于meta60在处理器性能、摄像头优化和电池续航方面都有显著提升。
用户
:
摄像头方面具体改进了什么?
--改写前
:
摄像头方面具体改进了什么?
--改写后
:
华为meta70手机的摄像头相比meta60有哪些具体改进?

2 关键词扩写 ¶

用户在搜索时常常输入的关键词较为简短,并且缺乏足够的上下文信息,这会影响语义检索(向量检索)的效果,导致召回的相关性较低。因此,需要对用户的原始关键词进行扩展和丰富。

用户输入
:
“机器学习 实践”
改写后的
Query: “机器学习在实际应用中的案例有哪些?哪些工具和方法适用于机器学习实践?”

3 伪答案改写 ¶

伪答案改写通过在原始查询中加入一种假设性答案,来增强查询的语义丰富性,从而提高检索或回应的精准度。伪答案并非真实的答案,而是一个设想的内容,用于帮助系统更好地理解并检索相关信息。
用户输入 : “如何提高企业的市场竞争力?” 改写后的 Query: “如何提高企业的市场竞争力?比如通过创新产品、优化营销策略或提升客户服务等手段。” 伪答案目的–> 通过提供假设性的提升方式,丰富查询的语义信息,从而增强系统在应对复杂问题时的检索能力 2.4 缩写词改写
用户在查询时常常使用缩写,而许多相关文档通常会使用完整的术语,因此需要对缩写进行扩展,以便更好地匹配相关内容。

用户输入
:
“如何提高企业的市场竞争力?”
改写后的
Query: “如何提高企业的市场竞争力?比如通过创新产品、优化营销策略或提升客户服务等手段。”
伪答案目的-->
通过提供假设性的提升方式,丰富查询的语义信息,从而增强系统在应对复杂问题时的检索能力
用户输入
:
“VR 技术在教育中的应用”
改写后的
Query: “虚拟现实(Virtual Reality)技术在教育中的应用有哪些?可以举一些实际的应用案例吗?”

11.4 噪声问题 ¶

1 般去噪改写 ¶

通过去除查询中的无关成分(如多余的修饰语、模糊表达或不相关的背景信息),简化并优化查询,使其更加精确和可操作。这种方法有助于提高检索的准确性和效率。

用户输入
:
“我最近在准备面试,但对于算法的理解还不太够,能推荐一些有效的学习资源吗?”
改写后的
Query: “有哪些有效的学习资源可以帮助提高算法理解?”
分析:去除与问题无关的背景信息
“我最近在准备面试,但对于算法的理解还不太够”。 直接提取核心意图 “帮助提高算法理解”。

2 关键词改写 ¶

这是一种专注于提取核心关键词并去除噪声的查询重写方法。通过识别查询中的关键内容,并排除冗余信息(如停用词、语气词和多余的描述),使查询更加简洁明了,从而提高检索效率和准确性。该方法特别适用于关键词检索召回,如 BM25 检索算法。

用户输入
:
“关于 Java 中的线程池,常见的实现方式有哪些?”
改写后
:
“Java 线程池 常见实现方式”

3 子查询改写 ¶

当查询涉及对比多个实体时,可能会产生相互干扰的情况。对比类查询通常包含多个元素,这些元素如果直接放在一个查询中,可能会导致信息重叠,影响检索的准确性。为了避免干扰,可以将对比类查询拆分成多个独立的查询,每个查询聚焦于其中一个实体,这样可以减少信息混淆,获得更准确的结果。

用户输入
:
“C++ 和 Go 哪个更适合做系统编程?”
拆分后的查询
:
“C++ 适合做系统编程的优点有哪些?”;“Go 适合做系统编程的优点有哪些?”
拆分原因:
• 直接对比 C++ 和 Go 的优劣,可能使得系统无法有效地提取每种语言的特点。拆分后,系统可以分别检索 C++ 和 Go 在系统编程中的优点,避免信息混乱

11.5 prompt示例 ¶

以下是一个开源rag系统,问题改写的示例,仅供参考:

您是查询扩展方面的专家,能够生成问题的释义。
我无法直接使用用户的问题从知识库中检索相关信息。
您需要通过多种方式扩展或释义用户的问题,例如使用同义词/短语、完整地写出缩写、添加一些额外的描述或解释、改变表达方式、将原始问题翻译成另一种语言(英语/中文)等。
并返回
5 个版本的问题,其中一个来自翻译。
只需列出问题。不需要其他单词。

11.6 本章小结 ¶

改写后的query质量在很大程度上依赖于提示(prompt)和大模型的能力。通常,查询改写可以提高召回的准确性,但有时仍可能发生改写失败的情况。因此,建议进行多次改写,或者将查询拆分为多个独立的查询,之后将召回结果输入到重排序(reranking)模型进行精确排序,再由大模型生成最终的答案。

12 数据集生成与优化(扩展资料)

12.1 学习目标 ¶

  • 1.掌握如何结合规则模板和 Qwen-Plus 模型生成高质量查询数据集。

  • 2.通过 tqdm 进度条监控生成进度,并实现分阶段数据保存。

generate_query_dataset_hybrid.py 是 EduRAG 系统中用于生成训练数据集的核心脚本,旨在为 QueryClassifier 提供 6000 条高质量数据(“通用知识”和“专业咨询”各 3000 条)。通过规则模板和 Qwen-Plus 模型的混合生成,结合进度条和分阶段保存,本脚本确保了生成效率和数据可靠性。本章节将详细讲解其功能、实现和应用。

12.2 数据生成与优化 ¶

1 功能概述 ¶

generate_query_dataset_hybrid.py 提供以下功能:

  • 规则生成 :基于模板和同义词替换生成 3000 条数据(“通用知识”和“专业咨询”各 1500 条)。

  • 大模型生成 :利用 Qwen-Plus 生成 3000 条自然语言查询(各 1500 条)。

  • 进度监控 :通过 tqdm 进度条可视化每个生成阶段。

  • 分阶段保存 :每生成 1500 条数据保存一次,最终合并保存完整数据集。

  • 数据集输出 :生成均衡的 6000 条数据,保存为 JSON 文件。

2 完整代码 ¶

# generate_query_dataset_hybrid.py
import json
import random
import re
import os
from openai import OpenAI
from dotenv import load_dotenv
from tqdm import tqdm
import time
# 加载环境变量
load_dotenv()
# 初始化Qwen-Plus客户端
client = OpenAI(
    api_key=os.getenv("DASHSCOPE_API_KEY"),
    base_url=os.getenv("DASHSCOPE_BASE_URL", "https://dashscope.aliyuncs.com/compatible-mode/v1"),
)
# 同义词词典
synonym_dict = {
    "什么": ["什么", "啥是", "如何解释", "具体是什么", "请解释"],
    "课程": ["课程", "培训课程", "课程安排", "教学内容", "课"],
    "学费": ["学费", "费用", "报名费", "学习费用", "价格", "收费"],
    "大纲": ["大纲", "课程内容", "教学计划", "讲义"],
    "师资": ["师资", "教师团队", "讲师阵容", "师资力量", "老师", "讲师"],
    "培训": ["培训", "辅导", "学习计划", "教育课程"],
    "在哪里": ["在哪里", "位于何处", "设在何地", "在哪"],
    "介绍": ["介绍", "说明", "讲解", "概述", "讲讲", "说说"],
    "请问": ["请问", "能否告知", "是否可以告诉我", "麻烦说下"],
    "原理": ["原理", "基本思想", "工作机制"],
    "写一个": ["写一个", "编写一个", "生成一个", "创建一个"],
    "等于多少": ["等于多少", "是多少", "结果是啥", "得多少"],
    "如何": ["如何", "怎样", "咋样"],
    "需要": ["需要", "要", "得有", "必须具备"],
    "基础": ["基础", "前提", "背景", "基本知识"]
}
def apply_synonym_variation(text, replace_prob=0.5):
    """对文本进行同义词随机替换"""
    for word, synonyms in synonym_dict.items():
        pattern = r"\b" + re.escape(word) + r"\b"
        def repl(match):
            if random.random() < replace_prob:
                return random.choice(synonyms)
            return match.group(0)
        text = re.sub(pattern, repl, text)
    return text
# 规则生成部分
def generate_generic_query_rule():
    """规则生成通用知识查询"""
    templates = [
        "什么{concept}?",
        "{concept}的定义是什么?",
        "请解释{concept}的原理。",
        "如何运用{concept}?",
        "计算{num1}+{num2}等于多少?",
        "写一个{lang}的{func}函数",
        "为什么{thing}是{state}?"
    ]
    concepts = ["AI", "Transformer模型", "Python", "递归", "算法复杂度", "数据结构", "机器学习"]
    langs = ["Python", "Java", "C++"]
    funcs = ["排序", "计算", "打印"]
    things = ["太阳", "水", "风"]
    states = ["热的", "流动的", "无形的"]
    nums = list(range(1, 100))
    t = random.choice(templates)
    if "{concept}" in t:
        replacements = {"concept": random.choice(concepts)}
    elif "{num1}" in t:
        replacements = {"num1": random.choice(nums), "num2": random.choice(nums)}
    elif "{lang}" in t:
        replacements = {"lang": random.choice(langs), "func": random.choice(funcs)}
    elif "{thing}" in t:
        replacements = {"thing": random.choice(things), "state": random.choice(states)}
    else:
        replacements = {}
    query = t.format(**replacements)
    return apply_synonym_variation(query)
def generate_professional_query_rule():
    """规则生成专业咨询查询"""
    templates = [
        "请问{subject}课程的学费是多少?",
        "{subject}的课程大纲是什么?",
        "{subject}培训的学习周期有多长?",
        "请介绍一下{subject}培训的主要项目内容。",
        "请问{subject}培训地点在哪里?"
    ]
    subjects = ["JAVA", "AI", "测试", "Web前端", "Python", "大数据", "DevOps"]
    t = random.choice(templates)
    query = t.format(subject=random.choice(subjects))
    return apply_synonym_variation(query)
# 大模型生成部分
def generate_with_qwen(prompt):
    """调用Qwen-Plus生成查询,添加超时控制"""
    try:
        completion = client.chat.completions.create(
            model="qwen-plus",
            messages=[{"role": "user", "content": prompt}],
            temperature=0.9,
            timeout=10
        )
        return completion.choices[0].message.content.strip()
    except Exception as e:
        print(f"Qwen-Plus调用失败: {e}")
        return None
def generate_generic_query_qwen():
    """使用Qwen-Plus生成通用知识查询"""
    prompt = """
    你是一个用户,生成一个“通用知识”类的查询,涉及数学计算、代码生成/纠错、概念与原理或常识性问题。
    示例:
    - “3+5等于多少?”
    - “写一个Python排序函数”
    - “什么是神经网络?”
    - “太阳为什么是热的?”
    请生成一个类似的查询,直接返回查询文本,不要多余说明。
    """
    return generate_with_qwen(prompt)
def generate_professional_query_qwen():
    """使用Qwen-Plus生成专业咨询查询"""
    prompt = """
    你是一个用户,生成一个“专业咨询”类的查询,涉及IT教育培训(如课程详情、师资、费用、周期、地点等)。
    示例:
    - “JAVA课程费用多少?”
    - “AI培训有哪些老师?”
    - “测试课程什么时候开课?”
    请生成一个类似的查询,直接返回查询文本,不要多余说明。
    """
    return generate_with_qwen(prompt)
# 保存数据集
def save_dataset(dataset, filename, stage_name):
    """保存数据集到指定文件"""
    with open(filename, "w", encoding="utf-8") as f:
        json.dump(dataset, f, ensure_ascii=False, indent=2)
    print(f"{stage_name}:已保存 {len(dataset)} 条数据到 {filename}")
def generate_training_dataset(total_samples=6000):
    """生成6000条训练数据,规则和大模型各占一半,每1500条保存"""
    num_per_category = total_samples // 2  # 3000
    num_rule = num_per_category // 2  # 1500
    num_qwen = num_per_category - num_rule  # 1500
    generic_samples = []
    professional_samples = []
    generic_set = set()
    professional_set = set()
    # 规则生成 - 通用知识
    print("生成规则通用知识数据...")
    with tqdm(total=num_rule, desc="Rule-based Generic") as pbar:
        while len(generic_samples) < num_rule:
            q = generate_generic_query_rule()
            if q not in generic_set:
                generic_set.add(q)
                generic_samples.append({"query": q, "label": "通用知识"})
                pbar.update(1)
    save_dataset(generic_samples, "rule_generic_1500.json", "规则通用知识")
    # 规则生成 - 专业咨询
    print("生成规则专业咨询数据...")
    with tqdm(total=num_rule, desc="Rule-based Professional") as pbar:
        while len(professional_samples) < num_rule:
            q = generate_professional_query_rule()
            if q not in professional_set:
                professional_set.add(q)
                professional_samples.append({"query": q, "label": "专业咨询"})
                pbar.update(1)
    save_dataset(professional_samples, "rule_professional_1500.json", "规则专业咨询")
    # Qwen-Plus生成 - 通用知识
    print("生成Qwen-Plus通用知识数据...")
    with tqdm(total=num_qwen, desc="Qwen-based Generic") as pbar:
        while len(generic_samples) < num_per_category:
            q = generate_generic_query_qwen()
            if q and q not in generic_set:
                generic_set.add(q)
                generic_samples.append({"query": q, "label": "通用知识"})
                pbar.update(1)
            time.sleep(0.5)  # 避免 API 限流
    save_dataset(generic_samples, "generic_3000.json", "通用知识(规则+Qwen)")
    # Qwen-Plus生成 - 专业咨询
    print("生成Qwen-Plus专业咨询数据...")
    with tqdm(total=num_qwen, desc="Qwen-based Professional") as pbar:
        while len(professional_samples) < num_per_category:
            q = generate_professional_query_qwen()
            if q and q not in professional_set:
                professional_set.add(q)
                professional_samples.append({"query": q, "label": "专业咨询"})
                pbar.update(1)
            time.sleep(0.5)  # 避免 API 限流
    save_dataset(professional_samples, "professional_3000.json", "专业咨询(规则+Qwen)")
    # 合并并混洗
    dataset = generic_samples + professional_samples
    random.shuffle(dataset)
    final_filename = "training_dataset_hybrid_6000.json"
    save_dataset(dataset, final_filename, "最终数据集")
    return dataset
if __name__ == "__main__":
    dataset = generate_training_dataset(total_samples=6000)
    print(f"成功生成 {len(dataset)} 条训练数据,保存在 training_dataset_hybrid_6000.json 文件中。")
    # 输出前10条作为示例
    for item in dataset[:10]:
        print(json.dumps(item, ensure_ascii=False))

3 实现细节 ¶

  • apply_synonym_variation :

  • 作用 :为规则生成查询增加多样性,通过同义词替换(如“什么” -> “啥是”)。

  • 逻辑 :正则表达式匹配词边界,50% 概率替换。

  • generate_generic_query_rule 和 generate_professional_query_rule :

  • 作用 :基于模板生成初始数据。

  • 设计 :覆盖数学计算、代码生成、概念问题和 IT 培训场景。

  • generate_with_qwen :

  • 作用 :封装 Qwen-Plus API 调用,生成自然语言查询。

  • 参数 : temperature=0.9 确保多样性, timeout=10 防止卡顿。

  • generate_generic_query_qwen 和 generate_professional_query_qwen :

  • 作用 :通过精心设计的 Prompt 指导 Qwen-Plus 生成符合类别的查询。

  • 逻辑 :提供示例,输出简洁的查询文本。

  • save_dataset :

  • 作用 :统一保存数据集,显示阶段名称和数据量。

  • 实现 :支持 JSON 格式,保存中间和最终结果。

  • generate_training_dataset :

  • 作用 :整合规则和大模型生成流程。

  • 流程 : 规则生成 1500 条“通用知识”,保存。 规则生成 1500 条“专业咨询”,保存。 Qwen-Plus 生成 1500 条“通用知识”,累计 3000 条保存。 Qwen-Plus 生成 1500 条“专业咨询”,累计 3000 条保存。 合并混洗 6000 条,保存最终数据集。

  • 规则生成 1500 条“通用知识”,保存。

  • 规则生成 1500 条“专业咨询”,保存。

  • Qwen-Plus 生成 1500 条“通用知识”,累计 3000 条保存。

  • Qwen-Plus 生成 1500 条“专业咨询”,累计 3000 条保存。

  • 合并混洗 6000 条,保存最终数据集。

  • 进度条 : tqdm 实时显示每个阶段的生成进度。

4 说明 ¶

  • 混合生成 :规则生成高效可控,Qwen-Plus 生成自然真实。

  • 分阶段保存 :每 1500 条保存一次,支持断点恢复。

  • 进度监控 : tqdm 提供直观反馈,优化用户体验。

5 执行示例 ¶

运行脚本时,输出类似: 生成规则通用知识数据…
Rule-based Generic: 100%|██████████| 1500/1500 [00:02<00:00, 750it/s]
规则通用知识:已保存 1500 条数据到 rule_generic_1500.json
生成规则专业咨询数据…
Rule-based Professional: 100%|██████████| 1500/1500 [00:02<00:00, 700it/s]
规则专业咨询:已保存 1500 条数据到 rule_professional_1500.json
生成Qwen-Plus通用知识数据…
Qwen-based Generic: 100%|██████████| 1500/1500 [05:00<00:00, 5it/s]
通用知识(规则+Qwen):已保存 3000 条数据到 generic_3000.json
生成Qwen-Plus专业咨询数据…
Qwen-based Professional: 100%|██████████| 1500/1500 [05:05<00:00, 4.9it/s]
专业咨询(规则+Qwen):已保存 3000 条数据到 professional_3000.json
最终数据集:已保存 6000 条数据到 training_dataset_hybrid_6000.json
成功生成 6000 条训练数据,保存在 training_dataset_hybrid_6000.json 文件中。

生成规则通用知识数据...
Rule-based Generic: 100%|██████████| 1500/1500 [00:02<00:00, 750it/s]
规则通用知识:已保存 1500 条数据到 rule_generic_1500.json
生成规则专业咨询数据...
Rule-based Professional: 100%|██████████| 1500/1500 [00:02<00:00, 700it/s]
规则专业咨询:已保存 1500 条数据到 rule_professional_1500.json
生成Qwen-Plus通用知识数据...
Qwen-based Generic: 100%|██████████| 1500/1500 [05:00<00:00, 5it/s]
通用知识(规则+Qwen):已保存 3000 条数据到 generic_3000.json
生成Qwen-Plus专业咨询数据...
Qwen-based Professional: 100%|██████████| 1500/1500 [05:05<00:00, 4.9it/s]
专业咨询(规则+Qwen):已保存 3000 条数据到 professional_3000.json
最终数据集:已保存 6000 条数据到 training_dataset_hybrid_6000.json
成功生成 6000 条训练数据,保存在 training_dataset_hybrid_6000.json 文件中。

12.3 数据集生成流程 ¶

1 生成流程 ¶

  • 规则生成(3000 条) :

  • 生成 1500 条“通用知识”数据,保存为 rule_generic_1500.json 。

  • 生成 1500 条“专业咨询”数据,保存为 rule_professional_1500.json .

  • 使用模板和同义词替换,确保多样性。

  • Qwen-Plus 生成(3000 条) :

  • 生成 1500 条“通用知识”数据,累计 3000 条保存为 generic_3000.json 。

  • 生成 1500 条“专业咨询”数据,累计 3000 条保存为 professional_3000.json .

  • 通过 Prompt 控制生成质量。

  • 数据整合 :

  • 合并 6000 条数据,随机混洗。

  • 保存为 training_dataset_hybrid_6000.json .

2 代码示例(使用数据集) ¶

# core/query_classifier.py
from sklearn.feature_extraction.text import TfidfVectorizer
from sklearn.naive_bayes import MultinomialNB
from sklearn.pipeline import Pipeline
import joblib
class QueryClassifier:
    def train_model(self):
        with open("training_dataset_hybrid_6000.json", "r", encoding="utf-8") as f:
            data = json.load(f)
        texts = [item["query"] for item in data]
        labels = [item["label"] for item in data]
        self.model = Pipeline([
            ("tfidf", TfidfVectorizer()),
            ("classifier", MultinomialNB()),
        ])
        self.model.fit(texts, labels)
        joblib.dump(self.model, "query_classifier_model.pkl")
        print("模型训练完成并保存")

12.4 数据集特点与优化 ¶

1 数据集特点 ¶

  • 总数 :6000 条(“通用知识” 3000 条,“专业咨询” 3000 条)。

  • 来源 :规则生成 50%(3000 条),Qwen-Plus 生成 50%(3000 条)。

  • 多样性 :

  • 规则生成:模板和同义词替换覆盖多种场景。

  • Qwen-Plus 生成:自然语言查询贴近真实用户输入。

  • 保存机制 :每 1500 条保存,支持断点续传。

  • 可视化 : tqdm 进度条提供实时反馈。

2 优化效果 ¶

  • 高效性 :简化生成流程,规则生成瞬时完成。

  • 可靠性 :分阶段保存确保数据安全。

  • 分类性能 :均衡数据集提升 QueryClassifier 准确性。

  • 用户体验 :进度条和保存日志增强交互性。

12.5 章节小结 ¶

本章节详细介绍了优化后的 generate_query_dataset_hybrid.py :

  • 功能 :混合生成 6000 条数据,规则和 Qwen-Plus 各占一半。

  • 实现 :通过模板生成 3000 条,Qwen-Plus 生成 3000 条,每 1500 条保存,进度条监控。

  • 作用 :为 QueryClassifier 提供高质量数据,优化 EduRAG 系统分类能力。

学习者掌握了高效的数据生成和保存方法,能够为 RAG 系统提供可靠支持。

本文为 程序员青阳 原创文章,遵循 CC BY-NC-SA 4.0 版权协议,转载请附上原文链接及本声明。

原文链接:https://heliufang.github.io/posts/35db8572/index.html