diff --git a/lib/actions/rand.py b/lib/actions/rand.py new file mode 100644 index 0000000..b063c55 --- /dev/null +++ b/lib/actions/rand.py @@ -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 + diff --git a/lib/parser/grammar.py b/lib/parser/grammar.py index c6d794f..15ff36e 100644 --- a/lib/parser/grammar.py +++ b/lib/parser/grammar.py @@ -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+ diff --git a/lib/parser/transformer.py b/lib/parser/transformer.py index 900906d..707f928 100644 --- a/lib/parser/transformer.py +++ b/lib/parser/transformer.py @@ -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))