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 关键要点回顾

  1. 环境配置:确保测试环境正确设置,包括Python路径和设备选择
  2. 核心功能测试:验证模型加载、基础推理等核心功能
  3. 边界测试:处理各种异常情况和边界条件
  4. 性能监控:测试推理速度和内存使用情况
  5. 集成测试:模拟真实RAG场景验证重排序效果

7.2 下一步建议

  • 将测试集成到CI/CD流程中,实现自动化测试
  • 扩展测试覆盖更多语言和特殊场景
  • 添加性能基准测试,监控模型性能变化
  • 考虑使用测试覆盖率工具,确保测试完整性

通过完善的测试体系,你可以更加自信地部署和使用BGE-Reranker-v2-m3,确保RAG系统在各种场景下都能提供准确可靠的检索结果。


获取更多AI镜像

想探索更多AI镜像和应用场景?访问 CSDN星图镜像广场,提供丰富的预置镜像,覆盖大模型推理、图像生成、视频生成、模型微调等多个领域,支持一键部署。

Logo

腾讯云面向开发者汇聚海量精品云计算使用和开发经验,营造开放的云计算技术生态圈。

更多推荐