diff --git a/src/more_math/helperFunc.py b/src/more_math/helperFunc.py index 3d8530f..41e447c 100644 --- a/src/more_math/helperFunc.py +++ b/src/more_math/helperFunc.py @@ -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) diff --git a/src/more_math/nodes.py b/src/more_math/nodes.py index d0ab2a3..c974dc5 100644 --- a/src/more_math/nodes.py +++ b/src/more_math/nodes.py @@ -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" } + + + diff --git a/src/more_math/parser.py b/src/more_math/parser.py index b38597c..a68ab20 100644 --- a/src/more_math/parser.py +++ b/src/more_math/parser.py @@ -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() diff --git a/tests/test_more_math.py b/tests/test_more_math.py index 0b4fb74..d602a05 100644 --- a/tests/test_more_math.py +++ b/tests/test_more_math.py @@ -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