发表机构
George Mason University(乔治梅森大学)
机构由 AI 辅助整理,请以论文原文为准。AI 中文总结
FlashBoB 提出一种 I/O 高效的精确反向-反向传播算法,利用 softmax 二阶反向的仿射结构,避免 N×N 中间张量,将序列长度扩展到 262K,比 FlashBack 快 6.3 倍。
AI 中文摘要
基于注意力机制的 Transformer 模型已成为现代深度学习中的核心构建模块,然而 softmax 注意力仍然是长上下文工作负载的主要瓶颈。尽管 FlashAttention 使前向传播和第一次反向传播具有 I/O 效率,但它不支持反向-反向传播(BoB),而 BoB 能够通过反向传播实现精确微分,从而支持二阶优化、测试时训练、基于梯度的记忆和元学习等应用。现有的 BoB 实现要么物化大型中间张量,要么在长序列长度下耗尽 GPU 内存。我们提出了 FlashBoB,一种用于 softmax 注意力中 BoB 的精确、I/O 高效算法,该算法将计算保持在片上块内,并避免所有 $N \times N$ 中间张量,其中 $N$ 是序列长度。关键洞察在于 softmax 二阶反向传播中的分层仿射结构:两个逐行标量通过仿射变换决定所有输出。这产生了一种具有有界片上静态随机存取存储器(SRAM)使用和最小片外高带宽存储器(HBM)流量的两遍调度。FlashBoB 实现了 $\Theta(N^2 d^2/M)$ 的 HBM 流量(其中 $d$ 是头维度,$M$ 是内存大小),并且在标准的 FlashAttention 风格分数重计算模型内,匹配了精确前向注意力的继承大缓存下界。实验上,它在单个 A100 80GB GPU 上将精确注意力 BoB 扩展到 $N=262\text{K}$,而先前的 PyTorch 精确基线在 $N=16\text{K}$ 时即失败,并且比 FlashBack 快达 $6.3\times$。这些结果使得精确二阶注意力在长上下文序列长度上变得实用,而先前的实现无法高效运行。
英文摘要
Transformer models built on the attention mechanism have become a central building block in modern deep learning, yet softmax attention remains a major bottleneck for long-context workloads. While FlashAttention makes the forward and first backward passes I/O-efficient, it does not support backward-over-backward (BoB), which enables exact differentiation through the backward pass for applications such as second-order optimization, test-time training, gradient-based memory, and meta-learning. Existing BoB implementations either materialize large intermediate tensors or exhaust GPU memory at long sequence lengths. We present FlashBoB, an exact, I/O-efficient algorithm for BoB in softmax attention that keeps computation within on-chip tiles and avoids all $N \times N$ intermediate tensors, where $N$ is the sequence length. The key insight is a hierarchical affine structure in the softmax double backward: two row-wise scalars determine all outputs through affine transformations. This yields a two-pass schedule with bounded on-chip static random-access memory (SRAM) usage and minimal off-chip high-bandwidth memory (HBM) traffic. FlashBoB achieves $Θ(N^2 d^2/M)$ HBM traffic ($d$ is the head dimension and $M$ is the memory size) and, within the standard FlashAttention-style score-recomputation model, matches the inherited large-cache lower bound for exact forward attention. Empirically, it scales exact attention BoB to $N=262\text{K}$ on a single A100 80GB GPU, where prior PyTorch exact baselines fail by $N=16\text{K}$, and is up to $6.3\times$ faster than FlashBack. These results make exact second-order attention practical at long-context sequence lengths where prior implementations cannot run efficiently.