From 2f0d8569d15c8d08a496baeee0375c7fdd9f359f Mon Sep 17 00:00:00 2001 From: Michael Poutre Date: Mon, 11 Sep 2023 19:40:44 -0700 Subject: [PATCH] feat(func/projection): Add function --- lib/__init__.py | 2 + lib/actions/action_utils.py | 9 +++ lib/actions/project.py | 109 ++++++++++++++++++++++++++++++++++++ 3 files changed, 120 insertions(+) create mode 100644 lib/actions/project.py diff --git a/lib/__init__.py b/lib/__init__.py index 5e463a7..f15c563 100644 --- a/lib/__init__.py +++ b/lib/__init__.py @@ -3,6 +3,7 @@ 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.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 @@ -20,3 +21,4 @@ register_action(SlerpAction) register_action(AverageAction) register_action(ScaleDims) register_action(SetDims) +register_action(ProjectAction) diff --git a/lib/actions/action_utils.py b/lib/actions/action_utils.py index 3f676f2..8624cd6 100644 --- a/lib/actions/action_utils.py +++ b/lib/actions/action_utils.py @@ -1,3 +1,6 @@ +from typing import List + +import torch from torch import Tensor from torch.nn import Embedding @@ -9,3 +12,9 @@ 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) diff --git a/lib/actions/project.py b/lib/actions/project.py new file mode 100644 index 0000000..b823e9f --- /dev/null +++ b/lib/actions/project.py @@ -0,0 +1,109 @@ +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 + ) + onto_embedding = get_embedding_for_segments( + self.onto_argument, embedding_module + ) + + # Perform the projection + return torch.mul( + torch.mul(source_embedding, onto_embedding) + / torch.mul(onto_embedding, onto_embedding), + onto_embedding, + ) + + # 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