make torch compile work better
This commit is contained in:
@@ -52,7 +52,7 @@ class CogVideoXAttnProcessor2_0:
|
||||
def __init__(self):
|
||||
if not hasattr(F, "scaled_dot_product_attention"):
|
||||
raise ImportError("CogVideoXAttnProcessor requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.")
|
||||
|
||||
@torch.compiler.disable()
|
||||
def __call__(
|
||||
self,
|
||||
attn: Attention,
|
||||
@@ -126,7 +126,7 @@ class FusedCogVideoXAttnProcessor2_0:
|
||||
def __init__(self):
|
||||
if not hasattr(F, "scaled_dot_product_attention"):
|
||||
raise ImportError("CogVideoXAttnProcessor requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.")
|
||||
|
||||
@torch.compiler.disable()
|
||||
def __call__(
|
||||
self,
|
||||
attn: Attention,
|
||||
|
||||
Reference in New Issue
Block a user