使用PyTorch nn.Sequential遇RuntimeError:矩阵形状无法相乘的解决方法
PyTorch中nn.Sequential实现维度展平的解决方案
问题背景
将可正常运行的PyTorch模型改用nn.Sequential重构后,前向传播触发如下错误:
RuntimeError: mat1 and mat2 shapes cannot be multiplied (64x49 and 3136x512)
原模型通过x = x.view(x.size()[0], -1)将卷积输出的4D张量展平为2D张量,以适配后续全连接层,但nn.Sequential无法直接嵌入原生的view操作,需要特殊处理。
原可运行代码
import torch.nn as nn import torch.nn.functional as F import torch class Net(nn.Module): def __init__(self, out_dims): super(Net, self).__init__() self.conv1 = nn.Conv2d(4, 32, 8, 4) self.relu1 = nn.ReLU(inplace=True) self.conv2 = nn.Conv2d(32, 64, 4, 2) self.relu2 = nn.ReLU(inplace=True) self.conv3 = nn.Conv2d(64, 64, 3, 1) self.relu3 = nn.ReLU(inplace=True) self.fc4 = nn.Linear(3136, 512) self.relu4 = nn.ReLU(inplace=True) self.fc5 = nn.Linear(512, out_dims) def forward(self, x): x = self.conv1(x) x = self.relu1(x) x = self.conv2(x) x = self.relu2(x) x = self.conv3(x) x = self.relu3(x) # torch.Size([1, 64, 7, 7]) x = x.view(x.size()[0], -1) # torch.Size([1, 3136]) x = self.fc4(x) x = self.relu4(x) x = self.fc5(x) return x net = Net(2) x = torch.rand(4, 84, 84).unsqueeze(0) net(x)
重构后报错的代码
import torch.nn as nn import torch.nn.functional as F import torch class Net(nn.Module): def __init__(self, out_dims): super(Net, self).__init__() self.layers = nn.Sequential( nn.Conv2d(4, 32, 8, 4), nn.ReLU(inplace=True), nn.Conv2d(32, 64, 4, 2), nn.ReLU(inplace=True), nn.Conv2d(64, 64, 3, 1), nn.ReLU(inplace=True), nn.Linear(3136, 512), # 此处触发维度不匹配错误 nn.ReLU(inplace=True), nn.Linear(512, out_dims), ) def forward(self, x): return self.layers(x) net = Net(2) x = torch.rand(4, 84, 84).unsqueeze(0) net(x)
解决方案
nn.Sequential仅接受继承自nn.Module的层,因此需要将维度展平操作封装为一个自定义层,或使用PyTorch内置的展平层。
方法1:使用PyTorch内置nn.Flatten(推荐,PyTorch 1.7+)
PyTorch 1.7及以上版本提供了nn.Flatten层,默认保留batch维度,将其余维度展平,完美替代view(x.size(0), -1):
修正后的完整代码:
import torch.nn as nn import torch class Net(nn.Module): def __init__(self, out_dims): super(Net, self).__init__() self.layers = nn.Sequential( nn.Conv2d(4, 32, 8, 4), nn.ReLU(inplace=True), nn.Conv2d(32, 64, 4, 2), nn.ReLU(inplace=True), nn.Conv2d(64, 64, 3, 1), nn.ReLU(inplace=True), nn.Flatten(), # 插入展平层 nn.Linear(3136, 512), nn.ReLU(inplace=True), nn.Linear(512, out_dims), ) def forward(self, x): return self.layers(x) net = Net(2) x = torch.rand(4, 84, 84).unsqueeze(0) net(x)
方法2:自定义展平层
如果使用低版本PyTorch,可手动实现一个展平层:
import torch.nn as nn import torch class Flatten(nn.Module): def forward(self, x): # 保留batch维度,展平其余维度 return x.view(x.size(0), -1) class Net(nn.Module): def __init__(self, out_dims): super(Net, self).__init__() self.layers = nn.Sequential( nn.Conv2d(4, 32, 8, 4), nn.ReLU(inplace=True), nn.Conv2d(32, 64, 4, 2), nn.ReLU(inplace=True), nn.Conv2d(64, 64, 3, 1), nn.ReLU(inplace=True), Flatten(), # 插入自定义展平层 nn.Linear(3136, 512), nn.ReLU(inplace=True), nn.Linear(512, out_dims), ) def forward(self, x): return self.layers(x) net = Net(2) x = torch.rand(4, 84, 84).unsqueeze(0) net(x)
原理说明
卷积层输出的是4D张量(形状为[batch_size, channels, height, width]),而全连接层nn.Linear要求输入为2D张量(形状为[batch_size, feature_num])。原代码中的view操作就是将4D张量展平为2D,在nn.Sequential中必须通过层的形式嵌入这个操作,才能保证数据流的连续性。
内容的提问来源于stack exchange,提问作者gameveloster
相关产品推荐
相关产品推荐

