arXivDaily arXiv每日学术速递 周一至周五更新
arXiv周末暂无论文更新,休息一下吧,周末愉快~~
arXiv 2609.34272cs.LGcs.DCcs.NAmath.NA

BF16注意力中的对称性破缺:为什么FlashAttention梯度在训练后期会爆炸

Broken Symmetry in BF16 Attention: Why FlashAttention Gradients Blow Up Late in Training

Junlin Chen, Daize Dong, Huanwei Di, Haolong Jia, Jiawei Wu, Haotian Xie, Mingkai Zheng, Yang Li, Leshang Chen, Huishu Wang, Eric P. Xing, Hongyi Wang

首次发表
浏览论文内容

中文总结 AI 辅助

针对BF16注意力训练后期梯度爆炸问题,提出GProj方法,通过恢复softmax梯度行和为零的守恒律,将梯度误差降至与FP32相当,仅增加4.7%时间,实现稳定训练。

中文摘要 AI 辅助

BF16现已成为大规模预训练的标准格式,包括在FlashAttention等融合注意力内核中,这些内核被广泛信任。然而,当我们使用FlashAttention-3在50B个token上预训练一个450M参数的Transformer时,遇到了一个问题:训练在前25B个token期间表现健康,随后梯度范数增长了千倍,损失最终比FP32注意力高出0.2 nats,且没有出现任何NaN。仅将两层的注意力反向计算重算为FP32,就消除了几乎所有的额外梯度。部分原因已知:前向softmax中的融合乘加(FMA)操作,迄今被视为极端输入NaN案例,且在FlashAttention-3中从未修复。修复该问题可阻止梯度爆炸,但查询梯度仍然偏差超过其自身大小,且在准确梯度下,训练仍会将注意力logits驱动到其规模的数千倍。剩余误差源于一个被破坏的守恒定律。softmax分数梯度沿每一行的和为零,这使得查询梯度对键(keys)作为整体的位置不敏感;将其舍入到BF16会留下一个小的非零和,将平均键泄漏到梯度中,且这种泄漏恰好随着训练后期键变大、注意力变尖锐而增长。我们引入了GProj(规范投影),在转换后通过每行两次秩一校正恢复零和。它将剩余的中位查询/键梯度误差从219%/13%降至0.34%/0.37%,与FP32注意力相当,而每个训练步骤仅增加4.7%的时间。在匹配的从头训练运行中,它达到了与FP32注意力相同的损失,而FlashAttention-3和键平滑(key smoothing)均导致不稳定。

英文摘要

BF16 is now standard in large-scale pretraining, including in fused attention kernels such as FlashAttention, and these kernels are widely trusted. When we used FlashAttention-3 to pretrain a 450M-parameter transformer on 50B tokens, however, we ran into a problem: training was healthy for 25B tokens, then the gradient norm grew a thousandfold and the loss ended 0.2 nats above FP32 attention, without a single NaN. Recomputing the attention backward of just two layers in FP32 removes almost all of the excess gradient. Part of the cause is known: a fused multiply-add in the forward softmax, so far treated as an extreme-input NaN case and never fixed in FlashAttention-3. Repairing it stops the blow-up, but the query gradient is still wrong by more than its own size, and training still drives attention logits to thousands of times their size under accurate gradients. The remaining error comes from a broken conservation law. The softmax score gradient sums to zero along every row, which makes the query gradient blind to where the keys sit as a group; rounding it to BF16 leaves a small nonzero sum that leaks the mean key into the gradient, and the leak grows exactly as late training makes keys large and attention sharp. We introduce GProj (gauge projection), which restores the zero sum after the cast with two rank-one corrections per row. It cuts the remaining median query/key gradient errors from 219%/13% to 0.34%/0.37%, on par with FP32 attention, for 4.7% more time per training step. In matched from-scratch runs it trains to the same loss as FP32 attention, while FlashAttention-3 and key smoothing both destabilize.

发表机构

  • Rutgers University(罗格斯大学)
  • Carnegie Mellon University(卡内基梅隆大学)
  • Oracle(甲骨文公司)
  • New York University(纽约大学)
  • MBZUAI(穆罕默德·本·扎耶德人工智能大学)

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

补充信息

↑