本文代码基于 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 | | |
learning_rate | | |
feature_fraction | | |
bagging_fraction | | |
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
夏天雨水多 ,大家多注意安全!!!