Compare commits
10
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
9df0382735 | ||
|
|
239113501a | ||
|
|
c3d45e306e | ||
|
|
611ee06761 | ||
|
|
97f48d0029 | ||
|
|
1059daf5e9 | ||
|
|
2a4b724d59 | ||
|
|
55918edc54 | ||
|
|
21ba1299d6 | ||
|
|
407724a64c |
+2
-6
@@ -1,11 +1,9 @@
|
||||
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
|
||||
from custom_nodes.KepPromptLang.lib.actions.pos_scale import PosScaleAction
|
||||
from custom_nodes.KepPromptLang.lib.actions.rand import RandAction
|
||||
from custom_nodes.KepPromptLang.lib.actions.scale_dims import ScaleDims
|
||||
from custom_nodes.KepPromptLang.lib.actions.set_dims import SetDims
|
||||
@@ -23,6 +21,4 @@ register_action(SlerpAction)
|
||||
register_action(AverageAction)
|
||||
register_action(ScaleDims)
|
||||
register_action(SetDims)
|
||||
register_action(MaxAction)
|
||||
register_action(AbsMaxAction)
|
||||
register_action(MinAction)
|
||||
register_action(PosScaleAction)
|
||||
|
||||
+9
-2
@@ -1,6 +1,6 @@
|
||||
from abc import ABC, abstractmethod
|
||||
from enum import Enum
|
||||
from typing import Union, List
|
||||
from typing import Union, List, TypedDict, Tuple
|
||||
|
||||
from torch import Tensor
|
||||
from torch.nn import Embedding
|
||||
@@ -13,6 +13,13 @@ class ActionArity(Enum):
|
||||
SINGLE = 1
|
||||
MULTI = 2
|
||||
|
||||
class PostModifiers(TypedDict):
|
||||
"""
|
||||
A dictionary of post modifiers for an action result.
|
||||
"""
|
||||
position_embed_scale: Union[float, None]
|
||||
|
||||
|
||||
class Action(ABC):
|
||||
@property
|
||||
@abstractmethod
|
||||
@@ -67,7 +74,7 @@ class Action(ABC):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def get_result(self, embedding_module: Embedding) -> Tensor:
|
||||
def get_result(self, embedding_module: Embedding) -> Union[Tensor, Tuple[Tensor, PostModifiers]]:
|
||||
"""
|
||||
Get the result of this action. This is called when the embeddings are being calculated.
|
||||
:param embedding_module: The embedding module to use to get the base embeddings for tokens.
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -0,0 +1,65 @@
|
||||
from typing import Tuple, List
|
||||
|
||||
import torch
|
||||
from torch.nn import Embedding
|
||||
|
||||
from custom_nodes.KepPromptLang.lib.action.base import (
|
||||
Action,
|
||||
PostModifiers,
|
||||
MultiArgAction,
|
||||
)
|
||||
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
|
||||
|
||||
|
||||
class PosScaleAction(MultiArgAction):
|
||||
grammar = 'posScale(" arg+ ")"'
|
||||
name = "posScale"
|
||||
chars = ["[", "]"]
|
||||
|
||||
def __init__(self, args: List[List[SegOrAction]]) -> None:
|
||||
super().__init__(args)
|
||||
if len(args) != 2:
|
||||
raise ValueError("PosScale action should have exactly two arguments")
|
||||
|
||||
self.target_arg = args[0]
|
||||
self._parse_multiplier(args[1])
|
||||
|
||||
def _parse_multiplier(self, arg: List[SegOrAction]) -> None:
|
||||
if len(arg) != 1:
|
||||
raise ValueError(
|
||||
"PosScale actions multiplier should have exactly one segment"
|
||||
)
|
||||
|
||||
multiplier_seg_or_action = arg[0]
|
||||
|
||||
if isinstance(multiplier_seg_or_action, Action):
|
||||
raise ValueError("PosScale actions multiplier must be a number")
|
||||
|
||||
try:
|
||||
self.parsed_multiplier = float(multiplier_seg_or_action.text)
|
||||
except ValueError:
|
||||
raise ValueError(
|
||||
"PosScale action should have an integer/float as the multiplier"
|
||||
)
|
||||
|
||||
def token_length(self) -> int:
|
||||
"""
|
||||
PosScale modifies the posional embeddings of the base segment, so the length is the length of the base segment
|
||||
:return:
|
||||
"""
|
||||
total_length = 0
|
||||
for seg_or_action in self.target_arg:
|
||||
total_length += seg_or_action.token_length()
|
||||
|
||||
return total_length
|
||||
|
||||
def get_result(self, embedding_module: Embedding) -> Tuple[torch.Tensor, PostModifiers]:
|
||||
all_embeddings = []
|
||||
for seg_or_action in self.target_arg:
|
||||
if isinstance(seg_or_action, Action):
|
||||
all_embeddings.append(seg_or_action.get_result(embedding_module))
|
||||
else:
|
||||
all_embeddings.append(seg_or_action.get_embeddings(embedding_module))
|
||||
|
||||
target_embeddings = torch.cat(all_embeddings, dim=1)
|
||||
return target_embeddings, {"position_embed_scale": self.parsed_multiplier}
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
+33
-1
@@ -7,6 +7,8 @@ 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
|
||||
from custom_nodes.KepPromptLang.lib.fun_clip_stuff import PromptLangTextModel
|
||||
@@ -14,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",
|
||||
@@ -209,3 +211,33 @@ 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 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 = 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(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"
|
||||
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)
|
||||
|
||||
+80
-23
@@ -1,10 +1,11 @@
|
||||
from typing import Optional, Tuple, Union, List
|
||||
from typing import Optional, Tuple, Union, List, TypedDict
|
||||
from importlib_metadata import version as import_version
|
||||
from packaging import version
|
||||
|
||||
import torch
|
||||
from transformers import CLIPTextConfig
|
||||
from transformers.modeling_outputs import BaseModelOutputWithPooling
|
||||
from transformers.models.clip.modeling_clip import (
|
||||
_expand_mask,
|
||||
CLIPTextEmbeddings,
|
||||
CLIPTextTransformer,
|
||||
CLIPTextModel,
|
||||
@@ -23,6 +24,17 @@ def slerp(val, low, high):
|
||||
res = (torch.sin((1.0-val)*omega)/so).unsqueeze(1)*low + (torch.sin(val*omega)/so).unsqueeze(1) * high
|
||||
return res
|
||||
|
||||
|
||||
class PosModifier(TypedDict):
|
||||
"""
|
||||
A dictionary of post modifiers for an action result.
|
||||
"""
|
||||
|
||||
position_embed_scale: Union[float]
|
||||
start_idx: Union[int]
|
||||
end_idx: Union[int]
|
||||
|
||||
|
||||
class PromptLangCLIPTextEmbeddings(CLIPTextEmbeddings):
|
||||
def __init__(self, config: CLIPTextConfig):
|
||||
super().__init__(config)
|
||||
@@ -39,14 +51,30 @@ class PromptLangCLIPTextEmbeddings(CLIPTextEmbeddings):
|
||||
raise ValueError("You have to specify input_dicts")
|
||||
|
||||
batches = []
|
||||
pos_modifiers: List[List[PosModifier]] = []
|
||||
for batch_idx, batch in enumerate(input_dicts):
|
||||
results = []
|
||||
batch_pos_modifiers = []
|
||||
token_idx = 0
|
||||
for seg_or_action in batch:
|
||||
if isinstance(seg_or_action, Action):
|
||||
results.append(seg_or_action.get_result(self.token_embedding))
|
||||
action_result = seg_or_action.get_result(self.token_embedding)
|
||||
if isinstance(action_result, tuple):
|
||||
result, post_modifiers = action_result
|
||||
if post_modifiers["position_embed_scale"] is not None:
|
||||
post_modifiers["start_idx"] = token_idx
|
||||
post_modifiers["end_idx"] = (
|
||||
token_idx + seg_or_action.token_length()
|
||||
)
|
||||
batch_pos_modifiers.append(post_modifiers)
|
||||
else:
|
||||
result = action_result
|
||||
else:
|
||||
results.append(seg_or_action.get_embeddings(self.token_embedding))
|
||||
result = seg_or_action.get_embeddings(self.token_embedding)
|
||||
results.append(result)
|
||||
token_idx += seg_or_action.token_length()
|
||||
batches.append(results)
|
||||
pos_modifiers.append(batch_pos_modifiers)
|
||||
|
||||
seq_length = batches[0][0].shape[-2]
|
||||
|
||||
@@ -60,8 +88,16 @@ class PromptLangCLIPTextEmbeddings(CLIPTextEmbeddings):
|
||||
else:
|
||||
embeds.append(torch.cat(batch, dim=-2))
|
||||
|
||||
position_embeddings = self.position_embedding(position_ids)
|
||||
embeddings = torch.cat(embeds, dim=0) + position_embeddings
|
||||
for idx, batch_pos_modifiers in enumerate(pos_modifiers):
|
||||
position_embeddings = self.position_embedding(position_ids)
|
||||
if len(batch_pos_modifiers) > 0:
|
||||
print(f"Found {len(batch_pos_modifiers)} pos modifiers for batch {idx}")
|
||||
for post_modifier in batch_pos_modifiers:
|
||||
position_embeddings[
|
||||
0, post_modifier["start_idx"] : post_modifier["end_idx"]
|
||||
] *= post_modifier["position_embed_scale"]
|
||||
embeds[idx] = embeds[idx] + position_embeddings
|
||||
embeddings = torch.cat(embeds, dim=0)
|
||||
|
||||
return embeddings
|
||||
|
||||
@@ -70,6 +106,40 @@ class PrompLangCLIPTextTransformer(CLIPTextTransformer):
|
||||
def __init__(self, config: CLIPTextConfig):
|
||||
super().__init__(config)
|
||||
self.embeddings = PromptLangCLIPTextEmbeddings(config)
|
||||
self.transformers_version = version.parse(import_version('transformers'))
|
||||
|
||||
def process_attention_mask(self, hidden_states, attention_mask, bsz, seq_len):
|
||||
# Parse the transformer version
|
||||
input_shape = torch.Size([bsz, seq_len])
|
||||
|
||||
v4_30 = version.parse('4.30.0')
|
||||
v4_35 = version.parse('4.35')
|
||||
if self.transformers_version < v4_30:
|
||||
print("Using transformers < 4.30.0")
|
||||
causal_attention_mask = self._build_causal_attention_mask(bsz, seq_len, hidden_states.dtype).to(
|
||||
hidden_states.device)
|
||||
elif v4_30 <= self.transformers_version < v4_35:
|
||||
print("Using transformers >= 4.30.0 and <= 4.34.*")
|
||||
from transformers.models.clip.modeling_clip import _make_causal_mask
|
||||
causal_attention_mask = _make_causal_mask(input_shape, hidden_states.dtype, device=hidden_states.device)
|
||||
else:
|
||||
print("Using transformers >= 4.35")
|
||||
from transformers.modeling_attn_mask_utils import _create_4d_causal_attention_mask
|
||||
causal_attention_mask = _create_4d_causal_attention_mask(
|
||||
input_shape, hidden_states.dtype, device=hidden_states.device
|
||||
)
|
||||
|
||||
# Expand attention_mask if it exists
|
||||
if attention_mask is not None:
|
||||
# Import _expand_mask or _prepare_4d_attention_mask based on version
|
||||
if self.transformers_version < v4_35:
|
||||
from transformers.models.clip.modeling_clip import _expand_mask
|
||||
attention_mask = _expand_mask(attention_mask, hidden_states.dtype)
|
||||
else:
|
||||
from transformers.modeling_attn_mask_utils import _prepare_4d_attention_mask
|
||||
attention_mask = _prepare_4d_attention_mask(attention_mask, hidden_states.dtype)
|
||||
|
||||
return causal_attention_mask, attention_mask
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -101,24 +171,8 @@ class PrompLangCLIPTextTransformer(CLIPTextTransformer):
|
||||
bsz = len(input_ids)
|
||||
# TODO: Properly gather this
|
||||
seq_len = 77
|
||||
input_shape = torch.Size([bsz, seq_len])
|
||||
# CLIP's text model uses causal mask, prepare it here.
|
||||
# https://github.com/openai/CLIP/blob/cfcffb90e69f37bf2ff1e988237a0fbe41f33c04/clip/model.py#L324
|
||||
## VERSION DIFF ##
|
||||
# transformers < 4.30.0
|
||||
if hasattr(self, "_build_causal_attention_mask"):
|
||||
print("Using transformers < 4.30.0")
|
||||
causal_attention_mask = self._build_causal_attention_mask(bsz, seq_len, hidden_states.dtype).to(hidden_states.device)
|
||||
else:
|
||||
# transformers >= 4.30.0
|
||||
print("Using transformers >= 4.30.0")
|
||||
from transformers.models.clip.modeling_clip import _make_causal_mask
|
||||
causal_attention_mask = _make_causal_mask(input_shape, hidden_states.dtype, device=hidden_states.device)
|
||||
|
||||
# expand attention_mask
|
||||
if attention_mask is not None:
|
||||
# [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]
|
||||
attention_mask = _expand_mask(attention_mask, hidden_states.dtype)
|
||||
causal_attention_mask, attention_mask = self.process_attention_mask(hidden_states, attention_mask, bsz, seq_len)
|
||||
|
||||
encoder_outputs = self.encoder(
|
||||
inputs_embeds=hidden_states,
|
||||
@@ -143,6 +197,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)
|
||||
|
||||
@@ -2,15 +2,10 @@ from typing import List
|
||||
|
||||
from lark import Transformer, Token
|
||||
|
||||
from comfy.sd1_clip import SD1Tokenizer
|
||||
from comfy.sd1_clip import 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
|
||||
from custom_nodes.KepPromptLang.lib.parser.registration import get_action_by_name
|
||||
from custom_nodes.KepPromptLang.lib.parser.utils import build_prompt_segment
|
||||
from custom_nodes.KepPromptLang.lib.actions.neg import NegAction
|
||||
from custom_nodes.KepPromptLang.lib.actions.norm import NormAction
|
||||
from custom_nodes.KepPromptLang.lib.actions.sum import SumAction
|
||||
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
|
||||
|
||||
|
||||
@@ -18,7 +13,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
|
||||
|
||||
|
||||
+2
-2
@@ -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,7 +11,7 @@ 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:
|
||||
|
||||
+27
-7
@@ -1,18 +1,18 @@
|
||||
from typing import List
|
||||
from typing import List, Dict
|
||||
|
||||
from lark import Tree
|
||||
|
||||
from comfy.sd1_clip import SD1Tokenizer
|
||||
from comfy.sd1_clip import SD1Tokenizer, SDTokenizer
|
||||
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
|
||||
|
||||
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, tokenizer_path=None, max_length=77, pad_with_end=True, embedding_directory=None, embedding_size=768, embedding_key='clip_l', special_tokens=None):
|
||||
super().__init__(tokenizer_path, max_length, pad_with_end, embedding_directory, embedding_size, embedding_key)
|
||||
|
||||
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
|
||||
@@ -43,9 +43,9 @@ class PromptLangTokenizer(SD1Tokenizer):
|
||||
|
||||
# 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 > self.max_length - 1:
|
||||
remaining_length = self.max_length - batch_size - 1 # -1 for end token
|
||||
remaining_length = self.max_length - batch_size
|
||||
# Pad batch
|
||||
batch.append(PromptSegment("__PAD__", [self.end_token] + [pad_token] * remaining_length - 1))
|
||||
batch.append(PromptSegment("__PAD__", [self.end_token] + [pad_token] * (remaining_length - 1))) # -1 for end token
|
||||
batched_segments.append(batch)
|
||||
|
||||
# start new batch
|
||||
@@ -66,3 +66,23 @@ class PromptLangTokenizer(SD1Tokenizer):
|
||||
# batch_size_info(batch)
|
||||
|
||||
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(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 = 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]]]:
|
||||
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
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
import random
|
||||
import os
|
||||
from typing import List, Tuple, Any
|
||||
|
||||
@@ -8,9 +7,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 (
|
||||
PromptLangSDXLClipModel,
|
||||
PromptLangSD1ClipModel,
|
||||
)
|
||||
|
||||
from custom_nodes.KepPromptLang.lib.tokenizer import PromptLangTokenizer
|
||||
from custom_nodes.KepPromptLang.lib.tokenizer import (
|
||||
PromptLangSDXLTokenizer,
|
||||
PromptLangSD1Tokenizer,
|
||||
)
|
||||
|
||||
|
||||
class EmptyClass:
|
||||
@@ -33,15 +41,22 @@ 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.clip_g,source_clip.cond_stage_model.clip_g.state_dict())
|
||||
comfy.sd.load_clip_weights(
|
||||
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(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.state_dict()
|
||||
)
|
||||
return (clip,)
|
||||
|
||||
|
||||
|
||||
@@ -1 +1,2 @@
|
||||
lark
|
||||
packaging
|
||||
|
||||
Reference in New Issue
Block a user