diff --git a/comfy/cli_args.py b/comfy/cli_args.py index ee9e1ce9f..9de244087 100644 --- a/comfy/cli_args.py +++ b/comfy/cli_args.py @@ -149,6 +149,7 @@ attn_group.add_argument("--use-quad-cross-attention", action="store_true", help= attn_group.add_argument("--use-pytorch-cross-attention", action="store_true", help="Use the new pytorch 2.0 cross attention function.") attn_group.add_argument("--use-sage-attention", action="store_true", help="Use sage attention.") attn_group.add_argument("--use-flash-attention", action="store_true", help="Use FlashAttention.") +attn_group.add_argument("--use-ck-attention", action="store_true", help="Use Comfy Kitchen attention.") parser.add_argument("--disable-xformers", action="store_true", help="Disable xformers.") diff --git a/comfy/ldm/minimax/model.py b/comfy/ldm/minimax/model.py index bc06288ab..76174483a 100644 --- a/comfy/ldm/minimax/model.py +++ b/comfy/ldm/minimax/model.py @@ -25,7 +25,7 @@ import comfy.model_prefetch import comfy.ops import comfy.patcher_extension import comfy.quant_ops -from comfy.ldm.modules.attention import optimized_attention +from comfy.ldm.modules.attention import AttentionTensorContainer, optimized_attention FRAME_PER_TOKEN = (1, 4, 4, 4, 4) FRAME_RESCALE = 5.0 / 3.0 @@ -165,9 +165,9 @@ class Attention(nn.Module): else: q = self.q_norm(q.view(s, self.heads, self.head_dim)) k = self.k_norm(k.view(s, self.heads, self.head_dim)) - q = q.transpose(0, 1).unsqueeze(0) - k = k.transpose(0, 1).unsqueeze(0) - v = v.transpose(0, 1).unsqueeze(0) + q = AttentionTensorContainer(q.transpose(0, 1).unsqueeze(0)) + k = AttentionTensorContainer(k.transpose(0, 1).unsqueeze(0)) + v = AttentionTensorContainer(v.transpose(0, 1).unsqueeze(0)) out = optimized_attention(q, k, v, self.heads, mask=None, skip_reshape=True, transformer_options=transformer_options) return self.out_proj(out.squeeze(0)) diff --git a/comfy/ldm/modules/attention.py b/comfy/ldm/modules/attention.py index 2c549e095..b22d03d77 100644 --- a/comfy/ldm/modules/attention.py +++ b/comfy/ldm/modules/attention.py @@ -10,6 +10,8 @@ from typing import Optional, Any, Callable, Union import logging import functools +import comfy_kitchen + from .diffusionmodules.util import AlphaBlender, timestep_embedding from .sub_quadratic_attention import efficient_dot_product_attention @@ -49,6 +51,8 @@ except ImportError: logging.error(f"\n\nTo use the `--use-flash-attention` feature, the `flash-attn` package must be installed first.\ncommand:\n\t{sys.executable} -m pip install flash-attn") exit(-1) +COMFY_KITCHEN_INT8_ATTENTION_IS_AVAILABLE = comfy_kitchen.int8_attention_is_available() + REGISTERED_ATTENTION_FUNCTIONS = {} def register_attention_function(name: str, func: Callable): # avoid replacing existing functions @@ -145,9 +149,34 @@ def Normalize(in_channels, dtype=None, device=None): return torch.nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True, dtype=dtype, device=device) +class AttentionTensorContainer: + """Single-owner tensor input consumed by an optimized attention backend.""" + + __slots__ = ("tensor",) + + def __init__(self, tensor: torch.Tensor): + self.tensor: torch.Tensor | None = tensor + + def peek(self) -> torch.Tensor: + if self.tensor is None: + raise RuntimeError("attention tensor container has already been consumed") + return self.tensor + + def take(self) -> torch.Tensor: + tensor = self.peek() + self.tensor = None + return tensor + + def wrap_attn(func): @functools.wraps(func) def wrapper(*args, **kwargs): + containers = None + if len(args) >= 3 and isinstance(args[0], AttentionTensorContainer): + if not isinstance(args[1], AttentionTensorContainer) or not isinstance(args[2], AttentionTensorContainer): + raise TypeError("q, k, and v must all be attention tensor containers") + containers = args[:3] + remove_attn_wrapper_key = False try: if "_inside_attn_wrapper" not in kwargs: @@ -156,11 +185,22 @@ def wrap_attn(func): kwargs["_inside_attn_wrapper"] = True if transformer_options is not None: if "optimized_attention_override" in transformer_options: - return transformer_options["optimized_attention_override"](func, *args, **kwargs) + optimized_attention_override = transformer_options["optimized_attention_override"] + if containers is not None: + if hasattr(optimized_attention_override, "container_function"): + return optimized_attention_override.container_function(*args, **kwargs) + args = tuple(container.take() for container in containers) + args[3:] + return optimized_attention_override(func, *args, **kwargs) + + if containers is not None: + if wrapper.container_function is not None: + return wrapper.container_function(*args, **kwargs) + args = tuple(container.take() for container in containers) + args[3:] return func(*args, **kwargs) finally: if remove_attn_wrapper_key: del kwargs["_inside_attn_wrapper"] + wrapper.container_function = None return wrapper @wrap_attn @@ -545,6 +585,63 @@ def attention_pytorch(q, k, v, heads, mask=None, attn_precision=None, skip_resha ).transpose(1, 2).reshape(-1, q.shape[2], heads * dim_head) return out +def _comfy_kitchen_int8_inputs(q, k, v, heads, mask, skip_reshape, enable_gqa): + dim_head = q.shape[-1] if skip_reshape else q.shape[-1] // heads + b = q.shape[0] + if not skip_reshape: + q, k, v = _reshape_qkv_to_heads(q, k, v, b, heads, dim_head, enable_gqa, expand_kv=False) + q, k, v = map(lambda t: t.transpose(1, 2), (q, k, v)) + + if mask is not None: + if mask.ndim == 2: + mask = mask.unsqueeze(0) + if mask.ndim == 3: + mask = mask.unsqueeze(1) + + return q, k, v, mask, b, dim_head + + +@wrap_attn +def attention_comfy_kitchen_int8(q, k, v, heads, mask=None, attn_precision=None, skip_reshape=False, skip_output_reshape=False, **kwargs): + q, k, v, mask, b, dim_head = _comfy_kitchen_int8_inputs( + q, k, v, heads, mask, skip_reshape, kwargs.get("enable_gqa", False) + ) + out = comfy_kitchen.int8_attention( + q, + k, + v, + scale=kwargs.get("scale", None), + attn_mask=mask, + ) + if not skip_output_reshape: + out = out.transpose(1, 2).reshape(b, -1, heads * dim_head) + return out + + +def _attention_comfy_kitchen_int8_containers(q, k, v, heads, mask=None, attn_precision=None, skip_reshape=False, skip_output_reshape=False, **kwargs): + q = q.take() + k = k.take() + v = v.take() + q, k, v, mask, b, dim_head = _comfy_kitchen_int8_inputs( + q, k, v, heads, mask, skip_reshape, kwargs.get("enable_gqa", False) + ) + quantized = comfy_kitchen.prequantize_int8_attention( + q, + k, + v, + scale=kwargs.get("scale", None), + attn_mask=mask, + ) + del q, k, v + out = comfy_kitchen.int8_attention_from_prequantized(quantized) + if not skip_output_reshape: + out = out.transpose(1, 2).reshape(b, -1, heads * dim_head) + return out + + +attention_comfy_kitchen_int8.container_function = _attention_comfy_kitchen_int8_containers + + @wrap_attn def attention_sage(q, k, v, heads, mask=None, attn_precision=None, skip_reshape=False, skip_output_reshape=False, **kwargs): if kwargs.get("low_precision_attention", True) is False or (mask is not None and not SAGE_ATTENTION_SUPPORTS_MASK): @@ -775,10 +872,20 @@ 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 +if model_management.comfy_kitchen_attention_enabled(): + if COMFY_KITCHEN_INT8_ATTENTION_IS_AVAILABLE: + logging.info("Using Comfy Kitchen attention") + optimized_attention = attention_comfy_kitchen_int8 + else: + logging.error("Comfy Kitchen attention is unavailable. Install a Comfy Kitchen build with attention support to use --use-ck-attention.") + exit(-1) + optimized_attention_masked = optimized_attention # register core-supported attention functions +if COMFY_KITCHEN_INT8_ATTENTION_IS_AVAILABLE: + register_attention_function("comfy_kitchen_int8", attention_comfy_kitchen_int8) if SAGE_ATTENTION_IS_AVAILABLE: register_attention_function("sage", attention_sage) if SAGE_ATTENTION3_IS_AVAILABLE: diff --git a/comfy/model_management.py b/comfy/model_management.py index 9f8e7f07b..65599424b 100644 --- a/comfy/model_management.py +++ b/comfy/model_management.py @@ -1658,6 +1658,9 @@ def unpin_memory(tensor): def sage_attention_enabled(): return args.use_sage_attention +def comfy_kitchen_attention_enabled(): + return args.use_ck_attention + def flash_attention_enabled(): return args.use_flash_attention diff --git a/comfy/model_patcher.py b/comfy/model_patcher.py index ae3f0191d..cb44e7394 100644 --- a/comfy/model_patcher.py +++ b/comfy/model_patcher.py @@ -685,6 +685,14 @@ class ModelPatcher: def set_model_attn2_output_patch(self, patch): self.set_model_patch(patch, "attn2_output_patch") + def set_model_optimized_attention(self, optimized_attention): + def optimized_attention_override(_, *args, **kwargs): + return optimized_attention(*args, **kwargs) + + if hasattr(optimized_attention, "container_function") and optimized_attention.container_function is not None: + optimized_attention_override.container_function = optimized_attention.container_function + self.model_options["transformer_options"]["optimized_attention_override"] = optimized_attention_override + def set_model_input_block_patch(self, patch): self.set_model_patch(patch, "input_block_patch") diff --git a/comfy_extras/nodes_model_advanced.py b/comfy_extras/nodes_model_advanced.py index a336ba079..21ea82148 100644 --- a/comfy_extras/nodes_model_advanced.py +++ b/comfy_extras/nodes_model_advanced.py @@ -1,6 +1,9 @@ +import logging + import comfy.sd import comfy.model_sampling import comfy.latent_formats +import comfy.ldm.modules.attention import nodes import torch import node_helpers @@ -346,6 +349,39 @@ class ModelComputeDtype: return (m, ) +class ModelAttentionBackend: + @classmethod + def INPUT_TYPES(s): + backends = ["pytorch attention"] + if comfy.ldm.modules.attention.COMFY_KITCHEN_INT8_ATTENTION_IS_AVAILABLE: + backends.append("comfy kitchen attention") + return {"required": {"model": ("MODEL",), + "attention": (backends,), + }} + + @classmethod + def VALIDATE_INPUTS(s, attention): + return True + + RETURN_TYPES = ("MODEL",) + FUNCTION = "patch" + + CATEGORY = "model/patch" + + def patch(self, model, attention): + attention_name = { + "comfy kitchen attention": "comfy_kitchen_int8", + "pytorch attention": "pytorch", + }.get(attention) + attention_function = comfy.ldm.modules.attention.get_attention_function(attention_name, None) + if attention_function is None: + logging.warning("Attention backend '%s' is unavailable; using PyTorch attention.", attention) + attention_function = comfy.ldm.modules.attention.get_attention_function("pytorch") + m = model.clone() + m.set_model_optimized_attention(attention_function) + return (m, ) + + NODE_CLASS_MAPPINGS = { "ModelSamplingDiscrete": ModelSamplingDiscrete, "ModelSamplingContinuousEDM": ModelSamplingContinuousEDM, @@ -357,4 +393,5 @@ NODE_CLASS_MAPPINGS = { "ModelNoiseScale": ModelNoiseScale, "RescaleCFG": RescaleCFG, "ModelComputeDtype": ModelComputeDtype, + "ModelAttentionBackend": ModelAttentionBackend, } diff --git a/requirements.txt b/requirements.txt index 94cd1c5eb..25eaf8bc7 100644 --- a/requirements.txt +++ b/requirements.txt @@ -22,7 +22,7 @@ alembic SQLAlchemy>=2.0.0 filelock av>=16.0.0 -comfy-kitchen==0.2.28 +comfy-kitchen==0.2.30 comfy-aimdo==0.4.13 requests simpleeval>=1.0.0