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, checkLazyNew ) from .Parser.UnifiedMathVisitor import UnifiedMathVisitor from comfy_api.latest import io import torch from .Stack import MrmthStack import copy from .ParseTree import MrmthParseTree class AudioMathNode(io.ComfyNode): """ Enables math expressions on Audio. Inputs: V: Autogrow audio inputs (V0, V1, ...) F: Autogrow float inputs (F0, F1, ...) Expression: 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="Expression to apply on weights", ), io.Combo.Input( id="length_mismatch", options=["do nothing","error","tile", "pad"], display_name="on size mismatch", default="error", tooltip="How to handle mismatched audio shapes. tile: repeat shorter inputs; error: raise error on mismatch; pad: treat missing samples as zero." ), io.Int.Input(id="batching", default=0), io.Boolean.Input( id="remember_stack", default=False, display_name="Remember stack across batch", tooltip=( "If enabled, stack is copied at output leading to changes being remembered during batch operations (node runs multiple times in sucession). If disabled each batch gets it's own copy of the stack." ), ), 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, remember_stack=False,stack={}): return checkLazyNew(Expression,V,F) @classmethod def execute(cls, V, F, Expression, length_mismatch="tile",batching=0, remember_stack=False, 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 = stack if remember_stack else (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), "batch": getIndexTensorAlongDim(a_w, 0), "C": getIndexTensorAlongDim(a_w, 1), "channel": getIndexTensorAlongDim(a_w, 1), "N": float(a_w.shape[1]), "channel_count": float(a_w.shape[1]), "S": getIndexTensorAlongDim(a_w, 2), "sample": getIndexTensorAlongDim(a_w, 2), "T": float(a_w.shape[2]), "sample_count": float(a_w.shape[2]), "R": sample_rate, "sample_rate": sample_rate, "batch_count": float(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: stack = stack if remember_stack else copy.deepcopy(stack) return ([{"waveform": result, "sample_rate": sample_rate}], stack)