Add tests

This commit is contained in:
mcDandy
2025-12-05 20:41:58 +01:00
parent fa8c415eb9
commit e673ba4349
4 changed files with 229 additions and 4 deletions
+34
View File
@@ -0,0 +1,34 @@
import sys
import os
import torch
sys.path.insert(0, os.path.abspath('src'))
# Add ComfyUI path for dependencies if needed
sys.path.insert(0, os.path.abspath('../../'))
from more_math.LatentMathNode import LatentMathNode
print("Testing 5D tensors...")
node = LatentMathNode()
# Shape: (1, 5, 4, 32, 32)
samples = torch.randn(1, 5, 4, 32, 32)
l_in = {"samples": samples}
try:
print("Test 1: Identity")
res = node.execute("a * 1.0", a=l_in)[0]["samples"]
print("Shape:", res.shape)
assert res.shape == (1, 5, 4, 32, 32)
res_t = node.execute("a + T", a=l_in)[0]["samples"]
print("Shape T:", res_t.shape)
assert torch.allclose(res_t, samples + 5.0), "T addition failed value check"
print("T addition passed strict check")
print("Test 3: FFT")
res_fft = node.execute("ifft(fft(a))", a=l_in)[0]["samples"]
assert torch.allclose(res_fft, samples, atol=1e-5), "FFT roundtrip failed"
print("FFT passed strict check")
except Exception as e:
import traceback
traceback.print_exc()
+1 -1
View File
@@ -121,7 +121,7 @@ class LatentMathNode(io.ComfyNode):
# scalar dims and counts
'W': width_val, 'width': width_val,
'H': height_val, 'height': height_val,
'T': batch_count, 'batch_count': batch_count,
'T': frame_count, 'batch_count': batch_count,
'N': channel_count, 'channel_count': channel_count,
}
+4 -2
View File
@@ -181,6 +181,8 @@ class NoiseExecutor():
if ndim >= 5:
time_dim = -4
frame_count = merged_samples.shape[time_dim] if time_dim is not None else merged_samples.shape[batch_dim]
B = getIndexTensorAlongDim(merged_samples, batch_dim)
W = getIndexTensorAlongDim(merged_samples, width_dim)
H = getIndexTensorAlongDim(merged_samples, height_dim)
@@ -190,13 +192,13 @@ class NoiseExecutor():
'w': self.w, 'x': self.x, 'y': self.y, 'z': self.z,
'B': B, 'X': W, 'Y': H, 'C': C,
'W': merged_samples.shape[width_dim], 'H': merged_samples.shape[height_dim],
'I': merged_samples, 'T': merged_samples.shape[0], 'N': merged_samples.shape[channel_dim],
'I': merged_samples, 'T': frame_count, 'N': merged_samples.shape[channel_dim],
'batch': B, 'width': merged_samples.shape[width_dim], 'height': merged_samples.shape[height_dim], 'channel': C,
'batch_count': merged_samples.shape[0], 'channel_count': merged_samples.shape[1], 'input_latent': merged_samples,
}
if time_dim is not None:
F = getIndexTensorAlongDim(merged_samples, time_dim)
variables.update({'frame': F, 'frame_count': merged_samples.shape[time_dim]})
variables.update({'frame': F, 'frame_count': frame_count})
visitor = TensorEvalVisitor(variables, variables['a'].shape)
merged_result = visitor.visit(self.tree)
+190 -1
View File
@@ -224,7 +224,196 @@ class TestMoreMath(unittest.TestCase):
self.assertAlmostEqual(res, 0.0)
# Test Smoothstep: smoothstep(0, 1, 0.5) -> 0.5
# 0.5*0.5*(3 - 2*0.5) = 0.25 * 2 = 0.5
res = node.execute("smoothstep(0, 1, a)", a=0.5)[0]
self.assertAlmostEqual(res, 0.5)
def test_float_math_extended(self):
node = FloatMathNode()
# Test Fract: fract(1.5) -> 0.5
res = node.execute("fract(a)", a=1.5)[0]
self.assertAlmostEqual(res, 0.5)
# Test Softplus: softplus(0) = log(1+exp(0)) = log(2) ~= 0.6931
res = node.execute("softplus(a)", a=0.0)[0]
self.assertAlmostEqual(res, 0.69314718, places=5)
# Test Sign: sign(-10) -> -1.0, sign(10) -> 1.0, sign(0) -> 0.0
self.assertEqual(node.execute("sign(a)", a=-10.0)[0], -1.0)
self.assertEqual(node.execute("sign(a)", a=10.0)[0], 1.0)
self.assertEqual(node.execute("sign(a)", a=0.0)[0], 0.0)
# Test Gelu: gelu(0) -> 0.0
self.assertEqual(node.execute("gelu(a)", a=0.0)[0], 0.0)
def test_latent_math_extended(self):
node = LatentMathNode()
l_a = {"samples": torch.zeros(1, 4, 32, 32)}
# Smoothstep
# smoothstep(0, 1, 0.5) -> 0.5
res = node.execute("smoothstep(0, 1, 0.5)", a=l_a)[0]["samples"]
self.assertTrue(torch.allclose(res, torch.full_like(res, 0.5)))
# Softplus
res = node.execute("softplus(0.0)", a=l_a)[0]["samples"]
self.assertTrue(torch.allclose(res, torch.full_like(res, 0.69314718)))
# Gelu
res = node.execute("gelu(0.0)", a=l_a)[0]["samples"]
self.assertTrue(torch.allclose(res, torch.zeros_like(res)))
def test_image_math_operations(self):
node = ImageMathNode()
# Image is (Batch, Height, Width, Channel)
# Create Red image [1, 0, 0] and Blue image [0, 0, 1]
img_red = torch.tensor([1.0, 0.0, 0.0]).view(1, 1, 1, 3)
img_blue = torch.tensor([0.0, 0.0, 1.0]).view(1, 1, 1, 3)
# Test Lerp (Blend)
# lerp(red, blue, 0.5) -> [0.5, 0.0, 0.5] (Purple)
res_blend = node.execute("lerp(a, b, 0.5)", a=img_red, b=img_blue)[0]
expected = torch.tensor([0.5, 0.0, 0.5]).view(1, 1, 1, 3)
self.assertTrue(torch.allclose(res_blend, expected))
# Test Swap Channels (Red <-> Blue)
# Input: Red [1, 0, 0]. Swap R(0) and B(2). Result: [0, 0, 1] (Blue)
# Note: In ImageMathNode, Channel is dim 1 (internally permuted to B,C,H,W)
res_swap = node.execute("swap(a, 1, 0, 2)", a=img_red)[0]
self.assertTrue(torch.allclose(res_swap, img_blue))
def test_nested_expressions(self):
node = FloatMathNode()
# lerp(0, 10, step(0.5, 0.8)) -> lerp(0, 10, 1) -> 10
res = node.execute("lerp(0, 10, step(0.5, 0.8))", a=0.0)[0]
self.assertEqual(res, 10.0)
# lerp(0, 10, step(0.5, 0.2)) -> lerp(0, 10, 0) -> 0
res2 = node.execute("lerp(0, 10, step(0.5, 0.2))", a=0.0)[0]
self.assertEqual(res2, 0.0)
def test_5d_tensors(self):
# Test support for 5D tensors: (Batch, Time, Channel, Height, Width)
# LatentMathNode expects this structure if ndim >= 5.
node = LatentMathNode()
# Create a 5D tensor: 1 Batch, 5 Frames, 4 Channels, 32 Height, 32 Width
# Shape: (1, 5, 4, 32, 32)
samples = torch.randn(1, 5, 4, 32, 32)
l_in = {"samples": samples}
# Test 1: Identity/Passthrough using a basic operation
# a * 1.0 should return same shape and values
res = node.execute("a * 1.0", a=l_in)[0]["samples"]
self.assertEqual(res.shape, (1, 5, 4, 32, 32))
self.assertTrue(torch.allclose(res, samples))
# Test 2: Accessing 'T' (batch_count/time) variable
# In 5D, T maps to dim -4 (size 5)
# Expression: "a + T" -> Adds 5.0 to every element
res_t = node.execute("a + T", a=l_in)[0]["samples"]
self.assertTrue(torch.allclose(res_t, samples + 5.0))
# Test 3: FFT on 5D tensor
# FFT logic operates on dims [2:] -> (Channel, Height, Width) for 5D input??
# Wait, LatentMathNode.eval_single_tensor sets:
# channel_dim = -3, height_dim = -2, width = -1.
# So for (B, T, C, H, W), it processes C, H, W as spatial?
# Let's verify standard roundtrip
res_fft = node.execute("ifft(fft(a))", a=l_in)[0]["samples"]
self.assertTrue(torch.allclose(res_fft, samples, atol=1e-5))
def test_noise_math_5d(self):
# NoiseMathNode takes objects that have a 'generate_noise' method.
# We need to mock this for testing.
from more_math.NoiseMathNode import NoiseMathNode
class MockNoise:
def __init__(self, tensor):
self.tensor = tensor
def generate_noise(self, input_latent):
# Return the stored tensor, ignoring input_latent for this test
return self.tensor
node = NoiseMathNode()
# Create 5D tensor (1, 5, 4, 32, 32)
samples = torch.randn(1, 5, 4, 32, 32)
# In NoiseMathNode, 'input_latent' is passed to generate_noise.
# The node execution logic calls generate_noise on inputs a,b,c,d.
# But wait, NoiseMathNode.execute returns a NoiseExecutor object.
# The actual calculation happens when NoiseExecutor.generate_noise() is called.
noise_a = MockNoise(samples)
# 1. Execute node to get executor
# Noise="a + T"
# We need to verify that T is 5 (frames) not 1 (batch)
result_executor = node.execute("a + T", a=noise_a)[0]
# 2. Call generate_noise on the result executor
# We need to pass a dummy input_latent because generate_noise expects it.
# The input_latent structure dictates dimensions if a,b,c,d don't override?
# Actually NoiseMathNode uses the 'samples' from input_latent as 'merged_samples' usually?
# Let's check logic:
# samples = input_latent["samples"]
# a_val = self.a.generate_noise()
# ...
# if samples is nested...
# else: merged_samples = samples.
#
# So 'merged_samples' comes from 'input_latent' passed to generate_noise!
# The 'a' input is just a generator.
# So to test 5D support, we must pass a 5D latent to generate_noise.
dummy_latent = {"samples": samples} # This sets the context dimensions
# If we pass a MockNoise that returns 'samples', then a_val will be 'samples'.
# And merged_samples will be 'samples'.
# So a + T -> samples + frame_count.
res = result_executor.generate_noise(dummy_latent)
self.assertEqual(res.shape, (1, 5, 4, 32, 32))
res = result_executor.generate_noise(dummy_latent)
self.assertEqual(res.shape, (1, 5, 4, 32, 32))
self.assertTrue(torch.allclose(res, samples + 5.0),
f"Expected samples + 5.0, got {res[0,0,0,0,0]} vs {samples[0,0,0,0,0]+5.0}")
def test_nested_tensor_support(self):
# Test compatibility with comfy.nested_tensor.NestedTensor
try:
from comfy.nested_tensor import NestedTensor
except ImportError:
# If comfy is not in path or fails to import, we might need to skip or mock
# But run_tests.py adds the root, so it should work.
# If it fails, fail the test to alert us.
self.fail("Could not import comfy.nested_tensor. Make sure ComfyUI root is in sys.path")
node = LatentMathNode()
# Create two tensors with different dimensions to simulate nesting
# Tensor 1: (1, 4, 32, 32)
t1 = torch.full((1, 4, 32, 32), 1.0)
# Tensor 2: (2, 4, 32, 32)
t2 = torch.full((2, 4, 32, 32), 2.0)
# Create NestedTensor
nt_in = NestedTensor([t1, t2])
l_in = {"samples": nt_in}
# Operation: a + 1.0
# Expected: t1 becomes 2.0s, t2 becomes 3.0s
res_lat = node.execute("a + 1.0", a=l_in)[0]["samples"]
# Verify result is NestedTensor
self.assertTrue(getattr(res_lat, 'is_nested', False), "Result should be a NestedTensor")
# Verify contents
res_list = res_lat.unbind()
self.assertEqual(len(res_list), 2)
self.assertTrue(torch.allclose(res_list[0], torch.full_like(t1, 2.0)))
self.assertTrue(torch.allclose(res_list[1], torch.full_like(t2, 3.0)))