Add sageattn mode that allows torch.compile
Latest wheel from woct0rdho includes the torch.compile fix: https://github.com/woct0rdho/SageAttention/releases Based on my quick testing this reduces peak VRAM usage a bit when running sageattn + torch.compile
This commit is contained in:
@@ -1002,6 +1002,7 @@ class WanVideoModelLoader:
|
||||
"sageattn",
|
||||
"sageattn_3",
|
||||
"radial_sage_attention",
|
||||
"sageattn_compiled",
|
||||
], {"default": "sdpa"}),
|
||||
"compile_args": ("WANCOMPILEARGS", ),
|
||||
"block_swap_args": ("BLOCKSWAPARGS", ),
|
||||
|
||||
@@ -26,6 +26,14 @@ try:
|
||||
return sageattn(q.to(torch.float16), k.to(torch.float16), v.to(torch.float16), attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, tensor_layout=tensor_layout).to(torch.float32)
|
||||
else:
|
||||
return sageattn(q, k, v, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, tensor_layout=tensor_layout)
|
||||
|
||||
def sageattn_func_compiled(q, k, v, attn_mask=None, dropout_p=0, is_causal=False, tensor_layout="HND"):
|
||||
if not (q.dtype == k.dtype == v.dtype):
|
||||
return sageattn(q, k.to(q.dtype), v.to(q.dtype), attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, tensor_layout=tensor_layout)
|
||||
elif q.dtype == torch.float32:
|
||||
return sageattn(q.to(torch.float16), k.to(torch.float16), v.to(torch.float16), attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, tensor_layout=tensor_layout).to(torch.float32)
|
||||
else:
|
||||
return sageattn(q, k, v, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, tensor_layout=tensor_layout)
|
||||
except Exception as e:
|
||||
log.warning(f"Warning: Could not load sageattention: {str(e)}")
|
||||
if isinstance(e, ModuleNotFoundError):
|
||||
@@ -227,5 +235,7 @@ def attention(
|
||||
max_seqlen_k=max_seqlen_k,
|
||||
max_seqlen_q=max_seqlen_q
|
||||
)
|
||||
elif attention_mode == 'sageattn_compiled':
|
||||
return sageattn_func_compiled(q, k, v, tensor_layout="NHD").contiguous()
|
||||
else:
|
||||
return sageattn_func(q, k, v, tensor_layout="NHD").contiguous()
|
||||
|
||||
Reference in New Issue
Block a user