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
相关产品推荐
相关产品推荐

