Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f2256932c8 | ||
|
|
94c46fe1ba | ||
|
|
ca11ea8e82 | ||
|
|
fb363aed7a |
@@ -1,9 +1,13 @@
|
||||
from .nodes import (
|
||||
BuildGif,
|
||||
SpecialClipLoader,
|
||||
MonacoPrompt,
|
||||
)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"Build Gif": BuildGif,
|
||||
"Special CLIP Loader": SpecialClipLoader,
|
||||
"Monaco Prompt": MonacoPrompt,
|
||||
}
|
||||
|
||||
WEB_DIRECTORY = ("./web/dist", ["app.bundle.js"])
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -17,6 +17,25 @@ class EmptyClass:
|
||||
pass
|
||||
|
||||
|
||||
class MonacoPrompt:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"clip": ("CLIP",),
|
||||
"prompt": ("MONACO",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING",)
|
||||
FUNCTION = "do_crap"
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
CATEGORY = "conditioning"
|
||||
|
||||
@staticmethod
|
||||
def do_crap(clip, prompt):
|
||||
return (clip,)
|
||||
|
||||
class SpecialClipLoader:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls): # type: ignore
|
||||
|
||||
+12963
File diff suppressed because one or more lines are too long
Generated
+2047
File diff suppressed because it is too large
Load Diff
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
BIN
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Vendored
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
BIN
Binary file not shown.
Binary file not shown.
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
BIN
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user