viable/strict/1787598738: [sdpa] Fix bottom-right efficient attention when Q exceeds K (#193632)
PyTorch 发布修复,针对 CUDA 内存高效注意力中当查询序列长度超过键序列长度时,底部右侧因果对角线偏移被存储为无符号值导致掩码失效的问题。修复将偏移存储为有符号值,并在可能包含完全掩码行时对输出和查询梯度进行零初始化。
Development
- First Reportviable/strict/1787598738: [sdpa] Fix bottom-right efficient attention when Q exceeds K (#193632)PyTorch Core
- Current AssessmentPyTorch 作为主流深度学习框架,其注意力实现的修复反映了 LLM 时代对掩码语义的重新思考。向底部右侧掩码的转变可能影响模型架构和推理优化,尤其是在处理长序列时。此修复可能推动框架间的一致性,并影响依赖高效注意力的下游库和模型。Agent Pulse · analysis
PyTorch 在 2026 年 8 月 24 日发布了一个修复,针对 CUDA 内存高效注意力(memory-efficient attention)在查询序列长度超过键序列长度时的边缘情况。此前,底部右侧因果对角线偏移被存储为无符号值,导致负偏移回绕并禁用混合查询块中的掩码。此外,完全掩码的查询块被跳过,未初始化前向输出或查询梯度。修复将偏移存储为有符号值,将每个块的键计数限制在零,并在底部右侧掩码可能包含完全掩码行时对输出和查询梯度进行零初始化。该修复定义了这些行在密集和打包可变长度注意力中为零,同时保留所有查询至少有一个有效键的形状的空分配路径。此更改是更大计划的一部分,旨在将 SDPA 和 varlen 对齐到右下角掩码语义,以更好地适应 LLM 时代。
此修复解决了 CUDA 内存高效注意力中的整数溢出问题,该问题在查询长度超过键长度时导致掩码失效。通过将有符号偏移和零初始化,确保了在混合查询块中正确应用因果掩码,并定义了完全掩码行的输出为零。这为长上下文推理中的注意力计算提供了更稳健的语义,可能影响未来 LLM 的掩码设计选择。
PyTorch 作为主流深度学习框架,其注意力实现的修复反映了 LLM 时代对掩码语义的重新思考。向底部右侧掩码的转变可能影响模型架构和推理优化,尤其是在处理长序列时。此修复可能推动框架间的一致性,并影响依赖高效注意力的下游库和模型。
此修复提高了 PyTorch 在长上下文场景下的可靠性,减少了因掩码错误导致的潜在错误结果。对于依赖 PyTorch 进行 LLM 训练和推理的企业,这降低了计算错误的风险,并可能提升模型质量。同时,它展示了 PyTorch 对 LLM 需求的响应,增强了其作为 AI 基础设施的竞争力。
此修复是更大计划的第一步,未来可能将 SDPA 和 varlen 对齐到右下角掩码语义。这可能导致 PyTorch 中注意力掩码的默认行为变化,并可能影响模型训练和推理的数值结果。开发者应关注后续版本中的相关更改。