← 返回文章列表
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 的批量大小 - 這正是批量和頭並行性導致 G​​PU 飢餓的情況。
  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 交易見解、市場分析和平台更新。

我們尊重您的隱私。您可以隨時退訂。