Team Ai
Apppublic

nyanko7/sd-diffusers-webui

sourceHugging Faceopenrailupdated 3y agoView on Hugging Face
141likes
model.py916 linesDownload Raw Back to modules
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