AGENT PULSESJCPal Special EditionAI 行业证据与趋势
2026年8月31日 · PyTorch

viable/strict/1788205825: [BE] Extend dropout support to complex numbers (#195373)

发生了什么

PyTorch 的 nn.functional.dropout 在 CPU 和 CUDA 上对复数输入抛出 NotImplementedError,而在 MPS 上静默工作。根因是 _dropout_impl 以输入 dtype 分配伯努利掩码并调用 bernoulli_,而 CPU/CUDA 不支持复数 bernoulli_。修复为将掩码分配为实数类型(c10::toRealValueType),使复数 dropout 在所有后端统一工作。

EVENT STORY

发展脉络

  1. 首次出现viable/strict/1788205825: [BE] Extend dropout support to complex numbers (#195373)PyTorch Core
  2. 行业反馈trunk/736985d0de6c552c8fe2324b101cc94bd21b245b: [BE] Extend dropout support to complex numbers (#195373)PyTorch Core
  3. 当前判断PyTorch 对复数 dropout 的支持统一了不同后端的行为,减少了 CPU/CUDA 与 MPS 之间的不一致。这有助于依赖复数运算的领域(如信号处理、量子计算模拟)在 PyTorch 中获得更一致的体验,可能推动这些领域在 PyTorch 上的采用。Agent Pulse · 分析
改变了什么

PyTorch 修复了 dropout 对复数输入的支持问题。此前,nn.functional.dropout 在 CPU 和 CUDA 上对复数输入抛出 NotImplementedError,而在 MPS 上静默工作,但 MPS 的行为是语义未定义的:Metal 内核将 {0,1} 掩码写入实部并将虚部置零。根因是 _dropout_impl 以输入 dtype 分配伯努利掩码并调用 bernoulli_,而 CPU/CUDA 的 AT_DISPATCH_ALL_TYPES_AND3 不支持复数。修复将掩码分配为实数类型(c10::toRealValueType),因为掩码是实值的,复数输入乘以实数掩码是良定义的,且与旧 MPS 行为数值一致。此外,复数输入被排除在 is_fused_kernel_acceptable 之外,因此在 CUDA/XPU/lazy/privateuseone 上走复合路径,与 MPS 一致。修复后,复数 dropout 在所有后端统一工作。

能力边界怎么变了

此修复揭示了 PyTorch 中复合操作(composite ops)与后端内核之间的 dtype 处理差异。掩码分配为输入 dtype 导致对复数输入调用不存在的复数 bernoulli_,而 MPS 的偶然实现掩盖了问题。修复通过将掩码强制为实数类型,确保了跨后端的语义一致性。这提示开发者:在实现涉及随机掩码的操作时,应明确掩码的 dtype 应为实数,而非输入 dtype。

为什么重要

PyTorch 对复数 dropout 的支持统一了不同后端的行为,减少了 CPU/CUDA 与 MPS 之间的不一致。这有助于依赖复数运算的领域(如信号处理、量子计算模拟)在 PyTorch 中获得更一致的体验,可能推动这些领域在 PyTorch 上的采用。

对谁有影响

此修复降低了 PyTorch 在不同硬件后端上的行为差异,提升了框架的可靠性,对依赖 PyTorch 进行科学计算和信号处理的企业用户具有价值,可能减少因后端不一致导致的调试成本。

接下来观察

后续可能进一步清理 MPS 对非浮点输入的排除逻辑,并可能扩展其他操作对复数输入的支持。可关注 PyTorch 是否在后续版本中为更多操作添加复数支持,以及是否引入复数专用的融合内核。