Possible bandaid for q k v dtype mismatch
Still no clue where or why it happens, but clearly is happening for some when using sage with Stand-In
This commit is contained in:
@@ -20,6 +20,8 @@ try:
|
||||
from sageattention import sageattn
|
||||
@torch.compiler.disable()
|
||||
def sageattn_func(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)
|
||||
if 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:
|
||||
|
||||
Reference in New Issue
Block a user