added fft/ifft compatibilty for other types
This commit is contained in:
+1
-1
@@ -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
@@ -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()}")
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user