4 Commits
10 changed files with 93 additions and 283 deletions
-6
View File
@@ -1,8 +1,5 @@
from custom_nodes.KepPromptLang.lib.actions.abs_max import AbsMaxAction
from custom_nodes.KepPromptLang.lib.actions.avg import AverageAction
from custom_nodes.KepPromptLang.lib.actions.diff import DiffAction
from custom_nodes.KepPromptLang.lib.actions.max import MaxAction
from custom_nodes.KepPromptLang.lib.actions.min import MinAction
from custom_nodes.KepPromptLang.lib.actions.mult import MultiplyAction
from custom_nodes.KepPromptLang.lib.actions.neg import NegAction
from custom_nodes.KepPromptLang.lib.actions.norm import NormAction
@@ -23,6 +20,3 @@ register_action(SlerpAction)
register_action(AverageAction)
register_action(ScaleDims)
register_action(SetDims)
register_action(MaxAction)
register_action(AbsMaxAction)
register_action(MinAction)
-104
View File
@@ -1,104 +0,0 @@
from typing import List, Union
import torch
from torch.nn import Embedding
from custom_nodes.KepPromptLang.lib.action.base import Action, MultiArgAction
from custom_nodes.KepPromptLang.lib.actions.action_utils import get_embedding
from custom_nodes.KepPromptLang.lib.actions.utils import is_broadcastable
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
class AbsMaxAction(MultiArgAction):
grammar = 'absMax(" arg ("|" arg)+ ")"'
name = "absMax"
chars = ["+", "+"]
def __init__(self, args: List[List[Union[PromptSegment, Action]]]) -> None:
super().__init__(args)
self.base_arg = args[0]
self.additional_args = args[1:]
def token_length(self) -> int:
# AbsMax modifies the base segment, so the length is the length of the base segment
return sum(seg_or_action.token_length() for seg_or_action in self.base_arg)
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
# Calculate the embeddings for the base segment
all_base_embeddings = [
get_embedding(seg_or_action, embedding_module)
for seg_or_action in self.base_arg
]
result = torch.cat(all_base_embeddings, dim=1)
for arg in self.additional_args:
all_arg_embeddings = [
get_embedding(seg_or_action, embedding_module) for seg_or_action in arg
]
arg_embedding = torch.cat(all_arg_embeddings, dim=1)
if is_broadcastable(result, arg_embedding):
positive_mask1 = result > 0
positive_mask2 = arg_embedding > 0
negative_mask1 = result < 0
negative_mask2 = arg_embedding < 0
zero_mask1 = result == 0
zero_mask2 = arg_embedding == 0
# For mixed signs, choose the one with the largest magnitude
mixed_mask_pos_neg = positive_mask1 & negative_mask2
mixed_mask_neg_pos = negative_mask1 & positive_mask2
mixed_mask = mixed_mask_pos_neg | mixed_mask_neg_pos
mixed_selection = torch.where(
result.abs() > arg_embedding.abs(), result, arg_embedding
)
# Apply max for positive dimensions, min for negative dimensions, and handle zeros and mixed signs
result = torch.where(
positive_mask1 & positive_mask2,
torch.max(result, arg_embedding),
torch.where(
negative_mask1 & negative_mask2,
torch.min(result, arg_embedding),
torch.where(
positive_mask1 & zero_mask2,
result,
torch.where(
zero_mask1 & positive_mask2,
arg_embedding,
torch.where(mixed_mask, mixed_selection, result),
),
),
),
)
else:
print(
"WARNING: shape mismatch when trying to apply absMax, arg will be averaged"
)
result = torch.max(result, torch.mean(arg_embedding, dim=1, keepdim=True))
return result
# def __repr__(self):
# return f"sum(\n\tbase_segment={self.base_segment},\n\targs={self.args}\n)"
def __repr__(self) -> str:
return f"sum({', '.join(map(str, self.additional_args))})"
def depth_repr(self, depth=1):
out = "NudgeAction(\n"
if isinstance(self.base_arg, Action):
base_segment_repr = self.base_arg.depth_repr(depth + 1)
out += "\t" * depth + f"base_segment={base_segment_repr}\n"
else:
out += "\t" * depth + f"base_segment={self.base_arg.depth_repr()},\n"
if isinstance(self.additional_args, Action):
target_repr = self.additional_args.depth_repr(depth + 1)
out += "\t" * depth + f"target={target_repr},\n"
else:
out += "\t" * depth + f"target={self.additional_args.depth_repr()},\n"
out += "\t" * depth + f"weight={self.weight},\n"
out += "\t" * (depth - 1) + ")"
return out
-72
View File
@@ -1,72 +0,0 @@
from typing import List, Union
import torch
from torch.nn import Embedding
from custom_nodes.KepPromptLang.lib.action.base import Action, MultiArgAction
from custom_nodes.KepPromptLang.lib.actions.action_utils import get_embedding
from custom_nodes.KepPromptLang.lib.actions.utils import is_broadcastable
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
class MaxAction(MultiArgAction):
grammar = 'max(" arg ("|" arg)+ ")"'
name = "max"
chars = ["+", "+"]
def __init__(self, args: List[List[Union[PromptSegment, Action]]]) -> None:
super().__init__(args)
self.base_arg = args[0]
self.additional_args = args[1:]
def token_length(self) -> int:
# Max modifies the base segment, so the length is the length of the base segment
return sum(seg_or_action.token_length() for seg_or_action in self.base_arg)
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
# Calculate the embeddings for the base segment
all_base_embeddings = [
get_embedding(seg_or_action, embedding_module)
for seg_or_action in self.base_arg
]
result = torch.cat(all_base_embeddings, dim=1)
for arg in self.additional_args:
all_arg_embeddings = [
get_embedding(seg_or_action, embedding_module) for seg_or_action in arg
]
arg_embedding = torch.cat(all_arg_embeddings, dim=1)
if is_broadcastable(result, arg_embedding):
result = torch.max(result, arg_embedding)
else:
print(
"WARNING: shape mismatch when trying to apply max, arg will be averaged"
)
result = torch.max(result, torch.mean(arg_embedding, dim=1, keepdim=True))
return result
# def __repr__(self):
# return f"sum(\n\tbase_segment={self.base_segment},\n\targs={self.args}\n)"
def __repr__(self) -> str:
return f"sum({', '.join(map(str, self.additional_args))})"
def depth_repr(self, depth=1):
out = "NudgeAction(\n"
if isinstance(self.base_arg, Action):
base_segment_repr = self.base_arg.depth_repr(depth + 1)
out += "\t" * depth + f"base_segment={base_segment_repr}\n"
else:
out += "\t" * depth + f"base_segment={self.base_arg.depth_repr()},\n"
if isinstance(self.additional_args, Action):
target_repr = self.additional_args.depth_repr(depth + 1)
out += "\t" * depth + f"target={target_repr},\n"
else:
out += "\t" * depth + f"target={self.additional_args.depth_repr()},\n"
out += "\t" * depth + f"weight={self.weight},\n"
out += "\t" * (depth - 1) + ")"
return out
-72
View File
@@ -1,72 +0,0 @@
from typing import List, Union
import torch
from torch.nn import Embedding
from custom_nodes.KepPromptLang.lib.action.base import Action, MultiArgAction
from custom_nodes.KepPromptLang.lib.actions.action_utils import get_embedding
from custom_nodes.KepPromptLang.lib.actions.utils import is_broadcastable
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
class MinAction(MultiArgAction):
grammar = 'min(" arg ("|" arg)+ ")"'
name = "min"
chars = ["+", "+"]
def __init__(self, args: List[List[Union[PromptSegment, Action]]]) -> None:
super().__init__(args)
self.base_arg = args[0]
self.additional_args = args[1:]
def token_length(self) -> int:
# Min modifies the base segment, so the length is the length of the base segment
return sum(seg_or_action.token_length() for seg_or_action in self.base_arg)
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
# Calculate the embeddings for the base segment
all_base_embeddings = [
get_embedding(seg_or_action, embedding_module)
for seg_or_action in self.base_arg
]
result = torch.cat(all_base_embeddings, dim=1)
for arg in self.additional_args:
all_arg_embeddings = [
get_embedding(seg_or_action, embedding_module) for seg_or_action in arg
]
arg_embedding = torch.cat(all_arg_embeddings, dim=1)
if is_broadcastable(result, arg_embedding):
result = torch.min(result, arg_embedding)
else:
print(
"WARNING: shape mismatch when trying to apply max, arg will be averaged"
)
result = torch.min(result, torch.mean(arg_embedding, dim=1, keepdim=True))
return result
# def __repr__(self):
# return f"sum(\n\tbase_segment={self.base_segment},\n\targs={self.args}\n)"
def __repr__(self) -> str:
return f"sum({', '.join(map(str, self.additional_args))})"
def depth_repr(self, depth=1):
out = "NudgeAction(\n"
if isinstance(self.base_arg, Action):
base_segment_repr = self.base_arg.depth_repr(depth + 1)
out += "\t" * depth + f"base_segment={base_segment_repr}\n"
else:
out += "\t" * depth + f"base_segment={self.base_arg.depth_repr()},\n"
if isinstance(self.additional_args, Action):
target_repr = self.additional_args.depth_repr(depth + 1)
out += "\t" * depth + f"target={target_repr},\n"
else:
out += "\t" * depth + f"target={self.additional_args.depth_repr()},\n"
out += "\t" * depth + f"weight={self.weight},\n"
out += "\t" * (depth - 1) + ")"
return out
-17
View File
@@ -36,20 +36,3 @@ def slerp(val: float, low: torch.Tensor, high: torch.Tensor, epsilon=1e-5):
scale_1 = torch.where(close_condition, val, scale_1)
return scale_0 * low + scale_1 * high
def is_broadcastable(tensor1, tensor2) -> bool:
"""
Check if two tensors are broadcastable.
Parameters:
- tensor1 (torch.Tensor): The target tensor against which broadcastability of tensor2 is checked.
- tensor2 (torch.Tensor): The tensor whose broadcastability is to be verified against tensor1.
Returns:
- bool: True if tensor2 is broadcastable to tensor1, False otherwise.
"""
try:
broadcasted_shape = torch.broadcast_shapes(tensor1.shape, tensor2.shape)
return True
except RuntimeError:
return False
+23
View File
@@ -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
}
+23
View File
@@ -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)
+3
View File
@@ -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
View File
@@ -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
+25 -10
View File
@@ -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,)