You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.26 12:54:57