Profiling in PyTorch (Part 3): Attention is all you profile
Profiling in PyTorch (Part 3): Attention is all you profile
Profiling in PyTorch (Part 3): Attention is all you profile
このブログ記事では、NVIDIA A100-SXM4-80GB GPUを使用してPyTorchにおけるアテンション・メカニズムのプロファイリングを行い、各バリアントがプロファイラ・トレースにどのように現れるかを実証します。
Naive attention
プリミティブ演算(matmul, mul, masked_fill, softmax, matmul)から構築されたナイーブなアテンション実装は、フォワードパスごとに6つのGPUカーネルを起動します。これには、out‑of‑placeな masked_fill によって引き起こされる予期しないメモリコピーが含まれています。
Naive attention with inplace causal masking
masked_fill をインプレースの masked_fill_ に置き換えることで、メモリコピー・カーネルが削除され、フォワードパスあたりのGPUカーネル数は6つから5つに減少します。
Scaled Dot Product Attention
PyTorchの F.scaled_dot_product_attention は、複数のバックエンドにディスパッチする単一のインターフェースを提供します。各バックエンドは異なるプロファイリング特性を示します。
Math backend
mathバックエンドに固定されている場合、SDPAはフォワードパスごとに約20のGPUカーネルを起動します。これはCUDAコア上でFP32で動作し、呼び出しのたびにcausal maskを再構築し、NaNを避けるために _safe_softmax を使用します。これは、ナイーブなインプレース版よりも約3.7倍遅いです。
Efficient backend
Efficientバックエンド(xformers)は、フォワードパスごとに単一の融合された fmha_cutlassF_bf16_aligned_64x64_rf_sm80 カーネルを起動します。これにより、計算はTensorコア上でbfloat16のまま維持され、HBMへの中間行列の書き込みが回避されます。
Flash backend
Flashバックエンドは、フォワードパスごとに単一の融合された pytorch_flash カーネル(FlashAttention‑2)を起動します。プロファイラでは低い推定占有率(約13%)を示しますが、レジスタと共有メモリを使用してアテンション・タイルをオンチップに保持するため、最も高速なバックエンドです。
cuDNN backend
cuDNNバックエンドは、フォワードパスごとに単一の生成されたカーネル(例:cudnn_generated_fort_native_sdpa_sm80_flash_fprop_wmma_f16_knob_6_128x64x64_4x1x1_cga1x1x1_kernel0_0)を起動します。CPU上での転置操作は回避されますが、実行時のプラン選択により高いCPU時間(約214 µs)が発生し、GPU時間はefficientバックエンドとflashバックエンドの中間に位置します。
Everything we covered, at a glance
| Variant | What we changed | Kernels / forward | What the trace revealed |
|---|---|---|---|
| Naive attention | Attention built by hand from primitives (matmul, mul, mask, softmax, matmul) | 6 | A hidden Memcpy from the out‑of‑place masked_fill. |
| Naive in-place | masked_fill → masked_fill_ |
5 | One line drops the Memcpy kernel entirely. |
| SDPA math | F.scaled_dot_product_attention pinned to the math backend |
20 | The reference: FP32 on CUDA cores, mask rebuilt every call, _safe_softmax. Correct but ~3.7× slower. |
| SDPA efficient | Efficient (xformers) backend | 1 | One fused fmha_cutlassF kernel, stays in bf16 on Tensor cores. |
| SDPA flash | Flash backend | 1 | One fused pytorch_flash kernel (FlashAttention‑2). Fastest, despite "wrong‑looking" 13% occupancy. |
| SDPA cuDNN | cuDNN backend | 1 | A per‑problem generated kernel: no transposes, cuLaunchKernelEx, but the cost moved to a fat CPU bar. |
Concluding the series
主な教訓は、プロファイラ・トレースが何を示すべきかについて予測を立て、トレースを調査し、不一致があればそれを最も興味深い洞察として扱うことです。この習慣によって、隠れた Memcpy、mathバックエンドの20のカーネル、flashの低い占有率、そしてcuDNNのCPU側のコストが明らかになりました。
これで、この「予測と検証」のアプローチを適用して、自身のモデルをプロファイリングできるようになります。