tests
This commit is contained in:
@@ -0,0 +1,258 @@
|
||||
import os
|
||||
import sys
|
||||
|
||||
_here = os.path.abspath(os.path.dirname(__file__))
|
||||
_project_root = os.path.abspath(os.path.join(_here, os.pardir))
|
||||
if _project_root not in sys.path:
|
||||
sys.path.insert(0, _project_root)
|
||||
|
||||
_comfy_root = os.path.abspath(os.path.join(_here, "../../.."))
|
||||
if _comfy_root not in sys.path:
|
||||
sys.path.insert(0, _comfy_root)
|
||||
|
||||
import torch
|
||||
import pytest
|
||||
from more_math.Parser.UnifiedMathVisitor import UnifiedMathVisitor
|
||||
|
||||
from more_math.LatentMathNode import LatentMathNode
|
||||
|
||||
def test_conv_1d():
|
||||
"""
|
||||
Test 1D convolution.
|
||||
Input: [Batch, Length, Channels] = [1, 10, 4]
|
||||
Kernel: 1D size 3
|
||||
"""
|
||||
print("\n--- Testing 1D Conv ---")
|
||||
node = LatentMathNode()
|
||||
shape = (1, 10, 4)
|
||||
a_val = torch.randn(*shape)
|
||||
|
||||
# conv(a, 3, 1.0) -> implies kernel of ones, size 3
|
||||
# Result should correspond to 1D conv
|
||||
try:
|
||||
# LatentMathNode expects latent dicts usually
|
||||
input_dict = {"samples": a_val}
|
||||
res = node.execute("conv(a, 3, 1.0)", a=input_dict)
|
||||
|
||||
# LatentMathNode returns list of dicts
|
||||
res_tensor = res[0]["samples"]
|
||||
|
||||
print(f"1D Conv Result Shape: {res_tensor.shape}")
|
||||
# Expect (1, 10, 4)
|
||||
assert res_tensor.shape == shape
|
||||
except Exception as e:
|
||||
print(f"1D Conv Failed: {e}")
|
||||
raise
|
||||
|
||||
def test_conv_3d():
|
||||
"""
|
||||
Test 3D convolution.
|
||||
Input: [Batch, Depth, Height, Width, Channels] = [1, 5, 32, 32, 4]
|
||||
Kernel: 3D size 3x3x3
|
||||
"""
|
||||
print("\n--- Testing 3D Conv ---")
|
||||
node = LatentMathNode()
|
||||
shape = (1, 5, 32, 32, 4)
|
||||
a_val = torch.randn(*shape)
|
||||
|
||||
try:
|
||||
input_dict = {"samples": a_val}
|
||||
res = node.execute("conv(a, 3, 3, 3, 1.0)", a=input_dict)
|
||||
res_tensor = res[0]["samples"]
|
||||
print(f"3D Conv Result Shape: {res_tensor.shape}")
|
||||
assert res_tensor.shape == shape
|
||||
except Exception as e:
|
||||
print(f"3D Conv Failed: {e}")
|
||||
raise
|
||||
|
||||
def test_conv_arbitrary_batch():
|
||||
"""
|
||||
Test generic tensor with extra batch dims.
|
||||
Input: [B1, B2, H, W, C] = [2, 2, 16, 16, 4] -> Should be treated as Batch=4
|
||||
"""
|
||||
print("\n--- Testing Arbitrary Batch ---")
|
||||
node = LatentMathNode()
|
||||
shape = (2, 2, 16, 16, 4)
|
||||
a_val = torch.randn(*shape)
|
||||
|
||||
try:
|
||||
# conv(a, 3, 3, 1.0) -> 2D conv on (16,16)
|
||||
input_dict = {"samples": a_val}
|
||||
res = node.execute("conv(a, 3, 3, 1.0)", a=input_dict)
|
||||
res_tensor = res[0]["samples"]
|
||||
print(f"Arbitrary Batch Result Shape: {res_tensor.shape}")
|
||||
assert res_tensor.shape == shape
|
||||
except Exception as e:
|
||||
print(f"Arbitrary Batch Failed: {e}")
|
||||
raise
|
||||
|
||||
def test_conv_list_kernel():
|
||||
"""
|
||||
Test conv with list kernel (Regression test for float64 mismatch).
|
||||
Kernel: 3x3x3 list of floats.
|
||||
"""
|
||||
print("\n--- Testing List Kernel Conv ---")
|
||||
node = LatentMathNode()
|
||||
shape = (1, 5, 10, 10, 4) # [B, D, H, W, C]
|
||||
a_val = torch.randn(*shape).float()
|
||||
|
||||
# 3x3x3 kernel = 27 elements
|
||||
# Using the user's example kernel
|
||||
kernel_list = [1,1,1,1,0,1,1,1,1, 0,0,0,0,1,0,0,0,0, 1,1,1,1,0,1,1,1,1]
|
||||
kernel_str = str(kernel_list)
|
||||
expr = f"conv(a, 3, 3, 3, {kernel_str})/8"
|
||||
|
||||
try:
|
||||
input_dict = {"samples": a_val}
|
||||
res = node.execute(expr, a=input_dict)
|
||||
res_tensor = res[0]["samples"]
|
||||
print(f"List Kernel Result Shape: {res_tensor.shape}")
|
||||
assert res_tensor.shape == shape
|
||||
assert res_tensor.dtype == torch.float32
|
||||
except Exception as e:
|
||||
print(f"List Kernel Failed: {e}")
|
||||
raise
|
||||
|
||||
def test_conv_audio():
|
||||
"""
|
||||
Test 1D conv on Audio [B, C, L].
|
||||
Input: [1, 2, 100]. Kernel: 3.
|
||||
Should be treated as Channels First -> [B, L, C].
|
||||
Output should preserve Channels First [B, 2, 100].
|
||||
"""
|
||||
print("\n--- Testing Audio Conv [B, C, L] ---")
|
||||
node = LatentMathNode()
|
||||
shape = (1, 2, 100) # [B, C, L] (L >> C)
|
||||
a_val = torch.randn(*shape).float()
|
||||
|
||||
# conv(a, 3, 1.0) on last dim (L)
|
||||
# Expected: result shape same as input
|
||||
try:
|
||||
input_dict = {"samples": a_val}
|
||||
res = node.execute("conv(a, 3, 1.0)", a=input_dict)
|
||||
res_tensor = res[0]["samples"]
|
||||
print(f"Audio Result Shape: {res_tensor.shape}")
|
||||
|
||||
if res_tensor.shape != shape:
|
||||
print(f"Likely interpreted as Channels Last [B, L, C] where C is small? No.")
|
||||
# If interpreted as Channels last [..., C].
|
||||
# [1, 2, 100]. Spatial=[2]. Channel=100.
|
||||
# Output [1, 2, 100] (but confusing channels).
|
||||
pass
|
||||
|
||||
assert res_tensor.shape == shape
|
||||
except Exception as e:
|
||||
print(f"Audio Conv Failed: {e}")
|
||||
raise
|
||||
|
||||
def test_conv_deep_latent():
|
||||
"""
|
||||
Test 3D conv on Deep Latent [B, 32, H, W] (User request).
|
||||
Input: [1, 32, 16, 16]. Kernel: 3x3x3.
|
||||
Should be treated as Channels First -> [B, 32, 16, 16, 1].
|
||||
Depth=32. H=16. W=16.
|
||||
"""
|
||||
print("\n--- Testing Deep Latent Conv [B, 32, H, W] ---")
|
||||
node = LatentMathNode()
|
||||
shape = (1, 32, 16, 16)
|
||||
a_val = torch.randn(*shape).float()
|
||||
|
||||
# conv(a, 3, 3, 3, 1.0)
|
||||
# 3D kernels need D,H,W.
|
||||
# D=32 (Channel). H=16. W=16.
|
||||
try:
|
||||
input_dict = {"samples": a_val}
|
||||
res = node.execute("conv(a, 3, 3, 3, 1.0)", a=input_dict)
|
||||
res_tensor = res[0]["samples"]
|
||||
print(f"Deep Latent Result Shape: {res_tensor.shape}")
|
||||
assert res_tensor.shape == shape
|
||||
|
||||
# Identity check (ensure D neighbors engaged)
|
||||
# Using simple kernel, center only vs ones.
|
||||
# But this test just checks shape and execution path.
|
||||
except Exception as e:
|
||||
print(f"Deep Latent Failed: {e}")
|
||||
raise
|
||||
|
||||
def test_conv_padding():
|
||||
"""
|
||||
Test padding consistency, especially for even kernels.
|
||||
Input: [1, 10, 10, 1]. Kernel: 4x4.
|
||||
Should produce [1, 10, 10, 1] output (Same padding).
|
||||
"""
|
||||
print("\n--- Testing Padding (Even Kernel Size 4) ---")
|
||||
node = LatentMathNode()
|
||||
shape = (1, 10, 10, 1)
|
||||
a_val = torch.randn(*shape).float()
|
||||
|
||||
# conv(a, 4, 4, 1.0)
|
||||
# If padding is symmetric 2, result is 11x11.
|
||||
# If padding is symmetric 1, result is 9x9.
|
||||
# We need asymmetric pad (1, 2) to get 10x10.
|
||||
try:
|
||||
input_dict = {"samples": a_val}
|
||||
res = node.execute("conv(a, 4, 4, 1.0)", a=input_dict)
|
||||
res_tensor = res[0]["samples"]
|
||||
print(f"Padding Test Result Shape: {res_tensor.shape}")
|
||||
assert res_tensor.shape == shape
|
||||
except Exception as e:
|
||||
print(f"Padding Test Failed: {e}")
|
||||
raise
|
||||
|
||||
def test_conv_complex_padding():
|
||||
"""
|
||||
Test asymmetric padding with mixed odd/even kernel sizes.
|
||||
Kernel: (3, 4). Input: (1, 10, 10, 1).
|
||||
Should produce (1, 10, 10, 1).
|
||||
"""
|
||||
print("\n--- Testing Complex Padding (3, 4) ---")
|
||||
node = LatentMathNode()
|
||||
shape = (1, 10, 10, 1)
|
||||
a_val = torch.randn(*shape).float()
|
||||
|
||||
try:
|
||||
input_dict = {"samples": a_val}
|
||||
res = node.execute("conv(a, 3, 4, 1.0)", a=input_dict)
|
||||
res_tensor = res[0]["samples"]
|
||||
print(f"Complex Padding Result Shape: {res_tensor.shape}")
|
||||
assert res_tensor.shape == shape
|
||||
except Exception as e:
|
||||
print(f"Complex Padding Failed: {e}")
|
||||
raise
|
||||
|
||||
def test_conv_3d_asymmetric():
|
||||
"""
|
||||
Test 3D conv with asymmetric spatial dims.
|
||||
Input: [1, 5, 10, 20, 1]. Kernel: 3x3x3.
|
||||
"""
|
||||
print("\n--- Testing 3D Asymmetric Input ---")
|
||||
node = LatentMathNode()
|
||||
shape = (1, 5, 10, 20, 1)
|
||||
a_val = torch.randn(*shape).float()
|
||||
|
||||
try:
|
||||
input_dict = {"samples": a_val}
|
||||
res = node.execute("conv(a, 3, 3, 3, 1.0)", a=input_dict)
|
||||
res_tensor = res[0]["samples"]
|
||||
print(f"3D Asymmetric Result Shape: {res_tensor.shape}")
|
||||
assert res_tensor.shape == shape
|
||||
except Exception as e:
|
||||
print(f"3D Asymmetric Failed: {e}")
|
||||
raise
|
||||
|
||||
if __name__ == "__main__":
|
||||
try:
|
||||
test_conv_1d()
|
||||
test_conv_3d()
|
||||
test_conv_arbitrary_batch()
|
||||
test_conv_list_kernel()
|
||||
test_conv_audio()
|
||||
test_conv_deep_latent()
|
||||
test_conv_padding()
|
||||
test_conv_complex_padding()
|
||||
test_conv_3d_asymmetric()
|
||||
print("All Conv tests passed!")
|
||||
except Exception as e:
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
sys.exit(1)
|
||||
@@ -0,0 +1,145 @@
|
||||
import torch
|
||||
import math
|
||||
import sys
|
||||
import os
|
||||
|
||||
# Add parent dir to sys.path
|
||||
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..")))
|
||||
|
||||
from more_math.helper_functions import eval_tensor_expr, eval_float_expr
|
||||
|
||||
def test_all_functions():
|
||||
# Setup some test data
|
||||
a_val = 2.0
|
||||
b_val = 3.0
|
||||
tensor_a = torch.tensor([1.0, 2.0, 3.0])
|
||||
tensor_b = torch.tensor([0.5, 1.5, 2.5])
|
||||
|
||||
variables = {
|
||||
'a': a_val, 'b': b_val,
|
||||
'ta': tensor_a, 'tb': tensor_b,
|
||||
'x': 0.5, 'y': 1.0, 'z': 2.0
|
||||
}
|
||||
|
||||
# helper for assertions
|
||||
def check(expr, expected_scalar=None, vars=variables):
|
||||
# Test scalar
|
||||
res_s = eval_float_expr(expr, vars)
|
||||
if expected_scalar is not None:
|
||||
if isinstance(res_s, (int, float)):
|
||||
assert abs(res_s - expected_scalar) < 1e-4, f"Scalar {expr} failed: {res_s} != {expected_scalar}"
|
||||
# if expected is tensor we check differently
|
||||
|
||||
# Test tensor
|
||||
res_t = eval_tensor_expr(expr, vars, (3,))
|
||||
assert torch.is_tensor(res_t) or isinstance(res_t, (list, int, float))
|
||||
return res_s, res_t
|
||||
|
||||
print("--- Testing Basic Unary Functions ---")
|
||||
check("sin(0)", 0.0)
|
||||
check("cos(0)", 1.0)
|
||||
check("tan(0)", 0.0)
|
||||
check("asin(0)", 0.0)
|
||||
check("acos(1)", 0.0)
|
||||
check("atan(0)", 0.0)
|
||||
check("sinh(0)", 0.0)
|
||||
check("cosh(0)", 1.0)
|
||||
check("tanh(0)", 0.0)
|
||||
check("asinh(0)", 0.0)
|
||||
# acosh(1) = 0
|
||||
check("acosh(1)", 0.0)
|
||||
check("atanh(0)", 0.0)
|
||||
|
||||
check("abs(-5)", 5.0)
|
||||
check("| -10 |", 10.0) # AbsExp
|
||||
check("sqrt(16)", 4.0)
|
||||
check("ln(e)", 1.0)
|
||||
check("log(100)", 2.0)
|
||||
check("exp(1)", math.e)
|
||||
|
||||
check("floor(1.9)", 1.0)
|
||||
check("ceil(1.1)", 2.0)
|
||||
check("round(1.5)", 2.0)
|
||||
check("gamma(3)", 2.0) # gamma(n) = (n-1)!
|
||||
check("sigm(0)", 0.5)
|
||||
|
||||
check("fract(1.25)", 0.25)
|
||||
check("relu(-5)", 0.0)
|
||||
check("relu(5)", 5.0)
|
||||
check("softplus(0)", math.log(2.0))
|
||||
# gelu(0) = 0
|
||||
check("gelu(0)", 0.0)
|
||||
check("sign(-10)", -1.0)
|
||||
check("sign(10)", 1.0)
|
||||
check("angle(ta)") # test complex angle? no, just ensuring it runs
|
||||
|
||||
print("--- Testing Two-Arg Functions ---")
|
||||
check("pow(2, 3)", 8.0)
|
||||
check("atan2(1, 1)", math.pi/4)
|
||||
check("tmin(5, 10)", 5.0)
|
||||
check("tmax(5, 10)", 10.0)
|
||||
check("step(0.5, 0.2)", 1.0) # step(x, edge) = 1 if x>=edge
|
||||
check("step(0.1, 0.2)", 0.0)
|
||||
|
||||
print("--- Testing Operators ---")
|
||||
check("1 + 2", 3.0)
|
||||
check("5 - 3", 2.0)
|
||||
check("2 * 4", 8.0)
|
||||
check("10 / 2", 5.0)
|
||||
check("7 % 3", 1.0)
|
||||
check("2 ^ 3", 8.0)
|
||||
|
||||
print("--- Testing Boolean/Comparison ---")
|
||||
check("5 > 3", 1.0)
|
||||
check("5 < 3", 0.0)
|
||||
check("5 >= 5", 1.0)
|
||||
check("5 <= 4", 0.0)
|
||||
check("2 == 2", 1.0)
|
||||
check("2 != 3", 1.0)
|
||||
|
||||
print("--- Testing Ternary/N-ary ---")
|
||||
check("clamp(5, 0, 10)", 5.0)
|
||||
check("clamp(-5, 0, 10)", 0.0)
|
||||
check("lerp(0, 10, 0.5)", 5.0)
|
||||
check("smoothstep(0.5, 0, 1)", 0.5) # smoothstep(x, edge0, edge1)
|
||||
|
||||
check("smin(1, 2, 3, 0)", 0.0)
|
||||
check("smax(1, 5, 2)", 5.0)
|
||||
|
||||
print("--- Testing Tensor Specifics (Norm, Map, Conv, FFT) ---")
|
||||
check("tnorm(ta)")
|
||||
check("snorm(ta)")
|
||||
check("map(ta, x)") # 1D map
|
||||
check("conv(ta, 3, 1)") # 1D conv, size 3, value 1
|
||||
check("permute(ta, [0])")
|
||||
|
||||
# FFT/IFFT
|
||||
# We need a shape for FFT usually
|
||||
res_fft = eval_tensor_expr("fft(ta)", variables, (3,))
|
||||
res_ifft = eval_tensor_expr("ifft(fft(ta))", variables, (3,))
|
||||
assert torch.allclose(res_ifft, tensor_a, atol=1e-4)
|
||||
|
||||
# Multi-dim Permute
|
||||
tensor_2d = torch.randn(2, 3)
|
||||
vars_2d = {'t2': tensor_2d}
|
||||
res_perm = eval_tensor_expr("permute(t2, [1, 0])", vars_2d, (3, 2))
|
||||
assert res_perm.shape == (3, 2)
|
||||
|
||||
print("--- Testing List and Constants ---")
|
||||
check("[1, 2, 3] + 1")
|
||||
check("pi", math.pi)
|
||||
check("e", math.e)
|
||||
check("print(1)", 1.0)
|
||||
check("pshp(ta)")
|
||||
|
||||
print("--- Testing Swap ---")
|
||||
# swap(tensor, dim, i, j)
|
||||
# ta = [1, 2, 3]
|
||||
# swap(ta, 0, 0, 2) -> [3, 2, 1]
|
||||
res_swap = eval_tensor_expr("swap(ta, 0, 0, 2)", variables, (3,))
|
||||
assert torch.equal(res_swap, torch.tensor([3.0, 2.0, 1.0]))
|
||||
|
||||
print("All coverage tests passed!")
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_all_functions()
|
||||
@@ -1,5 +1,6 @@
|
||||
import sys
|
||||
import os
|
||||
import comfy_api
|
||||
|
||||
# Ensure test runner (Visual Studio) can import the package regardless of working dir.
|
||||
# If repository uses `src/` layout, add that to sys.path; otherwise add project root.
|
||||
@@ -13,35 +14,14 @@ _comfy_root = os.path.abspath(os.path.join(_here, "../../.."))
|
||||
if _comfy_root not in sys.path:
|
||||
sys.path.insert(0, _comfy_root)
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
# Mock comfy_api
|
||||
try:
|
||||
import comfy_api
|
||||
except ImportError:
|
||||
mock_io = MagicMock()
|
||||
mock_io.ComfyNode = object
|
||||
mock_io.Schema = MagicMock()
|
||||
mock_io.Model = MagicMock()
|
||||
mock_io.Model.Input = MagicMock()
|
||||
mock_io.Model.Output = MagicMock()
|
||||
mock_io.Float = MagicMock()
|
||||
mock_io.Float.Input = MagicMock()
|
||||
mock_io.String = MagicMock()
|
||||
mock_io.String.Input = MagicMock()
|
||||
|
||||
mock_comfy = MagicMock()
|
||||
mock_comfy.latest.io = mock_io
|
||||
sys.modules["comfy_api"] = mock_comfy
|
||||
sys.modules["comfy_api.latest"] = mock_comfy.latest
|
||||
|
||||
import torch
|
||||
from more_math.ModelMathNode import ModelMathNode
|
||||
|
||||
class MockModelPatcher:
|
||||
def __init__(self, state_dict):
|
||||
self.model = MagicMock()
|
||||
self.model.state_dict.return_value = state_dict
|
||||
self.model = comfy_api.Model()
|
||||
self.model.state_dict = state_dict
|
||||
self.patches = {}
|
||||
|
||||
def clone(self):
|
||||
|
||||
+14
-6
@@ -431,17 +431,24 @@ def test_basic_utilities():
|
||||
|
||||
def test_advanced_activations():
|
||||
node = FloatMathNode()
|
||||
# sigmoid(0) = 0.5
|
||||
assert abs(node.execute("sigmoid(0)", a=0.0)[0] - 0.5) < 1e-5
|
||||
# sigm(0) = 0.5
|
||||
assert abs(node.execute("sigm(0)", a=0.0)[0] - 0.5) < 1e-5
|
||||
|
||||
if __name__ == "__main__":
|
||||
import sys
|
||||
try:
|
||||
#test_trig_functions()
|
||||
#test_pow_log_functions()
|
||||
test_conditioning_math_node_initialization()
|
||||
test_conditioning_math_node_metadata()
|
||||
test_latent_math_node_initialization()
|
||||
test_latent_math_node_metadata()
|
||||
test_image_math_node_initialization()
|
||||
test_image_math_node_metadata()
|
||||
test_trig_functions()
|
||||
test_inverse_trig_functions()
|
||||
test_pow_log_functions()
|
||||
test_min_max_functions()
|
||||
#test_basic_utilities()
|
||||
#test_activation_functions()
|
||||
test_basic_utilities()
|
||||
test_advanced_activations()
|
||||
print("All tests passed!")
|
||||
except Exception as e:
|
||||
import traceback
|
||||
@@ -449,3 +456,4 @@ if __name__ == "__main__":
|
||||
sys.exit(1)
|
||||
print("All tests in test_more_math.py passed!")
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,183 @@
|
||||
import sys
|
||||
import os
|
||||
import torch
|
||||
import math
|
||||
import pytest
|
||||
|
||||
# Ensure we can import the module
|
||||
_here = os.path.abspath(os.path.dirname(__file__))
|
||||
_project_root = os.path.abspath(os.path.join(_here, os.pardir))
|
||||
if _project_root not in sys.path:
|
||||
sys.path.insert(0, _project_root)
|
||||
|
||||
# Placeholder import - we will create this file next
|
||||
from more_math.Parser.UnifiedMathVisitor import UnifiedMathVisitor
|
||||
from more_math.Parser.MathExprLexer import MathExprLexer
|
||||
from more_math.Parser.MathExprParser import MathExprParser
|
||||
from antlr4 import InputStream, CommonTokenStream
|
||||
|
||||
def parse_and_visit(expr_str, variables):
|
||||
lexer = MathExprLexer(InputStream(expr_str))
|
||||
stream = CommonTokenStream(lexer)
|
||||
parser = MathExprParser(stream)
|
||||
tree = parser.expr()
|
||||
|
||||
# We might need to pass shape/device if UnifiedMathVisitor requires it for tensor creation
|
||||
# For now assuming it can infer or defaults.
|
||||
# The original TensorEvalVisitor required shape. Unified might need it for "1.0" -> Tensor promotion cases?
|
||||
# Or maybe "1.0" stays scalar till needed?
|
||||
# Let's assume we pass a default shape/device if needed, but for scalar tests we might not need it.
|
||||
|
||||
shape = (1, 1, 1, 1) # Dummy shape
|
||||
visitor = UnifiedMathVisitor(variables, shape)
|
||||
return visitor.visit(tree)
|
||||
|
||||
def test_scalar_ops():
|
||||
vars = {"a": 2.0, "b": 3.0}
|
||||
|
||||
assert parse_and_visit("a + b", vars) == 5.0
|
||||
assert parse_and_visit("a * b", vars) == 6.0
|
||||
assert parse_and_visit("sin(0)", vars) == 0.0
|
||||
assert parse_and_visit("smax(a, b)", vars) == 3.0
|
||||
|
||||
# Type check - ensure they are python float/int, not tensor
|
||||
res = parse_and_visit("a + b", vars)
|
||||
assert isinstance(res, (float, int))
|
||||
|
||||
def test_tensor_ops():
|
||||
t1 = torch.tensor([1.0, 2.0])
|
||||
t2 = torch.tensor([3.0, 4.0])
|
||||
vars = {"t1": t1, "t2": t2, "s": 2.0}
|
||||
|
||||
# Tensor + Tensor
|
||||
res = parse_and_visit("t1 + t2", vars)
|
||||
assert isinstance(res, torch.Tensor)
|
||||
assert torch.allclose(res, torch.tensor([4.0, 6.0]))
|
||||
|
||||
# Tensor + Scalar
|
||||
res2 = parse_and_visit("t1 * s", vars)
|
||||
assert isinstance(res2, torch.Tensor)
|
||||
assert torch.allclose(res2, torch.tensor([2.0, 4.0]))
|
||||
|
||||
# Scalar + Tensor
|
||||
res3 = parse_and_visit("s + t2", vars)
|
||||
assert isinstance(res3, torch.Tensor)
|
||||
assert torch.allclose(res3, torch.tensor([5.0, 6.0]))
|
||||
|
||||
def test_list_broadcasting():
|
||||
# Feature: List * Tensor -> Stack of Tensors
|
||||
t = torch.ones((2, 2)) # 2x2 ones
|
||||
l = [1.0, 2.0, 3.0]
|
||||
vars = {"t": t, "l": l}
|
||||
|
||||
# l * t should produce a stack of 3 tensors: 1*t, 2*t, 3*t
|
||||
# Expected shape: (3, 2, 2)
|
||||
res = parse_and_visit("l * t", vars)
|
||||
|
||||
assert isinstance(res, torch.Tensor)
|
||||
assert res.shape == (3, 2, 2)
|
||||
assert torch.allclose(res[0], t * 1.0)
|
||||
assert torch.allclose(res[1], t * 2.0)
|
||||
assert torch.allclose(res[2], t * 3.0)
|
||||
|
||||
def test_list_scalar_mapping():
|
||||
# Feature: List * Scalar -> List of results
|
||||
l = [1.0, 2.0, 3.0]
|
||||
vars = {"l": l}
|
||||
|
||||
res = parse_and_visit("l * 2", vars)
|
||||
assert isinstance(res, list)
|
||||
assert res == [2.0, 4.0, 6.0]
|
||||
|
||||
def test_func_dispatch():
|
||||
t = torch.tensor([0.0, math.pi/2])
|
||||
vars = {"t": t, "s": 0.0}
|
||||
|
||||
# sin(tensor) -> tensor
|
||||
res_t = parse_and_visit("sin(t)", vars)
|
||||
assert isinstance(res_t, torch.Tensor)
|
||||
assert torch.allclose(res_t, torch.tensor([0.0, 1.0]))
|
||||
|
||||
# sin(scalar) -> scalar
|
||||
res_s = parse_and_visit("sin(s)", vars)
|
||||
assert isinstance(res_s, float)
|
||||
assert abs(res_s) < 1e-6
|
||||
|
||||
# sin(list) -> list
|
||||
l = [0.0, math.pi/2]
|
||||
vars["l"] = l
|
||||
res_l = parse_and_visit("sin(l)", vars)
|
||||
assert isinstance(res_l, list)
|
||||
assert abs(res_l[0]) < 1e-6
|
||||
assert abs(res_l[1] - 1.0) < 1e-6
|
||||
|
||||
def test_power_ops():
|
||||
vars = {"a": 2.0, "b": 3.0}
|
||||
# Scalar ^ Scalar
|
||||
assert parse_and_visit("a ^ b", vars) == 8.0
|
||||
|
||||
# Tensor ^ Scalar
|
||||
t = torch.tensor([2.0, 3.0])
|
||||
vars["t"] = t
|
||||
res = parse_and_visit("t ^ 2", vars)
|
||||
assert torch.allclose(res, torch.tensor([4.0, 9.0]))
|
||||
|
||||
# Scalar ^ Tensor
|
||||
res2 = parse_and_visit("2 ^ t", vars)
|
||||
assert torch.allclose(res2, torch.tensor([4.0, 8.0]))
|
||||
|
||||
def test_hyperbolic_trig():
|
||||
vars = {"s": 0.0}
|
||||
assert parse_and_visit("sinh(s)", vars) == 0.0
|
||||
assert parse_and_visit("cosh(s)", vars) == 1.0
|
||||
assert parse_and_visit("tanh(s)", vars) == 0.0
|
||||
|
||||
t = torch.tensor([0.0])
|
||||
vars["t"] = t
|
||||
assert torch.allclose(parse_and_visit("sinh(t)", vars), torch.tensor([0.0]))
|
||||
assert torch.allclose(parse_and_visit("cosh(t)", vars), torch.tensor([1.0]))
|
||||
|
||||
def test_kernel_coords():
|
||||
# Simulate visitConvFunc context
|
||||
# Usually grid variables are provided by the visitor during visitConvFunc
|
||||
# We can test if they are correctly handled if present in variables
|
||||
grid = torch.linspace(-1, 1, 3)
|
||||
vars = {"kx": grid, "ky": grid}
|
||||
|
||||
# Test if expression using coordinates works
|
||||
res = parse_and_visit("kx^2 + ky^2", vars)
|
||||
assert isinstance(res, torch.Tensor)
|
||||
assert res.shape == grid.shape
|
||||
assert torch.allclose(res, grid**2 + grid**2)
|
||||
|
||||
def test_bool_ops():
|
||||
vars = {"a": 1, "b": 0}
|
||||
# Scalar bool
|
||||
assert parse_and_visit("a > b", vars) == 1
|
||||
assert parse_and_visit("a < b", vars) == 0
|
||||
|
||||
# Tensor bool
|
||||
t1 = torch.tensor([1.0, 0.0])
|
||||
t2 = torch.tensor([0.0, 1.0])
|
||||
vars = {"t1": t1, "t2": t2}
|
||||
res = parse_and_visit("t1 > t2", vars)
|
||||
assert isinstance(res, torch.Tensor)
|
||||
assert torch.all(res == torch.tensor([1.0, 0.0]))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
try:
|
||||
test_scalar_ops()
|
||||
test_tensor_ops()
|
||||
test_list_broadcasting()
|
||||
test_list_scalar_mapping()
|
||||
test_func_dispatch()
|
||||
test_power_ops()
|
||||
test_hyperbolic_trig()
|
||||
test_kernel_coords()
|
||||
test_bool_ops()
|
||||
print("All UnifiedMathVisitor tests passed!")
|
||||
except Exception as e:
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
sys.exit(1)
|
||||
Reference in New Issue
Block a user