add sageattention

Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
Vladimir Mandic
2024-10-12 15:43:34 -04:00
parent 73bd0816d1
commit 0c54c235cb
4 changed files with 17 additions and 10 deletions
+6
View File
@@ -130,6 +130,12 @@ And other goodies like multiple *XYZ grid* improvements, additional *Flux Contro
- added native load mode for qint8/qint4 models
- add additional controlnets: [JasperAI](https://huggingface.co/collections/jasperai/flux1-dev-controlnets-66f27f9459d760dcafa32e08) **Depth**, **Upscaler**, **Surface**, thanks @EnragedAntelope
- [SageAttention](https://github.com/thu-ml/SageAttention)
- new 8-bit attention implementation on top of SDP that can provide acceleration for some models, thanks @Disty0
- enable in *settings -> compute settings -> sdp options -> sage attention*
- compatible with DiT-based models: e.g. *Flux.1, AuraFlow, CogVideoX*
- not compatible with UNet-based models, e.g. *SD15, SDXL*
- **dtype**
- previously `cuda_dtype` in settings defaulted to `fp16` if available
- now `cuda_type` defaults to **Auto** which executes `bf16` and `fp16` tests on startup and selects best available dtype
+9 -7
View File
@@ -2,12 +2,14 @@ import os
import sys
import time
import contextlib
from functools import wraps
import torch
from modules.errors import log, display, install
from modules.errors import log, display, install as install_traceback
from installer import install
debug = os.environ.get('SD_DEVICE_DEBUG', None) is not None
install() # traceback handler
install_traceback() # traceback handler
opts = None # initialized in get_backend to avoid circular import
args = None # initialized in get_backend to avoid circular import
cuda_ok = torch.cuda.is_available()
@@ -331,7 +333,7 @@ def set_sdpa_params():
torch.backends.cuda.enable_flash_sdp('Flash attention' in opts.sdp_options)
torch.backends.cuda.enable_mem_efficient_sdp('Memory attention' in opts.sdp_options)
torch.backends.cuda.enable_math_sdp('Math attention' in opts.sdp_options)
global sdpa_original
global sdpa_original # pylint: disable=global-statement
if sdpa_original is not None:
torch.nn.functional.scaled_dot_product_attention = sdpa_original
else:
@@ -341,7 +343,6 @@ def set_sdpa_params():
try:
# https://github.com/huggingface/diffusers/discussions/7172
from flash_attn import flash_attn_func
from functools import wraps
sdpa_pre_flash_atten = torch.nn.functional.scaled_dot_product_attention
@wraps(sdpa_pre_flash_atten)
def sdpa_flash_atten(query, key, value, attn_mask=None, dropout_p=0.0, is_causal=False, scale=None):
@@ -355,12 +356,13 @@ def set_sdpa_params():
log.error(f'ROCm Flash Attention failed: {err}')
if 'Sage attention' in opts.sdp_options:
try:
install('sageattention')
from sageattention import sageattn
sdpa_pre_sage_atten = torch.nn.functional.scaled_dot_product_attention
@wraps(sdpa_pre_sage_atten)
def sdpa_sage_atten(query, key, value, attn_mask=None, dropout_p=0.0, is_causal=False, scale=None):
if query.shape[-1] in {128, 96, 64} and attn_mask is None and query.dtype != torch.float32:
return sageattn(q=query, k=key, v=value, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, scale=scale).transpose(1, 2)
if query.shape[-1] in {128, 96, 64} and attn_mask is None and query.dtype != torch.float32:
return sageattn(q=query, k=key, v=value, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, scale=scale)
else:
return sdpa_pre_sage_atten(query=query, key=key, value=value, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, scale=scale)
torch.nn.functional.scaled_dot_product_attention = sdpa_sage_atten
@@ -369,7 +371,7 @@ def set_sdpa_params():
log.error(f'SDPA Sage Attention failed: {err}')
if 'Dynamic attention' in opts.sdp_options:
try:
global sdpa_pre_dyanmic_atten
global sdpa_pre_dyanmic_atten # pylint: disable=global-statement
sdpa_pre_dyanmic_atten = torch.nn.functional.scaled_dot_product_attention
from modules.sd_hijack_dynamic_atten import sliced_scaled_dot_product_attention
torch.nn.functional.scaled_dot_product_attention = sliced_scaled_dot_product_attention
+1 -2
View File
@@ -1,9 +1,8 @@
import os
from functools import wraps
from functools import wraps, cache
import torch
import diffusers #0.29.1 # pylint: disable=import-error
from diffusers.models.attention_processor import Attention
from functools import cache
# pylint: disable=protected-access, missing-function-docstring, line-too-long
+1 -1
View File
@@ -449,7 +449,7 @@ options_templates.update(options_section(('cuda', "Compute Settings"), {
"cross_attention_sep": OptionInfo("<h2>Cross Attention</h2>", "", gr.HTML),
"cross_attention_optimization": OptionInfo(startup_cross_attention, "Attention optimization method", gr.Radio, lambda: {"choices": shared_items.list_crossattention(native) }),
"sdp_options": OptionInfo(startup_sdp_options, "SDP options", gr.CheckboxGroup, {"choices": ['Flash attention', 'Memory attention', 'Math attention', 'Dynamic attention'] }),
"sdp_options": OptionInfo(startup_sdp_options, "SDP options", gr.CheckboxGroup, {"choices": ['Flash attention', 'Memory attention', 'Math attention', 'Dynamic attention', 'Sage attention'] }),
"xformers_options": OptionInfo(['Flash attention'], "xFormers options", gr.CheckboxGroup, {"choices": ['Flash attention'] }),
"dynamic_attention_slice_rate": OptionInfo(4, "Dynamic Attention slicing rate in GB", gr.Slider, {"minimum": 0.1, "maximum": gpu_memory, "step": 0.1, "visible": native}),
"sub_quad_sep": OptionInfo("<h3>Sub-quadratic options</h3>", "", gr.HTML, {"visible": not native}),