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

在Sklearn中使用GridSearchCV后,如何获取MLPClassifier的.loss_curve_属性?

如何获取GridSearchCV中MLP每组参数训练的损失曲线?

完全可行,只需通过GridSearchCV的**return_estimator=True**参数保存所有训练的模型实例,再从中提取损失曲线即可。以下是具体实现步骤:

步骤1:修改GridSearchCV初始化参数

在创建GridSearchCV对象时,添加return_estimator=True,同时注意多进程可能导致模型序列化问题,暂时将n_jobs设为1(后续可根据运行环境调整):

def train(
    X_TV,
    y_TV,
    train_dict,
    param_grid,
    random_state,
    ):

    mlp_nickname = 'my_mlp'
    # Initialize learner pipeline
    pipeline = Pipeline([])
    learner = MLPClassifier(random_state=random_state)
    pipeline.steps.append((mlp_nickname, learner))
    
    # Set hard-coded parameters
    pipeline.set_params(**hf.pipeline_helper(train_dict, mlp_nickname))

    # Fit with return_estimator enabled
    print('\tGridSearch verbose:')
    grid = GridSearchCV(
        pipeline,
        param_grid,
        scoring='f1_weighted',
        n_jobs=1,  # 多进程可能导致estimator无法序列化,先设为1
        refit=True,
        cv=5,
        verbose=1,
        return_train_score=False,
        return_estimator=True  # 关键:保存每个训练的模型实例
    )
    grid.fit(X=X_TV, y=y_TV)
    
    # 后续提取损失曲线的代码放在这里

步骤2:提取并绘制每组参数的损失曲线

GridSearchCV完成训练后,cv_results_字典会包含estimator字段,每个元素对应一组参数在交叉验证中训练的所有折的模型。遍历这些模型即可获取损失曲线:

import matplotlib.pyplot as plt
import numpy as np

# 遍历所有参数组合
for idx, params in enumerate(grid.cv_results_['params']):
    print(f"当前参数组合: {params}")
    # 遍历该参数组合对应的所有交叉验证模型(5折)
    fold_loss_curves = []
    for fold_idx, estimator in enumerate(grid.cv_results_['estimator'][idx]):
        # 从pipeline中取出MLP模型
        mlp = estimator[mlp_nickname]
        loss_curve = mlp.loss_curve_
        fold_loss_curves.append(loss_curve)
        # 绘制单折曲线(半透明显示)
        plt.plot(loss_curve, alpha=0.6, label=f"折 {fold_idx+1}")
    
    # 可选:绘制平均损失曲线(统一长度后计算)
    max_epochs = max(len(curve) for curve in fold_loss_curves)
    padded_curves = [np.pad(curve, (0, max_epochs - len(curve)), mode='edge') for curve in fold_loss_curves]
    mean_curve = np.mean(padded_curves, axis=0)
    plt.plot(mean_curve, linewidth=2, color='black', label='平均损失')
    
    plt.title(f"损失曲线 - 参数组合: {params}")
    plt.xlabel("训练轮次(Epoch)")
    plt.ylabel("损失值(Loss)")
    plt.legend()
    plt.tight_layout()
    plt.show()

注意事项

  • 版本要求:return_estimator=True是scikit-learn 0.21及以上版本支持的参数,请确保你的库版本符合要求。
  • 多进程问题:如果必须使用n_jobs=-1,部分环境下可能会因为多进程的序列化限制导致estimator无法正常保存,此时建议改用单进程获取损失曲线。
  • 曲线处理:不同折的训练轮次可能不同(比如MLP触发提前停止),可以用补值的方式统一长度后计算平均曲线,便于对比。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.27 07:34:56