为 PyTorch 算子编写 Metal/MPS kernel。适用于给算子添加 MPS 设备支持、实现 Metal 着色器或将 CUDA kernel 移植到 Apple Silicon 的场景。涵盖 native_functions.yaml 分发配置、主机端算子实现以及 Metal kernel 的编写。
Metal Kernel 编写指南
本 Skill 将引导你如何在 Apple Silicon 上为 PyTorch 算子实现原生 Metal kernel。
重要提示: 本 Skill 的核心目标是利用 c10/metal/ 基础设施来发挥原生 Metal 的能力,而不是使用 MPSGraph。原生 Metal kernel 能带来更好的掌控度、性能和可维护性。
整体流程 (Overview)
本 Skill 涵盖两种常用工作流:
- 添加全新的 MPS 支持 — 从零开始实现一个新算子
- 从 MPSGraph 迁移 — 将基于 MPSGraph 的旧算子改写为原生 Metal 实现
两种工作流均需包含以下步骤:
- 在
aten/src/ATen/native/native_functions.yaml中更新分发配置(dispatch) - 在
aten/src/ATen/native/mps/kernels/中编写 Metal kernel - 在
aten/src/ATen/native/mps/operations/中实现主机端 stub 存根
第一步:更新 native_functions.yaml
位置: aten/src/ATen/native/native_functions.yaml
针对全新算子
在文件中找到对应的算子条目并添加 MPS 分发项:
# 简单的 MPS 专用实现
- func: my_op(Tensor self) -> Tensor
dispatch:
CPU: my_op_cpu
CUDA: my_op_cuda
MPS: my_op_mps
# 跨设备共享实现(结构化 kernel 的推荐方式)
- func: my_op.out(Tensor self, *, Tensor(a!) out) -> Tensor(a!)
dispatch:
CPU, CUDA, MPS: my_op_out
# 结构化 kernel(推荐用于新算子)
- func: my_op.out(Tensor self, *, Tensor(a!) out) -> Tensor(a!)
structured: True
structured_inherits: TensorIteratorBase
dispatch:
CPU, CUDA, MPS: my_op_out
针对从 MPSGraph 迁移的算子
当把现有算子从 MPSGraph 迁移到原生 Metal 时,需要合并分发配置条目:
# 迁移前(基于 MPSGraph,独立分发)
- func: atan2.out(Tensor self, Tensor other, *, Tensor(a!) out) -> Tensor(a!)
structured: True
structured_inherits: TensorIteratorBase
dispatch:
CPU, CUDA: atan2_out
MPS: atan2_out_mps # 独立的 MPS 实现
# 迁移后(原生 Metal,通过 stub 共享分发)
- func: atan2.out(Tensor self, Tensor other, *, Tensor(a!) out) -> Tensor(a!)
structured: True
structured_inherits: TensorIteratorBase
dispatch:
CPU, CUDA, MPS: atan2_out # MPS 现在复用同样的 stub 机制
核心变动: 将原来的 MPS: my_op_out_mps 替换掉,直接在共享分发行里加上 MPS(例如 CPU, CUDA, MPS: my_op_out)。
务必更新每一个重载(overload)。 单个算子通常在 native_functions.yaml 中包含多个条目 — 比如函数式、inplace 以及 .out 形式,再加上 Tensor 与 Scalar 变体。每个条目都有独立的 dispatch: 块,且都需要同步更新。如果有任何一个条目仍遗留为 MPS: my_op_mps,调用方在命中该特定重载时就会悄无声息地走到老旧的 MPSGraph 路径上。在确认迁移完成之前,请 grep 查一下旧函数名,确保没有任何条目还在引用它。
分发命名约定:
MPS: function_name_mps— 针对 MPS 独有的实现(旧版 MPSGraph 模式)CPU, CUDA, MPS: function_name— 共享 stub 的实现方式(原生 Metal 模式)
第二步:编写 Metal Kernel
位置: aten/src/ATen/native/mps/kernels/
一元 Kernel 模式 (Unary Kernel Pattern)
// MyKernel.metal
#include <c10/metal/indexing.h>
#include <c10/metal/utils.h>
#include <metal_stdlib>
using namespace metal;
using namespace c10::metal;
// 定义操作 functor 仿函数
struct my_op_functor {
template <typename T>
inline T operator()(const T x) {
return /* 具体操作逻辑 */;
}
};
// 注册支持的数据类型
REGISTER_UNARY_OP(my_op, float, float);
REGISTER_UNARY_OP(my_op, half, half);
REGISTER_UNARY_OP(my_op, bfloat, bfloat);
二元 Kernel 模式 (Binary Kernel Pattern)
struct my_binary_functor {
template <typename T>
inline T operator()(const T a, const T b) {
return /* 具体操作逻辑 */;
}
};
REGISTER_BINARY_OP(my_binary, float, float);
REGISTER_BINARY_OP(my_binary, half, half);
二元 Kernel 类型注册宏
对于二元运算,推荐使用 BinaryKernel.metal 中定义的便捷宏:
// 仅浮点类型(float, half, bfloat)
REGISTER_FLOAT_BINARY_OP(my_op);
// 整型输入、浮点输出(适用于 atan2、copysign 等数学算子)
// 注册类型:long->float, int->float, short->float, uchar->float, char->float, bool->float
REGISTER_INT2FLOAT_BINARY_OP(my_op);
// 整型输入、同类型输出(适用于按位/逻辑运算)
// 注册类型:long, int, short, uchar, char, bool
REGISTER_INTEGER_BINARY_OP(my_op);
// 带 opmath 精度的浮点运算(适用于需要更高精度的算子)
REGISTER_OPMATH_FLOAT_BINARY_OP(my_op);
常见搭配套路:
- 数学函数(atan2, copysign, logaddexp):同时使用
REGISTER_FLOAT_BINARY_OP和REGISTER_INT2FLOAT_BINARY_OP - 比较与逻辑算子(maximum, minimum):同时使用
REGISTER_FLOAT_BINARY_OP和REGISTER_INTEGER_BINARY_OP - 算术算子(add, sub, mul):同时使用
REGISTER_FLOAT_BINARY_OP和REGISTER_INTEGER_BINARY_OP
以 atan2 为例(同时支持浮点与整型输入):
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)));
}
template <typename T, enable_if_t<is_integral_v<T>, bool> = true>
inline float operator()(const T a, const T b) {
return precise::atan2(float(a), float(b));
}
};
REGISTER_FLOAT_BINARY_OP(atan2);
REGISTER_INT2FLOAT_BINARY_OP(atan2);
包含 Scalar 参数
struct my_alpha_functor {
template <typename T>
inline T operator()(const T a, const T b, const T alpha) {
return a + c10::metal::mul(alpha, b);
}
};
REGISTER_UNARY_ALPHA_OP(my_alpha, float, float, float);
REGISTER_UNARY_ALPHA_OP(my_alpha, half, half, half);
按类型特化的 Functor
struct special_functor {
// 浮点类型
template <typename T, enable_if_t<is_scalar_floating_point_v<T>, bool> = true>
inline T operator()(const T x) {
return precise::exp(x); // 使用高精度数学运算
}
// 整型
template <typename T, enable_if_t<is_scalar_integral_v<T>, bool> = true>
inline float operator()(const T x) {
return precise::exp(float(x));
}
// 复数类型 (cfloat 对应 float2,chalf 对应 half2)
template <typename T, enable_if_t<is_complex_v<T>, bool> = true>
inline T operator()(const T x) {
// x.x = 实部, x.y = 虚部
return T(/* 实部 */, /* 虚部 */);
}
};
关于复数的说明: Metal 中的复数是用向量类型表示的:
c10::complex<float>映射为float2(x = 实部,y = 虚部)c10::complex<half>映射为half2
在 functor 中可以使用 is_complex_v<T> 专门针对复数类型做特化。
可用的 c10/metal 工具库
utils.h:
opmath_t<T>— 运算数学精度类型(如 half->float)accum_t<T>— 规约(reduction)累加类型- 支持 NaN 传播的
max(),min()
special_math.h:
precise::exp(),precise::log(),precise::sqrt()precise::sin(),precise::cos(),precise::tan()erf(),erfc(),erfinv()
indexing.h:
REGISTER_UNARY_OP(name, in_type, out_type)REGISTER_BINARY_OP(name, in_type, out_type)REGISTER_UNARY_ALPHA_OP(name, in_type, alpha_type, out_type)
第三步:实现主机端 Stub
位置: aten/src/ATen/native/mps/operations/
根据算子类型选择或新建对应的文件:
UnaryKernel.mm— 通过 stub 分发的一元算子BinaryKernel.mm— 通过 stub 分发的二元算子UnaryOps.mm/BinaryOps.mm— 旧版的 MPSGraph 实现(供参考)ReduceOps.mm— 规约类算子(sum, mean, max 等)- 若为全新类别的算子,可以单独新建文件
Stub 注册模式(原生 Metal 推荐做法)
适用于使用 TensorIterator 模式的结构化算子(structured kernels):
// 在 BinaryKernel.mm(或对应的文件中)
static void my_op_mps_kernel(TensorIteratorBase& iter) {
lib.exec_binary_kernel(iter, "my_op"); // "my_op" 需与 .metal 中的 functor 名称保持一致
}
// 注册 MPS stub — 将其绑定到分发系统中
REGISTER_DISPATCH(my_op_stub, &my_op_mps_kernel)
一元算子写法:
static void my_unary_mps_kernel(TensorIteratorBase& iter) {
lib.exec_unary_kernel(iter, "my_unary");
}
REGISTER_DISPATCH(my_unary_stub, &my_unary_mps_kernel)
迁移步骤:清理旧版 MPSGraph 代码
在从 MPSGraph 迁移时,记得清理废弃的代码:
-
从 BinaryOps.mm(或 UnaryOps.mm)中删除:
- 删除
TORCH_IMPL_FUNC(my_op_out_mps)的具体实现 - 移除对应的
#include <ATen/ops/my_op_native.h>头文件引用
- 删除
-
添加到 BinaryKernel.mm(或 UnaryKernel.mm):
- 添加静态 kernel 函数
- 添加
REGISTER_DISPATCH调用
第四步:编译验证
修改完成后,执行编译以确认构建无误:
cd build && ninja torch_cpu
测试验证
基本算子支持已由 test/test_mps.py 中的 test_output_match 覆盖。完成算子实现后,移除预期失败标识以开启测试:
1. 从 common_mps.py 中移除
位置: torch/testing/_internal/common_mps.py
找到该算子并将其从跳过/预期失败列表中删掉:
# 移除形如以下内容的配置:
MPS_XFAILLIST = {
"my_op": ..., # 删除此行
}
MPS_SKIPLIST = {
"my_op": ..., # 删除此行
}
2. 从 OpInfo 装饰器中移除
位置: torch/testing/_internal/common_methods_invocations.py(或相关文件)
从 OpInfo 中移除针对 MPS 的装饰器:
OpInfo(
"my_op",
# 移除如下装饰器:
# decorators=[skipMPS, expectedFailureMPS("reason")],
...
)
3. 运行测试进行验证
# 运行指定算子的测试
python test/test_mps.py -k test_output_match_my_op
# 或运行完整的 MPS 测试集
python test/test_mps.py
使用 torch.mps.compile_shader 调试 Metal Kernel
可以通过 torch.mps.compile_shader 逐个 JIT 编译并单独测试 Metal kernel。在调试包含多个 kernel 的流水线(pipeline)时,该功能在独立验证每个阶段的逻辑时非常实用。
基本用法
import torch
source = '''
#include <metal_stdlib>
using namespace metal;
kernel void my_kernel(
const device float* input [[buffer(0)]],
device float* output [[buffer(1)]],
uint tid [[thread_position_in_grid]]) {
output[tid] = input[tid] * 2.0;
}
'''
lib = torch.mps.compile_shader(source)
inp = torch.tensor([1.0, 2.0, 3.0], device='mps')
out = torch.zeros(3, device='mps')
lib.my_kernel(inp, out, threads=[3, 1, 1], group_size=[3, 1, 1])
torch.mps.synchronize()
print(out) # tensor([2., 4., 6.], device='mps:0')
分发语义 (Dispatch Semantics)
compile_shader 使用的是 dispatchThreads 语义(与 PyTorch 内部的 mtl_dispatch1DJob 一致):
threads=[N, 1, 1]— 总线程数(不是 线程组数量)group_size=[G, 1, 1]— 每个线程组包含的线程数
这与部分主机端代码使用的 dispatchThreadgroups API 不同。如果要对标 dispatchThreadgroups:MTLSizeMake(num_tgs, num_slices, 1) threadsPerThreadgroup:MTLSizeMake(TG_SIZE, 1, 1):
# 等价的 compile_shader 调用方式:
lib.kernel(args...,
threads=[num_tgs * TG_SIZE, num_slices, 1],
group_size=[TG_SIZE, 1, 1])
常量缓冲区参数 (Constant Buffer Parameters)
将标量常量以单元素 Tensor 的形式传入:
slice_size = torch.tensor([1024], dtype=torch.int32, device='mps')
lib.my_kernel(data, output, slice_size, threads=[1024, 1, 1], group_size=[256, 1, 1])
多 Kernel 流水线的调试策略
当处理涉及多个 kernel 的流水线(如直方图 → 前缀和 → 散布写 scatter)时...
<!-- 翻译批次已截断;完整正文同源文件 -->






