PyTorch中CNN模型输出形状不符合预期的问题排查
DQN网络输出形状异常问题解决
测试自定义DQN网络时,最后一层线性层设置out_feature=6以匹配环境Discrete(6)的动作空间,但传入单个状态张量后,模型返回了形状为torch.Size([128,6])的输出,而非预期的torch.Size([1,6])。
问题代码
import torch import torch.nn as nn import gym class DQN(nn.Module): def __init__(self, n_channels, n_actions): super(DQN, self).__init__() self.network = nn.Sequential( nn.Conv2d(n_channels, 32, kernel_size = 3, padding = 1), nn.ReLU(), nn.Conv2d(32,64, kernel_size = 3, padding = 1), nn.MaxPool2d(2,2), nn.Conv2d(64,128, kernel_size = 3, padding = 1), nn.ReLU(), nn.Conv2d(128,128, kernel_size = 3, padding = 1), nn.MaxPool2d(2,2), nn.Flatten(), nn.Linear(2480,128), nn.ReLU(), nn.Linear(128,64), nn.ReLU(), nn.Linear(64, n_actions) ) def forward(self, x): return self.network(x) env = gym.make('AirRaid-v4', render_mode = 'rgb_array') # Observation space of the environment is Box(250,160,3) # Action space of the environment is Discrete(6) policy_net = DQN(3,6).to('cpu') state = env.observation_space.sample() state = torch.tensor(state, dtype = torch.float32, device = 'cpu').reshape((3,160,250)) output = policy_net.forward(state) print(output) print(output.shape)
模型输出
>> [[ 1.6506e+00, -5.5547e-01, -2.3225e+00, 5.3022e-01, 2.0827e-01, -8.5433e-02], [ 6.1949e-01, -6.5546e-01, -1.4047e+00, -1.7372e-01, -1.3691e-01, -3.6298e-01], ... [ 5.0159e-01, -5.9395e-01, -8.0273e-01, -6.5626e-01, 4.1781e-01, 7.8058e-01], [ 9.9379e-01, -3.2483e-01, -1.4414e+00, 6.7811e-02, -1.0157e-01, -7.2960e-01]]
输出形状
torch.Size([128,6])
问题原因
- 输入缺少batch维度:PyTorch的CNN要求输入格式为
(batch_size, channels, height, width),你传入的张量是(3,160,250),模型会把通道维度3误判为batch_size,后续卷积池化操作后,最终输出的batch维度被错误放大。 - 全连接层输入维度计算错误:经过两次
MaxPool2d(2,2)后,特征图尺寸应为(128, 40, 62),展平后维度是128*40*62=317440,而非代码中写的2480,维度不匹配进一步导致输出形状异常。
解决方案
1. 给输入添加batch维度
将reshape后的张量通过unsqueeze(0)增加batch维度,变成(1,3,160,250):
state = torch.tensor(state, dtype=torch.float32, device='cpu').reshape((3,160,250)).unsqueeze(0)
2. 修正全连接层输入维度
替换错误的2480为正确的展平维度:
# 计算正确的展平后维度 flatten_dim = 128 * 40 * 62 # 317440 nn.Linear(flatten_dim, 128),
修正后完整代码
import torch import torch.nn as nn import gym class DQN(nn.Module): def __init__(self, n_channels, n_actions): super(DQN, self).__init__() # 修正展平维度计算 flatten_dim = 128 * 40 * 62 self.network = nn.Sequential( nn.Conv2d(n_channels, 32, kernel_size=3, padding=1), nn.ReLU(), nn.Conv2d(32,64, kernel_size=3, padding=1), nn.MaxPool2d(2,2), nn.Conv2d(64,128, kernel_size=3, padding=1), nn.ReLU(), nn.Conv2d(128,128, kernel_size=3, padding=1), nn.MaxPool2d(2,2), nn.Flatten(), nn.Linear(flatten_dim,128), nn.ReLU(), nn.Linear(128,64), nn.ReLU(), nn.Linear(64, n_actions) ) def forward(self, x): return self.network(x) env = gym.make('AirRaid-v4', render_mode='rgb_array') policy_net = DQN(3,6).to('cpu') state = env.observation_space.sample() # 添加batch维度 state = torch.tensor(state, dtype=torch.float32, device='cpu').reshape((3,160,250)).unsqueeze(0) output = policy_net(state) print(output) print(output.shape) # 预期输出:torch.Size([1,6])
内容的提问来源于stack exchange,提问作者Toby
相关产品推荐
相关产品推荐

