十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

Python-sklearn-模型选择

Python-sklearn-模型选择 Sklearn 模型选择与调优sklearn.model_selection模块提供交叉验证、超参数搜索、数据分割等全部模型选择工具。✂️ 数据分割1.train_test_split()— 训练集/测试集分割 ⭐fromsklearn.model_selectionimporttrain_test_split X_train,X_test,y_train,y_testtrain_test_split(X,y,test_size0.3,# 测试集比例 (0.0 - 1.0)train_sizeNone,# 训练集比例与 test_size 二选一random_state42,# 随机种子复现用shuffleTrue,# 是否先打乱stratifyy# 分层抽样保持类别比例)重要: 分类数据集不均衡时务必使用stratifyy。2.StratifiedShuffleSplit— 分层随机分割fromsklearn.model_selectionimportStratifiedShuffleSplit sssStratifiedShuffleSplit(n_splits5,test_size0.2,train_sizeNone,random_state42)fortrain_idx,test_idxinsss.split(X,y):X_train,X_testX[train_idx],X[test_idx]y_train,y_testy[train_idx],y[test_idx]3.KFold/StratifiedKFold— K 折交叉验证fromsklearn.model_selectionimportKFold,StratifiedKFold# 标准 K 折kfKFold(n_splits5,shuffleTrue,random_state42)# 分层 K 折分类推荐⭐skfStratifiedKFold(n_splits5,shuffleTrue,random_state42)fortrain_idx,val_idxinskf.split(X,y):X_train,X_valX[train_idx],X[val_idx]y_train,y_valy[train_idx],y[val_idx]4. 其他分割器fromsklearn.model_selectionimport(ShuffleSplit,# 随机多次分割StratifiedShuffleSplit,# 分层随机分割RepeatedKFold,# 重复 K 折RepeatedStratifiedKFold,# 重复分层 K 折LeaveOneOut,# 留一法LeavePOut,# 留 P 法LeaveOneGroupOut,# 基于组留一法LeavePGroupsOut,# 基于组留 P 法GroupKFold,# 组 K 折GroupShuffleSplit,# 组随机分割StratifiedGroupKFold,# 分层组 K 折TimeSeriesSplit,# 时间序列分割PredefinedSplit,# 预定义折叠)# 重复分层 K 折更稳定的评估fromsklearn.model_selectionimportRepeatedStratifiedKFold rskfRepeatedStratifiedKFold(n_splits5,n_repeats3,random_state42)# 时间序列分割避免未来信息泄露fromsklearn.model_selectionimportTimeSeriesSplit tscvTimeSeriesSplit(n_splits5)# 留一法小数据集fromsklearn.model_selectionimportLeaveOneOut looLeaveOneOut() 交叉验证1.cross_val_score()— 交叉验证评分 ⭐fromsklearn.model_selectionimportcross_val_score scorescross_val_score(model,X,y,cv5,# 折叠数或分割器实例scoringaccuracy,# 评分指标n_jobs-1,# 并行核数verbose0,error_scoreraise# or float(nan))print(f{scores.mean():.3f}/-{scores.std():.3f})2.cross_validate()— 多指标交叉验证fromsklearn.model_selectionimportcross_validate resultscross_validate(model,X,y,cv5,scoring[accuracy,f1_macro],# 多个指标return_train_scoreTrue,# 同时返回训练集分数return_estimatorTrue,# 返回每个折叠的模型n_jobs-1)print(results.keys())# dict_keys([fit_time, score_time,# test_accuracy, test_f1_macro,# train_accuracy, train_f1_macro,# estimator])3.cross_val_predict()— 交叉验证预测fromsklearn.model_selectionimportcross_val_predictfromsklearn.metricsimportconfusion_matrix y_pred_cvcross_val_predict(model,X,y,cv5,methodpredict,# predict, predict_proba, predict_log_proba, decision_functionn_jobs-1)# 基于交叉验证的混淆矩阵cmconfusion_matrix(y,y_pred_cv)4.learning_curve()— 学习曲线fromsklearn.model_selectionimportlearning_curve train_sizes,train_scores,val_scoreslearning_curve(model,X,y,cv5,train_sizesnp.linspace(0.1,1.0,10),# 训练集大小scoringaccuracy,n_jobs-1,random_state42,shuffleTrue)# 计算均值和标准差train_meantrain_scores.mean(axis1)train_stdtrain_scores.std(axis1)val_meanval_scores.mean(axis1)val_stdval_scores.std(axis1)可视化:importmatplotlib.pyplotasplt plt.plot(train_sizes,train_mean,o-,labelTraining score)plt.plot(train_sizes,val_mean,o-,labelCross-validation score)plt.fill_between(train_sizes,train_mean-train_std,train_meantrain_std,alpha0.1)plt.fill_between(train_sizes,val_mean-val_std,val_meanval_std,alpha0.1)plt.xlabel(Training examples)plt.ylabel(Score)plt.legend()plt.show()5.validation_curve()— 验证曲线fromsklearn.model_selectionimportvalidation_curve train_scores,val_scoresvalidation_curve(model,X,y,param_nameC,# 参数名param_range[0.01,0.1,1,10],# 参数值范围cv5,scoringaccuracy,n_jobs-1)# 找到最佳参数best_idxval_scores.mean(axis1).argmax()best_paramparam_range[best_idx]6.permutation_test_score()— 置换检验fromsklearn.model_selectionimportpermutation_test_score score,permutation_scores,pvaluepermutation_test_score(model,X,y,cv5,n_permutations100,n_jobs-1,random_state42,scoringaccuracy)print(f真实分数:{score:.3f})print(fp-value:{pvalue:.4f})# p 0.05 表示显著 超参数搜索1.GridSearchCV— 网格搜索 ⭐fromsklearn.model_selectionimportGridSearchCV param_grid{C:[0.01,0.1,1,10,100],gamma:[scale,auto,0.01,0.1,1],kernel:[rbf,linear,poly]}grid_searchGridSearchCV(estimatorSVC(),param_gridparam_grid,scoringaccuracy,cv5,n_jobs-1,verbose1,refitTrue,# 用最优参数在全部训练集上重训return_train_scoreTrue,error_scoreraise)grid_search.fit(X_train,y_train)# 结果查看print(f最佳参数:{grid_search.best_params_})print(f最佳分数:{grid_search.best_score_:.3f})print(f最佳模型:{grid_search.best_estimator_})print(f最佳索引:{grid_search.best_index_})# 所有结果 DataFrameimportpandasaspd results_dfpd.DataFrame(grid_search.cv_results_)print(results_df[[params,mean_test_score,std_test_score,rank_test_score]])2.RandomizedSearchCV— 随机搜索 ⭐fromsklearn.model_selectionimportRandomizedSearchCVfromscipy.statsimportuniform,loguniform,randint param_distributions{C:loguniform(1e-3,1e3),gamma:loguniform(1e-4,1e1),kernel:[rbf,linear,poly],degree:randint(2,6)}random_searchRandomizedSearchCV(estimatorSVC(),param_distributionsparam_distributions,n_iter100,# 采样次数scoringaccuracy,cv5,n_jobs-1,random_state42,refitTrue,verbose1)random_search.fit(X_train,y_train)3.HalvingGridSearchCV— 减半网格搜索自适应分配计算资源逐步淘汰差的参数组合。fromsklearn.model_selectionimportHalvingGridSearchCV halving_gridHalvingGridSearchCV(estimatorSVC(),param_gridparam_grid,factor3,# 每轮保留 1/3resourcen_samples,# 或 n_estimatorsmax_resourcesauto,min_resourcesexhaust,aggressive_eliminationFalse,cv5,scoringaccuracy,random_state42)halving_grid.fit(X,y)4.HalvingRandomSearchCV— 减半随机搜索fromsklearn.model_selectionimportHalvingRandomSearchCV halving_randomHalvingRandomSearchCV(estimatorSVC(),param_distributionsparam_distributions,n_candidatesexhaust,factor3,resourcen_samples,max_resourcesauto,scoringaccuracy,cv5,random_state42)5.ParameterGrid/ParameterSampler手动生成参数组合。fromsklearn.model_selectionimportParameterGrid,ParameterSampler param_grid{C:[0.1,1,10],kernel:[rbf,linear]}# 遍历所有组合forparamsinParameterGrid(param_grid):print(params)# {C: 0.1, kernel: rbf}# {C: 0.1, kernel: linear}# {C: 1, kernel: rbf}# ...# 随机采样来自分布param_dist{C:loguniform(0.01,100),kernel:[rbf,linear]}forparamsinParameterSampler(param_dist,n_iter10,random_state42):print(params) 其他工具check_cv()— 验证 CV 参数fromsklearn.model_selectionimportcheck_cv cvcheck_cv(cv5,yy,classifierTrue) 完整调优模板分类任务完整流程fromsklearn.model_selectionimporttrain_test_split,StratifiedKFold,GridSearchCVfromsklearn.preprocessingimportStandardScalerfromsklearn.pipelineimportPipelinefromsklearn.svmimportSVC# 1. 分割X_train,X_test,y_train,y_testtrain_test_split(X,y,test_size0.2,random_state42,stratifyy)# 2. 管道pipelinePipeline([(scaler,StandardScaler()),(svc,SVC())])# 3. 超参数网格param_grid{svc__C:[0.1,1,10,100],svc__gamma:[scale,auto,0.01,0.1],svc__kernel:[rbf,linear]}# 4. 网格搜索cvStratifiedKFold(n_splits5,shuffleTrue,random_state42)gridGridSearchCV(pipeline,param_grid,cvcv,scoringaccuracy,n_jobs-1,verbose1)grid.fit(X_train,y_train)# 5. 评估print(fCV最佳分数:{grid.best_score_:.3f})print(f测试集分数:{grid.score(X_test,y_test):.3f})print(f最佳参数:{grid.best_params_})回归任务fromsklearn.model_selectionimportKFold,RandomizedSearchCVfromsklearn.ensembleimportRandomForestRegressorfromscipy.statsimportrandint,uniform param_dist{n_estimators:randint(50,500),max_depth:randint(3,20),min_samples_split:randint(2,20),min_samples_leaf:randint(1,10),max_features:uniform(0.1,0.9)}rfRandomForestRegressor(random_state42)cvKFold(n_splits5,shuffleTrue,random_state42)searchRandomizedSearchCV(rf,param_dist,n_iter100,cvcv,scoringneg_mean_squared_error,n_jobs-1,verbose1)search.fit(X_train,y_train)[[sklearn-总览|← 返回总览]]
返回列表