diff --git a/tests/reproduce_indexing.py b/tests/reproduce_indexing.py new file mode 100644 index 0000000..5e2c8a5 --- /dev/null +++ b/tests/reproduce_indexing.py @@ -0,0 +1,80 @@ +import torch +import unittest +import sys +import os + +# Add the project root to sys.path +sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), '..'))) + +from more_math.Parser.UnifiedMathVisitor import UnifiedMathVisitor +import antlr4 +from more_math.Parser.MathExprLexer import MathExprLexer +from more_math.Parser.MathExprParser import MathExprParser + +class TestIndexing(unittest.TestCase): + def evaluate(self, expr, variables): + input_stream = antlr4.InputStream(expr) + lexer = MathExprLexer(input_stream) + stream = antlr4.CommonTokenStream(lexer) + parser = MathExprParser(stream) + tree = parser.start() + visitor = UnifiedMathVisitor(variables) + return visitor.visit(tree) + + def test_tensor_indexing(self): + v = torch.randn(4, 4) + variables = {'v': v} + + # Single index + res = self.evaluate('v[0];', variables) + self.assertTrue(torch.allclose(res, v[0])) + + # Tuple index + res = self.evaluate('v[1, 2];', variables) + self.assertEqual(res, float(v[1, 2])) + + # List selection + res = self.evaluate('v[[0, 2]];', variables) + self.assertTrue(torch.allclose(res, v[[0, 2]])) + + def test_list_indexing(self): + l = [10, 20, 30, 40] + variables = {'l': l} + + # Single index + res = self.evaluate('l[0];', variables) + self.assertEqual(res, 10) + + # Negative index + res = self.evaluate('l[-1];', variables) + self.assertEqual(res, 40) + + # List selection + res = self.evaluate('l[[0, 2]];', variables) + self.assertEqual(res, [10, 30]) + + # Nested list indexing + nl = [[1, 2], [3, 4]] + variables['nl'] = nl + res = self.evaluate('nl[1, 0];', variables) + self.assertEqual(res, 3) + + def test_index_error(self): + v = torch.randn(4) + variables = {'v': v} + try: + self.evaluate('v[10];', variables) + except Exception as e: + print(f"Caught indexing error: {e}") + self.assertIn("1:0:", str(e)) + + l = [1, 2, 3] + variables = {'l': l} + try: + self.evaluate('l[10];', variables) + except Exception as e: + print(f"Caught list indexing error: {e}") + self.assertIn("1:0:", str(e)) + +if __name__ == '__main__': + unittest.main() diff --git a/tests/test_bitwise_comprehensive.py b/tests/test_bitwise_comprehensive.py new file mode 100644 index 0000000..d633a45 --- /dev/null +++ b/tests/test_bitwise_comprehensive.py @@ -0,0 +1,180 @@ +""" +Comprehensive integration test for fp16/int16 bitwise operations. +Demonstrates real-world usage of the enhanced bitwise operators. +""" + +import torch +import sys +import os + +sys.path.insert(0, os.path.dirname(__file__)) +from more_math.Parser.UnifiedMathVisitor import UnifiedMathVisitor + + +def test_mixed_dtype_operations(): + """Test operations with mixed dtypes.""" + print("=" * 70) + print("Testing Mixed Dtype Operations") + print("=" * 70) + + # int8 operations + print("\n1. INT8 Bitwise Operations:") + a_int8 = torch.tensor([15, 7, 3], dtype=torch.int8) + visitor = UnifiedMathVisitor({"a": a_int8}) + result = visitor._bitwise_not(a_int8) + print(f" Input: {a_int8} (dtype={a_int8.dtype})") + print(f" ~Input: {result} (dtype={result.dtype})") + assert result.dtype == torch.int8 + + # int16 operations + print("\n2. INT16 Bitwise Operations:") + a_int16 = torch.tensor([255, 127, 63], dtype=torch.int16) + b_int16 = torch.tensor([15, 31, 7], dtype=torch.int16) + visitor = UnifiedMathVisitor({"a": a_int16, "b": b_int16}) + + result_and = visitor._bitwise_op(a_int16, b_int16, torch.bitwise_and, lambda x, y: x & y) + result_or = visitor._bitwise_op(a_int16, b_int16, torch.bitwise_or, lambda x, y: x | y) + result_xor = visitor._bitwise_op(a_int16, b_int16, torch.bitwise_xor, lambda x, y: x ^ y) + + print(f" a: {a_int16} (dtype={a_int16.dtype})") + print(f" b: {b_int16} (dtype={b_int16.dtype})") + print(f" a & b: {result_and}") + print(f" a | b: {result_or}") + print(f" a ^ b: {result_xor}") + assert all(t.dtype == torch.int16 for t in [result_and, result_or, result_xor]) + + # fp16 operations + print("\n3. FP16 Bitwise Operations (bit-level manipulation):") + a_fp16 = torch.tensor([1.0, -2.0, 3.5], dtype=torch.float16) + visitor = UnifiedMathVisitor({"a": a_fp16}) + result = visitor._bitwise_not(a_fp16) + print(f" Input: {a_fp16} (dtype={a_fp16.dtype})") + print(f" Bit-flipped: {result} (dtype={result.dtype})") + assert result.dtype == torch.float16 + + # int32 operations (existing, should still work) + print("\n4. INT32 Bitwise Operations (backward compatibility):") + a_int32 = torch.tensor([65535, 32767, 16383], dtype=torch.int32) + visitor = UnifiedMathVisitor({"a": a_int32}) + result = visitor._bitwise_not(a_int32) + print(f" Input: {a_int32} (dtype={a_int32.dtype})") + print(f" ~Input: {result} (dtype={result.dtype})") + assert result.dtype == torch.int32 + + print("\n" + "=" * 70) + print("✓ All mixed dtype operations completed successfully!") + print("=" * 70) + + +def test_scalar_list_tensor_combinations(): + """Test operations with different input types.""" + print("\n" + "=" * 70) + print("Testing Scalar/List/Tensor Combinations") + print("=" * 70) + + visitor = UnifiedMathVisitor({}) + + # Tensor & Tensor + print("\n1. Tensor & Tensor (int16):") + a = torch.tensor([7, 14, 21], dtype=torch.int16) + b = torch.tensor([3, 5, 7], dtype=torch.int16) + result = visitor._bitwise_op(a, b, torch.bitwise_and, lambda x, y: x & y) + print(f" {a} & {b} = {result}") + + # Tensor & List + print("\n2. Tensor & List (mixed):") + a = torch.tensor([15, 14, 13], dtype=torch.int16) + b = [7, 3, 1] + result = visitor._bitwise_op(a, b, torch.bitwise_and, lambda x, y: x & y) + print(f" Tensor({a}) & List({b}) = {result}") + + # Scalar & List + print("\n3. Scalar & List:") + a = 15 + b = [7, 3, 1] + result = visitor._bitwise_op(a, b, torch.bitwise_and, lambda x, y: x & y) + print(f" {a} & {b} = {result}") + + # List & List + print("\n4. List & List:") + a = [15, 14, 13] + b = [7, 3, 1] + result = visitor._bitwise_op(a, b, torch.bitwise_and, lambda x, y: x & y) + print(f" {a} & {b} = {result}") + + print("\n" + "=" * 70) + print("✓ All scalar/list/tensor combinations work correctly!") + print("=" * 70) + + +def test_element_size_mapping(): + """Test the element size to dtype mapping for all supported types.""" + print("\n" + "=" * 70) + print("Testing Element Size Mapping") + print("=" * 70) + + visitor = UnifiedMathVisitor({}) + + test_dtypes = [ + (torch.int8, "int8", 1), + (torch.int16, "int16", 2), + (torch.int32, "int32", 4), + (torch.int64, "int64", 8), + (torch.float16, "float16", 2), + (torch.float32, "float32", 4), + (torch.float64, "float64", 8), + ] + + print("\nDtype -> Element Size -> View Dtype Mapping:") + for dtype, name, expected_elem_size in test_dtypes: + tensor = torch.zeros(1, dtype=dtype) + elem_size = tensor.element_size() + view_dtype = visitor._get_bitwise_view_dtype(elem_size) + + # Determine expected view dtype + if elem_size == 1: + expected_view = "int8" + elif elem_size == 2: + expected_view = "int16" + elif elem_size == 4: + expected_view = "int32" + else: + expected_view = "int64" + + print(f" {name:12} -> {elem_size} byte(s) -> {str(view_dtype):18}") + assert elem_size == expected_elem_size + + print("\n" + "=" * 70) + print("✓ Element size mapping is correct!") + print("=" * 70) + + +if __name__ == "__main__": + print("\n") + print("#" * 70) + print("# FP16/INT16 Bitwise Operations - Comprehensive Integration Test") + print("#" * 70) + + try: + test_element_size_mapping() + test_mixed_dtype_operations() + test_scalar_list_tensor_combinations() + + print("\n") + print("#" * 70) + print("# ✓ ALL TESTS PASSED SUCCESSFULLY") + print("#" * 70) + print("\nSummary:") + print(" ✓ fp16 (float16) bitwise operations work correctly") + print(" ✓ int16 (int16) bitwise operations work correctly") + print(" ✓ int8, int32, int64 operations maintained") + print(" ✓ Float dtypes supported through bit-pattern manipulation") + print(" ✓ Mixed tensor/list/scalar operations supported") + print(" ✓ Backward compatibility preserved") + print("\n") + + except Exception as e: + print(f"\n✗ Test failed: {e}") + import traceback + traceback.print_exc() + sys.exit(1) diff --git a/tests/test_bitwise_fp16.py b/tests/test_bitwise_fp16.py new file mode 100644 index 0000000..e513105 --- /dev/null +++ b/tests/test_bitwise_fp16.py @@ -0,0 +1,124 @@ +#!/usr/bin/env python3 +"""Test script to verify fp16 and int16 support in bitwise operations.""" + +import sys +import os +import torch + +# Add the custom_nodes/more_math directory to path +sys.path.insert(0, os.path.dirname(__file__)) + +from more_math.Parser.UnifiedMathVisitor import UnifiedMathVisitor + +def test_bitwise_fp16(): + """Test bitwise NOT operation with fp16 tensors.""" + print("Testing bitwise operations with fp16...") + + # Create fp16 tensors + a_fp16 = torch.tensor([1.5, 2.5, 3.5], dtype=torch.float16) + + variables = {"a": a_fp16} + visitor = UnifiedMathVisitor(variables, shape=(3,)) + + # Test bitwise NOT + result = visitor._bitwise_not(a_fp16) + print(f" fp16 tensor: {a_fp16}") + print(f" bitwise_not result dtype: {result.dtype}") + assert result.dtype == torch.float16, f"Expected float16, got {result.dtype}" + print(" ✓ fp16 bitwise NOT test passed") + +def test_bitwise_int16(): + """Test bitwise NOT operation with int16 tensors.""" + print("Testing bitwise operations with int16...") + + # Create int16 tensors + a_int16 = torch.tensor([1, 2, 3], dtype=torch.int16) + + variables = {"a": a_int16} + visitor = UnifiedMathVisitor(variables, shape=(3,)) + + # Test bitwise NOT + result = visitor._bitwise_not(a_int16) + print(f" int16 tensor: {a_int16}") + print(f" bitwise_not result dtype: {result.dtype}") + assert result.dtype == torch.int16, f"Expected int16, got {result.dtype}" + print(" ✓ int16 bitwise NOT test passed") + +def test_bitwise_and_int16(): + """Test bitwise AND operation with int16 tensors.""" + print("Testing bitwise AND with int16...") + + a = torch.tensor([7, 14, 21], dtype=torch.int16) + b = torch.tensor([3, 5, 7], dtype=torch.int16) + + variables = {"a": a, "b": b} + visitor = UnifiedMathVisitor(variables) + + # Test bitwise AND + result = visitor._bitwise_op(a, b, torch.bitwise_and, lambda x, y: x & y) + print(f" a: {a}, b: {b}") + print(f" bitwise_and result: {result}") + assert result.dtype in [torch.int16, torch.int32], f"Unexpected dtype {result.dtype}" + print(" ✓ int16 bitwise AND test passed") + +def test_bitwise_int8(): + """Test bitwise operations with int8 tensors.""" + print("Testing bitwise operations with int8...") + + a_int8 = torch.tensor([1, 2, 3], dtype=torch.int8) + + variables = {"a": a_int8} + visitor = UnifiedMathVisitor(variables, shape=(3,)) + + # Test bitwise NOT + result = visitor._bitwise_not(a_int8) + print(f" int8 tensor: {a_int8}") + print(f" bitwise_not result dtype: {result.dtype}") + assert result.dtype == torch.int8, f"Expected int8, got {result.dtype}" + print(" ✓ int8 bitwise NOT test passed") + +def test_get_bitwise_view_dtype(): + """Test the element size to dtype mapping function.""" + print("Testing _get_bitwise_view_dtype...") + + visitor = UnifiedMathVisitor({}) + + # Test various element sizes + test_cases = [ + (1, torch.int8), + (2, torch.int16), + (4, torch.int32), + (8, torch.int64), + ] + + for elem_size, expected_dtype in test_cases: + result_dtype = visitor._get_bitwise_view_dtype(elem_size) + print(f" element_size={elem_size} -> {result_dtype}") + assert result_dtype == expected_dtype, f"Expected {expected_dtype}, got {result_dtype}" + + print(" ✓ All element size mappings correct") + +if __name__ == "__main__": + print("=" * 60) + print("Testing fp16 and int16 support in bitwise operations") + print("=" * 60) + + try: + test_get_bitwise_view_dtype() + print() + test_bitwise_int8() + print() + test_bitwise_int16() + print() + test_bitwise_fp16() + print() + test_bitwise_and_int16() + print() + print("=" * 60) + print("✓ All tests passed successfully!") + print("=" * 60) + except Exception as e: + print(f"\n✗ Test failed with error: {e}") + import traceback + traceback.print_exc() + sys.exit(1) diff --git a/tests/test_error_messages.py b/tests/test_error_messages.py new file mode 100644 index 0000000..edd9942 --- /dev/null +++ b/tests/test_error_messages.py @@ -0,0 +1,104 @@ +import sys +import os +import torch + +# 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.start() + visitor = UnifiedMathVisitor(variables, (1, 64, 64)) + return visitor.visit(tree) + +def test_variable_not_found_error(): + try: + parse_and_visit("x + 1", {}) + except ValueError as e: + msg = str(e) + assert "1:0: Variable 'x' not found" in msg + else: + raise AssertionError("Expected ValueError not raised") + +def test_unknown_function_error(): + try: + parse_and_visit("unknown_func(1)", {}) + except ValueError as e: + msg = str(e) + assert "1:0: Unknown function: unknown_func" in msg + else: + raise AssertionError("Expected ValueError not raised") + +def test_get_value_invalid_arg_error(): + try: + parse_and_visit("get_value(1, [0])", {}) + except ValueError as e: + msg = str(e) + assert "1:0: get_value expects a tensor as first argument" in msg + else: + raise AssertionError("Expected ValueError not raised") + +def test_conv_arg_count_error(): + try: + parse_and_visit("conv(1, 1)", {}) + except ValueError as e: + msg = str(e) + assert "1:0: conv() requires at least 3 arguments" in msg + else: + raise AssertionError("Expected ValueError not raised") + +def test_pop_empty_slot_error(): + try: + parse_and_visit("stack_pop(0)", {}) + except ValueError as e: + msg = str(e) + assert "1:0: Pop from empty slot: 0" in msg + else: + raise AssertionError("Expected ValueError not raised") + +def test_indexed_assignment_not_found_error(): + try: + parse_and_visit("y[0] = 1", {}) + except ValueError as e: + msg = str(e) + assert "1:0: Variable 'y' not found for indexed assignment." in msg + else: + raise AssertionError("Expected ValueError not raised") + +def run_tests(): + tests = [ + test_variable_not_found_error, + test_unknown_function_error, + test_get_value_invalid_arg_error, + test_conv_arg_count_error, + test_pop_empty_slot_error, + test_indexed_assignment_not_found_error + ] + passed = 0 + for test in tests: + print(f"Running {test.__name__}...") + try: + test() + print(" PASSED") + passed += 1 + except Exception as e: + print(f" FAILED: {e}") + import traceback + traceback.print_exc() + + print(f"\n{passed}/{len(tests)} tests passed.") + if passed < len(tests): + sys.exit(1) + +if __name__ == "__main__": + run_tests() diff --git a/tests/test_permute_fix.py b/tests/test_permute_fix.py new file mode 100644 index 0000000..9f11372 --- /dev/null +++ b/tests/test_permute_fix.py @@ -0,0 +1,51 @@ +import torch +import sys +import os + +# Add paths +sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..", ".."))) # ComfyUI root +sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "more_math", "Parser"))) + +import optical_flow_utils as ofu +from comfy.nested_tensor import NestedTensor + +def test_permute_and_apply(): + print("\n--- Testing Permute and Flow Apply ---") + + # 1. Test Permute on NestedTensor + print("Testing NestedTensor.permute...") + nt = NestedTensor([torch.randn(1, 256, 256, 3) for _ in range(2)]) + try: + # Reorder [B, H, W, C] to [B, C, H, W] + nt_permuted = nt.permute(0, 3, 1, 2) + print(f"Original shape: {nt.shape}") + print(f"Permuted shape: {nt_permuted.shape}") + if nt_permuted.shape == (1, 3, 256, 256): + print("NestedTensor.permute SUCCESS!") + else: + print(f"NestedTensor.permute FAILED: Got {nt_permuted.shape}") + except Exception as e: + print(f"NestedTensor.permute FAILED with error: {e}") + import traceback + traceback.print_exc() + + # 2. Test apply_flow + print("\nTesting ofu.apply_flow...") + img = torch.randn(1, 128, 128, 3) + flow = torch.zeros(1, 128, 128, 2) # Zero flow should be identity + flow[..., 0] = 10.0 # Shift 10px right + + try: + warped = ofu.apply_flow(img, flow) + print(f"Warped shape: {warped.shape}") + if warped.shape == img.shape: + print("ofu.apply_flow SUCCESS!") + else: + print(f"ofu.apply_flow FAILED: Got {warped.shape}") + except Exception as e: + print(f"ofu.apply_flow FAILED with error: {e}") + import traceback + traceback.print_exc() + +if __name__ == "__main__": + test_permute_and_apply() diff --git a/test_procedural.py b/tests/test_procedural.py similarity index 100% rename from test_procedural.py rename to tests/test_procedural.py diff --git a/tests/test_signal_stack.py b/tests/test_signal_stack.py new file mode 100644 index 0000000..ead9a06 --- /dev/null +++ b/tests/test_signal_stack.py @@ -0,0 +1,104 @@ +import torch +import sys +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, variables, state_storage=None): + from antlr4.error.ErrorListener import ErrorListener + class ThrowingErrorListener(ErrorListener): + def syntaxError(self, recognizer, offendingSymbol, line, column, msg, e): + raise Exception(f"line {line}:{column}: {msg}") + + input_stream = InputStream(expr) + lexer = MathExprLexer(input_stream) + lexer.removeErrorListeners() + lexer.addErrorListener(ThrowingErrorListener()) + + stream = CommonTokenStream(lexer) + parser = MathExprParser(stream) + parser.removeErrorListeners() + parser.addErrorListener(ThrowingErrorListener()) + + tree = parser.start() + + visitor = UnifiedMathVisitor(variables, (1,), state_storage=state_storage) + return visitor.visit(tree) + +def test_all(): + with open("tests/test_debug.log", "w") as f_log: + def log(msg): + print(msg) + f_log.write(msg + "\n") + f_log.flush() + + # Mock the visitor's print if needed, but the trampoline prints to stdout. + # We'll just run the tests and then read the log. + log("Starting tests...") + + # 1. Return Bubbling + program1 = "f(x) -> { return (x + (1 * 2 / 0.5)); }; f(10);" + res1 = parse_and_visit(program1, {}) + log(f"Test 1 (Return): res={res1}, type={type(res1)}") + assert abs(res1 - 14.0) < 1e-6 + + # 2. Stack Storage + state2 = [] + program2 = "stack_push(0, 42); stack_get(0);" + res2 = parse_and_visit(program2, {}, state_storage=state2) + log(f"Test 2 (Stack): res={res2}, type={type(res2)}") + assert res2 == 42.0 + + # 3. Block Cleanup + variables3 = {"x": 1.0} + program3 = "f() -> { { y = 10; return y; } }; f();" + res3 = parse_and_visit(program3, variables3) + log(f"Test 3 (Cleanup): res={res3}, type={type(res3)}") + assert res3 == 10.0 + assert "y" not in variables3 + + # 4. Break Bubbling + program4 = "x = 0; while(x < 10) { z = (x > 2 ? break : 1); x = x + 1; } x;" + res4 = parse_and_visit(program4, {}) + log(f"Test 4 (Break): res={res4}, type={type(res4)}") + assert res4 == 3.0 + + # 5. Count Function + log("Test 5 (Count): list, tensor, scalar") + assert parse_and_visit("count([1, 2, 3])", {}) == 3.0 + assert parse_and_visit("cnt([1, 2, 3])", {}) == 3.0 + assert parse_and_visit("length([1, 2, 3])", {}) == 3.0 + + t = torch.randn(5, 10) + assert parse_and_visit("count(t)", {"t": t}) == 5.0 + assert parse_and_visit("count(42)", {}) == 1.0 + + # 6. Tensor Function + log("Test 6 (Tensor): creation") + t_empty = parse_and_visit("tensor([2, 3], 1.5)", {}) + assert t_empty.shape == (2, 3) + assert torch.all(t_empty == 1.5) + log(f"Test 6 (Tensor): shape={t_empty.shape}, value={t_empty[0,0]}") + + # 7. Batch Shuffle Function + log("Test 7 (Shuffle): reordering") + t_base = torch.tensor([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]]) # 3x2 + # shuffle(t, [0, 0, 2]) -> should be [[1,2], [1,2], [5,6]] + t_shuffled = parse_and_visit("shuffle(t, [0, 0, 2])", {"t": t_base}) + assert t_shuffled.shape == (3, 2) + assert t_shuffled[0, 0] == 1.0 + assert t_shuffled[1, 0] == 1.0 + assert t_shuffled[2, 0] == 5.0 + log(f"Test 7 (Shuffle): shape={t_shuffled.shape}") + + log("All tests passed successfully!") + +if __name__ == "__main__": + try: + test_all() + except Exception as e: + print(f"FAILED: {e}") + import traceback + traceback.print_exc() + sys.exit(1) diff --git a/tests/test_unified_math.py b/tests/test_unified_math.py index 2de4fec..2e1e0a3 100644 --- a/tests/test_unified_math.py +++ b/tests/test_unified_math.py @@ -273,8 +273,13 @@ def test_quartil(): 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)) + res_q2 = parse_and_visit("quartile(t, 2)", vars) + if not isinstance(res_q2, torch.Tensor): res_q2 = torch.tensor(res_q2) + assert torch.allclose(res_q2, torch.tensor(5.0)) + + res_q1 = parse_and_visit("quartile(t, 1)", vars) + if not isinstance(res_q1, torch.Tensor): res_q1 = torch.tensor(res_q1) + assert torch.allclose(res_q1, torch.tensor(2.5)) # Float inputs for quartil should be cast to int, so 0.5 -> 0 -> Min # verifying strict behavior or fallback @@ -290,7 +295,10 @@ def test_percentile(): # 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)) + + res_p75 = parse_and_visit("percentile(t, 75)", vars) + if not isinstance(res_p75, torch.Tensor): res_p75 = torch.tensor(res_p75) + assert torch.allclose(res_p75, torch.tensor(7.5)) # Aliases assert parse_and_visit("prcnt(l, 50)", vars) == 5.0