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

PyTorch胶囊网络适配自定义数据集时shape不匹配报错求助

问题分析与解决

错误根源

你遇到的RuntimeError: shape '[58, 2048, -1]' is invalid for input of size 534528本质是张量重塑时尺寸不匹配:

  • 输入张量总元素数为534528,按给定的[58,2048,-1]计算,58*2048=118784,534528/118784=4.5不是整数,无法完成重塑。
  • 核心原因是原胶囊网络代码针对MNIST的2828图像设计,换成3232自定义图像后,各卷积层输出的特征图尺寸变化,但代码中硬编码的胶囊数量、重塑参数未同步更新。

逐层修正步骤

1. 修正PrimaryCaps层的胶囊数量

原MNIST版本中,PrimaryCaps的特征图尺寸是66,胶囊数量为32*6*6=1152;换成3232图像后,经过ConvLayer(kernel=9, stride=1)输出24*24特征图,再经过PrimaryCaps的卷积层(kernel=9, stride=2),输出特征图尺寸为(24-9)//2 +1=8,胶囊数量应为32*8*8=2048。

修改PrimaryCaps的forward方法:

class PrimaryCaps(nn.Module):
    def __init__(self):
        super(PrimaryCaps, self).__init__()
        self.capsules = nn.ModuleList([
            nn.Conv2d(256, 8, kernel_size=9, stride=2) for _ in range(32)
        ])
    
    def squash(self, tensor, dim=-1):
        squared_norm = (tensor ** 2).sum(dim=dim, keepdim=True)
        scale = squared_norm / (1 + squared_norm)
        return scale * tensor / torch.sqrt(squared_norm + 1e-8)
    
    def forward(self, x):
        outputs = [capsule(x) for capsule in self.capsules]
        outputs = torch.cat(outputs, dim=1)
        # 替换硬编码的32*6*6,改为自动计算或显式32*8*8
        outputs = outputs.view(x.size(0), -1, 8)
        # 或显式写为 outputs = outputs.view(x.size(0), 32*8*8, 8)
        return self.squash(outputs)

2. 修正DigitCaps层的参数初始化

DigitCaps的权重矩阵依赖PrimaryCaps的胶囊数量,需同步更新:

class DigitCaps(nn.Module):
    def __init__(self, num_classes=10):
        super(DigitCaps, self).__init__()
        self.num_classes = num_classes
        self.num_primary_caps = 32*8*8  # 从32*6*6改为32*8*8
        self.in_dim = 8
        self.out_dim = 16
        self.W = nn.Parameter(torch.randn(self.num_classes, self.num_primary_caps, self.out_dim, self.in_dim))
    
    def forward(self, x):
        x = x.unsqueeze(1).unsqueeze(4)
        u_hat = torch.matmul(self.W, x)
        # 后续动态路由逻辑保持不变
        # ...

3. 修正Decoder层的输出尺寸

原Decoder针对MNIST的2828输出设计,需改为3232单通道图像:

class Decoder(nn.Module):
    def __init__(self, num_classes=10):
        super(Decoder, self).__init__()
        self.fc1 = nn.Linear(16*num_classes, 512)
        self.fc2 = nn.Linear(512, 1024)
        self.fc3 = nn.Linear(1024, 32*32*1)  # 从784改为1024
    
    def forward(self, x, y):
        mask = y.unsqueeze(2)
        x = x * mask
        x = x.view(x.size(0), -1)
        x = F.relu(self.fc1(x))
        x = F.relu(self.fc2(x))
        x = torch.sigmoid(self.fc3(x))
        return x.view(-1, 1, 32, 32)  # 重塑为32*32单通道图像

验证方法

在各层forward方法中加入打印语句,确认张量尺寸是否符合预期:

# 在ConvLayer的forward中
print("ConvLayer output shape:", x.shape)
# 在PrimaryCaps的forward中
print("PrimaryCaps after cat shape:", outputs.shape)
print("PrimaryCaps after view shape:", outputs.shape)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.11 21:40:41