move check to helper functions

This commit is contained in:
mcDandy
2026-02-21 16:43:33 +01:00
parent 4bc2b67a85
commit bc52d51900
2 changed files with 109 additions and 107 deletions
+3 -107
View File
@@ -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={}):
+106
View File
@@ -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