如何在PyTorch音频GAN模型中移除通道维度并适配LSTM层
解决GAN生成音频时Conv2d输出适配LSTM输入的维度问题
你遇到的核心问题是LSTM的输入维度顺序和当前张量的维度不匹配,而非单纯的移除维度操作失败。以下是具体的修正方案:
问题分析
你的Generator网络输出形状为(batch_size, 1, 256, 64),移除通道维度后得到(batch_size, 256, 64),但PyTorch的LSTM默认输入格式是(seq_len, batch_size, input_size),而你当前的张量是(batch_size, seq_len, input_size)的顺序,同时需要明确将256作为音频的时间步(序列长度),64作为每个时间步的特征维度。
解决方案
适配LSTM输入格式:
- 给LSTM添加
batch_first=True参数,让它直接接受(batch_size, seq_len, input_size)格式的输入,无需手动转置维度。 - 若不想修改
batch_first,则需手动将张量转置为(seq_len, batch_size, input_size)格式。
- 给LSTM添加
修正forward函数:
用squeeze(1)明确移除通道维度,再将调整后的张量送入LSTM即可。
修改后的完整代码
import torch import torch.nn as nn class Generator(nn.Module): def __init__(self, channels_noise, channels_img, features_g): super(Generator, self).__init__() self.net = nn.Sequential( # Input: N x channels_noise x 1 x 1 self._block(channels_noise, features_g * 16, (16,4), 1, 0), # 输出: N x 16*features_g x 16 x4 self._block(features_g * 16, features_g * 8, 4, 2, 1), # 输出: N x8*features_g x32x8 self._block(features_g * 8, features_g * 4, 4, 2, 1), # 输出: N x4*features_g x64x16 self._block(features_g * 4, features_g * 2, 4, 2, 1), # 输出: N x2*features_g x128x32 nn.ConvTranspose2d( features_g * 2, channels_img, kernel_size=4, stride=2, padding=1 ), # 最终输出: N x channels_img x256x64 nn.Tanh(), ) # 添加batch_first=True,适配(batch_size, seq_len, input_size)格式 self.lstm = nn.LSTM(input_size=64, hidden_size=10, batch_first=True) def _block(self, in_channels, out_channels, kernel_size, stride, padding): return nn.Sequential( nn.ConvTranspose2d( in_channels, out_channels, kernel_size, stride, padding, bias=False, ), nn.BatchNorm2d(out_channels), nn.ReLU(), ) def forward(self, x): x = self.net(x) # 移除第1维度(通道维度):从(batch,1,256,64) -> (batch,256,64) x = x.squeeze(1) # 此时x符合batch_first=True的LSTM输入格式:(batch, seq_len=256, input_size=64) lstm_out, _ = self.lstm(x) return lstm_out # 测试维度正确性 if __name__ == "__main__": noise = torch.randn(32, 100, 1, 1) # batch_size=32,噪声维度100 gen = Generator(channels_noise=100, channels_img=1, features_g=64) output = gen(noise) print(f"LSTM输出形状: {output.shape}") # 预期输出: (32,256,10),对应(batch, seq_len, hidden_size)
额外提示
- 若需要更复杂的序列建模,可以调整LSTM的
hidden_size,或叠加多层LSTM(例如nn.LSTM(input_size=64, hidden_size=64, num_layers=2, batch_first=True))。 - 调试时建议打印每一步的张量形状,快速定位维度问题。
内容的提问来源于stack exchange,提问作者Conweezy
相关产品推荐
相关产品推荐

