diff --git a/CHANGELOG.md b/CHANGELOG.md index 9895daefa..900f38912 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,6 +1,6 @@ # Change Log for SD.Next -## Update for 2024-02-18 +## Update for 2024-02-19 - **Improvements**: - **IP Adapter** major refactor @@ -68,6 +68,8 @@ - refactor txt2img/img2img api - enhanced theme loader - add additional debug env variables + - enhanced sdp cross-optimization control + see *settings -> compute settings* - **Fixes**: - add variation seed to diffusers txt2img, thanks @AI-Casanova - handle extensions that install conflicting versions of packages diff --git a/TODO.md b/TODO.md index 4d73062c2..a2863e1b2 100644 --- a/TODO.md +++ b/TODO.md @@ -7,7 +7,7 @@ Main ToDo list can be found at [GitHub projects](https://github.com/users/vladma - stable cascade - init latents, variations, tiling - lora sliders: -- x-adapoter: +- x-adapter: - diffusers public callbacks - image2video: pia and vgen pipelines - video2video diff --git a/cli/run-benchmark.py b/cli/run-benchmark.py index 9bc0721f0..afa8d5709 100755 --- a/cli/run-benchmark.py +++ b/cli/run-benchmark.py @@ -35,6 +35,7 @@ async def txt2img(): log.debug({ 'info': info }) if options['batch_size'] != len(data['images']): log.error({ 'requested': options['batch_size'], 'received': len(data['images']) }) + return 0 for i in range(len(data['images'])): data['images'][i] = Image.open(io.BytesIO(base64.b64decode(data['images'][i].split(',',1)[0]))) if args.save: @@ -113,7 +114,7 @@ if __name__ == '__main__': log.info({ 'run-benchmark' }) parser = argparse.ArgumentParser(description = 'run-benchmark') parser.add_argument("--steps", type=int, default=50, required=False, help="steps") - parser.add_argument("--sampler", type=str, default='Euler a', required=False, help="max batch size") + parser.add_argument("--sampler", type=str, default='Euler a', required=False, help="Use specific sampler") parser.add_argument("--prompt", type=str, default='photo of two dice on a table', required=False, help="prompt") parser.add_argument("--negative", type=str, default='foggy, blurry', required=False, help="prompt") parser.add_argument("--maxbatch", type=int, default=16, required=False, help="max batch size") diff --git a/modules/devices.py b/modules/devices.py index e4804df8d..ff8b5f80d 100644 --- a/modules/devices.py +++ b/modules/devices.py @@ -220,6 +220,7 @@ def set_cuda_params(): pass if torch.backends.cudnn.is_available(): try: + torch.backends.cudnn.deterministic = shared.opts.cudnn_deterministic torch.backends.cudnn.benchmark = True if shared.opts.cudnn_benchmark: log.debug('Torch enable cuDNN benchmark') @@ -227,6 +228,15 @@ def set_cuda_params(): torch.backends.cudnn.allow_tf32 = True except Exception: pass + try: + if shared.opts.cross_attention_optimization == "Scaled-Dot-Product": + torch.backends.cuda.enable_flash_sdp('Flash attention' in shared.opts.sdp_options) + torch.backends.cuda.enable_mem_efficient_sdp('Memory attention' in shared.opts.sdp_options) + torch.backends.cuda.enable_math_sdp('Math attention' in shared.opts.sdp_options) + except Exception: + pass + if shared.cmd_opts.profile: + shared.log.debug(f'Torch info: {torch.__config__.show()}') global dtype, dtype_vae, dtype_unet, unet_needs_upcast, inference_context # pylint: disable=global-statement if shared.opts.cuda_dtype == 'FP32': dtype = torch.float32 @@ -263,7 +273,7 @@ def set_cuda_params(): inference_context = torch.no_grad log_device_name = get_raw_openvino_device() if shared.cmd_opts.use_openvino else torch.device(get_optimal_device_name()) log.debug(f'Desired Torch parameters: dtype={shared.opts.cuda_dtype} no-half={shared.opts.no_half} no-half-vae={shared.opts.no_half_vae} upscast={shared.opts.upcast_sampling}') - log.info(f'Setting Torch parameters: device={log_device_name} dtype={dtype} vae={dtype_vae} unet={dtype_unet} context={inference_context.__name__} fp16={fp16_ok} bf16={bf16_ok}') + log.info(f'Setting Torch parameters: device={log_device_name} dtype={dtype} vae={dtype_vae} unet={dtype_unet} context={inference_context.__name__} fp16={fp16_ok} bf16={bf16_ok} optimization={shared.opts.cross_attention_optimization}') args = cmd_args.parser.parse_args() diff --git a/modules/sd_hijack.py b/modules/sd_hijack.py index d6b7ce18c..d2c31fffb 100644 --- a/modules/sd_hijack.py +++ b/modules/sd_hijack.py @@ -50,17 +50,17 @@ def apply_optimizations(): shared.log.warning("Cross-attention: xFormers is not available on CPU") shared.xformers_available = False - shared.log.info(f"Cross-attention: optimization={shared.opts.cross_attention_optimization} options={shared.opts.cross_attention_options}") + shared.log.info(f"Cross-attention: optimization={shared.opts.cross_attention_optimization}") if shared.opts.cross_attention_optimization == "Disabled": optimization_method = 'none' - if can_use_sdp and shared.opts.cross_attention_optimization == "Scaled-Dot-Product" and 'SDP disable memory attention' in shared.opts.cross_attention_options: - ldm.modules.attention.CrossAttention.forward = sd_hijack_optimizations.scaled_dot_product_no_mem_attention_forward - ldm.modules.diffusionmodules.model.AttnBlock.forward = sd_hijack_optimizations.sdp_no_mem_attnblock_forward - optimization_method = 'sdp-no-mem' - elif can_use_sdp and shared.opts.cross_attention_optimization == "Scaled-Dot-Product": - ldm.modules.attention.CrossAttention.forward = sd_hijack_optimizations.scaled_dot_product_attention_forward - ldm.modules.diffusionmodules.model.AttnBlock.forward = sd_hijack_optimizations.sdp_attnblock_forward + if can_use_sdp and shared.opts.cross_attention_optimization == "Scaled-Dot-Product": optimization_method = 'sdp' + if 'Memory attention' in shared.opts.sdp_options: + ldm.modules.attention.CrossAttention.forward = sd_hijack_optimizations.scaled_dot_product_no_mem_attention_forward + ldm.modules.diffusionmodules.model.AttnBlock.forward = sd_hijack_optimizations.sdp_no_mem_attnblock_forward + else: + ldm.modules.attention.CrossAttention.forward = sd_hijack_optimizations.scaled_dot_product_attention_forward + ldm.modules.diffusionmodules.model.AttnBlock.forward = sd_hijack_optimizations.sdp_attnblock_forward if shared.xformers_available and shared.opts.cross_attention_optimization == "xFormers": ldm.modules.attention.CrossAttention.forward = sd_hijack_optimizations.xformers_attention_forward ldm.modules.diffusionmodules.model.AttnBlock.forward = sd_hijack_optimizations.xformers_attnblock_forward diff --git a/modules/sd_hijack_optimizations.py b/modules/sd_hijack_optimizations.py index 6382ded59..6421de0de 100644 --- a/modules/sd_hijack_optimizations.py +++ b/modules/sd_hijack_optimizations.py @@ -2,19 +2,15 @@ from __future__ import annotations import sys import math import psutil -from functools import cache - import torch from torch import einsum - from ldm.util import default from einops import rearrange - from modules import shared, errors, devices from modules.hypernetworks import hypernetwork - from .sub_quadratic_attention import efficient_dot_product_attention # pylint: disable=relative-beyond-top-level + if shared.opts.cross_attention_optimization == "xFormers": try: import xformers.ops # pylint: disable=import-error @@ -313,9 +309,8 @@ def sub_quad_attention(q, k, v, q_chunk_size=1024, kv_chunk_size=None, kv_chunk_ def get_xformers_flash_attention_op(q, k, v): - if 'xFormers enable flash Attention' not in shared.opts.cross_attention_options: + if 'Flash attention' not in shared.opts.xformers_options: return None - try: flash_attention_op = xformers.ops.MemoryEfficientAttentionFlashAttentionOp # pylint: disable=used-before-assignment fw, _bw = flash_attention_op @@ -323,7 +318,6 @@ def get_xformers_flash_attention_op(q, k, v): return flash_attention_op except Exception as e: errors.display_once(e, "enabling flash attention") - return None diff --git a/modules/shared.py b/modules/shared.py index a192d9b9b..fb5e6aa4e 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -355,8 +355,9 @@ options_templates.update(options_section(('cuda', "Compute Settings"), { "cross_attention_sep": OptionInfo("

Attention

", "", gr.HTML), "cross_attention_optimization": OptionInfo(cross_attention_optimization_default, "Attention optimization method", gr.Radio, lambda: {"choices": shared_items.list_crossattention(diffusers=backend == Backend.DIFFUSERS) }), - "cross_attention_options": OptionInfo([], "Attention advanced options", gr.CheckboxGroup, {"choices": ['xFormers enable flash Attention', 'SDP disable memory attention'], "visible": False }), - "dynamic_attention_slice_rate": OptionInfo(4, "Slicing rate for Dynamic Attention Slicing in GB", gr.Slider, {"minimum": 0.1, "maximum": 16, "step": 0.1, "visible": backend == Backend.DIFFUSERS}), + "sdp_options": OptionInfo(['Flash attention', 'Memory attention', 'Math attention'], "SDP options", gr.CheckboxGroup, {"choices": ['Flash attention', 'Memory attention', 'Math 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": 16, "step": 0.1, "visible": backend == Backend.DIFFUSERS}), "sub_quad_sep": OptionInfo("

Sub-quadratic options

", "", gr.HTML, {"visible": backend == Backend.ORIGINAL}), "sub_quad_q_chunk_size": OptionInfo(512, "Attention query chunk size", gr.Slider, {"minimum": 16, "maximum": 8192, "step": 8, "visible": backend == Backend.ORIGINAL}), "sub_quad_kv_chunk_size": OptionInfo(512, "Attention kv chunk size", gr.Slider, {"minimum": 0, "maximum": 8192, "step": 8, "visible": backend == Backend.ORIGINAL}), @@ -365,6 +366,7 @@ options_templates.update(options_section(('cuda', "Compute Settings"), { "other_sep": OptionInfo("

Execution precision

", "", gr.HTML), "opt_channelslast": OptionInfo(False, "Use channels last "), "cudnn_benchmark": OptionInfo(False, "Full-depth cuDNN benchmark feature"), + "cudnn_deterministic": OptionInfo(False, "Use deterministic options for cuDNN"), "diffusers_fuse_projections": OptionInfo(False, "Fused projections"), "torch_gc_threshold": OptionInfo(80, "Memory usage threshold for GC", gr.Slider, {"minimum": 0, "maximum": 100, "step": 1}),