metal-kernel

metal-kernel

热门

为 PyTorch 算子编写 Metal/MPS kernel。适用于给算子添加 MPS 设备支持、实现 Metal 着色器或将 CUDA kernel 移植到 Apple Silicon 的场景。涵盖 native_functions.yaml 分发配置、主机端算子实现以及 Metal kernel 的编写。

10万Star
2.9万Fork
更新于 2026/8/4
SKILL.md
只读
名称
metal-kernel
描述

为 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 涵盖两种常用工作流:

  1. 添加全新的 MPS 支持 — 从零开始实现一个新算子
  2. 从 MPSGraph 迁移 — 将基于 MPSGraph 的旧算子改写为原生 Metal 实现

两种工作流均需包含以下步骤:

  1. aten/src/ATen/native/native_functions.yaml更新分发配置(dispatch)
  2. aten/src/ATen/native/mps/kernels/编写 Metal kernel
  3. 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 形式,再加上 TensorScalar 变体。每个条目都有独立的 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_OPREGISTER_INT2FLOAT_BINARY_OP
  • 比较与逻辑算子(maximum, minimum):同时使用 REGISTER_FLOAT_BINARY_OPREGISTER_INTEGER_BINARY_OP
  • 算术算子(add, sub, mul):同时使用 REGISTER_FLOAT_BINARY_OPREGISTER_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 迁移时,记得清理废弃的代码:

  1. BinaryOps.mm(或 UnaryOps.mm)中删除:

    • 删除 TORCH_IMPL_FUNC(my_op_out_mps) 的具体实现
    • 移除对应的 #include <ATen/ops/my_op_native.h> 头文件引用
  2. 添加到 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)时...

<!-- 翻译批次已截断;完整正文同源文件 -->