Team Ai
Apppublic

Covert1107/sd-diffusers-webui

sourceHugging Faceopenrailupdated 4y agoView on Hugging Face
2likes
model.py898 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 43def get_attention_scores(attn, query, key, attention_mask=None):44 45    if attn.upcast_attention:46        query = query.float()47        key = key.float()48 49    attention_scores = torch.baddbmm(50        torch.empty(51            query.shape[0],52            query.shape[1],53            key.shape[1],54            dtype=query.dtype,55            device=query.device,56        ),57        query,58        key.transpose(-1, -2),59        beta=0,60        alpha=attn.scale,61    )62 63    if attention_mask is not None:64        attention_scores = attention_scores + attention_mask65 66    if attn.upcast_softmax:67        attention_scores = attention_scores.float()68 69    return attention_scores70 71 72class CrossAttnProcessor(nn.Module):73    def __call__(74        self,75        attn,76        hidden_states,77        encoder_hidden_states=None,78        attention_mask=None,79    ):80        batch_size, sequence_length, _ = hidden_states.shape81        attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length)82 83        encoder_states = hidden_states84        is_xattn = False85        if encoder_hidden_states is not None:86            is_xattn = True87            img_state = encoder_hidden_states["img_state"]88            encoder_states = encoder_hidden_states["states"]89            weight_func = encoder_hidden_states["weight_func"]90            sigma = encoder_hidden_states["sigma"]91 92        query = attn.to_q(hidden_states)93        key = attn.to_k(encoder_states)94        value = attn.to_v(encoder_states)95 96        query = attn.head_to_batch_dim(query)97        key = attn.head_to_batch_dim(key)98        value = attn.head_to_batch_dim(value)99 100        if is_xattn and isinstance(img_state, dict):101            # use torch.baddbmm method (slow)102            attention_scores = get_attention_scores(attn, query, key, attention_mask)103            w = img_state[sequence_length].to(query.device)104            cross_attention_weight = weight_func(w, sigma, attention_scores)105            attention_scores += torch.repeat_interleave(106                cross_attention_weight, repeats=attn.heads, dim=0107            )108 109            # calc probs110            attention_probs = attention_scores.softmax(dim=-1)111            attention_probs = attention_probs.to(query.dtype)112            hidden_states = torch.bmm(attention_probs, value)113 114        elif xformers_available:115            hidden_states = xformers.ops.memory_efficient_attention(116                query.contiguous(),117                key.contiguous(),118                value.contiguous(),119                attn_bias=attention_mask,120            )121            hidden_states = hidden_states.to(query.dtype)122 123        else:124            q_bucket_size = 512125            k_bucket_size = 1024126 127            # use flash-attention128            hidden_states = FlashAttentionFunction.apply(129                query.contiguous(),130                key.contiguous(),131                value.contiguous(),132                attention_mask,133                False,134                q_bucket_size,135                k_bucket_size,136            )137            hidden_states = hidden_states.to(query.dtype)138 139        hidden_states = attn.batch_to_head_dim(hidden_states)140 141        # linear proj142        hidden_states = attn.to_out[0](hidden_states)143 144        # dropout145        hidden_states = attn.to_out[1](hidden_states)146 147        return hidden_states148 149class ModelWrapper:150    def __init__(self, model, alphas_cumprod):151        self.model = model152        self.alphas_cumprod = alphas_cumprod153 154    def apply_model(self, *args, **kwargs):155        if len(args) == 3:156            encoder_hidden_states = args[-1]157            args = args[:2]158        if kwargs.get("cond", None) is not None:159            encoder_hidden_states = kwargs.pop("cond")160        return self.model(161            *args, encoder_hidden_states=encoder_hidden_states, **kwargs162        ).sample163 164 165class StableDiffusionPipeline(DiffusionPipeline):166 167    _optional_components = ["safety_checker", "feature_extractor"]168 169    def __init__(170        self,171        vae,172        text_encoder,173        tokenizer,174        unet,175        scheduler,176    ):177        super().__init__()178 179        # get correct sigmas from LMS180        self.register_modules(181            vae=vae,182            text_encoder=text_encoder,183            tokenizer=tokenizer,184            unet=unet,185            scheduler=scheduler,186        )187        self.setup_unet(self.unet)188        self.setup_text_encoder()189 190    def setup_text_encoder(self, n=1, new_encoder=None):191        if new_encoder is not None:192            self.text_encoder = new_encoder193 194        self.prompt_parser = FrozenCLIPEmbedderWithCustomWords(self.tokenizer, self.text_encoder)195        self.prompt_parser.CLIP_stop_at_last_layers = n196 197    def setup_unet(self, unet):198        unet = unet.to(self.device)199        model = ModelWrapper(unet, self.scheduler.alphas_cumprod)200        if self.scheduler.prediction_type == "v_prediction":201            self.k_diffusion_model = CompVisVDenoiser(model)202        else:203            self.k_diffusion_model = CompVisDenoiser(model)204 205    def get_scheduler(self, scheduler_type: str):206        library = importlib.import_module("k_diffusion")207        sampling = getattr(library, "sampling")208        return getattr(sampling, scheduler_type)209 210    def encode_sketchs(self, state, scale_ratio=8, g_strength=1.0, text_ids=None):211        uncond, cond = text_ids[0], text_ids[1]212 213        img_state = []214        if state is None:215            return torch.FloatTensor(0)216 217        for k, v in state.items():218            if v["map"] is None:219                continue220 221            v_input = self.tokenizer(222                k,223                max_length=self.tokenizer.model_max_length,224                truncation=True,225                add_special_tokens=False,226            ).input_ids227 228            dotmap = v["map"] < 255229            out = dotmap.astype(float)230            if v["mask_outsides"]:231                out[out==0] = -1232                233            arr = torch.from_numpy(234                out * float(v["weight"]) * g_strength235            )236            img_state.append((v_input, arr))237 238        if len(img_state) == 0:239            return torch.FloatTensor(0)240 241        w_tensors = dict()242        cond = cond.tolist()243        uncond = uncond.tolist()244        for layer in self.unet.down_blocks:245            c = int(len(cond))246            w, h = img_state[0][1].shape247            w_r, h_r = w // scale_ratio, h // scale_ratio248 249            ret_cond_tensor = torch.zeros((1, int(w_r * h_r), c), dtype=torch.float32)250            ret_uncond_tensor = torch.zeros((1, int(w_r * h_r), c), dtype=torch.float32)251 252            for v_as_tokens, img_where_color in img_state:253                is_in = 0254 255                ret = (256                    F.interpolate(257                        img_where_color.unsqueeze(0).unsqueeze(1),258                        scale_factor=1 / scale_ratio,259                        mode="bilinear",260                        align_corners=True,261                    )262                    .squeeze()263                    .reshape(-1, 1)264                    .repeat(1, len(v_as_tokens))265                )266 267                for idx, tok in enumerate(cond):268                    if cond[idx : idx + len(v_as_tokens)] == v_as_tokens:269                        is_in = 1270                        ret_cond_tensor[0, :, idx : idx + len(v_as_tokens)] += ret271 272                for idx, tok in enumerate(uncond):273                    if uncond[idx : idx + len(v_as_tokens)] == v_as_tokens:274                        is_in = 1275                        ret_uncond_tensor[0, :, idx : idx + len(v_as_tokens)] += ret276 277                if not is_in == 1:278                    print(f"tokens {v_as_tokens} not found in text")279 280            w_tensors[w_r * h_r] = torch.cat([ret_uncond_tensor, ret_cond_tensor])281            scale_ratio *= 2282 283        return w_tensors284 285    def enable_attention_slicing(self, slice_size: Optional[Union[str, int]] = "auto"):286        r"""287        Enable sliced attention computation.288 289        When this option is enabled, the attention module will split the input tensor in slices, to compute attention290        in several steps. This is useful to save some memory in exchange for a small speed decrease.291 292        Args:293            slice_size (`str` or `int`, *optional*, defaults to `"auto"`):294                When `"auto"`, halves the input to the attention heads, so attention will be computed in two steps. If295                a number is provided, uses as many slices as `attention_head_dim // slice_size`. In this case,296                `attention_head_dim` must be a multiple of `slice_size`.297        """298        if slice_size == "auto":299            # half the attention head size is usually a good trade-off between300            # speed and memory301            slice_size = self.unet.config.attention_head_dim // 2302        self.unet.set_attention_slice(slice_size)303 304    def disable_attention_slicing(self):305        r"""306        Disable sliced attention computation. If `enable_attention_slicing` was previously invoked, this method will go307        back to computing attention in one step.308        """309        # set slice_size = `None` to disable `attention slicing`310        self.enable_attention_slicing(None)311 312    def enable_sequential_cpu_offload(self, gpu_id=0):313        r"""314        Offloads all models to CPU using accelerate, significantly reducing memory usage. When called, unet,315        text_encoder, vae and safety checker have their state dicts saved to CPU and then are moved to a316        `torch.device('meta') and loaded to GPU only when their specific submodule has its `forward` method called.317        """318        if is_accelerate_available():319            from accelerate import cpu_offload320        else:321            raise ImportError("Please install accelerate via `pip install accelerate`")322 323        device = torch.device(f"cuda:{gpu_id}")324 325        for cpu_offloaded_model in [326            self.unet,327            self.text_encoder,328            self.vae,329            self.safety_checker,330        ]:331            if cpu_offloaded_model is not None:332                cpu_offload(cpu_offloaded_model, device)333 334    @property335    def _execution_device(self):336        r"""337        Returns the device on which the pipeline's models will be executed. After calling338        `pipeline.enable_sequential_cpu_offload()` the execution device can only be inferred from Accelerate's module339        hooks.340        """341        if self.device != torch.device("meta") or not hasattr(self.unet, "_hf_hook"):342            return self.device343        for module in self.unet.modules():344            if (345                hasattr(module, "_hf_hook")346                and hasattr(module._hf_hook, "execution_device")347                and module._hf_hook.execution_device is not None348            ):349                return torch.device(module._hf_hook.execution_device)350        return self.device351 352    def decode_latents(self, latents):353        latents = latents.to(self.device, dtype=self.vae.dtype)354        latents = 1 / 0.18215 * latents355        image = self.vae.decode(latents).sample356        image = (image / 2 + 0.5).clamp(0, 1)357        # we always cast to float32 as this does not cause significant overhead and is compatible with bfloa16358        image = image.cpu().permute(0, 2, 3, 1).float().numpy()359        return image360 361    def check_inputs(self, prompt, height, width, callback_steps):362        if not isinstance(prompt, str) and not isinstance(prompt, list):363            raise ValueError(364                f"`prompt` has to be of type `str` or `list` but is {type(prompt)}"365            )366 367        if height % 8 != 0 or width % 8 != 0:368            raise ValueError(369                f"`height` and `width` have to be divisible by 8 but are {height} and {width}."370            )371 372        if (callback_steps is None) or (373            callback_steps is not None374            and (not isinstance(callback_steps, int) or callback_steps <= 0)375        ):376            raise ValueError(377                f"`callback_steps` has to be a positive integer but is {callback_steps} of type"378                f" {type(callback_steps)}."379            )380 381    def prepare_latents(382        self,383        batch_size,384        num_channels_latents,385        height,386        width,387        dtype,388        device,389        generator,390        latents=None,391    ):392        shape = (batch_size, num_channels_latents, height // 8, width // 8)393        if latents is None:394            if device.type == "mps":395                # randn does not work reproducibly on mps396                latents = torch.randn(397                    shape, generator=generator, device="cpu", dtype=dtype398                ).to(device)399            else:400                latents = torch.randn(401                    shape, generator=generator, device=device, dtype=dtype402                )403        else:404            # if latents.shape != shape:405            #     raise ValueError(f"Unexpected latents shape, got {latents.shape}, expected {shape}")406            latents = latents.to(device)407 408        # scale the initial noise by the standard deviation required by the scheduler409        return latents410 411    def preprocess(self, image):412        if isinstance(image, torch.Tensor):413            return image414        elif isinstance(image, PIL.Image.Image):415            image = [image]416 417        if isinstance(image[0], PIL.Image.Image):418            w, h = image[0].size419            w, h = map(lambda x: x - x % 8, (w, h))  # resize to integer multiple of 8420 421            image = [422                np.array(i.resize((w, h), resample=PIL_INTERPOLATION["lanczos"]))[423                    None, :424                ]425                for i in image426            ]427            image = np.concatenate(image, axis=0)428            image = np.array(image).astype(np.float32) / 255.0429            image = image.transpose(0, 3, 1, 2)430            image = 2.0 * image - 1.0431            image = torch.from_numpy(image)432        elif isinstance(image[0], torch.Tensor):433            image = torch.cat(image, dim=0)434        return image435 436    @torch.no_grad()437    def img2img(438        self,439        prompt: Union[str, List[str]],440        num_inference_steps: int = 50,441        guidance_scale: float = 7.5,442        negative_prompt: Optional[Union[str, List[str]]] = None,443        generator: Optional[torch.Generator] = None,444        image: Optional[torch.FloatTensor] = None,445        output_type: Optional[str] = "pil",446        latents=None,447        strength=1.0,448        pww_state=None,449        pww_attn_weight=1.0,450        sampler_name="",451        sampler_opt={},452        start_time=-1,453        timeout=180,454        scale_ratio=8.0,455    ):456        sampler = self.get_scheduler(sampler_name)457        if image is not None:458            image = self.preprocess(image)459            image = image.to(self.vae.device, dtype=self.vae.dtype)460 461            init_latents = self.vae.encode(image).latent_dist.sample(generator)462            latents = 0.18215 * init_latents463 464        # 2. Define call parameters465        batch_size = 1 if isinstance(prompt, str) else len(prompt)466        device = self._execution_device467        latents = latents.to(device, dtype=self.unet.dtype)468        # here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)469        # of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`470        # corresponds to doing no classifier free guidance.471        do_classifier_free_guidance = True472        if guidance_scale <= 1.0:473            raise ValueError("has to use guidance_scale")474 475        # 3. Encode input prompt476        text_ids, text_embeddings = self.prompt_parser([negative_prompt, prompt])477        text_embeddings = text_embeddings.to(self.unet.dtype)478 479        init_timestep = (480            int(num_inference_steps / min(strength, 0.999)) if strength > 0 else 0481        )482        sigmas = self.get_sigmas(init_timestep, sampler_opt).to(483            text_embeddings.device, dtype=text_embeddings.dtype484        )485 486        t_start = max(init_timestep - num_inference_steps, 0)487        sigma_sched = sigmas[t_start:]488 489        noise = randn_tensor(490            latents.shape,491            generator=generator,492            device=device,493            dtype=text_embeddings.dtype,494        )495        latents = latents.to(device)496        latents = latents + noise * sigma_sched[0]497 498        # 5. Prepare latent variables499        self.k_diffusion_model.sigmas = self.k_diffusion_model.sigmas.to(latents.device)500        self.k_diffusion_model.log_sigmas = self.k_diffusion_model.log_sigmas.to(501            latents.device502        )503 504        img_state = self.encode_sketchs(505            pww_state,506            g_strength=pww_attn_weight,507            text_ids=text_ids,508        )509 510        def model_fn(x, sigma):511 512            if start_time > 0 and timeout > 0:513                assert (time.time() - start_time) < timeout, "inference process timed out"514 515            latent_model_input = torch.cat([x] * 2)516            weight_func = lambda w, sigma, qk: w * math.log(1 + sigma) * qk.max()517            encoder_state = {518                "img_state": img_state,519                "states": text_embeddings,520                "sigma": sigma[0],521                "weight_func": weight_func,522            }523 524            noise_pred = self.k_diffusion_model(525                latent_model_input, sigma, cond=encoder_state526            )527            noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)528            noise_pred = noise_pred_uncond + guidance_scale * (529                noise_pred_text - noise_pred_uncond530            )531            return noise_pred532 533        sampler_args = self.get_sampler_extra_args_i2i(sigma_sched, sampler)534        latents = sampler(model_fn, latents, **sampler_args)535 536        # 8. Post-processing537        image = self.decode_latents(latents)538 539        # 10. Convert to PIL540        if output_type == "pil":541            image = self.numpy_to_pil(image)542 543        return (image,)544 545    def get_sigmas(self, steps, params):546        discard_next_to_last_sigma = params.get("discard_next_to_last_sigma", False)547        steps += 1 if discard_next_to_last_sigma else 0548 549        if params.get("scheduler", None) == "karras":550            sigma_min, sigma_max = (551                self.k_diffusion_model.sigmas[0].item(),552                self.k_diffusion_model.sigmas[-1].item(),553            )554            sigmas = k_diffusion.sampling.get_sigmas_karras(555                n=steps, sigma_min=sigma_min, sigma_max=sigma_max, device=self.device556            )557        else:558            sigmas = self.k_diffusion_model.get_sigmas(steps)559 560        if discard_next_to_last_sigma:561            sigmas = torch.cat([sigmas[:-2], sigmas[-1:]])562 563        return sigmas564 565    # https://github.com/AUTOMATIC1111/stable-diffusion-webui/blob/48a15821de768fea76e66f26df83df3fddf18f4b/modules/sd_samplers.py#L454566    def get_sampler_extra_args_t2i(self, sigmas, eta, steps, func):567        extra_params_kwargs = {}568 569        if "eta" in inspect.signature(func).parameters:570            extra_params_kwargs["eta"] = eta571 572        if "sigma_min" in inspect.signature(func).parameters:573            extra_params_kwargs["sigma_min"] = sigmas[0].item()574            extra_params_kwargs["sigma_max"] = sigmas[-1].item()575 576        if "n" in inspect.signature(func).parameters:577            extra_params_kwargs["n"] = steps578        else:579            extra_params_kwargs["sigmas"] = sigmas580 581        return extra_params_kwargs582 583    # https://github.com/AUTOMATIC1111/stable-diffusion-webui/blob/48a15821de768fea76e66f26df83df3fddf18f4b/modules/sd_samplers.py#L454584    def get_sampler_extra_args_i2i(self, sigmas, func):585        extra_params_kwargs = {}586 587        if "sigma_min" in inspect.signature(func).parameters:588            ## last sigma is zero which isn't allowed by DPM Fast & Adaptive so taking value before last589            extra_params_kwargs["sigma_min"] = sigmas[-2]590 591        if "sigma_max" in inspect.signature(func).parameters:592            extra_params_kwargs["sigma_max"] = sigmas[0]593 594        if "n" in inspect.signature(func).parameters:595            extra_params_kwargs["n"] = len(sigmas) - 1596 597        if "sigma_sched" in inspect.signature(func).parameters:598            extra_params_kwargs["sigma_sched"] = sigmas599 600        if "sigmas" in inspect.signature(func).parameters:601            extra_params_kwargs["sigmas"] = sigmas602 603        return extra_params_kwargs604 605    @torch.no_grad()606    def txt2img(607        self,608        prompt: Union[str, List[str]],609        height: int = 512,610        width: int = 512,611        num_inference_steps: int = 50,612        guidance_scale: float = 7.5,613        negative_prompt: Optional[Union[str, List[str]]] = None,614        eta: float = 0.0,615        generator: Optional[torch.Generator] = None,616        latents: Optional[torch.FloatTensor] = None,617        output_type: Optional[str] = "pil",618        callback_steps: Optional[int] = 1,619        upscale=False,620        upscale_x: float = 2.0,621        upscale_method: str = "bicubic",622        upscale_antialias: bool = False,623        upscale_denoising_strength: int = 0.7,624        pww_state=None,625        pww_attn_weight=1.0,626        sampler_name="",627        sampler_opt={},628        start_time=-1,629        timeout=180,630    ):631        sampler = self.get_scheduler(sampler_name)632        # 1. Check inputs. Raise error if not correct633        self.check_inputs(prompt, height, width, callback_steps)634 635        # 2. Define call parameters636        batch_size = 1 if isinstance(prompt, str) else len(prompt)637        device = self._execution_device638        # here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)639        # of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`640        # corresponds to doing no classifier free guidance.641        do_classifier_free_guidance = True642        if guidance_scale <= 1.0:643            raise ValueError("has to use guidance_scale")644 645        # 3. Encode input prompt646        text_ids, text_embeddings = self.prompt_parser([negative_prompt, prompt])647        text_embeddings = text_embeddings.to(self.unet.dtype)648 649        # 4. Prepare timesteps650        sigmas = self.get_sigmas(num_inference_steps, sampler_opt).to(651            text_embeddings.device, dtype=text_embeddings.dtype652        )653 654        # 5. Prepare latent variables655        num_channels_latents = self.unet.in_channels656        latents = self.prepare_latents(657            batch_size,658            num_channels_latents,659            height,660            width,661            text_embeddings.dtype,662            device,663            generator,664            latents,665        )666        latents = latents * sigmas[0]667        self.k_diffusion_model.sigmas = self.k_diffusion_model.sigmas.to(latents.device)668        self.k_diffusion_model.log_sigmas = self.k_diffusion_model.log_sigmas.to(669            latents.device670        )671 672        img_state = self.encode_sketchs(673            pww_state,674            g_strength=pww_attn_weight,675            text_ids=text_ids,676        )677 678        def model_fn(x, sigma):679 680            if start_time > 0 and timeout > 0:681                assert (time.time() - start_time) < timeout, "inference process timed out"682 683            latent_model_input = torch.cat([x] * 2)684            weight_func = lambda w, sigma, qk: w * math.log(1 + sigma) * qk.max()685            encoder_state = {686                "img_state": img_state,687                "states": text_embeddings,688                "sigma": sigma[0],689                "weight_func": weight_func,690            }691 692            noise_pred = self.k_diffusion_model(693                latent_model_input, sigma, cond=encoder_state694            )695            noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)696            noise_pred = noise_pred_uncond + guidance_scale * (697                noise_pred_text - noise_pred_uncond698            )699            return noise_pred700 701        extra_args = self.get_sampler_extra_args_t2i(702            sigmas, eta, num_inference_steps, sampler703        )704        latents = sampler(model_fn, latents, **extra_args)705 706        if upscale:707            target_height = height * upscale_x708            target_width = width * upscale_x709            vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1)710            latents = torch.nn.functional.interpolate(711                latents,712                size=(713                    int(target_height // vae_scale_factor),714                    int(target_width // vae_scale_factor),715                ),716                mode=upscale_method,717                antialias=upscale_antialias,718            )719            return self.img2img(720                prompt=prompt,721                num_inference_steps=num_inference_steps,722                guidance_scale=guidance_scale,723                negative_prompt=negative_prompt,724                generator=generator,725                latents=latents,726                strength=upscale_denoising_strength,727                sampler_name=sampler_name,728                sampler_opt=sampler_opt,729                pww_state=None,730                pww_attn_weight=pww_attn_weight / 2,731            )732 733        # 8. Post-processing734        image = self.decode_latents(latents)735 736        # 10. Convert to PIL737        if output_type == "pil":738            image = self.numpy_to_pil(image)739 740        return (image,)741 742 743class FlashAttentionFunction(Function):744    @staticmethod745    @torch.no_grad()746    def forward(ctx, q, k, v, mask, causal, q_bucket_size, k_bucket_size):747        """Algorithm 2 in the paper"""748 749        device = q.device750        max_neg_value = -torch.finfo(q.dtype).max751        qk_len_diff = max(k.shape[-2] - q.shape[-2], 0)752 753        o = torch.zeros_like(q)754        all_row_sums = torch.zeros((*q.shape[:-1], 1), device=device)755        all_row_maxes = torch.full((*q.shape[:-1], 1), max_neg_value, device=device)756 757        scale = q.shape[-1] ** -0.5758 759        if not exists(mask):760            mask = (None,) * math.ceil(q.shape[-2] / q_bucket_size)761        else:762            mask = rearrange(mask, "b n -> b 1 1 n")763            mask = mask.split(q_bucket_size, dim=-1)764 765        row_splits = zip(766            q.split(q_bucket_size, dim=-2),767            o.split(q_bucket_size, dim=-2),768            mask,769            all_row_sums.split(q_bucket_size, dim=-2),770            all_row_maxes.split(q_bucket_size, dim=-2),771        )772 773        for ind, (qc, oc, row_mask, row_sums, row_maxes) in enumerate(row_splits):774            q_start_index = ind * q_bucket_size - qk_len_diff775 776            col_splits = zip(777                k.split(k_bucket_size, dim=-2),778                v.split(k_bucket_size, dim=-2),779            )780 781            for k_ind, (kc, vc) in enumerate(col_splits):782                k_start_index = k_ind * k_bucket_size783 784                attn_weights = einsum("... i d, ... j d -> ... i j", qc, kc) * scale785 786                if exists(row_mask):787                    attn_weights.masked_fill_(~row_mask, max_neg_value)788 789                if causal and q_start_index < (k_start_index + k_bucket_size - 1):790                    causal_mask = torch.ones(791                        (qc.shape[-2], kc.shape[-2]), dtype=torch.bool, device=device792                    ).triu(q_start_index - k_start_index + 1)793                    attn_weights.masked_fill_(causal_mask, max_neg_value)794 795                block_row_maxes = attn_weights.amax(dim=-1, keepdims=True)796                attn_weights -= block_row_maxes797                exp_weights = torch.exp(attn_weights)798 799                if exists(row_mask):800                    exp_weights.masked_fill_(~row_mask, 0.0)801 802                block_row_sums = exp_weights.sum(dim=-1, keepdims=True).clamp(803                    min=EPSILON804                )805 806                new_row_maxes = torch.maximum(block_row_maxes, row_maxes)807 808                exp_values = einsum("... i j, ... j d -> ... i d", exp_weights, vc)809 810                exp_row_max_diff = torch.exp(row_maxes - new_row_maxes)811                exp_block_row_max_diff = torch.exp(block_row_maxes - new_row_maxes)812 813                new_row_sums = (814                    exp_row_max_diff * row_sums815                    + exp_block_row_max_diff * block_row_sums816                )817 818                oc.mul_((row_sums / new_row_sums) * exp_row_max_diff).add_(819                    (exp_block_row_max_diff / new_row_sums) * exp_values820                )821 822                row_maxes.copy_(new_row_maxes)823                row_sums.copy_(new_row_sums)824 825        lse = all_row_sums.log() + all_row_maxes826 827        ctx.args = (causal, scale, mask, q_bucket_size, k_bucket_size)828        ctx.save_for_backward(q, k, v, o, lse)829 830        return o831 832    @staticmethod833    @torch.no_grad()834    def backward(ctx, do):835        """Algorithm 4 in the paper"""836 837        causal, scale, mask, q_bucket_size, k_bucket_size = ctx.args838        q, k, v, o, lse = ctx.saved_tensors839 840        device = q.device841 842        max_neg_value = -torch.finfo(q.dtype).max843        qk_len_diff = max(k.shape[-2] - q.shape[-2], 0)844 845        dq = torch.zeros_like(q)846        dk = torch.zeros_like(k)847        dv = torch.zeros_like(v)848 849        row_splits = zip(850            q.split(q_bucket_size, dim=-2),851            o.split(q_bucket_size, dim=-2),852            do.split(q_bucket_size, dim=-2),853            mask,854            lse.split(q_bucket_size, dim=-2),855            dq.split(q_bucket_size, dim=-2),856        )857 858        for ind, (qc, oc, doc, row_mask, lsec, dqc) in enumerate(row_splits):859            q_start_index = ind * q_bucket_size - qk_len_diff860 861            col_splits = zip(862                k.split(k_bucket_size, dim=-2),863                v.split(k_bucket_size, dim=-2),864                dk.split(k_bucket_size, dim=-2),865                dv.split(k_bucket_size, dim=-2),866            )867 868            for k_ind, (kc, vc, dkc, dvc) in enumerate(col_splits):869                k_start_index = k_ind * k_bucket_size870 871                attn_weights = einsum("... i d, ... j d -> ... i j", qc, kc) * scale872 873                if causal and q_start_index < (k_start_index + k_bucket_size - 1):874                    causal_mask = torch.ones(875                        (qc.shape[-2], kc.shape[-2]), dtype=torch.bool, device=device876                    ).triu(q_start_index - k_start_index + 1)877                    attn_weights.masked_fill_(causal_mask, max_neg_value)878 879                p = torch.exp(attn_weights - lsec)880 881                if exists(row_mask):882                    p.masked_fill_(~row_mask, 0.0)883 884                dv_chunk = einsum("... i j, ... i d -> ... j d", p, doc)885                dp = einsum("... i d, ... j d -> ... i j", doc, vc)886 887                D = (doc * oc).sum(dim=-1, keepdims=True)888                ds = p * scale * (dp - D)889 890                dq_chunk = einsum("... i j, ... j d -> ... i d", ds, kc)891                dk_chunk = einsum("... i j, ... i d -> ... j d", ds, qc)892 893                dqc.add_(dq_chunk)894                dkc.add_(dk_chunk)895                dvc.add_(dv_chunk)896 897        return dq, dk, dv, None, None, None, None898