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

如何判断传入参数是否匹配类属性名并返回对应sklearn线性模型实例

实现方案

你需要补全依赖导入,同时增加自定义映射规则匹配你的特殊返回要求,完整可运行代码如下:

import numpy as np
import pandas as pd
import sklearn as sk
import inspect
from sklearn.model_selection import train_test_split
from sklearn.linear_model import *

# 自定义映射规则:按需求配置需要特殊替换的类
custom_mapping = {
    Lasso: Ridge
}

def model_type(linreg_model):
    # 先校验传入的类是否属于sk.linear_model下的成员
    all_linear_models = [m[1] for m in inspect.getmembers(sk.linear_model, inspect.isclass)]
    if linreg_model not in all_linear_models:
        raise ValueError("传入参数不是sklearn.linear_model下的有效类")
    
    # 匹配映射规则,无匹配则用原类
    target_class = custom_mapping.get(linreg_model, linreg_model)
    return target_class()

# 测试调用
model_type(ARDRegression) ## 返回sk.linear_model.ARDRegression()实例
model_type(Lasso)         ## 返回sk.linear_model.Ridge()实例

代码说明

  • 提前导入inspect模块用于读取sk.linear_model下的所有类成员,用于参数合法性校验
  • custom_mapping字典可以灵活配置特殊替换规则,你可以根据需求新增/修改映射关系
  • 校验逻辑可以根据你的实际需求删除,如果不需要做参数合法性判断可以直接去掉对应代码

内容的提问来源于stack exchange,提问作者Ali J

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.04 23:24:01