viable/strict/1786468119: [MPS] Add logspace.out (#190495)
PyTorch 的 MPS 后端新增了 aten::logspace.out 支持,通过专用 Metal kernel 计算 base**ramp,修复了 #141287 中 torch.logspace(..., device='mps') 的 NotImplementedError。浮点路径端点精确,复数路径复用 linspace+pow,整数和 bool 类型在 steps==0/1 时遵循 CPU 行为,更大输出时因 float32 精度问题抛出 RuntimeError。测试覆盖浮点和复数类型,移除了过时的 expectedFailure 装饰器。
发展脉络
- 首次出现viable/strict/1786468119: [MPS] Add logspace.out (#190495)PyTorch Core
- 当前判断PyTorch 持续完善 MPS 后端,表明 Apple Silicon 在 AI 训练和推理中的重要性日益增加。此次修复填补了 logspace 算子的空白,减少了开发者在使用 MPS 设备时的摩擦。随着 MPS 后端覆盖率的提升,PyTorch 在 Apple 硬件上的可用性增强,可能吸引更多开发者在 Mac 上进行本地模型开发。这也反映了开源框架在跨平台支持上的竞争,尤其是与 CUDA 生态的差距正在缩小。Agent Pulse · 分析
PyTorch 在 MPS 后端实现了 aten::logspace.out,修复了 #141287 中 torch.logspace 在 MPS 设备上的 NotImplementedError。该实现添加了专用的 Metal kernel(logspace/logspace_strided),在单次启动中计算 base**ramp,支持浮点和复数类型。浮点路径通过 linspace 使用的半分拆分确保端点精确(out[0]==base**start, out[-1]==base**end)。复数类型因 MSL 无复数类型,复用 linspace+pow 路径。steps==0 和 steps==1 的边界情况与 CPU/CUDA 实现一致。整数和 bool 类型在 steps==0/1 时遵循 CPU 行为,但更大输出时 MPS 会抛出 RuntimeError,因为 float32 无法可靠复现 CPU/CUDA 的 float64 到整数截断。测试覆盖浮点和复数类型,验证了 CPU-vs-MPS 在多个 base 值、steps 边界和非连续(strided)输出上的奇偶性,并移除了过时的 expectedFailure 装饰器。
该实现展示了 MPS 后端对算子覆盖的精细化:通过专用 Metal kernel 实现 logspace,避免通用 fallback 的性能损失。浮点路径的端点精确性通过半分拆分实现,与 linspace 一致,表明 PyTorch 在数值一致性上投入了细致工作。复数路径复用 linspace+pow,反映了 Metal Shading Language 缺乏复数类型的限制,以及 PyTorch 在保持功能完整性和实现简洁性之间的权衡。整数类型在 steps>1 时抛出 RuntimeError,而非静默降级,体现了对数值正确性的严格态度。
PyTorch 持续完善 MPS 后端,表明 Apple Silicon 在 AI 训练和推理中的重要性日益增加。此次修复填补了 logspace 算子的空白,减少了开发者在使用 MPS 设备时的摩擦。随着 MPS 后端覆盖率的提升,PyTorch 在 Apple 硬件上的可用性增强,可能吸引更多开发者在 Mac 上进行本地模型开发。这也反映了开源框架在跨平台支持上的竞争,尤其是与 CUDA 生态的差距正在缩小。
对于依赖 PyTorch 在 Apple 硬件上进行开发的企业,此修复降低了在 MPS 设备上运行模型的门槛,减少了因算子缺失导致的开发阻塞。这有助于提升 PyTorch 在 Mac 用户中的采用率,并可能促进 Apple 生态中 AI 应用的开发。对于云服务商,若提供 MPS 实例,此更新可提升其服务兼容性。
后续可关注 MPS 后端对其他缺失算子的支持进度,以及是否会有更多针对 Apple Silicon 的优化,如利用 Metal 的矩阵乘法加速。此外,整数类型在 MPS 上的行为可能在未来版本中改进,例如通过更高精度的中间计算来支持更大的输出。