如何在ColumnTransformer中删除列?自定义Pipeline删除失效问题
解决ColumnTransformer中drop转换器无法删除指定列的问题
问题原因
ColumnTransformer的每个转换器独立作用于原始输入数据的列,而非按顺序处理前一个转换器的输出结果。你定义的drop_out转换器是尝试删除原始数据中的to_drop列,但这些列已经被select转换器通过passthrough保留到最终输出里,所以删除操作不会生效。
解决方案
方案1:直接在passthrough阶段排除需删除的列
这是最简洁的方式,无需额外的drop转换器,直接在select步骤中过滤掉要删除的列:
def custom_pipeline(to_drop: list = [], features_out: bool = False) -> Pipeline: # Add 'Message Length' attribute based on the 'Raw Message' column attrib_adder = AttributeAdder(attribs_in=['Raw Message'], attribs_out=['Message Length'], func=get_message_length) # 过滤passthrough列,排除要删除的项 passthrough_cols = [col for col in ['Attachments', 'URLs', 'IPs', 'Images', 'Message Length'] if col not in to_drop] # Define the column transformer preprocessor = ColumnTransformer(transformers=[ ('virus_scanned', enumerate_virus_scanned, ['X-Virus-Scanned']), ('priority', enumerate_priority, ['X-Priority']), ('encoding', enumerate_encoding, ['Encoding']), ('flags', enumerate_bool, ['Is HTML', 'Is JavaScript', 'Is CSS']), ('select', 'passthrough', passthrough_cols) ]) # Define pipeline pipe = Pipeline(steps=[ ('attrib_adder', attrib_adder), ('preprocessor', preprocessor), ('scaler', MinMaxScaler()) ]) # Get features out if features_out: # 直接从preprocessor获取输出列名,更准确可靠 features = list(preprocessor.get_feature_names_out()) # Return pipeline and features return pipe, features # Return pipeline return pipe
方案2:添加独立的列删除步骤(适用于更复杂的场景)
如果需要在预处理完成后再删除列,可以在Pipeline中加入一个自定义的列过滤步骤,注意要确保preprocessor输出的是带列名的DataFrame:
from sklearn.preprocessing import FunctionTransformer def drop_columns(X, cols_to_drop): return X.drop(cols_to_drop, axis=1) def custom_pipeline(to_drop: list = [], features_out: bool = False) -> Pipeline: # Add 'Message Length' attribute based on the 'Raw Message' column attrib_adder = AttributeAdder(attribs_in=['Raw Message'], attribs_out=['Message Length'], func=get_message_length) # Define the column transformer,设置输出为DataFrame并保留列名 preprocessor = ColumnTransformer(transformers=[ ('virus_scanned', enumerate_virus_scanned, ['X-Virus-Scanned']), ('priority', enumerate_priority, ['X-Priority']), ('encoding', enumerate_encoding, ['Encoding']), ('flags', enumerate_bool, ['Is HTML', 'Is JavaScript', 'Is CSS']), ('select', 'passthrough', ['Attachments', 'URLs', 'IPs', 'Images', 'Message Length']) ], verbose_feature_names_out=True).set_output(transform="pandas") # Define pipeline,添加列删除步骤 pipe = Pipeline(steps=[ ('attrib_adder', attrib_adder), ('preprocessor', preprocessor), ('column_dropper', FunctionTransformer(drop_columns, kw_args={'cols_to_drop': to_drop})), ('scaler', MinMaxScaler()) ]) # Get features out if features_out: # 先获取preprocessor的输出列,再排除要删除的列 pre_features = list(preprocessor.get_feature_names_out()) features = [col for col in pre_features if col not in to_drop] # Return pipeline and features return pipe, features # Return pipeline return pipe
内容的提问来源于stack exchange,提问作者Filip Szczybura
相关产品推荐
相关产品推荐

