在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
相关产品推荐
相关产品推荐

