added fft/ifft compatibilty for other types

This commit is contained in:
mcDandy
2025-12-05 16:42:43 +01:00
parent bf77b39f4b
commit e4faaded4b
9 changed files with 196 additions and 227 deletions
+1 -1
View File
@@ -3,4 +3,4 @@ import sys
# Add the project root directory to Python path
# This allows the tests to import the project
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), '..')))
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), '../src')))
+62 -25
View File
@@ -2,10 +2,11 @@
"""Tests for `more_math` package."""
import pytest
from src.more_math.ConditioningMathNode import ConditioningMathNode
from src.more_math.LatentMathNode import LatentMathNode
from src.more_math.ImageMathNode import ImageMathNode
import unittest
import torch
from more_math.ConditioningMathNode import ConditioningMathNode
from more_math.LatentMathNode import LatentMathNode
from more_math.ImageMathNode import ImageMathNode
import tokenize
from io import StringIO
@@ -25,31 +26,67 @@ def tokenize_expression(expr):
filtered_tokens.append((token_name, tokval.strip()))
return filtered_tokens
def test_conditioning_math_node_initialization():
node = ConditioningMathNode()
assert isinstance(node, ConditioningMathNode)
def test_conditioning_math_node_metadata():
assert ConditioningMathNode.RETURN_TYPES == ("CONDITIONING",)
assert ConditioningMathNode.FUNCTION == "condMathNode"
assert ConditioningMathNode.CATEGORY == "More math"
class TestMoreMath(unittest.TestCase):
def test_conditioning_math_node_initialization(self):
node = ConditioningMathNode()
self.assertIsInstance(node, ConditioningMathNode)
def test_latent_math_node_initialization():
node = LatentMathNode()
assert isinstance(node, LatentMathNode)
def test_conditioning_math_node_metadata(self):
self.assertEqual(ConditioningMathNode.RETURN_TYPES, ["CONDITIONING"])
self.assertEqual(ConditioningMathNode.FUNCTION, "EXECUTE_NORMALIZED")
self.assertEqual(ConditioningMathNode.CATEGORY, "More math")
def test_latent_math_node_metadata():
assert LatentMathNode.RETURN_TYPES == ("LATENT",)
assert LatentMathNode.FUNCTION == "latMathNode"
assert LatentMathNode.CATEGORY == "More math"
def test_latent_math_node_initialization(self):
node = LatentMathNode()
self.assertIsInstance(node, LatentMathNode)
def test_image_math_node_initialization():
node = ImageMathNode()
assert isinstance(node, ImageMathNode)
def test_latent_math_node_metadata(self):
self.assertEqual(LatentMathNode.RETURN_TYPES, ["LATENT"])
self.assertEqual(LatentMathNode.FUNCTION, "EXECUTE_NORMALIZED")
self.assertEqual(LatentMathNode.CATEGORY, "More math")
def test_image_math_node_metadata():
assert ImageMathNode.RETURN_TYPES == ("IMAGE",)
assert ImageMathNode.FUNCTION == "imgMathNode"
assert ImageMathNode.CATEGORY == "More math"
def test_image_math_node_initialization(self):
node = ImageMathNode()
self.assertIsInstance(node, ImageMathNode)
def test_image_math_node_metadata(self):
self.assertEqual(ImageMathNode.RETURN_TYPES, ["IMAGE"])
self.assertEqual(ImageMathNode.FUNCTION, "EXECUTE_NORMALIZED")
self.assertEqual(ImageMathNode.CATEGORY, "More math")
def test_fft_invertibility(self):
# 1. Create random input latent (Batch, Channel, Height, Width)
input_tensor = torch.randn(1, 4, 32, 32, dtype=torch.float32)
input_dict = {"samples": input_tensor}
# 2. Execute ifft(fft(a))
# Note: execute is a classmethod
result = LatentMathNode.execute(
Latent="ifft(fft(a))",
a=input_dict
)
# 3. Get output tensor
output_tensor = result[0]["samples"]
# 4. Check correctness (approximate equality)
self.assertTrue(torch.allclose(input_tensor, output_tensor, atol=1e-5), \
f"Max difference: {(input_tensor - output_tensor).abs().max()}")
def test_image_fft_dims(self):
# Image input is (Batch, Height, Width, Channel)
# Verify 2D FFT works by doing a round trip
input_tensor = torch.randn(1, 32, 32, 3, dtype=torch.float32)
result = ImageMathNode.execute(
Image="ifft(fft(a))",
a=input_tensor
)
output_tensor = result[0]
self.assertEqual(input_tensor.shape, output_tensor.shape)
self.assertTrue(torch.allclose(input_tensor, output_tensor, atol=1e-5), \
f"Image FFT round trip failed. Max diff: {(input_tensor - output_tensor).abs().max()}")