AI: tests

This commit is contained in:
mcDandy
2026-02-13 18:33:46 +01:00
parent 0d65516cbb
commit d766172383
8 changed files with 654 additions and 3 deletions
+80
View File
@@ -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()
+180
View File
@@ -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)
+124
View File
@@ -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)
+104
View File
@@ -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()
+51
View File
@@ -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()
+104
View File
@@ -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)
+11 -3
View File
@@ -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