WGAN-LSTM带梯度惩罚训练报错:_cudnn_rnn_backward导数未实现
问题描述
尝试用LSTM作为判别器(Critic)和生成器的WGAN在MNIST数据集上训练图像生成,反复触发如下报错:
NotImplementedError: the derivative for '_cudnn_rnn_backward' is not implemented. Double backwards is not supported for CuDNN RNNs due to limitations in the CuDNN API. To run double backwards, please disable the CuDNN backend temporarily while running the forward pass of your RNN. For example: with torch.backends.cudnn.flags(enabled=False): output = model(inputs)
原以为未执行Double backwards操作,无法理解报错原因,求助解释及解决方法。
判别器(Critic)代码
class LSTM_Critic(nn.Module): def __init__(self, input_size, hidden_size, num_layers, num_classes): super(LSTM_Critic, self).__init__() self.hidden_size = hidden_size self.num_layers = num_layers self.lstm = nn.LSTM(IMG_SIZE, hidden_size, num_layers, batch_first=True) self.fc = nn.Linear(hidden_size, 1) def forward(self, x, labels): # Set initial hidden and cell states h0 = torch.zeros(self.num_layers, x.size(0), self.hidden_size).requires_grad_().to(device) c0 = torch.zeros(self.num_layers, x.size(0), self.hidden_size).requires_grad_().to(device) # Passing in the input and hidden state into the model and obtaining outputs x = x.reshape(BATCH_SIZE, IMG_SIZE, IMG_SIZE) out, hidden = self.lstm(x, (h0.detach(), c0.detach())) # out: tensor of shape (batch_size, seq_length, hidden_size) #Reshaping the outputs such that it can be fit into the fully connected layer out = self.fc(out[:, -1, :]) return out
初始化代码
gen = LSTM_Generator(200, 100, num_layers, num_classes).to(device) critic = LSTM_Critic(input_size, hidden_size, num_layers, num_classes).to(device) initialize_weights(gen) initialize_weights(critic) # initializate optimizer opt_gen = optim.Adam(gen.parameters(), lr=LEARNING_RATE, betas=(0.0, 0.9)) opt_critic = optim.Adam(critic.parameters(), lr=LEARNING_RATE, betas=(0.0, 0.9)) gen.train() critic.train()
训练代码
for epoch in range(NUM_EPOCHS): for batch_idx, (real, labels) in enumerate(tqdm(loader)): real = real.to(device) cur_batch_size = real.shape[0] labels = labels.to(device) # Train Critic: max E[critic(real)] - E[critic(fake)] # equivalent to minimizing the negative of that for _ in range(CRITIC_ITERATIONS): noise = torch.randn(cur_batch_size, Z_DIM, 1, 1).to(device) fake = gen(noise, labels) critic_real = critic(real, labels).reshape(-1) critic_fake = critic(fake, labels).reshape(-1) gp = gradient_penalty(critic, labels, real, fake, device=device) loss_critic = ( -(torch.mean(critic_real) - torch.mean(critic_fake)) + LAMBDA_GP * gp ) critic.zero_grad() loss_critic.backward(retain_graph=True) opt_critic.step()
梯度惩罚(Gradient Penalty)代码
def gradient_penalty(critic, labels, real, fake, device="cpu"): BATCH_SIZE, C, H, W = real.shape alpha = torch.rand((BATCH_SIZE, 1, 1, 1)).repeat(1, C, H, W).to(device) interpolated_images = real * alpha + fake * (1 - alpha) # Calculate critic scores mixed_scores = critic(interpolated_images, labels) # Take the gradient of the scores with respect to the images gradient = torch.autograd.grad( inputs=interpolated_images, outputs=mixed_scores, grad_outputs=torch.ones_like(mixed_scores), create_graph=True, retain_graph=True, )[0] gradient = gradient.reshape(gradient.shape[0], -1) gradient_norm = gradient.norm(2, dim=1) gradient_penalty = torch.mean((gradient_norm - 1) ** 2) return gradient_penalty
报错触发在loss_critic.backward(retain_graph=True)行。
报错原因
你确实触发了双重反向传播(Double backwards),只是并非显式调用两次backward(),而是来自WGAN-GP的梯度惩罚逻辑:
- 在
gradient_penalty函数中,torch.autograd.grad()设置了create_graph=True,这会构建梯度计算的计算图,目的是让梯度本身也能被求导(因为梯度惩罚项需要参与后续损失的反向传播)。 - 执行
loss_critic.backward()时,实际是对gp(梯度惩罚项)再次求导——也就是计算“梯度的梯度”,这就是报错中提到的Double backwards。 - 判别器使用了CuDNN加速的LSTM(PyTorch默认启用CuDNN加速RNN),但CuDNN的RNN实现不支持这种双重反向传播操作,因此抛出错误。
解决方法
有两种可行方案:
方案一:临时禁用CuDNN对LSTM的加速
按照报错提示,在判别器的LSTM前向传播时临时关闭CuDNN,修改判别器的forward方法:
def forward(self, x, labels): x = x.reshape(BATCH_SIZE, IMG_SIZE, IMG_SIZE) # 临时禁用CuDNN,规避双重反向传播兼容性问题 with torch.backends.cudnn.flags(enabled=False): # 无需手动初始化h0/c0,PyTorch会自动生成全零无梯度的隐藏状态 out, hidden = self.lstm(x) out = self.fc(out[:, -1, :]) return out
注意:禁用CuDNN会降低LSTM的前向/反向传播速度,但能直接解决报错问题。
方案二:替换为支持双重反向传播的LSTM实现
如果不想牺牲速度,可以手动实现LSTM的门控逻辑(用全连接层模拟),但这种方式代码改动较大,不如方案一简便。
内容的提问来源于stack exchange,提问作者Simon
相关产品推荐
相关产品推荐

