CoolFace
Apppublic

memef4rmer/edit_anything

sourceHugging Faceccupdated 3y agoView on Hugging Face
0likes
hack.py112 linesDownload Raw Back to cldm
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