From 081c8180029b1b5eb8f416e079456311ff467c83 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jukka=20Sepp=C3=A4nen?= <40791699+kijai@users.noreply.github.com> Date: Mon, 30 Sep 2024 14:49:09 +0300 Subject: [PATCH] more sdpa updates --- models/layers.py | 32 +++++++++++++------------------- nodes.py | 4 +--- 2 files changed, 14 insertions(+), 22 deletions(-) diff --git a/models/layers.py b/models/layers.py index df1bcde..e98fc8e 100644 --- a/models/layers.py +++ b/models/layers.py @@ -12,27 +12,17 @@ ops = comfy.ops.manual_cast logpy = logging.getLogger(__name__) +backends = [] + if version.parse(torch.__version__) >= version.parse("2.0.0"): SDP_IS_AVAILABLE = True - from torch.backends.cuda import SDPBackend, sdp_kernel + from torch.nn.attention import SDPBackend, sdpa_kernel BACKEND_MAP = { - SDPBackend.MATH: { - "enable_math": True, - "enable_flash": False, - "enable_mem_efficient": False, - }, - SDPBackend.FLASH_ATTENTION: { - "enable_math": False, - "enable_flash": True, - "enable_mem_efficient": False, - }, - SDPBackend.EFFICIENT_ATTENTION: { - "enable_math": False, - "enable_flash": False, - "enable_mem_efficient": True, - }, - None: {"enable_math": True, "enable_flash": True, "enable_mem_efficient": True}, + SDPBackend.MATH: [SDPBackend.MATH], + SDPBackend.FLASH_ATTENTION: [SDPBackend.FLASH_ATTENTION], + SDPBackend.EFFICIENT_ATTENTION: [SDPBackend.EFFICIENT_ATTENTION], + None: [SDPBackend.MATH, SDPBackend.FLASH_ATTENTION, SDPBackend.EFFICIENT_ATTENTION] } else: from contextlib import nullcontext @@ -143,7 +133,9 @@ class TemporalAttention_Masked(nn.Module): q, k, v = map(lambda t: rearrange(t, "b n (h d) -> b h n d", h=h), (q, k, v)) - with sdp_kernel(**BACKEND_MAP[self.backend]): + backends.extend(BACKEND_MAP[self.backend]) + + with sdpa_kernel(backends): # print("dispatching into backend", self.backend, "q/k/v shape: ", q.shape, k.shape, v.shape) out = F.scaled_dot_product_attention( q, k, v, attn_mask=mask @@ -244,7 +236,9 @@ class ReferenceAttention(nn.Module): q, k, v = map(lambda t: rearrange(t, "b n (h d) -> b h n d", h=h), (q, k, v)) - with sdp_kernel(**BACKEND_MAP[self.backend]): + backends.extend(BACKEND_MAP[self.backend]) + + with sdpa_kernel(backends): # print("dispatching into backend", self.backend, "q/k/v shape: ", q.shape, k.shape, v.shape) out = F.scaled_dot_product_attention( q, k, v, attn_mask=mask diff --git a/nodes.py b/nodes.py index 748d860..e548dc3 100644 --- a/nodes.py +++ b/nodes.py @@ -37,7 +37,6 @@ class LoadLVCDModel: def loadmodel(self, model, precision, use_xformers): device = mm.get_torch_device() - print(device) offload_device = mm.unet_offload_device() dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] mm.soft_empty_cache() @@ -60,7 +59,6 @@ class LoadLVCDModel: config = OmegaConf.load(config_path) config.model.params.drop_first_stage_model = False config.model.params.init_from_unet = False - print(config.model.params.conditioner_config.params.emb_models[0]) if use_xformers: config.model.params.network_config.params.spatial_transformer_attn_type = 'softmax-xformers' @@ -97,7 +95,7 @@ class LVCDSampler: "LVCD_pipe": ("LVCDPIPE",), "ref_images": ("IMAGE",), "sketch_images": ("IMAGE",), - "num_frames": ("INT", {"default": 19, "min": 1, "max": 100, "step": 1}), + "num_frames": ("INT", {"default": 19, "min": 15, "max": 100, "step": 1}), "num_steps": ("INT", {"default": 25, "min": 1, "max": 100, "step": 1}), "fps_id": ("INT", {"default": 6, "min": 1, "max": 100, "step": 1}), "motion_bucket_id": ("INT", {"default": 160, "min": 0, "max": 1000, "step": 1}),