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

训练Resnet/CAE时MSELoss张量尺寸不匹配问题求助

解决卷积自编码器(CAE)张量尺寸不匹配导致的MSELoss报错问题

问题根源

解码器输出[256, 3, 512, 512]和输入[256, 3, 32, 32]尺寸差了16倍,核心原因是编码器的下采样倍数和解码器的上采样倍数完全不对等——编码器把输入缩小的比例,解码器没有对应放大回去,反而过度放大了。

具体修复步骤

1. 先明确编码器的下采样路径

先给编码器的forward函数加打印,搞清楚输入是怎么从32x32变成最终特征图的:

def forward(self, x):
    print(f"Input: {x.shape}")
    x = self.conv1(x)
    print(f"After conv1: {x.shape}")
    # 依次打印后续每一层的输出尺寸
    # ...
    return x

比如用ResNet18处理32x32输入时,默认会经过4次2倍下采样(32→16→8→4→2),最终特征图尺寸为2x2。

2. 让解码器上采样完全匹配编码器下采样

编码器每做一次2倍下采样,解码器就要对应做一次2倍上采样,把尺寸拉回原大小。推荐用转置卷积,能同时完成上采样和特征变换,下面是对应2x2特征图回到32x32的解码器示例:

class Decoder(nn.Module):
    def __init__(self, in_channels=512):
        super().__init__()
        self.decoder_layers = nn.Sequential(
            # 2x2 → 4x4
            nn.ConvTranspose2d(in_channels, 256, kernel_size=4, stride=2, padding=1),
            nn.ReLU(inplace=True),
            # 4x4 → 8x8
            nn.ConvTranspose2d(256, 128, kernel_size=4, stride=2, padding=1),
            nn.ReLU(inplace=True),
            # 8x8 → 16x16
            nn.ConvTranspose2d(128, 64, kernel_size=4, stride=2, padding=1),
            nn.ReLU(inplace=True),
            # 16x16 → 32x32
            nn.ConvTranspose2d(64, 3, kernel_size=4, stride=2, padding=1),
            nn.Sigmoid()  # 若输入是归一化到0-1的图像,用此激活函数
        )
    
    def forward(self, x):
        return self.decoder_layers(x)

这里的转置卷积参数kernel_size=4, stride=2, padding=1是标准的2倍上采样配置,能保证尺寸刚好翻倍,无多余裁剪或填充。

3. 预训练ResNet编码器的适配要点

如果用预训练ResNet当编码器,必须去掉最后的全连接层和平均池化层,保留特征图输出:

import torchvision.models as models

encoder = models.resnet18(pretrained=True)
# 移除avgpool和fc层,得到特征图输出
encoder = nn.Sequential(*list(encoder.children())[:-2])
# 测试输入尺寸
test_input = torch.randn(256, 3, 32, 32)
feat_map = encoder(test_input)
print(f"Encoder output shape: {feat_map.shape}")  # 输出应为[256, 512, 2, 2]

拿到这个尺寸后,再对应设计解码器的上采样次数。

4. 预验证尺寸匹配

修改完网络后,先跑一次前向传播确认尺寸是否一致:

encoder = YourEncoder()
decoder = YourDecoder()
test_x = torch.randn(256, 3, 32, 32)
recon_x = decoder(encoder(test_x))
print(f"Reconstruction shape: {recon_x.shape}")
assert recon_x.shape == test_x.shape, "尺寸不匹配,请重新调整解码器"

确认没问题再启动训练,就不会触发MSELoss的尺寸错误了。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 15:39:52