Add trigger rate control to dyn atten

This commit is contained in:
Disty0
2025-01-26 03:48:01 +03:00
parent fad1edd8fa
commit bb07dd7a8f
3 changed files with 8 additions and 7 deletions
+1 -1
View File
@@ -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
+6 -6
View File
@@ -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):
+1
View File
@@ -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("<h3>Sub-quadratic options</h3>", "", 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}),