1 日のコンテキストは 10 分間のコンテキストに勝りますか?フラッシュ アテンションとシーケンス長の問題
この記事は次の質問に答えるために存在します。もし変圧器が 10 分の時間枠ではなく 1 秒の解像度で取引日全体を処理できたら、予測はより正確になるでしょうか?
最近までは尋ねることさえできませんでした。標準的な注意ニーズ そのため、float16 で 12 ヘッドで 23,400 ステップの 1 日を実行するには、スコア行列だけで約 12.9 GB が必要になります。これはモデル パラメーターよりも多く、ほとんどのカードが提供する量よりも多くなります。質問は誰かがテストする前に算数で終了しました。
Flash Attendant (Dao et al., 2022) がその始まりです。アテンションを近似することではなく、まったく同じ結果を計算します。IO 対応になるように計算を再構築し、GPU メモリ レベル間のトラフィックを最小限に抑えます。これがここでの本当に興味深い内容であり、この記事の大部分はそれがどのように機能するか、つまりタイリング、オンライン ソフトマックスの繰り返し、 IO バウンドおよびバックワードパス再計算。
しかし、メカニズムはそれを可能にするものであり、主張ではありません。 「コンテキストは長いほうが良い」というのは、市場とこのブログの立ち位置についての経験的な発言です — Temporal Fusion Transformers/en/blog/post/temporal-fusion-transformer-trading)、金融シリーズのオーバーフィットおよびショートルックバックリカレントモデルのバニラトランスは高周波数でも競争力を維持できることがわかり、逆の方向に切り込みました。したがって、この記事はメカニズムではなく測定について終了します。
なぜ注意力は記憶に縛られるのか
注意力を計算する — トレーディングのコンテキストにおけるプリミティブ自体は、マルチホライズン ポートフォリオ予測のための時間融合トランスフォーマー。問題全体はその 1 行です: 中間スコア行列 は 、それはメモリに書き込まれ、ソフトマックスのために読み戻され、再度書き込まれ、最終的な matmul のために再び読み取られます。そして、バックプロパゲーションのために保持する必要があります。
注意の演算強度は つまり、約 64 FLOP/バイト — A100 尾根点のかなり左。平坦なコンピューティング上限ではなく、傾斜した帯域幅上限に収まります。GPU は移動により多くの時間を費やします。 何かを掛けるよりも周りに。ここで使用されているルーフラインのフレームワーク (リッジ ポイント、傾斜天井と平らな天井、そしてそもそも GPU を購入する価値があるかどうかを同じ理由で決定する理由) は、When the GPU Pays Off。
アルゴリズムが利用するメモリ階層
| メモリレベル | サイズ | 帯域幅 | レイテンシ |
|---|---|---|---|
| HBM (高帯域幅メモリ) | 40~80GB | 2.0 TB/秒 | ~400ns |
| SRAM (オンチップ、共有メモリ) | 20MB | 19 TB/秒 | ~4ns |
SRAM は、容量が 1,000 分の 1 で、およそ 帯域幅が 10 倍、遅延が 100 倍低いです。 Flash Attend が行うすべてのことは、容量を放棄し、帯域幅と遅延を購入するという取引に基づいています。同じ「ハードウェアを購入するのではなくアルゴリズムを再構築する」という動きは、CPU バックテストで測定され、バックテスト速度ラダー。
フラッシュ アテンション アルゴリズム
フラッシュ アテンションは、SRAM に収まるサイズの タイル でアテンションを処理し、完全なアテンションを実現することはありません。 HBM のマトリックスはまったくありません。
パーティション の中へ 行ブロックと の中へ 柱ブロック、付き タイルとそのアキュムレータがチップ上に収まるように選択されます。クエリ ブロックごとに、すべてのキーと値のブロックを繰り返します。
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
オンラインソフトマックスの再発
タイリングを可能にするトリックは オンライン ソフトマックス です。単純なソフトマックスでは、行に対して 2 回のパスが必要です。1 回目は最大値を見つけるため (数値安定性のため)、もう 1 回目はべき乗と正規化を行います。保存を拒否した行を 2 回通過するのは矛盾です。そのため、Flash Attendance は統計を実行し続け、進行に合わせて再スケールします。
ブロック後 :
そして出力アキュムレータは同じ係数で補正されます。
新しいブロックが実行最大値を上げるたびに、以前に蓄積された出力が遡及的に再スケーリングされます。 — あたかも新しい最大値が最初からわかっていたかのように。結果は代数的に 2 パス ソフトマックスと同じです。厳密な算術では、これは近似ではありません。それは再結合です。 (有限精度では、これは 異なる 丸めパスですが、これが重要です。以下の正確性チェックを参照してください。)
IOの複雑さ
これが勝利の正式な声明です。フラッシュアテンションの実行
HBM はアクセスします。 SRAM サイズです。 標準実装の場合。ご了承ください 分母に表示されます。オンチップのスクラッチパッドが大きいほどラウンドトリップが少なくなります。そのため、アルゴリズムが FLOP カウントではなくメモリ階層の観点から記述されています。典型的な場合 そして KB、この比率では、アクセスが約 5 ~ 10 分の 1 少ない Flash アテンションが有利になります。
バックワードパス: 保存の代わりに再計算
注意による逆伝播には通常、 フォワードパスが保持することを拒否したばかりのマトリックス。フラッシュ アテンションはタイルを 再計算します バックワードパス中は出力のみを保存します そしてソフトマックス統計 - 両方 、 ない 。これにより、問題全体であった記憶用語と引き換えに、適度な量の冗長な演算が行われます。これは、単一のオペレーター内のタイル粒度で適用される、勾配チェックポイントと同じ取引です。
FA2: 並列処理
Flash Attendant 2 (Dao、2023) はアルゴリズムを維持し、スケジュールを修正しました。
- 非 matmul FLOP が減少しました。 FA1 は、再スケーリング、最大値の検出、およびべき乗 (tensor コアではなく CUDA コアで実行される操作) にリアルタイムで費やしました。 FA2 は、内側のループの終わりまで再スケーリングを延期します。
- シーケンス長に対する並列処理 FA1 はバッチとヘッドのみに対して並列処理を行います。 FA2 はクエリ ブロックに対しても並列化します。これは、資産ごとに 1 つの非常に長いシーケンス があり、バッチ サイズが 1 ~ 4 であることが多いトレーディングの場合に特に重要です。これはまさに、バッチとヘッドの並列処理が GPU を枯渇させる状況です。
- ワープ作業の分割。 各ワープは、スコア計算を分割してワープ間で削減するのではなく、クエリ ブロックの異なるサブセットを取得し、ワープ間の削減を削除します。
報告された結果: A100 では理論上のピーク FLOP の約 70% に対し、FA1 では約 35%。
FA3: ホッパー機構
Flash Attendant 3 (Dao、Shah、2024) は、H100 に固有のアーキテクチャです。
- 非同期ワープの特殊化。 Hopper の Tensor Memory Accelerator (TMA) は、HBM→SRAM を非同期に移動します。 FA3 はワープを、次の KV ブロックの TMA ロードを発行する プロデューサー と、現在の KV ブロックを計算する コンシューマー に分割するため、データの移動は演算の背後に隠れます。
- インターリーブされた matmul とソフトマックス 1 つのブロックのソフトマックスは tensor コアで実行され、前のブロックのソフトマックスは CUDA コアで実行されます。これは 2 つの異なるハードウェア ユニットであり、タイム スライスではなく真に同時実行されます。
- 非コヒーレント処理を伴う FP8。 H100 は 2x FP16 スループットで FP8 を実行します。 FP8 の素朴な注意は異常値によって打ち砕かれます。 FA3 は、ブロック単位の量子化の前にベクトルをランダムに回転して、外れ値の大きさを座標全体に分散させます。これは、単純な FP8 よりも数値誤差が 2.6 倍低いと報告されています。
| バージョン | GPU | 活用 | スピードアップと標準 |
|---|---|---|---|
| FA1 | A100 | ~35% | 2~4倍 |
| FA2 | A100 | ~70% | 5~7倍 |
| FA3(FP16) | H100 | ~75% | FA2 に対して 3 ~ 5 倍 |
| FA3(FP8) | H100 | ~75% | 1.6x 対 FA3 FP16 |
因果マスキングにより取引が割引される
因果マスキングは時系列では必須であり、モデルは将来を考慮してはなりません。また、タイリングでは、追加コストではなく節約コストとなります。キーがそのクエリに対して完全に将来のものであるタイルは、完全にスキップされ、決してロードされず、計算もされず、作業の約半分が削減されます。 PyTorch ではこれは is_causal=True;他には何も必要ありません。
統合は 8 行です
必要なコードはほとんど Flash アテンションに関するものではありません。融合カーネルの明示的なスコア マトリックス パスを交換します。
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、完全なトレーニング パイプラインについては、Temporal Fusion Transformers/en/blog/post/temporal-fusion-transformer-trading)。ここでその足場の 4 番目のコピーを構築しても、何も教えられません。
要件: 計算能力 >= 8.0 (A100、H100、RTX 3090+)、半精度入力、PyTorch >= 2.0。実際に使用されている高速パスを確認します。 torch.backends.cuda.sdp_kernel 診断と torch.cuda.max_memory_allocated() — SDPA は、前提条件が満たされない場合、サイレントに数学カーネルにフォールバックします。サイレント フォールバックは、単に遅いだけの作業モデルとまったく同じように見えます。
測定: コンテキストが長いほど効果があるでしょうか?
上記のすべては、32K または 128K コンテキストが 手頃な価格 になったことを示しています。それが役に立つかどうかについては何も述べていません。正直な実験:
同じアーキテクチャを次の場所でトレーニングします。 このシリーズの他の場所で使用されている BTC シリーズの SDPA-Flash では、パラメーター、オプティマイザー、ターゲットが固定されているため、シーケンスの長さが唯一の変数となります。次の 2 つのことを報告します。
- コスト エポックごとに測定された壁時計と
torch.cuda.max_memory_allocated()それぞれに . - 利点 サンプル外の予測パフォーマンスとの比較 、ウォークフォワードスプリットで。
この記事の初期の草稿には、 活性化式。これらの行は削除されます。これらの行は決して測定されておらず、記事自体のメモリ バジェットの計算と一致しませんでした。結果表に示されている導出された数値は捏造された結果であり、このブログではそれらを出荷しません。
この実験の興味深い特性は、どちらの方向にも公開できるということです。サンプル外のパフォーマンスが次のように単調に上昇する場合、 、これは長いコンテキストのプログラム全体を正当化します。数千ステップで頭打ちになったり劣化したりする場合、それはより強力な部分であり、正直な否定的 — そしてそれは、メモリの壁がトランスフォーマーの取引に対する拘束力を決して持たなかったことを意味します。
コンテキストが増えると容量が増えるため、表面の過剰適合も増加します
横ばいまたはマイナスの結果が予想されるのには、特別な理由があります。 時間融合トランスフォーマー バニラトランスフォーマーが金融系列のオーバーフィットに単純に適用されたことはすでに文書化されています。バニラトランスフォーマーには時間的誘導バイアスがなく、ショートルックバックリカレントモデルは高周波数でも競争力を維持します。コンテキストを 512 ステップから 32,768 ステップに拡張しても、長さに比例した情報は追加されません。ほぼ効率的な価格シリーズのわずか 32,000 番目の遅れはほとんど意味を持ちません。確実に追加されるのは、パラメーターに相当する適合するものです。
それで一掃 このブログが他のすべての検索に適用するのと同じ仕組みを使用して、モデル選択検索として扱う必要があります。 3 つのシーケンスの長さにその他の変化を乗算したものが試行回数となります。勝者は、単に隣接するシーケンスに勝つだけでなく、その試行回数に対して計算された デフレート シャープ レシオと PBO ゲートをクリアする必要があります。それ以外の場合、「長いコンテキストの勝利」は、ノイズの多い 3 つの実行のうち最良のものを選択することと区別がつきません。
正確性チェック。「正確」は多くの作業を行うため
フラッシュ アテンションは 正確な算術 で正確です。それに付随する推奨事項 (fp16 または bf16 で実行し、H100 では FP8 を検討すること) はそうではありません。これらは別の主張であり、実際には 2 番目の主張が優勢です。合計を再関連付けすることと半精度に下げることはどちらも摂動であり、順序付けの保証を導入した記事は精度の保証を手動で行うべきではありません。
ブログにはすでに適切な手段が用意されています。 GPU 精度の罠 は標準を確立します。低精度では警告はされず、もっともらしいガベージが返されます。そして、目で見る曲線によってではなく、下流の離散数量 (取引数) に関する **パリティ オラクル ** で正しさを証明します。ここで適用されます:
- bf16 の SDPA-Flash と、同一の入力に対する fp64 リファレンス実装を使用してアテンションを計算します。出力テンソルの 最大相対誤差 をレポートします。
- 決定まで押し込みます。上/平坦/下ラベルを出力するモデルの場合、2 つのパス間で ラベルが反転するラベルの数を、決定全体の一部として報告します。
小さく、境界があり、説明可能な不一致は、正しい高速パスの兆候です。無制限のものは、FP8 推奨がこのモデルにとって決して安全ではなかったことを意味します。どちらの数値も実行されるまでわかりません。
いつそれを手に入れるか
GPU 決定ガイド と同じ形式の決定に圧縮されています。/en/blog/post/when-gpu-pays-off-sweep-roofline):
- CUDA GPU で約 2,000 タイムステップを超える場合: はい、無条件です。 これは 1 行の変更で正確な出力が生成され、効果は時間とともに大きくなります。 。実現したいシナリオはない 代わりにパス。
- ~512 タイムステップ未満、CPU 上、または非アテンション アーキテクチャ (CNN、Mamba などの SSM) の場合: 無関係。 尾根の左側、固定オーバーヘッドが全体のコストであり、アテンションがボトルネックになることはありませんでした。
- 上記のしきい値は民間伝承であり、測定値ではありません - これらは一般文献に由来しており、独自のモデルとカードのクロスオーバーは 10 ラインのベンチマークです。概数を信頼するのではなく、実行してください。
結論
フラッシュ アテンションはクリーンで真に重要な結果です。メモリ階層を尊重し、ソフトマックスを再関連付けすることにより、正確なアテンションを計算します。 代わりに記憶 、そして IO バウンドはその理由を正確に説明しています。これをトレーディングトランスフォーマーに採用すると、1 行の変更で済み、正確な演算の精度コストは発生せず、メモリの大幅な節約になります。
それがしないのは、上部の質問に答えることです。 「1 日のコンテキストは不可能」を「1 日のコンテキストは安価」に変換します。これは実験の結果ではなく、実験のコストの変化です。記憶の壁が崩れるということは、測定への誘いであり、測定こそが、これを紙上の要約から発見に変えるものなのです。
参考文献
- Dao, T.、Fu, D.Y.、Ermon, S.、Rudra, A.、Re, C. 「FlashAttendant: IO 認識を備えた高速かつメモリ効率の高い正確なアテンション」。 NeurIPS (2022)。 arXiv:2205.14135
- Dao, T. 「FlashAttendant-2: より優れた並列処理と作業分割による高速なアテンション」。 ICLR (2024)。 arXiv:2307.08691
- Dao, T.、Shah, J. 「FlashAttendant-3: 非同期性と低精度による高速かつ正確なアテンション」。 NeurIPS (2024)。 arXiv:2407.08608
- Vaswani、A.、他。 「必要なのは注意力だけです。」 NeurIPS (2017)。
- Milakov, M.、Gimelshein, N. 「ソフトマックスのオンラインノーマライザー計算」 arXiv:1805.02867 (2018)。
Authors
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.