如何参数化编写单元测试,验证自定义模块对sklearn模型抛出NotImplemented异常
解决方案
核心思路
- 使用
pytest的参数化测试特性批量遍历多个sklearn模型,无需手动逐个定义测试用例 - 用sklearn内置的虚拟数据生成器快速创建测试数据
- 验证自研模块在处理这些模型时是否正确抛出
NotImplementedError
具体实现步骤
1. 导入依赖
import pytest from sklearn.datasets import make_classification, make_regression from sklearn.linear_model import LogisticRegression, LinearRegression from sklearn.tree import DecisionTreeClassifier, DecisionTreeRegressor from sklearn.ensemble import RandomForestClassifier, RandomForestRegressor # 导入你的自研模块 from your_module import your_custom_function
2. 定义参数化的模型列表
把需要测试的sklearn模型类整理成列表,作为参数化输入:
# 分类模型列表 CLASSIFICATION_MODELS = [ LogisticRegression, DecisionTreeClassifier, RandomForestClassifier ] # 回归模型列表(如果自研模块涉及回归场景) REGRESSION_MODELS = [ LinearRegression, DecisionTreeRegressor, RandomForestRegressor ]
3. 编写参数化测试用例
测试分类模型场景
@pytest.mark.parametrize("model_class", CLASSIFICATION_MODELS) def test_classification_models_raise_not_implemented(model_class): # 生成虚拟分类数据 X, y = make_classification(n_samples=100, n_features=5, random_state=42) # 默认参数初始化并拟合模型 model = model_class(random_state=42) model.fit(X, y) # 断言自研方法调用时抛出指定异常 with pytest.raises(NotImplementedError): your_custom_function(model)
测试回归模型场景
@pytest.mark.parametrize("model_class", REGRESSION_MODELS) def test_regression_models_raise_not_implemented(model_class): # 生成虚拟回归数据 X, y = make_regression(n_samples=100, n_features=5, random_state=42) # 默认参数初始化并拟合模型 model = model_class(random_state=42) model.fit(X, y) # 断言自研方法调用时抛出指定异常 with pytest.raises(NotImplementedError): your_custom_function(model)
关键说明
- 所有模型均使用默认参数初始化,无需关注模型合理性或效果
- 参数化测试会自动遍历列表中的每个模型类,生成独立测试用例
- 虚拟数据仅用于完成模型拟合流程,无需具备业务意义
- 如需扩展测试范围,直接在模型列表中添加新的sklearn模型类即可
内容的提问来源于stack exchange,提问作者Kyriacos Xanthos
相关产品推荐
相关产品推荐

