为何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
相关产品推荐
相关产品推荐

