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

如何将Sklearn Pipeline模型的参数保存至JSON文件?

解决Scikit-learn Pipeline参数转JSON序列化问题

问题场景

构建了包含Pipeline、StackingRegressor的嵌套机器学习模型:

from sklearn.datasets import load_diabetes
from sklearn.preprocessing import StandardScaler
from sklearn.pipeline import Pipeline
from sklearn.linear_model import RidgeCV
from sklearn.svm import LinearSVR
from sklearn.ensemble import RandomForestRegressor, StackingRegressor

X, y = load_diabetes(return_X_y=True)
estimators = [
    ('lr', RidgeCV()),
    ('svr', LinearSVR(random_state=42))
]
reg = StackingRegressor(
    estimators=estimators,
    final_estimator=RandomForestRegressor(n_estimators=10,
                                          random_state=42)
)
steps = [
    ("preprocessing", StandardScaler()),
    ("regression", reg)
]
pipe = Pipeline(steps)

想要将模型的完整参数信息保存为JSON文件,但直接使用json.dumps(pipe)会报错:Object of type Pipeline is not JSON serializable。

解决方案

Scikit-learn模型对象无法直接JSON序列化,需通过提取参数字典并处理特殊类型来实现:

1. 提取模型完整参数配置

使用模型的get_params()方法获取所有层级的参数嵌套字典,该方法会递归返回Pipeline、StackingRegressor及内部子模型的全部参数:

model_params = pipe.get_params()

2. 转换不可序列化的类型

get_params()返回的字典中可能包含numpy数组(如RidgeCV的alphas参数),需将其转换为Python原生列表:

import numpy as np

def convert_numpy_types(obj):
    if isinstance(obj, np.ndarray):
        return obj.tolist()
    elif isinstance(obj, dict):
        return {key: convert_numpy_types(value) for key, value in obj.items()}
    elif isinstance(obj, list):
        return [convert_numpy_types(item) for item in obj]
    return obj

processed_params = convert_numpy_types(model_params)

3. 序列化并保存为JSON文件

将处理后的字典转为JSON字符串并写入文件:

import json

with open("model_parameters.json", "w", encoding="utf-8") as f:
    json.dump(processed_params, f, indent=4)

结果说明

生成的JSON文件会包含所有模型的层级参数,例如:

  • preprocessing__with_std(StandardScaler的标准化参数)
  • regression__final_estimator__random_state(RandomForestRegressor的随机种子)
  • regression__estimators(StackingRegressor的子模型配置)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.23 03:24:25