XGBoost 入门指南:Python 梯度提升实战

0 阅读2分钟

做表格数据(结构化数据)的机器学习竞赛,XGBoost 常年是冠军选手。它基于梯度提升树(GBDT),通过多棵决策树叠加修正误差,在分类、回归任务上又快又准。本文用 Python 跑通一个完整的 XGBoost 训练流程。

一、环境准备

pip install xgboost scikit-learn pandas

要求 Python 3.8+。CPU 即可跑,大数据集建议有 GPU 版本(pip install xgboost 已含 GPU 支持,需编译)。

二、最小可运行示例

用经典鸢尾花数据集做分类:

import xgboost as xgb
from sklearn.datasets import load_iris
from sklearn.model_selection import train_test_split
from sklearn.metrics import accuracy_score

X, y = load_iris(return_X_y=True)
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42)

model = xgb.XGBClassifier(n_estimators=100, max_depth=3, learning_rate=0.1)
model.fit(X_train, y_train)

pred = model.predict(X_test)
print("准确率:", accuracy_score(y_test, pred))

几行代码就完成训练与预测,输出接近 1.0 的准确率。

三、核心概念

  • DMatrix:XGBoost 的高效数据格式,比 NumPy/Pandas 更快,支持权重、缺失值。
  • XGBClassifier / XGBRegressor:分类/回归封装,API 与 sklearn 一致(.fit/.predict)。
  • 关键超参n_estimators(树数量)、max_depth(树深)、learning_rate(步长)、subsample(样本采样)。

四、进阶用法:原生接口与早停

early_stopping 防止过拟合:

import xgboost as xgb
from sklearn.datasets import load_breast_cancer
from sklearn.model_selection import train_test_split

X, y = load_breast_cancer(return_X_y=True)
X_tr, X_val, y_tr, y_val = train_test_split(X, y, test_size=0.2, random_state=42)

dtrain = xgb.DMatrix(X_tr, label=y_tr)
dval = xgb.DMatrix(X_val, label=y_val)

params = {"objective": "binary:logistic", "max_depth": 4, "eta": 0.1}
bst = xgb.train(
    params, dtrain, num_boost_round=200,
    evals=[(dtrain, "train"), (dval, "val")],
    early_stopping_rounds=20, verbose_eval=False
)
print("最佳轮次:", bst.best_iteration)

验证集连续 20 轮不提升就停,省时又防过拟合。

五、实战场景:Kaggle 式房价回归

from sklearn.datasets import fetch_california_housing
from sklearn.metrics import mean_squared_error
import numpy as np

data = fetch_california_housing()
X_tr, X_te, y_tr, y_te = train_test_split(data.data, data.target, test_size=0.2, random_state=42)

reg = xgb.XGBRegressor(n_estimators=300, max_depth=5, learning_rate=0.05)
reg.fit(X_tr, y_tr)
rmse = np.sqrt(mean_squared_error(y_te, reg.predict(X_te)))
print("RMSE:", rmse)

objective="reg:squarederror"(回归默认)即可直接预测连续值。

六、常见坑 / 报错

1. ValueError: feature_names mismatch 训练与预测特征列不一致。统一用 DataFrame 并保列顺序,或训练/预测都用 DMatrix 且列对齐。

2. 过拟合(训练 100% / 验证差) 调小 max_depth、降低 learning_rate 并增大 n_estimators、加 early_stopping_rounds、用 subsample/colsample_bytree 做随机性。

3. 类别特征未编码 XGBoost 不吃字符串,需用 OneHotEncoderOrdinalEncoder 先编码;或设 enable_categorical=True(新版本支持直接传类别列)。

七、总结与下一步

XGBoost 是表格数据建模的“默认强者”:API 友好、速度快、可调性强。下一步建议:

  • GridSearchCV / Optuna 做超参搜索;
  • XGBClassifier(..., tree_method="gpu_hist") 开 GPU 加速;
  • 对比 LightGBM / CatBoost,三者常一起 baseline。

掌握 XGBoost,你就有了打榜和落地结构化 ML 的硬实力。