metal-kernel:把 PyTorch 在 Apple Silicon 上的内核,从 MPSGraph 拉回原生 Metal

用 Apple Silicon 跑 PyTorch 的人,多半对 MPS 后端又爱又恨。爱的是它确实让 M1/M2/M3 那几颗芯片有了存在感,恨的是它总在关键时刻掉链子,要么算子不支持,要么数值对不上,要么性能比 CPU 还拉胯。

问题的根子,PyTorch 官方自己也承认,出在 MPSGraph 这套实现方式上。MPSGraph 是 Apple 的高层图 API,PyTorch 早期为了快速铺开算子,用它在上面糊了一层。图 API 上手快,但控制力差,性能优化空间被框死了。

metal-kernel:把 PyTorch 在 Apple Silicon 上的内核,从 MPSGraph 拉回原生 Metal

于是有了 metal-kernel 这个 skill。它是 pytorch 官方在 Smithery 上发的,专门教你绕过 MPSGraph,用原生的 Metal 着色语言直接给算子写内核。一句话概括它的立场:走 c10/metal 这条基础设施,不走 MPSGraph。

这篇文章会带你走一遍它定义的两条工作流。一条是从零给算子加 MPS 支持,一条是把已经存在的 MPSGraph 实现迁移到原生 Metal。读完之后,你会明白这套东西为什么值得 PyTorch 维护者专门整理成 skill 发出来。

环境准备

这套工作流有个硬前提,你得从源码构建 PyTorch。pip 装出来的预编译包没有这些源码目录,改不了 dispatch,也加不了内核。所以第一步是把 pytorch/pytorch 仓库 clone 下来,配好 Xcode 和 CMake。

构建完成后的关键目录都在 aten 下面。dispatch 声明在 aten/src/ATen/native/native_functions.yaml,Metal 内核在 aten/src/ATen/native/mps/kernels/,宿主端 stub 在 aten/src/ATen/native/mps/operations/。这三个目录,就是整个工作流的地盘。

改完代码后的编译命令很简洁,就一条:

cd build && ninja torch_cpu

为什么是 torch_cpu 而不是别的 target,这里有个反直觉的点。MPS 内核的宿主端是 Objective-C++ 写的 .mm 文件,它挂在 ATen 的 dispatch 体系里,跟着 CPU 侧一起编。理解了这一层,后面 dispatch 的改动才说得通。

环境验证也很简单。跑一下 python test/test_mps.py,如果测试框架能正常发现 MPS 用例,说明你的构建环境是通的。很多人卡在这一步之前,其实不是代码问题,是根本没意识到 MPS 开发必须走源码构建这条路。

操作流程

skill 把整个流程收敛成三步,无论新增还是迁移,骨架都一样。第一步改 native_functions.yaml 里的 dispatch,第二步写 Metal 内核,第三步写宿主端 stub。三步走完,编译测试收尾。

dispatch 的改动是这套流程的入口,也是最容易改错的地方。新增算子时,给 dispatch 块加一行 MPS: my_op_mps;迁移时,则是把原本独立的 MPS: my_op_out_mps 删掉,把 MPS 并进共享的那一行,变成 CPU, CUDA, MPS: my_op_out

# 迁移后:MPS 走共享 stub,不再有独立实现
- func: atan2.out(Tensor self, Tensor other, *, Tensor(a!) out) -> Tensor(a!)
  structured: True
  structured_inherits: TensorIteratorBase
  dispatch:
    CPU, CUDA, MPS: atan2_out

这里有个 skill 反复强调的坑。一个算子通常在 yaml 里有多个重载,functional、inplace、.out,还有 Tensor 和 Scalar 变体。每个重载都有独立的 dispatch 块,漏改任何一个,那个重载就会继续走老的 MPSGraph 路径。迁移完成后得 grep 一遍旧的函数名,确认没有条目还引用它。

metal-kernel:把 PyTorch 在 Apple Silicon 上的内核,从 MPSGraph 拉回原生 Metal

写完 dispatch,Metal 内核本身反而简单。逐元素算子的内核就是一个 functor,把运算写在 operator() 里,再用宏注册支持的 dtype。这个设计把迭代器管线和类型分派全藏起来了,你只需要关心那一行真正的数学运算。

struct atan2_functor {
  template <typename T, enable_if_t<is_floating_point_v<T>, bool> = true>
  inline T operator()(const T a, const T b) {
    return static_cast<T>(precise::atan2(float(a), float(b)));
  }
};

REGISTER_FLOAT_BINARY_OP(atan2);
REGISTER_INT2FLOAT_BINARY_OP(atan2);

宿主端 stub 是连接 dispatch 系统和 Metal 内核的那根线。它用一个静态函数把 TensorIterator 交给 lib.exec_binary_kernel,再用 REGISTER_DISPATCH 挂进 dispatch 系统。函数名要和 Metal 里的 functor 名对得上,这是运行时靠字符串查找的约定。

static void atan2_mps_kernel(TensorIteratorBase& iter) {
  lib.exec_binary_kernel(iter, "atan2");
}

REGISTER_DISPATCH(atan2_stub, &atan2_mps_kernel)

关键设计

这套东西真正的价值,藏在 c10/metal 这套基础设施里。它不是随便写的几个头文件,而是一套设计意图很明确的抽象层,把 Metal 开发里最繁琐的部分全封装掉了。

最显眼的是那套注册宏。REGISTER_UNARY_OPREGISTER_BINARY_OPREGISTER_FLOAT_BINARY_OP,每一个都对应一种类型组合模式。比如数学函数用 REGISTER_FLOAT_BINARY_OP 加 REGISTER_INT2FLOAT_BINARY_OP,位运算用 REGISTER_INTEGER_BINARY_OP。选错宏,dtype 覆盖就会漏,或者输出类型对不上。

metal-kernel:把 PyTorch 在 Apple Silicon 上的内核,从 MPSGraph 拉回原生 Metal

类型系统是另一块容易忽略但很关键的拼图。opmath_t<T> 把 half 提升到 float 做中间运算,accum_t<T> 定义归约的累加类型,precise:: 命名空间里的 exp、log、sqrt 则提供高精度版本。这套约定保证了数值正确性,不用每个算子作者自己拍脑袋。

从架构推断,这个 skill 的真实意图,是要终结 MPSGraph 时代留下的技术债。MPSGraph 那套实现分散在各个 UnaryOps.mmBinaryOps.mm 里,每个算子一套 TORCH_IMPL_FUNC,维护成本高,性能上限低。原生 Metal 把实现收敛到统一的 kernel 注册机制,维护性和性能两条线同时改善。

使用场景

迁移是最能体现这套流程价值的场景。拿 atan2 来说,它原本在 MPSGraph 里有个独立实现 atan2_out_mps,迁移的过程就是删掉那个实现,把 dispatch 并进共享 stub,再补一个 Metal functor。改完之后,算子跟 CPU、CUDA 走同一条 dispatch 路径。

skill 还塞了一个很实用的调试技巧,用 torch.mps.compile_shader 对单个内核做 JIT 编译隔离测试。当多内核流水线出问题时,不用整个流程重跑,把每个内核单独拎出来跟 NumPy 参考实现对拍,能快速定位是哪个环节错了。

lib = torch.mps.compile_shader(source)
lib.my_kernel(inp, out, threads=[311], group_size=[311])
torch.mps.synchronize()

这里有个从文档里读出来的反直觉细节。compile_shader 的 threads 参数是总线程数,不是 threadgroup 数。5 个 threadgroup、每组 256 线程,应该传 threads=[1280, 1, 1]。搞反了这个,内核要么跑不满,要么越界。

metal-kernel:把 PyTorch 在 Apple Silicon 上的内核,从 MPSGraph 拉回原生 Metal

大张量场景也值得一提。Metal 内核默认用 32 位索引接收 numel,超过 INT32_MAX 的就得在宿主端拆分。skill 给的解法是 iter.with_32bit_indexing(),让 TensorIterator 自动切子迭代器,递归一层就终止。这个细节处理不好,大张量会静默算错,比直接崩溃更可怕。

那些不走逐元素模式的算子,就得自己驱动 TensorIterator,这里面的坑 skill 也点得很细。第一个是传参要传 TensorIteratorBase& 而不是 Tensor&,后者会丢掉 with_32bit_indexing() 产生的偏移信息。第二个是子迭代器里 iter.tensor(0) 返回的是整个张量,得用 bind_iter_tensors 绑定切片,不然会覆盖同一前缀还留下未初始化的尾部。

这两条规则都属于那种”第一次写必踩”的类型。skill 能把它们单独拎出来写清楚,说明作者真的在这上面栽过跟头。比起泛泛地讲一句”注意迭代器用法”,这种点名到具体 API 的提醒才有实操价值。

洞察与反思

翻完整个 skill,我最大的感受是它把维护者的隐性知识显性化了。这些东西平时散落在 PR 评论和 code review 里,没有谁会专门写下来,现在被整理成了可以照着执行的规则:

  • dispatch 该怎么改,哪些重载必须同步迁移
  • 注册宏该怎么选,不同 dtype 组合对应哪一套
  • threads 和 threadgroup 的区别,为什么参数含义容易搞反
  • 32 位索引的边界,超过之后宿主端怎么拆

这些点单拎出来每一个都不难,难的是把它们串成一条能落地的工作流。

但它的边界同样清晰。这套机制目前最顺滑的是逐元素算子,unary 和 binary 这两类。归约类的 sum、mean 虽然也提到了 ReduceOps.mm,但 skill 的笔墨明显偏轻。真要写一个高性能的归约内核,光靠这个 skill 还不够,得自己去啃 threadgroup 内存的用法。

错误报告那一节尤其让我觉得有点意思。skill 明确禁止为了检查错误把结果复制回 CPU,因为那会强制一次 GPU 同步,拖垮整个流水线。它给出的替代方案是让内核往共享错误缓冲区写消息,宿主在下次同步时抛出 AcceleratorError。这个设计巧就巧在,错误报告搭了现有同步的便车,没有引入额外的停顿。

往大了看,这类官方 skill 是 PyTorch 在把最难复制的领域知识做标准化分发。写 Metal 内核的门槛从来不是 Metal 语言本身,而是那套 dispatch 约定、类型系统、迭代器管线的组合拳。metal-kernel 把组合拳拆成了可以照着打的招式,这比任何教程都值钱。

资源地址

资源 地址
Smithery 页面 https://smithery.ai/skills/pytorch/metal-kernel
PyTorch 源码仓库 https://github.com/pytorch/pytorch

总结

metal-kernel 这个 skill 解决的是一个很具体的问题:怎么在 Apple Silicon 上给 PyTorch 写原生 Metal 内核。三步流程,dispatch 改声明、写 Metal functor、写宿主 stub,再加一套编译测试收尾。

它最值得称道的地方,是没把 MPSGraph 当成既成事实。用原生 Metal 换掉高层图 API,牺牲了一点上手速度,换来了性能和可维护性。从文档的措辞来看,这是 PyTorch 团队有意推动的方向,skill 只是把路线图变成了操作手册。这套路线图的终点,大概率是 MPS 后端跟 CUDA 一样,走同一条 dispatch 到内核的成熟链路。

如果你在维护 MPS 相关的 PyTorch 算子,或者正被某个算子的 MPS 实现折磨,这个 skill 值得认真读一遍。它不能替代你对 Metal 本身的理解,但能替你把那条从 dispatch 到内核的路,踩得明明白白。

skills资源

Aoti-debug:把 AOTI 崩溃排查,从玄学变成一张路由表

2026-8-17 14:00:02

实战分享

接入Seedance 2.0 后的 OiiOii,效果更好了

2026-4-7 12:09:05

0 条回复 A文章作者 M管理员
    暂无讨论,说说你的看法吧