PyTorch 官方 MPS 后端 API 完全指南:从设备管理、内存控制到 Metal Shader 与 Profiler
2026/9/10 1:25:05 网站建设 项目流程

PyTorch 官方 MPS 后端 API 完全指南:从设备管理、内存控制到 Metal Shader 与 Profiler

【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch

torch.mps是 PyTorch 中面向 Apple Metal GPU 的后端接口模块,它封装了 Metal Performance Shaders(MPS)框架,让张量运算可以调度到 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, device=device) 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(device="mps")返回当前 RNG 状态,类型为ByteTensor
set_rng_state(new_state, device="mps")恢复 RNG 状态,内部会先将状态张量clone(memory_format=torch.contiguous_format)转为连续内存再设置,避免非连续布局导致的问题

典型用法:

import torch torch.mps.manual_seed(230) a = torch.randn(3, 3, device="mps") 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 三个内存查询 API

API返回值含义
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,否则分别抛出TypeErrorValueError(见 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 Language(MSL)源码字符串,编译后返回一个 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, device="mps") lib.full(x, 3.14) # 用 3.14 填充 x

实现上(torch/mps/init.py)做了两件事:一是通过_embed_headerstorch/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, device="mps") lib.square(x)

该接口特别适合加载由外部工具(如 Triton、MetalASM)提前生成的 Metal 库,将"编译期"与"运行期"解耦。

六、MPS Profiler:OS Signpost 与 Metal GPU Capture

torch.mps.profiler子模块(实现见 torch/mps/profiler.py)提供两类性能剖析手段:OS Signpost 追踪和 Metal GPU Capture。

6.1 OS Signpost 追踪:start / stop / profile

OS Signpost 是 Apple 的日志追踪机制,产生的 trace 可以用 Xcode Instruments 的 Logging 工具查看。

  • start(mode="interval", wait_until_completed=False):开始生成 OS Signpost。
    • mode取值:"interval"(记录每个操作执行的持续时间)、"event"(标记执行完成时刻)、或"interval,event"(两者都记录);
    • wait_until_completed=True时,会等待 MPS 流完成每个已编码的 GPU 操作,让 trace 时间线上呈现单一 dispatch,但会明显降低性能;
    • 内部对 mode 做lower()与去空格归一化后传给底层_mps_profilerStartTrace
  • stop():停止生成 OS Signpost。
  • profile(mode, wait_until_completed)contextlib.contextmanager形式的上下文管理器,等价于start()+yield+finally: stop(),异常时也能保证停止追踪。
import torch with torch.mps.profiler.profile(mode="interval,event", wait_until_completed=False): a = torch.randn(1024, 1024, device="mps") b = a @ a torch.mps.synchronize()

6.2 Metal GPU Capture:metal_capture

metal_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_timing=True) end_event = torch.mps.event.Event(enable_timing=True) start_event.record() x = torch.randn(2048, 2048, device="mps") 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 接口:日常训练推理用device="mps"即可,配套的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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询