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

为何MPS设备在反向传播中的性能表现逊于CPU?

MPS反向传播性能劣于CPU的原因与优化方案

核心原因分析

  • 设备启动与单次操作开销
    MPS存在固定的设备初始化、数据调度开销,单次反向传播操作的场景下,这些开销会抵消GPU的计算优势。而Apple Silicon的CPU本身搭载了高度优化的向量指令集,处理单次大矩阵运算时效率反而更高。

  • 算子优化成熟度差异
    PyTorch对MPS的支持远不如CUDA成熟,部分算子的反向传播实现尚未完全优化。比如你测试的矩阵乘法反向路径,MPS版本的效率仍有瓶颈;而CPU对应的算子经过多年打磨,性能已经趋近最优。在MNIST模型训练中,小算子(如激活、池化的反向)的优化不足会叠加导致整体速度变慢。

  • 内存与调度机制差异
    MPS的GPU内存调度逻辑和CPU的缓存机制不同,大张量反向传播可能触发内存带宽瓶颈。Apple Silicon的CPU配备超大缓存,处理大张量时缓存命中率更高,反而在这类场景中占优。

优化方案

  • 消除启动开销,多次循环取平均
    单次测试无法反映真实性能,建议加入设备预热+多次循环的测试逻辑,避免启动开销干扰:

    import torch
    import time
    device = torch.device("mps")
    
    # 预热MPS设备
    torch.rand((100, 100), device=device).sum().backward()
    
    x = torch.rand((10000, 10000), device=device, requires_grad=True)
    y = torch.rand((10000, 10000), device=device, requires_grad=True)
    
    num_runs = 10
    total_gpu_time = 0.0
    for _ in range(num_runs):
        loss = (x @ y).sum()
        x.grad.zero_()
        y.grad.zero_()
        start = time.time()
        loss.backward()
        end = time.time()
        total_gpu_time += end - start
    avg_gpu_time = total_gpu_time / num_runs
    print(f"MPS平均反向传播时间: {avg_gpu_time}")
    
    # CPU测试同样调整
    a = torch.rand((10000, 10000), device='cpu', requires_grad=True)
    b = torch.rand((10000, 10000), device='cpu', requires_grad=True)
    total_cpu_time = 0.0
    for _ in range(num_runs):
        l = (a @ b).sum()
        a.grad.zero_()
        b.grad.zero_()
        start = time.time()
        l.backward()
        end = time.time()
        total_cpu_time += end - start
    avg_cpu_time = total_cpu_time / num_runs
    print(f"CPU平均反向传播时间: {avg_cpu_time}")
    
  • 适配MPS的模型与训练策略
    优先使用PyTorch官方标记为MPS优化的算子,避免自定义或冷门算子;训练前确保模型、数据全部转移到MPS设备,杜绝训练过程中的跨设备数据传输;对于MNIST这类模型,尝试增大批量大小(如256、512),充分利用GPU的并行计算能力。

  • 升级PyTorch版本
    PyTorch对MPS的性能优化一直在迭代,新版本会修复大量算子的性能瓶颈。建议升级到2.0及以上的稳定版本,多数反向传播的性能问题在新版本中已有改善。

内容的提问来源于stack exchange,提问作者Hesham Abdelgawad

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.27 20:42:05