From 200f6943e3f379eedfaf294b6d6ea180122f2a82 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Mon, 20 Oct 2025 15:16:43 +0300 Subject: [PATCH] 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 --- nodes_model_loading.py | 1 + wanvideo/modules/attention.py | 10 ++++++++++ 2 files changed, 11 insertions(+) diff --git a/nodes_model_loading.py b/nodes_model_loading.py index efe7365..124e6f8 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -1002,6 +1002,7 @@ class WanVideoModelLoader: "sageattn", "sageattn_3", "radial_sage_attention", + "sageattn_compiled", ], {"default": "sdpa"}), "compile_args": ("WANCOMPILEARGS", ), "block_swap_args": ("BLOCKSWAPARGS", ), diff --git a/wanvideo/modules/attention.py b/wanvideo/modules/attention.py index c6be1bc..d25a4d4 100644 --- a/wanvideo/modules/attention.py +++ b/wanvideo/modules/attention.py @@ -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()