PyTorch线性网络维度不匹配报错及输出尺寸异常排查
问题分析与解决
核心错误点
线性层定义位置错误
你把torch.nn.Linear的定义写在了forward方法里,这会导致每次前向传播都重新创建新的线性层,不仅权重无法被优化器更新,还会因为input的计算依赖batch大小(tensorSize[0]),导致线性层的维度随batch动态变化,最终引发维度不匹配报错。输入张量未展平
输入是4维张量(B,3,64,64),但线性层默认只处理最后一个维度的特征。你直接把4维张量喂给线性层,会让线性层把前三个维度都当成batch维度的一部分,最终输出保留前三个维度,得到(B,3,64,6)的结果,而不是你需要的(B,6)。线性层维度计算错误
input = tensorSize[0] * tensorSize[1] * tensorSize[2]这个计算完全错误,tensorSize[0]是batch大小,是动态变化的值,不能用来定义线性层的输入维度。线性层的输入维度应该是固定的特征数,即通道数×高×宽=3×64×64=12288。
修正后的代码示例
import torch class LinearNet(torch.nn.Module): def __init__(self): super().__init__() # 输入展平后特征数为3*64*64=12288 self.linear1 = torch.nn.Linear(12288, 32) self.linear2 = torch.nn.Linear(32, 6) def forward(self, x): # 展平输入:(B,3,64,64) -> (B, 12288) x_flat = x.view(x.size(0), -1) # 前向传播,可按需添加激活函数 out = torch.nn.functional.relu(self.linear1(x_flat)) out = self.linear2(out) return out # 输出尺寸为(B,6)
验证说明
- 展平操作
x.view(x.size(0), -1)会保留batch维度,将剩余的所有维度(3、64、64)合并成一个特征维度,得到标准的线性层输入格式(B, 特征数)。 - 线性层全部移到
__init__方法中定义,维度固定,权重可以被正常更新,不会再出现动态维度导致的匹配错误。
内容的提问来源于stack exchange,提问作者Aeryes
相关产品推荐
相关产品推荐

