You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在读取CSV文件时复现scikit-learn OneHotEncoder混合数据类型错误?

scikit-learn OneHotEncoder混合数据类型问题解析

官方文档关于categories参数的注意事项

list类型:categories[i] 存储第i列预期的类别。
单个特征中传入的类别不能同时包含字符串和数值类型,且如果是数值类型的类别需要排序。

两种测试方式

方式1:硬编码DataFrame

from sklearn.compose import ColumnTransformer
from sklearn.preprocessing import OneHotEncoder
import pandas as pd

X = pd.DataFrame(
    {'city': ['London', 'London', 'Paris', 'NewYork'],
     'country': ['UK', 0.2, 'FR', 'US'],
     'user_rating': [4, 5, 4, 3]}
)
categorical_features = ['city', 'country']
one_hot = OneHotEncoder()
transformer = ColumnTransformer([("one_hot", one_hot, categorical_features)], remainder="passthrough")
transformed_X = transformer.fit_transform(X)

执行transformed_X = transformer.fit_transform(X)时会触发TypeError,符合文档中混合类型的限制。

方式2:读取CSV文件

读取代码

from sklearn.compose import ColumnTransformer
from sklearn.preprocessing import OneHotEncoder
import pandas as pd

X = pd.read_csv('data.csv', header=0)
categorical_features = ['city', 'country']
one_hot = OneHotEncoder()
transformer = ColumnTransformer([("one_hot", one_hot, categorical_features)], remainder="passthrough")
transformed_X = transformer.fit_transform(X)

对应的CSV内容(data.csv)

city,country,user_rating
London,UK,4
London,0.2,5
Paris,FR,4
NewYork,US,3

此方式未触发错误,scikit-learn将0.2视为字符串处理。

问题与解答

问题

能否在读取CSV文件时复现该混合数据类型错误?是否因pandas列类型推断导致无法复现?

解答

能复现,但需要手动干预列类型,避免pandas自动将混合格式的列推断为纯字符串类型。

原因在于pandas的CSV读取逻辑:当某一列同时包含字符串和数值格式的内容时,pandas会默认将整列推断为object类型(实际是字符串数组),所有值都会被转为字符串,此时scikit-learn接收的是纯字符串类别,自然不会触发混合类型的限制错误。

要复现错误,只需手动让列中同时存在字符串和数值类型,比如修改读取后的列数据:

from sklearn.compose import ColumnTransformer
from sklearn.preprocessing import OneHotEncoder
import pandas as pd

# 读取CSV时指定country列为object类型
X = pd.read_csv('data.csv', header=0, dtype={'country': object})
# 将第二行的country值从字符串"0.2"改为数值0.2
X.loc[1, 'country'] = 0.2

categorical_features = ['city', 'country']
one_hot = OneHotEncoder()
transformer = ColumnTransformer([("one_hot", one_hot, categorical_features)], remainder="passthrough")
transformed_X = transformer.fit_transform(X)

此时执行fit_transform会触发和硬编码方式相同的TypeError,因为country列现在确实混合了字符串和数值类型。

总结:无法复现的核心原因是pandas自动消除了列中的类型混合;只要手动恢复混合类型,就能触发文档中描述的错误限制。

内容的提问来源于stack exchange,提问作者akalanka

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.18 16:13:18