From cf02031ba235dd2e0688f8bf2469731837b41dc4 Mon Sep 17 00:00:00 2001 From: mcDandy Date: Sat, 10 Jan 2026 15:32:04 +0100 Subject: [PATCH] (AI) add tests + forgot to change tests of functionalyty which changed --- tests/reproduce_conv_issues.py | 28 +++---- tests/reproduce_size_error.py | 40 ++++++++++ tests/test_convolution.py | 105 +++++++++++++++++++++++++ tests/test_grammar_coverage.py | 2 +- tests/test_large_quantile.py | 33 ++++++++ tests/test_quantile_ops.py | 136 +++++++++++++++++++++++++++++++++ tests/test_shorthands.py | 79 +++++++++++++++++++ tests/test_stats.py | 65 ++++++++++++++++ tests/test_unified_math.py | 59 +++++++++++--- 9 files changed, 522 insertions(+), 25 deletions(-) create mode 100644 tests/reproduce_size_error.py create mode 100644 tests/test_convolution.py create mode 100644 tests/test_large_quantile.py create mode 100644 tests/test_quantile_ops.py create mode 100644 tests/test_shorthands.py create mode 100644 tests/test_stats.py diff --git a/tests/reproduce_conv_issues.py b/tests/reproduce_conv_issues.py index 03f79dc..fceb959 100644 --- a/tests/reproduce_conv_issues.py +++ b/tests/reproduce_conv_issues.py @@ -26,12 +26,12 @@ def test_conv_1d(): shape = (1, 10, 4) a_val = torch.randn(*shape) - # conv(a, 3, 1.0) -> implies kernel of ones, size 3 + # ezconvolution(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) + res = node.execute("ezconvolution(a, 3, 1.0)", a=input_dict) # LatentMathNode returns list of dicts res_tensor = res[0]["samples"] @@ -57,7 +57,7 @@ def test_conv_3d(): try: input_dict = {"samples": a_val} - res = node.execute("conv(a, 3, 3, 3, 1.0)", a=input_dict) + res = node.execute("ezconvolution(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 @@ -77,9 +77,9 @@ def test_conv_arbitrary_batch(): a_val = torch.randn(*shape) try: - # conv(a, 3, 3, 1.0) -> 2D conv on (16,16) + # ezconvolution(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 = node.execute("ezconvolution(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 @@ -102,7 +102,7 @@ def test_conv_list_kernel(): # 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" + expr = f"ezconvolution(a, 3, 3, 3, {kernel_str})/8" try: input_dict = {"samples": a_val} @@ -128,11 +128,11 @@ def test_conv_audio(): shape = (1, 2, 100) # [B, C, L] (L >> C) a_val = torch.randn(*shape).float() - # conv(a, 3, 1.0) on last dim (L) + # ezconvolution(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 = node.execute("ezconvolution(a, 3, 1.0)", a=input_dict) res_tensor = res[0]["samples"] print(f"Audio Result Shape: {res_tensor.shape}") @@ -161,12 +161,12 @@ def test_conv_deep_latent(): shape = (1, 32, 16, 16) a_val = torch.randn(*shape).float() - # conv(a, 3, 3, 3, 1.0) + # ezconvolution(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 = node.execute("ezconvolution(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 @@ -190,13 +190,13 @@ def test_conv_padding(): shape = (1, 10, 10, 1) a_val = torch.randn(*shape).float() - # conv(a, 4, 4, 1.0) + # ezconvolution(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 = node.execute("ezconvolution(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 @@ -218,7 +218,7 @@ def test_conv_complex_padding(): try: input_dict = {"samples": a_val} - res = node.execute("conv(a, 3, 4, 1.0)", a=input_dict) + res = node.execute("ezconvolution(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 @@ -239,7 +239,7 @@ def test_conv_3d_asymmetric(): try: input_dict = {"samples": a_val} - res = node.execute("conv(a, 3, 3, 3, 1.0)", a=input_dict) + res = node.execute("ezconvolution(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 diff --git a/tests/reproduce_size_error.py b/tests/reproduce_size_error.py new file mode 100644 index 0000000..92fe3c5 --- /dev/null +++ b/tests/reproduce_size_error.py @@ -0,0 +1,40 @@ +import torch +import sys +import os + +# 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) + +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() + shape = (1, 1, 1, 1) + visitor = UnifiedMathVisitor(variables, shape) + return visitor.visit(tree) + +def reproduce(): + print("Reproducing percentile(list, tensor)...") + a = torch.rand(1, 4, 128, 128) * 100.0 # 0-100 range for percentile + l = [0.1, 1.0, 0.3] + + vars = {"a": a, "l": l} + # Note: percentile(l, a) -> a/100.0 will be 0-1. + try: + res = parse_and_visit("percentile(l, a)", vars) + print(f"Success! Result type: {type(res)}, shape: {res.shape}") + except Exception: + import traceback + traceback.print_exc() + +if __name__ == "__main__": + reproduce() diff --git a/tests/test_convolution.py b/tests/test_convolution.py new file mode 100644 index 0000000..159742b --- /dev/null +++ b/tests/test_convolution.py @@ -0,0 +1,105 @@ + +import torch +import sys +import os + +# 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) + +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() + shape = (1, 1, 1, 1) # Dummy shape + visitor = UnifiedMathVisitor(variables, shape) + return visitor.visit(tree) + +def test_ezconvolution_2d(): + # Test 2D convolution with "ezconvolution" (automagic) + # Case: (1, 10, 10, 3) - Standard ComfyUI Image (Batch, H, W, C) + # Kernel: 3x3 + + # Input: 1 batch, 10x10, 3 channels. + t = torch.zeros((1, 10, 10, 3)) + t[:, 5, 5, :] = 1.0 # Centered dot + + vars = {"t": t} + + # Kernel expr: 1.0 (box blur / sum) + # ezconvolution(tensor, kw, kh, k_expr) + # Expect output same shape (1, 10, 10, 3) because ezconvolution handles layout. + + expr = "ezconvolution(t, 3, 3, 1.0)" + res = parse_and_visit(expr, vars) + + assert isinstance(res, torch.Tensor) + assert res.shape == (1, 10, 10, 3) + + # Check value at 5,5. Neighboring 3x3 (9 pixels) should contribute. + # Kernel val is 1.0 everywhere. + # The dot is at 5,5. + # Convolution at 5,5 will sum up the 3x3 area around input 5,5. + # Input 5,5 is 1.0, others 0. + # So result at 5,5 should be 1.0 * 1.0 = 1.0? + # Wait, if kernel is 3x3 ones. Input has single 1. + # Conv is sum(I * K). + # If we are at 5,5. Area is 4,4 to 6,6. + # Input has 1 at 5,5. + # Result at 4,4 will include input 5,5? + # Yes. + # Result at 5,5 will include input 5,5. + + # ezconvolution preserves logic so it should work. + +def test_convolution_2d_strict(): + # Test 2D convolution with "convolution" (strict) + # Case: (1, 3, 10, 10) - PyTorch Standard (Batch, C, H, W) + + t = torch.zeros((1, 3, 10, 10)) + t[:, :, 5, 5] = 1.0 + + vars = {"t": t} + + # Kernel expr: 1.0 + # convolution(tensor, kw, kh, k_expr) + + expr = "convolution(t, 3, 3, 1.0)" + res = parse_and_visit(expr, vars) + + assert isinstance(res, torch.Tensor) + assert res.shape == (1, 3, 10, 10) + +def test_convolution_strict_fail_on_wrong_layout(): + # Test that strict convolution likely fails or produces weird shape on wrong layout + # Input: (1, 10, 10, 3) -> Batch=1, 10, 10, 3 + # Try 2D conv. + # Strict assumes (..., C, H, W). + # so C=10, H=10, W=3. Batch=1. + # Kernel 3x3. + # W=3. Padding for kernel 3 is 1. Padded W = 3+2=5. + # Conv 3 on 5 => valid. + # So it might RUN, but it interprets dimensions wrong. + + t = torch.randn((1, 10, 10, 3)) + vars = {"t": t} + + expr = "convolution(t, 3, 3, 1.0)" + res = parse_and_visit(expr, vars) + + # Result shape should be based on (1, 10, 10, 3) input. + # Batch=(1). C=10. H=10. W=3. + # Output: (1, 10, 10, 3). + # So shape matches input, but semantic is wrong. + # This confirms it didn't permute (if it permuted, it would crash or do something else). + + assert res.shape == (1, 10, 10, 3) + diff --git a/tests/test_grammar_coverage.py b/tests/test_grammar_coverage.py index 980ca03..7cd5a31 100644 --- a/tests/test_grammar_coverage.py +++ b/tests/test_grammar_coverage.py @@ -116,7 +116,7 @@ def test_all_functions(): check("snorm(ta)") check("map(ta, x)") # 1D map # Verify conv with kernel variables kW - check("conv(ta, 3, kW)") # 1D conv, size 3, value is kW (which is 3.0) + check("conv(reshape(ta, [1, 3]), 3, kW)") # 1D conv, Input [1, 3] (Channels=1, Width=3) check("permute(ta, [0])") check("reshape(ta, [3, 1])") diff --git a/tests/test_large_quantile.py b/tests/test_large_quantile.py new file mode 100644 index 0000000..8c143d7 --- /dev/null +++ b/tests/test_large_quantile.py @@ -0,0 +1,33 @@ +import torch +import sys +import os + +# 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) + +from more_math.Parser.UnifiedMathVisitor import UnifiedMathVisitor + +def test_large_quantile_fallback(): + print("Testing large quantile fallback in UnifiedMathVisitor...") + # 36 million elements + size = 36_000_000 + try: + val = torch.rand(size) + q = torch.tensor([0.0, 0.1, 0.5, 0.9, 1.0]) + + visitor = UnifiedMathVisitor({}) + + print(f"Calling _quartile_helper with size {size}...") + # This should trigger the fallback + res = visitor._quartile_helper(val, q) + print(f"Success! Result: {res}") + except Exception as e: + print(f"Failed with error: {e}") + import traceback + traceback.print_exc() + +if __name__ == "__main__": + test_large_quantile_fallback() diff --git a/tests/test_quantile_ops.py b/tests/test_quantile_ops.py new file mode 100644 index 0000000..e6b9b0e --- /dev/null +++ b/tests/test_quantile_ops.py @@ -0,0 +1,136 @@ +import torch +import sys +import os +import math + +# 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) + +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() + shape = (1, 1, 1, 1) + visitor = UnifiedMathVisitor(variables, shape) + return visitor.visit(tree) + +def test_quantile_basic(): + vars = {"t": torch.tensor([0.0, 1.0, 2.0, 3.0, 4.0]), "l": [0, 1, 2, 3, 4]} + + # Quantile (0-1) + assert parse_and_visit("quantile(t, 0.5)", vars) == 2.0 + assert parse_and_visit("quantile(l, 0.25)", vars) == 1.0 + assert parse_and_visit("quantile(t, 1.0)", vars) == 4.0 + assert parse_and_visit("quantile(l, 0)", vars) == 0.0 + +def test_percentile_basic(): + vars = {"t": torch.tensor([0.0, 1.0, 2.0, 3.0, 4.0]), "l": [0, 1, 2, 3, 4]} + + # Percentile (0-100) + assert parse_and_visit("percentile(t, 50)", vars) == 2.0 + assert parse_and_visit("percentile(l, 25)", vars) == 1.0 + assert parse_and_visit("percentile(t, 100)", vars) == 4.0 + assert parse_and_visit("percentile(l, 0)", vars) == 0.0 + +def test_quartile_basic(): + vars = {"t": torch.tensor([0.0, 1.0, 2.0, 3.0, 4.0]), "l": [0, 1, 2, 3, 4]} + + # Quartile (0-4) + assert parse_and_visit("quartile(t, 2)", vars) == 2.0 + assert parse_and_visit("quartile(l, 1)", vars) == 1.0 + assert parse_and_visit("quartile(t, 4)", vars) == 4.0 + assert parse_and_visit("quartile(l, 0)", vars) == 0.0 + + # Alias + assert parse_and_visit("quartil(t, 2)", vars) == 2.0 + +def test_tensor_queries(): + t = torch.tensor([0.0, 10.0, 20.0, 30.0, 40.0]) + q_tensor = torch.tensor([0.0, 0.5, 1.0]) + vars = {"t": t, "q": q_tensor} + + # Quantile with tensor q + res = parse_and_visit("quantile(t, q)", vars) + assert isinstance(res, torch.Tensor) + assert torch.allclose(res, torch.tensor([0.0, 20.0, 40.0])) + + # Percentile with tensor p + p_tensor = torch.tensor([0.0, 50.0, 100.0]) + vars["p"] = p_tensor + res_p = parse_and_visit("percentile(t, p)", vars) + assert torch.allclose(res_p, torch.tensor([0.0, 20.0, 40.0])) + +def test_list_with_tensor_query(): + # Optimization: list input promoted to tensor for tensor query + l = [0, 10, 20, 30, 40] + q = torch.tensor([0.25, 0.75]) + vars = {"l": l, "q": q} + + res = parse_and_visit("quantile(l, q)", vars) + assert isinstance(res, torch.Tensor) + assert torch.allclose(res, torch.tensor([10.0, 30.0])) + +def test_nd_tensor_query(): + # N-D tensor as query should reshape result + t = torch.tensor([0.0, 100.0]) + q_nd = torch.zeros((2, 2, 2)) + q_nd[0, 0, 1] = 1.0 # Max + q_nd[1, 1, 1] = 0.5 # Mid + vars = {"t": t, "q": q_nd} + + res = parse_and_visit("quantile(t, q)", vars) + assert res.shape == (2, 2, 2) + assert res[0, 0, 0] == 0.0 + assert res[0, 0, 1] == 100.0 + assert res[1, 1, 1] == 50.0 + +def test_large_tensor_fallback(): + # Test that the fallback doesn't crash (we can't easily verify it *fell back* + # but we can verify it works on larger ones) + # Using 1M for speed in regular tests, but enough to trust the logic + size = 1_000_000 + t = torch.rand(size) + vars = {"t": t} + + # Simple median + res = parse_and_visit("quantile(t, 0.5)", vars) + # torch.quantile(t, 0.5) should be very close to t.median() + assert math.isclose(res.item(), t.median().item(), abs_tol=1e-3) + +def test_list_query(): + # percentile(t, [0, 10, 20]) -> returns a list of results + t = torch.tensor([0.0, 10.0, 20.0, 30.0, 40.0]) + vars = {"t": t} + + res = parse_and_visit("percentile(t, [0, 50, 100])", vars) + assert isinstance(res, list) + assert len(res) == 3 + assert res[0] == 0.0 + assert res[1] == 20.0 + assert res[2] == 40.0 + +if __name__ == "__main__": + # If run as script + try: + test_quantile_basic() + test_percentile_basic() + test_quartile_basic() + test_tensor_queries() + test_list_with_tensor_query() + test_nd_tensor_query() + test_large_tensor_fallback() + test_list_query() + print("All quantile_ops tests passed!") + except Exception: + import traceback + traceback.print_exc() + sys.exit(1) diff --git a/tests/test_shorthands.py b/tests/test_shorthands.py new file mode 100644 index 0000000..b74187f --- /dev/null +++ b/tests/test_shorthands.py @@ -0,0 +1,79 @@ + +import torch +import sys +import os + +# 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) + +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() + shape = (1, 1, 1, 1) # Dummy shape + visitor = UnifiedMathVisitor(variables, shape) + return visitor.visit(tree) + +def test_shorthamd_ezconv(): + t = torch.zeros((1, 10, 10, 3)) + t[:, 5, 5, :] = 1.0 + vars = {"t": t} + + # ezconv should work exactly like ezconvolution + expr = "ezconv(t, 3, 3, 1.0)" + res = parse_and_visit(expr, vars) + assert isinstance(res, torch.Tensor) + assert res.shape == (1, 10, 10, 3) + +def test_shorthand_conv(): + t = torch.zeros((1, 3, 10, 10)) + t[:, :, 5, 5] = 1.0 + vars = {"t": t} + + # conv should work exactly like convolution (strictly) + expr = "conv(t, 3, 3, 1.0)" + res = parse_and_visit(expr, vars) + assert isinstance(res, torch.Tensor) + assert res.shape == (1, 3, 10, 10) + +def test_shorthand_print_shape(): + t = torch.zeros((1, 3)) + vars = {"t": t} + + # Test print_shape + # visitor prints to stdout, we just check it runs and returns t + res1 = parse_and_visit("print_shape(t)", vars) + assert torch.equal(res1, t) + + # Test pshp + res2 = parse_and_visit("pshp(t)", vars) + assert torch.equal(res2, t) + +def test_shorthand_perm(): + t = torch.randn((2, 3, 4)) + vars = {"t": t} + # permute(t, [2, 0, 1]) + res_full = parse_and_visit("permute(t, [2, 0, 1])", vars) + res_short = parse_and_visit("perm(t, [2, 0, 1])", vars) + + assert res_short.shape == (4, 2, 3) + assert torch.equal(res_full, res_short) + +def test_shorthand_rshp(): + t = torch.randn((2, 3, 4)) # 24 elements + vars = {"t": t} + # reshape(t, [24]) + res_full = parse_and_visit("reshape(t, [24])", vars) + res_short = parse_and_visit("rshp(t, [24])", vars) + + assert res_short.shape == (24,) + assert torch.equal(res_full, res_short) diff --git a/tests/test_stats.py b/tests/test_stats.py new file mode 100644 index 0000000..7bee093 --- /dev/null +++ b/tests/test_stats.py @@ -0,0 +1,65 @@ + +import torch +import sys +import os +import math + +# 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) + +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() + shape = (1, 1, 1, 1) # Dummy shape + visitor = UnifiedMathVisitor(variables, shape) + return visitor.visit(tree) + +def test_sum(): + vars = {"t": torch.tensor([1.0, 2.0, 3.0]), "l": [1.0, 2.0, 3.0], "s": 5.0} + assert parse_and_visit("sum(t)", vars) == 6.0 + assert parse_and_visit("sum(l)", vars) == 6.0 + assert parse_and_visit("sum(s)", vars) == 5.0 + +def test_mean(): + vars = {"t": torch.tensor([1.0, 2.0, 3.0]), "l": [1.0, 2.0, 3.0], "s": 5.0} + assert parse_and_visit("mean(t)", vars) == 2.0 + assert parse_and_visit("mean(l)", vars) == 2.0 + assert parse_and_visit("mean(s)", vars) == 5.0 + +def test_std(): + t = torch.tensor([1.0, 5.0]) # mean=3, var=((2^2 + 2^2)/1) = 8. std = sqrt(8)=2.828... + vars = {"t": t, "l": [1.0, 5.0]} + + assert math.isclose(parse_and_visit("std(t)", vars).item(), math.sqrt(8), rel_tol=1e-5) + assert math.isclose(parse_and_visit("std(l)", vars), math.sqrt(8), rel_tol=1e-5) + +def test_var(): + t = torch.tensor([1.0, 5.0]) + vars = {"t": t, "l": [1.0, 5.0]} + assert math.isclose(parse_and_visit("var(t)", vars).item(), 8.0, rel_tol=1e-5) + assert math.isclose(parse_and_visit("var(l)", vars), 8.0, rel_tol=1e-5) + +def test_dot(): + a = torch.tensor([1.0, 2.0]) + b = torch.tensor([3.0, 4.0]) + vars = {"a": a, "b": b} + # dot = 1*3 + 2*4 = 3 + 8 = 11 + assert parse_and_visit("dot(a, b)", vars) == 11.0 + + # Check flattening behavior (2D tensor) + a2 = torch.tensor([[1.0, 2.0]]) + b2 = torch.tensor([[3.0], [4.0]]) # 2x1 + vars = {"a2": a2, "b2": b2} + # a2 flat = [1, 2], b2 flat = [3, 4] -> dot=11 + assert parse_and_visit("dot(a2, b2)", vars) == 11.0 + diff --git a/tests/test_unified_math.py b/tests/test_unified_math.py index e6bda96..0833e00 100644 --- a/tests/test_unified_math.py +++ b/tests/test_unified_math.py @@ -9,7 +9,8 @@ _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 + +# Import visitor from more_math.Parser.UnifiedMathVisitor import UnifiedMathVisitor from more_math.Parser.MathExprLexer import MathExprLexer from more_math.Parser.MathExprParser import MathExprParser @@ -183,9 +184,6 @@ def test_bool_ops(): def test_topk(): - # Tensor topk - t = torch.tensor([1.0, 5.0, 2.0, 8.0, 3.0]) - vars = {"t": t} # Tensor topk masking t = torch.tensor([1.0, 5.0, 2.0, 8.0, 3.0]) vars = {"t": t} @@ -229,9 +227,9 @@ def test_botk(): def test_pinv(): # List permutation inverse - perm = [2, 0, 1] # 0->2, 1->0, 2->1 - vars = {"perm": perm} - res = parse_and_visit("pinv(perm)", vars) + p = [2, 0, 1] # 0->2, 1->0, 2->1 + vars = {"p": p} + res = parse_and_visit("pinv(p)", vars) assert isinstance(res, list) # Inverse: if perm[i]=j, then inv[j]=i # perm[0]=2 -> inv[2]=0 @@ -240,9 +238,9 @@ def test_pinv(): assert res == [1, 2, 0] # Tensor permutation inverse - perm_t = torch.tensor([2, 0, 1]) - vars["perm_t"] = perm_t - res_t = parse_and_visit("pinv(perm_t)", vars) + pt = torch.tensor([2, 0, 1]) + vars["pt"] = pt + res_t = parse_and_visit("pinv(pt)", vars) assert isinstance(res_t, torch.Tensor) assert torch.equal(res_t, torch.tensor([1, 2, 0])) @@ -253,6 +251,46 @@ def test_pinv_identity(): res = parse_and_visit("permute(permute(c,a),pinv(a))",varbl) assert torch.equal(tensor,res) + +def test_quartil(): + # Test Quartiles (Strict Integer Indices) + # List: [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10] (11 elements) + l = [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10] + vars = {"l": l} + + assert parse_and_visit("quartile(l, 0)", vars) == 0.0 + assert parse_and_visit("quartile(l, 1)", vars) == 2.5 + assert parse_and_visit("quartile(l, 2)", vars) == 5.0 + assert parse_and_visit("quartile(l, 3)", vars) == 7.5 + assert parse_and_visit("quartile(l, 4)", vars) == 10.0 + + # Tensor + t = torch.tensor(l, dtype=torch.float32) + vars["t"] = t + + assert torch.allclose(parse_and_visit("quartile(t, 2)", vars), torch.tensor(5.0)) + assert torch.allclose(parse_and_visit("quartile(t, 1)", vars), torch.tensor(2.5)) + + # Float inputs for quartil should be cast to int, so 0.5 -> 0 -> Min + # verifying strict behavior or fallback + assert parse_and_visit("quartile(l, 0.9)", vars) == 0.0 + + +def test_percentile(): + l = [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10] + vars = {"l": l} + t = torch.tensor(l, dtype=torch.float32) + vars["t"] = t + + # Percentile (0 - 100) + assert parse_and_visit("percentile(l, 50)", vars) == 5.0 # Median + assert parse_and_visit("percentile(l, 25)", vars) == 2.5 # Q1 + assert torch.allclose(parse_and_visit("percentile(t, 75)", vars), torch.tensor(7.5)) + + # Aliases + assert parse_and_visit("prcnt(l, 50)", vars) == 5.0 + + if __name__ == "__main__": try: test_scalar_ops() @@ -267,6 +305,7 @@ if __name__ == "__main__": test_topk() test_botk() test_pinv() + test_quartil() print("All UnifiedMathVisitor tests passed!") except Exception: import traceback