BGE-Reranker-v2-m3自动化测试:单元测试脚本编写指南
·
BGE-Reranker-v2-m3自动化测试:单元测试脚本编写指南
1. 引言
在RAG(检索增强生成)系统中,检索精度直接影响最终生成结果的质量。BGE-Reranker-v2-m3作为智源研究院开发的高性能重排序模型,通过Cross-Encoder架构深度分析查询与文档的逻辑匹配度,能够有效过滤检索噪音,提升系统整体表现。
本文将手把手教你如何为BGE-Reranker-v2-m3编写完整的单元测试脚本,确保模型在各种场景下都能稳定运行。无论你是刚接触RAG系统的新手,还是希望提升测试覆盖率的资深开发者,都能从本指南中获得实用价值。
2. 环境准备与基础配置
2.1 环境检查
在开始编写测试脚本前,首先确认你的环境已正确配置:
# 检查Python版本
python --version
# 建议使用Python 3.8及以上版本
# 检查必要的依赖包
pip list | grep -E "transformers|torch|numpy"
2.2 基础测试环境搭建
创建测试目录结构:
# test_setup.py
import os
import sys
import torch
from transformers import AutoModelForSequenceClassification, AutoTokenizer
def setup_test_environment():
"""设置测试环境"""
# 添加项目根目录到Python路径
project_root = os.path.dirname(os.path.abspath(__file__))
sys.path.insert(0, project_root)
# 检查GPU可用性
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"使用设备: {device}")
return device
# 初始化测试环境
test_device = setup_test_environment()
3. 核心功能单元测试
3.1 模型加载测试
# test_model_loading.py
import unittest
import time
from transformers import AutoModelForSequenceClassification, AutoTokenizer
class TestModelLoading(unittest.TestCase):
"""测试模型加载功能"""
def setUp(self):
self.model_name = "BAAI/bge-reranker-v2-m3"
self.start_time = time.time()
def test_model_loading_speed(self):
"""测试模型加载速度"""
start_time = time.time()
# 加载tokenizer
tokenizer = AutoTokenizer.from_pretrained(self.model_name)
# 加载模型
model = AutoModelForSequenceClassification.from_pretrained(
self.model_name,
torch_dtype=torch.float16 if torch.cuda.is_available() else torch.float32
)
loading_time = time.time() - start_time
print(f"模型加载时间: {loading_time:.2f}秒")
# 断言加载时间在合理范围内
self.assertLess(loading_time, 30, "模型加载时间过长")
# 检查模型和tokenizer是否成功加载
self.assertIsNotNone(tokenizer)
self.assertIsNotNone(model)
def test_model_device_placement(self):
"""测试模型设备放置"""
model = AutoModelForSequenceClassification.from_pretrained(
self.model_name,
torch_dtype=torch.float16
)
# 将模型移动到测试设备
model.to(test_device)
# 检查模型是否在正确的设备上
for param in model.parameters():
self.assertEqual(param.device, test_device)
if __name__ == '__main__':
unittest.main()
3.2 基础推理功能测试
# test_basic_inference.py
import unittest
import numpy as np
from transformers import AutoModelForSequenceClassification, AutoTokenizer
class TestBasicInference(unittest.TestCase):
"""测试基础推理功能"""
@classmethod
def setUpClass(cls):
"""一次性设置,避免重复加载模型"""
cls.model_name = "BAAI/bge-reranker-v2-m3"
cls.tokenizer = AutoTokenizer.from_pretrained(cls.model_name)
cls.model = AutoModelForSequenceClassification.from_pretrained(
cls.model_name,
torch_dtype=torch.float16
)
cls.model.eval()
cls.model.to(test_device)
def test_single_pair_scoring(self):
"""测试单对查询-文档打分"""
query = "人工智能的发展历程"
document = "人工智能从1956年达特茅斯会议开始发展,经历了多次寒冬和复兴"
# 编码输入
inputs = self.tokenizer.encode_plus(
query,
document,
max_length=512,
truncation=True,
padding=True,
return_tensors='pt'
)
inputs = {k: v.to(test_device) for k, v in inputs.items()}
# 推理
with torch.no_grad():
outputs = self.model(**inputs)
scores = outputs.logits.squeeze().cpu().numpy()
print(f"查询: {query}")
print(f"文档: {document}")
print(f"匹配分数: {scores:.4f}")
# 断言分数在合理范围内
self.assertIsInstance(scores, (float, np.floating))
self.assertGreater(scores, -10) # 分数不应该太低
self.assertLess(scores, 10) # 分数不应该太高
def test_multiple_pairs_scoring(self):
"""测试多对查询-文档批量打分"""
test_cases = [
{
"query": "机器学习的基本概念",
"documents": [
"机器学习是人工智能的一个分支,让计算机通过数据学习模式",
"深度学习是机器学习的一种,使用神经网络处理复杂模式",
"天气预报显示明天会下雨,记得带伞" # 不相关文档
]
}
]
for case in test_cases:
query = case["query"]
documents = case["documents"]
scores = []
for doc in documents:
inputs = self.tokenizer.encode_plus(
query,
doc,
max_length=512,
truncation=True,
padding=True,
return_tensors='pt'
)
inputs = {k: v.to(test_device) for k, v in inputs.items()}
with torch.no_grad():
outputs = self.model(**inputs)
score = outputs.logits.item()
scores.append(score)
print(f"\n查询: {query}")
for i, (doc, score) in enumerate(zip(documents, scores)):
print(f"文档{i+1}: {doc[:50]}... → 分数: {score:.4f}")
# 断言相关文档分数高于不相关文档
self.assertGreater(scores[0], scores[2], "相关文档分数应高于不相关文档")
self.assertGreater(scores[1], scores[2], "相关文档分数应高于不相关文档")
if __name__ == '__main__':
unittest.main()
4. 边界情况与异常处理测试
4.1 输入边界测试
# test_edge_cases.py
import unittest
from transformers import AutoModelForSequenceClassification, AutoTokenizer
class TestEdgeCases(unittest.TestCase):
"""测试边界情况和异常处理"""
@classmethod
def setUpClass(cls):
cls.model_name = "BAAI/bge-reranker-v2-m3"
cls.tokenizer = AutoTokenizer.from_pretrained(cls.model_name)
cls.model = AutoModelForSequenceClassification.from_pretrained(
cls.model_name,
torch_dtype=torch.float16
)
cls.model.eval()
cls.model.to(test_device)
def test_empty_input(self):
"""测试空输入处理"""
with self.assertRaises(ValueError):
self.tokenizer.encode_plus("", "", return_tensors='pt')
def test_long_text_truncation(self):
"""测试长文本截断"""
long_text = "人工智能 " * 200 # 创建超长文本
inputs = self.tokenizer.encode_plus(
"测试查询",
long_text,
max_length=512,
truncation=True,
padding=True,
return_tensors='pt'
)
# 检查是否成功截断
self.assertLessEqual(inputs['input_ids'].shape[1], 512)
inputs = {k: v.to(test_device) for k, v in inputs.items()}
# 应该能正常推理而不报错
try:
with torch.no_grad():
outputs = self.model(**inputs)
self.assertTrue(True) # 如果执行到这里,测试通过
except Exception as e:
self.fail(f"长文本处理失败: {e}")
def test_special_characters(self):
"""测试特殊字符处理"""
test_cases = [
("查询!!!", "文档@@@"),
("🌐 多语言测试", "🌍 国际化内容"),
("123数字", "456测试")
]
for query, doc in test_cases:
try:
inputs = self.tokenizer.encode_plus(
query,
doc,
max_length=512,
truncation=True,
padding=True,
return_tensors='pt'
)
inputs = {k: v.to(test_device) for k, v in inputs.items()}
with torch.no_grad():
outputs = self.model(**inputs)
# 如果有分数返回,说明处理成功
score = outputs.logits.item()
self.assertIsInstance(score, float)
except Exception as e:
self.fail(f"特殊字符处理失败: {e}")
if __name__ == '__main__':
unittest.main()
4.2 性能与稳定性测试
# test_performance.py
import unittest
import time
import numpy as np
from transformers import AutoModelForSequenceClassification, AutoTokenizer
class TestPerformance(unittest.TestCase):
"""测试性能与稳定性"""
@classmethod
def setUpClass(cls):
cls.model_name = "BAAI/bge-reranker-v2-m3"
cls.tokenizer = AutoTokenizer.from_pretrained(cls.model_name)
cls.model = AutoModelForSequenceClassification.from_pretrained(
cls.model_name,
torch_dtype=torch.float16
)
cls.model.eval()
cls.model.to(test_device)
def test_inference_speed(self):
"""测试推理速度"""
query = "测试查询"
document = "测试文档内容,用于性能测试"
# 预热
for _ in range(3):
inputs = self.tokenizer.encode_plus(
query, document, return_tensors='pt'
)
inputs = {k: v.to(test_device) for k, v in inputs.items()}
with torch.no_grad():
_ = self.model(**inputs)
# 正式测试
times = []
for _ in range(10):
start_time = time.time()
inputs = self.tokenizer.encode_plus(
query, document, return_tensors='pt'
)
inputs = {k: v.to(test_device) for k, v in inputs.items()}
with torch.no_grad():
outputs = self.model(**inputs)
end_time = time.time()
times.append(end_time - start_time)
avg_time = np.mean(times)
print(f"平均推理时间: {avg_time:.4f}秒")
print(f"每秒可处理查询数: {1/avg_time:.2f}")
# 断言性能在可接受范围内
self.assertLess(avg_time, 0.1, "推理时间过长")
def test_memory_usage(self):
"""测试内存使用"""
import psutil
import os
process = psutil.Process(os.getpid())
initial_memory = process.memory_info().rss / 1024 / 1024 # MB
# 执行多次推理
query = "内存测试查询"
document = "内存测试文档" * 50
for i in range(20):
inputs = self.tokenizer.encode_plus(
query, document, return_tensors='pt'
)
inputs = {k: v.to(test_device) for k, v in inputs.items()}
with torch.no_grad():
outputs = self.model(**inputs)
final_memory = process.memory_info().rss / 1024 / 1024
memory_increase = final_memory - initial_memory
print(f"内存增加: {memory_increase:.2f}MB")
# 断言内存使用在合理范围内
self.assertLess(memory_increase, 100, "内存泄漏嫌疑")
if __name__ == '__main__':
unittest.main()
5. 集成测试与实战场景
5.1 模拟真实RAG场景测试
# test_integration.py
import unittest
import numpy as np
from transformers import AutoModelForSequenceClassification, AutoTokenizer
class TestIntegration(unittest.TestCase):
"""集成测试:模拟真实RAG场景"""
@classmethod
def setUpClass(cls):
cls.model_name = "BAAI/bge-reranker-v2-m3"
cls.tokenizer = AutoTokenizer.from_pretrained(cls.model_name)
cls.model = AutoModelForSequenceClassification.from_pretrained(
cls.model_name,
torch_dtype=torch.float16
)
cls.model.eval()
cls.model.to(test_device)
def test_reranking_effectiveness(self):
"""测试重排序效果"""
# 模拟检索结果:3个相关文档,2个不相关文档
query = "如何学习深度学习"
documents = [
"深度学习是机器学习的一个分支,使用神经网络处理复杂模式", # 相关
"机器学习入门指南,包含基础概念和实践案例", # 相关
"天气预报显示本周将持续晴朗天气", # 不相关
"深度学习模型训练需要大量数据和计算资源", # 相关
"最新电影票房排行榜和影评分析" # 不相关
]
# 对每个文档进行打分
scores = []
for doc in documents:
inputs = self.tokenizer.encode_plus(
query,
doc,
max_length=512,
truncation=True,
padding=True,
return_tensors='pt'
)
inputs = {k: v.to(test_device) for k, v in inputs.items()}
with torch.no_grad():
outputs = self.model(**inputs)
score = outputs.logits.item()
scores.append(score)
# 按分数排序
ranked_indices = np.argsort(scores)[::-1] # 从高到低
ranked_docs = [documents[i] for i in ranked_indices]
ranked_scores = [scores[i] for i in ranked_indices]
print("重排序结果:")
for i, (doc, score) in enumerate(zip(ranked_docs, ranked_scores)):
print(f"{i+1}. 分数: {score:.4f} - 文档: {doc[:60]}...")
# 断言相关文档排名更高
relevant_indices = [0, 1, 3] # 相关文档的原始索引
irrelevant_indices = [2, 4] # 不相关文档的原始索引
# 检查相关文档的排名是否高于不相关文档
max_irrelevant_rank = min(ranked_indices.tolist().index(i) for i in irrelevant_indices)
min_relevant_rank = max(ranked_indices.tolist().index(i) for i in relevant_indices)
self.assertLess(min_relevant_rank, max_irrelevant_rank,
"相关文档应该排名高于不相关文档")
if __name__ == '__main__':
unittest.main()
5.2 批量处理测试
# test_batch_processing.py
import unittest
import torch
from transformers import AutoModelForSequenceClassification, AutoTokenizer
class TestBatchProcessing(unittest.TestCase):
"""测试批量处理能力"""
@classmethod
def setUpClass(cls):
cls.model_name = "BAAI/bge-reranker-v2-m3"
cls.tokenizer = AutoTokenizer.from_pretrained(cls.model_name)
cls.model = AutoModelForSequenceClassification.from_pretrained(
cls.model_name,
torch_dtype=torch.float16
)
cls.model.eval()
cls.model.to(test_device)
def test_batch_inference(self):
"""测试批量推理"""
# 创建批量测试数据
queries = ["机器学习", "人工智能", "深度学习"]
documents = [
"机器学习算法和应用案例",
"人工智能发展历史和未来趋势",
"深度学习模型和训练技巧"
]
# 准备批量输入
batch_inputs = []
for query, doc in zip(queries, documents):
inputs = self.tokenizer.encode_plus(
query,
doc,
max_length=512,
truncation=True,
padding=True,
return_tensors='pt'
)
batch_inputs.append(inputs)
# 手动批处理
batch_size = len(batch_inputs)
max_length = max(inputs['input_ids'].shape[1] for inputs in batch_inputs)
batched_inputs = {
'input_ids': torch.zeros(batch_size, max_length, dtype=torch.long),
'attention_mask': torch.zeros(batch_size, max_length, dtype=torch.long)
}
for i, inputs in enumerate(batch_inputs):
seq_len = inputs['input_ids'].shape[1]
batched_inputs['input_ids'][i, :seq_len] = inputs['input_ids'][0]
batched_inputs['attention_mask'][i, :seq_len] = inputs['attention_mask'][0]
# 移动到设备
batched_inputs = {k: v.to(test_device) for k, v in batched_inputs.items()}
# 批量推理
with torch.no_grad():
batch_outputs = self.model(**batched_inputs)
batch_scores = batch_outputs.logits.squeeze().cpu().numpy()
# 逐个推理对比
individual_scores = []
for inputs in batch_inputs:
inputs = {k: v.to(test_device) for k, v in inputs.items()}
with torch.no_grad():
outputs = self.model(**inputs)
score = outputs.logits.item()
individual_scores.append(score)
print("批量推理分数:", batch_scores)
print("逐个推理分数:", individual_scores)
# 断言结果一致(允许微小误差)
for batch_score, individual_score in zip(batch_scores, individual_scores):
self.assertAlmostEqual(batch_score, individual_score, delta=0.001)
if __name__ == '__main__':
unittest.main()
6. 测试执行与报告生成
6.1 完整的测试套件
# run_all_tests.py
import unittest
import sys
import os
# 添加测试模块
test_modules = [
'test_model_loading',
'test_basic_inference',
'test_edge_cases',
'test_performance',
'test_integration',
'test_batch_processing'
]
def run_all_tests():
"""运行所有测试"""
# 创建测试加载器
loader = unittest.TestLoader()
suite = unittest.TestSuite()
# 加载所有测试模块
for module_name in test_modules:
try:
module = __import__(module_name)
suite.addTests(loader.loadTestsFromModule(module))
except ImportError as e:
print(f"无法加载测试模块 {module_name}: {e}")
# 运行测试
runner = unittest.TextTestRunner(verbosity=2)
result = runner.run(suite)
# 生成测试报告
print("\n" + "="*50)
print("测试报告摘要:")
print(f"运行测试数: {result.testsRun}")
print(f"失败数: {len(result.failures)}")
print(f"错误数: {len(result.errors)}")
print(f"跳过数: {len(result.skipped)}")
if result.wasSuccessful():
print("所有测试通过! ✅")
return 0
else:
print("有测试未通过! ❌")
return 1
if __name__ == '__main__':
# 运行所有测试
exit_code = run_all_tests()
sys.exit(exit_code)
6.2 持续集成配置示例
# .github/workflows/test.yml
name: BGE-Reranker Tests
on:
push:
branches: [ main ]
pull_request:
branches: [ main ]
jobs:
test:
runs-on: ubuntu-latest
strategy:
matrix:
python-version: [3.8, 3.9, 3.10]
steps:
- uses: actions/checkout@v3
- name: Set up Python ${{ matrix.python-version }}
uses: actions/setup-python@v4
with:
python-version: ${{ matrix.python-version }}
- name: Install dependencies
run: |
python -m pip install --upgrade pip
pip install transformers torch numpy psutil
- name: Run tests
run: |
python run_all_tests.py
- name: Upload test results
if: always()
uses: actions/upload-artifact@v3
with:
name: test-results-${{ matrix.python-version }}
path: |
test-reports/
*.log
7. 总结
通过本指南,你学会了如何为BGE-Reranker-v2-m3编写完整的单元测试套件,涵盖了模型加载、基础推理、边界情况、性能测试、集成测试等多个方面。这些测试脚本不仅能够确保模型的稳定运行,还能在开发过程中快速发现问题。
7.1 关键要点回顾
- 环境配置:确保测试环境正确设置,包括Python路径和设备选择
- 核心功能测试:验证模型加载、基础推理等核心功能
- 边界测试:处理各种异常情况和边界条件
- 性能监控:测试推理速度和内存使用情况
- 集成测试:模拟真实RAG场景验证重排序效果
7.2 下一步建议
- 将测试集成到CI/CD流程中,实现自动化测试
- 扩展测试覆盖更多语言和特殊场景
- 添加性能基准测试,监控模型性能变化
- 考虑使用测试覆盖率工具,确保测试完整性
通过完善的测试体系,你可以更加自信地部署和使用BGE-Reranker-v2-m3,确保RAG系统在各种场景下都能提供准确可靠的检索结果。
获取更多AI镜像
想探索更多AI镜像和应用场景?访问 CSDN星图镜像广场,提供丰富的预置镜像,覆盖大模型推理、图像生成、视频生成、模型微调等多个领域,支持一键部署。
更多推荐
所有评论(0)