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

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)。

修正方案

将展平操作替换为以下两种方式之一:

  1. 显式指定正确形状:x = x.view(4, 128)
  2. 通用自适应写法(推荐,避免硬编码维度,后续修改卷积池化参数时无需改动):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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.09 03:31:04