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

使用Transformers训练模型时遇通道维度推断错误,输入为torch.Size([1,3,224,224])

解决Transformers训练中「无法推断通道维度格式」错误

错误核心原因

Hugging Face Feature Extractor需要明确识别图像通道维度的位置(通道在前[C,H,W]或在后[H,W,C])。你的输入张量为torch.Size([1,3,224,224])(批量格式),但预处理函数中对图像的处理逻辑不符合Feature Extractor的预期,导致无法自动识别通道维度。

具体解决方法

1. 修正预处理函数的张量处理逻辑

如果图像是张量格式,优先转换为PIL图像(通道位置明确),同时避免在单样本预处理阶段使用return_tensors='pt'(Trainer会自动组装批量张量):

def preprocess(batch):
    # 将张量转换为PIL图像(适配Feature Extractor的默认输入格式)
    images = []
    for img in batch['image']:
        if isinstance(img, torch.Tensor):
            # 通道在前转通道在后:[C,H,W] → [H,W,C]
            img = img.permute(1, 2, 0).numpy()
            img = Image.fromarray(img.astype('uint8'))
        images.append(img)
    
    inputs = feature_extractor(images)
    inputs['labels'] = batch['labels']
    return inputs

2. 直接指定通道维度格式

若坚持使用张量输入,可在调用Feature Extractor时明确指定data_format参数:

def preprocess(batch):
    inputs = feature_extractor(
        batch['image'],
        data_format='channels_first',  # 匹配你的[C,H,W]张量格式
        return_tensors='pt'
    )
    inputs['labels'] = batch['labels']
    return inputs

3. 清理数据集的图像维度

确保数据集返回的单张图像维度为[3,224,224]或[224,224,3],若存在多余的batch维度,需提前去除:

def preprocess(batch):
    # 去除单张图像的冗余batch维度
    images = [img.squeeze(0) if len(img.shape) == 4 else img for img in batch['image']]
    
    inputs = feature_extractor(
        images,
        data_format='channels_first',
        return_tensors='pt'
    )
    inputs['labels'] = batch['labels']
    return inputs

关键注意事项

  • 避免在单样本预处理阶段设置return_tensors='pt',Trainer会自动完成批量张量的组装。
  • 确保所有图像的通道数为3(RGB)或1(灰度),非标准通道数也会触发该错误。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 11:24:57