Covert1107/sd-diffusers-webui
2
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 