From 97f48d0029245fc9cf81dfd658e63cfc24735e4f Mon Sep 17 00:00:00 2001 From: Michael Poutre Date: Sun, 29 Oct 2023 17:32:55 -0700 Subject: [PATCH] refactor/fix: More fixes to align with clip changes in base --- lib/clip_model.py | 15 ++++++++++++--- lib/parser/transformer.py | 4 ++-- lib/parser/utils.py | 11 +++++------ lib/tokenizer.py | 38 ++++++++++++++++++++------------------ nodes.py | 11 ++++++----- 5 files changed, 45 insertions(+), 34 deletions(-) diff --git a/lib/clip_model.py b/lib/clip_model.py index fe080c1..3029cef 100644 --- a/lib/clip_model.py +++ b/lib/clip_model.py @@ -7,6 +7,7 @@ from transformers import CLIPTextConfig, modeling_utils from comfy import model_management import comfy.ops +from comfy.sd1_clip import SD1ClipModel from comfy.sdxl_clip import SDXLClipModel from custom_nodes.KepPromptLang.lib.action.base import Action from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction @@ -15,7 +16,7 @@ from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment # Methods with no comment can be assumed to be the same as comfy.sd1_clip.SD1ClipModel -class PromptLangClipModel(torch.nn.Module): +class PromptLangSDClipModel(torch.nn.Module): """Uses the CLIP transformer encoder for text (from huggingface)""" LAYERS = [ "last", @@ -211,15 +212,23 @@ class PromptLangClipModel(torch.nn.Module): return z_empty.cpu(), first_pooled.cpu() return torch.cat(output, dim=-2).cpu(), first_pooled.cpu() +class PromptLangSD1ClipModel(SD1ClipModel): + def __init__(self, device="cpu", dtype=None, clip_name="l", clip_model=PromptLangSDClipModel): + super().__init__() + self.clip_name = clip_name + self.clip = "clip_{}".format(self.clip_name) + setattr(self, self.clip, clip_model(device=device, dtype=dtype)) + + class PromptLangSDXLClipModel(SDXLClipModel): def __init__(self, device="cpu", dtype=None) -> None: # Skip SDXLClipModel's init super(SDXLClipModel, self).__init__() - self.clip_l = PromptLangClipModel(layer="hidden", layer_idx=11, device=device, dtype=dtype) + self.clip_l = PromptLangSDClipModel(layer="hidden", layer_idx=11, device=device, dtype=dtype) self.clip_l.layer_norm_hidden_state = False self.clip_g = PromptLangSDXLClipG(device, dtype) -class PromptLangSDXLClipG(PromptLangClipModel): +class PromptLangSDXLClipG(PromptLangSDClipModel): def __init__(self, device="cpu", max_length=77, freeze=True, layer="penultimate", layer_idx=None, textmodel_path=None, dtype=None): if layer == "penultimate": layer="hidden" diff --git a/lib/parser/transformer.py b/lib/parser/transformer.py index 71167cf..f078e6f 100644 --- a/lib/parser/transformer.py +++ b/lib/parser/transformer.py @@ -2,7 +2,7 @@ from typing import List from lark import Transformer, Token -from comfy.sd1_clip import SD1Tokenizer +from comfy.sd1_clip import SD1Tokenizer, SDTokenizer from custom_nodes.KepPromptLang.lib.action.base import Action, ActionArity from custom_nodes.KepPromptLang.lib.actions.diff import DiffAction from custom_nodes.KepPromptLang.lib.actions.rand import RandAction @@ -18,7 +18,7 @@ class PromptTransformer(Transformer): # def WORD(self, items): # return items - def __init__(self, tokenizer: SD1Tokenizer): + def __init__(self, tokenizer: SDTokenizer): super().__init__() self.tokenizer = tokenizer diff --git a/lib/parser/utils.py b/lib/parser/utils.py index fad5047..5dfa801 100644 --- a/lib/parser/utils.py +++ b/lib/parser/utils.py @@ -1,6 +1,6 @@ from lark import Token -from comfy.sd1_clip import SD1Tokenizer +from comfy.sd1_clip import SDTokenizer from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment @@ -11,13 +11,12 @@ def flatten_tree(tree): return [str(tree.data)] + sum([flatten_tree(child) for child in tree.children], []) -def build_prompt_segment(text: str, tokenizer: SD1Tokenizer) -> PromptSegment: +def build_prompt_segment(text: str, tokenizer: SDTokenizer) -> PromptSegment: split_text = text.split(" ") tokens = [] for word in split_text: - if word.startswith(tokenizer.clip_l.embedding_identifier) and tokenizer.clip_l.embedding_directory is not None: - embedding_name = word[len(tokenizer.clip_l.embedding_identifier):].strip('\n') - + if word.startswith(tokenizer.embedding_identifier) and tokenizer.embedding_directory is not None: + embedding_name = word[len(tokenizer.embedding_identifier):].strip('\n') get_embed_ret = tokenizer._try_get_embedding(embedding_name) embedding = get_embed_ret[0] @@ -34,6 +33,6 @@ def build_prompt_segment(text: str, tokenizer: SD1Tokenizer) -> PromptSegment: word = leftover else: continue - tokens.extend(tokenizer.clip_l.tokenizer(word)["input_ids"][1:-1]) + tokens.extend(tokenizer.tokenizer(word)["input_ids"][1:-1]) return PromptSegment(text, tokens) diff --git a/lib/tokenizer.py b/lib/tokenizer.py index ccb82df..536afcd 100644 --- a/lib/tokenizer.py +++ b/lib/tokenizer.py @@ -9,19 +9,18 @@ from custom_nodes.KepPromptLang.lib.parser import PromptParser from custom_nodes.KepPromptLang.lib.parser.transformer import PromptTransformer from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment -class PromptLangTokenizer(SD1Tokenizer): - def __init__(self, embedding_directory=None, clip_name='l', tokenizer=SDTokenizer) -> None: - super().__init__(embedding_directory, clip_name, tokenizer) +class PromptLangSDTokenizer(SDTokenizer): + def __init__(self, tokenizer_path=None, max_length=77, pad_with_end=True, embedding_directory=None, embedding_size=768, embedding_key='clip_l'): + super().__init__(tokenizer_path, max_length, pad_with_end, embedding_directory, embedding_size, embedding_key) """ 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[SegOrAction]]: - clip: SDTokenizer = getattr(self, self.clip) - if clip.pad_with_end: - pad_token = clip.end_token + if self.pad_with_end: + pad_token = self.end_token else: pad_token = 0 @@ -30,7 +29,7 @@ class PromptLangTokenizer(SD1Tokenizer): # reshape token array to CLIP input size batched_segments = [] - batch = [PromptSegment(text="[SOT]", tokens=[clip.start_token])] + batch = [PromptSegment(text="[SOT]", tokens=[self.start_token])] # batched_segments.append(batch) batch_size = 1 if isinstance(parsed_actions, Tree): @@ -40,17 +39,17 @@ class PromptLangTokenizer(SD1Tokenizer): for segment in segments_to_process: num_tokens = segment.token_length() # determine if we're going to try and keep the tokens in a single batch - is_large = num_tokens >= clip.max_word_length + is_large = num_tokens >= self.max_word_length # 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 > clip.max_length - 1: - remaining_length = clip.max_length - batch_size + if num_tokens + batch_size > self.max_length - 1: + remaining_length = self.max_length - batch_size # Pad batch - batch.append(PromptSegment("__PAD__", [clip.end_token] + [pad_token] * (remaining_length - 1))) # -1 for end token + batch.append(PromptSegment("__PAD__", [self.end_token] + [pad_token] * (remaining_length - 1))) # -1 for end token batched_segments.append(batch) # start new batch - batch = [PromptSegment(text="[SOT]", tokens=[clip.start_token]), segment] + batch = [PromptSegment(text="[SOT]", tokens=[self.start_token]), segment] batch_size = num_tokens + 1 # +1 for start token continue @@ -59,8 +58,8 @@ class PromptLangTokenizer(SD1Tokenizer): batch_size += num_tokens # Pad the last batch - remaining_length = clip.max_length - batch_size - 1 # -1 for end token - batch.append(PromptSegment("__PAD__", [clip.end_token] + [pad_token] * remaining_length)) + 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: @@ -68,15 +67,18 @@ class PromptLangTokenizer(SD1Tokenizer): return batched_segments +class PromptLangSD1Tokenizer(SD1Tokenizer): + def __init__(self, embedding_directory=None, clip_name='l', tokenizer=PromptLangSDTokenizer) -> None: + super().__init__(embedding_directory, clip_name, tokenizer) -class PromptLangSDXLClipGTokenizer(PromptLangTokenizer): - def __init__(self, tokenizer_path=None, embedding_directory=None) -> None: + +class PromptLangSDXLClipGTokenizer(PromptLangSDTokenizer): + def __init__(self, tokenizer_path=None, embedding_directory=None): super().__init__(tokenizer_path, pad_with_end=False, embedding_directory=embedding_directory, embedding_size=1280, embedding_key='clip_g') - class PromptLangSDXLTokenizer(SD1Tokenizer): def __init__(self, embedding_directory=None) -> None: - self.clip_l = PromptLangTokenizer(embedding_directory=embedding_directory) + self.clip_l = PromptLangSDTokenizer(embedding_directory=embedding_directory) self.clip_g = PromptLangSDXLClipGTokenizer(embedding_directory=embedding_directory) def tokenize_with_weights(self, text:str, return_word_ids=False) -> Dict[str, List[List[SegOrAction]]]: diff --git a/nodes.py b/nodes.py index 7261e85..0766aad 100644 --- a/nodes.py +++ b/nodes.py @@ -12,13 +12,13 @@ from comfy.sd2_clip import SD2ClipModel from comfy.sdxl_clip import SDXLClipModel from comfy.supported_models_base import ClipTarget from custom_nodes.KepPromptLang.lib.clip_model import ( - PromptLangClipModel, PromptLangSDXLClipModel, + PromptLangSD1ClipModel, ) from custom_nodes.KepPromptLang.lib.tokenizer import ( - PromptLangTokenizer, PromptLangSDXLTokenizer, + PromptLangSD1Tokenizer, ) @@ -46,16 +46,17 @@ class SpecialClipLoader: if isinstance(source_clip.cond_stage_model, SDXLClipModel): clip_target = ClipTarget(PromptLangSDXLTokenizer, PromptLangSDXLClipModel) clip = comfy.sd.CLIP(clip_target, embedding_directory=source_clip.tokenizer.clip_g.embedding_directory) + comfy.sd.load_clip_weights(clip.cond_stage_model.clip_g,source_clip.cond_stage_model.clip_g.state_dict()) comfy.sd.load_clip_weights( - clip.cond_stage_model, source_clip.cond_stage_model.state_dict() + clip.cond_stage_model.clip_l, source_clip.cond_stage_model.clip_l.state_dict() ) elif isinstance(source_clip, SD2ClipModel): raise ValueError("SD2 Clip model is not supported.") else: - clip_target = ClipTarget(PromptLangTokenizer, PromptLangClipModel) + clip_target = ClipTarget(PromptLangSD1Tokenizer, PromptLangSD1ClipModel) clip = comfy.sd.CLIP(clip_target, embedding_directory=source_clip.tokenizer.clip_l.embedding_directory) comfy.sd.load_clip_weights( - clip.cond_stage_model, source_clip.cond_stage_model.clip_l.state_dict() + clip.cond_stage_model, source_clip.cond_stage_model.state_dict() ) return (clip,)