为何需用transform处理XGBoost特征选择后的测试集?
XGBoost特征选择中测试集特征处理的疑问与解答
背景
我近期在开发XGBoost模型,采用基于特征重要性分数的特征选择方法,流程如下:
- 先基于全部特征构建模型,为每个特征赋予重要性分数;
- 迭代构建模型:依次基于Top1、Top2、Top3……的重要特征建模。
迭代建模的代码如下:
# load data dataset = loadtxt('pima-indians-diabetes.csv', delimiter=",") # split data into X and y X = dataset[:,0:8] Y = dataset[:,8] # split data into train and test sets X_train, X_test, y_train, y_test = train_test_split(X, Y, test_size=0.33, random_state=7) # fit model on all training data model = MyXGBClassifier() model.fit(X_train, y_train) # make predictions for test data and evaluate predictions = model.predict(X_test) accuracy = accuracy_score(y_test, predictions) print("Accuracy: %.2f%%" % (accuracy * 100.0)) # Fit model using each importance as a threshold thresholds = sort(model.feature_importances_, reverse=True) for thresh in thresholds: # select features using threshold selection = SelectFromModel(model, threshold=thresh, prefit=True) select_X_train = selection.transform(X_train) # train model selection_model = XGBClassifier() selection_model.fit(select_X_train, y_train) # eval model select_X_test = selection.transform(X_test) predictions = selection_model.predict(select_X_test) accuracy = accuracy_score(y_test, predictions) print("Thresh=%.3f, n=%d, Accuracy: %.2f%%" % (thresh, select_X_train.shape[1], accuracy*100.0))
疑问
为什么必须通过select_X_test = selection.transform(X_test)这行代码选择测试集特征?直接从model.feature_importances_中选取与select_X_train数量相同的TopN重要特征作为测试集子集来预测时,模型表现极差(几乎所有样本都被标记为正例),但使用transform方法时模型表现良好(约70%的精确率与召回率)?
解答
核心问题出在你手动选TopN特征时,大概率搞错了特征的对应原始列索引。
SelectFromModel的transform方法不是简单按重要性排序取前N个,它的逻辑是:
- 基于原模型输出的特征重要性,筛选出重要性≥设定阈值的所有特征,记录这些特征在原始数据中的列索引;
- 从训练/测试集中提取这些列索引对应的特征,确保训练集和测试集用的是完全一致的特征列。
而你手动操作时,很可能只把特征重要性分数排序后取前N个,然后错误地把“分数的排名”当成了“原始特征的列位置”——比如原始特征列0的重要性排第3,列3的重要性排第1,你取Top1时可能误拿了列0,而不是真正重要的列3。这种特征不匹配会导致测试集输入的完全是模型没见过的无效特征,模型只能输出数据中占比最高的类别(比如正例占比高就全标正例)。
另外,当多个特征重要性分数相同时,SelectFromModel会保留所有符合阈值条件的特征,这时候实际保留的特征数量可能和你预期的TopN不完全一致,手动选择会进一步偏离正确的特征集合。
总结两点:
selection.transform是基于原始特征的列索引筛选,保证训练、测试集特征完全对齐;- 手动选TopN容易混淆“重要性排名”和“原始列位置”,导致测试集特征错误,模型彻底失效。
内容的提问来源于stack exchange,提问作者Or Pickholz
相关产品推荐
相关产品推荐

