From e0a15a657cfbce50f3b41ceef2747b7d2557d90d Mon Sep 17 00:00:00 2001 From: Michael Poutre Date: Wed, 23 Aug 2023 00:29:44 -0700 Subject: [PATCH] Initial working commit --- lib/actions/arith.py | 51 ++++++++++++++++++++++ lib/actions/base.py | 27 +++++++++++- lib/actions/nudge.py | 43 +++++++++++++++++++ lib/clip_model.py | 94 +++++++++++++++++++++++----------------- lib/fun_clip_stuff.py | 99 +++++++++++++++++++++++++++---------------- lib/tokenizer.py | 72 +++++++++++++++---------------- 6 files changed, 272 insertions(+), 114 deletions(-) diff --git a/lib/actions/arith.py b/lib/actions/arith.py index cb813e5..2bcfd1a 100644 --- a/lib/actions/arith.py +++ b/lib/actions/arith.py @@ -1,5 +1,8 @@ from typing import Callable, Union +import torch +from torch.nn import Embedding + from comfy.sd1_clip import SD1Tokenizer from custom_nodes.ClipStuff.lib.actions.base import Action, PromptSegment @@ -37,6 +40,54 @@ class ArithAction(Action): out += "\n" + "\t" * (depth - 1) + ")" return out + def token_length(self): + # ArithAction modifies the embeddings of the base segment, so the length is the length of the base segment + if isinstance(self.base_segment, Action): + return self.base_segment.token_length() + + return len(self.base_segment.tokens) + + + def get_all_segments(self): + segments = [] + if isinstance(self.base_segment, Action): + segments += self.base_segment.get_all_segments() + else: + segments.append(self.base_segment) + + for op_key, ops in self.ops.items(): + for op in ops: + if isinstance(op, Action): + segments += op.get_all_segments() + else: + segments.append(op) + + return segments + + def get_result(self, embedding_module: Embedding): + if isinstance(self.base_segment, Action): + base_segment_result = self.base_segment.get_result(embedding_module) + else: + base_segment_result = self.base_segment.get_embeddings(embedding_module) + + for op_key, ops in self.ops.items(): + for op in ops: + if isinstance(op, Action): + op_result = op.get_result(embedding_module) + else: + op_result = op.get_embeddings(embedding_module) + + + if op_result.shape[1] > base_segment_result.shape[1]: + print('[WARN] ArithAction: op_result.shape[1] > base_segment_result.shape[1] - averaging op_result') + op_result = torch.mean(op_result, dim=1, keepdim=True) + + if op_key == "+": + base_segment_result.add(op_result) + elif op_key == "-": + base_segment_result.subtract(op_result) + + return base_segment_result @classmethod def parse_segment( diff --git a/lib/actions/base.py b/lib/actions/base.py index fec7129..e497e66 100644 --- a/lib/actions/base.py +++ b/lib/actions/base.py @@ -1,7 +1,9 @@ from abc import ABC, abstractmethod from typing import Callable, Union +import torch from torch import Tensor +from torch.nn import Embedding from comfy.sd1_clip import SD1Tokenizer @@ -17,6 +19,18 @@ class Action(ABC): def END_CHAR(self): pass + @abstractmethod + def token_length(self): + pass + + @abstractmethod + def get_all_segments(self): + pass + + @abstractmethod + def get_result(self, embedding_module: Embedding): + pass + @classmethod @abstractmethod def parse_segment( @@ -37,6 +51,14 @@ class PromptSegment: self.text = text self.tokens = tokens + def token_length(self): + return len(self.tokens) + + def get_embeddings(self, embedding_module: Embedding): + tensors = torch.LongTensor(self.tokens).to(torch.device('cpu')) + unsqueezed_tensors = tensors.unsqueeze(0) + return embedding_module(unsqueezed_tensors) + def depth_repr(self, depth=1): out = f'"{self.text}"(' @@ -59,7 +81,10 @@ def build_prompt_segment(text: str, tokenizer: SD1Tokenizer) -> PromptSegment: if embedding is None: print(f"warning, embedding:{embedding_name} does not exist, ignoring") else: - tokens.append(embedding) + if len(embedding.shape) == 1: + tokens.append(embedding) + else: + tokens.extend(embedding) if leftover != "": word = leftover diff --git a/lib/actions/nudge.py b/lib/actions/nudge.py index c906ce2..dacd842 100644 --- a/lib/actions/nudge.py +++ b/lib/actions/nudge.py @@ -1,10 +1,53 @@ from typing import Optional, Union, Callable +import torch +from torch.nn import Embedding + from comfy.sd1_clip import SD1Tokenizer from custom_nodes.ClipStuff.lib.actions.base import Action, PromptSegment class NudgeAction(Action): + def token_length(self): + # Nudge nudges the embeddings of the base segment, so the length is the length of the base segment + if isinstance(self.base_segment, Action): + return self.base_segment.token_length() + + return len(self.base_segment.tokens) + + def get_all_segments(self): + segments = [] + if isinstance(self.base_segment, Action): + segments += self.base_segment.get_all_segments() + else: + segments.append(self.base_segment) + + if isinstance(self.target, Action): + segments += self.target.get_all_segments() + else: + segments.append(self.target) + + return segments + + def get_result(self, embedding_module: Embedding): + if isinstance(self.base_segment, Action): + base_segment_result = self.base_segment.get_result(embedding_module) + else: + base_segment_result = self.base_segment.get_embeddings(embedding_module) + + if isinstance(self.target, Action): + target_segment_result = self.target.get_result(embedding_module) + else: + target_segment_result = self.target.get_embeddings(embedding_module) + + base_mean = torch.mean(base_segment_result, dim=1, keepdim=True) + if target_segment_result.shape[1] == 1: + translation_vector = target_segment_result - base_mean + else: + translation_vector = torch.mean(target_segment_result, dim=1, keepdim=True) - base_mean + + return base_segment_result.add(translation_vector, alpha=self.weight) + START_CHAR = "[" END_CHAR = "]" diff --git a/lib/clip_model.py b/lib/clip_model.py index 407a129..95b241e 100644 --- a/lib/clip_model.py +++ b/lib/clip_model.py @@ -1,5 +1,6 @@ import contextlib import os +from typing import Union import torch from transformers import CLIPTextConfig, modeling_utils @@ -7,6 +8,7 @@ from transformers import CLIPTextConfig, modeling_utils from comfy import model_management import comfy.ops from comfy.sd import CLIP +from custom_nodes.ClipStuff.lib.actions.base import PromptSegment, Action from custom_nodes.ClipStuff.lib.fun_clip_stuff import MyCLIPTextModel from custom_nodes.ClipStuff.lib.tokenizer import TokenDict @@ -66,50 +68,69 @@ class SD1FunClipModel(torch.nn.Module): self.layer = self.layer_default[0] self.layer_idx = self.layer_default[1] - def set_up_textual_embeddings(self, tokens: list[list[tuple[TokenDict]]], current_embeds): - out_tokens = [] + def set_up_textual_embeddings(self, tokens: list[list[PromptSegment | Action]], current_embeds): next_new_token = token_dict_size = current_embeds.weight.shape[0] - 1 embedding_weights = [] + # For each batch for batch in tokens: - tokens_temp = [] - for tokenDict in batch: - y = tokenDict[0].token_id - if isinstance(y, int): - if y == token_dict_size: # EOS token - y = -1 - tokens_temp += [y] + for seg_or_action in batch: + if isinstance(seg_or_action, Action): + segments = seg_or_action.get_all_segments() else: - if y.shape[0] == current_embeds.weight.shape[1]: - embedding_weights += [y] - tokens_temp += [next_new_token] - next_new_token += 1 - else: - print("WARNING: shape mismatch when trying to apply embedding, embedding will be ignored", - y.shape[0], current_embeds.weight.shape[1]) - while len(tokens_temp) < len(batch): - tokens_temp += [self.empty_tokens[0][-1]] - out_tokens += [tokens_temp] + segments = [seg_or_action] + + for segment in segments: + tokens_temp = [] + segment_length = segment.token_length() + for tid_or_tensor in segment.tokens: + if isinstance(tid_or_tensor, int): + if tid_or_tensor == token_dict_size: # Is EOS token + tid_or_tensor = -1 # Set to -1 so that it can be replaced with the EOS token later + tokens_temp += [tid_or_tensor] + else: + if tid_or_tensor.shape[0] == current_embeds.weight.shape[1]: + embedding_weights += [tid_or_tensor] + tokens_temp += [next_new_token] + next_new_token += 1 + else: + print("WARNING: shape mismatch when trying to apply embedding, embedding will be ignored", + tid_or_tensor.shape[0], current_embeds.weight.shape[1]) + if len(tokens_temp) < segment_length: + # Pretty sure this is only needed if the embedding is not the same size as the CLIP embedding + print("WARNING: segment length mismatch, padding with EOS token") + tokens_temp.extend([self.empty_tokens[0][-1] * (segment_length - len(tokens_temp))]) + segment.tokens = tokens_temp n = token_dict_size if len(embedding_weights) > 0: + # Create new embedding, with size of current embedding + number of new embeddings new_embedding = torch.nn.Embedding(next_new_token + 1, current_embeds.weight.shape[1], device=current_embeds.weight.device, dtype=current_embeds.weight.dtype) + # Copy current embedding weights to new embedding new_embedding.weight[:token_dict_size] = current_embeds.weight[:-1] + # Add new embeddings for embed in embedding_weights: new_embedding.weight[n] = embed n += 1 + + # Set re-add the EOS token new_embedding.weight[n] = current_embeds.weight[-1] # EOS embedding self.transformer.set_input_embeddings(new_embedding) - for i, out_batch in enumerate(out_tokens): - for tokenIdx in range(len(out_batch)): - if out_batch[tokenIdx] == -1: - tokens[i][tokenIdx][0].token_id = n # The EOS token should always be the largest one - else: - tokens[i][tokenIdx][0].token_id = out_batch[tokenIdx] - # return processed_tokens + for batch in tokens: + for seg_or_action in batch: + if isinstance(seg_or_action, Action): + segments = seg_or_action.get_all_segments() + else: + segments = [seg_or_action] + + for segment in segments: + for tokenIdx in range(len(segment.tokens)): + if segment.tokens[tokenIdx] == -1: + segment.tokens[tokenIdx] = n + def forward(self, tokens, **kwargs): backup_embeds = self.transformer.get_input_embeddings() device = backup_embeds.weight.device @@ -153,15 +174,10 @@ class SD1FunClipModel(torch.nn.Module): def load_sd(self, sd): return self.transformer.load_state_dict(sd, strict=False) - def encode_token_weights(self, token_dicts: list[list[tuple[TokenDict]]], **kwargs): - to_encode = [list( - map( - lambda id: (TokenDict(token_id=id, weight=1.0, nudge_id=None),), - self.empty_tokens[0] - ) - )] - for x in token_dicts: - to_encode.append(x) + def encode_token_weights(self, prompt_segments: list[list[Union[PromptSegment | Action]]], **kwargs): + to_encode = [[PromptSegment(text="_Empty Batch_", tokens=self.empty_tokens[0])]] + for batch in prompt_segments: + to_encode.append(batch) out, pooled = self.encode(to_encode, **kwargs) z_empty = out[0:1] @@ -173,10 +189,10 @@ class SD1FunClipModel(torch.nn.Module): output = [] for k in range(1, out.shape[0]): z = out[k:k + 1] - for i in range(len(z)): - for j in range(len(z[i])): - weight = token_dicts[k - 1][j][0].weight - z[i][j] = (z[i][j] - z_empty[0][j]) * weight + z_empty[0][j] + # for i in range(len(z)): + # for j in range(len(z[i])): + # weight = token_dicts[k - 1][j][0].weight + # z[i][j] = (z[i][j] - z_empty[0][j]) * weight + z_empty[0][j] output.append(z) if (len(output) == 0): diff --git a/lib/fun_clip_stuff.py b/lib/fun_clip_stuff.py index 40743c4..1f2cee9 100644 --- a/lib/fun_clip_stuff.py +++ b/lib/fun_clip_stuff.py @@ -7,6 +7,7 @@ from transformers.modeling_outputs import BaseModelOutputWithPooling from transformers.models.clip.modeling_clip import _expand_mask, CLIPTextEmbeddings, CLIPTextTransformer, \ CLIPTextModel +from custom_nodes.ClipStuff.lib.actions.base import PromptSegment, Action from custom_nodes.ClipStuff.lib.tokenizer import TokenDict def slerp(val, low, high): @@ -30,47 +31,57 @@ class MyCLIPTextEmbeddings(CLIPTextEmbeddings): position_ids: Optional[torch.LongTensor] = None, inputs_embeds: Optional[torch.FloatTensor] = None, ) -> torch.Tensor: - input_ids = [ - [ - tokenDict[0].token_id for tokenDict in batch - ] for batch in input_dicts - ] - tokens = torch.LongTensor(input_ids).to(torch.device('cpu')) - input_shape = tokens.size() - input_ids = tokens.view(-1, input_shape[-1]) - seq_length = input_ids.shape[-1] if input_ids is not None else inputs_embeds.shape[-2] + + batches = [] + for batch_idx, batch in enumerate(input_dicts): + results = [] + for seg_or_action in batch: + if isinstance(seg_or_action, Action): + results.append(seg_or_action.get_result(self.token_embedding)) + else: + results.append(seg_or_action.get_embeddings(self.token_embedding)) + batches.append(results) + + seq_length = batches[0][0].shape[-2] if position_ids is None: position_ids = self.position_ids[:, :seq_length] - if inputs_embeds is None: - inputs_embeds = self.token_embedding(input_ids) + # if inputs_embeds is None: + # inputs_embeds = self.token_embedding(input_ids) - for batch_idx, batch in enumerate(input_dicts): - for token_idx, token in enumerate(batch): - if token[0].nudge_id is not None: - nudged_embed = inputs_embeds[batch_idx, token_idx][:] + self.token_embedding(torch.LongTensor([token[0].nudge_id]).to(torch.device('cpu')))[0] - if token[0].nudge_index_start is not None and token[0].nudge_index_stop is not None: - nudge_start = token[0].nudge_index_start - nudge_end = token[0].nudge_index_stop - else: - nudge_start = 0 - nudge_end = 768 - inputs_embeds[batch_idx, token_idx][nudge_start:nudge_end] = (slerp(token[0].nudge_weight, inputs_embeds[batch_idx, token_idx][:], nudged_embed)[0][nudge_start:nudge_end]) - elif token[0].arith_ops is not None: - for op, id_list in token[0].arith_ops.items(): - if op == '+': - for this_id in id_list: - inputs_embeds[batch_idx, token_idx] += self.token_embedding(torch.LongTensor([this_id]).to(torch.device('cpu')))[0] - elif op == '-': - for this_id in id_list: - inputs_embeds[batch_idx, token_idx] -= self.token_embedding(torch.LongTensor([this_id]).to(torch.device('cpu')))[0] + # for batch_idx, batch in enumerate(input_dicts): + # for token_idx, token in enumerate(batch): + # if token[0].nudge_id is not None: + # nudged_embed = inputs_embeds[batch_idx, token_idx][:] + self.token_embedding(torch.LongTensor([token[0].nudge_id]).to(torch.device('cpu')))[0] + # if token[0].nudge_index_start is not None and token[0].nudge_index_stop is not None: + # nudge_start = token[0].nudge_index_start + # nudge_end = token[0].nudge_index_stop + # else: + # nudge_start = 0 + # nudge_end = 768 + # inputs_embeds[batch_idx, token_idx][nudge_start:nudge_end] = (slerp(token[0].nudge_weight, inputs_embeds[batch_idx, token_idx][:], nudged_embed)[0][nudge_start:nudge_end]) + # elif token[0].arith_ops is not None: + # for op, id_list in token[0].arith_ops.items(): + # if op == '+': + # for this_id in id_list: + # inputs_embeds[batch_idx, token_idx] += self.token_embedding(torch.LongTensor([this_id]).to(torch.device('cpu')))[0] + # elif op == '-': + # for this_id in id_list: + # inputs_embeds[batch_idx, token_idx] -= self.token_embedding(torch.LongTensor([this_id]).to(torch.device('cpu')))[0] + + embeds = [] + for batch in batches: + if len(batch) == 1: + embeds.append(batch[0]) + else: + embeds.append(torch.cat(batch, dim=-2)) position_embeddings = self.position_embedding(position_ids) - embeddings = inputs_embeds + position_embeddings + embeddings = torch.cat(embeds, dim=0) + position_embeddings - return embeddings, input_ids, input_shape + return embeddings class MyCLIPTextTransformer(CLIPTextTransformer): @@ -80,7 +91,7 @@ class MyCLIPTextTransformer(CLIPTextTransformer): def forward( self, - input_ids: Optional[list[list[tuple[TokenDict]]]] = None, + input_ids: Optional[list[list[PromptSegment | Action]]] = None, attention_mask: Optional[torch.Tensor] = None, position_ids: Optional[torch.Tensor] = None, output_attentions: Optional[bool] = None, @@ -103,9 +114,12 @@ class MyCLIPTextTransformer(CLIPTextTransformer): # input_shape = input_ids.size() # input_ids = input_ids.view(-1, input_shape[-1]) - hidden_states, input_ids, input_shape = self.embeddings(input_dicts=input_ids) + hidden_states = self.embeddings(input_dicts=input_ids) - bsz, seq_len = input_shape + bsz = len(input_ids) + # TODO: Properly gather this + seq_len = 77 + # bsz, seq_len = input_shape # CLIP's text model uses causal mask, prepare it here. # https://github.com/openai/CLIP/blob/cfcffb90e69f37bf2ff1e988237a0fbe41f33c04/clip/model.py#L324 causal_attention_mask = self._build_causal_attention_mask(bsz, seq_len, hidden_states.dtype).to( @@ -128,12 +142,25 @@ class MyCLIPTextTransformer(CLIPTextTransformer): last_hidden_state = encoder_outputs[0] last_hidden_state = self.final_layer_norm(last_hidden_state) + + # Hacky way to get idx of first EOT token + eot_idx = [1] + for batch in input_ids[1:]: + idx = 0 + for seg_or_action in batch: + if isinstance(seg_or_action, Action): + idx += seg_or_action.token_length() + else: + if seg_or_action.text == '__PAD__': + break + eot_idx.append(idx) # text_embeds.shape = [batch_size, sequence_length, transformer.width] # take features from the eot embedding (eot_token is the highest number in each sequence) # casting to torch.int for onnx compatibility: argmax doesn't support int64 inputs with opset 14 + # TODO: Get the index of the first EOT token pooled_output = last_hidden_state[ torch.arange(last_hidden_state.shape[0], device=last_hidden_state.device), - input_ids.to(dtype=torch.int, device=last_hidden_state.device).argmax(dim=-1), + eot_idx ] if not return_dict: diff --git a/lib/tokenizer.py b/lib/tokenizer.py index ca3e9e8..c46d659 100644 --- a/lib/tokenizer.py +++ b/lib/tokenizer.py @@ -18,7 +18,7 @@ from custom_nodes.ClipStuff.lib.actions.lib import ( is_any_action_segment, is_action_segment, ) - +from custom_nodes.ClipStuff.lib.actions.utils import batch_size_info arith_action = r'(<[a-zA-Z0-9\-_]+:[a-zA-Z0-9\-_]+>)' @@ -115,9 +115,11 @@ class MyTokenizer(SD1Tokenizer): super().__init__(tokenizer_path, max_length, pad_with_end, embedding_directory, embedding_size, embedding_key) """ - :return: list of tuples (tokenDict, word_id?) + Doesn't actually tokenize... + Returns batches of segments and actions + :return: List of list(batches) of segments and actions """ - def tokenize_with_weights(self, text:str, return_word_ids=False, **kwargs) -> list[list[Action | int]]: + def tokenize_with_weights(self, text:str, return_word_ids=False, **kwargs) -> list[list[PromptSegment | Action]]: if self.pad_with_end: pad_token = self.end_token else: @@ -139,44 +141,38 @@ class MyTokenizer(SD1Tokenizer): else: print(segment.depth_repr()) - tokens: list[list[Action | int ]] = [] - # reshape token array to CLIP input size - batched_tokens = [] - batch = [(TokenDict(token_id=self.start_token), 0)] - batched_tokens.append(batch) - for i, t_group in enumerate(tokens): + batched_segments = [] + batch = [PromptSegment(text="[SOT]", tokens=[self.start_token])] + # batched_segments.append(batch) + batch_size = 1 + for segment in parsed_actions: + num_tokens = segment.token_length() # determine if we're going to try and keep the tokens in a single batch - is_large = len(t_group) >= self.max_word_length + is_large = num_tokens >= self.max_word_length - while len(t_group) > 0: - if len(t_group) + len(batch) > self.max_length - 1: - remaining_length = self.max_length - len(batch) - 1 - # break word in two and add end token - if is_large: - batch.extend([(tokenDict, i+1) for tokenDict in t_group[:remaining_length]]) - batch.append((TokenDict(token_id=self.end_token), 0)) - t_group = t_group[remaining_length:] - # add end token and pad - else: - batch.append((TokenDict(token_id=self.end_token), 0)) - batch.extend([(TokenDict(token_id=pad_token), 0)] * remaining_length) - # start new batch - batch = [(TokenDict(token_id=self.start_token), 1.0, 0)] - batched_tokens.append(batch) - else: - batch.extend([(tokenDict, i+1) for tokenDict in t_group]) - t_group = [] + # If the segment is too large to fit in a single batch, pad the current batch and start a new one + if num_tokens + batch_size > self.max_length - 1: + remaining_length = self.max_length - batch_size - 1 # -1 for end token + # Pad batch + batch.append(PromptSegment("__PAD__", [self.end_token] + [pad_token] * remaining_length - 1)) + batched_segments.append(batch) - # fill last batch - batch.extend([(TokenDict(token_id=self.end_token), 0)] + [ - (TokenDict(token_id=pad_token), 0)] * (self.max_length - len(batch) - 1)) + # start new batch + batch = [PromptSegment(text="[SOT]", tokens=[self.start_token]), segment] + batch_size = num_tokens + 1 # +1 for start token + continue - if not return_word_ids: - batched_tokens = [ - [ - (tokenInfo[0],) for tokenInfo in batch - ] for batch in batched_tokens - ] + # If the segment is small enough to fit in the current batch, add it + batch.append(segment) + batch_size += num_tokens - return batched_tokens + # Pad the last batch + remaining_length = self.max_length - batch_size - 1 # -1 for end token + batch.append(PromptSegment("__PAD__", [self.end_token] + [pad_token] * remaining_length)) + batched_segments.append(batch) + + for batch in batched_segments: + batch_size_info(batch) + + return batched_segments