首页
/ PyTorch MPS 后端 Metal Kernel 开发指南:从 native_functions.yaml 分发到 Apple Silicon 算子实现

PyTorch MPS 后端 Metal Kernel 开发指南:从 native_functions.yaml 分发到 Apple Silicon 算子实现

2026-09-06 19:59:00作者:裴麒琰

本文围绕 PyTorch 仓库中的 Metal Kernel 编写技能文档(.claude/skills/metal-kernel/SKILL.md)展开,系统讲解如何在 PyTorch 中为算子实现基于原生 Metal 的 MPS 后端:包括 native_functions.yaml 的分发配置、c10/metal/ 基础设施下的 Kernel 编写范式、host 侧 stub 注册、从 MPSGraph 迁移到原生 Metal 的路径,以及用 torch.mps.compile_shader 调试独立 Kernel 的实战技巧。读完后你能够独立完成"新增一个 MPS 算子"或"将一个 MPSGraph 算子迁移到原生 Metal"的完整工作流,并理解 REGISTER_UNARY_OPREGISTER_DISPATCH 等机制背后的实现原理。

关键原则:本指南的目标是通过 c10/metal/ 基础设施使用原生 Metal 能力,而不是 MPSGraph。原生 Metal Kernel 在控制力、性能与可维护性上都更好,这也是当前 PyTorch MPS 后端的主推方向。

两条工作流:新增算子与迁移算子

该技能文档覆盖两类工作流:

  1. 新增 MPS 支持 —— 从零实现一个新算子;
  2. 从 MPSGraph 迁移 —— 将已有的 MPSGraph 算子转换为原生 Metal。

两条路径都涉及三个核心步骤:

  1. aten/src/ATen/native/native_functions.yaml更新分发
  2. aten/src/ATen/native/mps/kernels/编写 Metal Kernel
  3. aten/src/ATen/native/mps/operations/实现 host 侧 stub

从仓库源码结构看,MPS 后端的目录组织确实如此划分:kernels/ 目录下存放 .metal 着色器文件(如 UnaryKernel.metalBinaryKernel.metalReduceOps.metal),operations/ 目录下存放 Objective-C++ 的 host 侧实现(如 UnaryKernel.mmBinaryKernel.mm)。

第一步:更新 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 时,需要合并(consolidate)分发条目

# 迁移前(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)。仓库中的 atan2.out 条目目前正是这种合并形态(CPU, CUDA, MPS, XPU: atan2_out),可以作为已迁移算子的参照。

必须覆盖所有 overload。 一个算子通常有多个 native_functions.yaml 条目 —— functional / inplace / .out,以及 TensorScalar 变体。每个条目都有独立的 dispatch: 块,必须逐个迁移。任何仍指向 MPS: my_op_mps 的条目都会让该 overload 继续路由到 MPSGraph 代码,调用方命中哪个 overload 就会悄悄落入旧路径。在宣布迁移完成之前,应当用 grep 检索旧函数名,确认没有任何条目仍引用它。

分发命名约定

  • MPS: function_name_mps —— MPS 专属实现(旧 MPSGraph 模式)
  • CPU, CUDA, MPS: function_name —— 共享 stub 实现(原生 Metal 模式)

第二步:实现 Metal Kernel

位置: aten/src/ATen/native/mps/kernels/

一元(Unary)Kernel 范式

// 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 /* your operation */;
  }
};

// 为支持的类型注册
REGISTER_UNARY_OP(my_op, float, float);
REGISTER_UNARY_OP(my_op, half, half);
REGISTER_UNARY_OP(my_op, bfloat, bfloat);

从源码实现看(c10/metal/indexing.h),REGISTER_UNARY_OP(NAME, DTYPE0, DTYPE1) 宏做的事情远不止"注册":

  • 通过 static_assert 校验输出类型必须等于 result_of<NAME_functor, DTYPE0>,在编译期杜绝 functor 返回类型与注册类型不一致的错误;
  • 为同一 functor 实例化多个 kernel 变体,并各自赋予 host_name,例如 <NAME>_dense_<out>_<in><NAME>_dense_scalar_...<NAME>_strided_...<NAME>_inner_contiguous_... 以及跨 dtype 场景的 _castout 变体。host 侧的 exec_unary_kernel 会根据 tensor 的布局(稠密 / 带 stride / 内层连续)自动选择对应 kernel 名,未显式注册的跨 dtype 组合会回退到按输入 dtype 注册的 castout 变体(在输入精度下计算、存储时按用户 dtype 转换,与 CPU 语义对齐);
  • 稠密 kernel 内部还有 ILP(每线程处理 ILP_PER_THREAD 个元素)优化路径,见 unary_dense 的实现:批量加载到 thread-local 数组、展开循环执行 functor、批量写回,提高内存级并行度。

二元(Binary)Kernel 范式

struct my_binary_functor {
  template <typename T>
  inline T operator()(const T a, const T b) {
    return /* your operation */;
  }
};

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 示例(同时支持浮点与整数输入,通过 SFINAE 在 functor 内部做类型分派):

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);

REGISTER_BINARY_OP 宏本身(见 c10/metal/indexing.h#L1137-L1279)同样会实例化一大族 kernel:stridedinner_contiguousstrided_castdensedense_ilp4dense_castdense_broadcast(含 rhs 变体)、dense_scalar(含 lhs 变体)及其 cast 版本。这解释了为什么 host 侧只需调用 lib.exec_binary_kernel(iter, "atan2") 一个入口 —— 布局选择、广播处理、scalar 处理、dtype 转换全部由这套宏生成的 kernel 矩阵覆盖。om_t(opmath type)模板参数用于在输出为 half 时把 float 操作数提升到 float 精度计算,而非降回 half,避免精度损失。

带 Scalar 参数的 Kernel

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);

REGISTER_UNARY_ALPHA_OP(NAME, DTYPEI, DTYPEA, DTYPEO) 会生成 unary_alpha_dense(以及 ILP 变体)与 unary_alpha_strided 等 kernel,其中 alpha 作为 constant T2& 常量缓冲区参数传入(见 c10/metal/indexing.h#L398-L431)。

类型特化的 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(/* real */, /* imag */);
  }
};

复数类型说明:Metal 中复数以向量类型表示 —— c10::complex<float> 映射为 float2(x = 实部,y = 虚部),c10::complex<half> 映射为 half2。在 functor 中可用 is_complex_v<T> 对复数类型做特化。

可用的 c10/metal 工具库

  • utils.hc10/metal/utils.h):opmath_t<T>(运算数学类型,half->float)、accum_t<T>(归约的累加类型)、带 NaN 传播语义的 max()min()
  • special_math.hc10/metal/special_math.h):precise::exp()precise::log()precise::sqrt()precise::sin()precise::cos()precise::tan()erf()erfc()erfinv()
  • indexing.hc10/metal/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),以及索引原语 pos_from_thread_index / offset_from_coord 等。

第三步:实现 host 侧 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 模式的结构化 Kernel:

// 在 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)

仓库中的真实实现印证了这一模式。例如 UnaryKernel.mm 用一个宏批量生成一元算子的 stub:

#define REGISTER_UNARY_TI_DISPATCH(NAME)                    \
  static void NAME##_kernel_mps(TensorIteratorBase& iter) { \
    lib.exec_unary_kernel(iter, #NAME);                     \
  }                                                         \
  REGISTER_DISPATCH(NAME##_stub, NAME##_kernel_mps)

随后以 REGISTER_UNARY_TI_DISPATCH(exp); REGISTER_UNARY_TI_DISPATCH(tanh); ... 的形式一行注册一个算子。而 BinaryKernel.mmatan2 的 stub 则更简洁直接:

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

这里有一个关键的对应关系:native_functions.yamlCPU, CUDA, MPS: atan2_out 声明的函数名(去掉 MPS 行前缀后的 atan2_out)由 torchgen 生成 atan2_stub 分发点,MPS 端只需 REGISTER_DISPATCH(atan2_stub, ...) 把它绑定到 Metal 执行函数即可 —— 这就是"共享 stub 分发"能省掉独立 *_mps 实现的原理。

迁移时删除旧的 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

测试:接入 OpInfo 测试体系

算子的基础正确性已由 test/test_mps.py 中的 test_output_match 测试覆盖。实现完一个算子后,通过移除预期失败项即可启用测试。

1. 从 common_mps.py 中移除条目

位置: torch/testing/_internal/common_mps.py

找到并删除算子在 skip/xfail 列表中的条目(当前仓库中该列表名为 ON_MPS_XFAILLIST,按 dtype 组织预期失败):

# 删除形如:
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 可用于对单个 Metal Kernel 做 JIT 编译与独立测试,是调试多 Kernel 流水线时逐个验证每一阶段的利器。该接口实现在 torch/mps/__init__.py

基本用法

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)语义

compile_shader 使用 dispatchThreads 语义(与 PyTorch 中的 mtl_dispatch1DJob 相同):

  • threads=[N, 1, 1] —— 线程总数(不是 threadgroup 数);
  • group_size=[G, 1, 1] —— 每个 threadgroup 内的线程数。

这与部分 host 侧代码使用的 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])

常量缓冲区参数

标量常量以单元素张量传递:

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(例如 histogram → prefix_sum → scatter)输出错误结果时,逐个单独测试每个 kernel,并与 Python/NumPy 参考结果比对:

# 1. 运行 GPU kernel
lib.histogram(keys, hist, ..., threads=[N, 1, 1], group_size=[256, 1, 1])
torch.mps.synchronize()

# 2. 用 Python 计算参考结果
ref_hist = compute_histogram_cpu(keys.cpu().numpy(), ...)

# 3. 比对
assert np.array_equal(hist.cpu().numpy(), ref_hist), "Histogram mismatch!"

这样可以定位流水线中到底是哪个 kernel 出了问题,而不是整条流水线一起调试。

常见陷阱

  • threads 数量搞错 —— threads 是线程总数,不是 threadgroup 数。5 个 256 线程的 threadgroup 应写 threads=[1280, 1, 1]
  • Threadgroup 内存 —— compile_shader 不直接支持 [[threadgroup(N)]] 参数。若 kernel 需要 threadgroup 内存,改为在 kernel 函数体内声明 threadgroup 数组。

与 TensorIterator 协作的要点

REGISTER_UNARY_OP / REGISTER_BINARY_OP 隐藏了 iterator 的管线细节。带额外参数或非逐元素布局的 kernel 必须直接驱动 TensorIterator,其中几条不明显的规则值得注意:

  • TensorIteratorBase&,而不是 Tensor& Tensor& 会丢失 with_32bit_indexing() 产生的 offset/shape 信息。当 stub 交给你 const TensorBase&(例如 bernoulli_scalar_stub)时,用 at::TensorIterator::borrowing_nullary_op(self) 就地构建 iter,而不是做 const 转换;

  • 对 sub-iter,iter.tensor(0) 返回的是整个 tensor。 with_32bit_indexing() 之后,sub-iter 仍引用原始 storage,直接绑定 iter.tensor(0) 会让每个 sub-iter 覆写同一段前缀、尾部留空。应使用 bind_iter_tensors,它按每块的 offset 计算并把 buffer 0 绑定到切片上;分发用 iter.numel(),绝不用 iter.tensor(0).numel()

    bind_iter_tensors(computeEncoder, iter, /*ntensors=*/1);
    mtl_setArgs<1>(computeEncoder, params, ..., numel);
    mtl_dispatch1DJob(computeEncoder, pso, threads);
    
  • 优先用 mtl_setArgs<N>,而不是链式 mtl_setBytes mtl_setArgs<1>(encoder, a, b, c) 会在 slot 1/2/3 依次绑定 a/b/c,重载解析规则与宏一致(例如 std::array<long,2> 绑定为 constant long2&);

  • 使用 ceil_div host 侧用 ATen/ceil_div.hat::ceil_div;Metal 侧用 c10/metal/common.hc10::metal::ceil_divusing namespace c10::metal; 之后无需加限定)。两者都是 (a + b - 1) / b 的封装;

  • Metal 3/4 可移植性用 IF_CONSTEXPR Metal 4 有 if constexpr,Metal 3 没有。使用 c10/metal/common.h 中的宏,例如 if IF_CONSTEXPR (sizeof(T) == 8) { ... }

  • 通过 REGISTER_MPS_DISPATCH 共享 CPU/CUDA stub。 当算子在上游已有 DECLARE_DISPATCH stub(distributions、fused ops 等)时,用 REGISTER_MPS_DISPATCH(stub_name, &fn) 把 MPS 接入,而无需在 native_functions.yaml 中为 MPS 单开条目。stub 本身就接收 TensorIteratorBase&

处理超大 Tensor

大多数 Metal kernel 以 uint32_t 接收 numel、以 32 位索引寻址,因此超过 INT32_MAX 的元素数需要在 host 侧切分。这个切分应通过 TensorIterator 驱动,而不是手工切片。

通过 iter.with_32bit_indexing() 分解:

if (!iter.can_use_32bit_indexing()) {
  for (auto&& sub_iter : iter.with_32bit_indexing()) {
    my_kernel_impl(sub_iter, ...);
  }
  return;
}

每个产出的 sub-iter 都满足 can_use_32bit_indexing(),因此递归一层即终止。仓库中的 frexp_kernel_mpsUnaryKernel.mm)就是这一模式的实际用例。注意阈值是 INT32_MAX(TensorIterator 使用有符号 32 位 offset),不是 UINT32_MAX —— 略超过 INT32_MAX 的测试可以覆盖切分路径,但不覆盖 uint32_t 转换本身;转换只在 numel > 2^32 时才真正影响结果。

收窄类型时用带检查的转换:

const uint32_t numel = c10::checked_convert<uint32_t>(iter.numel(), "uint32_t");

使用 c10::checked_convert<c10/util/TypeCast.h>),让回绕(wraparound)变成 TORCH_CHECK 报错,而不是输出静默损坏。

Kernel 中的错误上报

绝不要为了检查错误而把结果拷回 CPU。 对 GPU tensor 的任何 .item().cpu() 或其他 host 侧读取都会强制一次完整的 GPU→CPU 同步,排空流上所有在途操作 —— 而不仅仅是你正在检查的那个归约。在真实流水线中这会卡住整个队列,代价远超检查本身。用 is_mps() 守护这次同步也没用;每次算子运行都会产生停顿。正确做法是在设备端做校验,并通过下面的机制上报错误。

GPU 代码无法抛异常,但 kernel 可以把错误写入共享的错误缓冲区,由 host 在下次同步时以 c10::AcceleratorError 抛出。使用 c10/metal/error.h 中的 TORCH_REPORT_ERROR(error_buf, ...):变参会被拼接成消息,整数按十进制格式化。投递是异步的 —— 出错线程继续运行,错误在 MPSStream::checkLastError() 下次运行时才暴露(synchronize() 之后或下一个排空 stream 的操作之后),因此不要依赖它做同一次 dispatch 内的控制流。关键优势是不引入任何强制同步:错误顺路搭乘用户代码本来就有的同步。

Kernel 侧:接收 device ErrorMessages* error_buf 参数,在错误路径上调用 TORCH_REPORT_ERROR,随后跳过出错元素,避免同时破坏内存:

#include <c10/metal/error.h>

kernel void index_set_1d(
    device float* self,
    constant float* values,
    constant long* indices,
    constant uint& self_numel,
    device ::c10::metal::ErrorMessages* error_buf,
    uint tid [[thread_position_in_grid]]) {
  long idx = indices[tid];
  if (idx < 0 || idx >= long(self_numel)) {
    TORCH_REPORT_ERROR(
        error_buf, "index ", idx, " out of bounds for size ", long(self_numel));
    return;
  }
  self[idx] = values[tid];
}

Host 侧:把当前 stream 的错误缓冲区绑定到对应参数位。mtl_setArgs 可直接接收它:

auto* stream = getCurrentMPSStream();
mtl_setArgs(encoder, self, values, indices, uint32_t(self.numel()),
            stream->getErrorBuffer());

缓冲区由 MPSStream 持有,容量为 30 条消息,每次 checkLastError() 排空后复位。只有第一条消息会上报;后续消息保留用于调试,但 AcceleratorError 只携带 msg[0]

完成检查清单

  • [ ] 已在 native_functions.yaml 中添加 MPS 分发
  • [ ] 已在 kernels/ 中实现 Metal Kernel
  • [ ] 已在 operations/ 中实现 host 侧算子
  • [ ] 能处理空 tensor
  • [ ] 能处理非连续 tensor
  • [ ] 支持所需 dtype(float32、float16、bfloat16,通常还需经 float2/half2 支持复数类型)
  • [ ] 已从 torch/testing/_internal/common_mps.py 中移除预期失败项
  • [ ] 已移除 OpInfo 的 skip/xfail 装饰器(如适用)

小结

PyTorch 的 MPS 后端为算子实现建立了一套"分发声明 + functor + stub 注册"的三层结构:native_functions.yaml 声明分发关系(迁移时以 CPU, CUDA, MPS 共享 stub 行收敛旧式的 MPS: xxx_mps 独立条目),c10/metal/indexing.hREGISTER_* 宏在编译期把单个 functor 展开为覆盖稠密/带 stride/内层连续/广播/标量/跨 dtype 的完整 kernel 矩阵,operations/*.mm 中的 REGISTER_DISPATCH 则把 stub 与 MetalShaderLibrary 的执行入口(exec_unary_kernel / exec_binary_kernel)接通。配合 with_32bit_indexing() 的超大 tensor 切分、TORCH_REPORT_ERROR 的异步设备端错误上报,以及 torch.mps.compile_shader 的独立调试能力,这套机制既能复用 CPU/CUDA 的 TensorIterator 语义,又保持了原生 Metal 的性能与可控性。

登录后查看全文
热门项目推荐
相关项目推荐