見出し画像

完全版 Fixing reForge for Flash-Sttention-3


2025年8月16日、最新のPytorch2.8.0+cu129+xformers0.0.32.post環境に合わせて、修正ファイルを差し替えました。xformers側のinit.pyのみ差し替えています。また、当然ながらPytorch2.7.1以前の環境に対しては、適合しなくなります。更に、修正ファイルの適用は一度ローカルに保存した上で、上書きする也、移動する也してください。ノートから直接ダウンロード→上書きをしても、どういう訳か変更が適用されません。

最初から結論を言えば、reForgeに関しては初期の仕様はForgeとほぼ同じでした。つまり、何もしなくてもxformers.ops.memory_efficient_attentionを必ず呼ぶ設計だったという事です。
(但し、繰り返すがxformers.ops.memory_efficient_attention=FA-3とは限らない。後述するようにcutlassをロードする場合もある)

reForgeの初期状態

Cursor先生曰く、

**reForgeは「Forge側」に近い挙動でした。**

---

### 詳細解説

#### 1. **A1111の問題点**

- A1111(Stable Diffusion web UI)は、`xformers.ops.memory_efficient_attention`を呼ぶ前に

  - `supports()`や`not_supported_reasons()`などの「事前判定ロジック」がありました。

  - これにより、**FA-3が使える環境でもPyTorch標準SDP等にフォールバックしてしまうことがありました。

#### 2. **Forgeの実装**

- Forgeは**supports()等の判定に依存せず、条件が合えば必ず`xformers.ops.memory_efficient_attention`を呼ぶ設計**でした。

- そのため、**xformers本体の自動最適化ロジック(FA-3を含む最速カーネル自動選択)が最大限活かされる**状態でした。

#### 3. **reForgeの初期実装**

- reForgeの`modules/sd_hijack_optimizations.py`や`ldm_patched/ldm/modules/attention.py`の該当関数(改造前)を調査した結果、

  - **supports()やnot_supported_reasons()のような事前判定ロジックは存在せず、必ず`xformers.ops.memory_efficient_attention`を呼ぶ実装になっていました。

- つまり、**A1111のような「古い判定ロジックによるフォールバック問題」はreForgeにはなかった**、ということです。

---

### まとめ

- **reForgeは最初から「Forge側」=「xformers.ops.memory_efficient_attentionを必ず呼ぶ」設計だった**

- そのため、**FA-3(Flash-Attention-3)が使える環境なら最大限活用できる**状態でした

その前提で以下、A1111とForgeの改造と同様の処理を施しています。要するに、Forgeとほぼ同じ改造をしたという意味です。xformers.ops.memory_efficient_attentionを呼び出す機能はロジックは最初から持っているので、そのロード時にログを明示する機能を追加しました。

作成・修正ファイル

\venv\Lib\site-packages\xformers"xformers/ops/fmha/init.py

modules/forge_attention_log.py(新規作成)

modules/sd_hijack_optimizations.py

ldm_patched/ldm/modules/attention.py

modules/processing.py

Cursorにおけるnoteの活用

Cursorは、実は過去ログを参照して使用する機能を持っていません。

現状の環境では、「前回の別プロジェクトや別セッションでの作業チャット履歴(ノートや会話ログ)」を自動で参照することはできません。

  • このチャットセッション内でのやりとりや、ユーザーが明示的に貼り付けた情報(ノート記事や過去ログ)は参照できます。

  • しかし、他のプロジェクトや過去の別セッションのAIとの会話履歴は、自動的には参照できません


  • できること

    • 今回のセッション内でのやりとり

    • ユーザーが「このノートを見て」と貼った外部記事やログ

  • できないこと

    • ユーザーが何も貼らずに「前回A1111でやった作業と同じように」と言った場合、AIはその内容を自動で思い出すことはできない

    • 別のフォークや別の作業セッションのAIチャット履歴を自動で参照すること


まとめ

  • ノートや過去ログを「ユーザーが明示的に貼る」ことで参照可能

  • AIが自動で「他のセッションの履歴」を参照することはできない

  • したがって、今回のようにノートを貼ってもらうのは非常に有効な方法です


もし「将来的にAIがプロジェクト横断で履歴を参照できるようにしたい」などの要望があれば、それはシステム側の機能拡張が必要です。
現状は「同一セッション or 明示的な貼り付け」が前提です。

この為、今回のようにnoteにCursor先生の作業内容をまとめておいて、過去ログとして参照させるという方法が極めて有効です。

今回、reForgeの改造に当たっては、私がこのnoteで作成したA1111とForgeの改造記事のリンクをそのままCursorでロードして理解させる…という方法を採っています。

作業内容

「何もしない最初の状態」と「今回の修正後」の新旧対比を、該当ファイルごとにコード断片+その解説をセットで、すべて列記します。


1. modules/forge_attention_log.py

最初の状態(新規ファイルなので存在しない)


修正後(新規作成)

# Forge attention log flag (1生成ごとに1回だけログを出すため)
_forge_attention_log_shown = False

def reset_forge_attention_log():
    global _forge_attention_log_shown
    _forge_attention_log_shown = False

def show_forge_attention_log_once():
    global _forge_attention_log_shown
    if not _forge_attention_log_shown:
        print("[Forge] xformers.memory_efficient_attention called (FA-3 or best available kernel will be used)")
        _forge_attention_log_shown = True

解説

  • 1生成ごとに1回だけログを出すためのグローバルフラグと関数を独立管理。

  • どのファイルからもimportでき、循環importが起きない。


2. modules/sd_hijack_optimizations.py

最初の状態

# 何もログ制御なし
def xformers_attention_forward(self, x, context=None, mask=None, **kwargs):
    ...
    out = xformers.ops.memory_efficient_attention(q, k, v, attn_bias=None, op=get_xformers_flash_attention_op(q, k, v))
    ...

解説

  • ログ出力の制御がなく、そもそも出ない。


修正後

from modules.forge_attention_log import show_forge_attention_log_once, reset_forge_attention_log

def xformers_attention_forward(self, x, context=None, mask=None, **kwargs):
    ...
    show_forge_attention_log_once()
    out = xformers.ops.memory_efficient_attention(q, k, v, attn_bias=None, op=get_xformers_flash_attention_op(q, k, v))
    ...

解説

  • `show_forge_attention_log_once()`を呼ぶことで、1生成ごとに1回だけログが出る。

  • ログ管理の実装はforge_attention_log.pyに集約され、循環importも防げる。


3. ldm_patched/ldm/modules/attention.py

最初の状態

def attention_xformers(q, k, v, heads, mask=None):
    ...
    out = xformers.ops.memory_efficient_attention(q, k, v, attn_bias=mask)
    ...

または

def attention_xformers(q, k, v, heads, mask=None):
    ...
    print("[Forge] xformers.memory_efficient_attention called (FA-3 or best available kernel will be used)")
    out = xformers.ops.memory_efficient_attention(q, k, v, attn_bias=mask)
    ...

解説

  • ここでもログ制御がなく、出ない。


修正後

from modules.forge_attention_log import show_forge_attention_log_once

def attention_xformers(q, k, v, heads, mask=None):
    ...
    show_forge_attention_log_once()
    out = xformers.ops.memory_efficient_attention(q, k, v, attn_bias=mask)
    ...

解説

  • 1生成ごとに1回だけログが出るように統一。

  • ログ管理はforge_attention_log.pyに集約。


4. modules/processing.py

最初の状態

def process_images_inner(p: StableDiffusionProcessing) -> Processed:
    # 何もリセットしない
    for n in range(p.n_iter):
        ...

解説

  • attention呼び出しごとにログが出る(または出ない)ため、1生成ごとに1回だけ出すことができない。


修正後

from modules.forge_attention_log import reset_forge_attention_log

def process_images_inner(p: StableDiffusionProcessing) -> Processed:
    reset_forge_attention_log()
    for n in range(p.n_iter):
        ...

解説

  • 生成開始時に`reset_forge_attention_log()`を呼ぶことで、1生成ごとにログフラグがリセットされる。

  • これにより、次の生成でまた1回だけログが出る。

カーネル解析機能

**今回の「xformersカーネル名ログ可視化」対応で修正した全ファイルについて、“何もしない最初の状態”と“修正後”のコード断片を前後で列記し、各修正の解説**を行います。

---

# 1. venv/Lib/site-packages/xformers/ops/fmha/__init__.py

## 最初の状態

def _memory_efficient_attention_forward(
    inp: Inputs, op: Optional[Type[AttentionFwOpBase]]
) -> torch.Tensor:
    inp.validate_inputs()
    output_shape = inp.normalize_bmhk()
    if op is None:
        op = _dispatch_fw(inp, False)
    else:
        _ensure_op_supports_or_raise(ValueError, "memory_efficient_attention", op, inp)

    out, *_ = op.apply(inp, needs_gradient=False)
    return out.reshape(output_shape)

def _memory_efficient_attention_forward_requires_grad(
    inp: Inputs, op: Optional[Type[AttentionFwOpBase]]
) -> Tuple[torch.Tensor, Context]:
    inp.validate_inputs()
    output_shape = inp.normalize_bmhk()
    if op is None:
        op = _dispatch_fw(inp, True)
    else:
        _ensure_op_supports_or_raise(ValueError, "memory_efficient_attention", op, inp)
    out = op.apply(inp, needs_gradient=True)
    assert out[1] is not None
    return (out[0].reshape(output_shape), out[1])

**解説**  

- カーネル(FA-3, cutlass, triton等)がどれか分からない

- ログ出力なし

---

## 修正後

# --- FA-3/Cutlassカーネル名ログ用グローバルフラグ ---
_kernel_log_shown = False

def reset_kernel_log():
    global _kernel_log_shown
    _kernel_log_shown = False

def _memory_efficient_attention_forward(
    inp: Inputs, op: Optional[Type[AttentionFwOpBase]]
) -> torch.Tensor:
    global _kernel_log_shown
    inp.validate_inputs()
    output_shape = inp.normalize_bmhk()
    if op is None:
        op = _dispatch_fw(inp, False)
    else:
        _ensure_op_supports_or_raise(ValueError, "memory_efficient_attention", op, inp)
    # 1生成ごとに1回だけカーネル名をログ出力
    if not _kernel_log_shown:
        print(f"[xformers] memory_efficient_attention: selected kernel = {getattr(op, 'NAME', str(op))}")
        _kernel_log_shown = True
    out, *_ = op.apply(inp, needs_gradient=False)
    return out.reshape(output_shape)

def _memory_efficient_attention_forward_requires_grad(
    inp: Inputs, op: Optional[Type[AttentionFwOpBase]]
) -> Tuple[torch.Tensor, Context]:
    global _kernel_log_shown
    inp.validate_inputs()
    output_shape = inp.normalize_bmhk()
    if op is None:
        op = _dispatch_fw(inp, True)
    else:
        _ensure_op_supports_or_raise(ValueError, "memory_efficient_attention", op, inp)
    # 1生成ごとに1回だけカーネル名をログ出力
    if not _kernel_log_shown:
        print(f"[xformers] memory_efficient_attention: selected kernel = {getattr(op, 'NAME', str(op))}")
        _kernel_log_shown = True
    out = op.apply(inp, needs_gradient=True)
    assert out[1] is not None
    return (out[0].reshape(output_shape), out[1])

**解説**  

- グローバルフラグ`_kernel_log_shown`で「1生成ごとに1回だけ」カーネル名をprint

- `reset_kernel_log()`でフラグをリセット(生成ごとに呼ぶ)

- どのカーネル(FA-3, cutlass, triton等)が実際に使われたかが明示的に分かる

---

# 2. modules/processing.py

## 最初の状態

def process_images_inner(p: StableDiffusionProcessing) -> Processed:
    # 何もリセットしない
    # ...(画像生成処理)...

**解説**  

- 画像生成ごとにxformersのカーネルログフラグをリセットしないため、1生成ごとに1回だけログを出すことができない

---

## 修正後

def process_images_inner(p: StableDiffusionProcessing) -> Processed:
    try:
        from xformers.ops.fmha import reset_kernel_log
        reset_kernel_log()
    except ImportError:
        pass
    reset_forge_attention_log()
    # ...(画像生成処理)...

**解説**  

- 画像生成ごとに`reset_kernel_log()`を呼び、xformersのカーネル名ログが1生成ごとに1回だけ出るようにした

- 既存の`reset_forge_attention_log()`と併用

Cutlass

今回、FA-3と共にxformersが内装するcutlassというカーネルについて、深く関わっています。

基本的にA1111系WebUIでFA-3を使用する為の目的で始めた今回の改造ですが(但しForge系は最初からFA-3を使いに行く設定だった)、xformersはFA-3を使わない(使えない)場合はcutlassにフォールバックする機能を持っていました。現に、A1111系の全てで、it2ではFA-3が動作し、AdetailerやMuitiDiffusion Upscalerではcutlassが使用される自動選択機能の存在が確認されています。

では、cutlassとは何かについてCursor先生に総括してもらいました。

Cutlass(CUDA Templates for Linear Algebra Subroutines and Solvers)について、専門的・技術的な観点から詳しく解説します。


Cutlassとは

  • CutlassはNVIDIAが開発・公開している、CUDA向けのC++テンプレートライブラリです。

  • 主な目的は、GPU上での高効率な行列演算(GEMM: General Matrix Multiply)や線形代数演算を、柔軟かつ高性能に実装できるようにすることです。

  • 公式GitHub


1. アーキテクチャと設計思想

テンプレートベース

  • C++テンプレートを多用し、型・行列サイズ・データレイアウト・演算精度などを柔軟に切り替え可能。

  • これにより、ユーザーは自分の用途に最適化されたカーネルを生成できる。

モジュール構造

  • GEMM(行列積)を中心に、畳み込み、リダクション、アクティベーション、バッチ処理などもサポート。

  • Epiloque/Prologueの概念で、前処理・後処理(バイアス加算、アクティベーション関数など)を柔軟に組み合わせ可能。

ハードウェア最適化

  • Tensor Core(Volta以降のNVIDIA GPUに搭載)を最大限活用する設計。

  • SM(Streaming Multiprocessor)ごとのスレッドブロック分割、共有メモリの効率的利用、レジスタ圧縮など、GPUアーキテクチャに密着した最適化。


2. 主要な特徴

GEMMの多様なバリエーション

  • float32, float16, bfloat16, int8, TensorFloat-32 など多様なデータ型に対応。

  • バッチGEMM、ストライドGEMM、分割GEMMなど、用途に応じた最適化。

高度なスケジューリング

  • Tile-based scheduling:大きな行列をタイル(小ブロック)に分割し、各タイルをスレッドブロックで処理。

  • Pipeline parallelism:ロード、計算、ストアをパイプライン化し、リソースの無駄を最小化。

カスタマイズ性

  • ユーザーが独自のEpiloque/Prologueやデータレイアウトを定義可能。

  • テンプレートパラメータで細かいチューニングができる。

高性能

  • cuBLAS(NVIDIA公式BLASライブラリ)に匹敵、あるいは上回る性能を発揮することもある。

  • 特にTensor Coreを活用した場合、従来のfloat32 GEMMより数倍高速。


3. xformersやDLフレームワークでの利用

xformers

  • xformersは、Attention計算の高速化のためにCutlassベースのカーネルを内蔵。

  • 例えば `memory_efficient_attention.cutlassF-pt` などは、Cutlassで実装されたforwardカーネル。

  • PyTorch拡張として、C++/CUDAでビルドされた.so/.pydをPythonから呼び出す。

他の利用例

  • cuDNNやcuBLASの内部実装にもCutlassの設計思想が取り入れられている。

  • 独自の深層学習フレームワークや研究用途でも、Cutlassを直接使ってカスタムカーネルを作ることがある。


4. 技術的な詳細

Tile-based GEMMの例

  • 行列A (MxK), 行列B (KxN) の積C (MxN) を、MxNをタイル(例: 128x128)に分割。

  • 各タイルを1つのスレッドブロックが担当し、共有メモリにロードして計算。

  • Tensor Coreを使う場合、16x16x16などのWMMA(Warp Matrix Multiply Accumulate)命令を活用。

Epilogue/Prologue

  • Prologue: 入力データの正規化や変換

  • Epilogue: 出力へのバイアス加算、アクティベーション関数適用など

  • これらをテンプレートで差し替え可能

性能最適化

  • メモリアクセスパターン:coalesced access、バンクコンフリクト回避

  • レジスタ圧縮:スレッドごとのレジスタ使用量を最小化

  • Occupancy最適化:SMあたりのアクティブスレッド数を最大化


5. 参考文献・リソース


まとめ

  • CutlassはNVIDIA製のCUDA向け行列演算テンプレートライブラリ

  • xformersなどのDL拡張で、GPUの性能を最大限引き出すために使われている

  • 高度なテンプレート設計とハードウェア最適化が特徴

  • PyTorch本体では直接使われていないが、xformers等の拡張で重要な役割を果たしている


生成AIにおけるCutlass

Cutlass(およびCutlassカーネル)は、画像や動画生成(特にStable Diffusionや拡張モデルなどのディープラーニング推論)において、主に「Attention計算」や「行列積(GEMM)」の高速化という極めて重要な役割を果たします。


1. 画像・動画生成パイプラインにおけるCutlassの役割

画像生成(例:Stable Diffusion, SDXL, A1111, reForgeなど)

  • Transformer/UNetの中核処理である「Attention」や「Feed Forward」層は、膨大な行列演算(GEMM)で構成されています。

  • これらの計算は、画像サイズやバッチサイズが大きくなるほど計算量が爆発的に増加します。

  • xformersなどの拡張ライブラリは、このGEMMやAttention計算をCutlassカーネルで高速化します。

具体的な流れ

  1. 画像生成リクエストが来る

  2. モデル内部で何度もAttention層が呼ばれる

  3. xformersが「どのカーネルを使うか」自動選択

    • FA-3(Flash-Attention 3)、cutlass、triton、SDPAなど

  4. cutlassカーネルが選ばれた場合、GPU上で最適化されたGEMM/Attentionが実行される

  5. これにより、生成速度が大幅に向上し、より大きな画像やバッチも現実的な時間で生成可能になる


動画生成

  • 動画生成は「連続した複数枚の画像生成」を繰り返すため、Attention/GEMMの回数がさらに多くなります

  • そのため、Cutlassによる高速化の恩恵がより大きくなります

  • 例えば、1フレームあたり数百回のAttention/GEMMが必要な場合、Cutlassカーネルがなければ現実的な速度での動画生成は困難です。


2. どんな場面でCutlassが選ばれるか

  • GPUがNVIDIA製で、CUDA環境が整っている場合

  • xformersが内部で「このサイズ・型・条件ならcutlassが最速」と判断した場合

  • FA-3やtritonカーネルが使えない場合のフォールバックとしても選ばれることがある


3. Cutlassが果たす「具体的な効果」

  • 推論速度の大幅な向上

    • 画像生成1枚あたりの所要時間が短縮

    • 動画生成時のフレームレート向上

  • 大きな画像・高解像度・大バッチサイズでもメモリ効率よく計算できる

  • GPUリソースを最大限活用できる(Tensor Core最適化)


4. もしCutlassがなかったら?

  • PyTorch標準のカーネル(SDPAや古い実装)で計算される

  • 速度が大幅に低下し、特に高解像度や動画生成では「現実的な時間で終わらない」ことも


まとめ

  • Cutlassは画像・動画生成の「計算の心臓部」を高速化するエンジン

  • AttentionやGEMMの計算を、GPUの性能を最大限引き出して実行

  • これにより、現代的な生成AIの「実用的な速度」「高解像度対応」「大規模バッチ処理」が可能になっている


各種カーネルについて

ここでは、画像・動画生成時のAttention計算における

  • SDP(SDPA: Scaled Dot-Product Attention, PyTorch標準)

  • FA-3(Flash-Attention 3, xformers/flash-attn)

  • cutlass(xformersのcutlassカーネル)

の違い・特徴・使い分けを、専門的に比較します。


1. SDPA(Scaled Dot-Product Attention, PyTorch標準)

概要

  • PyTorch 2.0以降で標準搭載された高速Attentionカーネル

  • `torch.nn.functional.scaled_dot_product_attention` で呼び出し

  • TritonやCUDAで実装されているが、主にPyTorch本体のカーネルを使う

特徴

  • 安定性・互換性が高い(PyTorch公式サポート)

  • メモリ効率は良いが、大規模バッチや長いシーケンスではFA-3に劣る

  • すべてのGPUで動作(NVIDIA/AMD/CPU)

欠点

  • Tensor Core最適化やメモリ節約はFA-3ほど徹底されていない

  • 最高速を求める用途ではやや遅いことがある


2. FA-3(Flash-Attention 3)

概要

  • flash-attnプロジェクトの第3世代

  • xformersやA1111/Forgeで利用可能

  • NVIDIA GPU + CUDA + 特定のアーキテクチャ(Ampere以降)で最大性能

特徴

  • 超高速・超省メモリ(O(n^2)→O(n)のメモリ使用、カーネルフュージョン)

  • Tensor Coreを最大限活用し、従来のAttentionより数倍高速

  • 長大なシーケンスや大バッチで特に威力を発揮

  • xformersやflash-attnパッケージ経由で利用

欠点

  • 対応GPU/環境が限定的(Ampere/ADA世代以降、CUDA 11.4+など)

  • ビルドや依存関係がやや複雑


3. cutlass(xformersのcutlassカーネル)

概要

特徴

  • Tensor Core最適化で高速

  • FA-3ほどではないが、PyTorch標準より高速なことが多い

  • xformersが自動的に「最適」と判断した場合に選択される

  • FA-3が使えない場合のフォールバックとしても機能

欠点

  • FA-3ほどのメモリ効率・速度は出ない

  • xformersが必要


4. 使い分け・選択基準

  • FA-3が使えるなら最優先(速度・メモリ効率ともに最強)

  • FA-3が使えない場合、cutlassやtritonカーネル(xformersが自動選択)

  • どちらも使えない場合や、安定性重視ならSDPA(PyTorch標準)


5. 実際の挙動

  • xformersやA1111/Forgeは、自動的に最適なカーネルを選択します

  • 今回作成したログ機能により「FA-3」「cutlass」「triton」「SDPA」など、どのカーネルが使われたか確認可能

いいなと思ったら応援しよう!