20 Commits
Author SHA1 Message Date
Michael Poutre 9df0382735 fix: Update to work with transformers changes to attention masks 2023-11-13 19:32:40 -08:00
Michael Poutre 239113501a feat(PosScale: Add node 2023-10-29 18:29:19 -07:00
Michael Poutre c3d45e306e feat: Add support for post embedding modifiers 2023-10-29 18:28:59 -07:00
Michael Poutre 611ee06761 chore: Cleanup imports 2023-10-29 17:33:37 -07:00
Michael Poutre 97f48d0029 refactor/fix: More fixes to align with clip changes in base 2023-10-29 17:32:55 -07:00
Michael Poutre 1059daf5e9 refactor/fix(SD1): Update for changes to clip handling
Fix end token calculation for long prompts
2023-10-29 15:50:08 -07:00
Michael Poutre 2a4b724d59 fix(Clip): Fix bug in EOT token detection 2023-09-13 20:29:37 -07:00
Michael Poutre 55918edc54 feat(SDXL): Initial stab at SDXL 2023-09-13 19:39:17 -07:00
Michael Poutre 21ba1299d6 refactor(Node): Add checks for clip model type 2023-09-13 19:02:21 -07:00
Michael Poutre 407724a64c refactor(Node): Use new comfy.supported_models_base.ClipTarget 2023-09-13 19:00:33 -07:00
Michael Poutre d67d8600ab Merge branch 'func/setDims' 2023-09-05 00:14:13 -07:00
Michael Poutre 6444e2890a feat(func/setDims): Add func 2023-09-04 23:35:24 -07:00
Michael Poutre 28141f2cbe refactor(nodes): Update Build Gif to show preview of Gif 2023-09-04 18:32:53 -07:00
Michael Poutre 70704f5e68 feat(func/scaleDims): Add scaleDims 2023-09-04 18:15:30 -07:00
Michael Poutre e59aaa499d fix(action reg): Don't allow multiple actions with same name 2023-09-04 18:00:35 -07:00
Michael Poutre c98df8d289 feat(fun/average): Add function 2023-08-31 23:50:45 -07:00
Michael Poutre a7c4bbe332 feat(nodes): Update saving method for build_gif 2023-08-31 23:05:41 -07:00
Michael Poutre ce795c52bc refactor(func/mult): Update error messages 2023-08-31 22:51:27 -07:00
Michael Poutre ef9693ec73 feat(nodes): Add frame_duration to build_gif 2023-08-31 22:51:16 -07:00
Michael Poutre 5362ac75fa feat(func/slerp): Add function 2023-08-31 22:50:47 -07:00
18 changed files with 700 additions and 71 deletions
+10
View File
@@ -1,8 +1,13 @@
from custom_nodes.KepPromptLang.lib.actions.avg import AverageAction
from custom_nodes.KepPromptLang.lib.actions.diff import DiffAction
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
from custom_nodes.KepPromptLang.lib.actions.slerp import SlerpAction
from custom_nodes.KepPromptLang.lib.actions.sum import SumAction
from custom_nodes.KepPromptLang.lib.parser.registration import register_action
@@ -12,3 +17,8 @@ register_action(NegAction)
register_action(NormAction)
register_action(RandAction)
register_action(SumAction)
register_action(SlerpAction)
register_action(AverageAction)
register_action(ScaleDims)
register_action(SetDims)
register_action(PosScaleAction)
+9 -2
View File
@@ -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.
+100
View File
@@ -0,0 +1,100 @@
from typing import List, Union
import torch
from torch.nn import Embedding
from custom_nodes.KepPromptLang.lib.action.base import MultiArgAction, Action
from custom_nodes.KepPromptLang.lib.actions.action_utils import get_embedding
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
from custom_nodes.KepPromptLang.lib.actions.utils import slerp
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
class AverageAction(MultiArgAction):
grammar = 'avg(" arg "|" arg "|" arg ")"'
name = "avg"
chars = ["+", "+"]
def __init__(self, args: List[List[Union[PromptSegment, Action]]]) -> None:
super().__init__(args)
if len(args) != 3:
raise ValueError("Average action should have exactly three arguments(2 vectors and a weight)")
self.first_arg = args[0]
self.second_arg = args[1]
self._parse_weight(args[2])
self._validate_args()
def _parse_weight(self, arg: List[SegOrAction]) -> None:
if len(arg) != 1:
raise ValueError("Average weight should have exactly one segment")
weight_seg_or_action = arg[0]
if isinstance(weight_seg_or_action, Action):
raise ValueError("Average weight should not have an action as an argument")
try:
self.parsed_weight = float(weight_seg_or_action.text)
except ValueError:
raise ValueError("Average should have an integer/float as the weight")
def _validate_args(self) -> None:
first_arg_token_length = sum(seg_or_action.token_length() for seg_or_action in self.first_arg)
second_arg_token_length = sum(seg_or_action.token_length() for seg_or_action in self.second_arg)
if first_arg_token_length != second_arg_token_length:
raise ValueError(f"Average start and end arguments should have the same length. Got {start_arg_token_length} and {end_arg_token_length}")
if self.parsed_weight < 0 or self.parsed_weight > 1:
print(f"WARNING: Average weight should be between 0 and 1. Got {self.parsed_weight}")
def token_length(self) -> int:
# Average interpolates between the embeddings of the start and end segments, so the length is the length of the start segment
return sum(seg_or_action.token_length() for seg_or_action in self.first_arg)
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
# Calculate the embeddings for the start segment
all_start_embeddings = [
get_embedding(seg_or_action, embedding_module)
for seg_or_action in self.first_arg
]
start_embedding = torch.cat(all_start_embeddings, dim=1)
# Calculate the embeddings for the end segment
all_end_embeddings = [
get_embedding(seg_or_action, embedding_module)
for seg_or_action in self.second_arg
]
end_embedding = torch.cat(all_end_embeddings, dim=1)
# Perform the weighted average
result = start_embedding * (1 - self.parsed_weight) + end_embedding * self.parsed_weight
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
+3 -3
View File
@@ -26,17 +26,17 @@ class MultiplyAction(MultiArgAction):
def _parse_multiplier(self, arg: List[SegOrAction]) -> None:
if len(arg) != 1:
raise ValueError("Multiply action first argument should have exactly one segment")
raise ValueError("Multiply actions multiplier should have exactly one segment")
multiplier_seg_or_action = arg[0]
if isinstance(multiplier_seg_or_action, Action):
raise ValueError("Multiply action should not have an action as an argument")
raise ValueError("Multiply actions multiplier must be a number")
try:
self.parsed_multiplier = float(multiplier_seg_or_action.text)
except ValueError:
raise ValueError("Multiply action should have an integer/float as the first argument")
raise ValueError("Multiply action should have an integer/float as the multiplier")
def token_length(self) -> int:
"""
+65
View File
@@ -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}
+72
View File
@@ -0,0 +1,72 @@
from typing import List, Union
import torch
from torch.nn import Embedding
from custom_nodes.KepPromptLang.lib.action.base import MultiArgAction, Action
from custom_nodes.KepPromptLang.lib.actions.action_utils import get_embedding
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
class ScaleDims(MultiArgAction):
grammar = 'scaleDims(" arg ("|" arg)* ")"'
name = "scaleDims"
description = "Scales the specified dimensions of the input embeddings by the specified amount"
example = "'The scaleDims(cat|4,1.5|76,1.2) is happy' scales the 4th dimension by 1.5 and the 76th dimension by 1.2 for the word 'cat'"
chars = ["-", "-"]
def __init__(self, args: List[List[Union[PromptSegment, Action]]]):
super().__init__(args)
self.base_arg = args[0]
self._parse_scale_args(args[1:])
def _parse_scale_args(self, args: List[List[Union[PromptSegment, Action]]]) -> None:
# scaleDims scale args should have format of "<dim>,<scale>" where dim is the dimension to scale and scale is the amount to scale it by
# scaleDims(some words|4,1.5|76,1.2)
self.scale_args = []
for arg in args:
if isinstance(arg, Action):
raise ValueError("ScaleDims scale args must be in the format of <dim>,<scale>(e.g. 4,1.5) but got an action")
if len(arg) != 1:
raise ValueError("ScaleDims scale args must be in the format of <dim>,<scale>(e.g. 4,1.5) but got multiple segments")
extracted_arg = arg[0]
assert isinstance(extracted_arg, PromptSegment)
if "," not in extracted_arg.text:
raise ValueError("ScaleDims scale args must be in the format of <dim>,<scale>(e.g. 4,1.5) but got a segment with no comma: " + extracted_arg.text)
# Split prompt segment into text and scale args
dim, scale = extracted_arg.text.split(",")
try:
# TODO: Check that dim is within the bounds of the embedding
parsed_dim = int(dim)
except ValueError:
raise ValueError("ScaleDims scale args must be in the format of <dim>,<scale>(e.g. 4,1.5) but got a segment with a non-integer dim: " + str(dim))
try:
parsed_scale = float(scale)
except ValueError:
raise ValueError("ScaleDims scale args must be in the format of <dim>,<scale>(e.g. 4,1.5) but got a segment with a non-float scale: " + str(scale))
self.scale_args.append((parsed_dim, parsed_scale))
def token_length(self) -> int:
# scaleDims modifies the embeddings of 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
]
base_embeddings = torch.cat(all_base_embeddings, dim=1)
for dim, scale in self.scale_args:
base_embeddings[0, :, dim] *= scale
return base_embeddings
+72
View File
@@ -0,0 +1,72 @@
from typing import List, Union
import torch
from torch.nn import Embedding
from custom_nodes.KepPromptLang.lib.action.base import MultiArgAction, Action
from custom_nodes.KepPromptLang.lib.actions.action_utils import get_embedding
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
class SetDims(MultiArgAction):
grammar = 'setDims(" arg ("|" arg)* ")"'
name = "setDims"
description = "Sets the specified dimensions of the input embeddings to the specified value"
example = "'The scaleDims(cat|4,1.5|76,1.2) is happy' scales the 4th dimension by 1.5 and the 76th dimension by 1.2 for the word 'cat'"
chars = ["-", "-"]
def __init__(self, args: List[List[Union[PromptSegment, Action]]]):
super().__init__(args)
self.base_arg = args[0]
self._parse_value_args(args[1:])
def _parse_value_args(self, args: List[List[Union[PromptSegment, Action]]]) -> None:
# setDims args should have format of "<dim>,<value>" where dim is the dimension to set and value is the value to set it to
# setDims(some words|4,-0.01254|76,1.2)
self.value_args = []
for arg in args:
if isinstance(arg, Action):
raise ValueError("SetDims value args must be in the format of <dim>,<value>(e.g. 4,1.5) but got an action")
if len(arg) != 1:
raise ValueError("SetDims value args must be in the format of <dim>,<value>(e.g. 4,1.5) but got multiple segments")
extracted_arg = arg[0]
assert isinstance(extracted_arg, PromptSegment)
if "," not in extracted_arg.text:
raise ValueError("SetDims value args must be in the format of <dim>,<value>(e.g. 4,1.5) but got a segment with no comma: " + extracted_arg.text)
# Split prompt segment into text and value args
dim, value = extracted_arg.text.split(",")
try:
# TODO: Check that dim is within the bounds of the embedding
parsed_dim = int(dim)
except ValueError:
raise ValueError("SetDims value args must be in the format of <dim>,<value>(e.g. 4,1.5) but got a segment with a non-integer dim: " + str(dim))
try:
parsed_value = float(value)
except ValueError:
raise ValueError("SetDims value args must be in the format of <dim>,<value>(e.g. 4,1.5) but got a segment with a non-float scale: " + str(value))
self.value_args.append((parsed_dim, parsed_value))
def token_length(self) -> int:
# setDims modifies the embeddings of 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
]
base_embeddings = torch.cat(all_base_embeddings, dim=1)
for dim, value in self.value_args:
base_embeddings[0, :, dim] = value
return base_embeddings
+100
View File
@@ -0,0 +1,100 @@
from typing import List, Union
import torch
from torch.nn import Embedding
from custom_nodes.KepPromptLang.lib.action.base import MultiArgAction, Action
from custom_nodes.KepPromptLang.lib.actions.action_utils import get_embedding
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
from custom_nodes.KepPromptLang.lib.actions.utils import slerp
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
class SlerpAction(MultiArgAction):
grammar = 'slerp(" arg "|" arg "|" arg ")"'
name = "slerp"
chars = ["+", "+"]
def __init__(self, args: List[List[Union[PromptSegment, Action]]]) -> None:
super().__init__(args)
if len(args) != 3:
raise ValueError("Slerp action should have exactly three arguments(2 vectors and a weight)")
self.start_argument = args[0]
self.end_argument = args[1]
self._parse_weight(args[2])
self._validate_args()
def _parse_weight(self, arg: List[SegOrAction]) -> None:
if len(arg) != 1:
raise ValueError("Slerp weight should have exactly one segment")
weight_seg_or_action = arg[0]
if isinstance(weight_seg_or_action, Action):
raise ValueError("Slerp weight should not have an action as an argument")
try:
self.parsed_weight = float(weight_seg_or_action.text)
except ValueError:
raise ValueError("Slerp should have an integer/float as the weight")
def _validate_args(self) -> None:
start_arg_token_length = sum(seg_or_action.token_length() for seg_or_action in self.start_argument)
end_arg_token_length = sum(seg_or_action.token_length() for seg_or_action in self.end_argument)
if start_arg_token_length != end_arg_token_length:
raise ValueError(f"Slerp start and end arguments should have the same length. Got {start_arg_token_length} and {end_arg_token_length}")
if self.parsed_weight < 0 or self.parsed_weight > 1:
print(f"WARNING: Slerp weight should be between 0 and 1. Got {self.parsed_weight}")
def token_length(self) -> int:
# Slerp interpolates between the embeddings of the start and end segments, so the length is the length of the start segment
return sum(seg_or_action.token_length() for seg_or_action in self.start_argument)
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
# Calculate the embeddings for the start segment
all_start_embeddings = [
get_embedding(seg_or_action, embedding_module)
for seg_or_action in self.start_argument
]
start_embedding = torch.cat(all_start_embeddings, dim=1)
# Calculate the embeddings for the end segment
all_end_embeddings = [
get_embedding(seg_or_action, embedding_module)
for seg_or_action in self.end_argument
]
end_embedding = torch.cat(all_end_embeddings, dim=1)
# Perform the slerp
result = slerp(self.parsed_weight, start_embedding, end_embedding)
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
+30
View File
@@ -1,8 +1,38 @@
from typing import List
import torch
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
def batch_size_info(batch: List[SegOrAction]):
for segment in batch:
print("Token Len: " + str(segment.token_length()))
print(segment.depth_repr())
def slerp(val: float, low: torch.Tensor, high: torch.Tensor, epsilon=1e-5):
# Convert val to tensor and clamp between 0 and 1
val = torch.tensor(val, dtype=torch.float32).clamp(0, 1)
# Normalize the vectors
low_norm = low / torch.norm(low, dim=-1, keepdim=True)
high_norm = high / torch.norm(high, dim=-1, keepdim=True)
# Calculate the cosine of the angle between the vectors
dot = (low_norm * high_norm).sum(-1, keepdim=True)
# Clamp to prevent numerical errors
dot = torch.clamp(dot, -1, 1)
omega = torch.acos(dot)
# Slerp formula
sin_omega = torch.sin(omega)
scale_0 = torch.sin((1.0 - val) * omega) / (sin_omega + epsilon)
scale_1 = torch.sin(val * omega) / (sin_omega + epsilon)
# Handle the case where omega is small (the vectors are close)
close_condition = sin_omega < epsilon
scale_0 = torch.where(close_condition, 1.0 - val, scale_0)
scale_1 = torch.where(close_condition, val, scale_1)
return scale_0 * low + scale_1 * high
+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
}
+33 -1
View File
@@ -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
View File
@@ -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)
+6
View File
@@ -5,6 +5,12 @@ from custom_nodes.KepPromptLang.lib.action.base import Action
action_registry: Dict[str, Type[Action]] = {}
def register_action(action: Type[Action]) -> None:
"""
:rtype: object
"""
if action.name in action_registry:
raise ValueError(f"Action {action.name} already registered")
action_registry[str(action.name)] = action
def get_action_by_name(name: str) -> Type[Action]:
+2 -7
View File
@@ -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
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,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
View File
@@ -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
+65 -26
View File
@@ -1,5 +1,5 @@
import random
from typing import List, Tuple
import os
from typing import List, Tuple, Any
import numpy as np
from PIL import Image
@@ -7,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:
@@ -32,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,)
@@ -49,9 +65,9 @@ def tensor2img(tensor_img) -> Image.Image:
i_np_arr = np.clip(i, 0, 255, out=i).astype(np.uint8, copy=False)
return Image.fromarray(i_np_arr)
class BuildGif:
def __init__(self) -> None:
self.output_dir = folder_paths.get_output_directory()
pass
@classmethod
@@ -60,6 +76,7 @@ class BuildGif:
"required": {
"images": ("IMAGE",),
"split_every": ("INT", {"default": -1}),
"frame_duration": ("INT", {"default": 125}),
"output_mode": (
["One Per Split", "Big Grid"],
{"default": "Big Grid"},
@@ -68,23 +85,31 @@ class BuildGif:
}
RELOAD_INST = True
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("Gifs",)
RETURN_TYPES = ()
# RETURN_NAMES = ("Gifs",)
INPUT_IS_LIST = True
FUNCTION = "build_gif"
OUTPUT_IS_LIST = (True,)
# OUTPUT_NODE = False
# OUTPUT_IS_LIST = (True,)
OUTPUT_NODE = True
CATEGORY = "List Stuff"
@staticmethod
def build_gif(images: list, split_every: List[int], output_mode: str):
def build_gif(self, images: List[Any], split_every: List[int], frame_duration: List[int], output_mode: List[str]):
print("Build GIF called!")
print(f"{type(images)}")
if len(split_every) > 1:
raise Exception("List input for split every is not supported.")
if len(output_mode) > 1:
raise Exception("List input for output_mode is not supported.")
output_mode = output_mode[0]
if len(frame_duration) > 1:
raise Exception("List input for frame_duration is not supported.")
frame_duration = frame_duration[0]
full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(filename_prefix="Gif", output_dir=self.output_dir, image_width=0, image_height=0)
split_every_val = split_every[0]
batch_size = images[0].size()[0]
if split_every_val == -1:
@@ -93,8 +118,6 @@ class BuildGif:
else:
split_chunks = int(len(images) / split_every_val)
out = []
num_wide = batch_size
num_tall = split_chunks
@@ -104,6 +127,7 @@ class BuildGif:
]
frames = []
results = list()
if output_mode == "Big Grid":
# For every image in gif
@@ -122,8 +146,9 @@ class BuildGif:
)
frames.append(img_frame)
file = f"{filename}_{counter:05}_"
save_path = (
f"{folder_paths.get_output_directory()}/{random.randint(1, 100)}"
f"{os.path.join(full_output_folder, file)}"
)
frames[0].save(
f"{save_path}.webp",
@@ -133,15 +158,24 @@ class BuildGif:
save_all=True,
append_images=frames[1:],
optimize=False,
duration=125,
duration=frame_duration,
loop=0,
)
results.append({
"filename": f"{file}.webp",
"subfolder": subfolder,
"type": "output"
})
elif output_mode == "One Per Split":
for split_idx in range(int(split_chunks)):
split_start = split_every_val * split_idx
split_end = split_every_val * (split_idx + 1)
for batch_idx in range(batch_size):
save_path = f"{folder_paths.get_output_directory()}/-{batch_idx}-{random.randint(1, 100)}"
file = f"{filename}_{counter:05}_"
save_path = (
f"{os.path.join(full_output_folder, file)}"
)
counter += 1
print(save_path)
tensor2img(images[split_start][batch_idx]).save(
f"{save_path}.webp",
@@ -151,7 +185,12 @@ class BuildGif:
for nested_batch in images[split_start + 1 : split_end]
],
optimize=False,
duration=125,
duration=frame_duration,
loop=0,
)
return (out,)
results.append({
"filename": f"{file}.webp",
"subfolder": subfolder,
"type": "output"
})
return { "ui": { "images": results } }
+1
View File
@@ -1 +1,2 @@
lark
packaging