見出し画像

Archive of fixed ComfyUI for Flash Attention-3&Sage-Attention


※2025年8月19日現在、以下記事はRTX5060Ti 16GB+Pytorch2.8.0+cu129以降環境では適用できませんが、RTX4070 12GB+Pytorch2.7.1+c128環境前提の情報として、アーカイヴ保存しておきます。

A1111系WebUIに続き、ComfyUIに対してもxfomrmersのカーネル解析機能とSAロード時のログ表示機能を実装しました。

A1111系WebUIでは、Adetailer周りは全てCutlassで動作していましたが、ComfyUIではFA-3で動作している事がわかります。

一方、Ultimate UpscalerはCutlassで動作するようです。

目的

  • 公式状態(ログなし)から、SAノード未使用時は xformers(FA‑3優先)、SAノード使用時は SageAttention、という最終形態を毎回確実に再現する

  • 変更は3ファイルのみ。余計な改造はしない

対象ファイル(3つ)

comfy\ldm\modules\attention.py

python_embeded\Lib\site-packages\xformers\ops\fmha\_init_.py

custom_nodes\ComfyUI-KJNodes\nodes\model_optimization_nodes.py

コード修正

attention.py(優先順位だけ変更)

目的: --use-sage-attention で起動しても、KJのSAノードを使っていない限り xformers(FA‑3)を優先する

変更箇所: optimized_attention を決める if-elif 連鎖(該当ブロックのみ差し替え)

[公式]

if model_management.sage_attention_enabled():
    logging.info("Using sage attention")
    optimized_attention = attention_sage
elif model_management.xformers_enabled():
    logging.info("Using xformers attention")
    optimized_attention = attention_xformers
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

[最終形態]

if model_management.xformers_enabled():
    logging.info("Using xformers attention")
    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

補足

  • SDPA(PyTorch attention)は無効化していない。条件次第で普通に選ばれる

  • 変更はこのブロックだけ。他は触らない


xformers/ops/fmha/init.py(dispatch直後に一回ログを追加)


目的: カーネル選択(FA‑3 / cutlass)が切り替わった時にだけ、1回だけ logger.info を出す

変更箇所: 関数 _memory_efficient_attention_forward(...) 内の op = _dispatch_fw(inp, False) 直後に下記を挿入
(明示 op 経路や _requires_grad 側には入れない。余計な handler 追加、propagate 変更、_set_use_fa3(True) 呼び出しはしない)

[差し込みブロック]

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

全体イメージ(該当部分のみ)

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)

補足

  • logger.info 方式のみ(環境のロガー設定で INFO をコンソールに出す想定)

  • 切替時だけ出すため、直近名を関数属性 _last_kernel で保持して重複抑止


KJNodes/nodes/model_optimization_nodes.py(世代ごと一回ログ)


目的: SAノード使用時に、世代冒頭で一度だけ SA 使用のログを出す(SAノード未使用時は一切出さない)

追加(存在しなければ追加。既にあれば編集不要)
A) フラグ(クラス初期化)

self._called_this_generation = False
self._called_sageattn_this_generation = False

B) auto 分岐の最初の一回だけ(SAノード使用時)

if sage_attention == "auto":
    def func(q, k, v, is_causal=False, attn_mask=None, tensor_layout="NHD"):
        if not self._called_sageattn_this_generation:
            print("[SageAttention][DEBUG] sageattn (auto) called")
            self._called_sageattn_this_generation = True
        return sageattn(q, k, v, is_causal=is_causal, attn_mask=attn_mask, tensor_layout=tensor_layout)
    return func

C) 世代ごとに一回だけ(SAノード使用時)

def attention_sage(...):
    if not self._called_this_generation:
        print("SageAttention kernel is being used for this generation.")
        self._called_this_generation = True
    ...

D) 世代開始時のリセット

def reset_attention_sage_flag(*args, **kwargs):
    self._called_this_generation = False
    self._called_sageattn_this_generation = False
model_clone.add_callback(CallbacksMP.ON_PRE_RUN, reset_attention_sage_flag)

補足

  • 置換の開始/終了は既存の ON_PRE_RUN(差し替え)と ON_CLEANUP(復帰)で行われる

  • SAノード未使用時は差し替え自体が起こらない=attention は xformers、ログも xformers 側のみ


動作確認(期待ログ)

SAノード未使用・通常生成:

  • FA‑3 条件: 冒頭で1回
    [xformers] memory_efficient_attention: selected kernel = fa3F@...

  • 非対応条件(マスク/タイル等): 切替時に1回
    [xformers] memory_efficient_attention: selected kernel = cutlassF-pt

SAノード使用:

  • 毎世代冒頭で1回
    SageAttention kernel is being used for this generation.

  • auto 指定の初回のみ
    [SageAttention][DEBUG] sageattn (auto) called


ロールバック

  • attention.py: if-elif 連鎖を元の順序(Sage → xformers → Flash → PyTorch)に戻す

  • xformers/init.py: 挿入ブロックを削除(またはコメントアウト)

  • KJNodes: 上記 A〜D を削除(またはコメントアウト)

以上。これ以外の変更は不要。これで公式から最終形態を毎回確実に再現できます。

Fixed CCSR

上を適用した時、CCSRロード時にFA-3ロードログが連続する不具合が発生した為、修正しました。

修正ファイル

ComfyUI-CCSR xformersログ制御修正内容

概要

CCSRノード実行中のxformersの連続ログを抑制し、進行バーを正常に動作させるための修正

修正ファイル

`ComfyUI/custom_nodes/comfyui-ccsr/nodes.py`

追加されたインポート

import sys
import logging

新規追加クラス: XFormersKernelOnce

クラス構造

class XFormersKernelOnce:
    """ロガー+stdout/errを収集モードにし、終了時に1行だけ出す"""
    def __init__(self):
        self._filter = self._XFormersKernelFilter()
        self._loggers = []
        self._saved_out = None
        self._saved_err = None
        self._proxy_out = None
        self._proxy_err = None
        self._agg = self._KernelAggregator()

内部クラス1: カーネル集約器

class _KernelAggregator:
    """カーネル選択を記録・集約"""
    def __init__(self):
        self._kernels = []
        self._fa2_seen = False
    
    def record(self, kernel):
        if "fa2" in kernel.lower():
            self._fa2_seen = True
        self._kernels.append(kernel)
    
    def selected(self):
        if self._fa2_seen:
            # FA2が含まれていればFA2を優先
            for k in self._kernels:
                if "fa2" in k.lower():
                    return k
        # 最後に観測したカーネルを返す
        return self._kernels[-1] if self._kernels else None

内部クラス2: ロガーフィルター

class _XFormersKernelFilter(logging.Filter):
    """xformersのカーネル選択ログを捕捉して抑止"""
    def filter(self, record):
        try:
            msg = record.getMessage()
        except Exception:
            return True
        if "memory_efficient_attention: selected kernel" in msg:
            return False  # ログを抑止
        return True

内部クラス3: 標準出力プロキシ

class _StdoutProxy:
    """stdout/stderrをラップして対象行のみ収集して抑止"""
    def __init__(self, underlying, agg):
        self._u = underlying
        self._agg = agg
    
    def write(self, s):
        try:
            text = str(s)
        except Exception:
            text = s
        if "memory_efficient_attention: selected kernel" in text:
            # カーネル名を抽出
            if "=" in text:
                kernel = text.split("=")[-1].strip()
                self._agg.record(kernel)
            return len(s)  # 書き込み長を返す(進行バー等の整合性を維持)
        return self._u.write(s)
    
    def flush(self):
        return self._u.flush()
    
    # 以降は透過委譲
    def fileno(self): return self._u.fileno() if hasattr(self._u, "fileno") else 1
    def isatty(self): return self._u.isatty() if hasattr(self._u, "isatty") else False
    def readable(self): return self._u.readable() if hasattr(self._u, "readable") else False
    def writable(self): return self._u.writable() if hasattr(self._u, "writable") else True
    def seekable(self): return self._u.seekable() if hasattr(self._u, "seekable") else False
    @property
    def encoding(self): return getattr(self._u, "encoding", "utf-8")
    @property
    def errors(self): return getattr(self._u, "errors", None)
    @property
    def buffer(self): return getattr(self._u, "buffer", None)
    def __getattr__(self, name): return getattr(self._u, name)

コンテキストマネージャー開始時

def __enter__(self):
    # よく使われるロガーに一括装着
    self._loggers = [
        logging.getLogger(),
        logging.getLogger("xformers"),
        logging.getLogger("xformers.ops"),
        logging.getLogger("xformers.ops.fmha"),
        logging.getLogger("xformers_attention_log"),
    ]
    for lg in self._loggers:
        try: 
            lg.addFilter(self._filter)
        except Exception: 
            pass
    
    # stdout/errをプロキシに差し替え(より強力に)
    self._saved_out, self._saved_err = sys.stdout, sys.stderr
    self._proxy_out = self._StdoutProxy(self._saved_out, self._agg)
    self._proxy_err = self._StdoutProxy(self._saved_err, self._agg)
    
    # グローバルに設定
    sys.stdout = self._proxy_out
    sys.stderr = self._proxy_err
    
    # さらに、xformersの内部ロガーも制御
    try:
        import xformers.ops.fmha
        if hasattr(xformers.ops.fmha, '_memory_efficient_attention_forward'):
            # 元の関数を保存
            if not hasattr(xformers.ops.fmha, '_original_memory_efficient_attention_forward'):
                xformers.ops.fmha._original_memory_efficient_attention_forward = xformers.ops.fmha._memory_efficient_attention_forward
            
            # ログ出力を無効化した関数で置き換え
            def _silent_memory_efficient_attention_forward(inp, op=None):
                inp.validate_inputs()
                output_shape = inp.normalize_bmhk()
                if op is None:
                    op = xformers.ops.fmha._dispatch_fw(inp, False)
                    # ログ出力を無効化
                    if not hasattr(_silent_memory_efficient_attention_forward, "_last_kernel"):
                        _silent_memory_efficient_attention_forward._last_kernel = None
                    last_kernel = _silent_memory_efficient_attention_forward._last_kernel
                    current = getattr(op, "NAME", str(op))
                    if current != last_kernel:
                        # ログ出力を無効化し、集約器に記録のみ
                        if "fa2" in str(current).lower():
                            self._agg._fa2_seen = True
                        self._agg._kernels.append(str(current))
                        _silent_memory_efficient_attention_forward._last_kernel = current
                else:
                    xformers.ops.fmha._ensure_op_supports_or_raise(ValueError, "memory_efficient_attention", op, inp)
                out, *_ = op.apply(inp, needs_gradient=False)
                return out.reshape(output_shape)
            
            xformers.ops.fmha._memory_efficient_attention_forward = _silent_memory_efficient_attention_forward
            
            # requires_grad版もパッチ
            if hasattr(xformers.ops.fmha, '_memory_efficient_attention_forward_requires_grad'):
                if not hasattr(xformers.ops.fmha, '_original_memory_efficient_attention_forward_requires_grad'):
                    xformers.ops.fmha._original_memory_efficient_attention_forward_requires_grad = xformers.ops.fmha._memory_efficient_attention_forward_requires_grad
                
                def _silent_memory_efficient_attention_forward_requires_grad(inp, op=None):
                    inp.validate_inputs()
                    output_shape = inp.normalize_bmhk()
                    if op is None:
                        op = xformers.ops.fmha._dispatch_fw(inp, True)
                        # ログ出力を無効化
                        if not hasattr(_silent_memory_efficient_attention_forward_requires_grad, "_last_kernel"):
                            _silent_memory_efficient_attention_forward_requires_grad._last_kernel = None
                        last_kernel = _silent_memory_efficient_attention_forward_requires_grad._last_kernel
                        current = getattr(op, "NAME", str(op))
                        if current != last_kernel:
                            # ログ出力を無効化し、集約器に記録のみ
                            if "fa2" in str(current).lower():
                                self._agg._fa2_seen = True
                            self._agg._kernels.append(str(current))
                            _silent_memory_efficient_attention_forward_requires_grad._last_kernel = current
                    else:
                        xformers.ops.fmha._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])
                
                xformers.ops.fmha._memory_efficient_attention_forward_requires_grad = _silent_memory_efficient_attention_forward_requires_grad
        except Exception as e:
            print(f"Warning: Could not patch xformers: {e}")
    
    return self

コンテキストマネージャー終了時

def __exit__(self, exc_type, exc_val, exc_tb):
    # まず復元(副作用を残さない)
    if self._saved_out is not None: 
        sys.stdout = self._saved_out
    if self._saved_err is not None: 
        sys.stderr = self._saved_err
    
    for lg in self._loggers:
        try: 
            lg.removeFilter(self._filter)
        except Exception: 
            pass
    
    # xformersのパッチを元に戻す
    try:
        import xformers.ops.fmha
        if hasattr(xformers.ops.fmha, '_original_memory_efficient_attention_forward'):
            xformers.ops.fmha._memory_efficient_attention_forward = xformers.ops.fmha._original_memory_efficient_attention_forward
        if hasattr(xformers.ops.fmha, '_original_memory_efficient_attention_forward_requires_grad'):
            xformers.ops.fmha._memory_efficient_attention_forward_requires_grad = xformers.ops.fmha._original_memory_efficient_attention_forward_requires_grad
    except Exception:
        pass
    
    # 観測結果から1行だけ表示(FA2優先→なければ最後に観測)
    selected = self._agg.selected()
    if selected:
        print(f"[xformers] memory_efficient_attention: selected kernel = {selected} (CCSR)")

CCSRノードでの適用箇所

修正前

autocast_condition = dtype == torch.float16 and not mm.is_device_mps(device)
with torch.autocast(mm.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
    for i in range(B):
        # ... サンプリング処理

修正後

autocast_condition = dtype == torch.float16 and not mm.is_device_mps(device)

# xformersのログを制御(CCSRブロック内限定)
with XFormersKernelOnce():
    with torch.autocast(mm.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
        for i in range(B):
            # ... サンプリング処理

技術的特徴

1. 多重制御

  • ロガーフィルター: `logging`経由のログを制御

  • 標準出力プロキシ: `print`文を直接制御

  • 関数パッチ: xformersの内部実装を置き換え

2. 安全性

  • 副作用なし: 処理完了後に必ず元の状態を復元

  • 例外処理: 各段階でエラーが発生しても安全に処理

  • 進行バー保護: 出力長を適切に返して進行バーの整合性を維持

3. 効率性

  • FA2優先: Flash Attention 2が使用された場合は優先的に記録

  • 重複排除: 同じカーネルが連続で選択された場合は記録しない

  • 集約出力: 処理完了後に1行でまとめて出力

効果

  • CCSRノード実行中のxformersの連続ログが完全に抑制

  • 進行バーが正常に動作

  • 最後に1行だけ「selected kernel = [カーネル名] (CCSR)」で出力



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