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

如何在运行时获取Sklearn模型的类型提示与参数默认值?

运行时获取Sklearn模型的参数类型与默认值

要实现运行时自动提取Sklearn模型的参数类型提示和默认值,有两种实用方案:

方案一:结合inspect与Sklearn内部参数验证

Sklearn大部分模型都内置了参数约束规则,结合Python的inspect模块可以直接在运行时获取参数签名和默认值,再从约束中提取类型信息:

import inspect
from sklearn.utils._param_validation import _NoConstraint

def extract_model_params(model_class):
    param_info = {}
    # 获取__init__方法的签名
    sig = inspect.signature(model_class.__init__)
    
    for param_name, param in sig.parameters.items():
        if param_name == "self":
            continue
        
        # 提取默认值
        default_val = param.default if param.default is not inspect.Parameter.empty else None
        param_info[param_name] = {"default": default_val}
        
        # 从模型的参数约束中提取类型信息
        if hasattr(model_class, "_parameter_constraints"):
            constraints = model_class._parameter_constraints.get(param_name, [_NoConstraint()])
            for constraint in constraints:
                # 处理基础类型约束
                if isinstance(constraint, type):
                    param_info[param_name]["type"] = constraint.__name__
                # 处理带type属性的约束(如Interval)
                elif hasattr(constraint, "type"):
                    param_info[param_name]["type"] = constraint.type.__name__
                # 处理包含None的可选类型
                elif hasattr(constraint, "choices") and None in constraint.choices:
                    non_none_choices = [c for c in constraint.choices if c is not None]
                    if non_none_choices:
                        base_type = type(non_none_choices[0]).__name__
                        param_info[param_name]["type"] = f"{base_type} | None"
    
    return param_info

# 测试示例
from sklearn.linear_model import LinearRegression
params = extract_model_params(LinearRegression)
for name, details in params.items():
    print(f"{name}: {details.get('type', 'unknown')} = {details['default']}")

优点:无需额外依赖,直接运行时获取,适配Sklearn自身的参数校验逻辑。
缺点:部分复杂类型(如联合类型)的提取可能不够精准。

方案二:解析Typeshed stub文件

VS Code的自动补全依赖Typeshed的stub文件,你可以在运行时读取这些stub并解析类型提示,实现和IDE一致的结果:

import os
import re
from importlib.metadata import distributions

def get_typeshed_stub(module_path):
    # 定位typeshed安装目录
    for dist in distributions():
        if dist.name == "typeshed":
            typeshed_root = dist.locate_file("")
            break
    else:
        raise ImportError("请先安装typeshed:pip install typeshed")
    
    # 构建stub文件路径
    stub_file = os.path.join(typeshed_root, "third_party", "3", module_path.replace(".", os.sep) + ".pyi")
    return stub_file

def parse_stub_params(stub_file, class_name):
    param_info = {}
    with open(stub_file, "r", encoding="utf-8") as f:
        content = f.read()
    
    # 匹配类的__init__方法
    init_pattern = re.compile(
        fr"class {class_name}\(.*?\):\s*def __init__\(self(.*?)\):",
        re.DOTALL
    )
    match = init_pattern.search(content)
    if not match:
        return param_info
    
    params_raw = [p.strip() for p in match.group(1).split(",") if p.strip()]
    for param in params_raw:
        # 匹配参数名、类型、默认值
        param_pattern = re.compile(r"(?P<name>\w+):\s*(?P<type>[^=]+)(?:\s*=\s*(?P<default>.+))?")
        param_match = param_pattern.match(param)
        if param_match:
            name = param_match.group("name")
            type_str = param_match.group("type").strip()
            default = param_match.group("default").strip() if param_match.group("default") else None
            param_info[name] = {"type": type_str, "default": default}
    
    return param_info

# 测试示例
stub_path = get_typeshed_stub("sklearn.linear_model")
params = parse_stub_params(stub_path, "LinearRegression")
for name, details in params.items():
    print(f"{name}: {details['type']} = {details['default']}")

优点:能获取和IDE完全一致的类型提示,精准度高。
缺点:需要额外安装typeshed,且依赖stub文件的结构稳定性。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 09:57:45