feat(func): Add rand(<token_length>)

This commit is contained in:
Michael Poutre
2023-08-30 20:29:46 -07:00
parent fee9c56abd
commit af761a5620
3 changed files with 49 additions and 0 deletions
+44
View File
@@ -0,0 +1,44 @@
import torch
from torch.nn import Embedding
from custom_nodes.KepPromptLang.lib.action.base import (
Action,
SingleArgAction,
)
class RandAction(SingleArgAction):
grammar = 'rand(" arg ")"'
name = "rand"
chars = None
parsed_token_length = 0
def __init__(self, arg):
super().__init__(arg)
if len(self.arg) != 1:
raise ValueError("Random action should have exactly one argument")
seg_or_action = self.arg[0]
if isinstance(seg_or_action, Action):
raise ValueError("Random action should not have an action as an argument")
try:
self.parsed_token_length = int(seg_or_action.text)
except ValueError:
raise ValueError("Random action should have an integer as an argument")
def token_length(self) -> int:
"""
Random returns a random embedding whose length is the number in the argument
:return:
"""
return self.parsed_token_length
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
# Create random tensor of size
result = torch.rand((1, self.parsed_token_length, embedding_module.embedding_dim))
return result
+2
View File
@@ -10,11 +10,13 @@ function: sum_function
| neg_function
| norm_function
| diff_function
| rand_function
sum_function: "sum(" arg ("|" arg)* ")"
neg_function: "neg(" arg ")"
norm_function: "norm(" arg ")"
diff_function: "diff(" arg ("|" arg)* ")"
rand_function: "rand(" arg ")"
arg: item+
+3
View File
@@ -5,6 +5,7 @@ from lark import Transformer, Token
from comfy.sd1_clip import SD1Tokenizer
from custom_nodes.KepPromptLang.lib.action.base import Action
from custom_nodes.KepPromptLang.lib.actions.diff import DiffAction
from custom_nodes.KepPromptLang.lib.actions.rand import RandAction
from custom_nodes.KepPromptLang.lib.parser.utils import build_prompt_segment
from custom_nodes.KepPromptLang.lib.actions.neg import NegAction
from custom_nodes.KepPromptLang.lib.actions.norm import NormAction
@@ -59,5 +60,7 @@ class PromptTransformer(Transformer):
return NormAction(item.children[0])
elif item.data == 'diff_function':
return DiffAction(item.children[0][:], item.children[1:][:])
elif item.data == 'rand_function':
return RandAction(item.children[0])
else:
raise Exception("Unknown function type: " + str(item.data))