Mac M1 MAX设备PyTorch中MPS任务同步函数咨询
MPS环境下替代torch.cuda.synchronize()的同步函数
在PyTorch的MPS(Metal Performance Shaders)后端中,对应CUDA环境torch.cuda.synchronize()的同步函数是torch.mps.synchronize()。
这个函数会让CPU等待所有MPS设备上的异步任务执行完成,确保你统计的耗时是真实的计算时间,而非任务提交到MPS后的立即返回时间。
修改后的测试代码
import torch if torch.has_mps: device = torch.device("mps") else: device = torch.device("cpu") print("using", device, "device") import time matrix_size = 32*512 x = torch.randn(matrix_size, matrix_size) y = torch.randn(matrix_size, matrix_size) print("***********cpu speed***************") start = time.time() result = torch.matmul(x,y) print(time.time()-start) print("verify device: ", result.device) x_mps = x.to(device) y_mps = y.to(device) for i in range(3): print("***********mps speed***************") start = time.time() result_mps = torch.matmul(x_mps,y_mps) torch.mps.synchronize() # 替换原CUDA同步语句,等待MPS任务完成 print(time.time()-start) print("verify device: ", result_mps.device)
关键说明
- MPS运算默认异步执行,若不调用
torch.mps.synchronize(),time.time()-start得到的只是任务提交到MPS设备的时间,远小于真实计算耗时。 - 调用该函数后,CPU会阻塞至MPS上的矩阵乘法任务完全结束,此时的计时结果才能准确反映M1 MAX神经引擎的计算性能。
内容的提问来源于stack exchange,提问作者Bin YE
相关产品推荐
相关产品推荐

