added more operators. Running to limits of tokenizer.
This commit is contained in:
+63
-14
@@ -3,10 +3,11 @@ from io import StringIO
|
||||
import tokenize
|
||||
|
||||
def to_tensor(val, ref):
|
||||
if isinstance(val, torch.Tensor):
|
||||
return val
|
||||
# Convert float/int, to tensor with the same shape as ref
|
||||
return torch.full_like(ref, float(val))
|
||||
if isinstance(val, torch.Tensor):
|
||||
return val
|
||||
# Convert float/int, to tensor with the same shape as ref
|
||||
return torch.full_like(ref, float(val))
|
||||
|
||||
def evaluate_tensor_expression(ast, variables):
|
||||
"""
|
||||
Evaluates AST using variables.
|
||||
@@ -14,7 +15,10 @@ def evaluate_tensor_expression(ast, variables):
|
||||
"""
|
||||
if not isinstance(ast, tuple):
|
||||
return ast # Return literal values directly (if any)
|
||||
|
||||
variables['w'] = to_tensor(variables['w'], variables['a'])
|
||||
variables['x'] = to_tensor(variables['x'], variables['a'])
|
||||
variables['y'] = to_tensor(variables['y'], variables['a'])
|
||||
variables['z'] = to_tensor(variables['z'], variables['a'])
|
||||
node_type = ast[0]
|
||||
|
||||
if node_type == 'VARIABLE':
|
||||
@@ -30,9 +34,12 @@ def evaluate_tensor_expression(ast, variables):
|
||||
op, operand_expr = ast[1]
|
||||
operand_val = evaluate_tensor_expression(operand_expr, variables)
|
||||
if op == '+':
|
||||
return to_tensor(operand_val,variables['a'])
|
||||
return to_tensor(operand_val, variables['a'])
|
||||
elif op == '-':
|
||||
return torch.neg(to_tensor(operand_val,variables['a']))
|
||||
return torch.neg(to_tensor(operand_val, variables['a']))
|
||||
elif op == '!':
|
||||
# Logická negace: !a
|
||||
return torch.logical_not(to_tensor(operand_val, variables['a']).bool())
|
||||
else:
|
||||
raise ValueError(f"Unsupported unary operator: {op}")
|
||||
|
||||
@@ -45,19 +52,55 @@ def evaluate_tensor_expression(ast, variables):
|
||||
if func_name == 'ABS':
|
||||
return torch.abs(arg_vals[0])
|
||||
elif func_name == 'NORM':
|
||||
return torch.nn.functional.normalize(to_tensor(arg_vals[0],variables['a']), p=2, dim=-1)
|
||||
return torch.nn.functional.normalize(to_tensor(arg_vals[0], variables['a']), p=2, dim=-1)
|
||||
elif func_name == 'SQRT':
|
||||
return torch.sqrt(to_tensor(arg_vals[0],variables['a']))
|
||||
return torch.sqrt(arg_vals[0])
|
||||
elif func_name == 'SIN':
|
||||
return torch.sin(to_tensor(arg_vals[0],variables['a']))
|
||||
return torch.sin(arg_vals[0])
|
||||
elif func_name == 'COS':
|
||||
return torch.cos(to_tensor(arg_vals[0],variables['a']))
|
||||
return torch.cos(arg_vals[0])
|
||||
elif func_name == 'TAN':
|
||||
return torch.tan(to_tensor(arg_vals[0],variables['a']))
|
||||
return torch.tan(arg_vals[0])
|
||||
elif func_name == 'MAX':
|
||||
return torch.max(*arg_vals)
|
||||
elif func_name == 'MIN':
|
||||
return torch.max(*arg_vals)
|
||||
elif func_name == 'FLOOR':
|
||||
return torch.floor(arg_vals[0])
|
||||
elif func_name == 'CEIL':
|
||||
return torch.ceil(arg_vals[0])
|
||||
elif func_name == 'ROUND':
|
||||
return torch.round(arg_vals[0])
|
||||
elif func_name == 'ASIN':
|
||||
return torch.asin(arg_vals[0])
|
||||
elif func_name == 'ACOS':
|
||||
return torch.acos(arg_vals[0])
|
||||
elif func_name == 'ATAN':
|
||||
return torch.atan(arg_vals[0])
|
||||
elif func_name == 'ATAN2':
|
||||
return torch.atan2(arg_vals[1], arg_vals[0])
|
||||
elif func_name == 'LN':
|
||||
return torch.log(arg_vals[0])
|
||||
elif func_name == 'SINH':
|
||||
return torch.sinh(arg_vals[0])
|
||||
elif func_name == 'COSH':
|
||||
return torch.cosh(arg_vals[0])
|
||||
elif func_name == 'TANH':
|
||||
return torch.tanh(arg_vals[0])
|
||||
elif func_name == 'ASINH':
|
||||
return torch.asinh(arg_vals[0])
|
||||
elif func_name == 'ACOSH':
|
||||
return torch.acosh(arg_vals[0])
|
||||
elif func_name == 'ATANH':
|
||||
return torch.atanh(arg_vals[0])
|
||||
elif func_name == 'EXP':
|
||||
return torch.exp(arg_vals[0])
|
||||
elif func_name == 'LOG':
|
||||
return torch.log10(arg_vals[0])
|
||||
elif func_name == 'GAMMA':
|
||||
return torch.gamma(arg_vals[0])
|
||||
elif func_name == 'XOR':
|
||||
return torch.logical_xor(arg_vals[0].bool(), arg_vals[1].bool())
|
||||
else:
|
||||
raise ValueError(f"Unsupported function: {func_name}")
|
||||
|
||||
@@ -73,10 +116,15 @@ def evaluate_tensor_expression(ast, variables):
|
||||
return torch.multiply(left_val, right_val)
|
||||
elif op == '/':
|
||||
return torch.divide(left_val, right_val)
|
||||
elif op == '^':
|
||||
return torch.pow(to_tensor(left_val,variables['a']).type(torch.complex64), to_tensor(right_val,variables['a']).type(torch.complex64)).real
|
||||
elif op == '%':
|
||||
return torch.fmod(left_val, right_val)
|
||||
elif op == '&':
|
||||
# Logický AND
|
||||
return torch.logical_and(to_tensor(left_val, variables['a']).bool(), to_tensor(right_val, variables['a']).bool())
|
||||
elif op == '|':
|
||||
return torch.logical_or(to_tensor(left_val, variables['a']).bool(), to_tensor(right_val, variables['a']).bool())
|
||||
elif op == '^':
|
||||
return torch.pow(to_tensor(left_val, variables['a']).type(torch.complex64), to_tensor(right_val, variables['a']).type(torch.complex64)).real
|
||||
else:
|
||||
raise ValueError(f"Unsupported operator: {op}")
|
||||
|
||||
@@ -86,6 +134,7 @@ def evaluate_tensor_expression(ast, variables):
|
||||
|
||||
else:
|
||||
raise ValueError(f"Unknown AST node type: {node_type}")
|
||||
|
||||
def tokenize_expression(expr):
|
||||
f = StringIO(expr)
|
||||
tokens_gen = tokenize.generate_tokens(f.readline)
|
||||
|
||||
+46
-1
@@ -2,18 +2,63 @@ from .ConditioningMathNode import ConditioningMathNode
|
||||
from .LatentMathNode import LatentMathNode
|
||||
from .ImageMathNode import ImageMathNode
|
||||
|
||||
class IntToFloatNode:
|
||||
"""
|
||||
Converts int to float.
|
||||
"""
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"value": ("INT", {"default": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("FLOAT",)
|
||||
FUNCTION = "convert"
|
||||
CATEGORY = "More math"
|
||||
|
||||
def convert(self, value):
|
||||
return (float(value),)
|
||||
|
||||
class FloatToIntNode:
|
||||
"""
|
||||
Converts float to int.
|
||||
"""
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"value": ("FLOAT", {"default": 0.0}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("INT",)
|
||||
FUNCTION = "convert"
|
||||
CATEGORY = "More math"
|
||||
|
||||
def convert(self, value):
|
||||
return (int(value),)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"mrmth_ConditioningMathNode": ConditioningMathNode,
|
||||
"mrmth_LatentMathNode": LatentMathNode,
|
||||
"mrmth_ImageMathNode": ImageMathNode
|
||||
"mrmth_ImageMathNode": ImageMathNode,
|
||||
"mrmth_IntToFloat": IntToFloatNode,
|
||||
"mrmth_FloatToInt": FloatToIntNode,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"mrmth_ConditioningMathNode": "Conditioning math node",
|
||||
"mrmth_LatentMathNode": "Latent math node",
|
||||
"mrmth_ImageMathNode": "Image math node",
|
||||
"mrmth_IntToFloat": "Int → Float",
|
||||
"mrmth_FloatToInt": "Float → Int",
|
||||
"Tensor": "Tensor expression",
|
||||
"Latent": "Latent expression",
|
||||
"Image": "Image expression",
|
||||
"pooled_output": "Pooled output tensor expression"
|
||||
}
|
||||
|
||||
|
||||
|
||||
|
||||
+59
-27
@@ -4,9 +4,9 @@ class Parser:
|
||||
self.pos = 0
|
||||
|
||||
def peek(self):
|
||||
# Skip INDENT/DEDENT/NEWLINE tokens
|
||||
# Přeskoč NEWLINE/INDENT/DEDENT
|
||||
while self.pos < len(self.tokens) and self.tokens[self.pos][0] in {"INDENT", "DEDENT", "NEWLINE"}:
|
||||
self.posF += 1
|
||||
self.pos += 1
|
||||
return self.tokens[self.pos] if self.pos < len(self.tokens) else ('ENDMARKER', '')
|
||||
|
||||
def consume(self):
|
||||
@@ -14,44 +14,83 @@ class Parser:
|
||||
self.pos += 1
|
||||
return token
|
||||
|
||||
def expect(self, expected_type=None, expected_value=None):
|
||||
tok_type, value = self.peek()
|
||||
if expected_type and tok_type != expected_type:
|
||||
raise SyntaxError(
|
||||
f"Expected token type {expected_type}, got {tok_type} ('{value}') at position {self.pos}."
|
||||
)
|
||||
if expected_value and value != expected_value:
|
||||
raise SyntaxError(
|
||||
f"Expected token value '{expected_value}', got '{value}' at position {self.pos}."
|
||||
)
|
||||
return self.consume()
|
||||
|
||||
def parse_expression(self):
|
||||
left = self.parse_term()
|
||||
# Logický OR
|
||||
left = self.parse_xor()
|
||||
while True:
|
||||
tok_type, value = self.peek()
|
||||
if value == '|':
|
||||
self.consume()
|
||||
right = self.parse_xor()
|
||||
left = ('BINOP', ('|', left, right))
|
||||
else:
|
||||
break
|
||||
return left
|
||||
|
||||
def parse_xor(self):
|
||||
# XOR
|
||||
left = self.parse_and()
|
||||
while True:
|
||||
tok_type, value = self.peek()
|
||||
if value == '^':
|
||||
self.consume()
|
||||
right = self.parse_and()
|
||||
left = ('BINOP', ('^', left, right))
|
||||
else:
|
||||
break
|
||||
return left
|
||||
|
||||
def parse_and(self):
|
||||
# Logický AND
|
||||
left = self.parse_add_sub()
|
||||
while True:
|
||||
tok_type, value = self.peek()
|
||||
if value == '&':
|
||||
self.consume()
|
||||
right = self.parse_add_sub()
|
||||
left = ('BINOP', ('&', left, right))
|
||||
else:
|
||||
break
|
||||
return left
|
||||
|
||||
def parse_add_sub(self):
|
||||
# Sčítání, odčítání
|
||||
left = self.parse_mul_div_mod()
|
||||
while True:
|
||||
tok_type, value = self.peek()
|
||||
if value in ['+', '-']:
|
||||
op = value
|
||||
self.consume()
|
||||
right = self.parse_term()
|
||||
right = self.parse_mul_div_mod()
|
||||
left = ('BINOP', (op, left, right))
|
||||
else:
|
||||
break
|
||||
return left
|
||||
|
||||
def parse_term(self):
|
||||
left = self.parse_factor()
|
||||
def parse_mul_div_mod(self):
|
||||
# Násobení, dělení, modulo
|
||||
left = self.parse_unary()
|
||||
while True:
|
||||
tok_type, value = self.peek()
|
||||
if value in ['*', '/', '^', '%']:
|
||||
if value in ['*', '/', '%']:
|
||||
op = value
|
||||
self.consume()
|
||||
right = self.parse_factor()
|
||||
right = self.parse_unary()
|
||||
left = ('BINOP', (op, left, right))
|
||||
else:
|
||||
break
|
||||
return left
|
||||
|
||||
def parse_unary(self):
|
||||
# Unární operátory: !, +, -
|
||||
tok_type, value = self.peek()
|
||||
if tok_type == 'OP' and value in ('+', '-', '!'):
|
||||
op = value
|
||||
self.consume()
|
||||
operand = self.parse_unary()
|
||||
return ('UNARYOP', (op, operand))
|
||||
return self.parse_factor()
|
||||
|
||||
def parse_arguments(self):
|
||||
args = []
|
||||
if self.peek()[0] == 'OP' and self.peek()[1] == ')':
|
||||
@@ -69,13 +108,6 @@ class Parser:
|
||||
def parse_factor(self):
|
||||
tok_type, value = self.peek()
|
||||
|
||||
# Podpora unárních operátorů + a -
|
||||
if tok_type == 'OP' and value in ('+', '-'):
|
||||
op = value
|
||||
self.consume()
|
||||
operand = self.parse_factor()
|
||||
return ('UNARYOP', (op, operand))
|
||||
|
||||
if tok_type == 'OP' and value == '(':
|
||||
self.consume() # Skip '('
|
||||
expr = self.parse_expression()
|
||||
|
||||
+138
-21
@@ -7,6 +7,24 @@ from src.more_math.ConditioningMathNode import ConditioningMathNode
|
||||
from src.more_math.LatentMathNode import LatentMathNode
|
||||
from src.more_math.ImageMathNode import ImageMathNode
|
||||
from src.more_math.parser import Parser
|
||||
import tokenize
|
||||
from io import StringIO
|
||||
|
||||
def tokenize_expression(expr):
|
||||
if not expr.endswith('\n'):
|
||||
expr = expr + '\n'
|
||||
f = StringIO(expr)
|
||||
tokens = tokenize.generate_tokens(f.readline)
|
||||
filtered_tokens = []
|
||||
for toktype, tokval, _, _, _ in tokens:
|
||||
token_name = tokenize.tok_name[toktype]
|
||||
# Opravíme ERRORTOKEN na OP pro !, &, |, ^
|
||||
if token_name == 'ERRORTOKEN' and tokval in {'!', '&', '|', '^'}:
|
||||
token_name = 'OP'
|
||||
if token_name in {'COMMENT', 'NL', 'NEWLINE', 'INDENT', 'DEDENT'}:
|
||||
continue
|
||||
filtered_tokens.append((token_name, tokval.strip()))
|
||||
return filtered_tokens
|
||||
|
||||
def test_conditioning_math_node_initialization():
|
||||
node = ConditioningMathNode()
|
||||
@@ -35,27 +53,6 @@ def test_image_math_node_metadata():
|
||||
assert ImageMathNode.FUNCTION == "imgMathNode"
|
||||
assert ImageMathNode.CATEGORY == "More math"
|
||||
|
||||
import tokenize
|
||||
from io import StringIO
|
||||
|
||||
def tokenize_expression(expr):
|
||||
# Always add a trailing newline to avoid TokenError on incomplete input
|
||||
if not expr.endswith('\n'):
|
||||
expr = expr + '\n'
|
||||
f = StringIO(expr)
|
||||
try:
|
||||
tokens = tokenize.generate_tokens(f.readline)
|
||||
filtered_tokens = []
|
||||
for toktype, tokval, _, _, _ in tokens:
|
||||
token_name = tokenize.tok_name[toktype]
|
||||
if token_name in {'COMMENT', 'NL', 'NEWLINE', 'INDENT', 'DEDENT'}:
|
||||
continue
|
||||
filtered_tokens.append((token_name, tokval.strip()))
|
||||
return filtered_tokens
|
||||
except tokenize.TokenError as e:
|
||||
# Return a special token to indicate tokenization error for incomplete input
|
||||
return [('TOKENIZE_ERROR', str(e))]
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"expr,expected_ast",
|
||||
[
|
||||
@@ -101,3 +98,123 @@ def test_parser_errors(expr, err_msg):
|
||||
with pytest.raises(SyntaxError) as excinfo:
|
||||
parser.parse_expression()
|
||||
assert err_msg in str(excinfo.value)
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"expr,expected_ast",
|
||||
[
|
||||
("a ^ b", ('BINOP', ('^', ('VARIABLE', 'a'), ('VARIABLE', 'b')))),
|
||||
("a & b", ('BINOP', ('&', ('VARIABLE', 'a'), ('VARIABLE', 'b')))),
|
||||
("a | b", ('BINOP', ('|', ('VARIABLE', 'a'), ('VARIABLE', 'b')))),
|
||||
("!a", ('UNARYOP', ('!', ('VARIABLE', 'a')))),
|
||||
("!a & b", ('BINOP', ('&', ('UNARYOP', ('!', ('VARIABLE', 'a'))), ('VARIABLE', 'b')))),
|
||||
("a | b & c", ('BINOP', ('|', ('VARIABLE', 'a'), ('BINOP', ('&', ('VARIABLE', 'b'), ('VARIABLE', 'c')))))),
|
||||
("a ^ b | c", ('BINOP', ('|', ('BINOP', ('^', ('VARIABLE', 'a'), ('VARIABLE', 'b'))), ('VARIABLE', 'c')))),
|
||||
]
|
||||
)
|
||||
def test_parser_logical_and_pow(expr, expected_ast):
|
||||
tokens = tokenize_expression(expr)
|
||||
parser = Parser(tokens)
|
||||
ast = parser.parse_expression()
|
||||
assert ast == expected_ast
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"expr,expected_ast",
|
||||
[
|
||||
# Priority: ! > * > + > & > ^ > |
|
||||
("a + b * c",
|
||||
('BINOP', ('+',
|
||||
('VARIABLE', 'a'),
|
||||
('BINOP', ('*', ('VARIABLE', 'b'), ('VARIABLE', 'c')))
|
||||
))
|
||||
),
|
||||
("a * b + c",
|
||||
('BINOP', ('+',
|
||||
('BINOP', ('*', ('VARIABLE', 'a'), ('VARIABLE', 'b'))),
|
||||
('VARIABLE', 'c')
|
||||
))
|
||||
),
|
||||
("a + b & c",
|
||||
('BINOP', ('&',
|
||||
('BINOP', ('+', ('VARIABLE', 'a'), ('VARIABLE', 'b'))),
|
||||
('VARIABLE', 'c')
|
||||
))
|
||||
),
|
||||
("a & b + c",
|
||||
('BINOP', ('&',
|
||||
('VARIABLE', 'a'),
|
||||
('BINOP', ('+', ('VARIABLE', 'b'), ('VARIABLE', 'c')))
|
||||
))
|
||||
),
|
||||
("a + b | c",
|
||||
('BINOP', ('|',
|
||||
('BINOP', ('+', ('VARIABLE', 'a'), ('VARIABLE', 'b'))),
|
||||
('VARIABLE', 'c')
|
||||
))
|
||||
),
|
||||
("a | b + c",
|
||||
('BINOP', ('|',
|
||||
('VARIABLE', 'a'),
|
||||
('BINOP', ('+', ('VARIABLE', 'b'), ('VARIABLE', 'c')))
|
||||
))
|
||||
),
|
||||
("a * b & c",
|
||||
('BINOP', ('&',
|
||||
('BINOP', ('*', ('VARIABLE', 'a'), ('VARIABLE', 'b'))),
|
||||
('VARIABLE', 'c')
|
||||
))
|
||||
),
|
||||
("a & b * c",
|
||||
('BINOP', ('&',
|
||||
('VARIABLE', 'a'),
|
||||
('BINOP', ('*', ('VARIABLE', 'b'), ('VARIABLE', 'c')))
|
||||
))
|
||||
),
|
||||
("a ^ b + c",
|
||||
('BINOP', ('^',
|
||||
('VARIABLE', 'a'),
|
||||
('BINOP', ('+', ('VARIABLE', 'b'), ('VARIABLE', 'c')))
|
||||
))
|
||||
),
|
||||
("a + b ^ c",
|
||||
('BINOP', ('^',
|
||||
('BINOP', ('+', ('VARIABLE', 'a'), ('VARIABLE', 'b'))),
|
||||
('VARIABLE', 'c')
|
||||
))
|
||||
),
|
||||
("a | b ^ c",
|
||||
('BINOP', ('|',
|
||||
('VARIABLE', 'a'),
|
||||
('BINOP', ('^', ('VARIABLE', 'b'), ('VARIABLE', 'c')))
|
||||
))
|
||||
),
|
||||
("a ^ b | c",
|
||||
('BINOP', ('|',
|
||||
('BINOP', ('^', ('VARIABLE', 'a'), ('VARIABLE', 'b'))),
|
||||
('VARIABLE', 'c')
|
||||
))
|
||||
),
|
||||
("!a * b",
|
||||
('BINOP', ('*',
|
||||
('UNARYOP', ('!', ('VARIABLE', 'a'))),
|
||||
('VARIABLE', 'b')
|
||||
))
|
||||
),
|
||||
("!(a + b) * c",
|
||||
('BINOP', ('*',
|
||||
('UNARYOP', ('!', ('PARENTHESIS', ('BINOP', ('+', ('VARIABLE', 'a'), ('VARIABLE', 'b')))))),
|
||||
('VARIABLE', 'c')
|
||||
))
|
||||
),
|
||||
("a + (b | c)",
|
||||
('BINOP', ('+',
|
||||
('VARIABLE', 'a'),
|
||||
('PARENTHESIS', ('BINOP', ('|', ('VARIABLE', 'b'), ('VARIABLE', 'c'))))
|
||||
))
|
||||
),
|
||||
]
|
||||
)
|
||||
def test_parser_operator_priority(expr, expected_ast):
|
||||
tokens = tokenize_expression(expr)
|
||||
parser = Parser(tokens)
|
||||
ast = parser.parse_expression()
|
||||
assert ast == expected_ast
|
||||
|
||||
Reference in New Issue
Block a user