memef4rmer/edit_anything
0
1import torch2import einops3 4import ldm.modules.encoders.modules5import ldm.modules.attention6 7from transformers import logging8from ldm.modules.attention import default9 10 11def disable_verbosity():12 logging.set_verbosity_error()13 print('logging improved.')14 return15 16 17def enable_sliced_attention():18 ldm.modules.attention.CrossAttention.forward = _hacked_sliced_attentin_forward19 print('Enabled sliced_attention.')20 return21 22 23def hack_everything(clip_skip=0):24 disable_verbosity()25 ldm.modules.encoders.modules.FrozenCLIPEmbedder.forward = _hacked_clip_forward26 ldm.modules.encoders.modules.FrozenCLIPEmbedder.clip_skip = clip_skip27 print('Enabled clip hacks.')28 return29 30 31# Written by Lvmin32def _hacked_clip_forward(self, text):33 PAD = self.tokenizer.pad_token_id34 EOS = self.tokenizer.eos_token_id35 BOS = self.tokenizer.bos_token_id36 37 def tokenize(t):38 return self.tokenizer(t, truncation=False, add_special_tokens=False)["input_ids"]39 40 def transformer_encode(t):41 if self.clip_skip > 1:42 rt = self.transformer(input_ids=t, output_hidden_states=True)43 return self.transformer.text_model.final_layer_norm(rt.hidden_states[-self.clip_skip])44 else:45 return self.transformer(input_ids=t, output_hidden_states=False).last_hidden_state46 47 def split(x):48 return x[75 * 0: 75 * 1], x[75 * 1: 75 * 2], x[75 * 2: 75 * 3]49 50 def pad(x, p, i):51 return x[:i] if len(x) >= i else x + [p] * (i - len(x))52 53 raw_tokens_list = tokenize(text)54 tokens_list = []55 56 for raw_tokens in raw_tokens_list:57 raw_tokens_123 = split(raw_tokens)58 raw_tokens_123 = [[BOS] + raw_tokens_i + [EOS] for raw_tokens_i in raw_tokens_123]59 raw_tokens_123 = [pad(raw_tokens_i, PAD, 77) for raw_tokens_i in raw_tokens_123]60 tokens_list.append(raw_tokens_123)61 62 tokens_list = torch.IntTensor(tokens_list).to(self.device)63 64 feed = einops.rearrange(tokens_list, 'b f i -> (b f) i')65 y = transformer_encode(feed)66 z = einops.rearrange(y, '(b f) i c -> b (f i) c', f=3)67 68 return z69 70 71# Stolen from https://github.com/basujindal/stable-diffusion/blob/main/optimizedSD/splitAttention.py72def _hacked_sliced_attentin_forward(self, x, context=None, mask=None):73 h = self.heads74 75 q = self.to_q(x)76 context = default(context, x)77 k = self.to_k(context)78 v = self.to_v(context)79 del context, x80 81 q, k, v = map(lambda t: einops.rearrange(t, 'b n (h d) -> (b h) n d', h=h), (q, k, v))82 83 limit = k.shape[0]84 att_step = 185 q_chunks = list(torch.tensor_split(q, limit // att_step, dim=0))86 k_chunks = list(torch.tensor_split(k, limit // att_step, dim=0))87 v_chunks = list(torch.tensor_split(v, limit // att_step, dim=0))88 89 q_chunks.reverse()90 k_chunks.reverse()91 v_chunks.reverse()92 sim = torch.zeros(q.shape[0], q.shape[1], v.shape[2], device=q.device)93 del k, q, v94 for i in range(0, limit, att_step):95 q_buffer = q_chunks.pop()96 k_buffer = k_chunks.pop()97 v_buffer = v_chunks.pop()98 sim_buffer = torch.einsum('b i d, b j d -> b i j', q_buffer, k_buffer) * self.scale99 100 del k_buffer, q_buffer101 # attention, what we cannot get enough of, by chunks102 103 sim_buffer = sim_buffer.softmax(dim=-1)104 105 sim_buffer = torch.einsum('b i j, b j d -> b i d', sim_buffer, v_buffer)106 del v_buffer107 sim[i:i + att_step, :, :] = sim_buffer108 109 del sim_buffer110 sim = einops.rearrange(sim, '(b h) n d -> b n (h d)', h=h)111 return self.to_out(sim)112 