From 4ac582c0452bce6a6bffe42ed11918ff93f889df Mon Sep 17 00:00:00 2001 From: mcDandy Date: Wed, 28 Jan 2026 17:53:56 +0100 Subject: [PATCH] fix + node conversion Fix conditioning not knowing about Float inputs; Convert VideoMathNode to new format --- more_math/ConditioningMathNode.py | 5 +- more_math/ImageMathNode.py | 8 +- more_math/ModelMathNode.py | 20 -- more_math/Parser/UnifiedMathVisitor.py | 2 +- more_math/VideoMathNode.py | 270 ++++++++++++++----------- more_math/deprecated/VideoMathNode.py | 158 +++++++++++++++ more_math/modelLikeCommon.py | 13 +- more_math/nodes.py | 2 + 8 files changed, 325 insertions(+), 153 deletions(-) create mode 100644 more_math/deprecated/VideoMathNode.py diff --git a/more_math/ConditioningMathNode.py b/more_math/ConditioningMathNode.py index 8627ac6..cbe0cf0 100644 --- a/more_math/ConditioningMathNode.py +++ b/more_math/ConditioningMathNode.py @@ -48,7 +48,7 @@ class ConditioningMathNode(io.ComfyNode): stream = CommonTokenStream(lexer) stream.fill() - input_stream = InputStream(Expression) + input_stream = InputStream(Expression_pi) lexer = MathExprLexer(input_stream) stream1 = CommonTokenStream(lexer) stream1.fill() @@ -131,6 +131,9 @@ class ConditioningMathNode(io.ComfyNode): "batch_count": a.shape[0], } | generate_dim_variables(a) | V_norm_tensors + for k, val in F.items(): + variables[k] = val if val is not None else 0.0 + # Execute Expression (Main Tensor) tree = parse_expr(Expression) visitor = UnifiedMathVisitor(variables, a.shape, state_storage=ss) diff --git a/more_math/ImageMathNode.py b/more_math/ImageMathNode.py index 5eff658..651e6c1 100644 --- a/more_math/ImageMathNode.py +++ b/more_math/ImageMathNode.py @@ -19,11 +19,7 @@ class ImageMathNode(io.ComfyNode): @classmethod def define_schema(cls) -> io.Schema: return io.Schema( - node_id="mrmth_ag_ImageMathNode", # New ID to avoid collision if necessary, or keep standard and user migrates? User asked to "switch", likely implies replacing functionality but maybe keeping ID? Usually replacing ID breaks workflows. - # Strategy: Use a NEW ID for the autogrow version if we want to allow side-by-side, but typically "Autogrow switch" implies replacing the main node. - # However, standard ComfyUI practice for breaking changes is often a new node or careful migration. - # Looking at AudioMathNode in step 6, it used "mrmth_ag_AudioMathNode". - # I will follow that pattern: mrmth_ag_ImageMathNode. + node_id="mrmth_ag_ImageMathNode", category="More math", display_name="Image math", inputs=[ @@ -75,7 +71,7 @@ class ImageMathNode(io.ComfyNode): return needed1 @classmethod - def execute(cls, V, F, Expression, length_mismatch="tile"): + def execute(cls, V, F, Expression, length_mismatch="error"): # I and F are Autogrow.Type which is dict[str, Any] # Identify all present tensors and their keys diff --git a/more_math/ModelMathNode.py b/more_math/ModelMathNode.py index 8e1fd23..15f2578 100644 --- a/more_math/ModelMathNode.py +++ b/more_math/ModelMathNode.py @@ -81,29 +81,9 @@ class ModelMathNode(io.ComfyNode): if a is None: raise ValueError("At least one input model is required.") - # Prepare variables - # V0..V3 map to a..d for backward compatibility in calculate_patches - - # Note: calculate_patches usually takes specific args. We might need to update it to support dynamic V/F or just pass everything. - # Looking at Step 34, calculate_patches signature: (Model, a, b, c, d, w, x, y, z) - # We need to verify if calculate_patches handles V/F. It probably doesn't. - # We should check 'modelLikeCommon.py' to see if update is needed. - - # Assume for now we pass a,b,c,d,w,x,y,z as standard. - # But for full autogrow support (more than 4 inputs), calculate_patches needs update. - # The prompt didn't explicitly ask to update modelLikeCommon, but "switch to Autogrow" implies full functionality. - # I'll check modelLikeCommon.py after this block. - # For now, I will pass V and F to calculate_patches if I modify it, or I will stick to legacy args if I don't modify it. - # However, to support V4+, I MUST modify calculate_patches. - - # Let's pass the V and F dicts to a modified calculate_patches, or overload it. - # I will update modelLikeCommon.py as part of this task. from .modelLikeCommon import calculate_patches_autogrow - # Map inputs to patchers if needed (Model.Input gives Model wrapper, need state_dict source?) - # ModelMathNode inputs are Model wrappers (comfy.model_patcher.ModelPatcher). - # So V items are ready to be used. aliases = {"a": "V0", "b": "V1", "c": "V2", "d": "V3", "w": "F0", "x": "F1", "y": "F2", "z": "F3"} patches = calculate_patches_autogrow(Expression, V=V, F=F, mapping=aliases) diff --git a/more_math/Parser/UnifiedMathVisitor.py b/more_math/Parser/UnifiedMathVisitor.py index ff2f4fc..4ae845d 100644 --- a/more_math/Parser/UnifiedMathVisitor.py +++ b/more_math/Parser/UnifiedMathVisitor.py @@ -162,7 +162,7 @@ class UnifiedMathVisitor(MathExprVisitor): return float(self.depth) if var_name in self.variables: return self.variables[var_name] - raise ValueError(f"Variable '{var_name}' not found") + raise ValueError(f"line {ctx.VARIABLE().getPayload().line}:{ctx.VARIABLE().getPayload().column}: Variable '{var_name}' not found") def visitListExp(self, ctx): res = [] diff --git a/more_math/VideoMathNode.py b/more_math/VideoMathNode.py index ae609ba..ab00235 100644 --- a/more_math/VideoMathNode.py +++ b/more_math/VideoMathNode.py @@ -1,158 +1,198 @@ import torch -import math -from .helper_functions import ( - generate_dim_variables, - getIndexTensorAlongDim, - parse_expr, - commonLazy, - normalize_to_common_shape, - as_tensor, - prepare_inputs -) - +from .helper_functions import generate_dim_variables, parse_expr, getIndexTensorAlongDim, as_tensor, normalize_to_common_shape, make_zero_like from .Parser.UnifiedMathVisitor import UnifiedMathVisitor from comfy_api.latest import io -from comfy_api.latest import VideoComponents, VideoFromComponents - +from antlr4 import InputStream, CommonTokenStream +from .Parser.MathExprLexer import MathExprLexer +from .Parser.MathExprParser import MathExprParser +import re class VideoMathNode(io.ComfyNode): """ - Enables math expressions on Videos (Images + Audio). - Video object is expected to have 'images' and 'audio' components. - Expressions for 'Images' and 'Audio' are handled separately. + 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_VideoMathNode", + node_id="mrmth_ag_VideoMathNode", category="More math", display_name="Video math", inputs=[ - io.Video.Input(id="a"), - io.Video.Input(id="b", optional=True, lazy=True), - io.Video.Input(id="c", optional=True, lazy=True), - io.Video.Input(id="d", optional=True, lazy=True), - io.Float.Input(id="w", default=0.0, optional=True, lazy=True, force_input=True), - io.Float.Input(id="x", default=0.0, optional=True, lazy=True, force_input=True), - io.Float.Input(id="y", default=0.0, optional=True, lazy=True, force_input=True), - io.Float.Input(id="z", default=0.0, optional=True, lazy=True, force_input=True), - io.String.Input(id="Images", default="a*(1-w)+b*w", tooltip="Expression for the image frames"), - io.String.Input(id="Audio", default="a*(1-w)+b*w", tooltip="Expression for the audio component"), + io.Autogrow.Input(id="V",template=io.Autogrow.TemplatePrefix(io.Video.Input("values"), 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.String.Input(id="Expression", default="I0*(1-F0)+I1*F0", tooltip="Expression to apply on tensor part of conditioning"), + io.String.Input(id="Expression_pi", default="I0*(1-F0)+I1*F0", tooltip="Expression to apply on pooled_input part of conditioning"), io.Combo.Input( id="length_mismatch", options=["tile", "error", "pad"], default="error", - tooltip="How to handle mismatched frame/sample counts. tile: repeat shorter inputs; error: raise error on mismatch; pad: treat missing as zero." + tooltip="How to handle mismatched image batch sizes. tile: repeat shorter inputs; error: raise error on mismatch; pad: treat missing frames as zero." ) ], outputs=[ - io.Video.Output(), + io.Conditioning.Output(), ], ) @classmethod - def check_lazy_status(cls, Images, Audio, a, b=[], c=[], d=[], w=0, x=0, y=0, z=0, length_mismatch="tile"): - return commonLazy(Images, a, b, c, d, w, x, y, z) + commonLazy(Audio, a, b, c, d, w, x, y, z) + def check_lazy_status(cls, Expression,Expression_pi, V, F, length_mismatch="tile"): + + input_stream = InputStream(Expression) + lexer = MathExprLexer(input_stream) + stream = CommonTokenStream(lexer) + stream.fill() + + input_stream = InputStream(Expression_pi) + lexer = MathExprLexer(input_stream) + stream1 = CommonTokenStream(lexer) + stream1.fill() + + # Support aliases + aliases_img = {"a": "V0", "b": "V1", "c": "V2", "d": "V3"} + aliases_flt = {"w": "F0", "x": "F1", "y": "F2", "z": "F3"} + + needed = set() + needed1 = set() + for token in filter(lambda t: t.type == MathExprParser.VARIABLE, stream.tokens + stream1.tokens): + var_name = token.text + + if re.match(r"[VF][0-9]+", var_name): + needed.add(var_name) + elif var_name in aliases_img: + needed.add(aliases_img[var_name]) + elif var_name in aliases_flt: + needed.add(aliases_flt[var_name]) + for v in needed: + if v.startswith("V"): + if v not in V or V[v] is None: + needed1.add(v) + elif v.startswith("F"): + if v not in F or F[v] is None: + needed1.add(v) + return needed1 @classmethod - def execute(cls, Images, Audio, a, b=None, c=None, d=None, w=0.0, x=0.0, y=0.0, z=0.0, length_mismatch="tile") -> io.NodeOutput: + def execute(cls, V, F, Expression, Expression_pi, length_mismatch="tile"): + ss = {} + tensor_keys = [k for k, v in V.items() if v is not None] + if not tensor_keys: + raise ValueError("At least one input is required.") - ac, bc, cc, dc = prepare_inputs(a,b,c,d) - imgs_ae = ac.images - imgs_be = bc.images - imgs_ce = cc.images - imgs_de = dc.images - imgs_ae, imgs_be, imgs_ce, imgs_de = normalize_to_common_shape(imgs_ae, imgs_be, imgs_ce, imgs_de, mode=length_mismatch) + tensors = [V[k][0] for k in tensor_keys] - img_vars = { - "a": imgs_ae, "b": imgs_be, "c": imgs_ce, "d": imgs_de, - "w": w, "x": x, "y": y, "z": z, - "B": getIndexTensorAlongDim(imgs_ae, 0), - "frame": getIndexTensorAlongDim(imgs_ae, 0), - "C": getIndexTensorAlongDim(imgs_ae, 1), - "channel": getIndexTensorAlongDim(imgs_ae, 1), - "H": imgs_ae.shape[2], - "height": imgs_ae.shape[2], - "W": imgs_ae.shape[3], - "width": imgs_ae.shape[3], - "T": imgs_ae.shape[0], - "frame_count": imgs_ae.shape[0], - "R": float(ac.frame_rate), - "frame_rate": float(ac.frame_rate), - "N": imgs_ae.shape[1], - "channel_count": imgs_ae.shape[1], - } | generate_dim_variables(imgs_ae) + # Normalize all tensors together to find the common target shape + normalized_tensors = normalize_to_common_shape(*tensors, mode=length_mismatch) + V_norm = dict(zip(tensor_keys, normalized_tensors)) - tree = parse_expr(Images); - visitor = UnifiedMathVisitor(img_vars, imgs_ae.shape) - result_tensor = visitor.visit(tree) - result_imgs = as_tensor(result_tensor, imgs_ae.shape) + # Use first normalized tensor to establish the reference shape + ref_tensor = normalized_tensors[0] + common_shape = ref_tensor.shape - # --- Process Audio --- - wav_a = ac.audio["waveform"] - sample_rate = ac.audio["sample_rate"] - wav_b = bc.audio["waveform"] if bc is not None else None - wav_c = cc.audio["waveform"] if cc is not None else None - wav_d = dc.audio["waveform"] if dc is not None else None + # Setup legacy variables a, b, c, d + ae = V_norm.get("V0", make_zero_like(ref_tensor)) + be = V_norm.get("V1", make_zero_like(ae)) + ce = V_norm.get("V2", make_zero_like(ae)) + de = V_norm.get("V3", make_zero_like(ae)) - # Determine target_wav_len as max of all provided audio - wav_lengths = [wav_a.shape[2]] - if bc is not None: wav_lengths.append(wav_b.shape[2]) - if cc is not None: wav_lengths.append(wav_c.shape[2]) - if dc is not None: wav_lengths.append(wav_d.shape[2]) - target_wav_len = max(wav_lengths) + # Ensure legacy variables are normalized in case they were zero-initialized + ae, be, ce, de = normalize_to_common_shape(ae, be, ce, de, mode=length_mismatch) - def resolve_wav(tensor, comp, name): - if comp is None: - return 0.0 + if(length_mismatch == "error"): + for name, tensor in V.items(): + if tensor is not None and tensor.shape[0] != common_shape[0]: + raise ValueError(f"Input '{name}' has shape {tensor.shape[0]}, expected {common_shape[0]} to match input.") - curr_len = tensor.shape[2] - if curr_len == target_wav_len: - return tensor + variables = { + "a": ae, "b": be, "c": ce, "d": de, + "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, + "X": getIndexTensorAlongDim(ae, 3), + "Y": getIndexTensorAlongDim(ae, 2), + "B": getIndexTensorAlongDim(ae, 0), + "batch": getIndexTensorAlongDim(ae, 0), + "C": getIndexTensorAlongDim(ae, 1), + "channel": getIndexTensorAlongDim(ae, 1), + "W": ae.shape[2], + "width": ae.shape[2], + "H": ae.shape[1], + "height": ae.shape[1], + "T": ae.shape[0], + "batch_count": ae.shape[0], + "N": ae.shape[3], + "channel_count": ae.shape[3], + } | generate_dim_variables(ae) - if length_mismatch == "tile": - return tensor.repeat(1, 1, math.ceil(target_wav_len / curr_len))[:, :, :target_wav_len] - elif length_mismatch == "pad": - out = torch.zeros((tensor.shape[0], tensor.shape[1], target_wav_len), device=tensor.device, dtype=tensor.dtype) - out[:, :, :curr_len] = tensor - return out - else: # error - raise ValueError(f"Audio samples mismatch in {name}: expected {target_wav_len} to match longest input, got {curr_len}. Set length_mismatch to 'broadcast' or 'pad' to handle this.") + # Add all dynamic inputs + variables.update(V_norm) - wav_ae = resolve_wav(wav_a, True, "input a") - wav_be = resolve_wav(wav_b, bc, "input b") - wav_ce = resolve_wav(wav_c, cc, "input c") - wav_de = resolve_wav(wav_d, dc, "input d") + for k, val in F.items(): + variables[k] = val if val is not None else 0.0 - wav_ae, wav_be, wav_ce, wav_de = normalize_to_common_shape(wav_ae, wav_be, wav_ce, wav_de, mode=length_mismatch) + tree = parse_expr(Expression); + visitor = UnifiedMathVisitor(variables, ae.shape,state_storage=ss) + result = visitor.visit(tree) + result = as_tensor(result, ae.shape) - wav_vars = { - "a": wav_ae, "b": wav_be, "c": wav_ce, "d": wav_de, - "w": w, "x": x, "y": y, "z": z, - "B": getIndexTensorAlongDim(wav_ae, 0), - "batch": getIndexTensorAlongDim(wav_ae, 0), - "C": getIndexTensorAlongDim(wav_ae, 1), - "channel": getIndexTensorAlongDim(wav_ae, 1), - "S": getIndexTensorAlongDim(wav_ae, 2), - "sample": getIndexTensorAlongDim(wav_ae, 2), + + + + 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, - "T": wav_ae.shape[2], - "sample_count": wav_ae.shape[2], - "N": wav_ae.shape[1], - "channel_count": wav_ae.shape[1], - } | generate_dim_variables(wav_ae) + "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 - tree = parse_expr(Audio); - visitor = UnifiedMathVisitor(wav_vars, wav_ae.shape) - result_tensor = visitor.visit(tree) - result_audio_wav = as_tensor(result_tensor, wav_ae.shape) + for k, val in F.items(): + variables[k] = val if val is not None else 0.0 - output = VideoFromComponents( - VideoComponents( - images=result_imgs, audio={"waveform": result_audio_wav, "sample_rate": sample_rate}, frame_rate=ac.frame_rate - ) - ) + tree = parse_expr(Expression); + visitor = UnifiedMathVisitor(variables, a_w.shape,state_storage=ss) + result1 = visitor.visit(tree) + result1 = as_tensor(result, a_w.shape) - return (output,) + return ([result,{"waveform":result1,"sample_rate":sample_rate}],) diff --git a/more_math/deprecated/VideoMathNode.py b/more_math/deprecated/VideoMathNode.py new file mode 100644 index 0000000..f2e86ca --- /dev/null +++ b/more_math/deprecated/VideoMathNode.py @@ -0,0 +1,158 @@ +import torch +import math +from ..helper_functions import ( + generate_dim_variables, + getIndexTensorAlongDim, + parse_expr, + commonLazy, + normalize_to_common_shape, + as_tensor, + prepare_inputs +) + +from ..Parser.UnifiedMathVisitor import UnifiedMathVisitor +from comfy_api.latest import io +from comfy_api.latest import VideoComponents, VideoFromComponents + + +class VideoMathNodeOLD(io.ComfyNode): + """ + Enables math expressions on Videos (Images + Audio). + Video object is expected to have 'images' and 'audio' components. + Expressions for 'Images' and 'Audio' are handled separately. + """ + + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id="mrmth_VideoMathNode", + is_deprecated=True, + display_name="Video math", + inputs=[ + io.Video.Input(id="a"), + io.Video.Input(id="b", optional=True, lazy=True), + io.Video.Input(id="c", optional=True, lazy=True), + io.Video.Input(id="d", optional=True, lazy=True), + io.Float.Input(id="w", default=0.0, optional=True, lazy=True, force_input=True), + io.Float.Input(id="x", default=0.0, optional=True, lazy=True, force_input=True), + io.Float.Input(id="y", default=0.0, optional=True, lazy=True, force_input=True), + io.Float.Input(id="z", default=0.0, optional=True, lazy=True, force_input=True), + io.String.Input(id="Images", default="a*(1-w)+b*w", tooltip="Expression for the image frames"), + io.String.Input(id="Audio", default="a*(1-w)+b*w", tooltip="Expression for the audio component"), + io.Combo.Input( + id="length_mismatch", + options=["tile", "error", "pad"], + default="error", + tooltip="How to handle mismatched frame/sample counts. tile: repeat shorter inputs; error: raise error on mismatch; pad: treat missing as zero." + ) + ], + outputs=[ + io.Video.Output(), + ], + ) + + @classmethod + def check_lazy_status(cls, Images, Audio, a, b=[], c=[], d=[], w=0, x=0, y=0, z=0, length_mismatch="tile"): + return commonLazy(Images, a, b, c, d, w, x, y, z) + commonLazy(Audio, a, b, c, d, w, x, y, z) + + @classmethod + def execute(cls, Images, Audio, a, b=None, c=None, d=None, w=0.0, x=0.0, y=0.0, z=0.0, length_mismatch="tile") -> io.NodeOutput: + + ac, bc, cc, dc = prepare_inputs(a,b,c,d) + imgs_ae = ac.images + imgs_be = bc.images + imgs_ce = cc.images + imgs_de = dc.images + imgs_ae, imgs_be, imgs_ce, imgs_de = normalize_to_common_shape(imgs_ae, imgs_be, imgs_ce, imgs_de, mode=length_mismatch) + + img_vars = { + "a": imgs_ae, "b": imgs_be, "c": imgs_ce, "d": imgs_de, + "w": w, "x": x, "y": y, "z": z, + "B": getIndexTensorAlongDim(imgs_ae, 0), + "frame": getIndexTensorAlongDim(imgs_ae, 0), + "C": getIndexTensorAlongDim(imgs_ae, 1), + "channel": getIndexTensorAlongDim(imgs_ae, 1), + "H": imgs_ae.shape[2], + "height": imgs_ae.shape[2], + "W": imgs_ae.shape[3], + "width": imgs_ae.shape[3], + "T": imgs_ae.shape[0], + "frame_count": imgs_ae.shape[0], + "R": float(ac.frame_rate), + "frame_rate": float(ac.frame_rate), + "N": imgs_ae.shape[1], + "channel_count": imgs_ae.shape[1], + } | generate_dim_variables(imgs_ae) + + tree = parse_expr(Images); + visitor = UnifiedMathVisitor(img_vars, imgs_ae.shape) + result_tensor = visitor.visit(tree) + result_imgs = as_tensor(result_tensor, imgs_ae.shape) + + # --- Process Audio --- + wav_a = ac.audio["waveform"] + sample_rate = ac.audio["sample_rate"] + wav_b = bc.audio["waveform"] if bc is not None else None + wav_c = cc.audio["waveform"] if cc is not None else None + wav_d = dc.audio["waveform"] if dc is not None else None + + # Determine target_wav_len as max of all provided audio + wav_lengths = [wav_a.shape[2]] + if bc is not None: wav_lengths.append(wav_b.shape[2]) + if cc is not None: wav_lengths.append(wav_c.shape[2]) + if dc is not None: wav_lengths.append(wav_d.shape[2]) + target_wav_len = max(wav_lengths) + + def resolve_wav(tensor, comp, name): + if comp is None: + return 0.0 + + curr_len = tensor.shape[2] + if curr_len == target_wav_len: + return tensor + + if length_mismatch == "tile": + return tensor.repeat(1, 1, math.ceil(target_wav_len / curr_len))[:, :, :target_wav_len] + elif length_mismatch == "pad": + out = torch.zeros((tensor.shape[0], tensor.shape[1], target_wav_len), device=tensor.device, dtype=tensor.dtype) + out[:, :, :curr_len] = tensor + return out + else: # error + raise ValueError(f"Audio samples mismatch in {name}: expected {target_wav_len} to match longest input, got {curr_len}. Set length_mismatch to 'broadcast' or 'pad' to handle this.") + + wav_ae = resolve_wav(wav_a, True, "input a") + wav_be = resolve_wav(wav_b, bc, "input b") + wav_ce = resolve_wav(wav_c, cc, "input c") + wav_de = resolve_wav(wav_d, dc, "input d") + + wav_ae, wav_be, wav_ce, wav_de = normalize_to_common_shape(wav_ae, wav_be, wav_ce, wav_de, mode=length_mismatch) + + wav_vars = { + "a": wav_ae, "b": wav_be, "c": wav_ce, "d": wav_de, + "w": w, "x": x, "y": y, "z": z, + "B": getIndexTensorAlongDim(wav_ae, 0), + "batch": getIndexTensorAlongDim(wav_ae, 0), + "C": getIndexTensorAlongDim(wav_ae, 1), + "channel": getIndexTensorAlongDim(wav_ae, 1), + "S": getIndexTensorAlongDim(wav_ae, 2), + "sample": getIndexTensorAlongDim(wav_ae, 2), + "R": sample_rate, + "sample_rate": sample_rate, + "T": wav_ae.shape[2], + "sample_count": wav_ae.shape[2], + "N": wav_ae.shape[1], + "channel_count": wav_ae.shape[1], + } | generate_dim_variables(wav_ae) + + tree = parse_expr(Audio); + visitor = UnifiedMathVisitor(wav_vars, wav_ae.shape) + result_tensor = visitor.visit(tree) + result_audio_wav = as_tensor(result_tensor, wav_ae.shape) + + output = VideoFromComponents( + VideoComponents( + images=result_imgs, audio={"waveform": result_audio_wav, "sample_rate": sample_rate}, frame_rate=ac.frame_rate + ) + ) + + return (output,) diff --git a/more_math/modelLikeCommon.py b/more_math/modelLikeCommon.py index 842ca76..8d4e256 100644 --- a/more_math/modelLikeCommon.py +++ b/more_math/modelLikeCommon.py @@ -34,7 +34,7 @@ def calculate_patches_autogrow(Expr, V, F, mapping=None): elif hasattr(m, "state_dict"): # VAE might have state_dict directly? all_keys.update(m.state_dict().keys()) elif hasattr(m, "patches"): # Mock object or raw patcher - # If it's just a patcher without underlying model access? + # If it's just a patcher without underlying model access? # Usually patcher.model.state_dict() is the way. pass @@ -112,15 +112,8 @@ def calculate_patches_autogrow(Expr, V, F, mapping=None): res = visitor.visit(tree) res = as_tensor(res, ref_tensor.shape) - # Calculate patch: Result - Original(V0) - # Assumption: We are patching V0. - # If V0 doesn't have the key, we assume V0 was zero? - # ComfyUI patching mechanism adds patch to original weights. - # If we output 'res', we need to return (res - original). - - # Get original weight for V0 (alias 'a' usually) - original = variables.get("V0") # Or strictly V.get("V0")'s weight - if original is None: + original = variables.get("V0") + if original is None: original = torch.zeros_like(res) diff = res - original diff --git a/more_math/nodes.py b/more_math/nodes.py index 55776fe..4a2986f 100644 --- a/more_math/nodes.py +++ b/more_math/nodes.py @@ -9,6 +9,7 @@ from .deprecated.MaskMathNode import MaskMathNodeOLD from .deprecated.ModelMathNode import ModelMathNodeOLD from .deprecated.VaeMathNode import VAEMathNodeOLD from .deprecated.ClipMathNode import CLIPMathNodeOLD +from .deprecated.VideoMathNode import VideoMathNodeOLD from .FloatMathNode import FloatMathNode from .ConditioningMathNode import ConditioningMathNode @@ -105,6 +106,7 @@ class MoreMathExtension(ComfyExtension): FloatToIntNode, AudioMathNodeOLD, VideoMathNode, + VideoMathNodeOLD, AudioToSpectrogram, SpectrogramToAudio, SigmasMathNode,