如何验证scikit-learn中cross_val_score的数据拆分在StandardScaler之前执行?
问题描述
我正在用scikit-learn构建决策树模型,要求先拆分数据再用StandardScaler()做缩放,同时想用cross_val_score()方法。我先用make_column_transformer()结合OneHotEncoding()对部分分类数据编码,代码如下:
transformer = sklearn.compose.make_column_transformer( (sklearn.preprocessing.OneHotEncoder(handle_unknown='ignore'), ['SoilDrainage', 'Geology', 'LU2016']), remainder='passthrough')
接着实例化模型和缩放器:
model = sklearn.tree.DecisionTreeClassifier() scalar = sklearn.preprocessing.StandardScaler()
把它们加入管道:
pipe = sklearn.pipeline.make_pipeline(transformer, scalar, model)
最后把管道传入cross_val_score():
sklearn.model_selection.cross_val_score(pipe, X, y, cv=5, scoring='accuracy').mean()
代码执行没报错,但因为拆分是在cross_val_score()内部完成的,我不确定怎么验证缩放是在数据拆分之后执行的。
验证方法
1. 手动模拟单轮交叉验证流程
手动拆分一次训练集和测试集,通过管道的分步操作对比缩放后的统计量,确认缩放仅基于训练集计算:
from sklearn.model_selection import train_test_split # 手动拆分数据 X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42) # 拟合管道的预处理步骤(编码+缩放) pipe[:-1].fit(X_train) # 计算训练集缩放后的特征统计量 train_scaled = pipe[:-1].transform(X_train) print("训练集缩放后前5列均值:", train_scaled[:, :5].mean(axis=0)) print("训练集缩放后前5列标准差:", train_scaled[:, :5].std(axis=0)) # 计算测试集缩放后的特征统计量 test_scaled = pipe[:-1].transform(X_test) print("测试集缩放后前5列均值:", test_scaled[:, :5].mean(axis=0)) print("测试集缩放后前5列标准差:", test_scaled[:, :5].std(axis=0))
如果训练集缩放后的均值接近0、标准差接近1,而测试集的统计量不会严格等于0和1,说明缩放是基于训练集计算的(即拆分后执行),符合无数据泄露的要求。
2. 查看管道的拟合参数
拟合管道后,直接查看StandardScaler的拟合属性,确认这些参数仅来自训练集:
# 用训练集拟合整个管道 pipe.fit(X_train, y_train) # 查看缩放器基于训练集计算的均值和标准差 print("缩放器训练集均值:", pipe.named_steps['standardscaler'].mean_) print("缩放器训练集标准差:", pipe.named_steps['standardscaler'].scale_)
这些属性是缩放器在训练集上拟合得到的,测试集处理时只会复用这些参数,不会重新计算,这也能证明缩放是在拆分后执行的。
3. 理解scikit-learn的内置逻辑
scikit-learn的Pipeline和cross_val_score从设计上就遵循训练阶段仅用训练集拟合预处理步骤的规则:
cross_val_score会自动拆分数据为k份,每轮用k-1份做训练集,1份做测试集- 对每轮训练集,管道会依次执行
transformer.fit_transform()、scalar.fit_transform(),最后训练模型 - 对测试集,管道仅执行
transformer.transform()、scalar.transform(),不会重新拟合预处理步骤,避免数据泄露
内容的提问来源于stack exchange,提问作者Zac Khan
相关产品推荐
相关产品推荐

