AI: tests
This commit is contained in:
@@ -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()
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user