随机森林和交叉验证

发布时间:2026/8/4 4:11:04
随机森林和交叉验证 一、随机森林随机森林是一种集成学习算法Ensemble Learning主要用于分类和回归任务。简单来说随机森林 多棵决策树 投票/平均1、随机森林的构成随机森林主要由两部分组成 (1). 决策树Decision Tree 每棵树都是一个分类器。邮件||--是否包含免费|||是|||是否大量包含链接|||是|||垃圾邮件(2). 多棵树组合 随机森林生成很多棵决策树随机森林 树1→ 垃圾邮件 树2→ 垃圾邮件 树3→ 正常邮件 树4→ 垃圾邮件 树5→ 垃圾邮件 最终 垃圾邮件2、随机森林分类流程(1)划分数据xtrain,xtest,ytrain,ytest\ train_test_split(x,y,test_size0.2,random_state100)(2)、SMOTE处理类别不平衡oversamplerSMOTE(random_state0)#创建 SMOTE 对象os_x_train,os_y_trainoversampler.fit_resample(xtrain,ytrain)使用 SMOTE 方法解决训练集中类别不平衡问题。 (3)、创建随机森林rfRandomForestClassifier(n_estimators100,max_features0.8,random_state0)n_estimators100表示建立100个树 max_features0.8表示每次分裂随机使用80%的特征 random_state0随机种子 (4)、训练模型rf.fit(os_x_train,os_y_train)(5)预测train_predictedrf.predict(xtrain)#训练集test_predictedrf.predict(xtest)#测试机代码展示importpandasaspdimportmatplotlib.pyplotaspltfromsklearn.model_selectionimporttrain_test_splitfromsklearn.ensembleimportRandomForestClassifierfromsklearn.metricsimportconfusion_matrix,classification_reportfromimblearn.over_samplingimportSMOTE# # 绘制混淆矩阵函数# defcm_plot(y_true,y_pred):cmconfusion_matrix(y_true,y_pred)plt.matshow(cm,cmapplt.cm.Blues)plt.colorbar()foriinrange(len(cm)):forjinrange(len(cm)):plt.annotate(cm[i,j],xy(j,i),horizontalalignmentcenter,verticalalignmentcenter)plt.ylabel(True label)plt.xlabel(Predicted label)returnplt# # 1.读取数据# dfpd.read_csv(spambase.csv)# 特征Xdf.iloc[:,:-1]# 标签ydf.iloc[:,-1]# # 2.划分训练集和测试集# X_train,X_test,y_train,y_testtrain_test_split(X,y,test_size0.2,random_state100)# # 3.SMOTE处理样本不平衡# smoteSMOTE(random_state0)X_train_smote,y_train_smotesmote.fit_resample(X_train,y_train)# # 4.建立随机森林模型# rfRandomForestClassifier(n_estimators100,max_features0.8,random_state0)# 训练rf.fit(X_train_smote,y_train_smote)# # 5.训练集预测# train_predrf.predict(X_train)print(训练集结果)print(classification_report(y_train,train_pred,digits4))# 混淆矩阵cm_plot(y_train,train_pred).show()# # 6.测试集预测# test_predrf.predict(X_test)print(测试集结果)print(classification_report(y_test,test_pred,digits4))# 测试集混淆矩阵cm_plot(y_test,test_pred).show()# # 7.特征重要性分析# # 获取特征重要性importancesrf.feature_importances_ importance_dfpd.DataFrame({feature:X.columns,importance:importances})# 排名前10importance_dfimportance_df.sort_values(byimportance,ascendingFalse).head(10)# 绘制柱状图plt.figure(figsize(8,5))plt.barh(importance_df[feature],importance_df[importance])plt.xlabel(Importance)plt.ylabel(Feature)plt.title(Top 10 Feature Importance)plt.gca().invert_yaxis()plt.show()