Compare commits
4
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
984d7f6f4d | ||
|
|
f625fb605c | ||
|
|
6af94de00a | ||
|
|
6237f0e503 |
+6
-2
@@ -1,9 +1,11 @@
|
||||
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.project import ProjectAction
|
||||
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
|
||||
@@ -21,4 +23,6 @@ register_action(SlerpAction)
|
||||
register_action(AverageAction)
|
||||
register_action(ScaleDims)
|
||||
register_action(SetDims)
|
||||
register_action(ProjectAction)
|
||||
register_action(MaxAction)
|
||||
register_action(AbsMaxAction)
|
||||
register_action(MinAction)
|
||||
|
||||
@@ -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
|
||||
@@ -1,6 +1,3 @@
|
||||
from typing import List
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
from torch.nn import Embedding
|
||||
|
||||
@@ -12,9 +9,3 @@ def get_embedding(seg_or_action: SegOrAction, embedding_module: Embedding) -> Te
|
||||
if isinstance(seg_or_action, Action):
|
||||
return seg_or_action.get_result(embedding_module)
|
||||
return seg_or_action.get_embeddings(embedding_module)
|
||||
|
||||
|
||||
def get_embedding_for_segments(
|
||||
segments: List[SegOrAction], embedding_module: Embedding
|
||||
) -> Tensor:
|
||||
return torch.cat([get_embedding(segment, embedding_module) for segment in segments], dim=1)
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -1,109 +0,0 @@
|
||||
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,
|
||||
get_embedding_for_segments,
|
||||
)
|
||||
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 ProjectAction(MultiArgAction):
|
||||
grammar = 'project(" arg "|" arg ("|" arg ")?)"'
|
||||
name = "project"
|
||||
chars = ["+", "+"]
|
||||
|
||||
weight = 1.0
|
||||
|
||||
def __init__(self, args: List[List[Union[PromptSegment, Action]]]) -> None:
|
||||
super().__init__(args)
|
||||
|
||||
num_args = len(args)
|
||||
if num_args != 3 and num_args != 2:
|
||||
raise ValueError(
|
||||
"Project action should have exactly three arguments(2 vectors and a weight)"
|
||||
)
|
||||
|
||||
self.source_argument = args[0]
|
||||
self.source_argument_token_length = sum(
|
||||
seg_or_action.token_length() for seg_or_action in self.source_argument
|
||||
)
|
||||
self.onto_argument = args[1]
|
||||
self.onto_argument_token_length = sum(
|
||||
seg_or_action.token_length() for seg_or_action in self.onto_argument
|
||||
)
|
||||
|
||||
if num_args == 3:
|
||||
self._parse_weight(args[2])
|
||||
|
||||
self._validate_args()
|
||||
|
||||
def _parse_weight(self, arg: List[SegOrAction]) -> None:
|
||||
if len(arg) != 1:
|
||||
raise ValueError("Project weight should have exactly one segment")
|
||||
|
||||
weight_seg_or_action = arg[0]
|
||||
|
||||
if isinstance(weight_seg_or_action, Action):
|
||||
raise ValueError("Project weight should not have an action as an argument")
|
||||
|
||||
try:
|
||||
self.weight = float(weight_seg_or_action.text)
|
||||
except ValueError:
|
||||
raise ValueError("Project should have an integer/float as the weight")
|
||||
|
||||
def _validate_args(self) -> None:
|
||||
if (
|
||||
self.source_argument_token_length != self.onto_argument_token_length
|
||||
and self.onto_argument_token_length != 1
|
||||
):
|
||||
raise ValueError(
|
||||
f"Project source and target arguments should have the same token lengths, or target should be one token. Got {self.source_argument_token_length} source tokens and {self.onto_argument_token_length} target tokens"
|
||||
)
|
||||
|
||||
def token_length(self) -> int:
|
||||
# Project projects the source onto the target, so the length of the result is the length of source
|
||||
return self.source_argument_token_length
|
||||
|
||||
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
|
||||
# Calculate the embeddings for the start segment
|
||||
source_embedding = get_embedding_for_segments(
|
||||
self.source_argument, embedding_module
|
||||
).to(dtype=torch.float32)
|
||||
onto_embedding = get_embedding_for_segments(
|
||||
self.onto_argument, embedding_module
|
||||
).to(dtype=torch.float32)
|
||||
|
||||
# Perform the projection
|
||||
return torch.mul(
|
||||
torch.mul(source_embedding, onto_embedding)
|
||||
/ torch.mul(onto_embedding, onto_embedding),
|
||||
onto_embedding,
|
||||
).to(dtype=torch.float16)
|
||||
|
||||
# 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,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
|
||||
|
||||
Reference in New Issue
Block a user