見出し画像

Archive of how to enable Flash Attention-3 on Pytorch2.8.0+cu129+xformers0.0.32.post2


How to fix xformers to RTX5060To 16GB(Balckwell) 

VGA環境をRTX5060Ti 16GBに変更しましたが、現状BlackwellではFA3を使用できないようで、xformersにおいてはFA-2又はcutlassにフォールバックする必要があります。

問題は、現在のxformersはFA-3を内装している上、デフォルトでFA-3をロードするロジックを持っている為、却ってそれがエラーの原因になります。

安定動作の為に、Blackwellでは寧ろ内装されているFA-3を無効化した方が良く、下記事でその方法は解説しています。

Ada…RTX4000系を想定した記述は以下になります。以下記事は、RTX4070 12GB環境にて作成しました。

For A1111

ComfyUIに先行して、A1111への適用を完成版として公開しました。A1111でのFA-3適用が、Pytorch2.8.0+cu129+sformers0.0.32.post2環境でも拍子抜けするほど再現できただけに(A1111本体側の改造が完成しているとは言え)、改めて、後述するComfyUIでの悪戦苦闘が理解できません。何で、この差になるかな。

本来、FA-3カーネルをロードする為に追加加工が必要だったのはA1111であって、ComfyUIは本来それがノーマルで出来た筈なのです。それが急に出来なくなった理由が真剣にわからん…

A1111と違って本体側も随時更新されているので、それが理由なのかどうか…それしか思いつかん。

For Forge

A1111に続き、Forgeに対しても修正ファイルを更新しました。こちらも、簡単に移行出来ました。更新したファイルも、A1111用に作成したxformers用のinit.pyがそのまま流用できるので、互換性を維持する改造が出来ました。

こういう形で使い回せるのが理想なのですが。

また、Forgeに関してはカーネル解析とログ表示機能のみの追加であり、FA-3カーネルのロード自体は改造なしでも適用される筈です。

更に、SA-2++適用に関しても、上記事適用により、そのまま使用できます。

For reForge

A1111に続き、Forgeに対しても修正ファイルを更新しました。こちらも、簡単に移行出来ました。更新したファイルも、A1111用に作成したxformers用のinit.pyがそのまま流用できました。

SA-2++の適用も、上記事適用を前提に、そのまま使用できます。

For ComfyUI

GPUをRTX5060Ti16GBに変更し、FA-3が使用できなくなった為、以下はアーカイヴとして保存しておきます。

ここで再現性を担保しておかないと、Cursorが再現できないので。あやつは、チャットを切り替えると完璧に過去を忘れやがるからです。繰り返しますが、歴史家としては許し難い仕様ですよ。

しかも、こいつは平然と大ウソをつきます。人間様舐めてんじゃねえ…こっちゃコードは無知でも、理屈を突き詰めて、矛盾点をひたすら責め続ける点には自信あんだよ。

危惧はしていましたが、やはり公式版whlを以てしても、FA-3が使えませんでした。以下の状態であっても、現実にはFA-3は使われないのです。

PS D:\userfiles\comfyui> python_embeded\python.exe -m xformers.info
xFormers 0.0.32.post1
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@0.0.0:             unavailable
memory_efficient_attention.fa2B@0.0.0:             unavailable
memory_efficient_attention.fa3F@2.8.0.post2-3-g3ba6f82: available
memory_efficient_attention.fa3B@2.8.0.post2-3-g3ba6f82: available
memory_efficient_attention.fa3F_splitKV@2.8.0.post2-3-g3ba6f82: available
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:                            8.9
gpu.name:                                          NVIDIA GeForce RTX 4070
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.post1
build.nvcc_version:                                12.8.93
source.privacy:                                    open source
PS D:\userfiles\comfyui>

0.0.31.post1から色々仕様そのものが変わっているようで、そこを力ずくで何とかしました。

勿論、SA-2++との共存が可能な形です。--use-sage-attention環境下で、KJ Nodesのオンオフにより、FA-3とSA-2++が切り替わります。

とりあえずComfyUIで成功させましたが、再現性の担保がこのノートの最重要目的です。

真剣に意味不明なのが、SM90にしてもfp32にしても、2.7.1+cu128+0.0.31.post1でも同じ条件だったという事です。それでも、ちゃんとFA-3は動作していました。にも関わらず、今回ComfyUIでFA-3カーネルをロードする為に、滅茶苦茶苦労した訳です。

前回の私の改造は、単にxformersのカーネルの解析であって、FA-3のカーネルのロード機能そのものを改造した訳ではありません。その改造が必要だったのは、A1111だけです。

なのに、Cursorがその差異が生じる合理的な理由を説明できません。

2.7.1+cu128+0.0.31.post1で出来た事が、何故2.8.0+cu129+0.0.32.post1で出来ないのか。

SM90もfp32も理由でないなら、別に何か2.8.0+cu129+0.0.32.post1で動作しなかった理由が存在する筈です

こっちが理解できないと思って舐めてんのか、推測や憶測でものを言いやがるので、その都度コード対比をさせるのですが、すぐ誤魔化そうとしやがるので、
(その種の噓を見抜く為に、高度なコードの知識は必要ない。初歩的な論理学の問題である。理屈に合う合わないは、すぐにわかる)

その都度「推測憶測でモノを言うな」「実際にコードを見ろ」と怒鳴りつけてやり直させてますが、未だに説明が出来ません。動作すると言えばするのでそれで良い…理屈もありますが、気持ち悪いんですよ、それじゃ。

PyTorch 2.7.1+cu128+xformers 0.0.31でFA-3が動作した理由

バックアップファイルの制限

fa3_available()関数

def fa3_available() -> bool:
    has_cuda = torch.version.cuda is not None
    is_90a = has_cuda and torch.cuda.get_device_capability() >= (9, 0)  # SM90以上
    has_valid_flash3 = flash3._C_flashattention3 is not None
    return is_90a and has_valid_flash3

SUPPORTED_DTYPES

SUPPORTED_DTYPES: Set[torch.dtype] = {
    torch.half,      # float16
    torch.bfloat16,  # bfloat16
}  # float32は含まれていない

実際のオペレーター選択メカニズム

_dispatch_fw_priority_list

def _dispatch_fw_priority_list(inp: Inputs, needs_gradient: bool):
    if torch.version.cuda:
        flash3_op = [flash3.FwOp] if _get_use_fa3() else []
        priority_list_ops = deque(
            flash3_op  # FA-3が最初に試される
            + [
                flash.FwOp,
                cutlass.FwOp,
            ]
        )

_get_use_fa3()の実装

_USE_FLASH_ATTENTION_3 = True  # 定数

def _get_use_fa3() -> bool:
    global _USE_FLASH_ATTENTION_3
    return _USE_FLASH_ATTENTION_3  # 常にTrue

flash3.FwOpの実際の制限

計算能力要件

class FwOp(AttentionFwOpBase):
    CUDA_MINIMUM_COMPUTE_CAPABILITY = (8, 0)  # SM80以上

実際の制限チェック

def _run_priority_list(name: str, priority_list: Sequence[T], inp: Inputs):
    for op in priority_list:
        not_supported = op.not_supported_reasons(inp)
        if not not_supported:  # サポートされている場合
            return op

PyTorch 2.7.1+cu128の影響

Flash Attention 3統合

if importlib.util.find_spec("...flash_attn_3._C", package=__package__):
    from ..._cpp_lib import _build_metadata
    from ...flash_attn_3 import _C
    _C_flashattention3 = torch.ops.flash_attn_3

ビルド情報

{
  "version": {
    "cuda": 1208,
    "torch": "2.7.1+cu126",
    "flash": "v2.8.0.post2-3-g3ba6f82"
  },
  "env": {
    "TORCH_CUDA_ARCH_LIST": "7.5 8.0+PTX 9.0a"
  }
}

実際の動作フロー

1. 初期化

  • `_C_flashattention3 = torch.ops.flash_attn_3` (利用可能)

  • `_USE_FLASH_ATTENTION_3 = True`

2. オペレーター選択

  • `_get_use_fa3()` → `True`

  • `flash3_op = [flash3.FwOp]` (FA-3が優先リストに追加)

  • `priority_list_ops = [flash3.FwOp, flash.FwOp, cutlass.FwOp]`

3. 制限チェック

  • `flash3.FwOp.not_supported_reasons(inp)` が呼ばれる

  • 計算能力: RTX 4070 (SM89) ≥ SM80 → 要件を満たす

  • データ型: PyTorch 2.7.1+cu128環境でfloat32が実際にサポート

矛盾の解決

fa3_available()の役割

  • バックアップ: SM90以上を要求

  • 実際: オペレーター選択では使用されない

SUPPORTED_DTYPESの制限

  • バックアップ: float32非サポート

  • 実際: PyTorch 2.7.1+cu128環境でfloat32サポート

計算能力要件

  • バックアップ: fa3_available()でSM90以上要求

  • 実際: flash3.FwOpでSM80以上要求(RTX 4070は満たす)

結論

動作した理由:

  1. `fa3_available()`は実際の選択に関係しない

  2. `flash3.FwOp`はSM80以上を要求(RTX 4070は満たす)

  3. PyTorch 2.7.1+cu128で実際にfloat32サポート

  4. バックアップファイルの制限は実際の動作と異なる

Flash Attention 3 (FA-3) を RTX 4070 (SM89) で動作させる完全な手順

概要

PyTorch 2.8.0+cu129 + xformers 0.0.32 環境で、RTX 4070 (SM89) にFlash Attention 3を動作させ、SageAttentionとの動的切り替えを実現する手順です。

環境情報

  • GPU: NVIDIA GeForce RTX 4070 (SM89)

  • PyTorch: 2.8.0+cu129

  • xformers: 0.0.32.post1

  • CUDA: 12.8.93

修正ファイル一覧

1. xformers 0.0.32 の修正

`python_embeded/Lib/site-packages/xformers/ops/fmha/dispatch.py`

修正内容: SM89対応 + カーネル選択ログ追加

def fa3_available() -> bool:
    has_cuda = torch.version.cuda is not None
    is_90a = has_cuda and torch.cuda.get_device_capability() >= (9, 0)
    has_valid_flash3 = flash3._C_flashattention3 is not None
    # Allow FA-3 on SM89 (Ada Lovelace) for testing
    is_89a = has_cuda and torch.cuda.get_device_capability() >= (8, 9)
    return (is_90a or is_89a) and has_valid_flash3
def _dispatch_fw(inp: Inputs, needs_gradient: bool) -> Type[AttentionFwOpBase]:
    """Computes the best operator for forward
    Raises:
        NotImplementedError: if not operator was found
    Returns:
        AttentionOp: The best operator for the configuration
    """
    op = _run_priority_list(
        "memory_efficient_attention_forward",
        _dispatch_fw_priority_list(inp, needs_gradient),
        inp,
    )

    # Log the selected kernel
    try:
        import logging
        logger = logging.getLogger("xformers_attention_log")
        if not hasattr(_dispatch_fw, "_last_kernel"):
            _dispatch_fw._last_kernel = None
        last_kernel = _dispatch_fw._last_kernel
        if op.NAME != last_kernel:
            logger.info(f"[xformers] memory_efficient_attention: selected kernel = {op.NAME}")
            _dispatch_fw._last_kernel = op.NAME
    except Exception:
        print(f"[xformers] memory_efficient_attention: selected kernel = {getattr(op, 'NAME', 'unknown')}")

    return op

`python_embeded/Lib/site-packages/xformers/ops/fmha/flash3.py`

修正内容: float32対応 + PyTorch 2.8.0互換性

SUPPORTED_DTYPES: Set[torch.dtype] = {
    torch.half,
    torch.bfloat16,
    torch.float,  # Add float32 support for compatibility with 0.0.31
} | ({torch.float8_e4m3fn} if FLASH3_HAS_FLOAT8 else set())
def _flash_attention3_incompatible_reason() -> Optional[str]:
    # PyTorch 2.8.0ではflash_attn_3のopsが利用できない場合があるため、xformersの内部実装を使用
    if not hasattr(torch.ops, 'flash_attn_3'):
        return None  # xformersの内部実装を使用するため、エラーを返さない

    if not hasattr(torch.ops.flash_attn_3, "fwd") or not hasattr(
        torch.ops.flash_attn_3, "bwd"
    ):
        return None  # xformersの内部実装を使用するため、エラーを返さない

    if not torch.ops.flash_attn_3.fwd.default._schema.is_backward_compatible_with(...):
        return None  # xformersの内部実装を使用するため、エラーを返さない

    if not torch.ops.flash_attn_3.bwd.default._schema.is_backward_compatible_with(...):
        return None  # xformersの内部実装を使用するため、エラーを返さない

    return None

2. ComfyUI の修正

`ComfyUI/comfy/ldm/modules/attention.py`

修正内容: 優先順位変更 + ログ追加

# デフォルトでxformers(FA-3)を使用、SageAttentionはKJノードが動的に制御
if model_management.xformers_enabled():
    logging.info("Using xformers attention (FA-3)")
    optimized_attention = attention_xformers
elif model_management.sage_attention_enabled():
    logging.info("Using sage attention")
    optimized_attention = attention_sage
elif model_management.flash_attention_enabled():
    logging.info("Using Flash Attention")
    optimized_attention = attention_flash
elif model_management.pytorch_attention_enabled():
    logging.info("Using pytorch attention")
    optimized_attention = attention_pytorch
else:
    if args.use_split_cross_attention:
        logging.info("Using split optimization for attention")
        optimized_attention = attention_split
    else:
        logging.info("Using sub quadratic optimization for attention, if you have memory or speed issues try using: --use-split-cross-attention")
        optimized_attention = attention_sub_quad
def attention_xformers(q, k, v, heads, mask=None, attn_precision=None, skip_reshape=False, skip_output_reshape=False):
    # ... 既存コード ...

    # Log xformers kernel selection
    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

        # Force xformers to select kernel and log it
        if not hasattr(attention_xformers, "_last_logged"):
            attention_xformers._last_logged = False
            logger.info("[xformers] attention_xformers called - kernel selection will be logged")
    except Exception:
        pass

    out = xformers.ops.memory_efficient_attention(q, k, v, attn_bias=mask)
    # ... 既存コード ...

`ComfyUI/comfy/ldm/modules/diffusionmodules/model.py`

修正内容: VAE用ログ追加

def xformers_attention(q, k, v):
    # ... 既存コード ...

    # Log xformers kernel selection for VAE
    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

        # Force xformers to select kernel and log it
        if not hasattr(xformers_attention, "_last_logged"):
            xformers_attention._last_logged = False
            logger.info("[xformers] VAE xformers_attention called - kernel selection will be logged")
    except Exception:
        pass

    try:
        out = xformers.ops.memory_efficient_attention(q, k, v, attn_bias=None)
        out = out.transpose(1, 2).reshape(orig_shape)
    # ... 既存コード ...

起動方法

SageAttentionとの動的切り替え用

python_embeded\python.exe ComfyUI\main.py --use-sage-attention

FA-3のみ使用

python_embeded\python.exe ComfyUI\main.py --use-xformers

動作確認

FA-3使用時(SAノードオフ)

Using xformers attention (FA-3)
[xformers] attention_xformers called - kernel selection will be logged
[xformers] memory_efficient_attention: selected kernel = fa3F@2.8.0.post2-3-g3ba6f82

SageAttention使用時(SAノードオン)

Patching comfy attention to use sageattn
SageAttention kernel is being used for this generation.
[SageAttention][DEBUG] sageattn (auto) called
Restoring initial comfy attention

技術的詳細

なぜこの修正が必要だったか

  1. SM89対応: xformers 0.0.32の`fa3_available()`はSM90+のみを許可

  2. float32対応: `SUPPORTED_DTYPES`に`torch.float`が含まれていない

  3. PyTorch 2.8.0互換性: `torch.ops.flash_attn_3`の可用性チェックが厳しすぎる

  4. 優先順位: ComfyUIでxformersをSageAttentionより優先させる必要

修正の効果

  • FA-3がSM89 (RTX 4070) で動作

  • float32データ型をサポート

  • PyTorch 2.8.0環境での安定動作

  • KJノードによる動的切り替えが可能

注意事項

  • この修正はxformers 0.0.32専用です

  • 他のバージョンでは動作しない可能性があります

  • 再インストール時は修正を再適用する必要があります

トラブルシューティング

循環インポートエラーが発生した場合

# xformersを再インストール
python_embeded\python.exe -m pip uninstall xformers
python_embeded\python.exe -m pip install xformers==0.0.32.post1
# 修正を再適用

FA-3が選択されない場合

  • `--use-sage-attention`オプションを外して起動

  • KJのSAノードがオフになっているか確認


この手順により、RTX 4070でFA-3とSageAttentionの動的切り替えが完全に動作します。

Fixed Florence-2 for Flash Attention-3

Pytorchとは無関係ですが、ノードそのものの更新に伴い、修正ファイルそのものを更新しました。


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