@@ -14,6 +14,7 @@ import torch.nn as nn
|
|||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
from einops import rearrange
|
from einops import rearrange
|
||||||
|
|
||||||
|
import comfy.ldm.common_dit
|
||||||
from .utils import to_2tuple
|
from .utils import to_2tuple
|
||||||
|
|
||||||
sdpa_32b = None
|
sdpa_32b = None
|
||||||
@@ -25,11 +26,14 @@ Q_4GB_LIMIT = 32000000
|
|||||||
|
|
||||||
from comfy import model_management
|
from comfy import model_management
|
||||||
if model_management.xformers_enabled():
|
if model_management.xformers_enabled():
|
||||||
import xformers
|
|
||||||
import xformers.ops
|
import xformers.ops
|
||||||
|
if int((xformers.__version__).split(".")[2]) >= 28:
|
||||||
|
block_diagonal_mask_from_seqlens = xformers.ops.fmha.attn_bias.BlockDiagonalMask.from_seqlens
|
||||||
|
else:
|
||||||
|
block_diagonal_mask_from_seqlens = xformers.ops.fmha.BlockDiagonalMask.from_seqlens
|
||||||
else:
|
else:
|
||||||
if model_management.xpu_available:
|
if model_management.xpu_available:
|
||||||
import intel_extension_for_pytorch as ipex
|
import intel_extension_for_pytorch as ipex # type: ignore
|
||||||
import os
|
import os
|
||||||
if not torch.xpu.has_fp64_dtype() and not os.environ.get('IPEX_FORCE_ATTENTION_SLICE', None):
|
if not torch.xpu.has_fp64_dtype() and not os.environ.get('IPEX_FORCE_ATTENTION_SLICE', None):
|
||||||
from ...utils.IPEX.attention import scaled_dot_product_attention_32_bit
|
from ...utils.IPEX.attention import scaled_dot_product_attention_32_bit
|
||||||
@@ -70,7 +74,7 @@ class MultiHeadCrossAttention(nn.Module):
|
|||||||
if model_management.xformers_enabled():
|
if model_management.xformers_enabled():
|
||||||
attn_bias = None
|
attn_bias = None
|
||||||
if mask is not None:
|
if mask is not None:
|
||||||
attn_bias = xformers.ops.fmha.BlockDiagonalMask.from_seqlens([N] * B, mask)
|
attn_bias = block_diagonal_mask_from_seqlens([N] * B, mask)
|
||||||
x = xformers.ops.memory_efficient_attention(
|
x = xformers.ops.memory_efficient_attention(
|
||||||
q, k, v,
|
q, k, v,
|
||||||
p=self.attn_drop.p,
|
p=self.attn_drop.p,
|
||||||
|
|||||||
@@ -39,8 +39,11 @@ Q_4GB_LIMIT = 32000000
|
|||||||
|
|
||||||
from comfy import model_management
|
from comfy import model_management
|
||||||
if model_management.xformers_enabled():
|
if model_management.xformers_enabled():
|
||||||
import xformers
|
|
||||||
import xformers.ops
|
import xformers.ops
|
||||||
|
if int((xformers.__version__).split(".")[2]) >= 28:
|
||||||
|
block_diagonal_mask_from_seqlens = xformers.ops.fmha.attn_bias.BlockDiagonalMask.from_seqlens
|
||||||
|
else:
|
||||||
|
block_diagonal_mask_from_seqlens = xformers.ops.fmha.BlockDiagonalMask.from_seqlens
|
||||||
else:
|
else:
|
||||||
if model_management.xpu_available:
|
if model_management.xpu_available:
|
||||||
import intel_extension_for_pytorch as ipex # type: ignore
|
import intel_extension_for_pytorch as ipex # type: ignore
|
||||||
@@ -94,7 +97,7 @@ class MultiHeadCrossAttention(nn.Module):
|
|||||||
if model_management.xformers_enabled():
|
if model_management.xformers_enabled():
|
||||||
attn_bias = None
|
attn_bias = None
|
||||||
if mask is not None:
|
if mask is not None:
|
||||||
attn_bias = xformers.ops.fmha.BlockDiagonalMask.from_seqlens([N] * B, mask)
|
attn_bias = block_diagonal_mask_from_seqlens([N] * B, mask)
|
||||||
x = xformers.ops.memory_efficient_attention(
|
x = xformers.ops.memory_efficient_attention(
|
||||||
q, k, v,
|
q, k, v,
|
||||||
p=self.attn_drop.p,
|
p=self.attn_drop.p,
|
||||||
|
|||||||
Reference in New Issue
Block a user