more sdpa updates
This commit is contained in:
+13
-19
@@ -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
|
||||
|
||||
@@ -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}),
|
||||
|
||||
Reference in New Issue
Block a user