arXivDaily arXiv每日学术速递 周一至周五更新
arXiv周末暂无论文更新,休息一下吧,周末愉快~~

$2^{32}$ 边界处的静默失败:关于 PyTorch 的 Apple MPS 后端中大张量矩阵乘法的技术报告

Silent Failures Beyond the 32-Bit Index Range: A Differential Characterization of Large-Tensor Matrix Multiplication in PyTorch's MPS Backend

Junichiro Niimi

arXiv 2609.22991首次发表:更新:

发表机构

Meijo University(名城大学)

机构由 AI 辅助整理,请以论文原文为准。

AI 中文总结

本报告发现 PyTorch MPS 后端在张量元素数超过 $2^{32}$ 时进行批量矩阵乘法会静默返回错误结果,并总结了错误规则及防护措施。

AI 中文摘要

配备 192GB 或更大统一内存的 Apple Silicon 机器使得在桌面 GPU 上放置包含超过 $2^{32}$ 个元素的张量成为常规操作。我们证明,PyTorch 的 Metal Performance Shaders (MPS) 后端在此规模下进行批量矩阵乘法时会静默返回错误结果。在 macOS 27.0 上,此 http URL,以及因此此 http URL 和 eager attention,在我们测试的从 2.4.1 到 2.14.0 的每个 PyTorch 版本中,都会返回相对误差大于 1 的结果,且不抛出异常或警告。在一台机器上,我们针对两种数据类型、四种内存布局、六种形状和 42 个介于 4096 和 65538 之间的批量大小对 bmm 进行了扫描(在 PyTorch 2.14.0 上进行了 1584 次运行,并在十个早期版本上进行了缩减扫描),并将每个结果与 CPU 上的 float64 计算进行对比。三条规则解释了 2.14.0 上的所有结果。当输出超过 $2^{32}$ 个元素且操作数是转置视图时,整个输出都是错误的,并且等于忽略该操作数步幅的计算结果。当连续输入超过 $2^{32}$ 个元素时,只有超过该点的批次是错误的,并且它们等于索引在 $2^{32}$ 处回绕的计算结果。具有至少 $2^{31}$ 个元素的视图操作数会抛出异常,因此更大的问题可以将显式错误转变为静默失败。在 CUDA 上,bmm 的对照是正确的,尽管此 http URL 在超过 $2^{32}$ 个元素时也会静默出错。在一个公开的情感分类器中,一个过大的批次破坏了三分之一的输出,并将它们折叠到单一类别上。该研究是黑盒式的:我们报告后端返回的结果,并与参考结果进行比较。我们在此 https URL 发布了扫描工具、原始结果以及一个阻止任何 MPS 操作触及 $2^{32}$ 或更多元素的防护措施。

英文摘要

Apple Silicon machines with large unified memory make it possible to hold large tensors on a desktop GPU. However, we found that PyTorch's Metal Performance Shaders (MPS) backend silently returns wrong results for batched matrix multiplication with more than $2^{32}$ elements. torch bmm, including its wrappers matmul and eager attention, returns relative errors above 1 without an exception or a warning in every PyTorch release tested (2.4.1 to 2.14.0). We sweep bmm over dtypes, memory layouts, shapes and batch sizes around $2^{31}$ and $2^{32}$ elements, and judge every result against a float64 computation on the CPU. Three rules account for every outcome on 2.14.0. When the output exceeds $2^{32}$ elements and an operand is a transposed view, the entire output is wrong and equals a computation that ignores that operand's strides. Otherwise, a view with at least $2^{31}$ elements raises an exception, and a contiguous input above $2^{32}$ elements makes exactly the batches beyond that point wrong, equal to a computation whose index wraps at $2^{32}$. A slightly larger problem can thus turn an explicit error into a silent failure. The rules extend to the backward pass, where a correct forward pass can return silently wrong gradients. A second machine with another chip, under two macOS versions, reproduces all 6156 results, including the wrong values, and the same sweeps on an NVIDIA A100 are correct in all 2530 runs. In a public sentiment classifier, one oversized batch corrupts a third of the outputs, which collapse onto one class. All findings come from observable behavior, without access to the backend's closed-source kernels; we release the harness, raw results and a guard that stops any MPS operation touching $2^{32}$ or more elements at jniimi/mps-silent-failures (https://github.com/jniimi/mps-silent-failures).

论文原文

arXiv 摘要页 · PDF 原文 · HTML 原文

↑