当前位置:首页>python>GBDT Python实战:5行代码搞定预测模型

GBDT Python实战:5行代码搞定预测模型

  • 2026-10-11 06:33:17
GBDT Python实战:5行代码搞定预测模型

本文代码基于 LightGBM,目前工业界最常用的 GBDT 库。


一、环境准备

# 安装 LightGBMpip install lightgbm scikit-learn pandas numpy matplotlib

二、完整代码:房价预测

import lightgbm as lgbfrom sklearn.datasets import fetch_california_housingfrom sklearn.model_selection import train_test_splitfrom sklearn.metrics import mean_squared_errorimport numpy as np# 1. 加载数据(加州房价数据集)data = fetch_california_housing()X, y = data.data, data.target# 2. 划分训练集和测试集X_train, X_test, y_train, y_test = train_test_split(    X, y, test_size=0.2, random_state=42)# 3. 创建 LightGBM 数据集train_data = lgb.Dataset(X_train, label=y_train)valid_data = lgb.Dataset(X_test, label=y_test, reference=train_data)# 4. 设置参数params = {    'objective': 'regression',      # 回归任务    'metric': 'rmse',               # 评估指标:均方根误差    'boosting_type': 'gbdt',        # 梯度提升树    'num_leaves': 31,               # 每棵树的叶子数(控制复杂度)    'learning_rate': 0.05,          # 学习率:每次修正的幅度    'feature_fraction': 0.9,        # 每次随机选 90% 特征    'bagging_fraction': 0.8,        # 每次随机选 80% 样本    'bagging_freq': 5,              # 每 5 轮迭代做一次 bagging    'verbose': -1                   # 不打印训练日志}# 5. 训练模型model = lgb.train(    params,    train_data,    num_boost_round=1000,           # 最多 1000 棵树    valid_sets=[valid_data],    callbacks=[lgb.early_stopping(50)]  # 50 轮不提升就停止)# 6. 预测y_pred = model.predict(X_test, num_iteration=model.best_iteration)# 7. 评估rmse = np.sqrt(mean_squared_error(y_test, y_pred))print(f"RMSE: {rmse:.4f}")# 8. 特征重要性importance = pd.DataFrame({    'feature': data.feature_names,    'importance': model.feature_importance()}).sort_values('importance', ascending=False)print("\n特征重要性 Top 5:")print(importance.head())

三、代码讲解

核心参数解析

参数
作用
调参建议
num_leaves
控制单棵树复杂度
默认 31,数据多可加大
learning_rate
学习率
0.01~0.1,配合更多树
feature_fraction
特征采样比例
防过拟合,默认 0.8~1
bagging_fraction
样本采样比例
防过拟合,默认 0.8
num_boost_round
树的数量
早停会自动选最优

早停机制(Early Stopping)

callbacks=[lgb.early_stopping(50)]

作用:如果连续 50 轮验证集指标没提升,自动停止训练。

好处:不用猜要多少棵树,模型自己决定。


四、分类任务代码

把上面代码改 3 处即可:

# 1. 加载分类数据(如鸢尾花)from sklearn.datasets import load_irisfrom sklearn.metrics import accuracy_scoredata = load_iris()X, y = data.data, data.target# 2. 改参数params = {    'objective': 'multiclass',      # 多分类    'num_class': 3,                 # 类别数    'metric': 'multi_logloss',    # ... 其他参数不变}# 3. 改评估y_pred = model.predict(X_test, num_iteration=model.best_iteration)y_pred_class = np.argmax(y_pred, axis=1)  # 取概率最大的类别accuracy = accuracy_score(y_test, y_pred_class)print(f"Accuracy: {accuracy:.4f}")

五、保存和加载模型

# 保存model.save_model('gbdt_model.txt')# 加载model = lgb.Booster(model_file='gbdt_model.txt')# 预测y_pred = model.predict(X_new)

六、可视化特征重要性

import matplotlib.pyplot as pltlgb.plot_importance(model, max_num_features=10, figsize=(10, 6))plt.title("特征重要性 Top 10")plt.tight_layout()plt.savefig('feature_importance.png', dpi=150)plt.show()

让数据分析变得简单有趣 🐍

#机器学习 #梯度提升树 #数据分析 #互联网 #AI

夏天雨水多 ,大家多注意安全!!!

最新文章

随机文章