RAG03-基于Mysql的FQA系统
基于Mysql的FQA系统
1 Python日志介绍
1.1 学习目标 ¶
理解日志记录的作用及其在程序开发中的重要性。
掌握Python logging 模块的基本用法。
学会通过示例配置日志级别、格式,并将日志存储到文件中。
在工程化项目中应用日志记录,追踪程序运行状态。
1.2 日志记录概述 ¶
1 概述 ¶
日志(Logging)是程序运行时记录关键信息的一种方式,例如操作成功、错误发生或调试信息。在开发和维护中非常重要:
调试 :帮助找到代码中的问题。
监控 :记录程序的运行状态。
审计 :追踪用户或系统的行为。
Python的 logging 模块是一个内置工具,提供灵活的日志记录功能,比 print 语句更强大。
2 核心概念 ¶
日志级别 :表示日志的重要性,常见级别从低到高:
DEBUG :最详细信息(最低级别)。
INFO :正常运行信息。
WARNING :警告,可能有问题。
ERROR :错误,已影响程序运行。
CRITICAL :严重错误(最高级别)。
日志处理器(Handler) :决定日志输出到哪里(如控制台或文件)。
日志格式(Formatter) :定义日志的显示样式(如 时间 - 级别 - 消息)。
1.3 基础日志记录 ¶
1 目标 ¶
通过简单示例展示如何记录日志到控制台。
2 代码 ¶
import logging
# 配置基本的日志设置
# DEBUG/INFO/WARNING/ERROR/CRITICAL
# 注意:实际生产一般是INFO;本地调试一般是DEBUG
logging.basicConfig(level=logging.INFO)
# 获取日志记录器
logger = logging.getLogger("Example1")
# 记录不同级别的日志,只显示INFO及以上的日志
logger.debug("这是调试信息,通常用于开发")
logger.info("程序运行正常")
logger.warning("注意,可能有小问题")
logger.error("发生错误")
logger.critical("严重错误,程序可能崩溃")
INFO:Example1:程序运行正常
WARNING:Example1:注意,可能有小问题
ERROR:Example1:发生错误
CRITICAL:Example1:严重错误,程序可能崩溃
3 分析 ¶
basicConfig(level=logging.INFO) 设置最低记录级别为 INFO ,因此 DEBUG 信息未显示。
日志默认输出到控制台,格式为 时间 级别 名称: 消息 。
1.4 自定义日志格式 ¶
1 目标 ¶
自定义日志输出格式,添加时间和级别。
2 代码 ¶
import logging
# 配置日志格式
logging.basicConfig(
level=logging.DEBUG,
# format='%(asctime)s - %(levelname)s - %(message)s'
format='%(asctime)s - %(levelname)s - %(module)s %(lineno)d - %(message)s'
)
"""
asctime: 时间
levelname: 日志级别
module: 当前脚本名称,去掉.py之后的文件名称
lineno: 行号
message: 日志信息
"""
# 获取日志记录器
logger = logging.getLogger("Example2")
# 记录日志
logger.debug("调试模式已开启")
logger.info("正在处理数据")
logger.error("数据处理失败")
INFO:Example2:正在处理数据
ERROR:Example2:数据处理失败
3 分析 ¶
format 参数使用占位符:
%(asctime)s :记录时间。
%(levelname)s :日志级别。
%(message)s :日志消息。
1.5 将日志存储到文件 ¶
1 目标 ¶
将日志保存到文件中,便于后续查看。
2 代码 ¶
import logging
# 配置日志,输出到文件
logging.basicConfig(
level=logging.INFO,
format='%(asctime)s - %(levelname)s - %(module)s %(lineno)d - %(message)s',
filename='app.log', # 日志文件路径、
encoding='utf-8',
# 没有特殊情况,不要使用w
filemode='a' # 'a'表示追加,'w'表示覆盖
)
# 获取日志记录器
logger = logging.getLogger("Example3")
# 记录日志
logger.info("程序启动")
logger.warning("内存使用率较高")
logger.error("无法连接数据库")
INFO:Example3:程序启动
WARNING:Example3:内存使用率较高
ERROR:Example3:无法连接数据库
3 分析 ¶
filename 指定日志文件路径。
filemode=’a’ 确保日志追加写入,不覆盖之前的内容。
1.6 同时输出到控制台和文件 ¶
1 目标 ¶
将日志同时记录到控制台和文件中。
2 代码 ¶
import logging
# 创建日志记录器
logger = logging.getLogger("Example4")
# TODO 全局配置,如果logger还有一些局部的配置,以局部配置为主。如果局部配置为空,使用全局配置
logger.setLevel(logging.DEBUG) # 设置记录器级别
# 创建控制台处理器
# TODO StreamHandler: 控制台处理器。 把日志打印到控制台
console_handler = logging.StreamHandler()
console_handler.setLevel(logging.INFO) # 控制台显示INFO及以上级别
# 创建文件处理器
# TODO FileHandler:文件处理器。 把日志打印到文件中
file_handler = logging.FileHandler('app.log', mode='a', encoding='utf-8')
file_handler.setLevel(logging.DEBUG) # 文件记录DEBUG及以上级别
# 定义日志格式
# TODO Formatter:日志格式定义。
formatter = logging.Formatter('%(asctime)s - %(levelname)s - %(module)s %(lineno)d - %(message)s')
# 为处理器设置格式
console_handler.setFormatter(formatter)
file_handler.setFormatter(formatter)
# 将处理器添加到记录器
logger.addHandler(console_handler)
logger.addHandler(file_handler)
# 记录日志
logger.debug("调试信息,仅写入文件")
logger.info("程序运行正常")
logger.error("发生错误")
DEBUG:Example4:调试信息,仅写入文件
2026-06-30 00:04:39,000 - INFO - 1179218226 32 - 程序运行正常
INFO:Example4:程序运行正常
2026-06-30 00:04:39,000 - ERROR - 1179218226 33 - 发生错误
ERROR:Example4:发生错误
- app.log文件内容 :
2025-04-01 10:00:00,123 - DEBUG - 调试信息,仅写入文件
2025-04-01 10:00:00,124 - INFO - 程序运行正常
2025-04-01 10:00:00,125 - ERROR - 发生错误
3 分析 ¶
使用 Handler 分别控制输出目标:
StreamHandler :输出到控制台。
FileHandler :输出到文件。
不同处理器可设置不同级别,灵活性更高。
1.7 代码实现 ¶
1 整体结构 ¶
logging_lesson/
├── logs/
│ └── app.log # 日志文件
├── base/
│ └── logger.py # 日志配置模块
└── main.py # 主程序入口
2 具体模块 ¶
日志配置模块 ( utils/logger.py ) ¶
import logging
import os
# logger.py是一个公共模块,用于在项目中所有的代码逻辑中,加入日志打印。
def setup_logger(name, log_file='logs/app.log'):
"""
构建一个通用的logger对象
:param name: logger的名字
:param log_file: 日志存放的位置。 log_file: 1.从配置文件读取(物理机部署) 2.根据项目启动的根目录计算的相对路径(容器启动)
:return: logger
"""
# 确保日志目录存在
os.makedirs(os.path.dirname(log_file), exist_ok=True)
# 创建日志记录器
logger = logging.getLogger(name)
logger.setLevel(logging.DEBUG) # 设置最低级别
# 创建控制台处理器
console_handler = logging.StreamHandler()
console_handler.setLevel(logging.INFO)
# 创建文件处理器
file_handler = logging.FileHandler(log_file, mode='a', encoding='utf-8')
file_handler.setLevel(logging.DEBUG)
# 定义日志格式
formatter = logging.Formatter('%(asctime)s - %(levelname)s - %(name)s - %(message)s')
# 设置处理器格式
console_handler.setFormatter(formatter)
file_handler.setFormatter(formatter)
# 添加处理器(避免重复添加)
if not logger.handlers:
logger.addHandler(console_handler)
logger.addHandler(file_handler)
return logger
# 单例对象
# 获取当前文件路径
# cur_path = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
# log_file_path = os.path.join(cur_path, 'logs/app.log')
logger = setup_logger("edu_rag", log_file='logs/app.log')
主程序 ( main.py ) ¶
# from base.logger import logger
# 初始化日志记录器
def process_data(data):
logger.debug(f"开始处理数据: {data}")
if not data:
logger.error("数据为空,无法处理")
return None
logger.info("数据处理完成")
return data.upper()
def main():
logger.info("程序启动")
result = process_data("hello")
if result:
logger.info(f"处理结果: {result}")
else:
logger.warning("处理失败")
logger.info("程序结束")
if __name__ == "__main__":
main()
- logs/app.log 文件内容 :
2025-04-01 10:00:00,123 - INFO - MainApp - 程序启动
2025-04-01 10:00:00,124 - DEBUG - MainApp - 开始处理数据: hello
2025-04-01 10:00:00,125 - INFO - MainApp - 数据处理完成
2025-04-01 10:00:00,126 - INFO - MainApp - 处理结果: HELLO
2025-04-01 10:00:00,127 - INFO - MainApp - 程序结束
1.8 本章小结 ¶
本课通过示例讲解了Python logging 模块的使用:
基础 :使用 basicConfig 快速配置日志。
自定义 :设置格式、级别和输出目标。
存储 :将日志保存到文件,同时支持控制台输出。
工程化 :封装日志配置为模块,便于复用。
1 应用场景 ¶
在QA系统中,记录数据库连接、检索结果等状态。
调试时使用 DEBUG 级别,生产环境调整为 INFO 或 ERROR 。
2 Redis数据库简介
2.1 学习目标 ¶
理解 Redis 数据库的基本原理及其在缓存和数据存储中的作用。
掌握如何使用 Python 的 redis 库进行数据存储和查询。
学会将 Redis 客户端集成到工程化代码中。
2.2 Redis 数据库概述 ¶
Redis(Remote Dictionary Server)是一个高性能的键值对数据库,常用于缓存、会话管理等场景。它支持多种数据结构(如字符串、哈希、列表等),并提供快速的内存操作。
1 Redis 的核心特性 ¶
高性能 :数据存储在内存中,读写速度极快。
持久化 :支持 RDB 和 AOF 两种持久化方式。
灵活性 :支持多种数据类型和丰富命令。
简单易用 :提供直观的 API,易于集成。
2 应用场景 ¶
缓存查询结果以减少数据库压力。
存储用户会话信息。
实现排行榜或计数器功能。
2.3 代码实现 ¶
1 整体结构 ¶
redis_lesson/
├── base/
│ └── logger.py # 日志配置模块
├── cache/
│ └── redis_client.py # Redis 客户端模块
├── logs/
│ └── app.log # 日志文件
|── main.py # 主程序入口
2 具体模块 ¶
Redis 客户端模块 ( redis_client.py ) ¶
import json
import redis
# from base import Config
class RedisClient:
def __init__(self):
self.logger = logger
config = Config()
try:
self.client = redis.StrictRedis(
host=config.REDIS_HOST,
port=config.REDIS_PORT,
password=config.REDIS_PASSWORD,
db=config.REDIS_DB,
decode_responses=True,
)
self.logger.info("Redis 连接成功")
except redis.RedisError as e:
self.logger.error(f"Redis 连接失败: {e}")
raise
def set_data(self, key, value):
try:
self.client.set(key, json.dumps(value, ensure_ascii=False))
self.logger.info(f"存储数据到 Redis: {key}")
except redis.RedisError as e:
self.logger.error(f"Redis 存储失败: {e}")
def get_data(self, key):
try:
data = self.client.get(key)
return json.loads(data) if data else None
except redis.RedisError as e:
self.logger.error(f"Redis 获取失败: {e}")
return None
def get_answer(self, query):
try:
answer = self.client.get(f"answer:{query}")
if answer:
self.logger.info(f"从 Redis 获取答案: {query}")
return answer
return None
except redis.RedisError as e:
self.logger.error(f"Redis 查询失败: {e}")
return None
def show_all_data(self):
# 展示 Redis 中的所有数据
try:
# 获取所有 keys
keys = self.client.keys("*")
if not keys:
print("Redis 中无数据")
self.logger.info("Redis 中无数据")
return
print(f"\n{'='*60}")
print(f"Redis 中共有 {len(keys)} 个键值对:")
print(f"{'='*60}")
for key in keys:
value = self.client.get(key)
# 尝试解析 JSON 格式的数据
try:
parsed_value = json.loads(value) if value else None
print(f"键: {key}")
print(f"值: {json.dumps(parsed_value, indent=2, ensure_ascii=False)}")
except (json.JSONDecodeError, TypeError):
# 如果不是 JSON 格式,直接显示
print(f"键: {key}")
print(f"值: {value}")
print("-" * 60)
self.logger.info(f"成功展示 Redis 数据,共 {len(keys)} 个键值对")
except redis.RedisError as e:
print(f"查看数据失败: {e}")
self.logger.error(f"查看 Redis 数据失败: {e}")
主程序 ( main.py ) ¶
# from redis_client import RedisClient
import logging
# 配置日志
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
logger = logging.getLogger(__name__)
def main():
# 初始化 Redis 客户端
redis_client = RedisClient()
# 示例数据
key = "user:1"
value = {"name": "Alice", "age": 25}
# 存储数据
redis_client.set_data(key, value)
# 展示redis_client中所有数据
redis_client.show_all_data()
# 获取数据
result = redis_client.get_data(key)
if result:
logger.info(f"查询结果: {result}")
else:
logger.info("未找到数据")
# 示例查询缓存
query = "test_query"
answer = redis_client.get_answer(query)
if answer:
logger.info(f"缓存答案: {answer}")
else:
logger.info("未找到缓存答案")
if __name__ == "__main__":
main()
依赖文件 ( requirements.txt ) ¶
redis
2.4 示例运行结果 ¶
1 运行 main.py ¶
假设 Redis 服务器运行在本地,输出如下:
2025-05-12 10:00:01,123 - INFO - Redis 连接成功
2025-05-12 10:00:01,124 - INFO - 存储数据到 Redis: user:1
2025-05-12 10:00:01,125 - INFO - 查询结果: {'name': 'Alice', 'age': 25}
2025-05-12 10:00:01,126 - INFO - 未找到缓存答案
2 分析 ¶
数据以 JSON 格式存储,适合复杂结构。
get_answer 方法用于查询缓存,减少重复计算。
异常处理确保代码鲁棒性,避免 Redis 连接或操作失败导致程序崩溃。
2.5 本章小结 ¶
本节主要介绍了 Redis 数据库的操作原理和应用:
原理 :高性能键值对存储,适合缓存和快速数据访问。
应用 :通过 redis 库实现数据存储、查询和缓存。
下一章将介绍如何将 Redis 与其他算法(如 BM25)结合,提升检索效率。
3 基于Mysql实现FQA问答系统
3.1 学习目标 ¶
理解FQA问答系统的整体流程。
掌握如何整合MySQL、Redis和BM25算法构建QA问答系统。
3.2 FQA系统概述 ¶
本系统从MySQL数据库检索问答对,使用BM25算法计算相似度,并通过Softmax归一化将得分转换为概率值,阈值0.85判断答案可靠性。Redis仅缓存高可靠性结果(相似度>0.85且有答案)。若MySQL无可靠答案,则调用RAG系统检索。
1 系统流程 ¶
flowchart LR;
%% =========================
%% 离线阶段
%% =========================
subgraph 数据准备-离线阶段
A[1.结构化FQA数据]
B[(2.MySQL数据库)]
C[3.jieba分词]
E[4.构建BM25检索器]
A --> B
B --> C
C --> E
end
style 数据准备-离线阶段 fill:none,stroke:#000000
%% =========================
%% 在线阶段
%% =========================
subgraph 问题检索-在线阶段
U[1.用户问题]
V[2.jieba分词]
W[3.BM25相似度检索]
X{4.双重阈值判断}
Y[5.1 Redis查询FQA答案]
Z[6.返回FQA答案]
N[5.2 调用RAG系统]
U --> V
V --> W
W --> X
X -- 是 --> Y
Y --> Z
X -- 否 --> N
end
style 问题检索-在线阶段 fill:none,stroke:#000000
%% =========================
%% 共享组件
%% =========================
R[(Redis缓存FQA)]
%% 在线读缓存
Y --> R
%% Redis未命中查MySQL
R -.未命中.-> B
%% 离线BM25提供在线检索
E -.BM25索引.-> W
数据准备-离线阶段
存储结构化FQA数据到MySQL数据库。
使用jieba分词后的数据构建BM25检索器。
将jieba分词后的数据写入Redis缓存。
问题检索-在线阶段
对用户问题进行jieba分词。
使用BM25计算相似度。
判断相似度是否超过阈值0.85。
若超过,查询Redis缓存FQA答案并返回,若未命中则查询MySQL数据库FQA答案并返回。
若未超过,返回无答案,后续调用RAG系统。
2 项目结构 ¶
integrated_qa_system/
├── config.ini # 配置文件,包含所有模块的配置
├── base/
│ ├── config.py # 配置管理,加载 config.ini
│ ├── logger.py # 日志设置
├── mysql_qa/
│ ├── data/
│ │ ├── JP学科知识问答.csv # FQA数据集
│ ├── db/
│ │ ├── mysql_client.py # MySQL 数据库操作
│ ├── cache/
│ │ ├── redis_client.py # Redis 缓存操作
│ ├── retrieval/
│ │ ├── bm25_search.py # BM25 搜索
│ ├── utils/
│ │ ├── preprocess.py # 文本预处理
│ ├── main.py # MySQL 系统独立入口,支持查询
├── requirements.txt # 依赖文件
└── logs/
└── app.log # 日志文件
3.3 代码实现 ¶
1 配置文件 ( config.ini ) ¶
# MySQL 配置
[mysql]
host = localhost
user = root
password = 123456
database = subjects_kg
# Redis 配置
[redis]
host = localhost
port = 6379
password = 1234
db = 0
# 日志配置
[logger]
log_file = /path/to/your/logs/app.log
2 配置管理 ¶
功能 ¶
config.py 文件定义了 Config 类,用于集中管理系统中的所有配置参数。这些参数包括数据库连接信息、模型选择、分块策略、API设置等。通过集中管理配置,系统可以方便地调整参数、适配不同环境,并支持通过环境变量进行灵活配置。
代码实现 ¶
# base/config.py
# 导入配置解析库
import configparser
# 导入路径操作库
import os
# __file__:当前文件路径
# 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, encoding='utf-8')
# MySQL 配置
# MySQL 主机地址
self.MYSQL_HOST = os.getenv('MYSQL_HOST', self.config.get('mysql', 'host', fallback='127.0.0.1'))
# 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='127.0.0.1'))
# 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='127.0.0.1'))
# 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.PROJECT_ROOT)
print(conf.MYSQL_USER)
print(conf.DASHSCOPE_API_KEY)
print(conf.LOG_FILE)
d:\EduRAG_20260412_dev\03-笔记
root
sk-b7314bb9c71a444293456ce5dcedf57e
d:\EduRAG_20260412_dev\03-笔记\logs\app.log
说明 ¶
默认值 :每个参数设有默认值,确保未配置环境变量时系统仍可运行。
参数分类 :按功能分类(如数据库、模型、分块等),便于管理和维护。
3 日志记录 ¶
功能 ¶
logger.py 文件定义了 setup_logging 函数,用于配置系统的日志记录器。日志记录器将运行信息、警告和错误输出到文件和控制台,便于开发、调试和运维人员监控系统状态。
代码实现 ¶
import logging
import os
# 导入配置类
# from base.config import config
# logger.py是一个公共模块,用于在项目中所有的代码逻辑中,加入日志打印。
def setup_logger(name, log_file='logs/app.log'):
"""
构建一个通用的logger对象
:param name: logger的名字
:param log_file: 日志存放的位置。 log_file: 1.从配置文件读取(物理机部署) 2.根据项目启动的根目录计算的相对路径(容器启动)
:return: logger
"""
# 确保日志目录存在
os.makedirs(os.path.dirname(log_file), exist_ok=True)
# 创建日志记录器
logger = logging.getLogger(name)
logger.setLevel(logging.DEBUG) # 设置最低级别
# 创建控制台处理器
console_handler = logging.StreamHandler()
console_handler.setLevel(logging.INFO)
# 创建文件处理器
file_handler = logging.FileHandler(log_file, mode='a', encoding='utf-8')
file_handler.setLevel(logging.DEBUG)
# 定义日志格式
formatter = logging.Formatter('%(asctime)s - %(levelname)s - %(name)s - %(module)s %(lineno)d- %(message)s')
# 设置处理器格式
console_handler.setFormatter(formatter)
file_handler.setFormatter(formatter)
# 添加处理器(避免重复添加)
if not logger.handlers:
logger.addHandler(console_handler)
logger.addHandler(file_handler)
return logger
# 单例对象
logger = setup_logger("edu_rag", log_file=config.LOG_FILE)
说明 ¶
日志级别 :默认设为 INFO ,记录关键运行信息。
双重输出 :同时输出到文件和控制台,便于实时监控和后续分析。
格式化 :日志包含时间戳、名称、级别和内容,便于问题定位。
4 MySQL操作模块 ¶
功能 ¶
mysql_client.py 是一个用于与 MySQL 交互的模块。模块通过读取配置文件连接数据库,支持创建表、从 CSV 文件插入数据、查询问题和答案,以及安全关闭连接。所有操作均通过日志记录,便于调试和监控系统状态。
代码实现 ¶
需要提前创建mysql database subjects_kg
命令行方法:
- 启动 mysql 服务
启动本机mysql
# windows,修改为实际版本号,比如MySQL81,MySQL82
net start MySQL80
# macos
brew services start mysql
# linux
sudo systemctl start mysql
进入mysql的docker容器
docker exec -it mysql bash
- 登录 mysql 服务
mysql -u root -p
- 查看数据库
SHOW DATABASES;
- 可选: 创建数据库 subjects_kg
CREATE DATABASE subjects_kg;
- 退出 mysql
EXIT;
- 关闭 mysql 服务
# windows
net stop MySQL80
# macos
brew services stop mysql
# linux
sudo systemctl stop mysql
"""
mysql_client模块
实现功能:
1.初始化mysql客户端
2.创建高频问答数据表table
3.写入csv到mysql
4.读取所有问题
5.根据问题获取对应的答案
"""
import pymysql # pip install pymysql
import pandas as pd
# from base.config import config
# from base.logger import logger
class MysqlClient:
# 1.初始化mysql客户端
def __init__(self):
try:
# 1.mysql连接
self.connect = pymysql.connect(
host=config.MYSQL_HOST, # mysql 地址
user=config.MYSQL_USER, # mysql 用户名
password=config.MYSQL_PASSWORD, # mysql 密码
port=3306, # mysql 端口号
)
# 2.创建mysql游标, 类比 mysql数据库的执行员,用于执行sql语句并获取结果
self.cursor = self.connect.cursor()
# 3.创建database
self.create_database()
logger.info("mysql初始化 成功")
except Exception as e:
logger.error(f"mysql初始化 失败: {e}")
raise
# 2.创建database
def create_database(self):
# 1.创建sql语句
db_name = config.MYSQL_DATABASE
# 数据库名来自配置文件,无注入风险
sql = f"CREATE DATABASE IF NOT EXISTS {db_name}"
try:
# 2.执行sql,创建 database
self.cursor.execute(sql)
# 3.切换到目标数据库
self.cursor.execute(f"USE {db_name}")
logger.info(f"创建database {db_name} 成功")
except Exception as e:
logger.error(f"创建database {db_name} 失败: {e}")
raise
# 3.创建高频问答数据表table
def create_table(self):
# 1.创建sql语句
# 要保证question字段是唯一的, 因为问题不能重复
sql = """
CREATE TABLE IF NOT EXISTS jpkb
(
id INT AUTO_INCREMENT PRIMARY KEY,
subject_name VARCHAR(20),
question VARCHAR(500) UNIQUE,
answer TEXT
)
"""
try:
# 2.执行sql,创建 数据表
self.cursor.execute(sql)
# 3.提交,会修改数据
self.connect.commit()
logger.info("创建高频问答数据表 成功")
except Exception as e:
logger.error(f"创建高频问答数据表 失败: {e}")
raise
# 4.写入csv到mysql
def insert_data(self, csv_path):
# 1.读取csv
df = pd.read_csv(csv_path)
try:
# 2.遍历每一行
for index, row in df.iterrows():
# 3.获取 学科名称,问题,答案
question = row["问题"]
answer = row["答案"]
subject_name = row["学科名称"]
# 4.插入数据到mysql
# 创建sql
# 方法1:严禁使用!字符串拼接,存在sql注入风险(直接把用户输入的内容拼接到sql语句中)
# sql = f"""
# INSERT INTO jpkb (question, answer, subject_name) VALUES ('{question}', '{answer}', '{subject_name}')
# """
# 比如 输入 question = "a", answer = "a", subject_name = "a; DROP TABLE jpkb;"
# 方法2:参数化查询,使用 %s 占位符,分离 sql语句 和 用户输入,可以防止sql注入攻击
# 如果 question已经存在,则update answer 和 subject_name
sql = """
INSERT INTO jpkb (question, answer, subject_name) VALUES (%s, %s, %s)
ON DUPLICATE KEY UPDATE
answer = VALUES(answer),
subject_name = VALUES(subject_name)
"""
# 执行sql
self.cursor.execute(sql, (question, answer, subject_name))
# 提交,一次性提交
# 事务:一组操作,确保要么全部成功,要么全部失败
# 类比:转账,A给B转了100元,A刚转出,A账户-100,B还没有收到,然后断电了;会导致A投诉;
# 将这个转账操作封装为一个事务,要么全部成功(A转出100元,A账户-100元,B收到100元,B账户+100元),要么全部失败(A和B的账户都不变)
self.connect.commit()
logger.info(f"写入高频问答数据表 成功, 写入了 {len(df)} 条数据")
except Exception as e:
# 事务回滚,如果事务执行一半出错,则回滚
self.connect.rollback()
logger.error(f"写入高频问答数据表 失败: {e}")
raise
# 5.读取所有问题
def fetch_questions(self):
try:
# 1.创建sql
sql = "SELECT question FROM jpkb"
# 2.执行sql
self.cursor.execute(sql)
# 3.获取结果
results = self.cursor.fetchall()
# 4.获取questions
questions = [result[0] for result in results]
logger.info(f"读取所有问题 成功, 共 {len(questions)} 条数据")
return questions
except Exception as e:
logger.error(f"读取所有问题 失败: {e}")
raise
# 6.根据问题获取对应的答案
def fetch_answer(self, question):
try:
# 1.创建sql
# 参数化查询,使用 %s 占位符,分离 sql语句 和 用户输入,可以防止sql注入攻击
sql = "SELECT answer FROM jpkb WHERE question = %s"
# 2.执行sql
self.cursor.execute(sql, (question,))
# 3.获取答案
results = self.cursor.fetchone()
if results:
answer = results[0]
logger.info(f"根据问题获取对应的答案 成功, 问题: {question}")
return answer
else:
logger.info(f"未找到问题对应的答案, 问题: {question}")
return None
except Exception as e:
logger.error(f"根据问题 {question} 获取对应的答案 失败: {e}")
raise
# 7.关闭连接
def close(self):
try:
self.cursor.close()
self.connect.close()
logger.info("关闭MySQL连接 成功")
except Exception as e:
logger.error(f"关闭MySQL连接 失败: {e}")
# 主程序
if __name__ == '__main__':
# 1.初始化mysql客户端
mysql_client = MysqlClient()
# 2.创建高频问答数据表table
mysql_client.create_table()
# 3.写入csv到mysql
mysql_client.insert_data("./data/JP学科知识问答.csv")
# 4.读取所有问题
questions=mysql_client.fetch_questions()
# print(questions)
# 5.根据问题获取对应的答案
answer = mysql_client.fetch_answer("windows如何安装redis")
print(answer)
# 6.关闭连接
mysql_client.close()
安装3.0.504版本 下载链接,https://github.com/MicrosoftArchive/redis/releases 如果安装其他一些版本建议卸载重现安装此版本,当前gitbub只维护64位版本,如果想使用32位版本请另行查找连接
安装redis过程中记得勾选把redis添加到环境变量,如果没有添加则需要自己把安装的redis添加到环境变量中
redis-server redis.windows.conf 启动redis server
说明 ¶
数据库连接 :通过 config.ini 配置文件读取 MySQL 参数,使用 pymysql 建立连接。
表管理 :创建 jpkb 表,包含字段 id(自增主键)、subject_name(学科名称)、question(问题)、answer(答案),使用 IF NOT EXISTS 避免重复创建。
异常处理 :每个方法均捕获异常,记录错误日志并根据需要回滚事务或抛出异常。
5 Redis 缓存操作模块 ¶
功能 ¶
redis_client.py 该模块用于与 Redis 数据库交互。模块通过配置文件连接 Redis,支持键值对存储与查询(使用 JSON 序列化)、答案缓存查询,并记录操作日志,便于调试和监控。
"""
定义redis客户端,实现以下功能:
1.初始化连接
2.get_data: 读取 json字符串并反序列化
3.set_data: 将python对象 序列化转为json字符串,写入redis
4.get_answer: 根据问题读取答案, key="answer:{question}"。
5.set_answer: 写入问题和答案{key: value}, 规定 key="answer:{question}"。
示例:问题和答案{key: value}的数据示例为 {"answer:大模型学什么": "大模型学大模型"}
"""
from redis import StrictRedis
# from base.logger import logger
# from base.config import config
import json
class RedisClient():
"""
Redis 缓存客户端,支持:
- 通用JSON的存取
- 问答对的存取(自动24h过期)
"""
# 问答对缓存的默认过期时间:24h
ANSWER_EXPIRE_SECONDS = 24 * 60 * 60
def __init__(self):
"""
初始化 Redis连接
参数使用默认值即可连接本地 Docker的 Redis
"""
self.redis = StrictRedis(
host=config.REDIS_HOST, # 默认连接本地
port=config.REDIS_PORT, # 默认端口
db=config.REDIS_DB, # 默认数据库
decode_responses=True, # 返回字符串,无需手动解码
encoding='utf-8', # 编码方式
password=config.REDIS_PASSWORD
)
logger.info("redis初始化 成功")
def get_data(self, key):
"""
读取 key 的 json数据并反序列化为 python对象
:param key: redis中的 key
:return: python对象
"""
try:
# 1.获取key对应的value
value = self.redis.get(key)
if value is None:
logger.info(f"读取数据未命中: key={key}")
return None
# 2.JSON 字符串 -> Python 对象
result = json.loads(value)
logger.info(f"读取数据成功: key={key}")
return result
except Exception as e:
logger.error(f"获取数据失败: key={key}, error={e}")
return None
def set_data(self, key, value):
"""
将 value 转为 字符串,写入 Redis
:param key: # redis中的 key
:param value: # python对象
:return: # None
"""
try:
# 1.Python 对象 -> JSON 字符串
value = json.dumps(value, ensure_ascii=False)
# 2.写入redis
self.redis.set(key, value)
logger.info(f"写入数据成功: key={key}")
except Exception as e:
logger.error(f"写入数据失败: key={key}, error={e}")
def get_answer(self, question):
"""
根据问题获取缓存答案
:param question: # 问题
:return: # 答案
"""
try:
# 1.构造key
key = f"answer:{question}"
# 2.获取key对应的answer
answer = self.redis.get(key)
if answer is None:
logger.info(f"读取答案未命中: key={key}")
return None
# 3.设置过期时长
self.redis.expire(key, self.ANSWER_EXPIRE_SECONDS)
logger.info(f"读取答案成功: key={key}, answer={answer}")
# 4.返回answer
return answer
except Exception as e:
logger.error(f"获取答案失败: question={question}, error={e}")
return None
def set_answer(self, question, answer):
"""
写入问题答案
:param question: # 问题
:param answer: # 答案
:return: # None
"""
try:
# 1.构造key
key = f"answer:{question}"
# 2.写入答案
self.redis.set(key, answer, ex=self.ANSWER_EXPIRE_SECONDS)
logger.info(f"写入答案成功: key={key}, answer={answer}")
except Exception as e:
logger.error(f"写入答案失败: question={question}, answer={answer}, error={e}")
# 主程序
if __name__ == '__main__':
redis = RedisClient()
# 获取问题的答案
answer = redis.get_answer("用上下文管理器实现函数运行时间的计算?")
print(answer)
说明 ¶
Redis 连接 :通过 config.ini 读取 Redis 配置,使用 redis.StrictRedis 建立连接。
数据操作 :
set_data:将键值对(值序列化为 JSON)存储到 Redis。
get_data:根据键获取值并反序列化 JSON。
get_answer:查询以 answer:{query} 格式存储的答案缓存。
6 文本预处理模块 ¶
功能 ¶
preprocess.py 是一个基于 jieba 分词库实现文本预处理的模块。该模块将输入文本转换为小写并进行分词,返回分词结果,支持日志记录以监控处理状态。
"""
文本预处理的流程:
1. 英文统一转小写
2. 把句子进行分词
什么时候调用:
1. 从MySQL中读取问题,进行分词,并写入分词后的问题 到redis,同时给bm25算法使用
2. 用户提交query进行查询,首先进行分词
"""
import jieba
# 导入日志
# from base.logger import logger
def preprocess_text(text):
# 预处理文本
logger.debug("开始预处理文本: {}".format(text))
try:
# 分词并转换为小写
return jieba.lcut(text.lower())
except AttributeError as e:
# 记录预处理失败
logger.error(f"文本预处理失败: {e}")
# 返回空列表
return []
说明 ¶
- 文本处理 :使用 jieba.lcut 对输入文本进行中文分词,并将文本转换为小写以规范化。
7 BM25+Softmax检索模块 ¶
功能 ¶
bm25_search.py 是一个基于 BM25 算法和 Softmax 归一化的文本检索模块,用于从问题库中检索与查询最匹配的答案。模块结合 Redis 缓存和 MySQL 数据库,支持问题加载、分词、BM25 评分、Softmax 归一化,并记录操作日志。
# retrieval/bm25_search.py
# 导入 BM25 算法
from rank_bm25 import BM25Okapi
# 导入数值计算库
import numpy as np
# # 导入文本预处理
# from mysql_qa.utils.preprocess import preprocess_text
# # 导入日志
# from base.logger import logger
# from mysql_qa.db.mysql_client import MysqlClient
# from mysql_qa.cache.redis_client import RedisClient
# 初始化问题名称
ORIGIN_QUESTION_KEY = "edurag:origin_questions"
QUESTION_KEY = "edurag:questions"
class BM25Search:
# 1.初始化
def __init__(self, redis_client: RedisClient, mysql_client: MysqlClient):
# 1.初始化日志
self.logger = logger
# 2.初始化 redis客户端
self.redis_client = redis_client
# 3.初始化 mysql客户端
self.mysql_client = mysql_client
# 4.初始化 bm25 模型
self.bm25 = None
# 5.初始化 原始问题列表(没有分词)
self.origin_questions = None
# 6.初始化 分词后的问题列表
self.questions = None
# 7.加载数据
self._load_data()
# 2.加载数据,mysql -> redis -> bm25
def _load_data(self):
"""
实现FQA 的加载数据功能
1.判断系统是不是第一次启动
2.如果是第一次启动
1.从mysql中读取所有问题
2.所有问题列表origin_questions 写入redis
3.分词后的问题列表questions 写入redis
3.用分词后的问题列表 构建bm25检索器
:return: None
"""
# 1.判断系统是不是第一次启动
# 判断 origin_questions, questions 是否都存在
origin_questions = self.redis_client.get_data(ORIGIN_QUESTION_KEY)
questions = self.redis_client.get_data(QUESTION_KEY)
# 2.如果是第一次启动
if not origin_questions or not questions:
self.logger.info("系统第一次启动,正在加载数据...")
# 1.从mysql中读取所有问题
origin_questions = self.mysql_client.fetch_questions()
if not origin_questions:
self.logger.error("mysql没有问题数据")
raise Exception("mysql没有问题数据")
else:
self.logger.info(f"从mysql中获取了问题,问题数量:{len(origin_questions)}")
# 2.所有问题列表origin_questions 写入redis
self.redis_client.set_data(ORIGIN_QUESTION_KEY, origin_questions)
self.logger.info(f"所有问题写入redis成功,问题数量:{len(origin_questions)}")
# 3.分词后的问题列表questions 写入redis
questions = [preprocess_text(question) for question in origin_questions]
self.redis_client.set_data(QUESTION_KEY, questions)
self.logger.info(f"分词后的问题写入redis成功,问题数量:{len(questions)}")
# 3.用分词后的问题列表 构建bm25检索器
self.origin_questions = origin_questions
self.questions = questions
self.bm25 = BM25Okapi(self.questions)
self.logger.info("BM25模型构建成功")
# 3.查询问题
def query(self, query, threshold=[0.85, 10.0]):
"""
实现 FQA 查询:用户输入一个query,进行bm25相似度检索,返回相似度超过双重阈值的问题对应的答案
1. 判断query是否合法,非空字符串
2. 先查redis中是否有一样的问题
3. 对query分词
4. 用bm25计算query和所有问题的相似度
5. 对相似度分数进行softmax归一化
6. 取最大相似度分数,包括 原始分数 和 归一化分数
7. 双重阈值判断。包括 相对阈值0.85 和 绝对阈值10.0
8. 根据索引找到对应的原始问题(查redis)
9. 查看redis中是否有该问题的答案
10. 查看mysql中是否有该问题的答案
11. 返回答案,并写入redis
:param query: 用户输入问题
:param threshold: 阈值,包含相对阈值和绝对阈值
:return: answer, True/False, 最相似的问题对应的答案, 是否调用RAG系统
"""
# 1. 判断query是否合法,非空字符串
if not isinstance(query, str) or not query.strip():
logger.info(f"用户输入的query非法: {query}")
# query非法,不需要进入RAG系统
return None, False
# 2. 先查redis中是否有一样的问题
answer = self.redis_client.get_answer(query)
if answer:
logger.info(f"在redis中找到了一样问题: {query}, 答案: {answer}")
# 在redis中找到了一样的问题对应的答案,不需要进入RAG系统
return answer, False
# 3. 对query分词
query_tokens = preprocess_text(query)
# 4. 用bm25计算query和所有问题的相似度
# 相似度分数形状 1D:(len(questions),)
scores = self.bm25.get_scores(query_tokens)
# 5. 对相似度分数进行softmax归一化
scores_softmax = self._soft_max(scores)
# 6. 取最大相似度分数,包括 原始分数 和 归一化分数
max_index = np.argmax(scores_softmax)
max_score = scores[max_index]
max_score_softmax = scores_softmax[max_index]
logger.info(f"最大相似度分数: 原始分数:{max_score}, 归一化分数: {max_score_softmax}")
# 8. 根据索引找到对应的原始问题(查redis)
origin_question = self.origin_questions[max_index]
logger.info(f"最相似问题: {origin_question}")
# 7. 双重阈值判断。包括 相对阈值0.85 和 绝对阈值10.0
if max_score_softmax > threshold[0] and max_score > threshold[1]:
# 9. 查看redis中是否有该问题的答案
answer = self.redis_client.get_answer(origin_question)
if answer:
logger.info(f"在redis中找到了答案: {answer}")
# 在redis中找到了相似度最高的问题的答案,不需要进入RAG系统
return answer, False
# 10. 查看mysql中是否有该问题的答案
answer = self.mysql_client.fetch_answer(origin_question)
if answer:
logger.info(f"在mysql中找到了答案: {answer}")
# 11.返回答案,并回写答案到redis
self.redis_client.set_answer(origin_question, answer)
# 在mysql中找到了相似度最高的问题的答案,不需要进入RAG系统
return answer, False
else:
logger.info(f"在mysql中未找到答案")
# 在mysql中未找到相似度最高问题的答案,需要进入RAG系统
return None, True
# 12.如果最大相似度分数小于阈值,则返回None,True,表示问题合法,需要进入RAG系统
return None, True
def _soft_max(self, scores):
# 1.指数运算前减去最大值,避免数值溢出
exp_scores = np.exp(scores - np.max(scores))
# 2.归一化
return exp_scores / np.sum(exp_scores)
# 主程序
if __name__ == '__main__':
bm25_search = BM25Search(RedisClient(), MysqlClient())
answer, is_rag = bm25_search.query("如何在 Ubuntu 快速创建Pycharm桌面快捷方式?", threshold=[0.85, 10.0])
print(f"答案: {answer}, 是否调用RAG系统: {is_rag}")
说明 ¶
数据加载 :优先从 Redis 获取问题和分词数据,若无则从 MySQL 加载并分词后缓存到 Redis。
BM25 检索 :使用 BM25Okapi 计算查询与问题库的相似度,结合 Softmax 归一化评分。
答案查询 :通过 Redis 缓存答案,若无缓存则从 MySQL 获取并缓存,阈值(默认 0.85)控制答案可靠性。
3.4 主程序 ( main.py ) ¶
1 代码 ¶
# # 导入 MySQL 客户端
# from db.mysql_client import MysqlClient
# # 导入 Redis 客户端
# from cache.redis_client import RedisClient
# # 导入 BM25 搜索
# from retrieval.bm25_search import BM25Search
# # 导入日志
# from base.logger import logger
# 导入时间库
import time
class MySQLQASystem:
def __init__(self):
# 初始化日志
self.logger = logger
# 初始化 MySQL 客户端
self.mysql_client = MysqlClient()
# 初始化 Redis 客户端
self.redis_client = RedisClient()
# 初始化 BM25 搜索
self.bm25_search = BM25Search(self.redis_client, self.mysql_client)
def query(self, query):
# 查询 MySQL 系统
start_time = time.time()
# 记录查询信息
self.logger.info(f"处理查询: '{query}'")
# 执行 BM25 搜索
answer, _ = self.bm25_search.query(query, threshold=[0.85, 10.0])
if answer:
# 记录 MySQL 答案
self.logger.info(f"MySQL 答案: {answer}")
else:
# 记录无答案
self.logger.info("SQL中未找到答案, 需要调用RAG系统")
# 设置默认答案
answer = "SQL未找到答案"
# 计算处理时间
processing_time = time.time() - start_time
# 记录处理时间
self.logger.info(f"查询处理耗时 {processing_time:.2f}秒")
# 返回答案
return answer
def main():
# 初始化 MySQL 系统
mysql_system = MySQLQASystem()
try:
# 打印欢迎信息
print("\n欢迎使用 MySQL 问答系统!")
print("输入查询进行问答,输入 'exit' 退出。")
while True:
# 获取用户输入
query = input("\n输入查询: ").strip()
if query.lower() == "exit":
# 记录退出日志
logger.info("退出 MySQL 系统")
# 打印退出信息
print("再见!")
break
# 执行查询
answer = mysql_system.query(query)
# 打印答案
print(f"\n答案: {answer}")
except Exception as e:
# 记录系统错误
logger.error(f"系统错误: {e}")
# 打印错误信息
print(f"发生错误: {e}")
finally:
# 关闭 MySQL 连接
mysql_system.mysql_client.close()
if __name__ == "__main__":
# 运行主程序
main()
2 示例运行结果 ¶
假设MySQL中有数据:
问题:”特殊符号如何切割”,答案:”使用split函数”
问题:”如何处理字符串”,答案:”使用字符串方法”
查询:”特殊符号的切割”
3.5 本章小结 ¶
本章整合Mysql和Redis功能,实现了基于余弦相似度问答的QA系统:
流程 :MySQL存储数据,Redis缓存优化,TF-IDF和余弦相似度匹配问题。
工程化 :模块化设计、配置文件、日志记录。
4 BM25算法简介
4.1 学习目标 ¶
理解BM25算法的基本原理及其在信息检索中的作用。
掌握如何使用BM25进行文本匹配。
学会将BM25算法集成到工程化代码中。
4.2 BM25算法概述 ¶
1 介绍
BM25(Best Matching 25)是经典的信息检索排序算法,用于衡量查询Q与文档D之间的相关性。改进了TF-IDF算法,引入文档长度归一化和词频饱和机制,解决了”长文档占优”和”词频无限增大”问题, 检索结果更准确。BM25 适合做召回/粗排。
2 公式 ¶
$$\text{BM25}(Q, D)=\sum_{q_i \in Q} \text{IDF}(q_i)\cdot \frac{f(q_i,D)\cdot (k_1+1)}{f(q_i,D)+k_1\left(1-b+b\cdot \frac{|D|}{\text{avgdl}}\right)}$$
其中 IDF(逆文档频率) 为:
$$\text{IDF}(q_i)=\log(\frac{N-n_i+0.5}{n_i+0.5}+1)$$
| 符号 | 含义 |
|---|---|
| $Q$ | 查询(Query) |
| $D$ | 文档(Document) |
| $f(q_i, D)$ | 词项 $q_i$ 在文档 $D$ 中出现次数(TF) |
| $N$ | 语料库中文档总数 |
| $n_i$ | 包含词项 $q_i$ 的文档数 |
| $|D|$ | 文档长度 |
| $\text{avgdl}$ | 语料库平均文档长度 |
| $k_1, b$ | 超参数(常用 $k_1=1.2,\ b=0.75$) |
优势
词频饱和:同一个词出现很多次,分数会增加,但增幅逐渐变小(更符合现实)。
长度归一化:避免长文档仅因词更多而得分偏高,提升公平性。
稀有词更重要:词越少见(IDF 越高),对相关性贡献越大。
3 简单示例 ¶
已知:
$D_1$:
我 喜欢 编程$D_2$:
编程 很 有趣$Q$:
他 喜欢 编程
设 $N=2,\ |D_1|=|D_2|=3,\ \text{avgdl}=3,\ k_1=1.5,\ b=0.75$。
计算过程
| 词项 | $n_q$ | IDF | 词频部分计算 $D_1$ | 词频部分计算 $D_2$ |
|---|---|---|---|---|
| 他 | 0 | $\ln(6)=1.7918$ | $f(q,D_1)=0$ | $f(q,D_2)=0$ |
| 喜欢 | 1 | $\ln(2)=0.6931$ | $f(q,D_1)=1$:$\frac{1\cdot(1.5+1)}{1+1.5\left(1-0.75+0.75\cdot\frac{3}{3}\right)}=\frac{2.5}{2.5}=1$ | $f(q,D_2)=0$ |
| 编程 | 2 | $\ln(1.2)=0.1823$ | $f(q,D_1)=1$:$\frac{1\cdot(1.5+1)}{1+1.5\left(1-0.75+0.75\cdot\frac{3}{3}\right)}=\frac{2.5}{2.5}=1$ | $f(q,D_2)=1$:$\frac{1\cdot(1.5+1)}{1+1.5\left(1-0.75+0.75\cdot\frac{3}{3}\right)}=\frac{2.5}{2.5}=1$ |
计算结果
- $D_1$:喜欢 + 编程
$$
\text{BM25}(Q,D_1)=0.6931+0.1823=0.8754
$$
- $D_2$:仅编程
$$
\text{BM25}(Q,D_2)=0.1823
$$
结论: $\text{BM25}(Q,D_1) > \text{BM25}(Q,D_2)$,查询更匹配 $D_1$。
BM25和TF-IDF的区别
| 对比项 | TF-IDF | BM25 |
|---|---|---|
| 输入 | 查询Q和文档集合{D} | 查询Q和文档集合{D} |
| 输出 | 每个文本 Q/D 对应一个TF-IDF特征向量 | 查询Q与每个文档的相关性分数 |
| 输出维度 | 每个句子输出一个向量,长度为词表大小 | Q和每个文档D得到一个分数,分数数量=文档数量 |
| 是否考虑文档长度 | 否 | 是 |
| 是否考虑词频饱和 | 否 | 是 |
| 主要用途 | 文本向量化、特征提取 | 文本检索、文档排序 |
| 结果示例 | TF-IDF向量[0, 0.6931, 0.1823, 0.0, 0.0] | Q与每个文档的BM25分数,假设有[D1,D2]:[0.8754, 0.1823] |
4.3 代码实现 ¶
1 整体结构 ¶
bm25_lesson/
├── retrieval/
│ └── bm25_search.py # BM25检索模块
├── main.py # 主程序入口
└── requirements.txt # 依赖文件
2 具体模块 ¶
检索模块 ( retrieval/bm25_search.py ) ¶
import jieba
from rank_bm25 import BM25L
import logging
# 配置日志
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
logger = logging.getLogger(__name__)
class BM25Search:
def __init__(self, documents):
# 初始化文档集合
self.documents = documents
# 分词后的文档
self.tokenized_docs = [jieba.lcut(doc) for doc in documents]
# 初始化BM25模型
self.bm25 = BM25L(self.tokenized_docs)
logger.info("BM25模型初始化完成")
def search(self, query):
# 分词查询
tokenized_query = jieba.lcut(query)
try:
# 计算每个文档的BM25得分
scores = self.bm25.get_scores(tokenized_query)
print(f'scores--》{scores}')
# 获取最高得分的文档索引
best_idx = scores.argmax()
best_score = scores[best_idx]
best_doc = self.documents[best_idx]
logger.info(f"查询: {query}, 最佳匹配: {best_doc}, 得分: {best_score}")
return best_doc, best_score
except Exception as e:
logger.error(f"检索失败: {e}")
return None, 0.0
主程序 ( main.py ) ¶
# from retrieval.bm25_search import BM25Search
import logging
# 配置日志
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
logger = logging.getLogger(__name__)
def main():
# 示例文档集合
documents = ["我们喜欢编程", "编程很有趣"]
# 初始化BM25检索器
bm25_search = BM25Search(documents)
# 示例查询
query = "他不喜欢编程"
# 执行检索
result, score = bm25_search.search(query)
if result:
logger.info(f"查询结果: {result}, 得分: {score}")
else:
logger.info("未找到匹配结果")
if __name__ == "__main__":
main()
scores--》[1.09433592 0.22790195]
依赖文件 ( requirements.txt ) ¶
jieba
rank_bm25
4.4 示例运行结果 ¶
运行 main.py ,输出如下: 2025-04-02 19:01:27,463 - INFO - BM25模型初始化完成
2025-04-02 19:01:27,464 - INFO - 查询: 他喜欢编程, 最佳匹配: 我喜欢编程, 得分: 1.094
2025-04-02 19:01:27,464 - INFO - 查询结果: 我喜欢编程, 得分: 1.094
2025-04-02 19:01:27,463 - INFO - BM25模型初始化完成
2025-04-02 19:01:27,464 - INFO - 查询: 他喜欢编程, 最佳匹配: 我喜欢编程, 得分: 1.094
2025-04-02 19:01:27,464 - INFO - 查询结果: 我喜欢编程, 得分: 1.094
分析:
“编程很有趣”得分高于”我喜欢编程”,因为前者词频更高且文档较短。
BM25通过长度归一化避免了长文档的过度优势。
4.5 本章小结 ¶
本节主要介绍了BM25算法的原理和应用:
原理 :结合TF、IDF、长度归一化和词频饱和。
应用 :通过 rank_bm25 库实现文本检索。
5 扩展1 - Redis的持久化方式
Redis 是一种内存数据库,为了防止宕机导致数据丢失,它提供了两种核心持久化机制:
RDB(Redis Database Backup)快照持久化
AOF(Append Only File)日志持久化
| 对比维度 | RDB(快照持久化) | AOF(日志持久化) |
|---|---|---|
| 原理 | 某一时刻内存数据快照 | 记录所有写操作命令到日志文件 |
| 文件类型 | 二进制 dump.rdb | 纯文本 appendonly.aof |
| 数据完整性 | 可能丢失最近数据 | 更完整,最多丢 1 秒数据 |
| 安全性 | 较低(取决于快照间隔) | 高(取决于刷盘策略) |
| 恢复速度 | 很快(直接加载) | 较慢(逐条执行命令) |
| 文件大小 | 小(压缩快照) | 大(日志累积) |
| 性能影响 | 较低(fork时有开销) | 较高(持续写磁盘) |
| 宕机恢复 | 丢失最近快照后数据 | 最多丢 1 秒数据 |
| 使用场景 | 备份、灾备、快速恢复 | 金融、订单、强一致业务 |
| 如何选择 Redis 持久化方式? |
| 选择方式 | 数据安全 | 恢复速度 | 适用场景 |
|---|---|---|---|
| RDB | 中等(可能丢数据) | 快 | 可以接受少量数据丢失,比如缓存、备份 |
| AOF | 高(最多丢1秒) | 慢 | 尽量不能丢数据,比如金融、订单 |
| RDB + AOF | 很高 | 快 + 稳定 | 既要安全,又要快,比如生产环境 |
6 扩展2 - 演示 SQL 注入
6.1 句话理解
SQL 注入:把用户输入直接拼进 SQL,导致用户输入被当成 SQL 命令执行。
6.2 错误写法(不要这样写)
下面代码把 username 和 password 直接拼接到 SQL 字符串中:
username = "admin' --"
password = "123456"
sql = f"SELECT * FROM users WHERE username='{username}' AND password='{password}'"
print(sql)
可能输出:
SELECT * FROM users WHERE username='admin' --' AND password='123456'
-- 后面的内容会被当作注释,password 条件失效,存在被绕过风险。
6.3 正确写法(推荐)
使用参数化查询,不要手动拼接 SQL:
import pymysql
username = "admin' --"
password = "123456"
conn = pymysql.connect(host='localhost', user='root', password='123456', database='test')
cursor = conn.cursor()
sql = "SELECT * FROM users WHERE username=%s AND password=%s"
cursor.execute(sql, (username, password))
result = cursor.fetchall()
print(result)
参数会作为“数据”处理,不会被当成 SQL 语句的一部分。
6.4 结论
不要字符串拼接 SQL。
所有 SQL 语句一律使用参数化查询。
7 扩展3 - 使用mysql的docker镜像
7.1 拉取mysql镜像
docker pull mysql:8.0
7.2 启动容器
docker run -d --name mysql-edurag -p 3307:3306 -e MYSQL_ROOT_PASSWORD=123456 mysql:8.0
7.3 修改config.ini
添加了port=3307
# MySQL 配置
[mysql]
host = 127.0.0.1
user = root
port = 3307
password = 123456
database = subjects_kg
8 扩展4 - SSLError: [ASN1: NOT_ENOUGH_DATA] not enough data (_ssl.c:4040)

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