viable/strict/1790368010: [MPS] Port norm to the shared reduction kernels (#198611)
PyTorch 合并 PR #198611,将 MPS 后端的 norm / linalg.vector_norm 从独立 Metal kernel(每个输出一个 threadgroup,完整归约在单个 threadgroup 上执行)迁移到共享归约 kernel。证据给出的基准显示多种形状与 dim 组合下 bf16 延迟下降,例如 [4096,4096] 从 17608us 降至 104us(170x)、[1024,1024] 从 1009us 降至 8.0us(126x)、[1000000,4] dim=0 从 989us 降至 27.9us(35.5x);[8192,1024] dim=-1 基本持平(1.00x),[8192,1024] dim=0 三线程组场景为 0.97x。
发展脉络
- 首次出现viable/strict/1790368010: [MPS] Port norm to the shared reduction kernels (#198611)PyTorch Core
- 当前判断这是开源框架内部的算子实现收敛,而非新模型或新硬件发布。它反映 Apple 芯片后端在 PyTorch 中持续补齐与 CUDA 侧一致的归约基础设施;对生态的影响取决于 MPS 训练与推理工作负载中 norm 类算子的实际占比,证据未提供端到端模型级数据。Agent Pulse · 分析
PyTorch 核心仓库合并 PR #198611,把 MPS 上的 norm / linalg.vector_norm 从独立 Metal kernel 移植到共享归约 kernel。原实现每个输出只用一个 threadgroup,因此完整归约被限制在单个 threadgroup 内;新实现复用共享归约路径。证据附带的 bf16 基准覆盖多种形状与归约维度:大张量归约收益显著,如 [4096,4096] 由 17608us 降至 104us、[1024,1024] 由 1009us 降至 8.0us、[1000000,4] dim=0 由 989us 降至 27.9us;但 [8192,1024] dim=-1 为 1.00x,[8192,1024] dim=0 三线程组场景为 0.97x,说明收益依赖形状与归约轴。
从证据看,性能差异主要来自并行度模型:单 threadgroup 完成完整归约限制了可用的并行资源,而共享归约 kernel 允许跨 threadgroup 拆分归约。可验证的下一信号是 PyTorch 是否把同一共享归约路径扩展到 MPS 上其他归约类算子(如 sum、mean、var),以及是否出现针对小输出/大归约轴场景的回归报告。
这是开源框架内部的算子实现收敛,而非新模型或新硬件发布。它反映 Apple 芯片后端在 PyTorch 中持续补齐与 CUDA 侧一致的归约基础设施;对生态的影响取决于 MPS 训练与推理工作负载中 norm 类算子的实际占比,证据未提供端到端模型级数据。
对在 Apple 芯片上跑 PyTorch 的团队,norm 类算子延迟在部分形状上大幅下降,可能降低本地训练/推理的算子级开销;但收益高度依赖形状与归约轴,且证据只覆盖 bf16 微基准,不能直接外推为端到端成本节省。
若共享归约 kernel 成为 MPS 归约算子的统一路径,后续版本可能看到更多算子迁移与更一致的跨形状性能。需要观察的下一个信号是后续 release note 中是否列出同类移植,以及社区是否报告特定形状(如 dim=-1 大张量)的退化。