Files
mcDandy-more_math/more_math/AudioMathNode.py
T
2026-02-21 16:30:52 +01:00

254 lines
9.7 KiB
Python

from tokenize import String
from .helper_functions import (
generate_dim_variables,
parse_expr,
getIndexTensorAlongDim,
as_tensor,
normalize_to_common_shape,
make_zero_like,
get_v_variable,
get_f_variable
)
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
from .ParseTree import MrmthParseTree
class AudioMathNode(io.ComfyNode):
"""
Enables math expressions on Audio.
Inputs:
I: Autogrow image inputs (I0, I1, ...)
F: Autogrow float inputs (F0, F1, ...)
Image: Expression
"""
@classmethod
def define_schema(cls) -> io.Schema:
return io.Schema(
node_id="mrmth_ag_AudioMathNode",
category="More math",
display_name="Audio math",
inputs=[
io.Autogrow.Input(id="V",template=io.Autogrow.TemplatePrefix(io.Audio.Input("values", optional=True), prefix="V", min=1, max=50)),
io.Autogrow.Input(id="F", template=io.Autogrow.TemplatePrefix(io.Float.Input("float", default=0.0, optional=True, lazy=True, force_input=True), prefix="F", min=1, max=50)),
io.MultiType.Input(
io.String.Input("Expression", default="", multiline=False),
types=[io.String,MrmthParseTree],
tooltip="3D model file or path string",
),
io.Combo.Input(
id="length_mismatch",
options=["do nothing","error","tile", "pad"],
display_name="on size mismatch",
default="error",
tooltip="How to handle mismatched image batch sizes. tile: repeat shorter inputs; error: raise error on mismatch; pad: treat missing frames as zero."
),
io.Int.Input(id="batching", default=0),
MrmthStack.Input(id="stack", tooltip="Access stack between nodes",optional=True)
],
outputs=[
io.Audio.Output(is_output_list=True),
MrmthStack.Output(),
],
)
@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
# 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={}):
# Identify all present audio inputs and their keys
tensor_keys = [k for k, v in V.items() if v is not None and isinstance(v, dict) and "waveform" in v]
if not tensor_keys:
raise ValueError("At least one audio input is required.")
stack = copy.deepcopy(stack) if stack is not None else {}
waveforms = {k: V[k]["waveform"] for k in tensor_keys}
sample_rates = {k + "sr": V[k].get("sample_rate", 44100) for k in tensor_keys}
# Normalize all waveforms together
normalized_waveforms = normalize_to_common_shape(*waveforms.values(), mode=length_mismatch)
V_norm_waveforms = dict(zip(tensor_keys, normalized_waveforms))
ref_waveform = normalized_waveforms[0]
common_shape = ref_waveform.shape
sample_rate = V[tensor_keys[0]].get("sample_rate", 44100)
if(length_mismatch == "error"):
for name in tensor_keys:
if waveforms[name].shape != common_shape:
raise ValueError(f"Input '{name}' has shape ({waveforms[name].shape[0]}, {waveforms[name].shape[2]}), expected ({common_shape[0]}, {common_shape[2]}) to match input.")
# Setup legacy variables a, b, c, d
a_w = V_norm_waveforms.get("V0", make_zero_like(ref_waveform))
b_w = V_norm_waveforms.get("V1", make_zero_like(a_w))
c_w = V_norm_waveforms.get("V2", make_zero_like(a_w))
d_w = V_norm_waveforms.get("V3", make_zero_like(a_w))
# Ensure legacy are normalized
a_w, b_w, c_w, d_w = normalize_to_common_shape(a_w, b_w, c_w, d_w, mode=length_mismatch)
variables = {
"a": a_w, "b": b_w, "c": c_w, "d": d_w,
"w": F.get("F0", 0.0) if F.get("F0") is not None else 0.0,
"x": F.get("F1", 0.0) if F.get("F1") is not None else 0.0,
"y": F.get("F2", 0.0) if F.get("F2") is not None else 0.0,
"z": F.get("F3", 0.0) if F.get("F3") is not None else 0.0,
"B": getIndexTensorAlongDim(a_w, 0),
"C": getIndexTensorAlongDim(a_w, 1),
"channel": getIndexTensorAlongDim(a_w, 1),
"S": getIndexTensorAlongDim(a_w, 2),
"sample": getIndexTensorAlongDim(a_w, 2),
"R": sample_rate,
"sample_rate": sample_rate,
"batch": getIndexTensorAlongDim(a_w, 0),
"T": a_w.shape[0],
"batch_count": a_w.shape[0],
} | generate_dim_variables(a_w) | V_norm_waveforms | sample_rates
v_stacked, v_cnt = get_v_variable(V_norm_waveforms, length_mismatch=length_mismatch)
if v_stacked is not None:
variables["V"] = v_stacked
variables["Vcnt"] = float(v_cnt)
variables["V_count"] = float(v_cnt)
f_stacked, f_cnt = get_f_variable(F)
if f_stacked is not None:
variables["F"] = f_stacked
variables["Fcnt"] = float(f_cnt)
variables["F_count"] = float(f_cnt)
for k, val in F.items():
variables[k] = val if val is not None else 0.0
tree = None
if isinstance(Expression,str):
tree = parse_expr(Expression)
else:
tree = Expression
visitor = UnifiedMathVisitor(variables, a_w.shape,a_w.device,state_storage=stack)
result = visitor.visit(tree)
result = as_tensor(result, a_w.shape)
if batching and batching > 0:
res = torch.split(result, batching, dim=0)
res_list = []
for result_chunk in res:
res_list.append({"waveform": result_chunk, "sample_rate": sample_rate})
return (res_list, stack)
else:
return ([{"waveform": result, "sample_rate": sample_rate}], stack)