Separate sage 1.x/2.x and sage 3.x on different functions and flags 1.

This commit is contained in:
Panchovix 2025-07-26 14:21:13 -04:00 committed by GitHub
parent 6476f0bac6
commit 4a9235007d
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194

View File

@ -18,25 +18,25 @@ if model_management.xformers_enabled():
import xformers.ops import xformers.ops
if model_management.sage_attention_enabled(): if model_management.sage_attention_enabled():
sage_attention_available = False try:
SAGE_ATTENTION_3_AVAILABLE = False from sageattention import sageattn
logging.info("Found SageAttention 1.x/2.x (sageattention package)")
except ModuleNotFoundError as e:
if e.name == "sageattention":
logging.error(f"\n\nTo use the `--use-sage-attention` feature, the `sageattention` package must be installed first.\ncommand:\n\t{sys.executable} -m pip install sageattention")
else:
raise e
exit(-1)
if model_management.sage_attention3_enabled():
try: try:
from sageattn import sageattn_blackwell from sageattn import sageattn_blackwell
SAGE_ATTENTION_3_AVAILABLE = True
sage_attention_available = True
logging.info("Found SageAttention3 (sageattn package)") logging.info("Found SageAttention3 (sageattn package)")
except ModuleNotFoundError as e:
except ImportError: if e.name == "sageattn":
try: logging.error(f"\n\nTo use the `--use-sage-attention3` feature, the `sageattn` package must be installed first.\ncommand:\n\t{sys.executable} -m pip install sageattn")
from sageattention import sageattn else:
sage_attention_available = True raise e
logging.info("Found SageAttention2 (sageattention package)")
except ModuleNotFoundError:
pass
if not sage_attention_available:
logging.error(f"\n\nTo use the `--use-sage-attention` feature, the `sageattention` package must be installed first.\ncommand:\n\t{sys.executable} -m pip install sageattention")
exit(-1) exit(-1)
if model_management.flash_attention_enabled(): if model_management.flash_attention_enabled():
@ -504,39 +504,52 @@ def attention_sage(q, k, v, heads, mask=None, attn_precision=None, skip_reshape=
mask = mask.unsqueeze(1) mask = mask.unsqueeze(1)
try: try:
if SAGE_ATTENTION_3_AVAILABLE and dim_head < 256: out = sageattn(q, k, v, attn_mask=mask, is_causal=False, tensor_layout=tensor_layout)
# SageAttention3 expects tensor layout as (batch, heads, seq_len, head_dim) except Exception as e:
if tensor_layout == "NHD": logging.error("Error running sage attention: {}, using pytorch attention instead.".format(e))
q_sa3, k_sa3, v_sa3 = map(lambda t: t.transpose(1, 2), (q, k, v)) if tensor_layout == "NHD":
else: q, k, v = map(
q_sa3, k_sa3, v_sa3 = q, k, v lambda t: t.transpose(1, 2),
(q, k, v),
out = sageattn_blackwell(q_sa3, k_sa3, v_sa3, attn_mask=mask, is_causal=False, per_block_mean=False) )
return attention_pytorch(q, k, v, heads, mask=mask, skip_reshape=True, skip_output_reshape=skip_output_reshape)
# Convert back to expected layout if tensor_layout == "HND":
if tensor_layout == "HND": if not skip_output_reshape:
if not skip_output_reshape: out = (
out = out.transpose(1, 2).reshape(b, -1, heads * dim_head) out.transpose(1, 2).reshape(b, -1, heads * dim_head)
else: )
if skip_output_reshape: else:
out = out.transpose(1, 2) if skip_output_reshape:
else: out = out.transpose(1, 2)
out = out.transpose(1, 2).reshape(b, -1, heads * dim_head)
elif not SAGE_ATTENTION_3_AVAILABLE:
# Fall back to SageAttention2 if available
out = sageattn(q, k, v, attn_mask=mask, is_causal=False, tensor_layout=tensor_layout)
if tensor_layout == "HND":
if not skip_output_reshape:
out = (
out.transpose(1, 2).reshape(b, -1, heads * dim_head)
)
else:
if skip_output_reshape:
out = out.transpose(1, 2)
else:
out = out.reshape(b, -1, heads * dim_head)
else: else:
# SageAttention3 is available but head_dim >= 256, fall back to pytorch out = out.reshape(b, -1, heads * dim_head)
return out
def attention_sage3(q, k, v, heads, mask=None, attn_precision=None, skip_reshape=False, skip_output_reshape=False):
if skip_reshape:
b, _, _, dim_head = q.shape
tensor_layout = "HND"
else:
b, _, dim_head = q.shape
dim_head //= heads
q, k, v = map(
lambda t: t.view(b, -1, heads, dim_head),
(q, k, v),
)
tensor_layout = "NHD"
if mask is not None:
# add a batch dimension if there isn't already one
if mask.ndim == 2:
mask = mask.unsqueeze(0)
# add a heads dimension if there isn't already one
if mask.ndim == 3:
mask = mask.unsqueeze(1)
try:
if dim_head >= 256:
# SageAttention3 doesn't support head_dim >= 256, fall back to pytorch
logging.warning(f"SageAttention3 doesn't support head_dim >= 256 (got {dim_head}), falling back to pytorch attention") logging.warning(f"SageAttention3 doesn't support head_dim >= 256 (got {dim_head}), falling back to pytorch attention")
if tensor_layout == "NHD": if tensor_layout == "NHD":
q, k, v = map( q, k, v = map(
@ -544,8 +557,26 @@ def attention_sage(q, k, v, heads, mask=None, attn_precision=None, skip_reshape=
(q, k, v), (q, k, v),
) )
return attention_pytorch(q, k, v, heads, mask=mask, skip_reshape=True, skip_output_reshape=skip_output_reshape) return attention_pytorch(q, k, v, heads, mask=mask, skip_reshape=True, skip_output_reshape=skip_output_reshape)
# SageAttention3 expects tensor layout as (batch, heads, seq_len, head_dim)
if tensor_layout == "NHD":
q_sa3, k_sa3, v_sa3 = map(lambda t: t.transpose(1, 2), (q, k, v))
else:
q_sa3, k_sa3, v_sa3 = q, k, v
out = sageattn_blackwell(q_sa3, k_sa3, v_sa3, attn_mask=mask, is_causal=False, per_block_mean=False)
# Convert back to expected layout
if tensor_layout == "HND":
if not skip_output_reshape:
out = out.transpose(1, 2).reshape(b, -1, heads * dim_head)
else:
if skip_output_reshape:
out = out.transpose(1, 2)
else:
out = out.transpose(1, 2).reshape(b, -1, heads * dim_head)
except Exception as e: except Exception as e:
logging.error("Error running sage attention: {}, using pytorch attention instead.".format(e)) logging.error("Error running sage attention 3: {}, using pytorch attention instead.".format(e))
if tensor_layout == "NHD": if tensor_layout == "NHD":
q, k, v = map( q, k, v = map(
lambda t: t.transpose(1, 2), lambda t: t.transpose(1, 2),
@ -614,8 +645,11 @@ def attention_flash(q, k, v, heads, mask=None, attn_precision=None, skip_reshape
optimized_attention = attention_basic optimized_attention = attention_basic
if model_management.sage_attention_enabled(): if model_management.sage_attention3_enabled():
logging.info("Using sage attention") print("Using sage attention 3")
optimized_attention = attention_sage3
elif model_management.sage_attention_enabled():
print("Using sage attention 1.x/2.x")
optimized_attention = attention_sage optimized_attention = attention_sage
elif model_management.xformers_enabled(): elif model_management.xformers_enabled():
logging.info("Using xformers attention") logging.info("Using xformers attention")