如何在nn.Sequential中添加Reshape层?GAN生成器报错求助
解决GAN生成器中nn.Sequential内的Reshape问题
错误根源
- torch.reshape无法直接放入nn.Sequential:
torch.reshape是张量操作函数,并非nn.Module子类,而nn.Sequential仅接受可调用的模块类实例。你直接传入整数rdim作为第一个参数,自然触发类型错误。 - BatchNorm层误用:Linear层输出是
(batch_size, rdim)的2D张量,而BatchNorm2d要求输入为4D(batch, channels, height, width),这里应该用BatchNorm1d处理2D特征。
可行解决方案
方案1:使用nn.Unflatten(PyTorch 1.7+)
nn.Unflatten是官方提供的模块,专门用于在序列中完成维度展开,完美适配nn.Sequential。同时修正BatchNorm层和后续ConvTranspose2d的输入通道:
class Generator(nn.Module): def __init__(self, z_dim=100, im_chan=1, hidden_dim=64, rdim=9216): super(Generator, self).__init__() self.z_dim = z_dim # rdim = 256*6*6 = 9216,对应目标通道数256 self.gen = nn.Sequential( nn.Linear(z_dim, rdim), nn.BatchNorm1d(rdim, momentum=0.9), # 替换为BatchNorm1d适配2D输入 nn.ReLU(inplace=True), nn.Unflatten(1, (256, 6, 6)), # 将第1维(特征维)拆为(通道, 高, 宽) self.make_gen_block(256, hidden_dim*2), # 输入通道改为256 self.make_gen_block(hidden_dim*2, hidden_dim), self.make_gen_block(hidden_dim, im_chan, final_layer=True), ) def make_gen_block(self, input_channels, output_channels, kernel_size=4, stride=2, final_layer=False): # 建议默认kernel_size设为4,否则stride=2可能导致输出尺寸异常,可根据你的架构调整 if not final_layer: return nn.Sequential( nn.ConvTranspose2d(input_channels, output_channels, kernel_size, stride), nn.BatchNorm2d(output_channels), nn.ReLU(inplace=True) ) else: return nn.Sequential( nn.ConvTranspose2d(input_channels, output_channels, kernel_size, stride), nn.Tanh() ) def forward(self, noise): return self.gen(noise) # Linear接受2D噪声,无需额外扩维
方案2:自定义Reshape模块(兼容低版本PyTorch)
如果你的PyTorch版本低于1.7,自定义一个简单的Reshape模块即可:
class Reshape(nn.Module): def __init__(self, shape): super().__init__() self.shape = shape def forward(self, x): # 保留batch维度,其余维度按指定形状调整 return x.view(x.size(0), *self.shape)
然后在nn.Sequential中替换为:
self.gen = nn.Sequential( nn.Linear(z_dim, rdim), nn.BatchNorm1d(rdim, momentum=0.9), nn.ReLU(inplace=True), Reshape((256, 6, 6)), self.make_gen_block(256, hidden_dim*2), # 其余部分同方案1 )
额外修正点
- 原代码中
unsqueeze_noise方法的self.zdim是拼写错误,应改为self.z_dim;且Linear层接受2D输入,无需将噪声扩维为4D。 - 测试代码中
test_hidden_block的输入通道需与生成器实际通道对应,避免尺寸不匹配错误。
内容的提问来源于stack exchange,提问作者Chris Christian
相关产品推荐
相关产品推荐

