1
0

Delete article

Deleted articles cannot be recovered.

Draft of this article would be also deleted.

Are you sure you want to delete this article?

PyTorch 2.13のMPS FlexAttentionをM1 Maxで検証――疎な注意で最大7.83倍

1
Last updated at Posted at 2026-07-21

こんにちは、皆さん。

長い文章を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.132026年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、miseuvが必要です。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の学習対応ではありません。

結果を簡単に読む

  1. 長い入力から、ごく一部だけを見るなら速くなりました。

    32,768 tokenで直近256 tokenだけを見る条件では7.83倍高速です。計算を省ける範囲が大きいほどFlexAttentionが活きます。

  2. 普通のAttentionをそのまま置き換える機能ではありません。

    過去全体を見るcausal attentionでは、標準のSDPAが約19倍高速でした。patternと密度を測って選ぶ必要があります。

  3. 同じ規則を繰り返し使うことが重要です。

    最初に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推論や、疎な関係だけを扱う実験には使えそうです。

1
0
0

Register as a new user and use Qiita more conveniently

  1. You get articles that match your needs
  2. You can efficiently read back useful information
  3. You can use dark theme
What you can do with signing up
1
0

Delete article

Deleted articles cannot be recovered.

Draft of this article would be also deleted.

Are you sure you want to delete this article?