You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何参数化编写单元测试,验证自定义模块对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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.08 05:50:21