Python使用XGBoost基于其他列预测目标列的问题排查
答复
首先给出明确结论:20列的数据集完全可以用前19列作为特征,通过XGBoost预测第20列的目标值,你遇到的预测值全为常数、无法绘制树结构的问题,均来自代码错误、参数设置不当,和数据集列数没有关系。
预测值全为常数的原因与修复方案
- 目标函数参数错误:你在
XGBRegressor中使用的objective='reg:linear'是XGBoost早已废弃的回归目标参数,新版本中回归任务的正确目标函数为reg:squarederror,调用废弃参数会导致模型训练逻辑异常,直接输出常数结果。 - 正则化强度过高:设置的
alpha=10为L1正则项权重,数值过大会将所有特征的分裂权重惩罚至0,模型退化为仅预测目标列均值的常数模型。 - 训练迭代轮数不足:
n_estimators=10搭配learning_rate=0.1的组合下,模型还未学习到特征和标签的关联就停止训练,极易输出常数预测。 - 代码缺失依赖导入:计算RMSE时调用了
np.sqrt,但代码中没有导入numpy库,运行时会直接报错,需要补上import numpy as np。
无法绘制树结构的原因与修复方案
- 模型被覆盖:你在可视化步骤前重新调用
xgb.train训练了一个新模型,覆盖了之前在训练集上拟合的XGBRegressor模型,且新模型训练时num_boost_round=10,你指定绘制num_trees=5时,如果因为正则过强没有生成足够的有效树,就会触发绘图错误。 - 绘图配置顺序错误:你将画布尺寸
plt.rcParams['figure.figsize'] = [50, 10]的设置放在了plot_tree调用之后,配置不会对已经生成的绘图对象生效,需要把尺寸设置放在绘图代码之前。 - 依赖缺失:XGBoost的树绘图依赖graphviz库,如果修正代码后仍无法绘图,需要确认已正确安装graphviz(不仅要安装python包,还要安装对应系统的graphviz二进制程序)。
修正后的参考代码
import xgboost as xgb import numpy as np import matplotlib.pyplot as plt from sklearn.metrics import mean_squared_error from sklearn.model_selection import train_test_split # 分离特征与目标列 X, y = f.iloc[:,:-1], f.iloc[:,-1] # 拆分训练、测试集 X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.33, random_state=123) # 初始化回归模型,修正错误参数 xg_reg = xgb.XGBRegressor( objective='reg:squarederror', colsample_bytree=0.3, learning_rate=0.1, max_depth=5, alpha=1, # 调低L1正则强度 n_estimators=100 # 适当增加迭代轮数 ) # 训练、预测 xg_reg.fit(X_train, y_train) preds = xg_reg.predict(X_test) # 计算RMSE rmse = np.sqrt(mean_squared_error(y_test, preds)) print(f"RMSE: {rmse:f}") # k折交叉验证 data_dmatrix = xgb.DMatrix(data=X, label=y) cv_params = { "objective":"reg:squarederror", "colsample_bytree": 0.3, "learning_rate": 0.1, "max_depth": 10, "alpha": 1 } cv_results = xgb.cv( dtrain=data_dmatrix, params=cv_params, nfold=3, num_boost_round=50, early_stopping_rounds=10, metrics="rmse", as_pandas=True, seed=123 ) print(cv_results["test-rmse-mean"].tail(1)) # 绘制树结构:先设置画布尺寸,再绘图,树索引从0开始避免越界 plt.rcParams['figure.figsize'] = [50, 10] xgb.plot_tree(xg_reg, num_trees=0) plt.show() # 绘制特征重要性 plt.rcParams['figure.figsize'] = [5, 5] xgb.plot_importance(xg_reg) plt.show()
额外排查项
如果修正代码后预测结果仍为常数,需要检查数据集本身:
- 确认目标列(第20列)本身不是几乎全为同一个值的常数列
- 确认前19列特征没有大量缺失值、全为常数列的情况,存在可供模型学习的有效信息
- 确认特征和目标列之间存在可学习的关联,不是完全独立的随机数据
内容的提问来源于stack exchange,提问作者SHAH
相关产品推荐
相关产品推荐

