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 = FalseB) 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 funcC) 世代ごとに一回だけ(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)」で出力
