PyTorch CNN张量形状不兼容错误求助:矩阵相乘失败
解决PyTorch CNN中张量维度不兼容的问题
错误的核心原因是展平张量时的形状逻辑错误:
你打印出的卷积池化后张量形状为torch.Size([4, 128, 1, 1]),其中第一维是批量大小(4),第二维是特征通道数(128),后两维是空间维度(1x1)。但你用x.view((128 * 1 * 1, 4))将张量强行改为(128,4),而线性层lin1的输入特征数是128,它期望的输入形状是[批量大小, 输入特征数](即[4,128]),这直接导致矩阵乘法维度不匹配(128x4 和 128x50 无法完成矩阵乘法,正确的维度应该是4x128 和 128x50)。
修正方案
将展平操作替换为以下两种方式之一:
- 显式指定正确形状:
x = x.view(4, 128) - 通用自适应写法(推荐,避免硬编码维度,后续修改卷积池化参数时无需改动):
x = x.view(x.size(0), -1),其中x.size(0)获取当前批量大小,-1让PyTorch自动计算剩余维度的总特征数。
修改后的完整forward方法
def forward(self, x): x = self.pool1(F.relu(self.conv1(x))) x = self.pool2(F.relu(self.conv2(x))) print(x.size()) # 正确展平:保持批量维度在第一位,合并剩余维度为特征维度 x = x.view(x.size(0), -1) x = F.relu(self.lin1(x)) # 注意:如果目标生长条件可能包含负数,最后一层用ReLU会截断输出,可能影响MSE拟合效果,可根据实际需求调整激活函数 x = F.relu(self.lin2(x)) return x
额外规范建议
你的__init__方法中super调用顺序不符合规范,建议先调用父类初始化再定义属性:
def __init__(self, image_length): super(My_Net, self).__init__() self.image_length = image_length # 后续卷积、线性层定义...
内容的提问来源于stack exchange,提问作者Christian Garvin
相关产品推荐
相关产品推荐

