Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
2a4b724d59 | ||
|
|
55918edc54 | ||
|
|
21ba1299d6 | ||
|
|
407724a64c |
@@ -0,0 +1,23 @@
|
||||
{
|
||||
"architectures": [
|
||||
"CLIPTextModel"
|
||||
],
|
||||
"attention_dropout": 0.0,
|
||||
"bos_token_id": 0,
|
||||
"dropout": 0.0,
|
||||
"eos_token_id": 2,
|
||||
"hidden_act": "gelu",
|
||||
"hidden_size": 1280,
|
||||
"initializer_factor": 1.0,
|
||||
"initializer_range": 0.02,
|
||||
"intermediate_size": 5120,
|
||||
"layer_norm_eps": 1e-05,
|
||||
"max_position_embeddings": 77,
|
||||
"model_type": "clip_text_model",
|
||||
"num_attention_heads": 20,
|
||||
"num_hidden_layers": 32,
|
||||
"pad_token_id": 1,
|
||||
"projection_dim": 1280,
|
||||
"torch_dtype": "float32",
|
||||
"vocab_size": 49408
|
||||
}
|
||||
@@ -7,6 +7,7 @@ from transformers import CLIPTextConfig, modeling_utils
|
||||
|
||||
from comfy import model_management
|
||||
import comfy.ops
|
||||
from comfy.sdxl_clip import SDXLClipModel
|
||||
from custom_nodes.KepPromptLang.lib.action.base import Action
|
||||
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
|
||||
from custom_nodes.KepPromptLang.lib.fun_clip_stuff import PromptLangTextModel
|
||||
@@ -209,3 +210,25 @@ class PromptLangClipModel(torch.nn.Module):
|
||||
if (len(output) == 0):
|
||||
return z_empty.cpu(), first_pooled.cpu()
|
||||
return torch.cat(output, dim=-2).cpu(), first_pooled.cpu()
|
||||
|
||||
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.layer_norm_hidden_state = False
|
||||
self.clip_g = PromptLangSDXLClipG(device, dtype)
|
||||
|
||||
class PromptLangSDXLClipG(PromptLangClipModel):
|
||||
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"
|
||||
layer_idx=-2
|
||||
|
||||
textmodel_json_config = os.path.join(os.path.dirname(os.path.realpath(__file__)), "clip_config_bigg.json")
|
||||
super().__init__(device=device, freeze=freeze, layer=layer, layer_idx=layer_idx, textmodel_json_config=textmodel_json_config, textmodel_path=textmodel_path, dtype=dtype)
|
||||
self.empty_tokens = [[49406] + [49407] + [0] * 75]
|
||||
self.layer_norm_hidden_state = False
|
||||
|
||||
def load_sd(self, sd):
|
||||
return super().load_sd(sd)
|
||||
|
||||
@@ -143,6 +143,9 @@ class PrompLangCLIPTextTransformer(CLIPTextTransformer):
|
||||
else:
|
||||
if seg_or_action.text == '__PAD__':
|
||||
break
|
||||
|
||||
# Is a segment, and isn't the pad segment
|
||||
idx += seg_or_action.token_length()
|
||||
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)
|
||||
|
||||
+19
-2
@@ -1,4 +1,4 @@
|
||||
from typing import List
|
||||
from typing import List, Dict
|
||||
|
||||
from lark import Tree
|
||||
|
||||
@@ -10,7 +10,7 @@ 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, tokenizer_path=None, max_length=77, pad_with_end=True, embedding_directory=None, embedding_size=768, embedding_key='clip_l', special_tokens=None):
|
||||
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) -> None:
|
||||
super().__init__(tokenizer_path, max_length, pad_with_end, embedding_directory, embedding_size, embedding_key)
|
||||
|
||||
"""
|
||||
@@ -66,3 +66,20 @@ class PromptLangTokenizer(SD1Tokenizer):
|
||||
# batch_size_info(batch)
|
||||
|
||||
return batched_segments
|
||||
|
||||
|
||||
class PromptLangSDXLClipGTokenizer(PromptLangTokenizer):
|
||||
def __init__(self, tokenizer_path=None, embedding_directory=None) -> 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_g = PromptLangSDXLClipGTokenizer(embedding_directory=embedding_directory)
|
||||
|
||||
def tokenize_with_weights(self, text:str, return_word_ids=False) -> Dict[str, List[List[SegOrAction]]]:
|
||||
out = {}
|
||||
out["g"] = self.clip_g.tokenize_with_weights(text, return_word_ids)
|
||||
out["l"] = self.clip_l.tokenize_with_weights(text, return_word_ids)
|
||||
return out
|
||||
|
||||
@@ -8,9 +8,18 @@ from PIL import Image
|
||||
import folder_paths
|
||||
import comfy.sd
|
||||
import comfy.ops
|
||||
from custom_nodes.KepPromptLang.lib.clip_model import PromptLangClipModel
|
||||
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,
|
||||
)
|
||||
|
||||
from custom_nodes.KepPromptLang.lib.tokenizer import PromptLangTokenizer
|
||||
from custom_nodes.KepPromptLang.lib.tokenizer import (
|
||||
PromptLangTokenizer,
|
||||
PromptLangSDXLTokenizer,
|
||||
)
|
||||
|
||||
|
||||
class EmptyClass:
|
||||
@@ -33,15 +42,21 @@ class SpecialClipLoader:
|
||||
|
||||
@staticmethod
|
||||
def load_clip(source_clip: comfy.sd.CLIP) -> Tuple[comfy.sd.CLIP]:
|
||||
clip_target = EmptyClass()
|
||||
clip_target.params = {}
|
||||
clip_target.clip = PromptLangClipModel
|
||||
clip_target.tokenizer = PromptLangTokenizer
|
||||
|
||||
clip = comfy.sd.CLIP(clip_target, embedding_directory=source_clip.tokenizer.embedding_directory)
|
||||
comfy.sd.load_clip_weights(
|
||||
clip.cond_stage_model, source_clip.cond_stage_model.state_dict()
|
||||
)
|
||||
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, source_clip.cond_stage_model.state_dict()
|
||||
)
|
||||
elif isinstance(source_clip, SD2ClipModel):
|
||||
raise ValueError("SD2 Clip model is not supported.")
|
||||
else:
|
||||
clip_target = ClipTarget(PromptLangTokenizer, PromptLangClipModel)
|
||||
clip = comfy.sd.CLIP(clip_target, embedding_directory=source_clip.tokenizer.embedding_directory)
|
||||
comfy.sd.load_clip_weights(
|
||||
clip.cond_stage_model, source_clip.cond_stage_model.state_dict()
|
||||
)
|
||||
return (clip,)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user