M2 Mac使用MPS运行WaveGAN时linear_backward未实现报错求助
问题分析与解决思路
核心原因
你遇到的derivative for aten::linear_backward is not implemented错误,本质是PyTorch的MPS后端对部分算子的反向传播支持不完善——尤其是在使用torch.autograd.grad计算WGAN-GP梯度惩罚时,线性层的反向传播路径触发了未实现的算子。M2芯片的MPS是PyTorch较新支持的后端,复杂反向传播场景的算子覆盖还不全。
另外代码里有两处过时写法,可能加重问题:
- 过度使用
.data属性(PyTorch 1.0+已不推荐,应替换为detach()) - 手动封装
Variable(新版本Tensor已整合Variable功能,无需手动处理)
可行解决方案
1. 调整梯度惩罚计算逻辑
修改calculate_discriminator_loss函数,优化计算图逻辑并修正过时写法:
def calculate_discriminator_loss(self, real, generated): # 替换.data为detach(),避免破坏计算图 real = real.detach() generated = generated.detach() disc_out_gen = self.discriminator(generated) disc_out_real = self.discriminator(real) # 直接用torch.rand生成alpha,简化代码 alpha = torch.rand((batch_size * 2, 1, 1), device=device).expand_as(real) # 生成插值样本,无需手动调用.data interpolated = (1 - alpha) * real + alpha * generated[:batch_size * 2] interpolated.requires_grad_(True) # 直接设置requires_grad,替代Variable prob_interpolated = self.discriminator(interpolated) ones = torch.ones_like(prob_interpolated, device=device) # 关闭create_graph减少计算图复杂度(不影响判别器梯度更新) gradients = torch.autograd.grad( outputs=prob_interpolated, inputs=interpolated, grad_outputs=ones, create_graph=False, retain_graph=True, only_inputs=True, )[0] grad_penalty = ( p_coeff * ((gradients.view(gradients.size(0), -1).norm(2, dim=1) - 1) ** 2).mean() ) assert not torch.isnan(grad_penalty) assert not torch.isnan(disc_out_gen.mean()) assert not torch.isnan(disc_out_real.mean()) cost_wd = disc_out_gen.mean() - disc_out_real.mean() cost = cost_wd + grad_penalty return cost, cost_wd
同时修改主训练代码的.data调用:
disc_cost, disc_wd = self.calculate_discriminator_loss( real_signal, generated # 去掉.data,直接传入Tensor )
2. 梯度惩罚部分回退到CPU计算
如果调整逻辑后仍报错,可将梯度惩罚相关代码单独放到CPU执行,其他部分保留MPS加速:
def calculate_discriminator_loss(self, real, generated): disc_out_gen = self.discriminator(generated) disc_out_real = self.discriminator(real) # 将插值样本与相关计算移到CPU real_cpu = real.detach().cpu() generated_cpu = generated.detach().cpu() alpha = torch.rand((batch_size * 2, 1, 1)).expand_as(real_cpu) interpolated = (1 - alpha) * real_cpu + alpha * generated_cpu[:batch_size * 2] interpolated.requires_grad_(True) # 判别器临时移到CPU计算插值样本输出 self.discriminator = self.discriminator.cpu() prob_interpolated = self.discriminator(interpolated) self.discriminator = self.discriminator.to(device) # 移回MPS ones = torch.ones_like(prob_interpolated) gradients = torch.autograd.grad( outputs=prob_interpolated, inputs=interpolated, grad_outputs=ones, create_graph=True, retain_graph=True, only_inputs=True, )[0] grad_penalty = ( p_coeff * ((gradients.view(gradients.size(0), -1).norm(2, dim=1) - 1) ** 2).mean() ) grad_penalty = grad_penalty.to(device) # 移回MPS参与损失计算 assert not torch.isnan(grad_penalty) assert not torch.isnan(disc_out_gen.mean()) assert not torch.isnan(disc_out_real.mean()) cost_wd = disc_out_gen.mean() - disc_out_real.mean() cost = cost_wd + grad_penalty return cost, cost_wd
3. 等待PyTorch版本更新
MPS后端一直在迭代完善,后续版本可能会补上linear_backward的支持,可关注PyTorch官方更新日志,升级到最新稳定版后再尝试。
总结
这不是你忽略了简单错误,而是当前MPS后端对WGAN-GP这种嵌套梯度计算的场景支持不足。优先尝试调整梯度惩罚的计算逻辑,不行就用部分CPU回退的方案,能在不放弃大部分MPS加速的前提下继续训练。
内容的提问来源于stack exchange,提问作者David Bangerter
相关产品推荐
相关产品推荐

