10月2日、PyTorchが「Optimizing Jagged Flash Attention with TLX: The Road Toward SOTA FA4 on Blackwell – PyTorch」と題した記事を公開した。この記事では、MetaのGEM(Generative Ads Model)を支えるJagged Flash Attentionカーネルを、TLX(Triton Low-level Extensions)を用いてNVIDIA Blackwell(B200)向けに最適化した取り組みが紹介されている。結論から言えば、フォワードパスでFA4比+13%、バックワードパスでFA4比+50%の性能を達成しながら、コード量はFA4の約10,000行に対し約3,200行(約3分の1)に抑えることに成功した。しかもCUDAではなくTritonベースのPythonコードで書かれており、カーネル専門家でなくとも読み書きできる。
FA4を上回る性能を3分の1のコード量で実現
最大の注目点は、このカーネルのコスト対性能比だ。
- コード量: 約3,200行(FA4の〜10,000行に対し約3分の1)
- フォワードパス性能: FA4比 +13%
- バックワードパス性能: FA4比 +50%
フォワードとバックワードの両パスを合算した総合性能においても、このカーネルはFA4を上回る。フォワード単体の+13%は控えめに見えるかもしれないが、バックワードの+50%が寄与した結果として、トレーニング全体のスループットへの改善効果は大きい。
この実装はCUDAやCuteDSLではなく、高レベルなTritonベースのPythonコードで書かれている。従来、Blackwellで最高性能を出すには手書きのCuteDSL/CUDAカーネルが必須とされていた。TLX(Triton Low-level Extensions)はMetaが開発したOSSのライブラリで、GitHub上でfacebookresearchリポジトリとして公開されている。Tritonの上にハードウェア制御の低レベルプリミティブ(明示的なSMEM/TMEM割り当て、非同期TMA/MMA、バリア、Cluster Launch Controlなど)を追加することで、「CUDAを書かずに最高性能を出す」というトレードオフを崩している。
なぜJagged(ギザギザ)なのか
Metaの広告モデルは、ユーザーの行動履歴のような可変長シーケンスを大量に処理する。これをパディングで固定長に揃えると、計算の最大50%が無駄になる。GEMはシーケンスをメモリ上に連続詰め込み(パック)し、各シーケンスの境界をオフセットテンソルで記録する方式を採用している。
Jagged Flash Attention(JFA)はこのパック表現に直接FlashAttentionアルゴリズムを適用するカーネルで、パディングされたトークンを一切生成しない。AttentionはGEMの中で最も遅いカーネルであり、ここを最適化することが全体性能に直結する。
実装の核心:2つのアプローチ
① 構造的変更(ワープ特化とメモリ管理)
TLXの明示的制御により、CTA(Cooperative Thread Array)内のワープを役割ごとに分割した:
- TMAロード専用ワープ
- テンソルコアMatMul専用ワープ
- Softmax/補正計算ワープ
- エピローグストアワープ
- バックワードではdQ削減専用ワープ
これによりテンソルコアが連続してMMAを発行できるようになり、Softmaxと並行実行される。またK/Vをトリプルバッファリングし、TMEMバッファのエイリアシングや明示的なプロデューサー/コンシューマーバリアで決定論的なパイプラインを構成している。
② 個別の最適化
ジャグドタイルのSM間スケジューリング
シーケンス長が桁違いにばらつくジャグド入力では、単純なタイル→SMマッピングでは一部のSMだけが高負荷になる。フォワードパスでは、全タイルをKVワークロード降順にソートし、ジグザグ(蛇行)パターンでSMに割り当てることで長短タイルをバランスよく配分する。この改善だけでフォワードカーネルで約20%の性能向上を回収した。
さらにCluster Launch Control(CLC)(Blackwellの新機能)を組み合わせ、実行時の残余分散を動的にスケジューリングする。静的バランシングとCLCは補完的な関係で、CLCは空タイルの判別ができないため、ホスト側での有効タイル事前計算と組み合わせることで真の効果を発揮する。
dQの多段ステージング(バックワード最大のボトルネック)
バックワードパスでは、dQをバッチ全体にわたって総和するため、全SMが同一のHBMアドレスにreduce-addを発行する——これがテンソルコア利用率の9〜11%を失う最大のボトルネックだった。
解決策は、SMEM経由のダブルバッファリングによって、あるスライスのHBMへのreduce-addと次のスライスのTMEM読み出しをオーバーラップさせることだ。
NCOL = BLOCK_D // (EPILOGUE_SUBTILE * 2)
STAGES = 2
dq_smem = tlx.local_alloc((BLOCK_M, NCOL), dq.dtype, STAGES)
for s in range(BLOCK_D // NCOL):
dq = tlx.local_load(dq_tmem[:, s * NCOL : (s + 1) * NCOL]) * LN2
tlx.local_store(dq_smem[s % STAGES], dq.to(dq.dtype))
tlx.fence_async_shared()
tlx.async_descriptor_store(desc_dq, dq_smem[s % STAGES],
offsets, store_reduce="add")
tlx.async_descriptor_store_wait(STAGES - 1)
コード中のLN2はlog(2)の値(≈0.6931)を表すスケーリング係数で、FlashAttentionが内部でexp2(底2の指数関数)を用いて効率的にSoftmaxを計算するための変換に使われる定数だ。
TMEMの早期解放
dQのTMEMバッファを全スライス消化前に解放することで、MMAワープが次のタイルのdQ Matmulを早期に開始できるようにする。最後の1〜2スライスをレジスタに先読みしてからTMEMを解放する設計で、解放するスライス数はオートチューニングで決定する(多すぎるとレジスタ圧力が高まり逆効果)。この改善でテンソルコア利用率が**+8〜11%**回復した。
ループピーリング
フォワードパスではKVループをマスクなしのバルクパスと末尾のマスク付き小ループに分割することで、イテレーションごとのマスク処理オーバーヘッドとそれに起因するレジスタスピルを排除する。
2-CTA協調MMA
FA4から採用した手法で、2つのCTAが単一のMatmulに協調することでMatmul重体のバックワードパスのテンソルコア利用率を向上させる。
開発効率とのトレードオフ
この取り組みが示す本質的なポイントは、カーネル開発の民主化だ。CuteDSLやCUDAによる10,000行のカーネルはカーネル専門家しか触れないが、TLXによる3,200行のTritonコードであれば、モデリングエンジニアが読んで拡張し、新しいAttentionバリアント(スライディングウィンドウ、ブロックスパースなど)を試せる。
TLXはMetaが社内で開発・活用しつつOSSとして公開しているツールであり、Tritonエコシステムの延長線上に位置する。Triton自体に不慣れな読者は、Triton公式ドキュメントやOpenAI Tritonの紹介記事を先に参照しておくと、本記事の実装の位置づけが把握しやすい。
コードはGitHubで公開されている: facebookresearch/ads_model_kernel_library
詳細はOptimizing Jagged Flash Attention with TLX: The Road Toward SOTA FA4 on Blackwell – PyTorchを参照していただきたい。