diff --git a/more_math/AudioMathNode.py b/more_math/AudioMathNode.py index 8ef0e87..9b72f17 100644 --- a/more_math/AudioMathNode.py +++ b/more_math/AudioMathNode.py @@ -1,5 +1,13 @@ import torch -from .helper_functions import generate_dim_variables, parse_expr, getIndexTensorAlongDim, as_tensor, prepare_inputs, normalize_to_common_shape +from .helper_functions import ( + generate_dim_variables, + parse_expr, + getIndexTensorAlongDim, + as_tensor, + prepare_inputs, + normalize_to_common_shape, + make_zero_like, +) from .Parser.UnifiedMathVisitor import UnifiedMathVisitor from comfy_api.latest import io from antlr4 import InputStream, CommonTokenStream @@ -73,50 +81,59 @@ class AudioMathNode(io.ComfyNode): @classmethod def execute(cls, V, F, Expression, length_mismatch="tile"): + # 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.") + + 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"): - max_lengths = V.get("V0")["waveform"].shape - for name, tensor in V.items(): - if tensor["waveform"] is not None and max_lengths!=tensor["waveform"].shape: - raise ValueError(f"Input '{name}' has shape ({tensor['waveform'].shape[0]}, {tensor['waveform'].shape[2]}), expected ({max_lengths[0]}, {max_lengths[2]}) to match input.") + 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) - waveforms={} - sample_rates={} - for key, audio in V.items(): - if audio is not None and isinstance(audio, dict) and "waveform" in audio: - waveforms[key] = audio["waveform"] - sample_rates[key+"sr"] = audio.get("sample_rate", 44100) - else: - waveforms[key] = torch.zeros(V["V0"]["waveform"].shape) - sample_rates[key] = 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, + "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, 0), - "C": getIndexTensorAlongDim(a, 1), - "channel": getIndexTensorAlongDim(a, 1), - "S": getIndexTensorAlongDim(a, 2), - "sample": getIndexTensorAlongDim(a, 2), + "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, 0), - "T": a.shape[0], - "batch_count": a.shape[0], - } | generate_dim_variables(a) | waveforms | sample_rates + "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 + + for k, val in F.items(): + variables[k] = val if val is not None else 0.0 tree = parse_expr(Expression); - visitor = UnifiedMathVisitor(variables, a.shape) + visitor = UnifiedMathVisitor(variables, a_w.shape) result = visitor.visit(tree) - result = as_tensor(result, a.shape) + result = as_tensor(result, a_w.shape) return ({"waveform":result,"sample_rate":sample_rate},) diff --git a/more_math/ConditioningMathNode.py b/more_math/ConditioningMathNode.py index 4346f77..da5fa80 100644 --- a/more_math/ConditioningMathNode.py +++ b/more_math/ConditioningMathNode.py @@ -1,5 +1,5 @@ import torch -from .helper_functions import generate_dim_variables, parse_expr, getIndexTensorAlongDim, as_tensor, prepare_inputs, normalize_to_common_shape +from .helper_functions import generate_dim_variables, parse_expr, getIndexTensorAlongDim, as_tensor, prepare_inputs, normalize_to_common_shape, make_zero_like from .Parser.UnifiedMathVisitor import UnifiedMathVisitor from comfy_api.latest import io from antlr4 import InputStream, CommonTokenStream @@ -79,40 +79,44 @@ class ConditioningMathNode(io.ComfyNode): @classmethod def execute(cls, V, F, Expression, Expression_pi, length_mismatch="tile"): + # Identify all present conditioning inputs + tensor_keys = [k for k, v in V.items() if v is not None and isinstance(v, list) and len(v) > 0] + if not tensor_keys: + raise ValueError("At least one input is required.") + # Extract tensors and pooled outputs - tensor = {} - pooled_output = {} + tensors = {} + pooled_outputs = {} - # Get shape reference from V0 (assumed to exist and be valid conditioning) - ref_cond = V.get("V0") - ref_tensor_shape = ref_cond[0][0].shape if ref_cond else None - ref_pooled_shape = ref_cond[0][1].get("pooled_output").shape if ref_cond and "pooled_output" in ref_cond[0][1] else None + for key in tensor_keys: + conditioning = V[key] + tensors[key] = conditioning[0][0] + # pooled_output is optional in the dict + pooled_outputs[key] = conditioning[0][1].get("pooled_output") - for key, conditioning in V.items(): - # Standard Conditioning is list of [tensor, dict] - if isinstance(conditioning, list) and len(conditioning) > 0 and isinstance(conditioning[0], (list, tuple)): - tensor[key] = conditioning[0][0] - pooled_output[key] = conditioning[0][1].get("pooled_output", torch.zeros(ref_pooled_shape) if ref_pooled_shape is not None else None) - else: - # Fallback to zeros if structure is unknown or empty - tensor[key] = torch.zeros(ref_tensor_shape) if ref_tensor_shape is not None else None - pooled_output[key] = torch.zeros(ref_pooled_shape) if ref_pooled_shape is not None else None + # Normalize main tensors + norm_tensors_batch = normalize_to_common_shape(*tensors.values(), mode=length_mismatch) + V_norm_tensors = dict(zip(tensor_keys, norm_tensors_batch)) - # Normalize shapes - new_values = normalize_to_common_shape(*tensor.values(), mode=length_mismatch) - tensor.update(zip(tensor.keys(), new_values)) + ref_tensor = norm_tensors_batch[0] + common_shape = ref_tensor.shape - if any(p is not None for p in pooled_output.values()): - # Filter out Nones for normalization if any - valid_pooled = {k:v for k,v in pooled_output.items() if v is not None} - if valid_pooled: - new_p_values = normalize_to_common_shape(*valid_pooled.values(), mode=length_mismatch) - pooled_output.update(zip(valid_pooled.keys(), new_p_values)) + # Normalize pooled outputs (if they exist) + valid_pooled_keys = [k for k, v in pooled_outputs.items() if v is not None] + if valid_pooled_keys: + norm_pooled_batch = normalize_to_common_shape(*[pooled_outputs[k] for k in valid_pooled_keys], mode=length_mismatch) + V_norm_pooled = dict(zip(valid_pooled_keys, norm_pooled_batch)) + ref_pooled = norm_pooled_batch[0] + else: + V_norm_pooled = {} + ref_pooled = torch.tensor([]) - a = tensor.get("V0") - b = tensor.get("V1") - c = tensor.get("V2") - d = tensor.get("V3") + # Setup legacy variables a, b, c, d (Main Tensor) + a = V_norm_tensors.get("V0", make_zero_like(ref_tensor)) + b = V_norm_tensors.get("V1", make_zero_like(a)) + c = V_norm_tensors.get("V2", make_zero_like(a)) + d = V_norm_tensors.get("V3", make_zero_like(a)) + a, b, c, d = normalize_to_common_shape(a, b, c, d, mode=length_mismatch) # variables for Main Tensor (Expression) variables = { @@ -125,7 +129,7 @@ class ConditioningMathNode(io.ComfyNode): "batch": getIndexTensorAlongDim(a, 0), "T": a.shape[0], "batch_count": a.shape[0], - } | generate_dim_variables(a) | tensor + } | generate_dim_variables(a) | V_norm_tensors # Execute Expression (Main Tensor) tree = parse_expr(Expression) @@ -135,28 +139,29 @@ class ConditioningMathNode(io.ComfyNode): # variables for Pooled Output (Expression_pi) - a_p = pooled_output.get("V0") - b_p = pooled_output.get("V1") - c_p = pooled_output.get("V2") - d_p = pooled_output.get("V3") + a_p = V_norm_pooled.get("V0", make_zero_like(ref_pooled)) + b_p = V_norm_pooled.get("V1", make_zero_like(a_p)) + c_p = V_norm_pooled.get("V2", make_zero_like(a_p)) + d_p = V_norm_pooled.get("V3", make_zero_like(a_p)) + a_p, b_p, c_p, d_p = normalize_to_common_shape(a_p, b_p, c_p, d_p, mode=length_mismatch) - variables = { + variables_pi = { "a": a_p, "b": b_p, "c": c_p, "d": d_p, "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_p, 0), - "batch": getIndexTensorAlongDim(a_p, 0), - "T": a_p.shape[0], - "batch_count": a_p.shape[0], - } | generate_dim_variables(a_p) | pooled_output + "B": getIndexTensorAlongDim(a_p, 0) if a_p.numel() > 0 else torch.tensor([]), + "batch": getIndexTensorAlongDim(a_p, 0) if a_p.numel() > 0 else torch.tensor([]), + "T": a_p.shape[0] if a_p.numel() > 0 else 0, + "batch_count": a_p.shape[0] if a_p.numel() > 0 else 0, + } | generate_dim_variables(a_p) | V_norm_pooled # Execute Expression_pi (Pooled Output) - tree = parse_expr(Expression_pi) - visitor = UnifiedMathVisitor(variables, a_p.shape) - rpooled = visitor.visit(tree) - rpooled = as_tensor(rpooled, a_p.shape) + tree_pi = parse_expr(Expression_pi) + visitor_pi = UnifiedMathVisitor(variables_pi, a_p.shape) + rpooled_raw = visitor_pi.visit(tree_pi) + rpooled = as_tensor(rpooled_raw, a_p.shape) # Clone result structure import copy diff --git a/more_math/ImageMathNode.py b/more_math/ImageMathNode.py index 14bb196..59a62fc 100644 --- a/more_math/ImageMathNode.py +++ b/more_math/ImageMathNode.py @@ -1,3 +1,4 @@ +import torch from .helper_functions import generate_dim_variables, parse_expr, getIndexTensorAlongDim, as_tensor, prepare_inputs, normalize_to_common_shape, make_zero_like from .Parser.UnifiedMathVisitor import UnifiedMathVisitor from comfy_api.latest import io @@ -78,34 +79,34 @@ class ImageMathNode(io.ComfyNode): 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: + # Identify all present tensors and their keys + 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.") - a = V.get("V0") - b = V.get("V1") - c = V.get("V2") - d = V.get("V3") + tensors = [V[k] for k in tensor_keys] - # Fallback for a if missing (unlikely if V0 is default but possible) - if a is None: - a = make_zero_like(ref_image) + # 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)) - ae, be, ce, de = prepare_inputs(a, b, c, d) + # Use first normalized tensor to establish the reference shape + ref_tensor = normalized_tensors[0] + common_shape = ref_tensor.shape + # 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)) + + # 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) if(length_mismatch == "error"): - max_length = ae.shape[0] for name, tensor in V.items(): - if tensor is not None and tensor.shape[0] != max_length: - raise ValueError(f"Input '{name}' has shape {tensor.shape[0]}, expected {max_length} to match largest input.") + 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.") variables = { "a": ae, "b": be, "c": ce, "d": de, @@ -130,13 +131,7 @@ class ImageMathNode(io.ComfyNode): } | generate_dim_variables(ae) # Add all dynamic inputs - for k, v in V.items(): - if v is not None: - # Normalize all images in V to match ae.shape - # Note: normalize_to_common_shape args are *tensors. - # We normalize individual V item against 'ae' (the reference shape) - norm_v = normalize_to_common_shape(ae, v, mode=length_mismatch)[1] - variables[k] = norm_v + variables.update(V_norm) for k, val in F.items(): variables[k] = val if val is not None else 0.0 diff --git a/more_math/LatentMathNode.py b/more_math/LatentMathNode.py index 518c621..48f4ceb 100644 --- a/more_math/LatentMathNode.py +++ b/more_math/LatentMathNode.py @@ -125,18 +125,26 @@ class LatentMathNode(io.ComfyNode): if a is None: a = make_zero_like(ref_latent) - a_c, b_c, c_c, d_c = prepare_inputs(a, b, c, d) - at,bt,ct,dt = a_c["samples"],b_c["samples"],c_c["samples"],d_c["samples"] + # Identify all present tensors and their keys + tensor_keys = [k for k, v in V.items() if v is not None] + at_list = [V[k]["samples"] for k in tensor_keys] + + # Normalize all together + normalized_samples = normalize_to_common_shape(*at_list, mode=length_mismatch) + V_norm_samples = dict(zip(tensor_keys, normalized_samples)) + + ae = V_norm_samples.get("V0", make_zero_like(normalized_samples[0])) + be = V_norm_samples.get("V1", make_zero_like(ae)) + ce = V_norm_samples.get("V2", make_zero_like(ae)) + de = V_norm_samples.get("V3", make_zero_like(ae)) + + # Ensure legacy are normalized + ae, be, ce, de = normalize_to_common_shape(ae, be, ce, de, mode=length_mismatch) if(length_mismatch == "error"): - max_length = at.shape[0] - for name, val in V.items(): - if val is not None: - tensor = val["samples"] - if tensor.shape[0] != max_length: - raise ValueError(f"Input '{name}' has shape {tensor.shape[0]}, expected {max_length} to match largest input.") - - ae, be, ce, de = normalize_to_common_shape(at, bt, ct, dt, mode=length_mismatch) + for name in tensor_keys: + if V[name]["samples"].shape[0] != ae.shape[0]: + raise ValueError(f"Input '{name}' has shape {V[name]['samples'].shape[0]}, expected {ae.shape[0]} to match input.") # parse expression once tree = parse_expr(Expression) @@ -179,11 +187,7 @@ class LatentMathNode(io.ComfyNode): variables.update({"frame_idx": F_idx, "frame": F_idx, "frame_count": frame_count}) # Add all dynamic inputs - for k, v in V.items(): - if v is not None: - v_tensor = v["samples"] - norm_v = normalize_to_common_shape(ae, v_tensor, mode=length_mismatch)[1] - variables[k] = norm_v + variables.update(V_norm_samples) for k, v in F.items(): variables[k] = v if v is not None else 0.0 @@ -191,7 +195,7 @@ class LatentMathNode(io.ComfyNode): visitor = UnifiedMathVisitor(variables, ae.shape) result_t = as_tensor(visitor.visit(tree), ae.shape) - result_latent = a_c.copy() + result_latent = ref_latent.copy() if stacked and orig_split_sizes is not None: from comfy.nested_tensor import NestedTensor # Restore original split sizes diff --git a/more_math/MaskMathNode.py b/more_math/MaskMathNode.py index acf8181..5b0e357 100644 --- a/more_math/MaskMathNode.py +++ b/more_math/MaskMathNode.py @@ -73,32 +73,33 @@ class MaskMathNode(io.ComfyNode): @classmethod def execute(cls, V, F, Expression, length_mismatch="tile"): - # Determine reference mask - ref_mask = None - for mask in V.values(): - if mask is not None: - ref_mask = mask - break - - if ref_mask is None: + # Identify all present tensors and their keys + 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.") - a = V.get("V0") - b = V.get("V1") - c = V.get("V2") - d = V.get("V3") + tensors = [V[k] for k in tensor_keys] - if a is None: - a = make_zero_like(ref_mask) + # Normalize all tensors together + normalized_tensors = normalize_to_common_shape(*tensors, mode=length_mismatch) + V_norm = dict(zip(tensor_keys, normalized_tensors)) - ae, be, ce, de = prepare_inputs(a, b, c, d) + # Establish reference shape + ref_tensor = normalized_tensors[0] + common_shape = ref_tensor.shape if(length_mismatch == "error"): - max_length = ae.shape[0] for name, tensor in V.items(): - if tensor is not None and tensor.shape[0] != max_length: - raise ValueError(f"Input '{name}' has shape {tensor.shape[0]}, expected {max_length} to match largest input.") + 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 largest input.") + # 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)) + + # Ensure legacy are normalized ae, be, ce, de = normalize_to_common_shape(ae, be, ce, de, mode=length_mismatch) variables = { @@ -120,13 +121,10 @@ class MaskMathNode(io.ComfyNode): } | generate_dim_variables(ae) # Add all dynamic inputs - for k, v in V.items(): - if v is not None: - norm_v = normalize_to_common_shape(ae, v, mode=length_mismatch)[1] - variables[k] = norm_v + variables.update(V_norm) - for k, v in F.items(): - variables[k] = v if v is not None else 0.0 + for k, val in F.items(): + variables[k] = val if val is not None else 0.0 tree = parse_expr(Expression); visitor = UnifiedMathVisitor(variables, ae.shape) diff --git a/more_math/Parser/UnifiedMathVisitor.py b/more_math/Parser/UnifiedMathVisitor.py index d53de38..e3d5f0e 100644 --- a/more_math/Parser/UnifiedMathVisitor.py +++ b/more_math/Parser/UnifiedMathVisitor.py @@ -1381,7 +1381,11 @@ class UnifiedMathVisitor(MathExprVisitor): def visitEdgeFunc(self, ctx): original_shape = tsr.shape tsr = tsr.float() + reshap = False + if len(ctx.expr()) >= 2: + reshap_val = self.visit(ctx.expr(1)) + reshap = bool(reshap_val.item()) if self._is_tensor(reshap_val) else bool(reshap_val) def sobel_op(x): kx = torch.tensor([[-1, 0, 1], [-2, 0, 2], [-1, 0, 1]], device=x.device, dtype=x.dtype) @@ -1394,11 +1398,17 @@ class UnifiedMathVisitor(MathExprVisitor): def visitGaussianFunc(self, ctx): tsr = self._promote_to_tensor(self.visit(ctx.expr(0))) - sigma = float(self.visit(ctx.expr(1))) + + sigma_val = self.visit(ctx.expr(1)) + sigma = float(sigma_val.item()) if self._is_tensor(sigma_val) else float(sigma_val) + if sigma <= 0: return tsr original_shape = tsr.shape tsr = tsr.float() reshap = False + if len(ctx.expr()) >= 3: + reshap_val = self.visit(ctx.expr(2)) + reshap = bool(reshap_val.item()) if self._is_tensor(reshap_val) else bool(reshap_val) def blur_op(x): kernel_size = int(6 * sigma + 1) if kernel_size % 2 == 0: kernel_size += 1 diff --git a/more_math/helper_functions.py b/more_math/helper_functions.py index e4dda99..3861251 100644 --- a/more_math/helper_functions.py +++ b/more_math/helper_functions.py @@ -18,12 +18,12 @@ def as_tensor(value, shape): if getattr(value, "is_nested", False): return value if isinstance(value, torch.Tensor): - if value.shape != shape: - return value.broadcast_to(shape).contiguous() + # Pass tensors through unchanged; the expression defines the output shape. return value.contiguous() if isinstance(value, (float, int)): value = (value,) - return torch.broadcast_to(torch.Tensor(value), shape) + # If it's a scalar or list, broadcast to the reference shape provided. + return torch.broadcast_to(torch.Tensor(value).to(dtype=torch.float32), shape).contiguous() def parse_expr(expr: str):