FlashAttentionの速度とGPUメモリを比較する方法|PyTorchの測定コードと注意点

同じ入力をMathとFlashで比較し時間とメモリを測るイメージ

同じ入力をMathとFlashで比較し時間とメモリを測るイメージ

FlashAttentionの効果を調べるには、同じ入力を使って計算方式だけを切り替え、処理時間と追加のピークメモリを別々に測ります。 PyTorchではSDPA(スケーリング付き内積Attentionを計算する機能)の実行方式を固定できるため、標準のMathとFlashを比較できます。 結果を読むときは、系列長・数値形式・GPU・推論か学習かをそろえ、非対応やメモリ不足を速度の数値に混ぜないことが大切です。

この記事は測定手順とコードの解説です。掲載コードは構文と設計を確認していますが、FlashAttention対応GPUでの実行検証は未実施で、実測の速度表は掲載していません。 手元の環境はGTX 1060 3GBで、今回の比較を実施できませんでした。対応環境で実行する際も、小さい入力から動作と出力を確認してください。

仕組みを先に知りたい方は、FlashAttentionの原理解説を参照してください。ここでは、測定結果を残せる形にすることへ進みます。

最初は、Attention一回のforwardに範囲を絞る

forward(入力から出力を計算する処理)だけを、torch.inference_mode()で測ります。モデルの重み、文章の前処理、学習の勾配計算は含めません。入力はQ・K・V(問い合わせ・照合対象・集約する値を表す配列)です。

項目 この記事の条件 そろえる理由
比較する実行方式 MATHFLASH_ATTENTION 自動選択の結果と混同しない
入力形状 B=1、H=8、N=512/1024/2048、D=64 系列長Nだけを変える
数値形式 FP16(16ビット浮動小数点) 数値形式による差を混ぜない
マスク QとKが同じ長さのcausal 未来の位置を参照しない条件を一致させる
dropout 0.0 ランダムな除外処理を入れない
実行範囲 推論時のforwardのみ 学習時の追加計算・保存量と区別する

モデル全体のうちSDPAのforwardだけを比較対象とする範囲

*AI生成の概念図。今回の測定対象はSDPAのforwardです。*

Bはバッチ数、HはAttentionのヘッド数、Nは系列長、Dはヘッド内の次元数です。同じNの比較では、生成済みの同一Q/K/Vを両方式へ渡します。

PyTorchのsdpa_kernelでFlashだけを許可すれば、非対応時に別方式へ黙って切り替わった結果をFlashの速度として記録することを避けられます。対応条件はGPU・PyTorchのビルド・入力に依存するので、警告やエラーも残します。

この記事はPyTorch内蔵SDPAを使い、外部のflash-attnパッケージは導入しません。外部パッケージの対応情報はDao-AILabのREADMEにありますが、それをそのままPyTorch内蔵方式の動作保証にはできません。

時間とメモリは別の区間で測る

GPUの処理完了を待ってから時間を読む

CUDA(NVIDIA GPUで計算する仕組み)の処理は、CPUからの呼び出しに対して非同期に進みます。Python関数を呼ぶ前後の時計だけを見ると、GPUへ仕事を渡す時間を測ってしまう場合があります。

ここでは、準備運転のwarm-upを10回行い、CUDA Event(GPUの処理列に記録する時刻の目印)で20回分の実行区間を測ります。終了Eventの完了を待ってから20で割り、これを5区間繰り返した中央値を記録します。PyTorchのCUDA解説でも、正しい計時には同期やEventが必要だと説明されています。

warm-up後のGPU Event間を計測し完了後に時間を読む流れ

*AI生成の概念図。warm-upを除き、終了Eventを待って時間を読み取ります。*

得られる値は、繰り返したforwardの一回当たり時間です。短い演算ではCPUからの投入の間隔も区間に影響するため、純粋なGPUカーネル一個の実行時間とは区別します。詳しい内訳が必要なら、次の段階でプロファイラー(処理別の時間を調べる道具)を使います。

メモリは、Q/K/Vを置いた後を基準にする

メモリ測定では、入力を用意してwarm-upを終えた時点のmemory_allocated()を基準値にします。ピーク記録をリセットして一回だけforwardを実行し、max_memory_allocated()から基準値を引きます。

追加ピーク = forward中の最大割当量 − forward直前の割当量

基準の入力メモリに出力と一時領域が加わる追加ピークの考え方

*AI生成の概念図。各領域の面積は実測のメモリ量を表しません。*

この差には出力テンソルと一時的な計算領域が含まれ、既に確保してある入力は差分から除かれます。MathではFP16入力でも中間値をfloatで保持するため、両方式の差にはその実装上の条件も含まれます。PyTorchの割当管理で把握する範囲なので、GPU全体の使用量でも、Attentionの中間行列だけの大きさでもありません。

memory_reserved()はキャッシュを含む管理領域で、nvidia-smiも別の観点の値です。まずは同じ指標で比較します。定義は最大割当量のAPICUDAのメモリ管理を確認してください。

前の出力を変数に残したまま次のforwardを実行すると、出力が二つ同時に存在する区間を作ってしまいます。コードでは各回の出力を解放し、正しさ確認用の参照出力も、両方式の計測が終わってから作ります。

測定コードを保存して実行する

CUDAを使えるPyTorch環境を前提にします。インストール方法はOSとGPU環境に合わせてPyTorch公式の導入案内で選び、使用した版を記録してください。以下をbenchmark.pyとして保存します。入力を大きく変更する前に、既定のN=512でエラーが出ないか確認します。

"""Compare two SDPA forward backends on the same CUDA tensors.

Run: python benchmark.py > results.jsonl
This example has not been executed on a Flash-capable GPU by the author.
"""
from __future__ import annotations

import json
import logging
import statistics
import sys

try:
    import torch
    import torch.nn.functional as F
    from torch.nn.attention import SDPBackend, sdpa_kernel
except ImportError:
    print('ERROR: Install a CUDA-enabled PyTorch with torch.nn.attention.sdpa_kernel.', file=sys.stderr)
    raise SystemExit(2)

logging.basicConfig(level=logging.INFO, format='%(levelname)s %(message)s')
LOGGER = logging.getLogger(__name__)
BACKENDS = {'math': SDPBackend.MATH, 'flash': SDPBackend.FLASH_ATTENTION}


def emit(row: dict[str, object]) -> None:
    """Write one record. Args: row, JSON-compatible fields. Returns: None.

    Raises: TypeError for non-JSON values. Example: emit({'status': 'ok'}).
    """
    print(json.dumps(row, ensure_ascii=False), flush=True)


def forward(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor:
    """Evaluate square causal attention. Args: q/k/v, CUDA B,H,N,D tensors.

    Returns: output tensor. Raises: RuntimeError for unsupported inputs.
    Example: output = forward(q, k, v).
    """
    # Explicit dropout avoids changing the operation between comparisons.
    return F.scaled_dot_product_attention(q, k, v, dropout_p=0.0, is_causal=True)


def measure(name: str, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> dict[str, object]:
    """Measure a warmed forward. Args: backend name and identical input tensors.

    Returns: latency and allocator statistics. Raises: RuntimeError/OOM.
    Example: row = measure('math', q, k, v).
    """
    with sdpa_kernel([BACKENDS[name]]):
        for _ in range(10):
            out = forward(q, k, v)
            del out  # Do not keep an earlier result alive during the next call.
        torch.cuda.synchronize()
        samples: list[float] = []
        start = torch.cuda.Event(enable_timing=True)
        end = torch.cuda.Event(enable_timing=True)
        for _ in range(5):
            start.record()
            for _ in range(20):
                out = forward(q, k, v)
                del out
            end.record()
            end.synchronize()
            samples.append(start.elapsed_time(end) / 20)

        # Measure memory separately so reference outputs cannot inflate the baseline.
        torch.cuda.synchronize()
        baseline = torch.cuda.memory_allocated()
        torch.cuda.reset_peak_memory_stats()
        out = forward(q, k, v)
        torch.cuda.synchronize()
        peak = torch.cuda.max_memory_allocated()
        del out

    return {
        'backend': name, 'status': 'ok',
        'median_ms': statistics.median(samples),
        'min_ms': min(samples), 'max_ms': max(samples),
        'baseline_MiB': baseline / 2**20,
        'peak_allocated_MiB': peak / 2**20,
        'extra_peak_MiB': (peak - baseline) / 2**20,
    }


@torch.inference_mode()
def run_case(n: int) -> None:
    """Compare one sequence length. Args: n, positive sequence length.

    Returns: None; writes JSON lines. Raises: RuntimeError on input allocation.
    Example: run_case(512).
    """
    shape = (1, 8, n, 64)
    q, k, v = [torch.randn(shape, device='cuda', dtype=torch.float16) for _ in range(3)]
    ok: set[str] = set()
    for name in BACKENDS:
        LOGGER.info('N=%d backend=%s', n, name)
        try:
            row = measure(name, q, k, v)
            ok.add(name)
        except torch.cuda.OutOfMemoryError as error:
            row = {'backend': name, 'status': 'oom', 'reason': str(error)}
        except RuntimeError as error:
            # Report the real error: not every RuntimeError means unsupported hardware.
            row = {'backend': name, 'status': 'error', 'reason': str(error)}
        emit({'kind': 'measurement', 'N': n, **row})
        torch.cuda.empty_cache()  # Only between backends, never inside the timed section.

    if ok == set(BACKENDS):
        try:
            # Accuracy checks follow both measurements to avoid retaining a reference tensor.
            with sdpa_kernel([SDPBackend.MATH]):
                reference = forward(q, k, v)
            with sdpa_kernel([SDPBackend.FLASH_ATTENTION]):
                actual = forward(q, k, v)
            diff = (actual.float() - reference.float()).abs()
            finite = bool(torch.isfinite(reference).all() & torch.isfinite(actual).all())
            emit({'kind': 'output_check', 'N': n,
                  'finite': finite,
                  'max_abs_diff': diff.max().item() if finite else None,
                  'mean_abs_diff': diff.mean().item() if finite else None})
        except RuntimeError as error:
            emit({'kind': 'output_check', 'N': n, 'status': 'error', 'reason': str(error)})


def main() -> int:
    """Run the example. Args: none. Returns: 0, or 2 without CUDA.

    Raises: unexpected non-RuntimeError failures. Example: raise SystemExit(main()).
    """
    if not torch.cuda.is_available():
        LOGGER.error('CUDA is unavailable; no performance measurements were made.')
        return 2
    torch.manual_seed(1234)
    emit({'kind': 'environment', 'torch': torch.__version__,
          'cuda_runtime': torch.version.cuda, 'gpu': torch.cuda.get_device_name(0),
          'capability': torch.cuda.get_device_capability(0),
          'B': 1, 'H': 8, 'D': 64, 'dtype': 'float16', 'causal': True,
          'dropout': 0.0, 'mode': 'inference_forward', 'warmup': 10,
          'repeats_per_interval': 20, 'intervals': 5})
    for n in (512, 1024, 2048):
        try:
            run_case(n)
        except RuntimeError as error:
            emit({'kind': 'case', 'N': n, 'status': 'error', 'reason': str(error)})
        torch.cuda.empty_cache()
    return 0


if __name__ == '__main__':
    raise SystemExit(main())

実行すると、環境情報と各条件の記録をJSONL(1行ごとに独立したJSONを保存する形式)で出力します。進捗ログとPyTorchの警告は標準エラー出力へ残します。

python benchmark.py > results.jsonl 2> benchmark.log

NVIDIAのドライバー版は、同じ環境のnvidia-smiでも記録します。その表示のCUDA Versionはドライバー側の対応情報で、コードが出すtorch.version.cudaのビルド情報とは分けて扱ってください。

出力のどこを比較するか

出力項目 読み方 注意点
median_ms 5区間の一回当たり平均時間の中央値 全forward個別の中央値ではない
min_ms / max_ms 同じ5区間の平均時間の最小・最大 ばらつきが大きければ負荷や温度も確認
extra_peak_MiB 基準から増えた最大割当量 MiBは1024×1024バイト。出力も含む
peak_allocated_MiB 測定区間の割当量ピーク 入力などの基準値を含む
max_abs_diff / mean_abs_diff Math出力とFlash出力の絶対差 Mathを数学的な真値とはみなさない
finite 両出力にNaNや無限大がないか trueだけでは正しさや品質を保証しない

速度比は、同じNのMathの時間÷Flashの時間で計算します。両方がstatus: okの行だけを使い、エラーの行へ0ミリ秒を入れないようにします。status: oomはGPUメモリ不足、status: errorは理由を確認すべき実行エラーです。後者はGPU非対応に限らず、入力や環境の不一致も含みます。

コードの終了コードが0でも、全条件の成功を意味しません。非有限値があると差の欄はnullになるため、速度比の採用前に確認します。個々のstatusoutput_checkを読み、出力差の確認に失敗した条件を採用しないでください。CUDAが使えない場合や必要なPyTorch機能を読み込めない場合は、測定せず終了コード2で止まります。

融合された演算では計算の順序が変わり、出力がビット単位で一致するとは限りません。SDPAの公式仕様を踏まえ、絶対差を記録したうえで、実際のモデルで許容できる誤差や品質を別途確認します。この例のランダム入力で差が小さくても、文章生成の品質検査の代わりにはなりません。

条件を広げるときの順番

まず系列長Nだけを変え、時間と追加ピークがどう動くかを見ます。その後にバッチ数、ヘッド数、ヘッド次元、数値形式を一つずつ変えます。条件を変えた記録は、環境情報のB/H/D/dtypeも合わせて更新してください。

変更するもの 確認したいこと 混同しやすい点
系列長 長い入力で時間・メモリ差が広がるか 長くすれば必ず一定倍率になるとは限らない
バッチ数・ヘッド数 並列に処理する量による変化 総入力サイズも増える
数値形式 対応可否、時間、誤差の変化 形式変更と方式変更を同時に比較しない
実行順 Math先行・Flash先行で傾向が変わらないか 温度やクロック、他処理の影響

学習を調べるなら、逆伝播も含む別の測定を用意します。入力の勾配や保存される中間値が増えるため、この記事のforwardだけのメモリ値を学習時の必要量として使うことはできません。実モデルの生成速度を知りたい場合も、モデル全体の処理を別に計測します。

一文字ずつの生成へ、そのまま流用しない

decode(過去の情報を参照して次のトークンを生成する段階)では、Qの長さが1、K/Vが過去を含む長さになることがあります。この記事の正方形入力を単にQ=1へ変え、is_causal=Trueのまま測ると、意図と違うマスクになるおそれがあります。

PyTorch SDPAの非正方形causalは左上揃えです。Qの長さ1、Kの長さ4なら、左上の三角形として許されるのは先頭の位置になります。K/Vの4位置がすべて「現在までの有効な履歴」なら、最後のQが参照してほしい範囲とは一致しません。paddingや未来の位置を含むかによって必要なマスクが変わるため、decodeの意味を決めてから実装してください。SDPAのis_causal仕様

Qが1でKが4の左上揃えcausalと全履歴が有効な場合の参照範囲

*AI生成の概念図。下段は4位置すべてが有効な履歴である場合の例です。*

自分の用途に近い条件で両方式が動き、時間・追加ピーク・出力差を確認できたら、実モデルで比較する段階へ進めます。測定対象と条件を残しておくと、PyTorchやGPUを変えたときも同じ基準で見直せます。

参考資料

コメント

タイトルとURLをコピーしました