refactor/fix: More fixes to align with clip changes in base
This commit is contained in:
+12
-3
@@ -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,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
@@ -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
@@ -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]]]:
|
||||
|
||||
@@ -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,)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user