Compare commits
4
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"])
|
||||
|
||||
@@ -3,7 +3,6 @@ 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
|
||||
@@ -21,4 +20,3 @@ register_action(SlerpAction)
|
||||
register_action(AverageAction)
|
||||
register_action(ScaleDims)
|
||||
register_action(SetDims)
|
||||
register_action(ProjectAction)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
@@ -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.
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