PyTorch搭建口罩识别CNN报输入通道数不匹配RuntimeError
问题背景
- 开发目标:编写CNN模型,实现人像照片中人物是否佩戴口罩、佩戴口罩类型的识别分类任务
- 数据集配置:训练集包含约1500张图片,各类别样本量均衡;测试集包含约450张图片
- 触发场景:完成训练数据加载器、测试逻辑的代码编写后运行程序,抛出运行时错误
报错详情
RuntimeError: Given groups=1, weight of size [6, 3, 3, 3], expected input[4, 224, 3, 224] to have 3 channels, but got 224 channels instead
错误根因
报错本质是卷积层要求的输入张量维度顺序,和实际传入的张量维度顺序不匹配:
- PyTorch的
Conv2d层强制要求输入张量遵循(batch_size, 通道数, 图像高度, 图像宽度)的(N,C,H,W)格式 - 你定义的第一层卷积权重形状为
[6, 3, 3, 3],对应输出6个特征通道、输入3个通道(匹配RGB三通道图像)、卷积核尺寸3*3,权重本身定义无问题 - 实际传入模型的输入张量形状为
[4, 224, 3, 224]:其中batch size为4,第二维的224实际是图像高度,第三维的3才是RGB通道数,第四维的224是图像宽度。也就是你传入的是(N,H,W,C)格式的张量,框架误将第二维的224识别为输入通道数,才会抛出"期望3通道、实际得到224通道"的错误。
修复方案
根据常见的触发场景,对应修复方式如下:
- 场景1:手动读取图片后未调整维度顺序
OpenCV、PIL读取的图片默认存储格式为(H,W,C),传入模型前需要手动调换维度顺序,参考代码:import torch # 假设img为读取得到的numpy数组,形状为(H,W,C) img_tensor = torch.from_numpy(img).permute(2, 0, 1) # 转换为(C,H,W)格式 # 组batch后最终形状自动符合(N,C,H,W)要求 - 场景2:数据变换流程配置错误
若使用torchvision.transforms做数据预处理,确保流程中包含ToTensor()操作,该算子会自动将PIL图像、(H,W,C)格式的numpy数组转换为(C,H,W)格式、取值范围在[0,1]的张量。注意不要在ToTensor()执行后手动调换维度打乱顺序。 - 场景3:自定义DataLoader的批处理逻辑错误
检查自定义collate_fn的批拼接逻辑,确保拼接后的batch张量维度顺序为(N,C,H,W),不要将通道维放到高度、宽度维度之后。
内容的提问来源于stack exchange,提问作者christ
相关产品推荐
相关产品推荐

