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

ResNet101版U-Net训练灰度图报通道数不匹配错误求解

问题背景
  • 基于U-Net架构在德国沥青路面病害(GAPs)数据集训练裂缝分割模型,基于公开裂缝分割仓库代码做适配,自定义训练脚本为train_unet_GAPs.py,在Colab环境执行训练命令如下:
!python /content/drive/Othercomputers/My\ Laptop/crack_segmentation_khanhha/crack_segmentation-master/train_unet_GAPs.py -data_dir "/content/drive/Othercomputers/My Laptop/crack_segmentation_khanhha/crack_segmentation-master/GAPs/" -model_dir /content/drive/Othercomputers/My\ Laptop/crack_segmentation_khanhha/crack_segmentation-master/model/ -model_type resnet101
  • 训练启动阶段抛出通道数不匹配的RuntimeError,核心报错信息:

RuntimeError: Given groups=1, weight of size [64, 64, 1, 1], expected input[4, 1, 1080, 1920] to have 64 channels, but got 1 channels instead

  • 根因确认:GAPs数据集为单通道灰度图像,而代码中调用的torchvision内置ResNet101骨干默认接收3通道RGB输入,输入张量通道数和模型第一层卷积的权重通道数不匹配。
解决方案

提供两种可直接落地的修改方式,优先选择第一种:

方案1:数据预处理层将单通道灰度图转为3通道(改动最小,推荐)

  • 不需要修改模型结构,可直接复用ImageNet预训练的ResNet101权重,训练收敛更稳定,分割效果不受影响
  • 找到数据集加载模块的图像预处理transform配置,在ToTensor()操作之后添加通道复制逻辑,将形状为[1, H, W]的单通道张量复制3次,拼接为[3, H, W]的3通道张量,和ResNet默认输入格式对齐
  • 若使用torchvision实现transform,直接在Compose列表中添加如下代码即可:
from torchvision import transforms
transform = transforms.Compose([
    transforms.ToTensor(),
    # 新增下面这行
    transforms.Lambda(lambda x: x.repeat(3, 1, 1)),
    # 其他原有预处理逻辑,比如Resize、Normalize等
])
  • 修改完成后可先打印单个训练batch的输入张量形状,确认维度为[batch_size, 3, 1080, 1920]后再启动训练。

方案2:修改ResNet骨干第一层卷积适配单通道输入

  • 不需要改动数据处理逻辑,直接调整模型输入层适配灰度图输入
  • 打开仓库中unet/unet_transfer.py文件,找到ResNet101模型加载完成的位置,添加如下代码替换原有第一层卷积:
import torch
import torch.nn as nn
from torchvision.models import resnet101

# 原有加载预训练模型的代码,示例
model = resnet101(pretrained=True)

# 新增以下代码:替换第一层卷积为单通道输入
old_conv1 = model.conv1
model.conv1 = nn.Conv2d(
    in_channels=1, 
    out_channels=old_conv1.out_channels,
    kernel_size=old_conv1.kernel_size,
    stride=old_conv1.stride,
    padding=old_conv1.padding,
    bias=old_conv1.bias
)
# 复用预训练权重:将原3通道卷积的权重沿通道维度取均值,赋值给新的单通道卷积,避免随机初始化导致训练不稳定
with torch.no_grad():
    model.conv1.weight.copy_(old_conv1.weight.mean(dim=1, keepdim=True))
  • 该修改会解决后续所有层的通道数连锁不匹配问题,无需再调整U-Net其他模块的参数。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.30 11:15:30