more sdpa updates

This commit is contained in:
Jukka Seppänen
2024-09-30 14:49:09 +03:00
parent e39ecb6f88
commit 081c818002
2 changed files with 14 additions and 22 deletions
+13 -19
View File
@@ -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
+1 -3
View File
@@ -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}),