如何使用XGB高效处理150+分类变量以计算特征重要性
分类变量场景下特征重要性计算流程优化方案
编码阶段优化(解决维度爆炸、运行卡顿问题)
- 放弃全量
pd.get_dummies编码方案:你当前的流程不仅存在先编码再拆分数据集导致的数据泄露问题,150+分类变量独热编码后维度指数级上升也是内核崩溃的核心原因。XGBoost 1.3版本之后原生支持无序分类变量输入,不需要做独热编码,也不会触发哑变量陷阱,运行效率会提升数倍,仅需提前把分类列转为category类型即可:
模型定义时新增参数# 仅把object类型列标记为分类类型,不需要额外编码 cat_cols = df.select_dtypes(include='object').columns df[cat_cols] = df[cat_cols].astype('category')enable_categorical=True即可直接适配分类特征。 - 若必须使用独热编码,可改用
OneHotEncoder配合ColumnTransformer放在Pipeline中处理,设置参数过滤低频分类取值、避免哑变量陷阱,大幅降低生成的特征数量:from sklearn.compose import ColumnTransformer from sklearn.preprocessing import OneHotEncoder, StandardScaler num_cols = df.select_dtypes(include='number').columns cat_cols = df.select_dtypes(include='object').columns preprocessing_pipeline = ColumnTransformer( transformers=[ ('num', StandardScaler(), num_cols), # drop='first'避免哑变量陷阱,min_frequency过滤出现占比低于1%的取值,合并为同一类 ('cat', OneHotEncoder(drop='first', min_frequency=0.01, sparse_output=False), cat_cols) ])
流程逻辑修正(避免数据泄露、提升结果准确率)
- 调整数据集拆分顺序:必须先拆分训练集和测试集,再拟合预处理管道和模型,否则会将测试集的数据信息泄露到训练过程,导致特征重要性结果失真、模型泛化能力下降。
- 简化Pipeline配置:使用XGB原生分类特征支持的方案可直接删除编码相关的预处理步骤,进一步提升运行效率:
# 先拆分数据集,再做后续处理 X_train, X_test, y_train, y_test = train_test_split(df, y, test_size=0.2, random_state=42, stratify=y_closed) # 定义支持分类变量的XGB模型 model = XGBClassifier(enable_categorical=True, use_label_encoder=False, eval_metric='logloss') pipeline = Pipeline([ ("classifier", model) ]) trained_pipeline = pipeline.fit(X_train, y_train)
特征重要性绘图代码修正
你原有的特征名获取逻辑存在错误,修正后代码如下:
%matplotlib inline import matplotlib.pyplot as plt import numpy as np N_FEATURES = 10 importances = trained_pipeline.named_steps['classifier'].feature_importances_ # 原生分类特征方案可直接使用原始列名,可读性极强 feature_names = df.columns indices = np.argsort(importances)[-N_FEATURES:] plt.figure(figsize=(15, 8)) plt.title('Feature Importances') plt.barh(range(N_FEATURES), importances[indices], align='center') plt.yticks(range(N_FEATURES), feature_names[indices]) plt.xlabel('Importance Score') plt.show()
如果使用了OneHotEncoder方案,可通过preprocessing_pipeline.get_feature_names_out()获取编码后的全量特征名。
内容的提问来源于stack exchange,提问作者Tara-S1983
相关产品推荐
相关产品推荐

