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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 06:41:34