Better dyn atten sdpa usage estimation logic

This commit is contained in:
Disty0
2025-01-26 16:57:26 +03:00
parent a21bc9d9cd
commit b0ecdf2f1c
3 changed files with 40 additions and 57 deletions
+20 -29
View File
@@ -7,38 +7,29 @@ from functools import cache, wraps
# ARC GPUs can't allocate more than 4GB to a single block so we slice the attetion layers
sdpa_slice_trigger_rate = float(os.environ.get('IPEX_SDPA_SLICE_TRIGGER_RATE', 4))
attention_slice_rate = float(os.environ.get('IPEX_ATTENTION_SLICE_RATE', 4))
sdpa_slice_trigger_rate = float(os.environ.get('IPEX_SDPA_SLICE_TRIGGER_RATE', 3))
attention_slice_rate = float(os.environ.get('IPEX_ATTENTION_SLICE_RATE', 2))
# Find something divisible with the input_tokens
@cache
def find_split_size(split_size, slice_block_size):
while (split_size * slice_block_size) > attention_slice_rate:
split_size = split_size // 2
def find_split_size(original_size, slice_block_size, slice_rate=2):
split_size = original_size
while True:
if (split_size * slice_block_size) <= slice_rate and original_size % split_size == 0:
return split_size
split_size = split_size - 1
if split_size <= 1:
split_size = 1
break
return 1
return split_size
@cache
def find_query_size(query_size, slice_query_size):
while (math.sqrt(query_size) * slice_query_size) > attention_slice_rate:
query_size = query_size // 2
if query_size <= 1:
query_size = 1
break
return query_size
# Find slice sizes for SDPA
@cache
def find_sdpa_slice_sizes(query_shape, key_shape, value_shape, query_element_size):
def find_sdpa_slice_sizes(query_shape, key_shape, query_element_size, slice_rate=2, trigger_rate=3):
batch_size, attn_heads, query_len, _ = query_shape
_, _, key_len, _ = key_shape
_, _, _, head_dim = value_shape
slice_batch_size = attn_heads * math.sqrt(query_len * key_len) * math.sqrt(head_dim) * 2 * query_element_size / 1024 / 1024
slice_batch_size = attn_heads * (query_len * key_len) * query_element_size / 1024 / 1024 / 1024
split_batch_size = batch_size
split_head_size = attn_heads
@@ -48,19 +39,19 @@ 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 > sdpa_slice_trigger_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)
split_batch_size = find_split_size(batch_size, slice_batch_size, slice_rate=slice_rate)
if split_batch_size * slice_batch_size > attention_slice_rate:
slice_head_size = split_batch_size * math.sqrt(query_len * key_len) * math.sqrt(head_dim) * 2 * query_element_size / 1024 / 1024
if split_batch_size * slice_batch_size > slice_rate:
slice_head_size = split_batch_size * (query_len * key_len) * query_element_size / 1024 / 1024 / 1024
do_head_split = True
split_head_size = find_split_size(split_head_size, slice_head_size)
split_head_size = find_split_size(attn_heads, slice_head_size, slice_rate=slice_rate)
if split_batch_size * slice_batch_size > attention_slice_rate:
slice_query_size = split_batch_size * attn_heads * math.sqrt(key_len) * math.sqrt(head_dim) * 2 * query_element_size / 1024 / 1024
if split_head_size * slice_head_size > slice_rate:
slice_query_size = split_batch_size * split_head_size * (key_len) * query_element_size / 1024 / 1024 / 1024
do_query_split = True
split_query_size = find_query_size(split_query_size, slice_query_size)
split_query_size = find_split_size(query_len, slice_query_size, slice_rate=slice_rate)
return do_batch_split, do_head_split, do_query_split, split_batch_size, split_head_size, split_query_size
@@ -78,7 +69,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())
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, query.element_size(), slice_rate=attention_slice_rate, trigger_rate=sdpa_slice_trigger_rate)
# Slice SDPA
if do_batch_split:
+18 -26
View File
@@ -6,34 +6,26 @@ from diffusers.utils import USE_PEFT_BACKEND # pylint: disable=unused-import
from modules import shared, devices
# Find something divisible with the input_tokens
@cache
def find_split_size(split_size, slice_block_size, slice_rate=4):
while (split_size * slice_block_size) > slice_rate:
split_size = split_size // 2
def find_split_size(original_size, slice_block_size, slice_rate=2):
split_size = original_size
while True:
if (split_size * slice_block_size) <= slice_rate and original_size % split_size == 0:
return split_size
split_size = split_size - 1
if split_size <= 1:
split_size = 1
break
return 1
return split_size
@cache
def find_query_size(query_size, slice_query_size, slice_rate=4):
while (math.sqrt(query_size) * slice_query_size) > slice_rate:
query_size = query_size // 2
if query_size <= 1:
query_size = 1
break
return query_size
# Find slice sizes for SDPA
@cache
def find_sdpa_slice_sizes(query_shape, key_shape, value_shape, query_element_size, slice_rate=4, trigger_rate=6):
def find_sdpa_slice_sizes(query_shape, key_shape, query_element_size, slice_rate=2, trigger_rate=3):
batch_size, attn_heads, query_len, _ = query_shape
_, _, key_len, _ = key_shape
_, _, _, head_dim = value_shape
slice_batch_size = attn_heads * math.sqrt(query_len * key_len) * head_dim * query_element_size / 1024 / 1024 / 2
slice_batch_size = attn_heads * (query_len * key_len) * query_element_size / 1024 / 1024 / 1024
split_batch_size = batch_size
split_head_size = attn_heads
@@ -43,19 +35,19 @@ 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 > trigger_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)
split_batch_size = find_split_size(batch_size, slice_batch_size, slice_rate=slice_rate)
if split_batch_size * slice_batch_size > slice_rate:
slice_head_size = split_batch_size * math.sqrt(query_len * key_len) * head_dim * query_element_size / 1024 / 1024 / 2
slice_head_size = split_batch_size * (query_len * key_len) * query_element_size / 1024 / 1024 / 1024
do_head_split = True
split_head_size = find_split_size(split_head_size, slice_head_size, slice_rate=slice_rate)
split_head_size = find_split_size(attn_heads, slice_head_size, slice_rate=slice_rate)
if split_batch_size * slice_batch_size > slice_rate:
slice_query_size = split_batch_size * attn_heads * math.sqrt(key_len) * head_dim * query_element_size / 1024 / 1024 / 2
if split_head_size * slice_head_size > slice_rate:
slice_query_size = split_batch_size * split_head_size * (key_len) * query_element_size / 1024 / 1024 / 1024
do_query_split = True
split_query_size = find_query_size(split_query_size, slice_query_size, slice_rate=slice_rate)
split_query_size = find_split_size(query_len, slice_query_size, slice_rate=slice_rate)
return do_batch_split, do_head_split, do_query_split, split_batch_size, split_head_size, split_query_size
@@ -72,7 +64,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, trigger_rate=shared.opts.dynamic_attention_trigger_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, 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:
+2 -2
View File
@@ -514,8 +514,8 @@ options_templates.update(options_section(('cuda', "Compute Settings"), {
"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', '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}),
"dynamic_attention_slice_rate": OptionInfo(2, "Dynamic Attention slicing rate in GB", gr.Slider, {"minimum": 0.1, "maximum": gpu_memory, "step": 0.1, "visible": native}),
"dynamic_attention_trigger_rate": OptionInfo(3, "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}),