使用Scikit-learn训练ExtraTrees回归模型

举报
yd_37369233 发表于 2026/08/05 21:42:13 2026/08/05
【摘要】 引言在机器学习回归任务中,集成学习方法因其出色的预测性能和鲁棒性而备受青睐。随机森林(Random Forest)作为Bagging家族的明星算法广为人知,但你可能不知道,它的“亲戚”——极端随机树(Extremely Randomized Trees,简称Extra Trees) 在很多时候表现同样出色,甚至在某些场景下更胜一筹。Extra Trees由Pierre Geurts等人于2...

引言

在机器学习回归任务中,集成学习方法因其出色的预测性能和鲁棒性而备受青睐。随机森林(Random Forest)作为Bagging家族的明星算法广为人知,但你可能不知道,它的“亲戚”——极端随机树(Extremely Randomized Trees,简称Extra Trees) 在很多时候表现同样出色,甚至在某些场景下更胜一筹。

Extra Trees由Pierre Geurts等人于2006年提出,与随机森林高度相似,但核心区别在于它更“极端”的随机化策略。本文将带你从原理到实战,全面掌握如何使用Scikit-learn训练ExtraTrees回归模型。

一、Extra Trees核心原理

1.1 什么是Extra Trees?

Extra Trees是一种基于决策树的集成学习算法。它通过构建大量随机决策树,并将它们的预测结果进行平均(回归任务)来得到最终预测。与随机森林一样,它属于并行集成算法——所有决策树独立训练,互不依赖。

1.2 与随机森林的两大区别

Extra Trees与随机森林有两点核心区别:

对比项 随机森林 Extra Trees(极端随机树)
样本采样 使用Bootstrap采样(有放回抽样) 使用全部训练样本
分裂方式 随机选特征,找最优切分点 随机选特征,随机选切分点

简单来说,随机森林是“随机选特征,找最优切分” ,而 Extra Trees是“随机选特征 + 随机切分” 。这种“极端”的随机化策略带来了几个显著特点:

  • 训练速度更快:无需计算最优切分点,直接随机选择
  • 方差更低:更强的随机性降低了过拟合风险
  • 偏差稍高:单棵树的拟合精度略低于随机森林
  • 泛化能力更强:整体稳定性更好

用一句通俗的话概括:Extra Trees是一群“随机切菜”的决策树,单棵树不聪明,但靠数量多、随机性强,集体决策的结果既快又靠谱

二、Scikit-learn中的ExtraTreesRegressor

2.1 基本导入与初始化

from sklearn.ensemble import ExtraTreesRegressor
from sklearn.datasets import make_regression
from sklearn.model_selection import train_test_split
from sklearn.metrics import mean_squared_error, r2_score

2.2 核心参数详解

ExtraTreesRegressor的核心参数如下:

参数 默认值 说明
n_estimators 100 森林中决策树的数量
criterion ‘mse’ 分裂质量评估标准,可选’mse’(均方误差)或’mae’(平均绝对误差)
max_depth None 树的最大深度,None表示不限制
min_samples_split 2 内部节点再划分所需的最小样本数
min_samples_leaf 1 叶节点所需的最小样本数
max_features ‘auto’ 寻找最佳分裂时考虑的特征数量
bootstrap False 是否使用Bootstrap采样(Extra Trees默认不使用)
random_state None 随机种子,用于结果复现
n_jobs None 并行运行的作业数,-1表示使用所有处理器

2.3 关键参数调优建议

n_estimators(树的数量) :树越多,模型通常越稳定,但训练和预测时间也会增加。默认100是一个不错的起点。

max_features(特征采样数) :控制每棵树分裂时可用的特征数量。对于回归任务,默认'auto'等价于max_features = n_features

max_depth(树的最大深度) :限制树的深度可以防止过拟合。如果数据量较大或特征较多,建议设置一个合理的深度值。

min_samples_splitmin_samples_leaf :控制树的生长粒度,值越大,树越简单,泛化能力越强。

三、实战:训练ExtraTrees回归模型

下面我们通过一个完整的示例来演示如何使用ExtraTreesRegressor

3.1 生成示例数据

# 生成一个合成的回归数据集
X, y = make_regression(
    n_samples=1000,      # 样本数
    n_features=20,       # 特征数
    n_informative=15,    # 有效特征数
    noise=0.1,           # 噪声水平
    random_state=42
)

# 划分训练集和测试集
X_train, X_test, y_train, y_test = train_test_split(
    X, y, test_size=0.2, random_state=42
)

print(f"训练集大小: {X_train.shape[0]}")
print(f"测试集大小: {X_test.shape[0]}")

3.2 训练模型

# 创建ExtraTrees回归器
etr = ExtraTreesRegressor(
    n_estimators=100,
    max_depth=None,
    min_samples_split=2,
    min_samples_leaf=1,
    random_state=42,
    n_jobs=-1  # 使用所有CPU核心
)

# 训练模型
etr.fit(X_train, y_train)

3.3 模型评估

# 在测试集上进行预测
y_pred = etr.predict(X_test)

# 计算评估指标
mse = mean_squared_error(y_test, y_pred)
r2 = r2_score(y_test, y_pred)

print(f"均方误差 (MSE): {mse:.4f}")
print(f"决定系数 (R²): {r2:.4f}")

3.4 特征重要性分析

ExtraTreesRegressor提供了feature_importances_属性,可以评估每个特征对预测的贡献程度:

import numpy as np
import matplotlib.pyplot as plt

# 获取特征重要性
importances = etr.feature_importances_
indices = np.argsort(importances)[::-1]

# 可视化
plt.figure(figsize=(10, 6))
plt.title("Extra Trees 特征重要性")
plt.bar(range(len(importances)), importances[indices])
plt.xticks(range(len(importances)), indices)
plt.xlabel("特征索引")
plt.ylabel("重要性")
plt.show()

# 打印Top 5重要特征
for i in range(5):
    print(f"特征 {indices[i]}: {importances[indices[i]]:.4f}")

四、超参数调优

4.1 使用GridSearchCV进行网格搜索

from sklearn.model_selection import GridSearchCV

# 定义参数网格
param_grid = {
    'n_estimators': [50, 100, 200],
    'max_depth': [None, 10, 20, 30],
    'min_samples_split': [2, 5, 10],
    'min_samples_leaf': [1, 2, 4],
    'max_features': ['auto', 'sqrt', 'log2']
}

# 创建网格搜索对象
grid_search = GridSearchCV(
    estimator=ExtraTreesRegressor(random_state=42, n_jobs=-1),
    param_grid=param_grid,
    cv=5,                # 5折交叉验证
    scoring='r2',        # 评估指标
    n_jobs=-1,
    verbose=1
)

# 执行搜索
grid_search.fit(X_train, y_train)

# 输出最佳参数
print(f"最佳参数: {grid_search.best_params_}")
print(f"最佳交叉验证得分: {grid_search.best_score_:.4f}")

4.2 使用最佳参数重新训练

# 使用最佳参数训练模型
best_etr = grid_search.best_estimator_
y_pred_best = best_etr.predict(X_test)

print(f"优化后的R²: {r2_score(y_test, y_pred_best):.4f}")
【版权声明】本文为华为云社区用户原创内容,未经允许不得转载,如需转载请自行联系原作者进行授权。如果您发现本社区中有涉嫌抄袭的内容,欢迎发送邮件进行举报,并提供相关证据,一经查实,本社区将立刻删除涉嫌侵权内容,举报邮箱: cloudbbs@huaweicloud.com
  • 点赞
  • 收藏
  • 关注作者

评论(0

0/1000
抱歉,系统识别当前为高风险访问,暂不支持该操作

全部回复

上滑加载中

设置昵称

在此一键设置昵称,即可参与社区互动!

*长度不超过10个汉字或20个英文字符,设置后3个月内不可修改。

*长度不超过10个汉字或20个英文字符,设置后3个月内不可修改。