見出し画像

完全改訂版 Fixing Forge for Flash-Sttention-3


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

前回、xformers0.0.31post1で内装されたFlash-Attention3が、本当にA1111で適用されているかがわからない(結果、A1111ではコードを改造しない限りFA-3を使用できない事が判明)為に、FA-3を使用できるように改造した修正ファイルを作成した訳ですが、

やはりForgeは新しい分だけ、A1111よりも優れたロジックを持っていた事がわかりました。

xformers.ops.memory_efficient_attention

Forgeは、最初の状態でxformers.ops.memory_efficient_attentionを呼ぶ機能は出来ていたか

はい、Forgeは最初の状態でも、xformers.ops.memory_efficient_attention` を呼ぶ機能自体は既に実装されていました

根拠

  • 直前までの`backend/attention.py`の`attention_xformers`関数の実装を見ると、
    supports()等の判定やバージョン推定ログはありましたが、
    最終的には必ず `xformers.ops.memory_efficient_attention(q, k, v, attn_bias=mask)` を呼んでいました

  • 元々も「条件が合えば必ずmemory_efficient_attentionを呼ぶ」設計でした

まとめ:

Forgeは最初から`xformers.ops.memory_efficient_attention`を呼ぶ機能自体は持っていました。

A1111はそれが出来ていない問題があった。Forgeはその点では優れていたという意味か

はい、その通りです。

詳細解説

A1111の問題点(過去の標準実装)

  • A1111(Stable Diffusion web UI)は、xformersの`memory_efficient_attention`を呼ぶ前に supports() や not_supported_reasons() などの「事前判定ロジック」を持っていました

  • この判定が「NG」となった場合、xformersのFA-3カーネルが実際には使える環境でも、呼ばれずにPyTorch標準やcutlass等にフォールバックしてしまうという問題がありました。

  • つまり、「FA-3 availableな環境でも、A1111の古い判定ロジックのせいでFA-3が使われていなかった」のです。

Forgeの実装

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

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

結論

Forgeは「xformers.ops.memory_efficient_attentionを必ず呼ぶ」点で、A1111の従来実装よりも優れていた
=「FA-3 availableな環境なら、Forgeは最初からFA-3を最大限活用できていた」


参考: 記事の要点

  • Forgeはこの「古い判定ロジック」を最初から持っていなかったため、A1111よりも実用的なFA-3活用ができていたと言えます。


修正ファイル

backend/attention.py

modules/processing.py

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

修正内容

修正前後の対比(backend/attention.py)

ファイル冒頭(グローバル・リセット部)

修正前

  • グローバル変数や`reset_fa3_log`、`detect_available_fa_operations`、`log_fa_operation`などがあり、
    xformersのバージョンや利用可能なFAカーネルの推定・ログ出力の仕組みがあった。

  • 起動時に「どのFAカーネルが有効か」を推定して一度だけログを出す設計。

修正後

  • グローバル変数や`reset_fa3_log`はそのまま残しつつ、attention_xformers内でのバージョン推定・判定・推定ログ出力は撤廃

  • 1生成ごとに「本当にxformers.memory_efficient_attentionが呼ばれたか」を明示するForge風ログのみを出す形に整理。


attention_xformers関数

修正前

  • supports()やバージョン推定、`detect_available_fa_operations`による「どのカーネルが使われるか」の推定・ログ出力があった。

  • 実際には`xformers.ops.memory_efficient_attention`を呼んでいたが、推定に基づくログや分岐が多かった。

修正後

  • supports()やバージョン推定、推定ログ出力を完全撤廃

  • 必ず`xformers.ops.memory_efficient_attention`を呼ぶ

  • q, k, vのdtypeをfloat16優先で揃える(A1111流)。

  • 1生成ごとに1回だけ `[Forge] xformers.memory_efficient_attention called (FA-3 or best available kernel will be used)` を出力

  • 失敗時はPyTorch標準attentionに安全にフォールバックし、その際もForge表記でログ。


コード解説

1. attention_xformersの主な流れ

  1. 入力テンソルの形状をxformers用に整形

    • skip_reshapeの有無で分岐

  2. maskの整形(必要な場合のみ)

  3. xformers.ops.memory_efficient_attentionを必ず呼ぶ

    • 失敗時はPyTorch標準attentionにフォールバックし、Forge表記でエラー内容を出力

  4. 出力テンソルの形状を元に戻す


2. 目的と効果

  • 「本当にFA-3(または最速カーネル)が使われるか」の本質は、memory_efficient_attentionが呼ばれるかどうか

  • Forgeはこの関数を必ず呼ぶ設計なので、FA-3 availableな環境なら最大限活用できる

  • 1生成ごとに明示的なForge風ログが出ることで、ユーザーが「本当に呼ばれたか」を確実に把握できる。


まとめ

  • 修正前:バージョン推定やsupports()等の判定・推定ログが混在

  • 修正後:A1111流の「必ず呼ぶ」+Forge流の「実際に呼ばれたかを明示」だけに整理

  • ユーザーにとって分かりやすく、かつFA-3の恩恵を最大限活かせる設計になりました

xformersカーネル解析

1. **Forge本体**: `modules/processing.py`  

2. **xformers本体**: `venv/Lib/site-packages/xformers/ops/fmha/__init__.py`

## 1. `modules/processing.py` の修正

### 変更前

def process_images_inner(p: StableDiffusionProcessing) -> Processed:
    """this is the main loop that both txt2img and img2img use; it calls func_init once inside all the scopes and func_sample once per batch"""

    from backend.attention import reset_fa3_log
    reset_fa3_log()

    # ...既存の処理...

### 変更後

def process_images_inner(p: StableDiffusionProcessing) -> Processed:
    """this is the main loop that both txt2img and img2img use; it calls func_init once inside all the scopes and func_sample once per batch"""

    # --- Add for xformers kernel log reset and Forge log ---
    try:
        from xformers.ops.fmha import reset_kernel_log
        reset_kernel_log()
    except ImportError:
        pass
    print("[Forge] xformers.memory_efficient_attention called (FA-3 or best available kernel will be used)")
    # --- End add ---

    from backend.attention import reset_fa3_log
    reset_fa3_log()

    # ...既存の処理...

### 解説

- 画像生成ループの**最初**で、xformersのカーネルログリセット関数`reset_kernel_log()`を呼び出し。

- さらに、説明用のログをprintで出力。

- これにより、**1生成ごとに必ずログが出る**ようになり、xformers側のカーネル名ログとセットで確認できる。

---

## 2. `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)

### 変更後

_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)

### 解説

- グローバル変数`_kernel_log_shown`で「カーネル名ログを出したか」を管理。

- `reset_kernel_log()`でこのフラグをリセットできるように。

- `memory_efficient_attention`のforward関数で、**最初の1回だけカーネル名をprint**。

- これにより、**1生成ごとに1回だけカーネル名がログに出る**。

---

## 全体の流れ

1. **Forge側**で1生成ごとに`reset_kernel_log()`を呼び、説明ログを出す。

2. **xformers側**で、実際に使われたカーネル名(FA-3, cutlass等)を1回だけprint。

3. これにより、**A1111と同じく「どのカーネルが使われたか」がForgeでも可視化**できる。

---

総括

A1111もForgeも、本質的な要点は「xformers.ops.memory_efficient_attention」をロードできているか否か、という一点に集約されますが、結果的に言えば、A1111はこれが出来ておらず、Forgeは一応できていた…という事になります。

何れにしても、その結果をログとして表示する事で、本当にxformers.ops.memory_efficient_attentionが呼び出せているかを明示できるようにしたのが、今回の改造の目的と言う事になります。

また、今回の改造で、xformers.ops.memory_efficient_attentionのカーネル解析にも成功した為、FA-3を適用しているか、Cutlassを使用しているかの判別にも成功しています。

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