見出し画像

How to fix Flash Attention-2 on ComfyUI

2025年8月19日、Pytorch2.8.0+cu129環境でビルドしたv2.8.2のwindows用whlを、HuggingFaceにて公開しました。また、以下ビルド方法についても更新しています。


ここまでの私の努力を全否定するかのように…RTX5000系…つまりBlackwellでは、FA-3が使えないとかいう事態に直面しましたが、

元々、H100に特化して最適化されているFA-3は一般向けGPUに於いては効果が疑問視されているコードではありました。

ならば、より一般的な用途を想定しているFA-2を優先して適用し、カーネル使用を明示化する改造を、まずはComfyUIに適用しました。

まあ、ここまで散々FA-3適用に苦労したことも無駄ではなく、要は此処迄構築したロジックにFA-2を追加するだけの話です

今回の改造では、xformersは以下の状態になっている事が前提です。当然ながら、使用するFA2用whlも、Python3.12環境下でPytorch2.8.0+cu129に合わせてビルドする必要があります。

PS D:\userfiles\comfyui> python_embeded\python.exe -m xformers.info
xFormers 0.0.32.post2
memory_efficient_attention.ckF:                    unavailable
memory_efficient_attention.ckB:                    unavailable
memory_efficient_attention.ck_splitKF:             unavailable
memory_efficient_attention.cutlassF-pt:            available
memory_efficient_attention.cutlassB-pt:            available
memory_efficient_attention.fa2F@2.7.4.post1:       available
memory_efficient_attention.fa2B@2.7.4.post1:       available
memory_efficient_attention.fa3F@0.0.0:             unavailable
memory_efficient_attention.fa3B@0.0.0:             unavailable
memory_efficient_attention.fa3F_splitKV@0.0.0:     unavailable
memory_efficient_attention.triton_splitKF:         available
indexing.scaled_index_addF:                        available
indexing.scaled_index_addB:                        available
indexing.index_select:                             available
sp24.sparse24_sparsify_both_ways:                  available
sp24.sparse24_apply:                               available
sp24.sparse24_apply_dense_output:                  available
sp24._sparse24_gemm:                               available
sp24._cslt_sparse_mm_search@0.0.0:                 available
sp24._cslt_sparse_mm@0.0.0:                        available
swiglu.dual_gemm_silu:                             available
swiglu.gemm_fused_operand_sum:                     available
swiglu.fused.p.cpp:                                available
is_triton_available:                               True
pytorch.version:                                   2.8.0+cu129
pytorch.cuda:                                      available
gpu.compute_capability:                            12.0
gpu.name:                                          NVIDIA GeForce RTX 5060 Ti
dcgm_profiler:                                     unavailable
build.info:                                        available
build.cuda_version:                                1208
build.hip_version:                                 None
build.python_version:                              3.9.13
build.torch_version:                               2.8.0+cu128
build.env.TORCH_CUDA_ARCH_LIST:                    7.5 8.0+PTX 9.0a
build.env.PYTORCH_ROCM_ARCH:                       None
build.env.XFORMERS_BUILD_TYPE:                     Release
build.env.XFORMERS_ENABLE_DEBUG_ASSERTIONS:        None
build.env.NVCC_FLAGS:                              -allow-unsupported-compiler
build.env.XFORMERS_PACKAGE_FROM:                   wheel-v0.0.32.post2
build.nvcc_version:                                12.8.93
source.privacy:                                    open source
PS D:\userfiles\comfyui>

尚、以下の改造は以前公開したFA-3適用改造の実施を前提にしています。また、以下で解説している「FA-3の無効化」も必要です。
(仮にFA-3が有効ならば、FA-3が優先される仕様で開発した)

この改造がなくとも、FA-2ライブラリーがインストールされていて、条件が合致すれば、xformersはFA-2カーネルをロードする…理屈にはなっていたようです。但し、当然それは可視化されていませんでした。

Cursorによれば、改造前はこのような状態だったようです。

FA-2の自動判別について

FA-2ライブラリのインストール状況

はい、FA-2ライブラリがインストールされていれば、xformersは自動的に判別できました。

1. xformersのカーネル検出メカニズム

xformersは起動時に以下の処理を行います:

# xformers内部でのカーネル検出(簡略化)
def _detect_available_kernels():
    available_kernels = []
    
    # FA-2ライブラリの存在確認
    try:
        import flash_attn
        if hasattr(flash_attn, 'flash_attn_func'):
            available_kernels.append('fa2')
    except ImportError:
        pass
    
    # Cutlassの確認
    try:
        import cutlass
        available_kernels.append('cutlass')
    except ImportError:
        pass
    
    return available_kernels

2. FA-2ライブラリのインストール確認

FA-2がインストールされている場合:

# インストール確認
pip list | grep flash-attn
# 出力例: flash-attn 2.7.4

# または
python -c "import flash_attn; print('FA-2 available')"
# 出力例: FA-2 available

3. 自動判別の流れ

A. 起動時の検出

# xformers起動時
1. flash_attnライブラリを検索
2. 利用可能なカーネルをリストアップ
3. 優先順位を決定
4. ディスパッチテーブルを構築

B. 実行時の選択

# 実際のattention実行時
1. 入力サイズをチェック
2. 利用可能なカーネルから最適なものを選択
3. FA-2が利用可能ならFA-2を優先
4. フォールバックとしてCutlassを使用

4. 改造前の実際の動作

FA-2ライブラリがインストールされていれば:

# 改造前のattention_xformers
def attention_xformers(q, k, v, heads, mask=None, attn_precision=None, skip_reshape=False, skip_output_reshape=False):
    # xformers内部で自動的に:
    # 1. FA-2カーネルを検出
    # 2. パフォーマンスを評価
    # 3. FA-2を選択して実行
    out = xformers.ops.memory_efficient_attention(q, k, v, attn_bias=mask)
    return out

5. 改造の必要性

改造が必要だった理由:

A. 確実性の向上

# 改造前:xformersの内部アルゴリズムに依存
# 改造後:明示的にFA-2を最優先に設定
dispatch_fw._kernels = fa2_kernels + other_kernels

B. ログによる可視化

# 改造前:どのカーネルが使用されているか不明
# 改造後:詳細なログで確認可能
logger.info(f"[xformers] FA2 kernels prioritized: {len(fa2_kernels)} kernels moved to front")

6. インストール状況の確認方法

A. コマンドライン確認

# FA-2ライブラリの確認
pip show flash-attn

# 出力例:
# Name: flash-attn
# Version: 2.7.4
# Summary: Flash Attention 2

B. Python確認

# Pythonでの確認
try:
    import flash_attn
    print(f"FA-2 version: {flash_attn.__version__}")
except ImportError:
    print("FA-2 not installed")

7. 自動判別の条件

FA-2が自動判別される条件:

| 条件 | 必須 | 説明 |
|------|------|------|
| flash-attnライブラリ | ✅ | `pip install flash-attn` |
| CUDA対応 | ✅ | CUDA 11.8+ |
| GPU互換性 | ✅ | RTX 5060 Ti対応 |
| PyTorch互換性 | ✅ | PyTorch 2.0+ |

8. 結論

FA-2ライブラリがインストールされていれば、xformersは自動的に判別して使用していました。

改造の目的:

  1. より確実なFA-2選択

  2. 詳細なログ表示

  3. 将来のFA-3対応準備

  4. デバッグの容易さ向上

基本的な動作は変わらず、改造は「保険的な措置」として機能しています。

9.ForgeとreForgeの場合

上のCursorの説を裏付けるように、Forge及びreForgeにおいてはFA-3を無効化していれば、以下のxformersカーネル解析ログ出力によって、FA-2が使用されている事が確認されています。

10.A1111の場合

A1111の場合、ノーマル状態ではFA-2ロードする機能を持っていません。以下記事の改造が必要になります。

ComfyUI xformers改造 - ログ強化とFA-2優先設定

改造概要

実施した改造:

  1. xformersログ強化 - カーネル選択の詳細表示

  2. FA-2優先設定 - RTX 5060 Ti対応

  3. 動的優先順位制御 - FA-3/FA-2/Cutlassの自動選択

ファイル情報

改造ファイル: `ComfyUI/comfy/ldm/modules/attention.py`

改造関数: `attention_xformers()`
改造日: 2025-08-17
対象GPU: RTX 5060 Ti (SM120)

1. ログ強化の改造

改造前

def attention_xformers(q, k, v, heads, mask=None, attn_precision=None, skip_reshape=False, skip_output_reshape=False):
    # ... 既存のコード ...
    out = xformers.ops.memory_efficient_attention(q, k, v, attn_bias=mask)
    # ... 既存のコード ...

改造後

def attention_xformers(q, k, v, heads, mask=None, attn_precision=None, skip_reshape=False, skip_output_reshape=False):
    # ... 既存のコード ...
    
    # Log xformers kernel selection and set FA2 priority
    try:
        import logging
        logger = logging.getLogger("xformers_attention_log")
        if not logger.handlers:
            handler = logging.StreamHandler()
            formatter = logging.Formatter('%(message)s')
            handler.setFormatter(formatter)
            logger.addHandler(handler)
            logger.setLevel(logging.INFO)
            logger.propagate = False
        
        # FA3/FA2優先設定(RTX 5060 Ti対応)
        if not hasattr(attention_xformers, "_fa_priority_set"):
            try:
                import xformers.ops.fmha.dispatch as dispatch_module
                # dispatch.pyのログフラグもリセット
                if hasattr(dispatch_module, '_dispatch_fw'):
                    if hasattr(dispatch_module._dispatch_fw, '_last_kernel'):
                        dispatch_module._dispatch_fw._last_kernel = None
                # より積極的なFA優先設定
                if hasattr(dispatch_module, '_dispatch_fw'):
                    dispatch_fw = dispatch_module._dispatch_fw
                    if hasattr(dispatch_fw, '_kernels'):
                        # FA-3が利用可能かチェック
                        fa3_kernels = [k for k in dispatch_fw._kernels if 'fa3' in str(k).lower() and 'unavailable' not in str(k).lower()]
                        fa2_kernels = [k for k in dispatch_fw._kernels if 'fa2' in str(k).lower() and 'unavailable' not in str(k).lower()]
                        other_kernels = [k for k in dispatch_fw._kernels if 'fa3' not in str(k).lower() and 'fa2' not in str(k).lower()]
                        
                        if fa3_kernels:
                            # FA-3が利用可能: FA-3を最優先
                            dispatch_fw._kernels = fa3_kernels + fa2_kernels + other_kernels
                            logger.info(f"[xformers] FA3 kernels prioritized: {len(fa3_kernels)} kernels moved to front")
                        elif fa2_kernels:
                            # FA-3が利用不可: FA-2を最優先
                            dispatch_fw._kernels = fa2_kernels + other_kernels
                            logger.info(f"[xformers] FA2 kernels prioritized: {len(fa2_kernels)} kernels moved to front")
                        else:
                            logger.info("[xformers] No FA kernels available, using default priority")
            except Exception as e:
                logger.info(f"[xformers] FA priority setting failed: {e}")
            attention_xformers._fa_priority_set = True
        
        # 1行だけログを表示
        if not hasattr(attention_xformers, "_logged_once"):
            logger.info("[xformers] attention_xformers called - kernel selection will be logged")
            attention_xformers._logged_once = True
    except Exception:
        pass
    
    out = xformers.ops.memory_efficient_attention(q, k, v, attn_bias=mask)
    # ... 既存のコード ...

2. 改造の詳細解説

A. ログシステムの構築

logger = logging.getLogger("xformers_attention_log")
if not logger.handlers:
    handler = logging.StreamHandler()
    formatter = logging.Formatter('%(message)s')
    handler.setFormatter(formatter)
    logger.addHandler(handler)
    logger.setLevel(logging.INFO)
    logger.propagate = False

目的:

  • 専用ロガーを作成してxformersログを分離

  • 重複ログを防ぐためのハンドラー管理

  • シンプルなフォーマットで読みやすく表示

B. フラグリセット機能

if hasattr(dispatch_module, '_dispatch_fw'):
    if hasattr(dispatch_module._dispatch_fw, '_last_kernel'):
        dispatch_module._dispatch_fw._last_kernel = None

目的:

  • `_last_kernel`フラグをリセット

  • 各処理でログを再表示するため

  • 拡大処理やFaceDetailerでもログが表示される

C. 動的カーネル優先順位制御

# FA-3が利用可能かチェック
fa3_kernels = [k for k in dispatch_fw._kernels if 'fa3' in str(k).lower() and 'unavailable' not in str(k).lower()]
fa2_kernels = [k for k in dispatch_fw._kernels if 'fa2' in str(k).lower() and 'unavailable' not in str(k).lower()]
other_kernels = [k for k in dispatch_fw._kernels if 'fa3' not in str(k).lower() and 'fa2' not in str(k).lower()]

動作:

  1. 利用可能なカーネルを検索(`'unavailable' not in str(k).lower()`)

  2. FA-3、FA-2、その他に分類

  3. 優先順位に応じてリストを再構築

D. 条件分岐による優先順位制御

if fa3_kernels:
    # FA-3が利用可能: FA-3を最優先
    dispatch_fw._kernels = fa3_kernels + fa2_kernels + other_kernels
    logger.info(f"[xformers] FA3 kernels prioritized: {len(fa3_kernels)} kernels moved to front")
elif fa2_kernels:
    # FA-3が利用不可: FA-2を最優先
    dispatch_fw._kernels = fa2_kernels + other_kernels
    logger.info(f"[xformers] FA2 kernels prioritized: {len(fa2_kernels)} kernels moved to front")
else:
    logger.info("[xformers] No FA kernels available, using default priority")

優先順位:

  1. FA-3利用可能 → FA-3 → FA-2 → Cutlass

  2. FA-3利用不可 → FA-2 → Cutlass

  3. 両方利用不可 → Cutlass(デフォルト)

E. 一度だけログ表示

if not hasattr(attention_xformers, "_logged_once"):
    logger.info("[xformers] attention_xformers called - kernel selection will be logged")
    attention_xformers._logged_once = True

目的:

  • 最初の呼び出し時のみログを表示

  • 過度なログ出力を防ぐ

  • 処理の開始を明確に示す

3. 改造の効果

A. 詳細なカーネル選択ログ

[xformers] attention_xformers called - kernel selection will be logged
[xformers] FA2 kernels prioritized: 2 kernels moved to front
[xformers] memory_efficient_attention: selected kernel = fa2F@2.7.4.post1

B. 動的優先順位制御

  • RTX 5060 Ti: FA-2が自動選択

  • 将来のBlackwell: FA-3が自動選択

  • FA-2削除時: Cutlassが自動選択

C. 各処理でのログ再表示

  • 最初の生成: ログ表示

  • 拡大処理: ログ再表示

  • FaceDetailer: ログ再表示

4. 技術的なポイント

A. 安全な実装

  • try-exceptでエラーをキャッチ

  • hasattrで属性の存在確認

  • 段階的なフォールバック

B. パフォーマンスへの配慮

  • 一度だけ設定(`_fa_priority_set`フラグ)

  • 軽量なログ出力

  • 既存機能への影響なし

C. 将来対応

  • FA-3対応の準備完了

  • 新しいカーネルへの拡張性

  • ハードウェア変更への自動適応

5. テスト結果の検証

正常動作の確認

  • FA-2優先選択: `fa2F@2.7.4.post1`

  • 動的切り替え: SAオフ/オン

  • ログ再表示: 各処理で表示

  • フォールバック: FA-2削除時Cutlass選択

改造は完璧に成功しています!

ただ、身も蓋もない事を言えば、より簡単に実装できるSA-2++の方が、効果が薄いSDXLベースであっても高速です。

以下で公開している、フル仕様のSDXLプログラムを走らせた時、FA-2の約85秒前後に対しSA-2++は平均75秒前後と有意に高速です。cutlassとFA-2の比較はしていませんが。

CCSR

CCSRでxformersのカーネル解析ログが連続する問題に関しては、以下記事を参照して下さい。

今日のBGMは

鞘師里保 - Super Red (Behind The Scenes)

鞘師里保 - Super Red (Performance Video)


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