diff --git a/modules/attention/dispatcher.py b/modules/attention/dispatcher.py index 82cbc323c..19d58e99b 100644 --- a/modules/attention/dispatcher.py +++ b/modules/attention/dispatcher.py @@ -5,40 +5,20 @@ from installer import install, torch_info def set_diffusers_attention(pipe, quiet = False): from modules import shared, devices - import diffusers.models.attention_processor as p - - def set_attn(pipe, attention, name: str | None = None): - if attention is None: - return - # other models uses their own attention processor - if getattr(pipe, "unet", None) is not None and hasattr(pipe.unet, "set_attn_processor"): - try: - pipe.unet.set_attn_processor(attention) - except Exception as e: - if 'Nunchaku' in pipe.unet.__class__.__name__: - pass - else: - log.error(f'Torch attention: type="{name}" cls={attention.__class__.__name__} pipe={pipe.__class__.__name__} {e}') log.quiet(quiet, f'Setting model: attention="{shared.opts.cross_attention_optimization}"') if shared.opts.cross_attention_optimization == "Disabled": torch_info.set(attention="disabled") elif shared.opts.cross_attention_optimization == "Scaled-Dot-Product": # The default set by Diffusers devices.set_sdpa_params() - # set_attn(pipe, p.AttnProcessor2_0(), name="Scaled-Dot-Product") elif shared.opts.cross_attention_optimization == "xFormers": if hasattr(pipe, 'enable_xformers_memory_efficient_attention'): torch_info.set(attention="xformers") pipe.enable_xformers_memory_efficient_attention() else: log.warning(f"Attention: xFormers is not compatible with {pipe.__class__.__name__}") - elif shared.opts.cross_attention_optimization == "Batch matrix-matrix": - torch_info.set(attention="bmm") - set_attn(pipe, p.AttnProcessor(), name="Batch matrix-matrix") - elif shared.opts.cross_attention_optimization == "Dynamic Attention BMM": - from modules.sd_hijack_dynamic_atten import DynamicAttnProcessorBMM - torch_info.set(attention="dynamic_bmm") - set_attn(pipe, DynamicAttnProcessorBMM(), name="Dynamic Attention BMM") + else: + log.warning(f'Torch attention: method="{shared.opts.cross_attention_optimization}" unknown, pipe={pipe.__class__.__name__} keeps its own attention processor') if shared.opts.attention_slicing != "Default" and hasattr(pipe, "enable_attention_slicing") and hasattr(pipe, "disable_attention_slicing"): if shared.opts.attention_slicing == "Enabled": diff --git a/modules/options_handler.py b/modules/options_handler.py index d0bb8d454..2400a2069 100644 --- a/modules/options_handler.py +++ b/modules/options_handler.py @@ -18,9 +18,22 @@ if TYPE_CHECKING: cmd_opts = cmd_args.parse_args() compatibility_opts = ['clip_skip', 'uni_pc_lower_order_final', 'uni_pc_order', 'xformers_options'] +removed_values = { # a stored choice that no longer exists is kept by validate, so it has to be rewritten or it selects nothing + 'cross_attention_optimization': (['Batch matrix-matrix', 'Dynamic Attention BMM'], 'Scaled-Dot-Product'), +} secrets_pattern = ['_version', '_token', '_key', '_secret', '_password'] +def migrate_removed_values(data: dict) -> list: + """Rewrite stored settings whose choice was removed, returning what changed.""" + migrated = [] + for key, (removed, replacement) in removed_values.items(): + if data.get(key, None) in removed: + migrated.append(f'{key}={data[key]} replaced={replacement}') + data[key] = replacement + return migrated + + class Options: data_labels: dict[str, OptionInfo | LegacyOption] data: dict[str, Any] @@ -203,6 +216,9 @@ class Options: self.secrets = readfile(secretsfn, lock=True, as_type="dict") if self.data.get('quicksettings') is not None and self.data.get('quicksettings_list') is None: self.data['quicksettings_list'] = [i.strip() for i in self.data.get('quicksettings', '').split(',')] + migrated = migrate_removed_values(self.data) + if len(migrated) > 0: + log.warning(f"Setting migration: {migrated}") unknown_settings = [] for k, v in self.data.items(): info = self.data_labels.get(k, None) diff --git a/modules/sd_hijack_dynamic_atten.py b/modules/sd_hijack_dynamic_atten.py index 0c0c4771d..cf932496d 100644 --- a/modules/sd_hijack_dynamic_atten.py +++ b/modules/sd_hijack_dynamic_atten.py @@ -113,160 +113,3 @@ def dynamic_scaled_dot_product_attention(query: torch.FloatTensor, key: torch.Fl if is_unsqueezed: hidden_states = hidden_states.squeeze(0) return hidden_states - - -@cache -def find_bmm_slice_sizes(query_shape, query_element_size, slice_rate=2, trigger_rate=4): - if len(query_shape) == 3: - batch_size_attention, query_tokens, shape_three = query_shape - shape_four = 1 - else: - batch_size_attention, query_tokens, shape_three, shape_four = query_shape - - slice_block_size = query_tokens * shape_three * shape_four / 1024 / 1024 * query_element_size - block_size = batch_size_attention * slice_block_size - - split_slice_size = batch_size_attention - split_2_slice_size = query_tokens - split_3_slice_size = shape_three - - do_split = False - do_split_2 = False - do_split_3 = False - - 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: - slice_2_block_size = split_slice_size * shape_three * shape_four / 1024 / 1024 * query_element_size - do_split_2 = True - split_2_slice_size = find_split_size(split_2_slice_size, slice_2_block_size, slice_rate=slice_rate) - if split_2_slice_size * slice_2_block_size > slice_rate: - slice_3_block_size = split_slice_size * split_2_slice_size * shape_four / 1024 / 1024 * query_element_size - do_split_3 = True - split_3_slice_size = find_split_size(split_3_slice_size, slice_3_block_size, slice_rate=slice_rate) - - return do_split, do_split_2, do_split_3, split_slice_size, split_2_slice_size, split_3_slice_size - - -class DynamicAttnProcessorBMM: - r""" - dynamically slices attention queries in order to keep them under the slice rate - slicing will not get triggered if the query size is smaller than the slice rate to gain performance - - slice rate is in GB - based on AttnProcessor V1 - """ - - def __call__(self, attn, hidden_states: torch.Tensor, encoder_hidden_states=None, attention_mask=None, temb=None, *args, **kwargs) -> torch.Tensor: # pylint: disable=too-many-statements, too-many-locals, too-many-branches, keyword-arg-before-vararg - - residual = hidden_states - - if attn.spatial_norm is not None: - hidden_states = attn.spatial_norm(hidden_states, temb) - - input_ndim = hidden_states.ndim - - if input_ndim == 4: - batch_size, channel, height, width = hidden_states.shape - hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) - - batch_size, sequence_length, _ = ( - hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape - ) - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - - if attn.group_norm is not None: - hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) - - query = attn.to_q(hidden_states) - - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - elif attn.norm_cross: - encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) - - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - query = attn.head_to_batch_dim(query) - key = attn.head_to_batch_dim(key) - value = attn.head_to_batch_dim(value) - - #################################################################### - # 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*4, trigger_rate=shared.opts.dynamic_attention_trigger_rate*4) - - if do_split: - for i in range(batch_size_attention // split_slice_size): - start_idx = i * split_slice_size - end_idx = (i + 1) * split_slice_size - if do_split_2: - for i2 in range(query_tokens // split_2_slice_size): # pylint: disable=invalid-name - start_idx_2 = i2 * split_2_slice_size - end_idx_2 = (i2 + 1) * split_2_slice_size - if do_split_3: - for i3 in range(shape_three // split_3_slice_size): # pylint: disable=invalid-name - start_idx_3 = i3 * split_3_slice_size - end_idx_3 = (i3 + 1) * split_3_slice_size - - query_slice = query[start_idx:end_idx, start_idx_2:end_idx_2, start_idx_3:end_idx_3] - key_slice = key[start_idx:end_idx, start_idx_2:end_idx_2, start_idx_3:end_idx_3] - attn_mask_slice = attention_mask[start_idx:end_idx, start_idx_2:end_idx_2, start_idx_3:end_idx_3] if attention_mask is not None else None - - attn_slice = attn.get_attention_scores(query_slice, key_slice, attn_mask_slice) - del query_slice - del key_slice - del attn_mask_slice - attn_slice = torch.bmm(attn_slice, value[start_idx:end_idx, start_idx_2:end_idx_2, start_idx_3:end_idx_3]) - - hidden_states[start_idx:end_idx, start_idx_2:end_idx_2, start_idx_3:end_idx_3] = attn_slice - del attn_slice - else: - query_slice = query[start_idx:end_idx, start_idx_2:end_idx_2] - key_slice = key[start_idx:end_idx, start_idx_2:end_idx_2] - attn_mask_slice = attention_mask[start_idx:end_idx, start_idx_2:end_idx_2] if attention_mask is not None else None - - attn_slice = attn.get_attention_scores(query_slice, key_slice, attn_mask_slice) - del query_slice - del key_slice - del attn_mask_slice - attn_slice = torch.bmm(attn_slice, value[start_idx:end_idx, start_idx_2:end_idx_2]) - - hidden_states[start_idx:end_idx, start_idx_2:end_idx_2] = attn_slice - del attn_slice - else: - query_slice = query[start_idx:end_idx] - key_slice = key[start_idx:end_idx] - attn_mask_slice = attention_mask[start_idx:end_idx] if attention_mask is not None else None - - attn_slice = attn.get_attention_scores(query_slice, key_slice, attn_mask_slice) - del query_slice - del key_slice - del attn_mask_slice - attn_slice = torch.bmm(attn_slice, value[start_idx:end_idx]) - - hidden_states[start_idx:end_idx] = attn_slice - del attn_slice - else: - attention_probs = attn.get_attention_scores(query, key, attention_mask) - hidden_states = torch.bmm(attention_probs, value) - #################################################################### - hidden_states = attn.batch_to_head_dim(hidden_states) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - if input_ndim == 4: - hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) - - if attn.residual_connection: - hidden_states = hidden_states + residual - - hidden_states = hidden_states / attn.rescale_output_factor - - return hidden_states diff --git a/modules/shared_items.py b/modules/shared_items.py index e7d8dab8e..ad37125db 100644 --- a/modules/shared_items.py +++ b/modules/shared_items.py @@ -142,8 +142,6 @@ def list_crossattention(): "Disabled", "Scaled-Dot-Product", "xFormers", - "Batch matrix-matrix", - "Dynamic Attention BMM" ] diff --git a/test/test-attention-router.py b/test/test-attention-router.py index 245467326..d100e6168 100644 --- a/test/test-attention-router.py +++ b/test/test-attention-router.py @@ -443,6 +443,27 @@ def test_attention_slicing_follows_the_choice(): return True +def test_removed_attention_methods_are_gone(): + from modules import shared_items + from modules import options_handler + from modules import sd_hijack_dynamic_atten + + removed = ['Batch matrix-matrix', 'Dynamic Attention BMM'] + choices = shared_items.list_crossattention() + assert not [name for name in removed if name in choices], choices + for name in removed: + data = {'cross_attention_optimization': name} + migrated = options_handler.migrate_removed_values(data) + assert data['cross_attention_optimization'] == 'Scaled-Dot-Product', data + assert len(migrated) == 1, migrated + kept = {'cross_attention_optimization': 'xFormers'} + assert options_handler.migrate_removed_values(kept) == [], 'a live choice is left alone' + assert kept['cross_attention_optimization'] == 'xFormers', kept + assert not hasattr(sd_hijack_dynamic_atten, 'DynamicAttnProcessorBMM'), 'the bmm processor is removed' + assert hasattr(sd_hijack_dynamic_atten, 'dynamic_scaled_dot_product_attention'), 'the sliced sdpa path stays' + return True + + def test_escape_hatch_bypasses_the_router(): from modules import devices saved_sdpa = torch.nn.functional.scaled_dot_product_attention @@ -516,6 +537,7 @@ def run_all(): test_debug_observe_logs_each_route_once, test_reapply_options_cover_declared_backend_options, test_attention_slicing_follows_the_choice, + test_removed_attention_methods_are_gone, test_escape_hatch_bypasses_the_router, ]: run_test(cat, fn) diff --git a/ui/locale/locale_en.json b/ui/locale/locale_en.json index 089b239b6..d9260dc01 100644 --- a/ui/locale/locale_en.json +++ b/ui/locale/locale_en.json @@ -189,7 +189,6 @@ {"id":"","label":"block_level","localized":"","hint":"","ui":"settings_offload"}, {"id":"","label":"Backend storage","localized":"","hint":"","ui":"settings_quantization"}, {"id":"","label":"BF16","localized":"","hint":"Use modified 16-bit floating point precision for calculations","ui":"settings_cuda"}, - {"id":"","label":"Batch matrix-matrix","localized":"","hint":"Standard batched matrix multiplication for attention. Reliable but not VRAM-efficient.","ui":"settings_cuda"}, {"id":"","label":"BCFHW","localized":"","hint":"","ui":"settings_advanced"}, {"id":"","label":"BFCHW","localized":"","hint":"","ui":"settings_advanced"}, {"id":"","label":"BCHW","localized":"","hint":"","ui":"settings_advanced"}, @@ -397,7 +396,6 @@ {"id":"","label":"Dequantize using torch.compile","localized":"","hint":"Compiles the dequantization step with torch.compile for faster inference. Requires Triton.

Changing this needs a full restart to take effect.

Enabled by default when Triton is available.","reload":"server","ui":"settings_quantization"}, {"id":"","label":"Dequantize using full precision","localized":"","hint":"Uses FP32 for the dequantization step for better numerical accuracy, at a small speed cost.

Enabled by default.","reload":"model","ui":"settings_quantization"}, {"id":"","label":"Disabled","localized":"","hint":"","ui":"settings_cuda"}, - {"id":"","label":"Dynamic Attention BMM","localized":"","hint":"Performs attention computation in steps instead of all at once. Slower inference times, but greatly reduced memory usage","ui":"settings_cuda"}, {"id":"","label":"Dynamic attention","localized":"","hint":"Adjusts attention computation dynamically per step. Saves VRAM but slows generation.","ui":"settings_cuda"}, {"id":"","label":"Dynamic Attention slicing rate","localized":"","hint":"","ui":"settings_cuda"}, {"id":"","label":"Dynamic Attention trigger rate","localized":"","hint":"","ui":"settings_cuda"},