← 返回文章列表
August 4, 2026
5 分钟阅读

一整天的内容能胜过十分钟的内容吗? Flash 注意力和序列长度问题

一整天的内容能胜过十分钟的内容吗? Flash 注意力和序列长度问题
#deep-learning
#attention
#Flash-Attention
#GPU
#optimization

本文要回答的问题是:如果变压器能够以一秒的分辨率而不是十分钟的窗口处理整个交易日,它的预测会更好吗?

直到最近你甚至不能问。标准注意力需求 O(N2)\mathcal{O}(N^2) 内存,所以一天 23,400 步,12 个头,float16 仅分数矩阵就需要大约 12.9 GB——比模型参数还多,也比大多数卡给你的要多。在有人测试之前,这个问题已经通过算术结束了。

Flash Attention (Dao et al., 2022) 打开它。不是通过近似注意力——它计算出完全相同的结果,而是通过将计算重构为IO感知,最大限度地减少GPU内存级别之间的流量。这是这里真正有趣的内容,本文的大部分内容都花在它的工作原理上:平铺、在线 softmax 递归、 Θ(N2d2/M)\Theta(N^2 d^2 / M) IO 限制和向后传递重新计算。

但机制是推动者,而不是主张。 “上下文越长越好”是关于市场的实证陈述,以及本博客的立场——来自时间融合变压器,它发现金融系列过度拟合和短回溯循环模型上的普通变压器在高频下保持竞争力 - 反过来说。因此,本文结束于测量,而不是机制。

为什么注意力是受记忆限制的

注意力计算 softmax(QK/dk)V\text{softmax}(\mathbf{Q}\mathbf{K}^\top / \sqrt{d_k})\mathbf{V} - 在交易环境中,原语本身包含在[用于多水平投资组合预测的时间融合变压器](/en/blog/post/temporal-fusion-transformer-trading)。整个问题就是一行:中间分数矩阵 S=QK/dk\mathbf{S} = \mathbf{Q}\mathbf{K}^\top/\sqrt{d_k}N×NN \times N,它被写入内存,读回 softmax,再次写入,然后再次读取最终的 matmul——并且必须保留它以用于反向传播。

Attention的算术强度为 min(d,N)\approx \min(d, N),所以大约 64 FLOP/字节 d=64d = 64 — A100 山脊点左侧。它位于倾斜的带宽上限,而不是平坦的计算上限:GPU 花费更多时间移动 S\mathbf{S} 周围比乘以任何东西。它使用的屋顶线框架——山脊点、倾斜与平坦的天花板,以及为什么同样的推理决定 GPU 是否值得购买——是根据当 GPU 发挥作用

算法利用的内存层次结构

内存级别 尺寸 带宽 延迟
HBM(高带宽内存) 40-80 GB 2.0 TB/秒 ~400 纳秒
SRAM(片上、共享内存) 20 MB 19TB/秒 ~4 纳秒

SRAM 的容量大约为10 倍带宽和 100 倍低延迟,容量仅为千分之一。 Flash Attention 所做的一切都遵循这一交易:放弃容量,购买带宽和延迟。同样的“重组算法而不是购买硬件”的举措,在 CPU 回测中测量,是回测速度阶梯

Flash 注意力算法

Flash Attention 在大小适合 SRAM 的 tiles 中处理注意力,并且永远不会实现完整的注意力 N×NN \times N HBM 中的矩阵。

分割 Q\mathbf{Q} 进入 Tr=N/BrT_r = \lceil N/B_r \rceil 行块和 K,V\mathbf{K}, \mathbf{V} 进入 Tc=N/BcT_c = \lceil N/B_c \rceil 柱块,具有 Br,BcB_r, B_c 选择一块瓷砖及其累加器以适合芯片。对于每个查询块,迭代所有键值块:

For each query block Q_i:
    Initialize: O_i = 0, m_i = -inf, l_i = 0    # output, running max, running sum
    For each KV block (K_j, V_j):
        1. Load Q_i, K_j, V_j from HBM to SRAM
        2. Compute S_ij = Q_i @ K_j^T / sqrt(d)  # in SRAM
        3. Compute local max: m_ij = rowmax(S_ij)
        4. Compute P_ij = exp(S_ij - m_ij)        # in SRAM
        5. Compute local sum: l_ij = rowsum(P_ij)
        6. Update running statistics:
           m_new = max(m_i, m_ij)
           l_new = l_i * exp(m_i - m_new) + l_ij * exp(m_ij - m_new)
           O_i = O_i * (l_i * exp(m_i - m_new) / l_new)
                + P_ij @ V_j * (exp(m_ij - m_new) / l_new)
           m_i = m_new, l_i = l_new
    Write O_i to HBM

在线softmax递归

使平铺成为可能的技巧是在线softmax。朴素的 softmax 需要对行进行两次遍历:一次查找最大值(为了数值稳定性),一次求幂和归一化。对您拒绝存储的行进行两次传递是一个矛盾 - 因此 Flash Attention 会不断运行统计数据并在运行过程中重新调整大小。

块后 1,,j1, \ldots, j:

m(j)=max(m(j1),max(Sij))m^{(j)} = \max(m^{(j-1)}, \max(\mathbf{S}_{ij})) (j)=(j1)em(j1)m(j)+keSijkm(j)\ell^{(j)} = \ell^{(j-1)} \cdot e^{m^{(j-1)} - m^{(j)}} + \sum_k e^{S_{ijk} - m^{(j)}}

并且输出累加器通过相同的因子进行校正:

Oi(j)=(j1)(j)em(j1)m(j)Oi(j1)+1(j)emijm(j)PijVj\mathbf{O}_i^{(j)} = \frac{\ell^{(j-1)}}{\ell^{(j)}} \cdot e^{m^{(j-1)} - m^{(j)}} \cdot \mathbf{O}_i^{(j-1)} + \frac{1}{\ell^{(j)}} \cdot e^{m_{ij} - m^{(j)}} \cdot \mathbf{P}_{ij}\mathbf{V}_j

每次新块提高运行最大值时,先前累积的输出都会追溯性地重新调整 em(j1)m(j)e^{m^{(j-1)} - m^{(j)}} ——就好像新的最大值从一开始就已经知道了。结果在代数上与两遍 softmax 相同。在精确的算术中,这不是一个近似值;而是一个近似值。这是一种重新关联。 (在有限精度下,这是一个不同的舍入路径,这很重要 - 请参阅下面的准确性检查。)

IO 复杂度

这是获胜的正式声明。 Flash Attention 执行

Θ(N2d2M)\Theta\left(\frac{N^2 d^2}{M}\right)

HBM 访问,其中 MM 是SRAM大小,反对 Θ(Nd+N2)\Theta(Nd + N^2) 为标准的实施。注意 MM 出现在分母中:片上暂存器越大,往返次数就越少,这就是为什么该算法是用内存层次结构而不是 FLOP 计数来表示的。对于典型的 d=64d = 64M100M \approx 100 KB,该比率有利于 Flash Attention,访问次数减少大约 5-10 倍。

向后传递:重新计算而不是存储

通过注意力进行反向传播通常需要 P\mathbf{P} 前向传递拒绝保留的矩阵。 Flash Attention 重新计算来自 Q,K\mathbf{Q}, \mathbf{K} 在向后传递期间,仅存储输出 O\mathbf{O} 和softmax统计 (m,)(m, \ell) - 两个都 O(N)\mathcal{O}(N), 不是 O(N2)\mathcal{O}(N^2)。它用适量的冗余算术换取了整个问题的记忆项。这与梯度检查点相同,在单个算子内以图块粒度应用。

FA2:并行性

Flash Attention 2(Dao,2023)保留了算法并修复了调度:

  1. 更少的非 matmul FLOP。 FA1 实时花费在重新缩放、求最大值和取幂上——在 CUDA 核心而不是张量核心上运行的操作。 FA2 将重新缩放推迟到内循环结束时。
  2. 序列长度上的并行性。 FA1 仅在批次和头上进行并行化。 FA2 还对查询块进行并行化。这对于交易案例尤其重要,在这种情况下,您通常有每个资产一个非常长的序列和 1-4 的批量大小 - 这正是批量和头并行性导致 GPU 饥饿的情况。
  3. Warp 工作分区。 每个 warp 都采用不同的查询块子集,而不是分割分数计算并减少跨 warp,从而消除跨 warp 减少。

报告结果:A100 的理论峰值 FLOP 约为 70%,而 FA1 约为 35%。

FA3:料斗力学

Flash Attention 3 (Dao, Shah, 2024) 是特定于 H100 的架构:

  1. 异步扭曲专门化。 Hopper 的张量内存加速器 (TMA) 异步移动 HBM→SRAM。 FA3 将扭曲分为生产者,为下一个 KV 块发出 TMA 负载,而消费者在当前 KV 块上进行计算,因此数据移动隐藏在算术后面。
  2. 交错 matmul 和 softmax。 QiKj\mathbf{Q}_i\mathbf{K}_j^\top 其中一个块在张量核心上运行,而前一个块的 softmax 在 CUDA 核心上运行——两个不同的硬件单元,真正并发而不是时间切片。
  3. 具有不相干处理的 FP8。 H100 以 2 倍 FP16 吞吐量执行 FP8。天真的 FP8 注意力会被异常值破坏; FA3 在按块量化之前随机旋转向量,以将离群值分布在坐标上,据报道,数值误差比朴素 FP8 低 2.6 倍。
版本 图形处理器 利用率 加速与标准
FA1 A100 〜35% 2-4x
FA2 A100 〜70% 5-7x
FA3 (FP16) H100 〜75% 3-5 倍 vs FA2
FA3 (FP8) H100 〜75% 1.6 倍 vs FA3 FP16

因果掩盖是交易获得折扣的地方

因果屏蔽对于时间序列是强制性的——模型不能考虑未来——并且在平铺下它不是增加成本,而是节省成本。任何其键完全在未来相对于其查询的图块都会被“直接跳过”,永远不会加载,也永远不会计算,从而减少了大约一半的工作。在 PyTorch 中这是 is_causal=True;不需要其他任何东西。

集成是八行

您需要的代码几乎都与 Flash Attention 无关。交换融合内核的显式分数矩阵路径:

scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(d_head)
scores = scores.masked_fill(causal_mask, float("-inf"))
out = torch.matmul(torch.softmax(scores, dim=-1), v)

out = torch.nn.functional.scaled_dot_product_attention(
    q, k, v,
    dropout_p=self.dropout if self.training else 0.0,
    is_causal=True,
)

这就是整个变化。 q, k, v 有形状的 (batch, heads, seq, head_dim);因果掩码消失了,因为内核构建了它。要将完整的带注释的 PyTorch 交易模型放入输入投影、块和 3 类向上/平坦/向下头中,请使用 [DeepLOB](/en/blog/post/deeplob-deep-learning-order-book),有关完整的训练流程,请参阅 [Temporal Fusion Transformers](/en/blog/post/temporal-fusion-transformer-trading)。在这里建造这个脚手架的第四个副本不会有任何帮助。

要求:计算能力 >= 8.0(A100、H100、RTX 3090+)、半精度输入、PyTorch >= 2.0。验证实际参与的快速路径 torch.backends.cuda.sdp_kernel 诊断和 torch.cuda.max_memory_allocated() — 如果任何先决条件失败,SDPA 会默默地回退到数学内核,并且静默回退看起来与工作模型完全一样,只是速度很慢。

衡量标准:较长的上下文是否值得?

以上所有内容都表明 32K 或 128K 上下文现在“负担得起”。它没有说明它是否“有用”。诚实的实验:

训练相同的架构 N{512,4096,32768}N \in \{512, 4096, 32768\} 在本系列其他地方使用的 BTC 系列上的 SDPA-Flash,保持参数、优化器和目标固定,因此序列长度是唯一的变量。报告两件事:

  1. 成本。 每个时期测量的挂钟和 torch.cuda.max_memory_allocated() 在每个 NN.
  2. 好处。 样本外预测性能与 NN,在向前推进的分裂中。

本文的早期草稿包含了从分析得出的每个序列长度的内存数字表 O(N)\mathcal{O}(N) 激活公式。这些行被删除:它们从未被测量过,并且它们不同意文章自己的内存预算算术。结果表中显示的派生数字是捏造的结果,本博客不提供这些结果。

这个实验的有趣之处在于它可以在任一方向发布。如果样本外性能单调上升 NN,这证明了整个长上下文程序的合理性。如果它在几千步处停滞不前或降级,那么这是一个更强大的部分 - 与诚实的负面——这意味着记忆墙永远不会成为交易变压器的约束力。

更多的上下文意味着更多的容量,因此更多的过度拟合表面

预期结果持平或负面是有特定原因的。 时间融合变压器已经记录了普通变压器天真地应用于金融序列过度拟合——它们缺乏时间归纳偏差,并且短回溯循环模型在高频下保持竞争力。将上下文从 512 步扩展到 32,768 步并不会增加与长度成比例的信息;接近有效的价格序列的第 32,000 个边际滞后影响很小。它可靠地添加了适合的参数值。

于是扫了一遍 NN 必须按其本质对待:模型选择搜索,其机制与本博客适用于所有其他搜索的机制相同。三序列长度乘以其他变化就是试验计数,获胜者必须清除根据该试验计数和 PBO 门计算的缩减夏普比率,而不仅仅是击败其邻居。否则,“长上下文获胜”与选择三个噪声运行中最好的一个没有什么区别。

精确性检查,因为“精确”需要做很多工作

Flash Attention 是精确的在精确的算术中。附加的建议 - 在 fp16 或 bf16 中运行,在 H100 上考虑 FP8 - 不是。这些是单独的主张,而第二个主张在实践中占主导地位:重新关联总和并降低到一半精度都是扰动,并且介绍排序保证的文章不应该手动波动精度。

该博客已经拥有合适的工具。 GPU精度陷阱上的平价预言机来证明正确性,而不是通过目视曲线。此处应用:

  • 使用 bf16 中的 SDPA-Flash 和相同输入上的 fp64 参考实现来计算注意力;报告输出张量的最大相对误差
  • 将其推向决策:对于发出向上/平坦/向下标签的模型,报告两条路径之间有多少标签翻转**,作为总决策的一部分。

小的、有限的、可解释的分歧是正确的快速路径的标志。无界意味着 FP8 建议对于该模型来说从来都不安全。在运行之前这两个数字都是未知的。

何时伸手去拿它

压缩为决策,其形状与GPU决策指南

  • 在 CUDA GPU 上超过约 2K 时间步:是的,无条件的。 这是一个产生精确输出的单行更改,并且胜利随着 NN。不存在您想要具体化的场景 N×NN \times N 路径代替。
  • 低于约 512 个时间步长,在 CPU 上,或使用非注意力架构(CNN、SSM,如 Mamba):无关紧要。 在山脊左侧,固定开销是全部成本,注意力从来都不是瓶颈。
  • 上面的阈值是民间传说,而不是测量 - 它们来自一般文献,您自己的模型和卡上的交叉是十线基准。运行它而不是相信整数。

结论

Flash Attention 是一个干净且真正重要的结果:通过尊重内存层次结构并重新关联 softmax,它可以计算出精确的注意力 O(N)\mathcal{O}(N) 内存而不是 O(N2)\mathcal{O}(N^2),以及 Θ(N2d2/M)\Theta(N^2 d^2 / M) IO 限制准确地解释了原因。在交易变压器中采用它是一项单行更改,无需精确算术的准确性成本,并且可以节省大量内存。

它“不”做的是回答顶部的问题。它将“全天环境是不可能的”转换为“全天环境很便宜”,这是实验成本的变化,而不是结果的变化。记忆之墙的倒塌是对测量的邀请,而测量正是将其从纸质摘要转变为发现的原因。

参考文献

  1. Dao, T.、Fu, D.Y.、Ermon, S.、Rudra, A.、Re, C.“FlashAttention:具有 IO 感知的快速、内存高效的精确注意力。” NeurIPS (2022)。 arXiv:2205.14135
  2. Dao, T.“FlashAttention-2:更快的注意力,更好的并行性和工作分区。” ICLR (2024)。 arXiv:2307.08691
  3. Dao, T., Shah, J.“FlashAttention-3:具有异步和低精度的快速准确的注意力。” NeurIPS (2024)。 arXiv:2407.08608
  4. Vaswani, A. 等人。 “你所需要的就是注意力。” NeurIPS (2017)。
  5. Milakov, M., Gimelshein, N.“softmax 的在线标准化器计算”。 arXiv:1805.02867(2018)。
免责声明:本文提供的信息仅用于教育和参考目的,不构成财务、投资或交易建议。加密货币交易涉及重大损失风险。

Authors

Eugen Soloviov
Eugen Soloviov

Trading-systems engineer

Trading-systems engineer building bots since 2017: cross-exchange arbitrage (connected up to 30 venues), cointegration-based pairs arbitrage across spot and futures, scalping, news and sentiment-driven strategies, trend algorithms, and portfolio management and balancing algorithms. Also builds sub-millisecond order execution, big-data warehouses, backtesting engines, AI agents, and trading interfaces (incl. open-source profitmaker.cc). Stack: JS/TS, Python, Rust/Zig/Go, DevOps, backend, frontend, architecture.

Newsletter

紧跟市场步伐

订阅我们的时事通讯,获取独家 AI 交易见解、市场分析和平台更新。

我们尊重您的隐私。您可以随时退订。