
PyTorch 官方 MPS 后端 API 完全指南从设备管理、内存控制到 Metal Shader 与 Profiler【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorchtorch.mps是 PyTorch 中面向 Apple Metal GPU 的后端接口模块它封装了 Metal Performance ShadersMPS框架让张量运算可以调度到 Mac 的 GPU 上执行从而获得加速效果。本文基于仓库中的官方文档 docs/source/mps.md 及其对应实现源码 torch/mps/init.py、torch/mps/profiler.py、torch/mps/event.py 展开系统讲解 MPS 后端的全部公开 API从设备可用性检测、全局同步、随机数种子管理到显存分配策略、自定义 Metal Shader 编译加载再到 OS Signpost 性能剖析与事件计时帮助你完整掌握在 Apple Silicon 上使用 PyTorch GPU 加速的实践方法。一、模块概览torch.mps 提供哪些能力从文档的 API 清单与源码的__all__导出列表见 torch/mps/init.py可以看到torch.mps模块的能力可以归纳为四大类类别公开 API用途设备与可用性is_available()、device_count()检测 MPS 后端是否可用及设备数量同步与流控制synchronize()等待 MPS 设备上所有流的所有 kernel 执行完毕随机数状态get_rng_state()、set_rng_state()、manual_seed()、seed()管理 MPS 设备上的随机数生成器RNG内存管理empty_cache()、set_per_process_memory_fraction()、current_allocated_memory()、driver_allocated_memory()、recommended_max_memory()控制 MPS 缓存分配器的显存使用自定义 Shadercompile_shader()、load_metallib()在 Python 中直接编译并调用 Metal compute kernel剖析与计时profiler子模块start/stop/profile/metal_capture等、event.Event生成 OS Signpost 追踪、捕获 GPU trace、事件计时此外模块还导出了一个内部辅助函数_host_alias_storage()用于让 CPU 侧torch.UntypedStorage直接别名 MPS 分配器分配的宿主可见MTLBuffer内存主要用于与 safetensors 等批量加载器做高级互操作。二、设备可用性检测与全局同步2.1 is_available() 与 device_count()is_available()用于判断当前环境macOS Apple Silicon/AMD GPU 编译时启用 MPS是否支持 MPS 后端其实现直接复用了device_count()def is_available() - bool: return device_count() 0 def device_count() - int: return int(torch._C._has_mps and torch._C._mps_is_available())这两层检查_has_mps为编译期宏开关、_mps_is_available()为运行时检测确保在非 Apple 平台或未启用 MPS 的 PyTorch 构建上安全返回。因此在写跨平台代码时惯用写法是import torch if torch.backends.mps.is_available(): device torch.device(mps) x torch.ones(4, 4, devicedevice) else: print(MPS 不可用将回退到 CPU)2.2 synchronize()MPS 与 CUDA 类似kernel 的执行是异步的。synchronize()会阻塞 CPU 线程等待 MPS 设备上所有流stream中的所有 kernel 完成def synchronize() - None: return torch._C._mps_deviceSynchronize()底层对应 C 绑定中的_mps_deviceSynchronize见 torch/csrc/mps/Module.cpp。它最常见的用途是在基准测试前后插入以获得准确的计时结果在测试代码中test/test_mps.py 大量使用torch.mps.synchronize()与torch.mps.empty_cache()配对确保内存统计不被异步执行中的 kernel 干扰例如该文件 L959-L976、L4926-L4947 处的内存与同步用例。三、随机数生成器RNG状态管理MPS 后端为每个设备维护一个默认的随机数生成器。torch.mps提供了一组与torch.cuda对称的 RNG API源码中通过模块级缓存的_get_default_mps_generator()获取底层torch._C.Generator对象见 torch/mps/init.py。API说明manual_seed(seed: int)设置随机数种子保证可复现性。内部先检查_has_mps不可用时直接返回而不报错因此可以被全局的torch.manual_seed()安全调用seed()用系统随机数设置种子get_rng_state(devicemps)返回当前 RNG 状态类型为ByteTensorset_rng_state(new_state, devicemps)恢复 RNG 状态内部会先将状态张量clone(memory_formattorch.contiguous_format)转为连续内存再设置避免非连续布局导致的问题典型用法import torch torch.mps.manual_seed(230) a torch.randn(3, 3, devicemps) state torch.mps.get_rng_state() # 保存状态 torch.mps.seed() # 打乱种子 torch.mps.set_rng_state(state) # 恢复之后生成的随机数序列与保存点一致 torch.mps.synchronize()在 test/test_mps.py 的 L9619-L9675 附近有完整的manual_seed → get_rng_state → set_rng_state → seed往返测试验证状态保存与恢复的正确性。四、MPS 内存管理缓存分配器与显存上限MPS 后端自带一个缓存分配器MPSAllocator通过内存池复用MTLBuffer避免频繁向 Metal 驱动申请/释放内存。4.1 三个内存查询 APIAPI返回值含义current_allocated_memory()当前张量实际占用的 GPU 内存字节不包含MPSAllocator 内存池中的缓存driver_allocated_memory()进程从 Metal 驱动申请到的总 GPU 内存字节包含MPSAllocator 池中的缓存以及 MPS/MPSGraph 框架自身的分配recommended_max_memory()Metal 设备推荐的 GPU 工作集大小上限字节对应 Metal API 的device.recommendedMaxWorkingSetSize三者关系可以理解为current_allocated_memory() ≤ driver_allocated_memory()差值主要来自缓存池中的闲置块。在 test/test_mps.py L109-L141 的TestMPSAllocator用例中正是通过比较这两个数值来验证empty_cache()是否释放了缓存。4.2 empty_cache()释放缓存分配器中所有未被占用的缓存内存使其可供其他 GPU 应用使用torch.mps.empty_cache()该函数不会强制释放仍被张量占用的内存只回收空闲缓存。它是排查 Mac 上内存告警的常用手段。4.3 set_per_process_memory_fraction(fraction)限制当前进程在 MPS 设备上的最大可分配内存。允许的内存 fraction × recommended_max_memory()。需要注意的关键约束fraction必须是float类型且取值范围为0 ~ 2否则分别抛出TypeError与ValueError见 torch/mps/init.py传入0表示不限制若内存不足可能导致系统级故障传入大于1.0的值允许突破recommendedMaxWorkingSetSize的限制一旦进程尝试分配超过该上限的内存分配器会抛出内存不足OOM错误。import torch # 将进程的 MPS 显存上限设为推荐工作集大小的 0.5 倍 torch.mps.set_per_process_memory_fraction(0.5) # 查询当前上限对应的字节数 print(torch.mps.recommended_max_memory())五、在 Python 中编译与加载 Metal Shader这是torch.mps最具特色的能力无需编写 C 扩展直接在 Python 运行时编译 Metal compute shader并像调用普通函数一样调用其中定义的 kernel。5.1 compile_shader(source)接收一段 Metal Shading LanguageMSL源码字符串编译后返回一个 shader 库对象。文档给出的完整示例lib torch.mps.compile_shader( kernel void full(device float* out, constant float val, uint idx [[thread_position_in_grid]]) { out[idx] val; } ) x torch.zeros(16, devicemps) lib.full(x, 3.14) # 用 3.14 填充 x实现上torch/mps/init.py做了两件事一是通过_embed_headers将torch/include目录下的头文件内联进源码保证可编辑安装场景下也能正确解析头文件路径二是调用 C 绑定_mps_compileShader完成编译见 torch/csrc/mps/Module.cpp。5.2 load_metallib(source)加载预编译的.metallib库文件返回可调用其中 kernel 的 shader 库对象。source参数支持两种形式bytes/bytearray直接传入 metallib 的原始字节内容走_mps_loadMetalllib绑定str/os.PathLike传入.metallib文件路径走_mps_loadMetallibFromPath绑定其他类型会抛出TypeError。# 从文件加载预编译库 lib torch.mps.load_metallib(kernels.metallib) x torch.ones(16, devicemps) lib.square(x)该接口特别适合加载由外部工具如 Triton、MetalASM提前生成的 Metal 库将编译期与运行期解耦。六、MPS ProfilerOS Signpost 与 Metal GPU Capturetorch.mps.profiler子模块实现见 torch/mps/profiler.py提供两类性能剖析手段OS Signpost 追踪和 Metal GPU Capture。6.1 OS Signpost 追踪start / stop / profileOS Signpost 是 Apple 的日志追踪机制产生的 trace 可以用 Xcode Instruments 的 Logging 工具查看。start(modeinterval, wait_until_completedFalse)开始生成 OS Signpost。mode取值interval记录每个操作执行的持续时间、event标记执行完成时刻、或interval,event两者都记录wait_until_completedTrue时会等待 MPS 流完成每个已编码的 GPU 操作让 trace 时间线上呈现单一 dispatch但会明显降低性能内部对 mode 做lower()与去空格归一化后传给底层_mps_profilerStartTrace。stop()停止生成 OS Signpost。profile(mode, wait_until_completed)contextlib.contextmanager形式的上下文管理器等价于start()yieldfinally: stop()异常时也能保证停止追踪。import torch with torch.mps.profiler.profile(modeinterval,event, wait_until_completedFalse): a torch.randn(1024, 1024, devicemps) b a a torch.mps.synchronize()6.2 Metal GPU Capturemetal_capturemetal_capture(fname)是一个上下文管理器用于把上下文内所有 Metal 调用捕获到一份.gputrace文件中之后可以用 Xcode 打开逐条查看 GPU 命令with torch.mps.profiler.metal_capture(my_trace.gputrace): c torch.mm(a, b)其配套的两个查询函数is_metal_capture_enabled()返回metal_capture上下文管理器是否可用。需要在启动进程前设置环境变量MTL_CAPTURE_ENABLED否则不可用is_capturing_metal()返回当前是否正在捕获 Metal 调用。注意底层_mps_stopCapture在结束捕获前会等待 MPS 流上已入队的工作完成即使上下文体内抛出了异常也会执行见 torch/mps/profiler.py。七、MPS Event流同步与计时torch.mps.event.Event实现见 torch/mps/event.py是对 MPS 事件的封装与 CUDA 事件的用法一致用于监测设备进度、测量耗时和同步流。import torch start_event torch.mps.event.Event(enable_timingTrue) end_event torch.mps.event.Event(enable_timingTrue) start_event.record() x torch.randn(2048, 2048, devicemps) y x x torch.mps.synchronize() end_event.record() end_event.synchronize() print(f耗时: {start_event.elapsed_time(end_event):.2f} ms)Event 的完整方法集如下方法说明record()在默认流中记录该事件wait()让默认流上之后提交的所有工作等待该事件完成query()返回事件捕获的所有工作是否已完成布尔值非阻塞synchronize()阻塞 CPU 线程直到事件完成elapsed_time(end_event)返回本事件记录到end_event记录之间的毫秒耗时事件 ID 由底层_mps_acquireEvent/_mps_releaseEvent管理析构时若torch._C尚未销毁且事件 ID 有效会自动释放见 torch/mps/event.py。八、底层实现与测试验证torch.mps的所有 Python API 最终都收敛到 C 绑定层。在 torch/csrc/mps/Module.cpp 中可以看到完整的绑定表_mps_deviceSynchronize、_mps_is_available、_mps_emptyCache、_mps_setMemoryFraction、_mps_currentAllocatedMemory、_mps_driverAllocatedMemory、_mps_recommendedMaxMemory、_mps_profilerStartTrace、_mps_acquireEvent/_mps_recordEvent/_mps_waitForEvent/_mps_synchronizeEvent/_mps_queryEvent/_mps_elapsedTimeOfEvents等一应俱全而 shader 相关能力_mps_compileShader、_mps_loadMetallibFromPath、_mps_isCaptureEnabled、_mps_isCapturing、_mps_startCapture、_mps_stopCapture则在同文件 L552-L567 处以 pybind 方式注册。仓库中的 test/test_mps.py 是验证这些 API 行为的最佳参考涵盖内存分配器测试TestMPSAllocator用current_allocated_memory()/driver_allocated_memory()对比验证empty_cache()的回收效果RNG 状态往返测试L9619-L9675验证manual_seed/seed/get_rng_state/set_rng_state的序列可复现性同步与内存统计用例L959-L1010验证synchronize()与缓存释放的配合使用显存上限与recommended_max_memory()的查询用例L9694 附近。九、小结torch.mps为 PyTorch 在 Apple 平台上的 GPU 加速提供了一套完整且自洽的 Python 接口日常训练推理用devicemps即可配套的synchronize()与 RNG 状态 API 保证正确性与可复现性内存紧张的场景用set_per_process_memory_fraction()/empty_cache()主动管控显存需要深度定制时可用compile_shader()/load_metallib()直接调用 Metal kernel性能调优阶段则借助profiler的 OS Signpost 与.gputrace捕获配合Event做精细计时。从 Python 封装到 C 绑定再到测试用例整个模块链路清晰、行为有据可查是理解 PyTorch 后端抽象与 Apple GPU 编程之间关系的绝佳范例。【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考