如何用Matplotlib可视化GridSearchCV训练的岭回归性能与alpha的关系?
Hey there! Let's get your alpha vs. performance visualization sorted out—this is a super useful plot to understand how regularization affects your Ridge model.
First, let's recap the key pieces we need from your GridSearchCV object: the list of alpha values you tested, and the corresponding cross-validation (CV) scores for each. GridSearchCV stores all this in its cv_results_ attribute, which we'll leverage to build our plot.
Step 1: Extract Data from GridSearchCV
Assuming your code follows a common structure like this:
from sklearn.linear_model import Ridge from sklearn.model_selection import GridSearchCV import matplotlib.pyplot as plt import numpy as np # 你的训练数据准备(示例) X_train, y_train = ... # 替换为你的实际训练数据 # 定义模型与参数网格 ridge = Ridge() param_grid = {'alpha': [0.001, 0.01, 0.1, 1, 10, 100, 1000]} # 可根据你的参数调整 # 运行网格搜索 grid_search = GridSearchCV(ridge, param_grid, cv=5, scoring='neg_mean_squared_error') grid_search.fit(X_train, y_train)
We'll pull the critical data from the GridSearchCV results:
# 提取alpha参数和对应的CV得分(注意:如果用了neg_mean_squared_error,要取反转为正的MSE) alphas = grid_search.cv_results_['param_alpha'].data cv_scores = -grid_search.cv_results_['mean_test_score']
Step 2: Sort Data for a Smooth Plot
GridSearchCV doesn't guarantee results are ordered by alpha, so we'll sort both arrays to make the plot logical:
# 按alpha从小到大排序 sorted_indices = np.argsort(alphas) sorted_alphas = alphas[sorted_indices] sorted_cv_scores = cv_scores[sorted_indices]
Step 3: Build the Visualization
Now we can create the plot, plus highlight the best alpha found by GridSearchCV:
plt.figure(figsize=(10, 6)) # 绘制alpha与CV得分的曲线 plt.plot(sorted_alphas, sorted_cv_scores, marker='o', linestyle='-', color='#1f77b4') # 标注最佳参数与对应得分 best_alpha = grid_search.best_params_['alpha'] best_score = -grid_search.best_score_ plt.scatter(best_alpha, best_score, color='#ff4b5c', s=120, zorder=5, label=f'Best alpha: {best_alpha}\nBest MSE: {best_score:.4f}') # 优化图表可读性 plt.xscale('log') # alpha通常跨数量级,对数轴更清晰 plt.xlabel('Regularization Parameter (alpha)') plt.ylabel('Cross-Validation Mean Squared Error (MSE)') plt.title('Ridge Regression: Performance vs. Regularization Strength') plt.legend() plt.grid(True, alpha=0.3) plt.show()
Optional Enhancements
- Add Error Bars: Show CV score variability with standard deviation:
sorted_std = grid_search.cv_results_['std_test_score'][sorted_indices] plt.errorbar(sorted_alphas, sorted_cv_scores, yerr=sorted_std, fmt='o-', color='#1f77b4', capsize=5) - Adjust Scoring Metric: If you used
r2instead of MSE, skip the score negation and usegrid_search.cv_results_['mean_test_score']directly.
This should give you a clean, informative plot that matches the visualization you're aiming for—clearly showing how regularization strength impacts model performance, with a clear marker for the optimal alpha.
内容的提问来源于stack exchange,提问作者Edward Lin

