Clearer error on attention mode imports

This commit is contained in:
kijai
2025-12-05 02:00:31 +02:00
parent e2333d0f04
commit 4e7b8dd92c
+10 -6
View File
@@ -26,6 +26,8 @@ try:
# Return tensor with same shape as q
return q.clone()
sageattn_func = torch.ops.wanvideo.sageattn
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)
@@ -39,7 +41,7 @@ except Exception as e:
log.warning("sageattention package is not installed, sageattention will not be available")
elif isinstance(e, ImportError) and "DLL" in str(e):
log.warning("sageattention DLL loading error, sageattention will not be available")
sageattn_func = None
sageattn_func = attention_func_error
try:
from sageattention import sageattn_varlen
@@ -60,7 +62,8 @@ try:
def _(q, k, v, q_lens, k_lens, max_seqlen_q, max_seqlen_k, dropout_p=0.0, is_causal=False):
# Return tensor with same shape as q
return q.clone()
except Exception as e:
sageattn_varlen_func = torch.ops.wanvideo.sageattn_varlen
except:
sageattn_varlen_func = attention_func_error
# sage3
@@ -82,8 +85,9 @@ try:
@sageattn_func_ultravico.register_fake
def _(qkv, attn_mask=None, dropout_p=0.0, is_causal=False, multi_factor=0.9):
torch.empty_like(qkv[0]).contiguous()
sageattn_func_ultravico = torch.ops.wanvideo.sageattn_ultravico
except:
sageattn_ultravico = attention_func_error
sageattn_func_ultravico = attention_func_error
def attention(q, k, v, q_lens=None, k_lens=None, max_seqlen_q=None, max_seqlen_k=None, dropout_p=0.,
@@ -96,13 +100,13 @@ def attention(q, k, v, q_lens=None, k_lens=None, max_seqlen_q=None, max_seqlen_k
elif attention_mode == 'sageattn_3':
return sageattn_blackwell(q.transpose(1,2), k.transpose(1,2), v.transpose(1,2), per_block_mean=False).transpose(1,2).contiguous()
elif attention_mode == 'sageattn_varlen':
return torch.ops.wanvideo.sageattn_varlen(q,k,v, q_lens=q_lens, k_lens=k_lens, max_seqlen_k=max_seqlen_k, max_seqlen_q=max_seqlen_q)
return sageattn_varlen_func(q,k,v, q_lens=q_lens, k_lens=k_lens, max_seqlen_k=max_seqlen_k, max_seqlen_q=max_seqlen_q)
elif attention_mode == 'sageattn_compiled': # for sage versions that allow torch.compile, may be redundant now as other sageattn ops are wrapper in custom ops
return sageattn_func_compiled(q, k, v, tensor_layout="NHD").contiguous()
elif attention_mode == 'sageattn':
return torch.ops.wanvideo.sageattn(q, k, v, tensor_layout="NHD").contiguous()
return sageattn_func(q, k, v, tensor_layout="NHD").contiguous()
elif attention_mode == 'sageattn_ultravico':
return torch.ops.wanvideo.sageattn_ultravico([q, k, v], multi_factor=multi_factor).contiguous()
return sageattn_func_ultravico([q, k, v], multi_factor=multi_factor).contiguous()
else: # sdpa
if not (q.dtype == k.dtype == v.dtype):
return torch.nn.functional.scaled_dot_product_attention(q.transpose(1, 2), k.transpose(1, 2).to(q.dtype), v.transpose(1, 2).to(q.dtype), attn_mask=attn_mask).transpose(1, 2).contiguous()