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:
kijai
2025-10-20 15:16:43 +03:00
parent 8081e1337c
commit 200f6943e3
2 changed files with 11 additions and 0 deletions
+1
View File
@@ -1002,6 +1002,7 @@ class WanVideoModelLoader:
"sageattn",
"sageattn_3",
"radial_sage_attention",
"sageattn_compiled",
], {"default": "sdpa"}),
"compile_args": ("WANCOMPILEARGS", ),
"block_swap_args": ("BLOCKSWAPARGS", ),
+10
View File
@@ -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()