Fix audio math node and make it visible

This commit is contained in:
mcDandy
2026-01-22 15:21:22 +01:00
parent 947393c6cf
commit d3921204ae
3 changed files with 36 additions and 61 deletions
+29 -56
View File
@@ -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},)
+4 -4
View File
@@ -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)
+3 -1
View File
@@ -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
]