Author SHA1 Message Date
Michael Poutre 984d7f6f4d feat(func/min): Add function 2023-09-07 18:18:05 -07:00
Michael Poutre f625fb605c feat(func/absmax): Add function 2023-09-07 18:18:05 -07:00
Michael Poutre 6af94de00a feat(func/max): Add function 2023-09-06 18:33:25 -07:00
Michael Poutre 6237f0e503 feat(util): Add is_broadcastable 2023-09-06 18:22:55 -07:00
5 changed files with 271 additions and 0 deletions
+6
View File
@@ -1,5 +1,8 @@
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
@@ -20,3 +23,6 @@ register_action(SlerpAction)
register_action(AverageAction)
register_action(ScaleDims)
register_action(SetDims)
register_action(MaxAction)
register_action(AbsMaxAction)
register_action(MinAction)
+104
View File
@@ -0,0 +1,104 @@
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
@@ -0,0 +1,72 @@
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
@@ -0,0 +1,72 @@
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,3 +36,20 @@ 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