From 0c54c235cb2645b0db7e2897830cbf7b0e718983 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sat, 12 Oct 2024 15:43:34 -0400 Subject: [PATCH] add sageattention Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 6 ++++++ modules/devices.py | 16 +++++++++------- modules/intel/ipex/diffusers.py | 3 +-- modules/shared.py | 2 +- 4 files changed, 17 insertions(+), 10 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 5e77b46eb..f709bba5f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -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 diff --git a/modules/devices.py b/modules/devices.py index 1274bf82b..a572252f0 100644 --- a/modules/devices.py +++ b/modules/devices.py @@ -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 diff --git a/modules/intel/ipex/diffusers.py b/modules/intel/ipex/diffusers.py index 9f17b266c..f742fe5c0 100644 --- a/modules/intel/ipex/diffusers.py +++ b/modules/intel/ipex/diffusers.py @@ -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 diff --git a/modules/shared.py b/modules/shared.py index 152060e5b..175b8e7a8 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -449,7 +449,7 @@ options_templates.update(options_section(('cuda', "Compute Settings"), { "cross_attention_sep": OptionInfo("

Cross Attention

", "", 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("

Sub-quadratic options

", "", gr.HTML, {"visible": not native}),