From bc52d51900102d25bc0412c7aaa52ba9cf968102 Mon Sep 17 00:00:00 2001 From: mcDandy Date: Sat, 21 Feb 2026 16:43:33 +0100 Subject: [PATCH] move check to helper functions --- more_math/AudioMathNode.py | 110 +--------------------------------- more_math/helper_functions.py | 106 ++++++++++++++++++++++++++++++++ 2 files changed, 109 insertions(+), 107 deletions(-) diff --git a/more_math/AudioMathNode.py b/more_math/AudioMathNode.py index 583f7e5..2bd9077 100644 --- a/more_math/AudioMathNode.py +++ b/more_math/AudioMathNode.py @@ -7,14 +7,14 @@ from .helper_functions import ( normalize_to_common_shape, make_zero_like, get_v_variable, - get_f_variable + get_f_variable, + checkLazyNew ) from .Parser.UnifiedMathVisitor import UnifiedMathVisitor from comfy_api.latest import io from antlr4 import InputStream, CommonTokenStream from .Parser.MathExprLexer import MathExprLexer from .Parser.MathExprParser import MathExprParser -import re import torch from .Stack import MrmthStack import copy @@ -62,112 +62,8 @@ class AudioMathNode(io.ComfyNode): @classmethod def check_lazy_status(cls, Expression, V, F, length_mismatch="tile",batching=0,stack={}): - tree = None - parser = None - if isinstance(Expression,str): - tree = parse_expr(Expression) - parser = tree.parser - else: - tree = Expression - parser = tree.parser + return checkLazyNew(Expression,V,F) - # Support aliases - aliases = {"a": "V0", "b": "V1", "c": "V2", "d": "V3", - "w": "F0", "x": "F1", "y": "F2", "z": "F3"} - - assigned_vars = set() - needed_vars = set() - - # Process all top-level statements - for child in tree.children: - if not hasattr(child, 'getRuleIndex'): - continue - - rule_name = parser.ruleNames[child.getRuleIndex()] if child.getRuleIndex() < len(parser.ruleNames) else None - - # Process function definitions: scan for reads but ignore writes - if rule_name == 'funcDef': - func_params = set() - if child.paramList(): - for param in child.paramList().VARIABLE(): - func_params.add(param.getText()) - - cls._collect_reads_only(child, needed_vars, assigned_vars, func_params) - - # Top-level assignments - elif rule_name == 'varDef': - var_name = child.VARIABLE().getText() - cls._collect_vars_from_node(child, needed_vars, assigned_vars, set()) - assigned_vars.add(var_name) - - # Track other top-level statements - else: - cls._collect_vars_from_node(child, needed_vars, assigned_vars, set()) - - # Normalize variable names through aliases - needed = set() - for var in needed_vars: - norm = aliases.get(var, var) - if var == "V": - needed.update(V.keys()) - if var == "F": - needed.update(F.keys()) - if re.match(r"[VF][0-9]+", norm): - needed.add(norm) - - return needed - - @classmethod - def _collect_reads_only(cls, node, needed_vars, assigned_vars, shadowed_vars): - if node is None: - return - - node_type = type(node).__name__ - - if node_type == 'VariableExpContext': - var_name = node.VARIABLE().getText() - if var_name in shadowed_vars: - return - if var_name not in assigned_vars: - needed_vars.add(var_name) - return - - if node_type == 'FunctionDefContext': - return - - if node_type == 'VarDefContext': - for expr in node.expr(): - cls._collect_reads_only(expr, needed_vars, assigned_vars, shadowed_vars) - return - - for i in range(node.getChildCount()): - cls._collect_reads_only(node.getChild(i), needed_vars, assigned_vars, shadowed_vars) - - @classmethod - def _collect_vars_from_node(cls, node, needed_vars, assigned_vars, shadowed_vars): - """Recursively collect variable reads from an AST node""" - if node is None: - return - - node_type = type(node).__name__ - - # Found a variable read - if node_type == 'VariableExpContext': - var_name = node.VARIABLE().getText() - # Skip if shadowed - if var_name in shadowed_vars: - return - if var_name not in assigned_vars: - needed_vars.add(var_name) - return - - # Skip - if node_type == 'FunctionDefContext': - return - - # Recursively visit children - for i in range(node.getChildCount()): - cls._collect_vars_from_node(node.getChild(i), needed_vars, assigned_vars, shadowed_vars) @classmethod def execute(cls, V, F, Expression, length_mismatch="tile",batching=0,stack={}): diff --git a/more_math/helper_functions.py b/more_math/helper_functions.py index b60e23c..8721d77 100644 --- a/more_math/helper_functions.py +++ b/more_math/helper_functions.py @@ -6,6 +6,7 @@ import torch from .Parser.MathExprLexer import MathExprLexer from .Parser.MathExprParser import MathExprParser +import re class ThrowingErrorListener(ErrorListener): """Error listener that raises ValueError on syntax errors.""" @@ -234,3 +235,108 @@ def get_f_variable(f_dict): except Exception: return None, len(ordered_values) +def _collect_reads_only(cls, node, needed_vars, assigned_vars, shadowed_vars): + if node is None: + return + + node_type = type(node).__name__ + + if node_type == 'VariableExpContext': + var_name = node.VARIABLE().getText() + if var_name in shadowed_vars: + return + if var_name not in assigned_vars: + needed_vars.add(var_name) + return + + if node_type == 'FunctionDefContext': + return + + if node_type == 'VarDefContext': + for expr in node.expr(): + cls._collect_reads_only(expr, needed_vars, assigned_vars, shadowed_vars) + return + + for i in range(node.getChildCount()): + cls._collect_reads_only(node.getChild(i), needed_vars, assigned_vars, shadowed_vars) + +def _collect_vars_from_node(cls, node, needed_vars, assigned_vars, shadowed_vars): + """Recursively collect variable reads from an AST node""" + if node is None: + return + + node_type = type(node).__name__ + + # Found a variable read + if node_type == 'VariableExpContext': + var_name = node.VARIABLE().getText() + # Skip if shadowed + if var_name in shadowed_vars: + return + if var_name not in assigned_vars: + needed_vars.add(var_name) + return + + # Skip + if node_type == 'FunctionDefContext': + return + + # Recursively visit children + for i in range(node.getChildCount()): + cls._collect_vars_from_node(node.getChild(i), needed_vars, assigned_vars, shadowed_vars) + +def checkLazyNew(Expression, V, F): + tree = None + parser = None + if isinstance(Expression,str): + tree = parse_expr(Expression) + parser = tree.parser + else: + tree = Expression + parser = tree.parser + + # Support aliases + aliases = {"a": "V0", "b": "V1", "c": "V2", "d": "V3", + "w": "F0", "x": "F1", "y": "F2", "z": "F3"} + + assigned_vars = set() + needed_vars = set() + + # Process all top-level statements + for child in tree.children: + if not hasattr(child, 'getRuleIndex'): + continue + + rule_name = parser.ruleNames[child.getRuleIndex()] if child.getRuleIndex() < len(parser.ruleNames) else None + + # Process function definitions: scan for reads but ignore writes + if rule_name == 'funcDef': + func_params = set() + if child.paramList(): + for param in child.paramList().VARIABLE(): + func_params.add(param.getText()) + + _collect_reads_only(child, needed_vars, assigned_vars, func_params) + + # Top-level assignments + elif rule_name == 'varDef': + var_name = child.VARIABLE().getText() + _collect_vars_from_node(child, needed_vars, assigned_vars, set()) + assigned_vars.add(var_name) + + # Track other top-level statements + else: + _collect_vars_from_node(child, needed_vars, assigned_vars, set()) + + # Normalize variable names through aliases + needed = set() + for var in needed_vars: + norm = aliases.get(var, var) + if var == "V": + needed.update(V.keys()) + if var == "F": + needed.update(F.keys()) + if re.match(r"[VF][0-9]+", norm): + needed.add(norm) + + return needed \ No newline at end of file