nyanko7/sd-diffusers-webui
141
1import importlib2import inspect3import math4from pathlib import Path5import re6from collections import defaultdict7from typing import List, Optional, Union8 9import time10import k_diffusion11import numpy as np12import PIL13import torch14import torch.nn as nn15import torch.nn.functional as F16from einops import rearrange17from k_diffusion.external import CompVisDenoiser, CompVisVDenoiser18from modules.prompt_parser import FrozenCLIPEmbedderWithCustomWords19from torch import einsum20from torch.autograd.function import Function21 22from diffusers import DiffusionPipeline23from diffusers.utils import PIL_INTERPOLATION, is_accelerate_available24from diffusers.utils import logging, randn_tensor25 26import modules.safe as _27from safetensors.torch import load_file28 29xformers_available = False30try:31 import xformers32 33 xformers_available = True34except ImportError:35 pass36 37EPSILON = 1e-638exists = lambda val: val is not None39default = lambda val, d: val if exists(val) else d40logger = logging.get_logger(__name__) # pylint: disable=invalid-name41 42# from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.rescale_noise_cfg43def rescale_noise_cfg(noise_cfg, noise_pred_text, guidance_rescale=0.0):44 """45 Rescale `noise_cfg` according to `guidance_rescale`. Based on findings of [Common Diffusion Noise Schedules and46 Sample Steps are Flawed](https://arxiv.org/pdf/2305.08891.pdf). See Section 3.447 """48 std_text = noise_pred_text.std(dim=list(range(1, noise_pred_text.ndim)), keepdim=True)49 std_cfg = noise_cfg.std(dim=list(range(1, noise_cfg.ndim)), keepdim=True)50 # rescale the results from guidance (fixes overexposure)51 noise_pred_rescaled = noise_cfg * (std_text / std_cfg)52 # mix with the original results from guidance by factor guidance_rescale to avoid "plain looking" images53 noise_cfg = guidance_rescale * noise_pred_rescaled + (1 - guidance_rescale) * noise_cfg54 return noise_cfg55 56 57def get_attention_scores(attn, query, key, attention_mask=None):58 59 if attn.upcast_attention:60 query = query.float()61 key = key.float()62 63 attention_scores = torch.baddbmm(64 torch.empty(65 query.shape[0],66 query.shape[1],67 key.shape[1],68 dtype=query.dtype,69 device=query.device,70 ),71 query,72 key.transpose(-1, -2),73 beta=0,74 alpha=attn.scale,75 )76 77 if attention_mask is not None:78 attention_scores = attention_scores + attention_mask79 80 if attn.upcast_softmax:81 attention_scores = attention_scores.float()82 83 return attention_scores84 85 86class CrossAttnProcessor(nn.Module):87 def __call__(88 self,89 attn,90 hidden_states,91 encoder_hidden_states=None,92 attention_mask=None,93 ):94 batch_size, sequence_length, _ = hidden_states.shape95 attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size=batch_size)96 97 encoder_states = hidden_states98 is_xattn = False99 if encoder_hidden_states is not None:100 is_xattn = True101 img_state = encoder_hidden_states["img_state"]102 encoder_states = encoder_hidden_states["states"]103 weight_func = encoder_hidden_states["weight_func"]104 sigma = encoder_hidden_states["sigma"]105 106 query = attn.to_q(hidden_states)107 key = attn.to_k(encoder_states)108 value = attn.to_v(encoder_states)109 110 query = attn.head_to_batch_dim(query)111 key = attn.head_to_batch_dim(key)112 value = attn.head_to_batch_dim(value)113 114 if is_xattn and isinstance(img_state, dict):115 # use torch.baddbmm method (slow)116 attention_scores = get_attention_scores(attn, query, key, attention_mask)117 w = img_state[sequence_length].to(query.device)118 cross_attention_weight = weight_func(w, sigma, attention_scores)119 attention_scores += torch.repeat_interleave(120 cross_attention_weight, repeats=attn.heads, dim=0121 )122 123 # calc probs124 attention_probs = attention_scores.softmax(dim=-1)125 attention_probs = attention_probs.to(query.dtype)126 hidden_states = torch.bmm(attention_probs, value)127 128 elif xformers_available:129 hidden_states = xformers.ops.memory_efficient_attention(130 query.contiguous(),131 key.contiguous(),132 value.contiguous(),133 attn_bias=attention_mask,134 )135 hidden_states = hidden_states.to(query.dtype)136 137 else:138 q_bucket_size = 512139 k_bucket_size = 1024140 141 # use flash-attention142 hidden_states = FlashAttentionFunction.apply(143 query.contiguous(),144 key.contiguous(),145 value.contiguous(),146 attention_mask,147 False,148 q_bucket_size,149 k_bucket_size,150 )151 hidden_states = hidden_states.to(query.dtype)152 153 hidden_states = attn.batch_to_head_dim(hidden_states)154 155 # linear proj156 hidden_states = attn.to_out[0](hidden_states)157 158 # dropout159 hidden_states = attn.to_out[1](hidden_states)160 161 return hidden_states162 163class ModelWrapper:164 def __init__(self, model, alphas_cumprod):165 self.model = model166 self.alphas_cumprod = alphas_cumprod167 168 def apply_model(self, *args, **kwargs):169 if len(args) == 3:170 encoder_hidden_states = args[-1]171 args = args[:2]172 if kwargs.get("cond", None) is not None:173 encoder_hidden_states = kwargs.pop("cond")174 return self.model(175 *args, encoder_hidden_states=encoder_hidden_states, **kwargs176 ).sample177 178 179class StableDiffusionPipeline(DiffusionPipeline):180 181 _optional_components = ["safety_checker", "feature_extractor"]182 183 def __init__(184 self,185 vae,186 text_encoder,187 tokenizer,188 unet,189 scheduler,190 ):191 super().__init__()192 193 # get correct sigmas from LMS194 self.register_modules(195 vae=vae,196 text_encoder=text_encoder,197 tokenizer=tokenizer,198 unet=unet,199 scheduler=scheduler,200 )201 self.setup_unet(self.unet)202 self.setup_text_encoder()203 204 def setup_text_encoder(self, n=1, new_encoder=None):205 if new_encoder is not None:206 self.text_encoder = new_encoder207 208 self.prompt_parser = FrozenCLIPEmbedderWithCustomWords(self.tokenizer, self.text_encoder)209 self.prompt_parser.CLIP_stop_at_last_layers = n210 211 def setup_unet(self, unet):212 unet = unet.to(self.device)213 model = ModelWrapper(unet, self.scheduler.alphas_cumprod)214 if self.scheduler.prediction_type == "v_prediction":215 self.k_diffusion_model = CompVisVDenoiser(model)216 else:217 self.k_diffusion_model = CompVisDenoiser(model)218 219 def get_scheduler(self, scheduler_type: str):220 library = importlib.import_module("k_diffusion")221 sampling = getattr(library, "sampling")222 return getattr(sampling, scheduler_type)223 224 def encode_sketchs(self, state, scale_ratio=8, g_strength=1.0, text_ids=None):225 uncond, cond = text_ids[0], text_ids[1]226 227 img_state = []228 if state is None:229 return torch.FloatTensor(0)230 231 for k, v in state.items():232 if v["map"] is None:233 continue234 235 v_input = self.tokenizer(236 k,237 max_length=self.tokenizer.model_max_length,238 truncation=True,239 add_special_tokens=False,240 ).input_ids241 242 dotmap = v["map"] < 255243 out = dotmap.astype(float)244 if v["mask_outsides"]:245 out[out==0] = -1246 247 arr = torch.from_numpy(248 out * float(v["weight"]) * g_strength249 )250 img_state.append((v_input, arr))251 252 if len(img_state) == 0:253 return torch.FloatTensor(0)254 255 w_tensors = dict()256 cond = cond.tolist()257 uncond = uncond.tolist()258 for layer in self.unet.down_blocks:259 c = int(len(cond))260 w, h = img_state[0][1].shape261 w_r, h_r = w // scale_ratio, h // scale_ratio262 263 ret_cond_tensor = torch.zeros((1, int(w_r * h_r), c), dtype=torch.float32)264 ret_uncond_tensor = torch.zeros((1, int(w_r * h_r), c), dtype=torch.float32)265 266 for v_as_tokens, img_where_color in img_state:267 is_in = 0268 269 ret = (270 F.interpolate(271 img_where_color.unsqueeze(0).unsqueeze(1),272 scale_factor=1 / scale_ratio,273 mode="bilinear",274 align_corners=True,275 )276 .squeeze()277 .reshape(-1, 1)278 .repeat(1, len(v_as_tokens))279 )280 281 for idx, tok in enumerate(cond):282 if cond[idx : idx + len(v_as_tokens)] == v_as_tokens:283 is_in = 1284 ret_cond_tensor[0, :, idx : idx + len(v_as_tokens)] += ret285 286 for idx, tok in enumerate(uncond):287 if uncond[idx : idx + len(v_as_tokens)] == v_as_tokens:288 is_in = 1289 ret_uncond_tensor[0, :, idx : idx + len(v_as_tokens)] += ret290 291 if not is_in == 1:292 print(f"tokens {v_as_tokens} not found in text")293 294 w_tensors[w_r * h_r] = torch.cat([ret_uncond_tensor, ret_cond_tensor])295 scale_ratio *= 2296 297 return w_tensors298 299 def enable_attention_slicing(self, slice_size: Optional[Union[str, int]] = "auto"):300 r"""301 Enable sliced attention computation.302 303 When this option is enabled, the attention module will split the input tensor in slices, to compute attention304 in several steps. This is useful to save some memory in exchange for a small speed decrease.305 306 Args:307 slice_size (`str` or `int`, *optional*, defaults to `"auto"`):308 When `"auto"`, halves the input to the attention heads, so attention will be computed in two steps. If309 a number is provided, uses as many slices as `attention_head_dim // slice_size`. In this case,310 `attention_head_dim` must be a multiple of `slice_size`.311 """312 if slice_size == "auto":313 # half the attention head size is usually a good trade-off between314 # speed and memory315 slice_size = self.unet.config.attention_head_dim // 2316 self.unet.set_attention_slice(slice_size)317 318 def disable_attention_slicing(self):319 r"""320 Disable sliced attention computation. If `enable_attention_slicing` was previously invoked, this method will go321 back to computing attention in one step.322 """323 # set slice_size = `None` to disable `attention slicing`324 self.enable_attention_slicing(None)325 326 def enable_sequential_cpu_offload(self, gpu_id=0):327 r"""328 Offloads all models to CPU using accelerate, significantly reducing memory usage. When called, unet,329 text_encoder, vae and safety checker have their state dicts saved to CPU and then are moved to a330 `torch.device('meta') and loaded to GPU only when their specific submodule has its `forward` method called.331 """332 if is_accelerate_available():333 from accelerate import cpu_offload334 else:335 raise ImportError("Please install accelerate via `pip install accelerate`")336 337 device = torch.device(f"cuda:{gpu_id}")338 339 for cpu_offloaded_model in [340 self.unet,341 self.text_encoder,342 self.vae,343 self.safety_checker,344 ]:345 if cpu_offloaded_model is not None:346 cpu_offload(cpu_offloaded_model, device)347 348 @property349 def _execution_device(self):350 r"""351 Returns the device on which the pipeline's models will be executed. After calling352 `pipeline.enable_sequential_cpu_offload()` the execution device can only be inferred from Accelerate's module353 hooks.354 """355 if self.device != torch.device("meta") or not hasattr(self.unet, "_hf_hook"):356 return self.device357 for module in self.unet.modules():358 if (359 hasattr(module, "_hf_hook")360 and hasattr(module._hf_hook, "execution_device")361 and module._hf_hook.execution_device is not None362 ):363 return torch.device(module._hf_hook.execution_device)364 return self.device365 366 def decode_latents(self, latents):367 latents = latents.to(self.device, dtype=self.vae.dtype)368 latents = 1 / 0.18215 * latents369 image = self.vae.decode(latents).sample370 image = (image / 2 + 0.5).clamp(0, 1)371 # we always cast to float32 as this does not cause significant overhead and is compatible with bfloa16372 image = image.cpu().permute(0, 2, 3, 1).float().numpy()373 return image374 375 def check_inputs(self, prompt, height, width, callback_steps):376 if not isinstance(prompt, str) and not isinstance(prompt, list):377 raise ValueError(378 f"`prompt` has to be of type `str` or `list` but is {type(prompt)}"379 )380 381 if height % 8 != 0 or width % 8 != 0:382 raise ValueError(383 f"`height` and `width` have to be divisible by 8 but are {height} and {width}."384 )385 386 if (callback_steps is None) or (387 callback_steps is not None388 and (not isinstance(callback_steps, int) or callback_steps <= 0)389 ):390 raise ValueError(391 f"`callback_steps` has to be a positive integer but is {callback_steps} of type"392 f" {type(callback_steps)}."393 )394 395 def prepare_latents(396 self,397 batch_size,398 num_channels_latents,399 height,400 width,401 dtype,402 device,403 generator,404 latents=None,405 ):406 shape = (batch_size, num_channels_latents, height // 8, width // 8)407 if latents is None:408 if device.type == "mps":409 # randn does not work reproducibly on mps410 latents = torch.randn(411 shape, generator=generator, device="cpu", dtype=dtype412 ).to(device)413 else:414 latents = torch.randn(415 shape, generator=generator, device=device, dtype=dtype416 )417 else:418 # if latents.shape != shape:419 # raise ValueError(f"Unexpected latents shape, got {latents.shape}, expected {shape}")420 latents = latents.to(device)421 422 # scale the initial noise by the standard deviation required by the scheduler423 return latents424 425 def preprocess(self, image):426 if isinstance(image, torch.Tensor):427 return image428 elif isinstance(image, PIL.Image.Image):429 image = [image]430 431 if isinstance(image[0], PIL.Image.Image):432 w, h = image[0].size433 w, h = map(lambda x: x - x % 8, (w, h)) # resize to integer multiple of 8434 435 image = [436 np.array(i.resize((w, h), resample=PIL_INTERPOLATION["lanczos"]))[437 None, :438 ]439 for i in image440 ]441 image = np.concatenate(image, axis=0)442 image = np.array(image).astype(np.float32) / 255.0443 image = image.transpose(0, 3, 1, 2)444 image = 2.0 * image - 1.0445 image = torch.from_numpy(image)446 elif isinstance(image[0], torch.Tensor):447 image = torch.cat(image, dim=0)448 return image449 450 @torch.no_grad()451 def img2img(452 self,453 prompt: Union[str, List[str]],454 num_inference_steps: int = 50,455 guidance_scale: float = 7.5,456 negative_prompt: Optional[Union[str, List[str]]] = None,457 generator: Optional[torch.Generator] = None,458 image: Optional[torch.FloatTensor] = None,459 output_type: Optional[str] = "pil",460 latents=None,461 strength=1.0,462 pww_state=None,463 pww_attn_weight=1.0,464 sampler_name="",465 sampler_opt={},466 start_time=-1,467 timeout=180,468 scale_ratio=8.0,469 ):470 sampler = self.get_scheduler(sampler_name)471 if image is not None:472 image = self.preprocess(image)473 image = image.to(self.vae.device, dtype=self.vae.dtype)474 475 init_latents = self.vae.encode(image).latent_dist.sample(generator)476 latents = 0.18215 * init_latents477 478 # 2. Define call parameters479 batch_size = 1 if isinstance(prompt, str) else len(prompt)480 device = self._execution_device481 latents = latents.to(device, dtype=self.unet.dtype)482 # here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)483 # of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`484 # corresponds to doing no classifier free guidance.485 do_classifier_free_guidance = True486 if guidance_scale <= 1.0:487 raise ValueError("has to use guidance_scale")488 489 # 3. Encode input prompt490 text_ids, text_embeddings = self.prompt_parser([negative_prompt, prompt])491 text_embeddings = text_embeddings.to(self.unet.dtype)492 493 init_timestep = (494 int(num_inference_steps / min(strength, 0.999)) if strength > 0 else 0495 )496 sigmas = self.get_sigmas(init_timestep, sampler_opt).to(497 text_embeddings.device, dtype=text_embeddings.dtype498 )499 500 t_start = max(init_timestep - num_inference_steps, 0)501 sigma_sched = sigmas[t_start:]502 503 noise = randn_tensor(504 latents.shape,505 generator=generator,506 device=device,507 dtype=text_embeddings.dtype,508 )509 latents = latents.to(device)510 latents = latents + noise * sigma_sched[0]511 512 # 5. Prepare latent variables513 self.k_diffusion_model.sigmas = self.k_diffusion_model.sigmas.to(latents.device)514 self.k_diffusion_model.log_sigmas = self.k_diffusion_model.log_sigmas.to(515 latents.device516 )517 518 img_state = self.encode_sketchs(519 pww_state,520 g_strength=pww_attn_weight,521 text_ids=text_ids,522 )523 524 def model_fn(x, sigma):525 526 if start_time > 0 and timeout > 0:527 assert (time.time() - start_time) < timeout, "inference process timed out"528 529 latent_model_input = torch.cat([x] * 2)530 weight_func = lambda w, sigma, qk: w * math.log(1 + sigma) * qk.max()531 encoder_state = {532 "img_state": img_state,533 "states": text_embeddings,534 "sigma": sigma[0],535 "weight_func": weight_func,536 }537 538 noise_pred = self.k_diffusion_model(539 latent_model_input, sigma, cond=encoder_state540 )541 noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)542 noise_pred = noise_pred_uncond + guidance_scale * (543 noise_pred_text - noise_pred_uncond544 )545 546 # noise_pred = rescale_noise_cfg(noise_pred, noise_pred_text, guidance_rescale=0.7)547 return noise_pred548 549 sampler_args = self.get_sampler_extra_args_i2i(sigma_sched, sampler)550 latents = sampler(model_fn, latents, **sampler_args)551 552 # 8. Post-processing553 image = self.decode_latents(latents)554 555 # 10. Convert to PIL556 if output_type == "pil":557 image = self.numpy_to_pil(image)558 559 return (image,)560 561 def get_sigmas(self, steps, params):562 discard_next_to_last_sigma = params.get("discard_next_to_last_sigma", False)563 steps += 1 if discard_next_to_last_sigma else 0564 565 if params.get("scheduler", None) == "karras":566 sigma_min, sigma_max = (567 self.k_diffusion_model.sigmas[0].item(),568 self.k_diffusion_model.sigmas[-1].item(),569 )570 sigmas = k_diffusion.sampling.get_sigmas_karras(571 n=steps, sigma_min=sigma_min, sigma_max=sigma_max, device=self.device572 )573 else:574 sigmas = self.k_diffusion_model.get_sigmas(steps)575 576 if discard_next_to_last_sigma:577 sigmas = torch.cat([sigmas[:-2], sigmas[-1:]])578 579 return sigmas580 581 # https://github.com/AUTOMATIC1111/stable-diffusion-webui/blob/48a15821de768fea76e66f26df83df3fddf18f4b/modules/sd_samplers.py#L454582 def get_sampler_extra_args_t2i(self, sigmas, eta, steps, func):583 extra_params_kwargs = {}584 585 if "eta" in inspect.signature(func).parameters:586 extra_params_kwargs["eta"] = eta587 588 if "sigma_min" in inspect.signature(func).parameters:589 extra_params_kwargs["sigma_min"] = sigmas[0].item()590 extra_params_kwargs["sigma_max"] = sigmas[-1].item()591 592 if "n" in inspect.signature(func).parameters:593 extra_params_kwargs["n"] = steps594 else:595 extra_params_kwargs["sigmas"] = sigmas596 597 return extra_params_kwargs598 599 # https://github.com/AUTOMATIC1111/stable-diffusion-webui/blob/48a15821de768fea76e66f26df83df3fddf18f4b/modules/sd_samplers.py#L454600 def get_sampler_extra_args_i2i(self, sigmas, func):601 extra_params_kwargs = {}602 603 if "sigma_min" in inspect.signature(func).parameters:604 ## last sigma is zero which isn't allowed by DPM Fast & Adaptive so taking value before last605 extra_params_kwargs["sigma_min"] = sigmas[-2]606 607 if "sigma_max" in inspect.signature(func).parameters:608 extra_params_kwargs["sigma_max"] = sigmas[0]609 610 if "n" in inspect.signature(func).parameters:611 extra_params_kwargs["n"] = len(sigmas) - 1612 613 if "sigma_sched" in inspect.signature(func).parameters:614 extra_params_kwargs["sigma_sched"] = sigmas615 616 if "sigmas" in inspect.signature(func).parameters:617 extra_params_kwargs["sigmas"] = sigmas618 619 return extra_params_kwargs620 621 @torch.no_grad()622 def txt2img(623 self,624 prompt: Union[str, List[str]],625 height: int = 512,626 width: int = 512,627 num_inference_steps: int = 50,628 guidance_scale: float = 7.5,629 negative_prompt: Optional[Union[str, List[str]]] = None,630 eta: float = 0.0,631 generator: Optional[torch.Generator] = None,632 latents: Optional[torch.FloatTensor] = None,633 output_type: Optional[str] = "pil",634 callback_steps: Optional[int] = 1,635 upscale=False,636 upscale_x: float = 2.0,637 upscale_method: str = "bicubic",638 upscale_antialias: bool = False,639 upscale_denoising_strength: int = 0.7,640 pww_state=None,641 pww_attn_weight=1.0,642 sampler_name="",643 sampler_opt={},644 start_time=-1,645 timeout=180,646 ):647 sampler = self.get_scheduler(sampler_name)648 # 1. Check inputs. Raise error if not correct649 self.check_inputs(prompt, height, width, callback_steps)650 651 # 2. Define call parameters652 batch_size = 1 if isinstance(prompt, str) else len(prompt)653 device = self._execution_device654 # here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)655 # of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`656 # corresponds to doing no classifier free guidance.657 do_classifier_free_guidance = True658 if guidance_scale <= 1.0:659 raise ValueError("has to use guidance_scale")660 661 # 3. Encode input prompt662 text_ids, text_embeddings = self.prompt_parser([negative_prompt, prompt])663 text_embeddings = text_embeddings.to(self.unet.dtype)664 665 # 4. Prepare timesteps666 sigmas = self.get_sigmas(num_inference_steps, sampler_opt).to(667 text_embeddings.device, dtype=text_embeddings.dtype668 )669 670 # 5. Prepare latent variables671 num_channels_latents = self.unet.in_channels672 latents = self.prepare_latents(673 batch_size,674 num_channels_latents,675 height,676 width,677 text_embeddings.dtype,678 device,679 generator,680 latents,681 )682 latents = latents * sigmas[0]683 self.k_diffusion_model.sigmas = self.k_diffusion_model.sigmas.to(latents.device)684 self.k_diffusion_model.log_sigmas = self.k_diffusion_model.log_sigmas.to(685 latents.device686 )687 688 img_state = self.encode_sketchs(689 pww_state,690 g_strength=pww_attn_weight,691 text_ids=text_ids,692 )693 694 def model_fn(x, sigma):695 696 if start_time > 0 and timeout > 0:697 assert (time.time() - start_time) < timeout, "inference process timed out"698 699 latent_model_input = torch.cat([x] * 2)700 weight_func = lambda w, sigma, qk: w * math.log(1 + sigma) * qk.max()701 encoder_state = {702 "img_state": img_state,703 "states": text_embeddings,704 "sigma": sigma[0],705 "weight_func": weight_func,706 }707 708 noise_pred = self.k_diffusion_model(709 latent_model_input, sigma, cond=encoder_state710 )711 noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)712 noise_pred = noise_pred_uncond + guidance_scale * (713 noise_pred_text - noise_pred_uncond714 )715 716 # noise_pred = rescale_noise_cfg(noise_pred, noise_pred_text, guidance_rescale=0.7)717 return noise_pred718 719 extra_args = self.get_sampler_extra_args_t2i(720 sigmas, eta, num_inference_steps, sampler721 )722 latents = sampler(model_fn, latents, **extra_args)723 724 if upscale:725 target_height = height * upscale_x726 target_width = width * upscale_x727 vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1)728 latents = torch.nn.functional.interpolate(729 latents,730 size=(731 int(target_height // vae_scale_factor),732 int(target_width // vae_scale_factor),733 ),734 mode=upscale_method,735 antialias=upscale_antialias,736 )737 return self.img2img(738 prompt=prompt,739 num_inference_steps=num_inference_steps,740 guidance_scale=guidance_scale,741 negative_prompt=negative_prompt,742 generator=generator,743 latents=latents,744 strength=upscale_denoising_strength,745 sampler_name=sampler_name,746 sampler_opt=sampler_opt,747 pww_state=None,748 pww_attn_weight=pww_attn_weight / 2,749 )750 751 # 8. Post-processing752 image = self.decode_latents(latents)753 754 # 10. Convert to PIL755 if output_type == "pil":756 image = self.numpy_to_pil(image)757 758 return (image,)759 760 761class FlashAttentionFunction(Function):762 @staticmethod763 @torch.no_grad()764 def forward(ctx, q, k, v, mask, causal, q_bucket_size, k_bucket_size):765 """Algorithm 2 in the paper"""766 767 device = q.device768 max_neg_value = -torch.finfo(q.dtype).max769 qk_len_diff = max(k.shape[-2] - q.shape[-2], 0)770 771 o = torch.zeros_like(q)772 all_row_sums = torch.zeros((*q.shape[:-1], 1), device=device)773 all_row_maxes = torch.full((*q.shape[:-1], 1), max_neg_value, device=device)774 775 scale = q.shape[-1] ** -0.5776 777 if not exists(mask):778 mask = (None,) * math.ceil(q.shape[-2] / q_bucket_size)779 else:780 mask = rearrange(mask, "b n -> b 1 1 n")781 mask = mask.split(q_bucket_size, dim=-1)782 783 row_splits = zip(784 q.split(q_bucket_size, dim=-2),785 o.split(q_bucket_size, dim=-2),786 mask,787 all_row_sums.split(q_bucket_size, dim=-2),788 all_row_maxes.split(q_bucket_size, dim=-2),789 )790 791 for ind, (qc, oc, row_mask, row_sums, row_maxes) in enumerate(row_splits):792 q_start_index = ind * q_bucket_size - qk_len_diff793 794 col_splits = zip(795 k.split(k_bucket_size, dim=-2),796 v.split(k_bucket_size, dim=-2),797 )798 799 for k_ind, (kc, vc) in enumerate(col_splits):800 k_start_index = k_ind * k_bucket_size801 802 attn_weights = einsum("... i d, ... j d -> ... i j", qc, kc) * scale803 804 if exists(row_mask):805 attn_weights.masked_fill_(~row_mask, max_neg_value)806 807 if causal and q_start_index < (k_start_index + k_bucket_size - 1):808 causal_mask = torch.ones(809 (qc.shape[-2], kc.shape[-2]), dtype=torch.bool, device=device810 ).triu(q_start_index - k_start_index + 1)811 attn_weights.masked_fill_(causal_mask, max_neg_value)812 813 block_row_maxes = attn_weights.amax(dim=-1, keepdims=True)814 attn_weights -= block_row_maxes815 exp_weights = torch.exp(attn_weights)816 817 if exists(row_mask):818 exp_weights.masked_fill_(~row_mask, 0.0)819 820 block_row_sums = exp_weights.sum(dim=-1, keepdims=True).clamp(821 min=EPSILON822 )823 824 new_row_maxes = torch.maximum(block_row_maxes, row_maxes)825 826 exp_values = einsum("... i j, ... j d -> ... i d", exp_weights, vc)827 828 exp_row_max_diff = torch.exp(row_maxes - new_row_maxes)829 exp_block_row_max_diff = torch.exp(block_row_maxes - new_row_maxes)830 831 new_row_sums = (832 exp_row_max_diff * row_sums833 + exp_block_row_max_diff * block_row_sums834 )835 836 oc.mul_((row_sums / new_row_sums) * exp_row_max_diff).add_(837 (exp_block_row_max_diff / new_row_sums) * exp_values838 )839 840 row_maxes.copy_(new_row_maxes)841 row_sums.copy_(new_row_sums)842 843 lse = all_row_sums.log() + all_row_maxes844 845 ctx.args = (causal, scale, mask, q_bucket_size, k_bucket_size)846 ctx.save_for_backward(q, k, v, o, lse)847 848 return o849 850 @staticmethod851 @torch.no_grad()852 def backward(ctx, do):853 """Algorithm 4 in the paper"""854 855 causal, scale, mask, q_bucket_size, k_bucket_size = ctx.args856 q, k, v, o, lse = ctx.saved_tensors857 858 device = q.device859 860 max_neg_value = -torch.finfo(q.dtype).max861 qk_len_diff = max(k.shape[-2] - q.shape[-2], 0)862 863 dq = torch.zeros_like(q)864 dk = torch.zeros_like(k)865 dv = torch.zeros_like(v)866 867 row_splits = zip(868 q.split(q_bucket_size, dim=-2),869 o.split(q_bucket_size, dim=-2),870 do.split(q_bucket_size, dim=-2),871 mask,872 lse.split(q_bucket_size, dim=-2),873 dq.split(q_bucket_size, dim=-2),874 )875 876 for ind, (qc, oc, doc, row_mask, lsec, dqc) in enumerate(row_splits):877 q_start_index = ind * q_bucket_size - qk_len_diff878 879 col_splits = zip(880 k.split(k_bucket_size, dim=-2),881 v.split(k_bucket_size, dim=-2),882 dk.split(k_bucket_size, dim=-2),883 dv.split(k_bucket_size, dim=-2),884 )885 886 for k_ind, (kc, vc, dkc, dvc) in enumerate(col_splits):887 k_start_index = k_ind * k_bucket_size888 889 attn_weights = einsum("... i d, ... j d -> ... i j", qc, kc) * scale890 891 if causal and q_start_index < (k_start_index + k_bucket_size - 1):892 causal_mask = torch.ones(893 (qc.shape[-2], kc.shape[-2]), dtype=torch.bool, device=device894 ).triu(q_start_index - k_start_index + 1)895 attn_weights.masked_fill_(causal_mask, max_neg_value)896 897 p = torch.exp(attn_weights - lsec)898 899 if exists(row_mask):900 p.masked_fill_(~row_mask, 0.0)901 902 dv_chunk = einsum("... i j, ... i d -> ... j d", p, doc)903 dp = einsum("... i d, ... j d -> ... i j", doc, vc)904 905 D = (doc * oc).sum(dim=-1, keepdims=True)906 ds = p * scale * (dp - D)907 908 dq_chunk = einsum("... i j, ... j d -> ... i d", ds, kc)909 dk_chunk = einsum("... i j, ... i d -> ... j d", ds, qc)910 911 dqc.add_(dq_chunk)912 dkc.add_(dk_chunk)913 dvc.add_(dv_chunk)914 915 return dq, dk, dv, None, None, None, None916 