refactor/fix: More fixes to align with clip changes in base

This commit is contained in:
Michael Poutre
2023-11-13 19:47:56 -08:00
parent 2f115352ec
commit ef3f66ef91
5 changed files with 45 additions and 34 deletions
+12 -3
View File
@@ -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"
+2 -2
View File
@@ -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
+5 -6
View File
@@ -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)
+20 -18
View File
@@ -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]]]:
+6 -5
View File
@@ -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,)