diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..1b05740 --- /dev/null +++ b/.gitignore @@ -0,0 +1,2 @@ + +__pycache__/ diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..95f7034 --- /dev/null +++ b/__init__.py @@ -0,0 +1,10 @@ +from .nodes import ClipInjectedCheckpointLoader, FunCLIPTextEncode, BuildGif + +NODE_CLASS_MAPPINGS = { + "ClipInjectedCheckpointLoader": ClipInjectedCheckpointLoader, + "FunCLIPTextEncode": FunCLIPTextEncode, + "Build Gif": BuildGif +} +# +# EXTENSION_NAME = "ComfyLiterals" +# symlink_web_dir("js", EXTENSION_NAME) diff --git a/lib/__init__.py b/lib/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/lib/clip_config.json b/lib/clip_config.json new file mode 100644 index 0000000..0158a1f --- /dev/null +++ b/lib/clip_config.json @@ -0,0 +1,25 @@ +{ + "_name_or_path": "openai/clip-vit-large-patch14", + "architectures": [ + "CLIPTextModel" + ], + "attention_dropout": 0.0, + "bos_token_id": 0, + "dropout": 0.0, + "eos_token_id": 2, + "hidden_act": "quick_gelu", + "hidden_size": 768, + "initializer_factor": 1.0, + "initializer_range": 0.02, + "intermediate_size": 3072, + "layer_norm_eps": 1e-05, + "max_position_embeddings": 77, + "model_type": "clip_text_model", + "num_attention_heads": 12, + "num_hidden_layers": 12, + "pad_token_id": 1, + "projection_dim": 768, + "torch_dtype": "float32", + "transformers_version": "4.24.0", + "vocab_size": 49408 +} diff --git a/lib/clip_model.py b/lib/clip_model.py new file mode 100644 index 0000000..efe48b5 --- /dev/null +++ b/lib/clip_model.py @@ -0,0 +1,203 @@ +import contextlib +import os + +import torch +from transformers import CLIPTextConfig, modeling_utils + +from comfy import model_management +import comfy.ops +from comfy.sd import CLIP +from custom_nodes.ClipStuff.lib.fun_clip_stuff import MyCLIPTextModel +from custom_nodes.ClipStuff.lib.tokenizer import TokenDict + + +class FunCLIP(CLIP): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + + def encode_from_tokens(self, tokens, return_pooled=False, position_ids=None): + if self.layer_idx is not None: + self.cond_stage_model.clip_layer(self.layer_idx) + else: + self.cond_stage_model.reset_clip_layer() + model_management.load_model_gpu(self.patcher) + cond, pooled = self.cond_stage_model.encode_token_weights(tokens, position_ids=position_ids) + if return_pooled: + return cond, pooled + return cond + + +class FunClipTokenWeightEncoder: + def encode_token_weights(self, token_dicts: list[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) + + out, pooled = self.encode(to_encode, **kwargs) + z_empty = out[0:1] + if pooled.shape[0] > 1: + first_pooled = pooled[1:2] + else: + first_pooled = pooled[0:1] + + 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] + output.append(z) + + if (len(output) == 0): + return z_empty.cpu(), first_pooled.cpu() + return torch.cat(output, dim=-2).cpu(), first_pooled.cpu() + + +class SD1FunClipModel(torch.nn.Module, FunClipTokenWeightEncoder): + """Uses the CLIP transformer encoder for text (from huggingface)""" + LAYERS = [ + "last", + "pooled", + "hidden" + ] + + def __init__(self, version="openai/clip-vit-large-patch14", device="cpu", max_length=77, + freeze=True, layer="last", layer_idx=None, textmodel_json_config=None, + textmodel_path=None): # clip-vit-base-patch32 + super().__init__() + assert layer in self.LAYERS + self.num_layers = 12 + if textmodel_path is not None: + self.transformer = MyCLIPTextModel.from_pretrained(textmodel_path) + else: + if textmodel_json_config is None: + textmodel_json_config = os.path.join(os.path.dirname(os.path.realpath(__file__)), "clip_config.json") + config = CLIPTextConfig.from_json_file(textmodel_json_config) + self.num_layers = config.num_hidden_layers + with comfy.ops.use_comfy_ops(): + with modeling_utils.no_init_weights(): + self.transformer = MyCLIPTextModel(config) + + self.max_length = max_length + if freeze: + self.freeze() + self.layer = layer + self.layer_idx = None + self.empty_tokens = [[49406] + [49407] * 76] + self.text_projection = None + self.layer_norm_hidden_state = True + if layer == "hidden": + assert layer_idx is not None + assert abs(layer_idx) <= self.num_layers + self.clip_layer(layer_idx) + self.layer_default = (self.layer, self.layer_idx) + + def freeze(self): + self.transformer = self.transformer.eval() + # self.train = disabled_train + for param in self.parameters(): + param.requires_grad = False + + def clip_layer(self, layer_idx): + if abs(layer_idx) >= self.num_layers: + self.layer = "last" + else: + self.layer = "hidden" + self.layer_idx = layer_idx + + def reset_clip_layer(self): + self.layer = self.layer_default[0] + self.layer_idx = self.layer_default[1] + + def set_up_textual_embeddings(self, tokens, current_embeds): + out_tokens = [] + next_new_token = token_dict_size = current_embeds.weight.shape[0] - 1 + embedding_weights = [] + + 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] + 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] + + n = token_dict_size + if len(embedding_weights) > 0: + new_embedding = torch.nn.Embedding(next_new_token + 1, current_embeds.weight.shape[1], + device=current_embeds.weight.device, dtype=current_embeds.weight.dtype) + new_embedding.weight[:token_dict_size] = current_embeds.weight[:-1] + for embed in embedding_weights: + new_embedding.weight[n] = embed + n += 1 + 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 + def forward(self, tokens, **kwargs): + backup_embeds = self.transformer.get_input_embeddings() + device = backup_embeds.weight.device + self.set_up_textual_embeddings(tokens, backup_embeds) + # tokens = torch.LongTensor(tokens).to(device) + + if backup_embeds.weight.dtype != torch.float32: + precision_scope = torch.autocast + else: + precision_scope = contextlib.nullcontext + + + if (kwargs.get("position_ids", None) is not None): + position_ids = torch.LongTensor(kwargs["position_ids"]).to(device) + else: + position_ids = None + + + with precision_scope(model_management.get_autocast_device(device)): + outputs = self.transformer(input_ids=tokens, output_hidden_states=self.layer == "hidden", + position_ids=position_ids) + self.transformer.set_input_embeddings(backup_embeds) + + if self.layer == "last": + z = outputs.last_hidden_state + elif self.layer == "pooled": + z = outputs.pooler_output[:, None, :] + else: + z = outputs.hidden_states[self.layer_idx] + if self.layer_norm_hidden_state: + z = self.transformer.text_model.final_layer_norm(z) + + pooled_output = outputs.pooler_output + if self.text_projection is not None: + pooled_output = pooled_output.to(self.text_projection.device) @ self.text_projection + return z.float(), pooled_output.float() + + def encode(self, tokens, **kwargs): + return self(tokens, **kwargs) + + def load_sd(self, sd): + return self.transformer.load_state_dict(sd, strict=False) diff --git a/lib/fun_clip_stuff.py b/lib/fun_clip_stuff.py new file mode 100644 index 0000000..6e38f5d --- /dev/null +++ b/lib/fun_clip_stuff.py @@ -0,0 +1,182 @@ +from typing import Optional, Tuple, Union + +import torch +from torch import device +from transformers import CLIPTextConfig +from transformers.modeling_outputs import BaseModelOutputWithPooling +from transformers.models.clip.modeling_clip import _expand_mask, CLIPTextEmbeddings, CLIPTextTransformer, \ + CLIPTextModel + +from custom_nodes.ClipStuff.lib.tokenizer import TokenDict + +def slerp(val, low, high): + low = low.unsqueeze(0) + high = high.unsqueeze(0) + low_norm = low/torch.norm(low, dim=1, keepdim=True) + high_norm = high/torch.norm(high, dim=1, keepdim=True) + omega = torch.acos((low_norm*high_norm).sum(1)) + so = torch.sin(omega) + res = (torch.sin((1.0-val)*omega)/so).unsqueeze(1)*low + (torch.sin(val*omega)/so).unsqueeze(1) * high + return res + +class MyCLIPTextEmbeddings(CLIPTextEmbeddings): + def __init__(self, config: CLIPTextConfig): + super().__init__(config) + + def forward( + self, + input_dicts: Optional[list[list[tuple[TokenDict]]]] = None, + input_ids: Optional[torch.LongTensor] = None, + 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] + + if position_ids is None: + position_ids = self.position_ids[:, :seq_length] + + 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_indx_start'] is not None and token[0]['nudge_index_stop'] is not None: + nudge_start = int(token[0]['nudge_indx_start']) + nudge_end = int(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]) + + position_embeddings = self.position_embedding(position_ids) + embeddings = inputs_embeds + position_embeddings + + return embeddings, input_ids, input_shape + + +class MyCLIPTextTransformer(CLIPTextTransformer): + def __init__(self, config: CLIPTextConfig): + super().__init__(config) + self.embeddings = MyCLIPTextEmbeddings(config) + + def forward( + self, + input_ids: Optional[list[list[tuple[TokenDict]]]] = None, + attention_mask: Optional[torch.Tensor] = None, + position_ids: Optional[torch.Tensor] = None, + output_attentions: Optional[bool] = None, + output_hidden_states: Optional[bool] = None, + return_dict: Optional[bool] = None, + ) -> Union[Tuple, BaseModelOutputWithPooling]: + r""" + Returns: + + """ + output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions + output_hidden_states = ( + output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states + ) + return_dict = return_dict if return_dict is not None else self.config.use_return_dict + + if input_ids is None: + raise ValueError("You have to specify input_ids") + + # 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) + + 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( + hidden_states.device + ) + # expand attention_mask + if attention_mask is not None: + # [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len] + attention_mask = _expand_mask(attention_mask, hidden_states.dtype) + + encoder_outputs = self.encoder( + inputs_embeds=hidden_states, + attention_mask=attention_mask, + causal_attention_mask=causal_attention_mask, + output_attentions=output_attentions, + output_hidden_states=output_hidden_states, + return_dict=return_dict, + ) + + last_hidden_state = encoder_outputs[0] + last_hidden_state = self.final_layer_norm(last_hidden_state) + + # 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 + 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), + ] + + if not return_dict: + return (last_hidden_state, pooled_output) + encoder_outputs[1:] + + return BaseModelOutputWithPooling( + last_hidden_state=last_hidden_state, + pooler_output=pooled_output, + hidden_states=encoder_outputs.hidden_states, + attentions=encoder_outputs.attentions, + ) + + +class MyCLIPTextModel(CLIPTextModel): + def __init__(self, config: CLIPTextConfig): + super().__init__(config) + self.text_model = MyCLIPTextTransformer(config) + + def forward( + self, + input_ids: Optional[list[list[tuple[TokenDict]]]] = None, + attention_mask: Optional[torch.Tensor] = None, + position_ids: Optional[torch.Tensor] = None, + output_attentions: Optional[bool] = None, + output_hidden_states: Optional[bool] = None, + return_dict: Optional[bool] = None, + ) -> Union[Tuple, BaseModelOutputWithPooling]: + r""" + Returns: + + Examples: + + ```python + >>> from transformers import AutoTokenizer, CLIPTextModel + + >>> model = CLIPTextModel.from_pretrained("openai/clip-vit-base-patch32") + >>> tokenizer = AutoTokenizer.from_pretrained("openai/clip-vit-base-patch32") + + >>> inputs = tokenizer(["a photo of a cat", "a photo of a dog"], padding=True, return_tensors="pt") + + >>> outputs = model(**inputs) + >>> last_hidden_state = outputs.last_hidden_state + >>> pooled_output = outputs.pooler_output # pooled (EOS token) states + ```""" + return_dict = return_dict if return_dict is not None else self.config.use_return_dict + + return self.text_model( + input_ids=input_ids, + attention_mask=attention_mask, + position_ids=position_ids, + output_attentions=output_attentions, + output_hidden_states=output_hidden_states, + return_dict=return_dict, + ) diff --git a/lib/tokenizer.py b/lib/tokenizer.py new file mode 100644 index 0000000..ecf1a55 --- /dev/null +++ b/lib/tokenizer.py @@ -0,0 +1,191 @@ +from typing import Union, TypedDict, Optional + +from comfy.sd1_clip import SD1Tokenizer, escape_important, unescape_important, parse_parentheses + + +def token_weights(string, current_weight): + a = parse_parentheses(string) + out = [] + for x in a: + weight = current_weight + if len(x) >= 2 and x[-1] == ')' and x[0] == '(': + x = x[1:-1] + xx = x.rfind(":") + weight *= 1.1 + if xx > 0: + try: + weight = float(x[xx+1:]) + x = x[:xx] + except: + pass + out += token_weights(x, weight) + else: + out += [(x, current_weight)] + return out + +def parse_brackets(string): + out = [] + current = "" + for char in string: + if char == '[': + out += [current] + current = "[" + elif char == ']': + out += [current + ']'] + current = "" + else: + current += char + out += [current] + return out + +def parse_nudges(string) -> list[tuple[str, Union[str, None]]]: + out = [] + for nudge_segment in parse_brackets(string): + if nudge_segment == "": + continue + + if nudge_segment[0] != '[' and nudge_segment[-1] != ']': + out += [(nudge_segment, None, None)] + continue + + nudge_segment = nudge_segment[1:-1] + sep_idx = nudge_segment.find(":") + if sep_idx < 0: + out += [(nudge_segment, None, None)] + continue + + nudge_to = nudge_segment[sep_idx+1:] + weight = None + + weight_sep_idx = nudge_to.find(":") + if weight_sep_idx >= 0: + [nudge_to, weight] = nudge_to.split(":") + weight = float(weight) + + out += [(nudge_segment[:sep_idx], nudge_to, weight)] + return out + +# class TokenDict(TypedDict): +# token_id: int +# weight: float +# nudge_id: Optional[int] +# nudge_weight: Optional[float] +# + +class TokenDict: + def __init__(self, token_id: int, weight: float, nudge_id = None, nudge_weight = None, nudge_start: int = None, nudge_end: int = None): + self.token_id = token_id + self.weight = weight + self.nudge_id = nudge_id + self.nudge_weight = nudge_weight + self.nudge_start = nudge_start + self.nudge_end = nudge_end + + + +class MyTokenizer(SD1Tokenizer): + def __init__(self, tokenizer_path=None, max_length=77, pad_with_end=True, embedding_directory=None, embedding_size=768, embedding_key='clip_l', special_tokens=None): + super().__init__(tokenizer_path, max_length, pad_with_end, embedding_directory, embedding_size, embedding_key) + + """ + :return: list of tuples (tokenDict, word_id?) + """ + def tokenize_with_weights(self, text:str, return_word_ids=False, **kwargs): + if self.pad_with_end: + pad_token = self.end_token + else: + pad_token = 0 + + parsed_nudges = parse_nudges(text) + + nudge_start = None + nudge_end = None + if kwargs.get('nudge_start', None) is not None and kwargs.get('nudge_end', None) is not None: + nudge_start = int(kwargs.get('nudge_start')) + nudge_end = int(kwargs.get('nudge_end')) + + + #tokenize words + tokens: list[list[TokenDict]] = [] + weight = 1.0 + for token_segment, nudge_to_token, nudge_weight in parsed_nudges: + to_tokenize = token_segment.split(' ') + to_tokenize = [x for x in to_tokenize if x != ""] + # if token_segment == ' ': + # continue + + if nudge_weight is None: + nudge_weight = .5 + + nudge_to_id = None + if nudge_to_token is not None: + # self.convert_tokens_to_ids + nudge_to_id = self.tokenizer(nudge_to_token)["input_ids"][1:-1][0] + + for word in to_tokenize: + #if we find an embedding, deal with the embedding + if word.startswith(self.embedding_identifier) and self.embedding_directory is not None: + embedding_name = word[len(self.embedding_identifier):].strip('\n') + embed, leftover = self._try_get_embedding(embedding_name) + if embed is None: + print(f"warning, embedding:{embedding_name} does not exist, ignoring") + else: + if len(embed.shape) == 1: + tokens.append([TokenDict(token_id=embed, weight=weight, nudge_id=None, nudge_weight=None)]) + else: + tokens.append([ + TokenDict(token_id=embed[x], weight=weight, nudge_id=None, nudge_weight=nudge_weight) + for x in range(embed.shape[0]) + ]) + #if we accidentally have leftover text, continue parsing using leftover, else move on to next word + if leftover != "": + word = leftover + else: + continue + #parse word + tokens.append([TokenDict( + token_id=t, + weight=weight, + nudge_id=nudge_to_id, + nudge_weight=nudge_weight + ) for t in self.tokenizer(word)["input_ids"][1:-1]]) + + #reshape token array to CLIP input size + batched_tokens = [] + batch = [(TokenDict(token_id=self.start_token, weight=1.0, nudge_id=None, nudge_weight=None), 0)] + batched_tokens.append(batch) + for i, t_group in enumerate(tokens): + #determine if we're going to try and keep the tokens in a single batch + is_large = len(t_group) >= 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, weight=1.0, nudge_id=None, nudge_weight=None), 0)) + t_group = t_group[remaining_length:] + #add end token and pad + else: + batch.append((TokenDict(token_id=self.end_token, weight=1.0, nudge_id=None, nudge_weight=None), 0)) + batch.extend([(TokenDict(token_id=pad_token, weight=1.0, nudge_id=None, nudge_weight=None), 0)] * (remaining_length)) + #start new batch + batch = [(TokenDict(token_id=self.start_token, weight=1.0, nudge_id=None, nudge_weight=None), 1.0, 0)] + batched_tokens.append(batch) + else: + batch.extend([(tokenDict,i+1) for tokenDict in t_group]) + t_group = [] + + #fill last batch + batch.extend([(TokenDict(token_id=self.end_token, weight=1.0, nudge_id=None, nudge_weight=None), 0)] + [ + (TokenDict(token_id=pad_token, weight=1.0, nudge_id=None, nudge_weight=None), 0)] * (self.max_length - len(batch) - 1)) + + if not return_word_ids: + batched_tokens = [ + [ + (tokenInfo[0],) for tokenInfo in batch + ] for batch in batched_tokens + ] + + return batched_tokens diff --git a/nodes.py b/nodes.py new file mode 100644 index 0000000..0faa170 --- /dev/null +++ b/nodes.py @@ -0,0 +1,208 @@ +import random + +import numpy as np +import torch +from PIL import Image + +import folder_paths +import comfy.sd +from comfy import model_management +from custom_nodes.ClipStuff.lib.clip_model import SD1FunClipModel, FunCLIP +from custom_nodes.ClipStuff.lib.tokenizer import MyTokenizer + + +class ClipInjectedCheckpointLoader: + @classmethod + def INPUT_TYPES(s): + return {"required": {"config_name": (folder_paths.get_filename_list("configs"),), + "ckpt_name": (folder_paths.get_filename_list("checkpoints"),)}} + + RETURN_TYPES = ("MODEL", "CLIP", "VAE") + FUNCTION = "load_checkpoint" + + CATEGORY = "advanced/loaders" + + def load_checkpoint(self, config_name, ckpt_name, output_vae=True, output_clip=True): + config_path = folder_paths.get_full_path("configs", config_name) + ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name) + return comfy.sd.load_checkpoint(config_path, ckpt_path, output_vae=True, output_clip=True, + embedding_directory=folder_paths.get_folder_paths("embeddings"), + clip_model=SD1FunClipModel, + clip_class=FunCLIP, + clip_tokenizer=MyTokenizer) + + +class FunCLIPTextEncode: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "text": ("STRING", {"multiline": True}), "clip": ("CLIP",), + # "slerp_power": ("FLOAT", {"min": 0.0, "max": 1.0}), + }} + + RETURN_TYPES = ("CONDITIONING",) + FUNCTION = "encode" + OUTPUT_IS_LIST = (True,) + CATEGORY = "conditioning" + + def encode(self, clip, text): + ret = [] + for prompt in text.split("\n"): + if prompt.strip() == "": + continue + tokens = clip.tokenizer.tokenize_with_weights(text, return_word_ids=False,) + cond, pooled = clip.encode_from_tokens(tokens, return_pooled=True, position_ids=[0] * 77) + cond = [[cond, {"pooled_output": pooled}]] + ret.append(cond) + # if clip.layer_idx is not None: + # clip.cond_stage_model.clip_layer(clip.layer_idx) + # else: + # clip.cond_stage_model.reset_clip_layer() + # + # model_management.load_model_gpu(clip.patcher) + # position_ids = [0] * 77 + # cond, pooled = clip.cond_stage_model.encode_token_weights(tokens, position_ids=position_ids) + # # return cond, pooled + + # return ([[cond, {"pooled_output": pooled}]],) + return (ret,) + + +# +# def buildGif(processing_res, path='./outputs/gif/'): +# width = processing_res.width +# height = processing_res.height +# +# x_batch[0].save(f"{path}{speed}_{processing_res.seed}_batch_{y}.gif", save_all=True, append_images=x_batch[1:], +# optimize=False, duration=speed, loop=0) +# +# count_x = int(processing_res.images[0].width / width) +# count_y = int(processing_res.images[0].height / height) +# +# gif_interval = sharedObj.Config['gif_interval'] +# if gif_interval.find(",") > 0: +# speeds = list(map(int, gif_interval.split(","))) +# else: +# speeds = [int(gif_interval)] +# +# gif_axis = sharedObj.Config['gif_axis'] +# # ax1 = count_y if gif_axis == "X" else count_x +# # ax2 = count_x if gif_axis == "X" else count_y +# if gif_axis == "X": +# for y in range(0, count_y): +# x_batch = [] +# for x in range(0, count_x): +# bbox = (x * width, y * height, (x + 1) * width, (y + 1) * height) +# print(bbox) +# x_batch.append(processing_res.images[0].crop(bbox)) +# # images.save_image(processed.images[g], p.outpath_grids, "xyz_grid" +# # working_slice.show() +# if sharedConfig.Config.get('gif_boomerang', False): +# boomerang_in_place(x_batch) +# for speed in speeds: +# +# else: +# print("Y!") +# for x in range(0, count_x): +# y_batch = [] +# for y in range(0, count_y): +# bbox = (x * width, y * height, (x + 1) * width, (y + 1) * height) +# print(bbox) +# y_batch.append(processing_res.images[0].crop(bbox)) +# # images.save_image(processed.images[g], p.outpath_grids, "xyz_grid" +# # working_slice.show() +# if sharedConfig.Config.get('gif_boomerang', False): +# boomerang_in_place(y_batch) +# for speed in speeds: +# y_batch[0].save(f"{path}{speed}_{processing_res.seed}_batch_{x}.gif", save_all=True, append_images=y_batch[1:], optimize=False, duration=speed, loop=0) + + +def tensor2img(tensor_img): + i = 255. * tensor_img.cpu().numpy() + i_np_arr = np.clip(i, 0, 255, out=i).astype(np.uint8, copy=False) + return Image.fromarray(i_np_arr) + + +class BuildGif: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "images": ("IMAGE",), + }, + } + + RELOAD_INST = True + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("Gifs",) + INPUT_IS_LIST = True + FUNCTION = "build_gif" + OUTPUT_IS_LIST = (True,) + # OUTPUT_NODE = False + + CATEGORY = "List Stuff" + + def build_gif(self, images: list): + print('Build GIF called!') + print(f"{type(images)}") + + batch_size = images[0].size()[0] + cell_width = images[0].size()[1] + cell_height = images[0].size()[2] + # x_batch[0].save(f"{path}{speed}_{processing_res.seed}_batch_{y}.gif", save_all=True, append_images=x_batch[1:], + # optimize=False, duration=speed, loop=0) + + out = [] + + for batch_idx in range(batch_size): + save_path = f"{folder_paths.get_output_directory()}/-{batch_idx}-{random.randint(1, 100)}" + print(save_path) + # tensor2img(images[0][batch_idx]).save( + # f"{save_path}.gif", + # save_all=True, + # append_images=[ + # tensor2img(nested_batch[batch_idx]) for nested_batch in images[1:] + # ], + # optimize=False, + # duration=100, + # loop=0 + # ) + tensor2img(images[0][batch_idx]).save( + f"{save_path}.webp", + save_all=True, + append_images=[ + tensor2img(nested_batch[batch_idx]) for nested_batch in images[1:] + ], + optimize=False, + duration=250, + loop=0 + ) + + + # for idx, img_batch in enumerate(images): + # for batch_idx, img in enumerate(img_batch): + # print("Stuff") + # img = tensor2img(img) + # pil_img = np.array(img.convert("RGB")).astype(np.float32, copy=False) / 255 + # out.append(torch.from_numpy(pil_img).unsqueeze(0)) + + + # box = (x * cell_width + margin + row_label_size, y * cell_height + margin + column_label_size) + # print(f"Box: {box}") + # print(f"Image: {type(img)}") + # grid_image.paste(img, box) + # + # if y == 0: + # draw.text((box[0] + cell_width / 2, box[1] - column_label_size), str(Y_Labels[x]), fill='white', + # font=font) + # if x == 0: + # draw.text((box[0] - row_label_size, box[1] + cell_width / 2), str(X_Labels[y]), fill='white', font=font) + # + # np_grid_image = np.array(grid_image.convert("RGB")).astype(np.float32, copy=False) / 255 + # torch_image = torch.from_numpy(np_grid_image)[None,] + # np_grid_image = None + # print(f"GridImage Shape: {grid_image.}") + return (out,)