diff --git a/CHANGELOG.md b/CHANGELOG.md index 18e0de6ee..77ac8ec0a 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -56,7 +56,7 @@ - **loader**: ability to run in-memory models - **schedulers**: ability to create model-less schedulers - **quantiation**: code refactor into dedicated module - - **dynamic attention sdpa**: more correct implementation + - **dynamic attention sdpa**: more correct implementation and new trigger rate control - **Authentication**: - perform auth check on ui startup - unified standard and modern-ui authentication method diff --git a/modules/sd_hijack_dynamic_atten.py b/modules/sd_hijack_dynamic_atten.py index cce8c13f0..a3e63dd6f 100644 --- a/modules/sd_hijack_dynamic_atten.py +++ b/modules/sd_hijack_dynamic_atten.py @@ -28,7 +28,7 @@ def find_query_size(query_size, slice_query_size, slice_rate=4): # Find slice sizes for SDPA @cache -def find_sdpa_slice_sizes(query_shape, key_shape, value_shape, query_element_size, slice_rate=4): +def find_sdpa_slice_sizes(query_shape, key_shape, value_shape, query_element_size, slice_rate=4, trigger_rate=6): batch_size, attn_heads, query_len, _ = query_shape _, _, key_len, _ = key_shape _, _, _, head_dim = value_shape @@ -43,7 +43,7 @@ def find_sdpa_slice_sizes(query_shape, key_shape, value_shape, query_element_siz do_head_split = False do_query_split = False - if batch_size * slice_batch_size > slice_rate: + if batch_size * slice_batch_size > trigger_rate: do_batch_split = True split_batch_size = find_split_size(split_batch_size, slice_batch_size, slice_rate=slice_rate) @@ -72,7 +72,7 @@ def dynamic_scaled_dot_product_attention(query, key, value, attn_mask=None, drop key = key.unsqueeze(0) if len(value.shape) == 3: value = value.unsqueeze(0) - do_batch_split, do_head_split, do_query_split, split_batch_size, split_head_size, split_query_size = find_sdpa_slice_sizes(query.shape, key.shape, value.shape, query.element_size(), slice_rate=shared.opts.dynamic_attention_slice_rate) + do_batch_split, do_head_split, do_query_split, split_batch_size, split_head_size, split_query_size = find_sdpa_slice_sizes(query.shape, key.shape, value.shape, query.element_size(), slice_rate=shared.opts.dynamic_attention_slice_rate, trigger_rate=shared.opts.dynamic_attention_trigger_rate) # Slice SDPA if do_batch_split: @@ -125,7 +125,7 @@ def dynamic_scaled_dot_product_attention(query, key, value, attn_mask=None, drop @cache -def find_bmm_slice_sizes(query_shape, query_element_size, slice_rate=4): +def find_bmm_slice_sizes(query_shape, query_element_size, slice_rate=4, trigger_rate=6): if len(query_shape) == 3: batch_size_attention, query_tokens, shape_three = query_shape shape_four = 1 @@ -143,7 +143,7 @@ def find_bmm_slice_sizes(query_shape, query_element_size, slice_rate=4): do_split_2 = False do_split_3 = False - if block_size > slice_rate: + if block_size > trigger_rate: do_split = True split_slice_size = find_split_size(split_slice_size, slice_block_size, slice_rate=slice_rate) if split_slice_size * slice_block_size > slice_rate: @@ -206,7 +206,7 @@ class DynamicAttnProcessorBMM: # Slicing parts: batch_size_attention, query_tokens, shape_three = query.shape[0], query.shape[1], query.shape[2] hidden_states = torch.zeros(query.shape, device=query.device, dtype=query.dtype) - do_split, do_split_2, do_split_3, split_slice_size, split_2_slice_size, split_3_slice_size = find_bmm_slice_sizes(query.shape, query.element_size(), slice_rate=shared.opts.dynamic_attention_slice_rate) + do_split, do_split_2, do_split_3, split_slice_size, split_2_slice_size, split_3_slice_size = find_bmm_slice_sizes(query.shape, query.element_size(), slice_rate=shared.opts.dynamic_attention_slice_rate, trigger_rate=shared.opts.dynamic_attention_trigger_rate) if do_split: for i in range(batch_size_attention // split_slice_size): diff --git a/modules/shared.py b/modules/shared.py index da9fa3f6d..25bab6fd4 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -515,6 +515,7 @@ options_templates.update(options_section(('cuda', "Compute Settings"), { "sdp_options": OptionInfo(startup_sdp_options, "SDP options", gr.CheckboxGroup, {"choices": ['Flash attention', 'Memory attention', 'Math attention', 'Dynamic attention', 'Sage attention'], "visible": native}), "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}), + "dynamic_attention_trigger_rate": OptionInfo(6, "Dynamic Attention trigger rate in GB", gr.Slider, {"minimum": 0.1, "maximum": gpu_memory*2, "step": 0.1, "visible": native}), "sub_quad_sep": OptionInfo("

Sub-quadratic options

", "", gr.HTML, {"visible": not native}), "sub_quad_q_chunk_size": OptionInfo(512, "Attention query chunk size", gr.Slider, {"minimum": 16, "maximum": 8192, "step": 8, "visible": not native}), "sub_quad_kv_chunk_size": OptionInfo(512, "Attention kv chunk size", gr.Slider, {"minimum": 0, "maximum": 8192, "step": 8, "visible": not native}),