feat(func): Add rand(<token_length>)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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+
|
||||
|
||||
|
||||
@@ -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))
|
||||
|
||||
Reference in New Issue
Block a user