From 49e6c1564c6821808e5ba49c93dd06feadd23979 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Tue, 26 Nov 2024 13:13:04 -0500 Subject: [PATCH] add style aligned Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 47 +++-- modules/processing_class.py | 4 +- modules/processing_helpers.py | 4 +- modules/sd_samplers.py | 2 + modules/style_aligned/inversion.py | 124 ++++++++++++ modules/style_aligned/sa_handler.py | 281 ++++++++++++++++++++++++++++ scripts/style_aligned.py | 117 ++++++++++++ 7 files changed, 559 insertions(+), 20 deletions(-) create mode 100644 modules/style_aligned/inversion.py create mode 100644 modules/style_aligned/sa_handler.py create mode 100644 scripts/style_aligned.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 3ff4f0944..8183fbbc8 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -2,11 +2,13 @@ ## Update for 2024-11-26 -- [Flux Tools](https://blackforestlabs.ai/flux-1-tools/): +### New models and integrations + +- [Flux Tools](https://blackforestlabs.ai/flux-1-tools/) **Redux** is actually a tool, **Fill** is inpaint/outpaint optimized version of *Flux-dev* **Canny** & **Depth** are optimized versions of *Flux-dev* for their respective tasks: they are *not* ControlNets that work on top of a model - To use, go to image or control interface and select *Flux Tools* in scripts - All models are auto-downloaded on first use + to use, go to image or control interface and select *Flux Tools* in scripts + all models are auto-downloaded on first use *note*: All models are [gated](https://github.com/vladmandic/automatic/wiki/Gated) and require acceptance of terms and conditions via web page *recommended*: Enable on-the-fly [quantization](https://github.com/vladmandic/automatic/wiki/Quantization) or [compression](https://github.com/vladmandic/automatic/wiki/NNCF-Compression) to reduce resource usage *todo*: support for Canny/Depth LoRAs @@ -19,16 +21,23 @@ *recommended*: guidance scale 30 - [Depth](https://huggingface.co/black-forest-labs/FLUX.1-Depth-dev): ~23.8GB, replaces currently loaded model *recommended*: guidance scale 10 -- Model loader improvements: +- [Style Aligned Image Generation](https://style-aligned-gen.github.io/) + enable in scripts, compatible with sd-xl + enter multiple prompts in prompt field separated by new line + style-aligned applies selected attention layers uniformly to all images to achive consistency + can be used with or without input image in which case first prompt is used to establish baseline + *note:* all prompts are processes as a single batch, so vram is limiting factor + +### UI and workflow improvements + +- **Model loader** improvements: - detect model components on model load fail - Flux, SD35: force unload model - Flux: apply `bnb` quant when loading *unet/transformer* - Flux: all-in-one safetensors example: - Flux: do not recast quants -- Sampler improvements - - update DPM FlowMatch samplers -- UI: +- **UI**: - improved stats on generate completion - improved live preview display and performance - improved accordion behavior @@ -37,16 +46,20 @@ - control: optionn to hide input column - control: add stats - browser->server logging framework -- Fixes: - - update `diffusers` - - fix README links - - fix sdxl controlnet single-file loader - - relax settings validator - - improve js progress calls resiliency - - fix text-to-video pipeline - - avoid live-preview if vae-decode is running - - allow xyz-grid with multi-axis s&r - - fix xyz-grid with lora +- **Sampler** improvements + - update DPM FlowMatch samplers + +### Fixes: + +- update `diffusers` +- fix README links +- fix sdxl controlnet single-file loader +- relax settings validator +- improve js progress calls resiliency +- fix text-to-video pipeline +- avoid live-preview if vae-decode is running +- allow xyz-grid with multi-axis s&r +- fix xyz-grid with lora ## Update for 2024-11-21 diff --git a/modules/processing_class.py b/modules/processing_class.py index 79f51576f..21e86c1b0 100644 --- a/modules/processing_class.py +++ b/modules/processing_class.py @@ -31,8 +31,8 @@ class StableDiffusionProcessing: n_iter: int = 1, steps: int = 50, clip_skip: int = 1, - width: int = 512, - height: int = 512, + width: int = 1024, + height: int = 1024, # samplers sampler_index: int = None, # pylint: disable=unused-argument # used only to set sampler_name sampler_name: str = None, diff --git a/modules/processing_helpers.py b/modules/processing_helpers.py index ec7fbf048..22acf296c 100644 --- a/modules/processing_helpers.py +++ b/modules/processing_helpers.py @@ -561,7 +561,9 @@ def save_intermediate(p, latents, suffix): def update_sampler(p, sd_model, second_pass=False): sampler_selection = p.hr_sampler_name if second_pass else p.sampler_name if hasattr(sd_model, 'scheduler'): - if sampler_selection is None or sampler_selection == 'None': + if sampler_selection == 'None': + return + if sampler_selection is None: sampler = sd_samplers.all_samplers_map.get("UniPC") else: sampler = sd_samplers.all_samplers_map.get(sampler_selection, None) diff --git a/modules/sd_samplers.py b/modules/sd_samplers.py index e560744dd..d8416e5d9 100644 --- a/modules/sd_samplers.py +++ b/modules/sd_samplers.py @@ -47,6 +47,8 @@ def visible_sampler_names(): def create_sampler(name, model): + if name is None or name == 'None': + return model.scheduler try: current = model.scheduler.__class__.__name__ except Exception: diff --git a/modules/style_aligned/inversion.py b/modules/style_aligned/inversion.py new file mode 100644 index 000000000..8c91cc02a --- /dev/null +++ b/modules/style_aligned/inversion.py @@ -0,0 +1,124 @@ +# Copyright 2023 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + + +from __future__ import annotations +from typing import Callable, TYPE_CHECKING +from diffusers import StableDiffusionXLPipeline +import torch +from tqdm import tqdm +if TYPE_CHECKING: + import numpy as np + + +T = torch.Tensor +TN = T +InversionCallback = Callable[[StableDiffusionXLPipeline, int, T, dict[str, T]], dict[str, T]] + + +def _get_text_embeddings(prompt: str, tokenizer, text_encoder, device): + # Tokenize text and get embeddings + text_inputs = tokenizer(prompt, padding='max_length', max_length=tokenizer.model_max_length, truncation=True, return_tensors='pt') + text_input_ids = text_inputs.input_ids + + with torch.no_grad(): + prompt_embeds = text_encoder( + text_input_ids.to(device), + output_hidden_states=True, + ) + + pooled_prompt_embeds = prompt_embeds[0] + prompt_embeds = prompt_embeds.hidden_states[-2] + if prompt == '': + negative_prompt_embeds = torch.zeros_like(prompt_embeds) + negative_pooled_prompt_embeds = torch.zeros_like(pooled_prompt_embeds) + return negative_prompt_embeds, negative_pooled_prompt_embeds + return prompt_embeds, pooled_prompt_embeds + + +def _encode_text_sdxl(model: StableDiffusionXLPipeline, prompt: str) -> tuple[dict[str, T], T]: + device = model._execution_device # pylint: disable=protected-access + prompt_embeds, pooled_prompt_embeds, = _get_text_embeddings(prompt, model.tokenizer, model.text_encoder, device) # pylint: disable=unused-variable + prompt_embeds_2, pooled_prompt_embeds2, = _get_text_embeddings( prompt, model.tokenizer_2, model.text_encoder_2, device) + prompt_embeds = torch.cat((prompt_embeds, prompt_embeds_2), dim=-1) + text_encoder_projection_dim = model.text_encoder_2.config.projection_dim + add_time_ids = model._get_add_time_ids((1024, 1024), (0, 0), (1024, 1024), model.text_encoder.dtype, # pylint: disable=protected-access + text_encoder_projection_dim).to(device) + added_cond_kwargs = {"text_embeds": pooled_prompt_embeds2, "time_ids": add_time_ids} + return added_cond_kwargs, prompt_embeds + + +def _encode_text_sdxl_with_negative(model: StableDiffusionXLPipeline, prompt: str) -> tuple[dict[str, T], T]: + added_cond_kwargs, prompt_embeds = _encode_text_sdxl(model, prompt) + added_cond_kwargs_uncond, prompt_embeds_uncond = _encode_text_sdxl(model, "") + prompt_embeds = torch.cat((prompt_embeds_uncond, prompt_embeds, )) + added_cond_kwargs = {"text_embeds": torch.cat((added_cond_kwargs_uncond["text_embeds"], added_cond_kwargs["text_embeds"])), + "time_ids": torch.cat((added_cond_kwargs_uncond["time_ids"], added_cond_kwargs["time_ids"])),} + return added_cond_kwargs, prompt_embeds + + +def _encode_image(model: StableDiffusionXLPipeline, image: np.ndarray) -> T: + image = torch.from_numpy(image).float() / 255. + image = (image * 2 - 1).permute(2, 0, 1).unsqueeze(0) + latent = model.vae.encode(image.to(model.vae.device, model.vae.dtype))['latent_dist'].mean * model.vae.config.scaling_factor + return latent + + +def _next_step(model: StableDiffusionXLPipeline, model_output: T, timestep: int, sample: T) -> T: + timestep, next_timestep = min(timestep - model.scheduler.config.num_train_timesteps // model.scheduler.num_inference_steps, 999), timestep + alpha_prod_t = model.scheduler.alphas_cumprod[int(timestep)] if timestep >= 0 else model.scheduler.final_alpha_cumprod + alpha_prod_t_next = model.scheduler.alphas_cumprod[int(next_timestep)] + beta_prod_t = 1 - alpha_prod_t + next_original_sample = (sample - beta_prod_t ** 0.5 * model_output) / alpha_prod_t ** 0.5 + next_sample_direction = (1 - alpha_prod_t_next) ** 0.5 * model_output + next_sample = alpha_prod_t_next ** 0.5 * next_original_sample + next_sample_direction + return next_sample + + +def _get_noise_pred(model: StableDiffusionXLPipeline, latent: T, t: T, context: T, guidance_scale: float, added_cond_kwargs: dict[str, T]): + latents_input = torch.cat([latent] * 2) + noise_pred = model.unet(latents_input, t, encoder_hidden_states=context, added_cond_kwargs=added_cond_kwargs)["sample"] + noise_pred_uncond, noise_prediction_text = noise_pred.chunk(2) + noise_pred = noise_pred_uncond + guidance_scale * (noise_prediction_text - noise_pred_uncond) + # latents = next_step(model, noise_pred, t, latent) + return noise_pred + + +def _ddim_loop(model: StableDiffusionXLPipeline, z0, prompt, guidance_scale) -> T: + all_latent = [z0] + added_cond_kwargs, text_embedding = _encode_text_sdxl_with_negative(model, prompt) + latent = z0.clone().detach().to(model.text_encoder.dtype) + for i in tqdm(range(model.scheduler.num_inference_steps)): + t = model.scheduler.timesteps[len(model.scheduler.timesteps) - i - 1] + noise_pred = _get_noise_pred(model, latent, t, text_embedding, guidance_scale, added_cond_kwargs) + latent = _next_step(model, noise_pred, t, latent) + all_latent.append(latent) + return torch.cat(all_latent).flip(0) + + +def make_inversion_callback(zts, offset: int = 0): + + def callback_on_step_end(pipeline: StableDiffusionXLPipeline, i: int, t: T, callback_kwargs: dict[str, T]) -> dict[str, T]: # pylint: disable=unused-argument + latents = callback_kwargs['latents'] + latents[0] = zts[max(offset + 1, i + 1)].to(latents.device, latents.dtype) + return {'latents': latents} + return zts[offset], callback_on_step_end + + +@torch.no_grad() +def ddim_inversion(model: StableDiffusionXLPipeline, x0: np.ndarray, prompt: str, num_inference_steps: int, guidance_scale,) -> T: + z0 = _encode_image(model, x0) + model.scheduler.set_timesteps(num_inference_steps, device=z0.device) + zs = _ddim_loop(model, z0, prompt, guidance_scale) + return zs diff --git a/modules/style_aligned/sa_handler.py b/modules/style_aligned/sa_handler.py new file mode 100644 index 000000000..ee4b1ca79 --- /dev/null +++ b/modules/style_aligned/sa_handler.py @@ -0,0 +1,281 @@ +# Copyright 2023 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + + +from __future__ import annotations +from typing import TYPE_CHECKING +if TYPE_CHECKING: + from diffusers import StableDiffusionXLPipeline +from dataclasses import dataclass +import torch +import torch.nn as nn +from torch.nn import functional as nnf +from diffusers.models import attention_processor # pylint: disable=ungrouped-imports +import einops + +T = torch.Tensor + + +@dataclass(frozen=True) +class StyleAlignedArgs: + share_group_norm: bool = True + share_layer_norm: bool = True + share_attention: bool = True + adain_queries: bool = True + adain_keys: bool = True + adain_values: bool = False + full_attention_share: bool = False + shared_score_scale: float = 1. + shared_score_shift: float = 0. + only_self_level: float = 0. + + +def expand_first(feat: T, scale=1.,) -> T: + b = feat.shape[0] + feat_style = torch.stack((feat[0], feat[b // 2])).unsqueeze(1) + if scale == 1: + feat_style = feat_style.expand(2, b // 2, *feat.shape[1:]) + else: + feat_style = feat_style.repeat(1, b // 2, 1, 1, 1) + feat_style = torch.cat([feat_style[:, :1], scale * feat_style[:, 1:]], dim=1) + return feat_style.reshape(*feat.shape) + + +def concat_first(feat: T, dim=2, scale=1.) -> T: + feat_style = expand_first(feat, scale=scale) + return torch.cat((feat, feat_style), dim=dim) + + +def calc_mean_std(feat, eps: float = 1e-5) -> tuple[T, T]: + feat_std = (feat.var(dim=-2, keepdims=True) + eps).sqrt() + feat_mean = feat.mean(dim=-2, keepdims=True) + return feat_mean, feat_std + + +def adain(feat: T) -> T: + feat_mean, feat_std = calc_mean_std(feat) + feat_style_mean = expand_first(feat_mean) + feat_style_std = expand_first(feat_std) + feat = (feat - feat_mean) / feat_std + feat = feat * feat_style_std + feat_style_mean + return feat + + +class DefaultAttentionProcessor(nn.Module): + + def __init__(self): + super().__init__() + self.processor = attention_processor.AttnProcessor2_0() + + def __call__(self, attn: attention_processor.Attention, hidden_states, encoder_hidden_states=None, + attention_mask=None, **kwargs): + return self.processor(attn, hidden_states, encoder_hidden_states, attention_mask) + + +class SharedAttentionProcessor(DefaultAttentionProcessor): + + def shifted_scaled_dot_product_attention(self, attn: attention_processor.Attention, query: T, key: T, value: T) -> T: + logits = torch.einsum('bhqd,bhkd->bhqk', query, key) * attn.scale + logits[:, :, :, query.shape[2]:] += self.shared_score_shift + probs = logits.softmax(-1) + return torch.einsum('bhqk,bhkd->bhqd', probs, value) + + def shared_call( # pylint: disable=unused-argument + self, + attn: attention_processor.Attention, + hidden_states, + encoder_hidden_states=None, + attention_mask=None, + **kwargs + ): + + residual = hidden_states + 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 + ) + + if attention_mask is not None: + attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) + # scaled_dot_product_attention expects attention_mask shape to be + # (batch, heads, source_length, target_length) + attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) + + 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) + key = attn.to_k(hidden_states) + value = attn.to_v(hidden_states) + inner_dim = key.shape[-1] + head_dim = inner_dim // attn.heads + + query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + # if self.step >= self.start_inject: + if self.adain_queries: + query = adain(query) + if self.adain_keys: + key = adain(key) + if self.adain_values: + value = adain(value) + if self.share_attention: + key = concat_first(key, -2, scale=self.shared_score_scale) + value = concat_first(value, -2) + if self.shared_score_shift != 0: + hidden_states = self.shifted_scaled_dot_product_attention(attn, query, key, value,) + else: + hidden_states = nnf.scaled_dot_product_attention( + query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False + ) + else: + hidden_states = nnf.scaled_dot_product_attention( + query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False + ) + # hidden_states = adain(hidden_states) + hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) + hidden_states = hidden_states.to(query.dtype) + + # 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 + + def __call__(self, attn: attention_processor.Attention, hidden_states, encoder_hidden_states=None, + attention_mask=None, **kwargs): + if self.full_attention_share: + _b, n, _d = hidden_states.shape + hidden_states = einops.rearrange(hidden_states, '(k b) n d -> k (b n) d', k=2) + hidden_states = super().__call__(attn, hidden_states, encoder_hidden_states=encoder_hidden_states, + attention_mask=attention_mask, **kwargs) + hidden_states = einops.rearrange(hidden_states, 'k (b n) d -> (k b) n d', n=n) + else: + hidden_states = self.shared_call(attn, hidden_states, hidden_states, attention_mask, **kwargs) + + return hidden_states + + def __init__(self, style_aligned_args: StyleAlignedArgs): + super().__init__() + self.share_attention = style_aligned_args.share_attention + self.adain_queries = style_aligned_args.adain_queries + self.adain_keys = style_aligned_args.adain_keys + self.adain_values = style_aligned_args.adain_values + self.full_attention_share = style_aligned_args.full_attention_share + self.shared_score_scale = style_aligned_args.shared_score_scale + self.shared_score_shift = style_aligned_args.shared_score_shift + + +def _get_switch_vec(total_num_layers, level): + if level <= 0: + return torch.zeros(total_num_layers, dtype=torch.bool) + if level >= 1: + return torch.ones(total_num_layers, dtype=torch.bool) + to_flip = level > .5 + if to_flip: + level = 1 - level + num_switch = int(level * total_num_layers) + vec = torch.arange(total_num_layers) + vec = vec % (total_num_layers // num_switch) + vec = vec == 0 + if to_flip: + vec = ~vec + return vec + + +def init_attention_processors(pipeline: StableDiffusionXLPipeline, style_aligned_args: StyleAlignedArgs | None = None): + attn_procs = {} + unet = pipeline.unet + number_of_self, number_of_cross = 0, 0 + num_self_layers = len([name for name in unet.attn_processors.keys() if 'attn1' in name]) + if style_aligned_args is None: + only_self_vec = _get_switch_vec(num_self_layers, 1) + else: + only_self_vec = _get_switch_vec(num_self_layers, style_aligned_args.only_self_level) + for i, name in enumerate(unet.attn_processors.keys()): + is_self_attention = 'attn1' in name + if is_self_attention: + number_of_self += 1 + if style_aligned_args is None or only_self_vec[i // 2]: + attn_procs[name] = DefaultAttentionProcessor() + else: + attn_procs[name] = SharedAttentionProcessor(style_aligned_args) + else: + number_of_cross += 1 + attn_procs[name] = DefaultAttentionProcessor() + + unet.set_attn_processor(attn_procs) + + +def register_shared_norm(pipeline: StableDiffusionXLPipeline, + share_group_norm: bool = True, + share_layer_norm: bool = True, + ): + def register_norm_forward(norm_layer: nn.GroupNorm | nn.LayerNorm) -> nn.GroupNorm | nn.LayerNorm: + if not hasattr(norm_layer, 'orig_forward'): + setattr(norm_layer, 'orig_forward', norm_layer.forward) # noqa + orig_forward = norm_layer.orig_forward + + def forward_(hidden_states: T) -> T: + n = hidden_states.shape[-2] + hidden_states = concat_first(hidden_states, dim=-2) + hidden_states = orig_forward(hidden_states) + return hidden_states[..., :n, :] + + norm_layer.forward = forward_ + return norm_layer + + def get_norm_layers(pipeline_, norm_layers_: dict[str, list[nn.GroupNorm | nn.LayerNorm]]): + if isinstance(pipeline_, nn.LayerNorm) and share_layer_norm: + norm_layers_['layer'].append(pipeline_) + if isinstance(pipeline_, nn.GroupNorm) and share_group_norm: + norm_layers_['group'].append(pipeline_) + else: + for layer in pipeline_.children(): + get_norm_layers(layer, norm_layers_) + + norm_layers = {'group': [], 'layer': []} + get_norm_layers(pipeline.unet, norm_layers) + return [register_norm_forward(layer) for layer in norm_layers['group']] + [register_norm_forward(layer) for layer in + norm_layers['layer']] + + +class Handler: + + def register(self, style_aligned_args: StyleAlignedArgs): + self.norm_layers = register_shared_norm(self.pipeline, style_aligned_args.share_group_norm, + style_aligned_args.share_layer_norm) + init_attention_processors(self.pipeline, style_aligned_args) + + def remove(self): + for layer in self.norm_layers: + layer.forward = layer.orig_forward + self.norm_layers = [] + init_attention_processors(self.pipeline, None) + + def __init__(self, pipeline: StableDiffusionXLPipeline): + self.pipeline = pipeline + self.norm_layers = [] diff --git a/scripts/style_aligned.py b/scripts/style_aligned.py new file mode 100644 index 000000000..25feb49bc --- /dev/null +++ b/scripts/style_aligned.py @@ -0,0 +1,117 @@ +import gradio as gr +import torch +import numpy as np +import diffusers +from modules import scripts, processing, shared, devices + + +handler = None +zts = None +supported_model_list = ['sdxl'] +orig_prompt_attention = None + + +class Script(scripts.Script): + def title(self): + return 'Style Aligned Image Generation' + + def show(self, is_img2img): + return shared.native + + def reset(self): + global handler, zts # pylint: disable=global-statement + handler = None + zts = None + shared.log.info('SA: image upload') + + def preset(self, preset): + if preset == 'text': + return [['attention', 'adain_queries', 'adain_keys'], 1.0, 0, 0.0] + if preset == 'image': + return [['group_norm', 'layer_norm', 'attention', 'adain_queries', 'adain_keys'], 1.0, 2, 0.0] + if preset == 'all': + return [['group_norm', 'layer_norm', 'attention', 'adain_queries', 'adain_keys', 'adain_values', 'full_attention_share'], 1.0, 1, 0.5] + + def ui(self, _is_img2img): # ui elements + with gr.Row(): + gr.HTML('  Style Aligned Image Generation

') + with gr.Row(): + preset = gr.Dropdown(label="Preset", choices=['text', 'image', 'all'], value='text') + scheduler = gr.Checkbox(label="Override scheduler", value=False) + with gr.Row(): + shared_opts = gr.Dropdown(label="Shared options", + multiselect=True, + choices=['group_norm', 'layer_norm', 'attention', 'adain_queries', 'adain_keys', 'adain_values', 'full_attention_share'], + value=['attention', 'adain_queries', 'adain_keys'], + ) + with gr.Row(): + shared_score_scale = gr.Slider(label="Scale", minimum=0.0, maximum=2.0, step=0.01, value=1.0) + shared_score_shift = gr.Slider(label="Shift", minimum=0, maximum=10, step=1, value=0) + only_self_level = gr.Slider(label="Level", minimum=0.0, maximum=1.0, step=0.01, value=0.0) + with gr.Row(): + prompt = gr.Textbox(lines=1, label='Optional image description', placeholder='use the style from the image') + with gr.Row(): + image = gr.Image(label='Optional image', source='upload', type='pil') + + image.change(self.reset) + preset.change(self.preset, inputs=[preset], outputs=[shared_opts, shared_score_scale, shared_score_shift, only_self_level]) + + return [image, prompt, scheduler, shared_opts, shared_score_scale, shared_score_shift, only_self_level] + + def run(self, p: processing.StableDiffusionProcessing, image, prompt, scheduler, shared_opts, shared_score_scale, shared_score_shift, only_self_level): # pylint: disable=arguments-differ + global handler, zts, orig_prompt_attention # pylint: disable=global-statement + if shared.sd_model_type not in supported_model_list: + shared.log.warning(f'SA: class={shared.sd_model.__class__.__name__} model={shared.sd_model_type} required={supported_model_list}') + return None + + from modules.style_aligned import sa_handler, inversion + + handler = sa_handler.Handler(shared.sd_model) + sa_args = sa_handler.StyleAlignedArgs( + share_group_norm='group_norm' in shared_opts, + share_layer_norm='layer_norm' in shared_opts, + share_attention='attention' in shared_opts, + adain_queries='adain_queries' in shared_opts, + adain_keys='adain_keys' in shared_opts, + adain_values='adain_values' in shared_opts, + full_attention_share='full_attention_share' in shared_opts, + shared_score_scale=float(shared_score_scale), + shared_score_shift=np.log(shared_score_shift) if shared_score_shift > 0 else 0, + only_self_level=1 if only_self_level else 0, + ) + handler.register(sa_args) + + if scheduler: + shared.sd_model.scheduler = diffusers.DDIMScheduler(beta_start=0.00085, beta_end=0.012, beta_schedule="scaled_linear", clip_sample=False, set_alpha_to_one=False) + p.sampler_name = 'None' + + if image is not None and zts is None: + shared.log.info(f'SA: inversion image={image} prompt="{prompt}"') + image = image.resize((1024, 1024)) + x0 = np.array(image).astype(np.float32) / 255.0 + shared.sd_model.scheduler = diffusers.DDIMScheduler(beta_start=0.00085, beta_end=0.012, beta_schedule="scaled_linear", clip_sample=False, set_alpha_to_one=False) + zts = inversion.ddim_inversion(shared.sd_model, x0, prompt, num_inference_steps=50, guidance_scale=2) + + p.prompt = p.prompt.splitlines() + p.batch_size = len(p.prompt) + orig_prompt_attention = shared.opts.prompt_attention + shared.opts.data['prompt_attention'] = 'fixed' # otherwise need to deal with class_tokens_mask + + if zts is not None: + processing.fix_seed(p) + zT, inversion_callback = inversion.make_inversion_callback(zts, offset=0) + generator = torch.Generator(device='cpu') + generator.manual_seed(p.seed) + latents = torch.randn(p.batch_size, 4, 128, 128, device='cpu', generator=generator, dtype=devices.dtype,).to(devices.device) + latents[0] = zT + p.task_args['latents'] = latents + p.task_args['callback_on_step_end'] = inversion_callback + + shared.log.info(f'SA: batch={p.batch_size} type={"image" if zts is not None else "text"} config={sa_args.__dict__}') + + def after(self, p: processing.StableDiffusionProcessing, *args): # pylint: disable=unused-argument + global handler # pylint: disable=global-statement + if handler is not None: + handler.remove() + handler = None + shared.opts.data['prompt_attention'] = orig_prompt_attention