mirror of
https://github.com/vladmandic/automatic
synced 2026-09-17 08:19:11 +02:00
Fix dyn atten not registering slice rate changes
This commit is contained in:
@@ -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):
|
||||
def find_sdpa_slice_sizes(query_shape, key_shape, value_shape, query_element_size, slice_rate=4):
|
||||
batch_size, attn_heads, query_len, _ = query_shape
|
||||
_, _, key_len, _ = key_shape
|
||||
_, _, _, head_dim = value_shape
|
||||
@@ -43,19 +43,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 > shared.opts.dynamic_attention_slice_rate:
|
||||
if batch_size * slice_batch_size > slice_rate:
|
||||
do_batch_split = True
|
||||
split_batch_size = find_split_size(split_batch_size, slice_batch_size, slice_rate=shared.opts.dynamic_attention_slice_rate)
|
||||
split_batch_size = find_split_size(split_batch_size, slice_batch_size, slice_rate=slice_rate)
|
||||
|
||||
if split_batch_size * slice_batch_size > shared.opts.dynamic_attention_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
|
||||
do_head_split = True
|
||||
split_head_size = find_split_size(split_head_size, slice_head_size, slice_rate=shared.opts.dynamic_attention_slice_rate)
|
||||
split_head_size = find_split_size(split_head_size, slice_head_size, slice_rate=slice_rate)
|
||||
|
||||
if split_batch_size * slice_batch_size > shared.opts.dynamic_attention_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
|
||||
do_query_split = True
|
||||
split_query_size = find_query_size(split_query_size, slice_query_size, slice_rate=shared.opts.dynamic_attention_slice_rate)
|
||||
split_query_size = find_query_size(split_query_size, 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 +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())
|
||||
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)
|
||||
|
||||
# Slice SDPA
|
||||
if do_batch_split:
|
||||
|
||||
Reference in New Issue
Block a user