こんにちは、皆さん。
長い文章をAIへ渡すと、Attentionの計算量は急に大きくなります。では、「直近の単語だけを見る」という制限を付ければ、MacのGPUでも速くできるのでしょうか。
さて、今日はPyTorch 2.13でApple Siliconへ対応したFlexAttentionを、M1 Max上で標準のSDPAと比較します。
先に結果を書くと、32,768 tokenから直近256 tokenだけを見る条件では、FlexAttentionが75.27 ms、SDPAが589.05 msで、7.83倍高速でした。一方、通常のcausal attentionではSDPAの方が約19倍高速です。FlexAttentionは常に速いのではなく、「長く、極端に疎なAttention」で効く機能でした。
FlexAttentionとMPS対応
FlexAttentionは、PyTorch 2.5とともに2024年10月にprototypeとして発表されました。Attentionのルールを短いPython関数で表し、torch.compileが専用の高速kernelへ変換するAPIです。
Attentionは、入力中のどのtokenを重視するか計算する仕組みです。通常はtokenが増えるほど比較の組み合わせが急増します。FlexAttentionでは「過去だけを見る」「直近256 tokenだけを見る」などの規則を指定し、不要な組み合わせを省けます。この省かれた部分が多い状態を**疎(sparse)**と呼びます。
PyTorch 2.13は2026年7月8日に発表され、FlexAttentionのApple Silicon向けMetal/MPS kernelが追加されました。公式benchmarkでは、疎なpatternでSDPAより最大約12倍高速と報告されています。APIとkernel optionは2.13時点でもunstableです。
今回は学習済みAIモデルを使っていません。乱数で作ったAttention入力を、2つの実装へ同じように渡すkernel benchmarkです。
使用技術の役割
| 技術 | 役割 |
|---|---|
| PyTorch 2.13.0 / FlexAttention | customなAttention規則をcompiled kernelで実行 |
| SDPA | 比較対象となるPyTorch標準のAttention実装 |
| MPS / Metal | Apple SiliconのGPUで計算 |
役割とデータの流れ
2つの経路は、同じquery、key、valueを受け取ります。これらは簡単にいえば「何を探すか」「何と照合するか」「取り出す内容」です。
同じ乱数入力(query / key / value)
├─ FlexAttention
│ Pythonのmask規則 -> BlockMask -> torch.compile -> Metal kernel
│
└─ SDPA
dense mask --------------------> MPS backend
-> 出力誤差とforward時間を比較
BlockMaskは、計算する範囲を128×128 tokenのblock単位でまとめた情報です。FlexAttentionは不要なblockを飛ばせます。SDPA側には同じ規則を表す通常のboolean maskを渡しました。
今回検証する内容
- 公式と同じ代表shapeで、M1 Maxでも疎なAttentionが速くなるか
- 8,192 tokenではwindowをどこまで広げるとSDPAが逆転するか
- 初回compileとBlockMask生成を含めても利点があるか
- FlexAttentionとSDPAの出力が一致するか
- MPSでbackward(学習時の勾配計算)が使えるか
完全なコードとJSON reportは、kiarina/labsのpytorch-2-13-flexattention-mps labで公開しています。
検証環境の再現
Apple Silicon Mac、mise、uvが必要です。default taskは32,768×32,768のboolean maskも生成します。これは単体で1 GiBになるため、ほかのGPU workloadを終了し、十分なmemoryがある環境で実行してください。
git clone --depth 1 --filter=blob:none --sparse \
https://github.com/kiarina/labs.git
cd labs
git sparse-checkout set .gitignore .mise/tasks Makefile mise.toml \
2026/07/21/pytorch-2-13-flexattention-mps
mise -C 2026/07/21/pytorch-2-13-flexattention-mps run
長いcaseを省く場合は、lab内で次を実行できます。
uv run python benchmark.py --quick
検証条件
query、key、valueは独立した乱数です。両実装へ同じtensorを渡し、mask生成とcompileを除いたforwardだけを測りました。MPSは非同期に動くため、各計測の前後で処理完了を待っています。3回のwarm-up後、10回の中央値を採用しました。
machine: MacBook Pro (Apple M1 Max, 32 GPU cores, 64 GB)
OS: macOS 26.5.2
Python: 3.13.7
PyTorch: 2.13.0
shape: batch 1、8 heads、head dimension 64
dtype: bfloat16
FlexAttention: torch.compile(..., dynamic=False)
BlockMask: 128×128 block
CPU fallback: disabled
sliding windowは、現在位置から過去W tokenまでを見るpatternです。たとえばwindow 256なら、非常に長い入力でも各tokenが見る範囲を直近256 tokenへ限定します。
検証結果
SDPA / Flexが1より大きければFlexAttentionが高速です。密度は、全組み合わせのうち実際に見る割合です。
| Pattern | Sequence | Window | Token / block密度 | Flex中央値 | SDPA中央値 | SDPA / Flex |
|---|---|---|---|---|---|---|
| causal | 8,192 | — | 50.01% / 50.78% | 231.22 ms | 12.33 ms | 0.05× |
| local | 8,192 | 64 | 0.78% / 3.10% | 11.75 ms | 25.23 ms | 2.15× |
| local | 8,192 | 256 | 3.08% / 4.61% | 19.14 ms | 25.25 ms | 1.32× |
| local | 8,192 | 1,024 | 11.72% / 13.18% | 56.24 ms | 25.30 ms | 0.45× |
| local | 8,192 | 4,096 | 37.50% / 38.67% | 169.24 ms | 25.30 ms | 0.15× |
| local | 32,768 | 256 | 0.78% / 1.17% | 75.27 ms | 589.05 ms | 7.83× |
8,192 tokenでは、window 256まではFlexAttentionが勝ち、window 1,024ではSDPAが2.22倍高速に逆転しました。この条件の境界はtoken密度3.08%から11.72%の間です。
causal attentionは各tokenが過去全体を見るため、密度が約50%あります。SDPAにはこの一般的なpatternの専用pathがあり、FlexAttentionより18.75倍高速でした。通常のcausal attentionを置き換えるだけでは逆効果です。
32,768 / window 256では、10回中1回だけFlexAttentionが147.89 msまで遅くなりました。ただしFlexAttentionの範囲74.96〜147.89 msとSDPAの586.08〜590.03 msは重ならず、優劣は変わりません。
公式benchmarkとの差
PyTorch公式値は8,192 / window 64で4.15倍、32,768 / window 256で約12.3倍です。M1 Maxではそれぞれ2.15倍、7.83倍でした。
疎なほど速く、sequenceが長いほど差が広がる傾向は再現できましたが、倍率は公式値に届きませんでした。公式blogには測定したApple Siliconの機種が書かれていないため、差をhardwareだけの影響とは断定できません。
初回準備コスト
FlexAttentionには、定常的なforwardとは別にBlockMask生成と初回compiled callが必要です。
| Pattern | Sequence / window | BlockMask生成 | 初回compiled call |
|---|---|---|---|
| causal | 8,192 / — | 564.54 ms | 546.57 ms |
| local | 8,192 / 64 | 95.90 ms | 94.31 ms |
| local | 8,192 / 256 | 100.54 ms | 102.34 ms |
| local | 8,192 / 1,024 | 95.02 ms | 130.75 ms |
| local | 8,192 / 4,096 | 96.65 ms | 244.62 ms |
| local | 32,768 / 256 | 1,120.85 ms | 162.57 ms |
32,768 / window 256は1回あたり約514 ms短縮するため、同じmaskを再利用すれば約3 forwardで準備コストを回収できます。8,192 / window 64では約14 forward必要です。1回しか使わないmaskなら、定常時の7.83倍や2.15倍だけで選べません。
この初回値にはhost上のcompile cache状態が含まれます。またBlockMaskはeagerに生成しており、mask生成自体をcompileする方法は測っていません。
出力の一致と対応範囲
同じMPS bfloat16入力に対するFlexAttentionとSDPAの最大絶対誤差は0.0078125〜0.015625、平均絶対誤差は0.000079〜0.000363でした。
小さい入力をCPU float32のSDPAとも比較しました。
| MPS実装 | 最大絶対誤差 | 平均絶対誤差 |
|---|---|---|
| FlexAttention bfloat16 | 0.012440 | 0.000553 |
| SDPA bfloat16 | 0.012440 | 0.000677 |
両MPS実装の最大誤差は同じで、FlexAttention固有の大きなずれは観測しませんでした。ただし、これはモデル全体の品質評価ではなく、合成入力に対する数値比較です。
requires_grad=TrueのprobeはFlexAttention does not support backward on MPSで失敗しました。PyTorch 2.13のMPS版はforward inference専用です。2.13で追加されたdeterministic backwardはCUDA向けで、MPSの学習対応ではありません。
結果を簡単に読む
-
長い入力から、ごく一部だけを見るなら速くなりました。
32,768 tokenで直近256 tokenだけを見る条件では7.83倍高速です。計算を省ける範囲が大きいほどFlexAttentionが活きます。
-
普通のAttentionをそのまま置き換える機能ではありません。
過去全体を見るcausal attentionでは、標準のSDPAが約19倍高速でした。patternと密度を測って選ぶ必要があります。
-
同じ規則を繰り返し使うことが重要です。
最初にmask作成とcompileの待ち時間があります。複数layerや複数回の推論で再利用できる処理に向きます。
制限
- M1 Max 1台、1 process、合成乱数入力だけを測定した
- bfloat16、batch 1、8 heads、head dimension 64だけを測定した
- 3回warm-up後の10回という短時間の結果で、長期的な発熱やほかのGPU負荷を制御していない
- forward prefillだけを対象とし、decode、GQA、captured buffer、score modificationを試していない
- peak memory、消費電力、Metal kernel単位のprofileを測っていない
- official benchmarkと機種、OS、測定processの全条件が同一ではない
- FlexAttention APIとkernel optionは2.13時点でunstable
検証後の感想
M1 Maxでも、極端に疎な長いAttentionで7.83倍まで差が開いたのは良い結果でした。Pythonでmask規則を書くだけでMetalの専用kernelへ変換できるため、独自のAttentionをMac上で試す敷居はかなり下がっています。
一方で、windowを広げると早い段階でSDPAに逆転され、初回準備にも無視できない時間がかかりました。「Flex」という名前だけで万能な高速化を期待せず、密度と再利用回数を実際に測る必要があります。長い文書から近傍だけを参照するローカルLLM推論や、疎な関係だけを扱う実験には使えそうです。