Did I forgot commit these? + AI refactor

This commit is contained in:
mcDandy
2026-01-25 22:43:13 +01:00
parent a395bbf3b4
commit 5574fe6627
7 changed files with 177 additions and 148 deletions
+51 -34
View File
@@ -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},)
+49 -44
View File
@@ -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
+21 -26
View File
@@ -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
+20 -16
View File
@@ -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
+22 -24
View File
@@ -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)
+11 -1
View File
@@ -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
+3 -3
View File
@@ -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):