make torch compile work better

This commit is contained in:
kijai
2024-10-26 02:33:29 +03:00
parent 25f16462aa
commit dcca095743
3 changed files with 35 additions and 37 deletions
+2 -2
View File
@@ -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,