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

使用Skorch训练ResNet分类模型时遇TypeError:含Categorical数据

修复PyTorch default_collate无法处理Categorical类型的问题

方案1:将Categorical特征转换为整数编码

直接把pandas Categorical列转成整数编码,PyTorch的Embedding层可直接接收这类数据。注意要保证训练集和测试集编码一致,避免数据泄露:

# 提取训练集的类别映射
cat_categories = X_train['分类特征列名'].cat.categories
# 训练集转整数编码
X_train['分类特征列名'] = X_train['分类特征列名'].cat.codes
# 测试集复用训练集的映射编码
X_test['分类特征列名'] = X_test['分类特征列名'].astype('category').cat.set_categories(cat_categories).cat.codes

方案2:自定义collate函数

给Skorch模型传入自定义的collate_fn,在数据加载阶段处理Categorical类型:

import torch

def custom_collate(batch):
    X_batch, y_batch = zip(*batch)
    processed_X = []
    for x in X_batch:
        # 将Categorical列转为长整型张量
        x['分类特征列名'] = torch.tensor(x['分类特征列名'].cat.codes, dtype=torch.long)
        processed_X.append(x)
    # 调用默认collate处理其余数据
    return torch.utils.data.default_collate(processed_X), torch.utils.data.default_collate(y_batch)

# 初始化Skorch模型时传入自定义collate函数
from skorch import NeuralNetClassifier
model = NeuralNetClassifier(
    YourResNetModel,
    collate_fn=custom_collate,
    # 其他参数如optimizer、lr等
)
model.fit(X_train, y_train)

方案3:用预处理管道统一处理

结合sklearn的ColumnTransformer和OrdinalEncoder,提前处理分类特征,再传入模型训练:

from sklearn.compose import ColumnTransformer
from sklearn.preprocessing import OrdinalEncoder

# 定义预处理规则:仅处理分类列,其余列保持原样
preprocessor = ColumnTransformer(
    transformers=[
        ('cat_encoder', OrdinalEncoder(), ['分类特征列名'])
    ],
    remainder='passthrough'
)

# 训练集拟合并转换,测试集直接转换
X_train_processed = preprocessor.fit_transform(X_train)
X_test_processed = preprocessor.transform(X_test)

# 传入Skorch模型训练
model = NeuralNetClassifier(YourResNetModel)
model.fit(X_train_processed, y_train)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 10:15:40