AI: fix most tests

This commit is contained in:
mcDandy
2026-02-27 18:17:41 +01:00
parent 5398b651ed
commit 64a2f6358b
14 changed files with 1489 additions and 1364 deletions
+13 -13
View File
@@ -43,7 +43,7 @@ class ConditioningMathNode(io.ComfyNode):
default="error",
tooltip="How to handle mismatched image batch sizes. tile: repeat shorter inputs; error: raise error on mismatch; pad: treat missing frames as zero."
),
io.Int.Input(id="batching"),
io.Int.Input(id="batching", default=0),
MrmthStack.Input(id="stack",optional=True)
],
outputs=[
@@ -53,14 +53,14 @@ class ConditioningMathNode(io.ComfyNode):
)
@classmethod
def check_lazy_status(cls, Expression,Expression_pi, V, F,batching, length_mismatch="tile",stack={}):
def check_lazy_status(cls, Expression,Expression_pi, V, F, length_mismatch="tile", batching=0, stack={}):
d = checkLazyNew(Expression,V,F)
b = checkLazyNew(Expression_pi,V,F)
return d|b
@classmethod
def execute(cls, V, F, Expression, Expression_pi,batching, length_mismatch="tile",stack={}):
def execute(cls, V, F, Expression, Expression_pi, length_mismatch="tile", batching=0, stack={}):
# 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:
@@ -198,20 +198,20 @@ class ConditioningMathNode(io.ComfyNode):
result_tensor = rt_chunks[i] if i < len(rt_chunks) else torch.zeros([1])
result_pooled = rp_chunks[i] if i < len(rp_chunks) else torch.zeros([1])
base = copy.deepcopy(V["V0"])
base[0][0] = result_tensor
if len(base[0]) == 1:
base[0].append({"pooled_output": result_pooled})
else:
base[0][1]["pooled_output"] = result_pooled
# base[0] is a tuple (tensor, dict), need to reconstruct
old_dict = base[0][1] if len(base[0]) > 1 else {}
new_dict = old_dict.copy()
new_dict["pooled_output"] = result_pooled
base[0] = (result_tensor, new_dict)
res_list.append(base)
else:
# Single output (no batching)
base = copy.deepcopy(V["V0"])
base[0][0] = rtensor
if len(base[0]) == 1:
base[0].append({"pooled_output": rpooled})
else:
base[0][1]["pooled_output"] = rpooled
# base[0] is a tuple (tensor, dict), need to reconstruct
old_dict = base[0][1] if len(base[0]) > 1 else {}
new_dict = old_dict.copy()
new_dict["pooled_output"] = rpooled
base[0] = (rtensor, new_dict)
res_list = [base]
return (res_list,stack)
+1 -1
View File
@@ -83,7 +83,7 @@ class ImageMathNode(io.ComfyNode):
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.")
raise ValueError(f"Input '{name}' has shape {tensor.shape[0]}, expected {common_shape[0]} to match largest input.")
variables = {
"a": ae, "b": be, "c": ce, "d": de,
+4 -2
View File
@@ -123,7 +123,8 @@ func1:
| SIGM LPAREN expr RPAREN # sigmoidFunc
| ANGL LPAREN expr RPAREN # anglFunc
| PRNT LPAREN expr RPAREN # printFunc
| FRACT LPAREN expr RPAREN # ReluFunc
| FRACT LPAREN expr RPAREN # FractFunc
| RELU LPAREN expr RPAREN # ReluFunc
| SOFTPLUS LPAREN expr RPAREN # SoftplusFunc
| GELU LPAREN expr RPAREN # GeluFunc
| SIGN LPAREN expr RPAREN # SignFunc
@@ -162,6 +163,7 @@ func1:
| LOWER LPAREN expr RPAREN # LowerFunc
| TRIM LPAREN expr RPAREN # TrimFunc
| ENTROPY LPAREN expr RPAREN # EntropyFunc
| SFFT LPAREN expr RPAREN # SfftFunc
| DILATE LPAREN expr (COMMA expr)? RPAREN # DilateFunc
| ERODE LPAREN expr (COMMA expr)? RPAREN # ErodeFunc
| MORPH_OPEN LPAREN expr (COMMA expr)? RPAREN # MorphOpenFunc
@@ -217,7 +219,7 @@ func3:
| SINE_EASE LPAREN expr COMMA expr COMMA expr RPAREN # SineEaseFunc
| SMOOTHERSTEP LPAREN expr COMMA expr COMMA expr RPAREN # SmootherstepFunc
| CROP LPAREN expr COMMA expr COMMA expr RPAREN # CropFunc
| SIFFT LPAREN expr (COMMA expr)? RPAREN # sifftFunc
| SIFFT LPAREN expr (COMMA expr)? RPAREN # SifftFunc
| OVERLAY LPAREN expr COMMA expr COMMA expr RPAREN # OverlayFunc
| RGB_TO_HSV LPAREN expr (COMMA expr COMMA expr)? (COMMA expr)? RPAREN # RgbToHsvFunc
| HSV_TO_RGB LPAREN expr (COMMA expr COMMA expr)? (COMMA expr)? RPAREN # HsvToRgbFunc;
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large Load Diff
+11 -1
View File
@@ -474,6 +474,11 @@ class MathExprVisitor(ParseTreeVisitor):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#FractFunc.
def visitFractFunc(self, ctx:MathExprParser.FractFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#ReluFunc.
def visitReluFunc(self, ctx:MathExprParser.ReluFuncContext):
return self.visitChildren(ctx)
@@ -669,6 +674,11 @@ class MathExprVisitor(ParseTreeVisitor):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#SfftFunc.
def visitSfftFunc(self, ctx:MathExprParser.SfftFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#DilateFunc.
def visitDilateFunc(self, ctx:MathExprParser.DilateFuncContext):
return self.visitChildren(ctx)
@@ -919,7 +929,7 @@ class MathExprVisitor(ParseTreeVisitor):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#sifftFunc.
# Visit a parse tree produced by MathExprParser#SifftFunc.
def visitSifftFunc(self, ctx:MathExprParser.SifftFuncContext):
return self.visitChildren(ctx)
+9 -7
View File
@@ -1096,7 +1096,7 @@ class UnifiedMathVisitor(MathExprVisitor):
self.variables = self.variables | generate_dim_variables(k_sq_sum)
try:
val = self._promote_to_tensor((yield ctx.expr()))
val = self._promote_to_tensor((yield ctx.expr(0)))
dims = tuple(range(val.ndim))
return torch.fft.ifftn(val, dim=dims).real
finally:
@@ -2947,13 +2947,15 @@ class UnifiedMathVisitor(MathExprVisitor):
if off >= base_size:
return base # Overlay outside of base, return original
# Determine overlay crop region (what part of overlay to use)
crop_start = max(0, -off) # Crop from overlay if offset is negative
crop_end = min(overlay_size, base_size - off) # Crop if overlay extends beyond base
if off < 0:
overlay = overlay[-off:]
off = 0
# Determine base paste region (where to place overlay in base)
paste_start = max(0, off) # Start position in base
paste_end = min(base_size, off + overlay_size) # End position in base
end = min(base_size, off + overlay_size)
paste_start = off
paste_end = end
crop_start = 0
crop_end = paste_end - paste_start
crop_slices.append(slice(crop_start, crop_end))
paste_slices.append(slice(paste_start, paste_end))
+1 -1
View File
@@ -72,7 +72,7 @@ class VAEMathNode(io.ComfyNode):
aliases = {"a": "V0", "b": "V1", "c": "V2", "d": "V3", "w": "F0", "x": "F1", "y": "F2", "z": "F3"}
layer_count = V.get("V0").model.state_dict().__len__() if hasattr(V.get("V0"), "model") and hasattr(V.get("V0").model, "state_dict") else 0
pbar = comfy.utils.ProgressBar(layer_count)
patches = calculate_patches_autogrow(Expression, V=patchers_V, F=F, mapping=aliases,stack=stack)
patches = calculate_patches_autogrow(Expression, V=patchers_V, F=F, pbar=pbar, mapping=aliases, stack=stack)
# VAE does not have a clone method, so we shallow copy and clone the patcher
out_vae = copy.copy(a)
+12 -5
View File
@@ -5,9 +5,9 @@ import torch
def calculate_patches(Model, a, b=None, c=None, d=None, w=0.0, x=0.0, y=0.0, z=0.0):
"""Legacy calculate_patches for backward compatibility."""
return calculate_patches_autogrow(Model, V={"V0": a, "V1": b, "V2": c, "V3": d}, F={"F0": w, "F1": x, "F2": y, "F3": z}, mapping={"a": "V0", "b": "V1", "c": "V2", "d": "V3", "w": "F0", "x": "F1", "y": "F2", "z": "F3"})
return calculate_patches_autogrow(Model, V={"V0": a, "V1": b, "V2": c, "V3": d}, F={"F0": w, "F1": x, "F2": y, "F3": z}, pbar=None, mapping={"a": "V0", "b": "V1", "c": "V2", "d": "V3", "w": "F0", "x": "F1", "y": "F2", "z": "F3"})
def calculate_patches_autogrow(Expr, V, F,pbar, mapping=None,stack = []):
def calculate_patches_autogrow(Expr, V, F, pbar=None, mapping=None, stack=[]):
"""
Calculate patches for model-like objects (Model, VAE, CLIP) using Autogrow inputs.
Iterates over the UNION of keys from all input models to support merging disjoint architectures/patches.
@@ -113,8 +113,14 @@ def calculate_patches_autogrow(Expr, V, F,pbar, mapping=None,stack = []):
for alias, target in mapping.items():
if target in variables:
variables[alias] = variables[target]
elif target in V: # V exists but key missing
variables[alias] = torch.zeros_like(ref_tensor)
elif target in V:
# V key exists in input dict but this specific layer key is missing
variables[alias] = torch.zeros_like(ref_tensor)
else:
# Target doesn't exist at all (e.g., V1 not provided)
# Check if it's a V-key pattern and zero-fill
if target.startswith("V") and target[1:].isdigit():
variables[alias] = torch.zeros_like(ref_tensor)
v_stacked, v_cnt = get_v_variable(variables)
if v_stacked is not None:
@@ -146,6 +152,7 @@ def calculate_patches_autogrow(Expr, V, F,pbar, mapping=None,stack = []):
if not torch.all(diff == 0):
patches[key] = (diff,)
pbar.update(1)
if pbar is not None:
pbar.update(1)
return patches
+23 -7
View File
@@ -17,6 +17,14 @@ def test_mixed_dtype_operations():
print("Testing Mixed Dtype Operations")
print("=" * 70)
# Create a dummy context for testing
class DummyContext:
class Start:
line = 0
column = 0
start = Start()
ctx = DummyContext()
# int8 operations
print("\n1. INT8 Bitwise Operations:")
a_int8 = torch.tensor([15, 7, 3], dtype=torch.int8)
@@ -32,9 +40,9 @@ def test_mixed_dtype_operations():
b_int16 = torch.tensor([15, 31, 7], dtype=torch.int16)
visitor = UnifiedMathVisitor({"a": a_int16, "b": b_int16})
result_and = visitor._bitwise_op(a_int16, b_int16, torch.bitwise_and, lambda x, y: x & y)
result_or = visitor._bitwise_op(a_int16, b_int16, torch.bitwise_or, lambda x, y: x | y)
result_xor = visitor._bitwise_op(a_int16, b_int16, torch.bitwise_xor, lambda x, y: x ^ y)
result_and = visitor._bitwise_op(a_int16, b_int16, torch.bitwise_and, lambda x, y: x & y, ctx)
result_or = visitor._bitwise_op(a_int16, b_int16, torch.bitwise_or, lambda x, y: x | y, ctx)
result_xor = visitor._bitwise_op(a_int16, b_int16, torch.bitwise_xor, lambda x, y: x ^ y, ctx)
print(f" a: {a_int16} (dtype={a_int16.dtype})")
print(f" b: {b_int16} (dtype={b_int16.dtype})")
@@ -72,34 +80,42 @@ def test_scalar_list_tensor_combinations():
print("Testing Scalar/List/Tensor Combinations")
print("=" * 70)
# Create a dummy context for testing
class DummyContext:
class Start:
line = 0
column = 0
start = Start()
ctx = DummyContext()
visitor = UnifiedMathVisitor({})
# Tensor & Tensor
print("\n1. Tensor & Tensor (int16):")
a = torch.tensor([7, 14, 21], dtype=torch.int16)
b = torch.tensor([3, 5, 7], dtype=torch.int16)
result = visitor._bitwise_op(a, b, torch.bitwise_and, lambda x, y: x & y)
result = visitor._bitwise_op(a, b, torch.bitwise_and, lambda x, y: x & y, ctx)
print(f" {a} & {b} = {result}")
# Tensor & List
print("\n2. Tensor & List (mixed):")
a = torch.tensor([15, 14, 13], dtype=torch.int16)
b = [7, 3, 1]
result = visitor._bitwise_op(a, b, torch.bitwise_and, lambda x, y: x & y)
result = visitor._bitwise_op(a, b, torch.bitwise_and, lambda x, y: x & y, ctx)
print(f" Tensor({a}) & List({b}) = {result}")
# Scalar & List
print("\n3. Scalar & List:")
a = 15
b = [7, 3, 1]
result = visitor._bitwise_op(a, b, torch.bitwise_and, lambda x, y: x & y)
result = visitor._bitwise_op(a, b, torch.bitwise_and, lambda x, y: x & y, ctx)
print(f" {a} & {b} = {result}")
# List & List
print("\n4. List & List:")
a = [15, 14, 13]
b = [7, 3, 1]
result = visitor._bitwise_op(a, b, torch.bitwise_and, lambda x, y: x & y)
result = visitor._bitwise_op(a, b, torch.bitwise_and, lambda x, y: x & y, ctx)
print(f" {a} & {b} = {result}")
print("\n" + "=" * 70)
+9 -1
View File
@@ -48,6 +48,14 @@ def test_bitwise_and_int16():
"""Test bitwise AND operation with int16 tensors."""
print("Testing bitwise AND with int16...")
# Create a dummy context for testing
class DummyContext:
class Start:
line = 0
column = 0
start = Start()
ctx = DummyContext()
a = torch.tensor([7, 14, 21], dtype=torch.int16)
b = torch.tensor([3, 5, 7], dtype=torch.int16)
@@ -55,7 +63,7 @@ def test_bitwise_and_int16():
visitor = UnifiedMathVisitor(variables)
# Test bitwise AND
result = visitor._bitwise_op(a, b, torch.bitwise_and, lambda x, y: x & y)
result = visitor._bitwise_op(a, b, torch.bitwise_and, lambda x, y: x & y, ctx)
print(f" a: {a}, b: {b}")
print(f" bitwise_and result: {result}")
assert result.dtype in [torch.int16, torch.int32], f"Unexpected dtype {result.dtype}"
+4 -3
View File
@@ -15,10 +15,11 @@ def test_conditioning_token_mismatch_padding():
# a + b -> result should have 154 tokens
# tokens 0-76: 1 + 0.5 = 1.5
# tokens 77-153: 0 + 0.5 = 0.5
result, = ConditioningMathNode.execute(V={"V0": ca, "V1": cb}, F={}, Expression="a + b", Expression_pi="a + b", length_mismatch="pad")
result_list, stack = ConditioningMathNode.execute(V={"V0": ca, "V1": cb}, F={}, Expression="a + b", Expression_pi="a + b", length_mismatch="pad")
result = result_list[0]
res_tensor = result[0][0]
res_dict = result[0][1]["pooled_output"]
res_tensor = result[0]
res_dict = result[1]["pooled_output"]
assert res_tensor.shape == (1, 154, 1024)
assert torch.allclose(res_tensor[0, :77, :], torch.full((77, 1024), 1.5))
assert torch.allclose(res_tensor[0, 77:, :], torch.full((77, 1024), 0.5))
+6 -5
View File
@@ -28,7 +28,7 @@ class TestMathGuider(unittest.TestCase):
# Expression: Average V0 and V1
expr = "V0 * 0.5 + V1 * 0.5"
math_guider = MathGuider(V, F, expr)
math_guider = MathGuider(V, F, expr, expr) # Add expression1 parameter
# Pseudo input
x = torch.zeros((1, 4, 16, 16))
@@ -48,7 +48,7 @@ class TestMathGuider(unittest.TestCase):
# a = V0, w = F0
expr = "a + w"
math_guider = MathGuider(V, F, expr)
math_guider = MathGuider(V, F, expr, expr) # Add expression1 parameter
x = torch.zeros((1, 4, 8, 8))
sigma = torch.tensor(1.0)
@@ -60,7 +60,7 @@ class TestMathGuider(unittest.TestCase):
g0 = MockGuider(1.0)
V = {"V0": g0}
F = {}
math_guider = MathGuider(V, F, "V0")
math_guider = MathGuider(V, F, "V0", "V0") # Add expression1 parameter
# Check if the property exists and matches g0's patcher
self.assertIsNotNone(math_guider.model_patcher)
@@ -72,7 +72,7 @@ class TestMathGuider(unittest.TestCase):
del g0.model_patcher # force remove
V = {"V0": g0}
F = {}
math_guider = MathGuider(V, F, "V0")
math_guider = MathGuider(V, F, "V0", "V0") # Add expression1 parameter
self.assertIsNone(math_guider.model_patcher)
def test_math_guider_steps_context(self):
@@ -81,7 +81,8 @@ class TestMathGuider(unittest.TestCase):
g0 = MockGuider(1.0)
V = {"V0": g0}
math_guider = MathGuider(V, {}, "current_step / steps")
expr = "current_step / steps"
math_guider = MathGuider(V, {}, expr, expr) # Add expression1 parameter
math_guider.sigmas = sigmas # sets sigmas directly for testing
# Step 0: sigma = 10.0
+6 -3
View File
@@ -10,7 +10,8 @@ def test_image_mismatch_broadcast_a_longer():
b = torch.ones((1, 64, 64, 3)) * 0.5
# a + b -> [1+0.5, 1+0.5]
result, = ImageMathNode.execute(V={"V0": a, "V1": b}, F={}, Expression="a + b", length_mismatch="tile")
result_list, stack = ImageMathNode.execute(V={"V0": a, "V1": b}, F={}, Expression="a + b", length_mismatch="tile")
result = result_list[0]
assert result.shape[0] == 2
assert torch.allclose(result, torch.ones((2, 64, 64, 3)) * 1.5)
@@ -20,7 +21,8 @@ def test_image_mismatch_broadcast_b_longer():
b = torch.ones((2, 64, 64, 3)) * 0.5
# a + b -> [1+0.5, 1+0.5] (a broadcasted to length 2)
result, = ImageMathNode.execute(V={"V0": a, "V1": b}, F={}, Expression="a + b", length_mismatch="tile")
result_list, stack = ImageMathNode.execute(V={"V0": a, "V1": b}, F={}, Expression="a + b", length_mismatch="tile")
result = result_list[0]
assert result.shape[0] == 2
assert torch.allclose(result, torch.ones((2, 64, 64, 3)) * 1.5)
@@ -37,7 +39,8 @@ def test_image_mismatch_pad_b_longer():
b = torch.ones((2, 64, 64, 3)) * 0.5
# a + b -> [1+0.5, 0+0.5] = [1.5, 0.5]
result, = ImageMathNode.execute(V={"V0": a, "V1": b}, F={}, Expression="a + b", length_mismatch="pad")
result_list, stack = ImageMathNode.execute(V={"V0": a, "V1": b}, F={}, Expression="a + b", length_mismatch="pad")
result = result_list[0]
assert result.shape[0] == 2
assert torch.allclose(result[0], torch.ones((64, 64, 3)) * 1.5)
assert torch.allclose(result[1], torch.ones((64, 64, 3)) * 0.5)