OneHotEncoder指定字符串类别触发Shape mismatch错误,求排查方案
问题解答:OneHotEncoder categories参数报错处理
问题背景
数据集结构如下:
| type_of_bicycle_lane | month | other variables |
|---|---|---|
| 'on the street' | 1 | ... |
| 'fully seperated' | 4 | ... |
| 'fully seperated' | 8 | ... |
| 'fully seperated' | 1 | ... |
| 'other_1' | 12 | ... |
| 'other_2' | 8 | ... |
尝试通过scikit-learn的ColumnTransformer搭配OneHotEncoder,仅对type_of_bicycle_lane列的['on the street', 'fully seperated']生成独热编码,编写代码如下:
full_pipeline = ColumnTransformer([ ("bicycle_lane", OneHotEncoder(categories = ['on the street', 'fully seperated']), ["type_of_bicycle_lane"]), ])
但触发错误:Shape mismatch: if categories is an array, it has to be of shape (n_features,)
此前处理month列时使用("month", OneHotEncoder(categories = [range(1,13)]), ["month"])运行正常,不清楚当前错误原因。
错误原因与修正
报错核心是OneHotEncoder的categories参数要求为嵌套列表结构,每个子列表对应一个输入特征的类别集合。
处理month列时用的[range(1,13)]是外层套列表的结构,符合参数形状要求(对应1个特征的类别);而处理自行车道类型时直接传入一维列表['on the street', 'fully seperated'],不符合(n_features,)的形状要求(此处n_features=1,需在外层再包一层列表)。
修正后的代码:
full_pipeline = ColumnTransformer([ ("bicycle_lane", OneHotEncoder(categories = [['on the street', 'fully seperated']]), ["type_of_bicycle_lane"]), ])
如果需要将other_1、other_2这类未指定的取值归为“其他”类别,可添加handle_unknown='ignore'参数,避免编码器因未知类别报错:
full_pipeline = ColumnTransformer([ ("bicycle_lane", OneHotEncoder(categories = [['on the street', 'fully seperated']], handle_unknown='ignore'), ["type_of_bicycle_lane"]), ])
内容的提问来源于stack exchange,提问作者silke
相关产品推荐
相关产品推荐

