From d3921204aebd2d3cceb6511b43aa1ac75bdca958 Mon Sep 17 00:00:00 2001 From: mcDandy Date: Thu, 22 Jan 2026 15:21:22 +0100 Subject: [PATCH] Fix audio math node and make it visible --- more_math/AudioMathNode.py | 85 ++++++++++++----------------------- more_math/helper_functions.py | 8 ++-- more_math/nodes.py | 4 +- 3 files changed, 36 insertions(+), 61 deletions(-) diff --git a/more_math/AudioMathNode.py b/more_math/AudioMathNode.py index 3cbb34b..0ee3ce1 100644 --- a/more_math/AudioMathNode.py +++ b/more_math/AudioMathNode.py @@ -40,9 +40,9 @@ class AudioMathNode(io.ComfyNode): ) @classmethod - def check_lazy_status(cls, Image, V, F, length_mismatch="tile"): + def check_lazy_status(cls, Expression, V, F, length_mismatch="tile"): - input_stream = InputStream(Image) + input_stream = InputStream(Expression) lexer = MathExprLexer(input_stream) stream = CommonTokenStream(lexer) stream.fill() @@ -73,77 +73,50 @@ class AudioMathNode(io.ComfyNode): @classmethod def execute(cls, V, F, Expression, length_mismatch="tile"): - # I and F are Autogrow.Type which is dict[str, Any] - - # Determine reference image for zero-initialization (fallback for a,b,c,d) - ref_image = None - for img in V.values(): - if img is not None: - ref_image = img - break - - if ref_image is None: - raise ValueError("At least one input is required.") - - aa = V.get("V0") - ba = V.get("V1") - ca = V.get("V2") - da = V.get("V3") - - if aa is None: - aa = make_zero_like(ref_image) - - ae, be, ce, de = prepare_inputs(aa, ba, ca, da) if(length_mismatch == "error"): - max_lengths = a.shape + max_lengths = V.get("V0")["waveform"].shape for name, tensor in V.items(): - if tensor is not None and max_lengths==tensor.shape: - raise ValueError(f"Input '{name}' has shape {tensor.shape}, expected {max_lengths} to match input.") + if tensor["waveform"] is not None and max_lengths!=tensor["waveform"].shape: + raise ValueError(f"Input '{name}' has shape {tensor["waveform"].shape}, expected {max_lengths} to match input.") - a, b, c, d = normalize_to_common_shape(ae["waweform"], be["waweform"], ce["waweform"], de, mode=length_mismatch) - waweforms={} + waveforms={} sample_rates={} for key, audio in V.items(): if audio is not None and isinstance(audio, dict) and "waveform" in audio: - waweforms[key] = audio["waveform"] - sample_rates[key] = audio.get("sample_rate", 44100) + waveforms[key] = audio["waveform"] + sample_rates[key+"sr"] = audio.get("sample_rate", 44100) else: - waweforms[key] = torch.zeros(V["V0"]["waveform"].shape) + waveforms[key] = torch.zeros(V["V0"]["waveform"].shape) sample_rates[key] = 44100 - sample_rate = sample_rates["V0"] if len(sample_rates) > 0 else 44100 - - + sample_rate = sample_rates["V0sr"] if len(sample_rates) > 0 else 44100 + new_values = normalize_to_common_shape(*waveforms.values(), mode=length_mismatch) + waveforms.update(zip(waveforms.keys(), new_values)) + a,b,c,d = prepare_inputs(V.get("V0"),V.get("V1"),V.get("V2"),V.get("V3")) + a = a["waveform"] + b = b["waveform"] + c = c["waveform"] + d = d["waveform"] variables = { "a": a, "b": b, "c": c, "d": d, "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(ae, 0), - "C": getIndexTensorAlongDim(ae, 1), - "channel": getIndexTensorAlongDim(ae, 1), - "S": getIndexTensorAlongDim(ae, 2), - "sample": getIndexTensorAlongDim(ae, 2), + "B": getIndexTensorAlongDim(a, 0), + "C": getIndexTensorAlongDim(a, 1), + "channel": getIndexTensorAlongDim(a, 1), + "S": getIndexTensorAlongDim(a, 2), + "sample": getIndexTensorAlongDim(a, 2), "R": sample_rate, "sample_rate": sample_rate, - "batch": getIndexTensorAlongDim(ae, 0), - "T": ae.shape[0], - "batch_count": ae.shape[0], - } | generate_dim_variables(ae) | waweforms | sample_rates - - # Add all dynamic inputs - for k, v in V.items(): - if v is not None: - # Normalize all images in I to match ae.shape - norm_v = normalize_to_common_shape(ae, v, mode=length_mismatch)[1] - variables[k] = norm_v - - for k, v in F.items(): - variables[k] = v if v is not None else 0.0 + "batch": getIndexTensorAlongDim(a, 0), + "T": a.shape[0], + "batch_count": a.shape[0], + } | generate_dim_variables(a) | waveforms | sample_rates tree = parse_expr(Expression); - visitor = UnifiedMathVisitor(variables, ae.shape) + visitor = UnifiedMathVisitor(variables, a.shape) result = visitor.visit(tree) - result = as_tensor(result, ae.shape) - return (result,) + result = as_tensor(result, a.shape) + return ({"waveform":result,"sample_rate":sample_rate},) diff --git a/more_math/helper_functions.py b/more_math/helper_functions.py index 73a8953..24feecb 100644 --- a/more_math/helper_functions.py +++ b/more_math/helper_functions.py @@ -16,7 +16,7 @@ class ThrowingErrorListener(ErrorListener): def as_tensor(value, shape): if isinstance(value, torch.Tensor) or getattr(value, "is_nested", False): - return value + return value.contiguous() if isinstance(value, (float, int)): value = (value,) return torch.broadcast_to(torch.Tensor(value), shape) @@ -155,7 +155,7 @@ def normalize_to_common_shape(*tensors, mode="pad"): def normalize_one(t, shape): if t.shape == shape: - return t + return t.contiguous() # Match ndim first curr_t = t @@ -181,12 +181,12 @@ def normalize_to_common_shape(*tensors, mode="pad"): out = torch.zeros(shape, dtype=curr_t.dtype, device=curr_t.device) slices = tuple(slice(0, d) for d in curr_t.shape) out[slices] = curr_t - return out + return out.contiguous() result = [] for t in tensors: if torch.is_tensor(t): result.append(normalize_one(t, target_shape)) else: - result.append(t) + result.append(t.contiguous()) return tuple(result) diff --git a/more_math/nodes.py b/more_math/nodes.py index 76a1302..2a801f0 100644 --- a/more_math/nodes.py +++ b/more_math/nodes.py @@ -15,6 +15,7 @@ from .SpectrogramToAudioNode import SpectrogramToAudio from .AudioToSpectrogramNode import AudioToSpectrogram from .NoiseMathNode import NoiseMathNode +from .AudioMathNode import AudioMathNode from comfy_api.latest import ComfyExtension, io @@ -93,7 +94,8 @@ class MoreMathExtension(ComfyExtension): SpectrogramToAudio, SigmasMathNode, GuiderMathNode, - NoiseMathNode + NoiseMathNode, + AudioMathNode ]