京公网安备 11010802034615号
经营许可证编号:京B2-20210330
以下使用scikit-learn中数据集进行分享。
如果选用随机森林作为最终的模型,那么找出它的最佳参数可能有1000多种组合的可能,你可以使用使用穷尽的网格搜索(Exhaustive Grid Seaarch)方法,但时间成本将会很高(运行很久...),或者使用随机搜索(Randomized Search)方法,仅分析超参数集合中的子集合。
该例子以手写数据集为例,使用支持向量机的方法对数据进行建模,然后调用scikit-learn中validation_surve方法将模型交叉验证的结果进行可视化。需要注意的是,在使用validation_curve方法时,只能验证一个超参数与模型训练集和验证集得分的关系(即二维的可视化),而不能实现多参数与得分间关系的可视化。以下搜索的参数是gamma,需要给定参数范围,用param_range进行传递,评分策略用scoring参数进行传递。其代码示例如下所示:
print(__doc__) import matplotlib.pyplot as plt import numpy as np from sklearn.datasets import load_digits from sklearn.svm import SVC from sklearn.model_selection import validation_curve X, y = load_digits(return_X_y=True) param_range = np.logspace(-6, -1, 5) train_scores, test_scores = validation_curve( SVC(), X, y, param_name="gamma", param_range=param_range, scoring="accuracy", n_jobs=1) train_scores_mean = np.mean(train_scores, axis=1) train_scores_std = np.std(train_scores, axis=1) test_scores_mean = np.mean(test_scores, axis=1) test_scores_std = np.std(test_scores, axis=1) plt.title("Validation Curve with SVM") plt.xlabel(r"$\gamma$") plt.ylabel("Score") plt.ylim(0.0, 1.1) lw = 2 plt.semilogx(param_range, train_scores_mean, label="Training score", color="darkorange", lw=lw) plt.fill_between(param_range, train_scores_mean - train_scores_std, train_scores_mean + train_scores_std, alpha=0.2, color="darkorange", lw=lw) plt.semilogx(param_range, test_scores_mean, label="Cross-validation score", color="navy", lw=lw) plt.fill_between(param_range, test_scores_mean - test_scores_std, test_scores_mean + test_scores_std, alpha=0.2, color="navy", lw=lw) plt.legend(loc="best") plt.show();
代码中:
X, y = load_digits(return_X_y=True) # 等价于 digits = load_digits() X_digits = digits.data y_digits = digits.target
以下是支持向量机的验证曲线,调节的超参数gamma共有5个值,每一个点的分数是五折交叉验证(cv=5)的均值。
当想看模型多个超参数与模型评分之间的关系时,使用scikit-learn中validation curve就难以实现,因此可以考虑绘制三维坐标图。
主要用plotly的库绘制3D Scatter(3d散点图)。以下的例子使用scikit-learn中的莺尾花的数据集(iris)。以下例子选用随机森林模型(RandomForestRegressor),利用scikit-learn中的GridSearchCV方法调试最佳超参(tuning hyper-parameters),分别设置超参数"n_estimators","max_features","min_samples_split"的参数范围,详见代码如下:
import numpy as np from sklearn.model_selection import validation_curve from sklearn.datasets import load_iris from sklearn.ensemble import RandomForestRegressor from plotly.offline import iplot from plotly.graph_objs as go model = RandomForestRegressor(n_jobs=-1, random_state=2, verbose=2) grid = {'n_estimators': [10,110,200], 'max_features': [0.05, 0.07, 0.09, 0.11, 0.13], 'min_samples_split': [2, 3, 5, 8]} rf_gridsearch = GridSearchCV(estimator=model, param_grid=grid, n_jobs=4, cv=5, verbose=2, return_train_score=True) rf_gridsearch.fit(X, y) # and after some hours... df_gridsearch = pd.DataFrame(rf_gridsearch.cv_results_) trace = go.Scatter3d( x=df_gridsearch['param_max_features'], y=df_gridsearch['param_n_estimators'], z=df_gridsearch['param_min_samples_split'], mode='markers', marker=dict( # size=df_gridsearch.mean_fit_time ** (1 / 3), size = 10, color=df_gridsearch.mean_test_score, opacity=0.99, colorscale='Viridis', colorbar=dict(title = 'Test score'), line=dict(color='rgb(140, 140, 170)'), ), text=df_gridsearch.Text, hoverinfo='text' ) data = [trace] layout = go.Layout( title='3D visualization of the grid search results', margin=dict( l=30, r=30, b=30, t=30 ), scene = dict( xaxis = dict( title='max_features', nticks=10 ), yaxis = dict( title='n_estimators', ), zaxis = dict( title='min_samples_split', ), ), ) fig = go.Figure(data=data, layout=layout) iplot(fig)
其运行结果如果,是一个三维散点图(3D Scatter)。
可以看到颜色越浅,分数越高。n_estimators(子估计器)越多,分数越高,max_features的变化对模型分数的影响较小,在图中看不到变化,min_samples_split的个数并不是越高越好,但与模型分数并不呈单调关系,在min_samples_split取2时(此时,其它条件不变),模型分数最高。
除了使用scikit-learn中validation curve绘制超参数与得分的可视化,还可以利用seaborn库中heatmap方法来实现两个超参数之间的关系图,如下代码示例:
import seaborn as sns title = '''Maximum R2 score on test set VS max_features, min_samples_split''' sns.heatmap(max_scores.mean_test_score, annot=True, fmt='.4g'); plt.title(title); plt.savefig("heatmap_test.png", dpi = 300);
import seaborn as sns title = '''Maximum R2 score on train set VS max_features, min_samples_split''' sns.heatmap(max_scores.mean_train_score, annot=True, fmt='.4g'); plt.title(title); plt.savefig("heatmap_train.png", dpi = 300);
max_features和min_samples与模型得分关系的可视化如下图所示(分别为网格搜索中测试集和训练集的得分):
由于一般人很难迅速的在大量数据中找到隐藏的关系,因此,可以考虑绘图,将数据关系以图表的形式,清晰的显现出来。
综上,当关注单个超参数的学习曲线时,可以使用scikit-learn中validation curve,找到拐点,作为模型的最佳参数。
当关注两个超参数的共同变化对模型分数的影响时,可以使用seaborn库中的heatmap方法,制作“热图”,以找到超参数协同变化对分数影响的趋势。
当关注三个参数的协同变化与模型得分的关系时,可以使用poltly库中的iplot和go方法,绘制3d散点图(3D Scatter),将其协同变化对模型分数的影响展现在高维图中。
数据分析咨询请扫描二维码
若不方便扫码,搜微信号:CDAshujufenxi
在企业数字化运营、业务流程管理与精细化管控体系中,流程运营是串联各项业务环节、保障工作落地、提升运转效率的核心载体。无论 ...
2026-09-16在数据分析与统计学研究中,卡方检验是分析分类变量关联性与差异性的重要方法,广泛应用于市场调研、行为统计、社会调查、商业数 ...
2026-09-16 很多数据分析师每天盯着GMV、DAU、转化率,但当被问到“什么是指标”“指标和维度有什么区别”“如何定义指标值的计算规则和 ...
2026-09-16CDA数据分析师 出品 作者:李诗怡 定义: 用户增长核心分析框架,刻画用户从接触产品到自发推荐的全生命周期,五个递进环节构建 ...
2026-09-15在数字化营销与精细化用户运营时代,企业传统的广撒网式营销模式成本高、转化率低,已无法适配精准商业竞争需求。客户画像作为大 ...
2026-09-15 很多数据分析师精通描述性统计,能熟练计算均值、中位数、标准差,但当被问到“用500个样本如何推断10万用户的真实满意度” ...
2026-09-15在MySQL数据库数据查询与数据分析中,GROUP BY与ORDER BY是使用频率极高的核心关键字。二者语法结构相似,常搭配使用,但核心功 ...
2026-09-14随着数字化治理、智慧运营、数字孪生技术的普及,数字体征成为衡量业务状态、系统运行、城市治理与企业经营健康度的核心体系。数 ...
2026-09-14 很多数据分析师沉迷于复杂的模型和算法,却忽略了数据分析的一项基础能力——描述性统计。事实上,大量商业分析问题,用描述 ...
2026-09-14在MySQL数据库运维与开发实践中,经常出现一种典型现象:数据库实际存储的数据量很小,数据表条数少、文件体积低,但服务器整体 ...
2026-09-11 很多数据分析师能熟练计算均值、标准差,但当被问到“总体和样本有什么区别”“参数和统计量有什么关系”“数据级别的高低如 ...
2026-09-11CDA数据分析师 出品 作者:李诗怡 定义: 将同一时间段内因具备相同属性或共同经历的用户划分为群体,分析其留存与生命周期价值 ...
2026-09-11在零售、商超、餐饮、线下门店等实体商业运营中,客流与销售额是衡量门店经营状态的两大核心指标。销售额是门店经营的最终结果, ...
2026-09-10在数据可视化体系中,柱形图是最基础、应用最广泛的图表类型,其中**累计柱形图(堆积柱状图)**是兼顾整体总量与内部结构的核心 ...
2026-09-10 许多数据分析师精通Excel函数和SQL查询,但当面对一张上万行的销售明细表,要快速回答“哪个地区销量最高”“哪款产品增长最 ...
2026-09-10在Python Pandas数据分析中,DataFrame是承载结构化数据的核心载体,数据清洗、数据修正、条件赋值、字段更新等实操场景,都离不 ...
2026-09-09 很多数据分析师掌握了Excel函数、会写SQL查询,但当被问到“数据从哪里来”“数据加工有哪些步骤”“如何使用分析工具连接数 ...
2026-09-09卡方检验(Chi-Square Test)是统计学中针对分类数据的经典显著性检验方法,核心用于判断两个离散分类变量是否相互独立、数据实 ...
2026-09-09CDA数据分析师 出品 作者:李诗怡 1. 销售漏斗阶段判断 题目:销售漏斗模型中,通过广告、社交媒体等方式触达品牌信息(如浏览品 ...
2026-09-07在Python数据分析中,Pandas库的DataFrame是最核心、最常用的结构化数据表对象,类似于Excel的二维表格,具备规整的行列结构、字 ...
2026-09-07