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

如何为ResNet50添加额外通道?解决通道不匹配报错

为ResNetV2(适配TransUNet)添加4通道输入并解决通道不匹配报错

问题原因

报错RuntimeError: Given groups=1, weight of size [64, 3, 7, 7], expected input[2, 4, 256, 256] to have 3 channels, but got 4 channels instead的核心原因:

  • 你修改了根卷积层StdConv2d(4, width, ...)的输入通道为4,但加载预训练权重时,原权重的输入通道是3,直接复制后权重维度仍为[64,3,7,7],和输入的4通道不匹配。
  • 代码中使用了conv4x4函数但未定义,会引发额外报错。
  • 自定义的PreActBottleneck新增了conv4和gn4,但load_from方法直接复用原权重的conv3参数,逻辑错误。

解决方案

1. 补全conv4x4函数

在conv3x3函数下方添加:

def conv4x4(cin, cout, stride=1, groups=1, bias=False):
    return StdConv2d(cin, cout, kernel_size=4, stride=stride,
                     padding=1, bias=bias, groups=groups)

2. 修改根卷积层的权重加载逻辑

将原3通道的预训练权重扩展为4通道,可复制原3通道权重作为第4通道,或随机初始化第4通道。

3. 修正PreActBottleneck的load_from方法

新增的conv4和gn4没有对应预训练权重,改为随机初始化,避免加载错误。

修改后的完整代码

import math
from os.path import join as pjoin
from collections import OrderedDict
import torch
import torch.nn as nn
import torch.nn.functional as F


def np2th(weights, conv=False):
    """Possibly convert HWIO to OIHW."""
    if conv:
        weights = weights.transpose([3, 2, 0, 1])
    return torch.from_numpy(weights)


class StdConv2d(nn.Conv2d):
    def forward(self, x):
        w = self.weight
        v, m = torch.var_mean(w, dim=[1, 2, 3], keepdim=True, unbiased=False)
        w = (w - m) / torch.sqrt(v + 1e-5)
        return F.conv2d(x, w, self.bias, self.stride, self.padding,
                        self.dilation, self.groups)


def conv3x3(cin, cout, stride=1, groups=1, bias=False):
    return StdConv2d(cin, cout, kernel_size=3, stride=stride,
                     padding=1, bias=bias, groups=groups)


def conv4x4(cin, cout, stride=1, groups=1, bias=False):
    return StdConv2d(cin, cout, kernel_size=4, stride=stride,
                     padding=1, bias=bias, groups=groups)


def conv1x1(cin, cout, stride=1, bias=False):
    return StdConv2d(cin, cout, kernel_size=1, stride=stride,
                     padding=0, bias=bias)


class PreActBottleneck(nn.Module):
    """Pre-activation (v2) bottleneck block."""
    def __init__(self, cin, cout=None, cmid=None, stride=1):
        super().__init__()
        cout = cout or cin
        cmid = cmid or cout//4

        self.gn1 = nn.GroupNorm(32, cmid, eps=1e-6)
        self.conv1 = conv1x1(cin, cmid, bias=False)
        self.gn2 = nn.GroupNorm(32, cmid, eps=1e-6)
        self.conv2 = conv4x4(cmid, cmid, stride, bias=False)
        self.gn3 = nn.GroupNorm(32, cmid, eps=1e-6)
        self.conv3 = conv4x4(cmid, cmid, stride, bias=False)
        self.gn4 = nn.GroupNorm(32, cout, eps=1e-6)
        self.conv4 = conv1x1(cmid, cout, bias=False)
        self.relu = nn.ReLU(inplace=True)

        if (stride != 1 or cin != cout):
            self.downsample = conv1x1(cin, cout, stride, bias=False)
            self.gn_proj = nn.GroupNorm(cout, cout)

    def forward(self, x):
        residual = x
        if hasattr(self, 'downsample'):
            residual = self.downsample(x)
            residual = self.gn_proj(residual)

        y = self.relu(self.gn1(self.conv1(x)))
        y = self.relu(self.gn2(self.conv2(y)))
        y = self.relu(self.gn3(self.conv3(y)))
        y = self.gn4(self.conv4(y))

        y = self.relu(residual + y)
        return y

    def load_from(self, weights, n_block, n_unit):
        # 加载原有层的权重
        conv1_weight = np2th(weights[pjoin(n_block, n_unit, "conv1/kernel")], conv=True)
        conv2_weight = np2th(weights[pjoin(n_block, n_unit, "conv2/kernel")], conv=True)
        conv3_weight = np2th(weights[pjoin(n_block, n_unit, "conv3/kernel")], conv=True)

        gn1_weight = np2th(weights[pjoin(n_block, n_unit, "gn1/scale")])
        gn1_bias = np2th(weights[pjoin(n_block, n_unit, "gn1/bias")])

        gn2_weight = np2th(weights[pjoin(n_block, n_unit, "gn2/scale")])
        gn2_bias = np2th(weights[pjoin(n_block, n_unit, "gn2/bias")])

        gn3_weight = np2th(weights[pjoin(n_block, n_unit, "gn3/scale")])
        gn3_bias = np2th(weights[pjoin(n_block, n_unit, "gn3/bias")])

        self.conv1.weight.copy_(conv1_weight)
        self.conv2.weight.copy_(conv2_weight)
        self.conv3.weight.copy_(conv3_weight)
        # 新增的conv4随机初始化
        nn.init.kaiming_normal_(self.conv4.weight, mode='fan_out', nonlinearity='relu')

        self.gn1.weight.copy_(gn1_weight.view(-1))
        self.gn1.bias.copy_(gn1_bias.view(-1))

        self.gn2.weight.copy_(gn2_weight.view(-1))
        self.gn2.bias.copy_(gn2_bias.view(-1))

        self.gn3.weight.copy_(gn3_weight.view(-1))
        self.gn3.bias.copy_(gn3_bias.view(-1))
        # 新增的gn4随机初始化
        nn.init.constant_(self.gn4.weight, 1)
        nn.init.constant_(self.gn4.bias, 0)

        if hasattr(self, 'downsample'):
            proj_conv_weight = np2th(weights[pjoin(n_block, n_unit, "conv_proj/kernel")], conv=True)
            proj_gn_weight = np2th(weights[pjoin(n_block, n_unit, "gn_proj/scale")])
            proj_gn_bias = np2th(weights[pjoin(n_block, n_unit, "gn_proj/bias")])

            self.downsample.weight.copy_(proj_conv_weight)
            self.gn_proj.weight.copy_(proj_gn_weight.view(-1))
            self.gn_proj.bias.copy_(proj_gn_bias.view(-1))


class ResNetV2(nn.Module):
    """Implementation of Pre-activation (v2) ResNet mode."""
    def __init__(self, block_units, width_factor):
        super().__init__()
        width = int(64 * width_factor)
        self.width = width

        self.root = nn.Sequential(OrderedDict([
            ('conv', StdConv2d(4, width, kernel_size=7, stride=2, bias=False, padding=3)),
            ('gn', nn.GroupNorm(32, width, eps=1e-6)),
            ('relu', nn.ReLU(inplace=True)),
        ]))

        self.body = nn.Sequential(OrderedDict([
            ('block1', nn.Sequential(OrderedDict(
                [('unit1', PreActBottleneck(cin=width, cout=width*4, cmid=width))] +
                [(f'unit{i:d}', PreActBottleneck(cin=width*4, cout=width*4, cmid=width)) for i in range(2, block_units[0] + 1)],
                ))),
            ('block2', nn.Sequential(OrderedDict(
                [('unit1', PreActBottleneck(cin=width*4, cout=width*8, cmid=width*2, stride=2))] +
                [(f'unit{i:d}', PreActBottleneck(cin=width*8, cout=width*8, cmid=width*2)) for i in range(2, block_units[1] + 1)],
                ))),
            ('block3', nn.Sequential(OrderedDict(
                [('unit1', PreActBottleneck(cin=width*8, cout=width*16, cmid=width*4, stride=2))] +
                [(f'unit{i:d}', PreActBottleneck(cin=width*16, cout=width*16, cmid=width*4)) for i in range(2, block_units[2] + 1)],
                ))),
        ]))

    def forward(self, x):
        features = []
        b, c, in_size, _ = x.size()
        x = self.root(x)
        features.append(x)
        x = nn.MaxPool2d(kernel_size=3, stride=2, padding=0)(x)
        for i in range(len(self.body)-1):
            x = self.body[i](x)
            right_size = int(in_size / 4 / (i+1))
            if x.size()[2] != right_size:
                pad = right_size - x.size()[2]
                assert pad < 3 and pad > 0, f"x {x.size()} should match {right_size}"
                feat = torch.zeros((b, x.size()[1], right_size, right_size), device=x.device)
                feat[:, :, :x.size()[2], :x.size()[3]] = x
            else:
                feat = x
            features.append(feat)
        x = self.body[-1](x)
        return x, features[::-1]

    def load_from(self, weights):
        # 处理根卷积层的权重:将3通道扩展为4通道
        root_conv_weight = np2th(weights[pjoin("root", "conv", "kernel")], conv=True)
        # 复制原3通道权重到第4通道
        new_root_weight = torch.cat([root_conv_weight, root_conv_weight[:, -1:, :, :]], dim=1)
        self.root.conv.weight.copy_(new_root_weight)

        root_gn_weight = np2th(weights[pjoin("root", "gn", "scale")])
        root_gn_bias = np2th(weights[pjoin("root", "gn", "bias")])
        self.root.gn.weight.copy_(root_gn_weight.view(-1))
        self.root.gn.bias.copy_(root_gn_bias.view(-1))

        # 加载各block的权重
        for block_idx, block_name in enumerate(['block1', 'block2', 'block3']):
            for unit_idx in range(1, len(self.body[block_idx])+1):
                self.body[block_idx][f'unit{unit_idx}'].load_from(weights, block_name, f'unit{unit_idx}')

关键修改说明

  1. 补全conv4x4函数:解决未定义函数的报错。
  2. 根卷积层权重扩展:将原3通道权重复制一份作为第4通道,保证输入4通道时权重维度匹配。
  3. 修正bottleneck的load_from:新增的conv4和gn4没有预训练权重,改为随机初始化,避免加载错误。
  4. 新增ResNetV2的load_from方法:统一处理根层和各block的权重加载逻辑。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.24 06:48:09