From b7f3f9e533a20a4affafce94ed964a2276bceaa6 Mon Sep 17 00:00:00 2001 From: mcDandy Date: Tue, 3 Mar 2026 22:32:19 +0100 Subject: [PATCH] formatting --- more_math/AudioMathNode.py | 4 -- more_math/ConditioningMathNode.py | 1 - more_math/Parser/MathExprParser.py | 76 +++++++++++++------------- more_math/Parser/UnifiedMathVisitor.py | 2 +- more_math/ScriptTextWindow.py | 1 - more_math/helper_functions.py | 20 +++---- test_bitwise_fix.py | 44 +++++++-------- test_bitwise_shifts.py | 48 ++++++++-------- tests/reproduce_indexing.py | 14 ++--- tests/test_bitwise_comprehensive.py | 48 ++++++++-------- tests/test_bitwise_fp16.py | 36 ++++++------ tests/test_conditioning_mismatch.py | 2 +- tests/test_error_messages.py | 3 +- tests/test_guider_math.py | 18 +++--- tests/test_more_math.py | 32 +++++------ tests/test_permute_fix.py | 4 +- tests/test_signal_stack.py | 4 +- tests/test_unified_math.py | 28 +++++----- 18 files changed, 189 insertions(+), 196 deletions(-) diff --git a/more_math/AudioMathNode.py b/more_math/AudioMathNode.py index ee7f88a..3c9530d 100644 --- a/more_math/AudioMathNode.py +++ b/more_math/AudioMathNode.py @@ -1,4 +1,3 @@ -from tokenize import String from .helper_functions import ( generate_dim_variables, parse_expr, @@ -12,9 +11,6 @@ from .helper_functions import ( ) from .Parser.UnifiedMathVisitor import UnifiedMathVisitor from comfy_api.latest import io -from antlr4 import InputStream, CommonTokenStream -from .Parser.MathExprLexer import MathExprLexer -from .Parser.MathExprParser import MathExprParser import torch from .Stack import MrmthStack import copy diff --git a/more_math/ConditioningMathNode.py b/more_math/ConditioningMathNode.py index 6996e51..db93ac3 100644 --- a/more_math/ConditioningMathNode.py +++ b/more_math/ConditioningMathNode.py @@ -1,4 +1,3 @@ -from tkinter import E import torch from .helper_functions import checkLazyNew, generate_dim_variables, parse_expr, getIndexTensorAlongDim, as_tensor, normalize_to_common_shape, make_zero_like, get_v_variable, get_f_variable from .Parser.UnifiedMathVisitor import UnifiedMathVisitor diff --git a/more_math/Parser/MathExprParser.py b/more_math/Parser/MathExprParser.py index ad775cc..60801a1 100644 --- a/more_math/Parser/MathExprParser.py +++ b/more_math/Parser/MathExprParser.py @@ -1024,7 +1024,7 @@ class MathExprParser ( Parser ): self.stmt() pass - + self.state = 69 self._errHandler.sync(self) _alt = self._interp.adaptivePredict(self._input,1,self._ctx) @@ -1061,7 +1061,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_funcDef - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -1325,7 +1325,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_stmt - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -2064,7 +2064,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_ternaryExpr - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -2135,7 +2135,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_compExpr - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -2395,7 +2395,7 @@ class MathExprParser ( Parser ): self.addExpr(0) pass - + self.state = 211 self._errHandler.sync(self) _alt = self._interp.adaptivePredict(self._input,14,self._ctx) @@ -2420,7 +2420,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_addExpr - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -2540,7 +2540,7 @@ class MathExprParser ( Parser ): self.mulExpr(0) pass - + self.state = 225 self._errHandler.sync(self) _alt = self._interp.adaptivePredict(self._input,16,self._ctx) @@ -2565,7 +2565,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_mulExpr - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -2720,7 +2720,7 @@ class MathExprParser ( Parser ): self.shiftExpr(0) pass - + self.state = 242 self._errHandler.sync(self) _alt = self._interp.adaptivePredict(self._input,18,self._ctx) @@ -2745,7 +2745,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_shiftExpr - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -2865,7 +2865,7 @@ class MathExprParser ( Parser ): self.powExpr() pass - + self.state = 256 self._errHandler.sync(self) _alt = self._interp.adaptivePredict(self._input,20,self._ctx) @@ -2890,7 +2890,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_powExpr - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -2983,7 +2983,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_unaryExpr - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -3098,7 +3098,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_indexExpr - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -3226,7 +3226,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_atom - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -3837,7 +3837,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_func0 - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -3897,7 +3897,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_func1 - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -6511,7 +6511,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_func2 - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -8262,7 +8262,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_func3 - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -9046,7 +9046,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_func4 - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -9249,7 +9249,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_func5 - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -9338,7 +9338,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_funcN - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -9746,7 +9746,7 @@ class MathExprParser ( Parser ): def getRuleIndex(self): return MathExprParser.RULE_funcNoise - + def copyFrom(self, ctx:ParserRuleContext): super().copyFrom(ctx) @@ -10835,63 +10835,63 @@ class MathExprParser ( Parser ): def compExpr_sempred(self, localctx:CompExprContext, predIndex:int): if predIndex == 0: return self.precpred(self._ctx, 7) - + if predIndex == 1: return self.precpred(self._ctx, 6) - + if predIndex == 2: return self.precpred(self._ctx, 5) - + if predIndex == 3: return self.precpred(self._ctx, 4) - + if predIndex == 4: return self.precpred(self._ctx, 3) - + if predIndex == 5: return self.precpred(self._ctx, 2) - + def addExpr_sempred(self, localctx:AddExprContext, predIndex:int): if predIndex == 6: return self.precpred(self._ctx, 3) - + if predIndex == 7: return self.precpred(self._ctx, 2) - + def mulExpr_sempred(self, localctx:MulExprContext, predIndex:int): if predIndex == 8: return self.precpred(self._ctx, 4) - + if predIndex == 9: return self.precpred(self._ctx, 3) - + if predIndex == 10: return self.precpred(self._ctx, 2) - + def shiftExpr_sempred(self, localctx:ShiftExprContext, predIndex:int): if predIndex == 11: return self.precpred(self._ctx, 3) - + if predIndex == 12: return self.precpred(self._ctx, 2) - + def indexExpr_sempred(self, localctx:IndexExprContext, predIndex:int): if predIndex == 13: return self.precpred(self._ctx, 2) - + diff --git a/more_math/Parser/UnifiedMathVisitor.py b/more_math/Parser/UnifiedMathVisitor.py index 231f675..efbdbea 100644 --- a/more_math/Parser/UnifiedMathVisitor.py +++ b/more_math/Parser/UnifiedMathVisitor.py @@ -2448,7 +2448,7 @@ class UnifiedMathVisitor(MathExprVisitor): if a.ndim < 1 or b.ndim < 1: raise ValueError("Cross product requires at least 1D tensors") if a.shape[-1] != 3 or b.shape[-1] != 3: - raise ValueError(f"Cross product requires last dimension size = 3") + raise ValueError("Cross product requires last dimension size = 3") # Float8 handling float8_dtypes = { diff --git a/more_math/ScriptTextWindow.py b/more_math/ScriptTextWindow.py index 86848c0..3920cdd 100644 --- a/more_math/ScriptTextWindow.py +++ b/more_math/ScriptTextWindow.py @@ -1,5 +1,4 @@ from comfy_api.latest import io -import torch from .helper_functions import parse_expr from .ParseTree import MrmthParseTree diff --git a/more_math/helper_functions.py b/more_math/helper_functions.py index 4d41492..a0442f8 100644 --- a/more_math/helper_functions.py +++ b/more_math/helper_functions.py @@ -294,40 +294,40 @@ def checkLazyNew(Expression, V, F): else: tree = Expression parser = tree.parser - + # Support aliases aliases = {"a": "V0", "b": "V1", "c": "V2", "d": "V3", "w": "F0", "x": "F1", "y": "F2", "z": "F3"} - + assigned_vars = set() needed_vars = set() - + # Process all top-level statements for child in tree.children: if not hasattr(child, 'getRuleIndex'): continue - + rule_name = parser.ruleNames[child.getRuleIndex()] if child.getRuleIndex() < len(parser.ruleNames) else None - + # Process function definitions: scan for reads but ignore writes if rule_name == 'funcDef': func_params = set() if child.paramList(): for param in child.paramList().VARIABLE(): func_params.add(param.getText()) - + _collect_reads_only(child, needed_vars, assigned_vars, func_params) - + # Top-level assignments elif rule_name == 'varDef': var_name = child.VARIABLE().getText() _collect_vars_from_node(child, needed_vars, assigned_vars, set()) assigned_vars.add(var_name) - + # Track other top-level statements else: _collect_vars_from_node(child, needed_vars, assigned_vars, set()) - + # Normalize variable names through aliases needed = set() for var in needed_vars: @@ -338,5 +338,5 @@ def checkLazyNew(Expression, V, F): needed.update(F.keys()) if re.match(r"[VF][0-9]+", norm): needed.add(norm) - + return needed \ No newline at end of file diff --git a/test_bitwise_fix.py b/test_bitwise_fix.py index 9c57a91..1067e81 100644 --- a/test_bitwise_fix.py +++ b/test_bitwise_fix.py @@ -9,17 +9,17 @@ from custom_nodes.more_math.more_math.Parser.UnifiedMathVisitor import UnifiedMa def test_bitwise_xor_float(): """Test XOR with float tensors (the reported bug)""" print("Testing bitwise XOR with float tensors...") - + # Create float tensors (this was causing the error) a = torch.tensor([1.0, 2.0, 3.0], dtype=torch.float32) b = torch.tensor([4.0, 5.0, 6.0], dtype=torch.float32) - + visitor = UnifiedMathVisitor({"a": a, "b": b}) - + try: # This should convert to int64, perform XOR, then convert back result = visitor._bitwise_op(a, b, torch.bitwise_xor, lambda x, y: x ^ y) - print(f"✓ XOR succeeded!") + print("✓ XOR succeeded!") print(f" Input a (float32): {a}") print(f" Input b (float32): {b}") print(f" Result: {result}") @@ -32,15 +32,15 @@ def test_bitwise_xor_float(): def test_bitwise_and_float(): """Test AND with float tensors""" print("\nTesting bitwise AND with float tensors...") - + a = torch.tensor([15.0, 14.0, 13.0], dtype=torch.float32) b = torch.tensor([7.0, 3.0, 1.0], dtype=torch.float32) - + visitor = UnifiedMathVisitor({"a": a, "b": b}) - + try: result = visitor._bitwise_op(a, b, torch.bitwise_and, lambda x, y: x & y) - print(f"✓ AND succeeded!") + print("✓ AND succeeded!") print(f" Input a (float32): {a}") print(f" Input b (float32): {b}") print(f" Result: {result}") @@ -53,15 +53,15 @@ def test_bitwise_and_float(): def test_bitwise_or_float(): """Test OR with float tensors""" print("\nTesting bitwise OR with float tensors...") - + a = torch.tensor([15.0, 14.0, 13.0], dtype=torch.float32) b = torch.tensor([7.0, 3.0, 1.0], dtype=torch.float32) - + visitor = UnifiedMathVisitor({"a": a, "b": b}) - + try: result = visitor._bitwise_op(a, b, torch.bitwise_or, lambda x, y: x | y) - print(f"✓ OR succeeded!") + print("✓ OR succeeded!") print(f" Input a (float32): {a}") print(f" Input b (float32): {b}") print(f" Result: {result}") @@ -74,14 +74,14 @@ def test_bitwise_or_float(): def test_bitwise_not_float(): """Test NOT with float tensors""" print("\nTesting bitwise NOT with float tensors...") - + a = torch.tensor([1.0, 2.0, 3.0], dtype=torch.float32) - + visitor = UnifiedMathVisitor({"a": a}) - + try: result = visitor._bitwise_not(a) - print(f"✓ NOT succeeded!") + print("✓ NOT succeeded!") print(f" Input a (float32): {a}") print(f" Result: {result}") print(f" Result dtype: {result.dtype}") @@ -93,16 +93,16 @@ def test_bitwise_not_float(): def test_int_tensors_preserved(): """Ensure int tensors still work as before""" print("\nTesting that int tensor dtypes are preserved...") - + a = torch.tensor([15, 14, 13], dtype=torch.int16) b = torch.tensor([7, 3, 1], dtype=torch.int16) - + visitor = UnifiedMathVisitor({"a": a, "b": b}) - + try: result = visitor._bitwise_op(a, b, torch.bitwise_and, lambda x, y: x & y) assert result.dtype == torch.int16, f"Expected int16, got {result.dtype}" - print(f"✓ Int16 dtype preserved!") + print("✓ Int16 dtype preserved!") print(f" Input a (int16): {a}") print(f" Input b (int16): {b}") print(f" Result (int16): {result}") @@ -115,7 +115,7 @@ if __name__ == "__main__": print("=" * 70) print("Bitwise Operations Fix Verification") print("=" * 70) - + results = [ test_bitwise_xor_float(), test_bitwise_and_float(), @@ -123,7 +123,7 @@ if __name__ == "__main__": test_bitwise_not_float(), test_int_tensors_preserved(), ] - + print("\n" + "=" * 70) if all(results): print("✓ All tests passed!") diff --git a/test_bitwise_shifts.py b/test_bitwise_shifts.py index 9c6c672..4faed22 100644 --- a/test_bitwise_shifts.py +++ b/test_bitwise_shifts.py @@ -14,13 +14,13 @@ def parse_and_evaluate(expression, variables=None): """Parse and evaluate a math expression""" if variables is None: variables = {} - + input_stream = InputStream(expression) lexer = MathExprParser(input_stream).lexer stream = CommonTokenFactory() parser = MathExprParser(input_stream) tree = parser.start() - + visitor = UnifiedMathVisitor(variables, device='cpu') result = visitor.visit(tree) return result @@ -30,24 +30,24 @@ def test_bit_shifts(): print("=" * 60) print("Testing Bitwise Shift Operators") print("=" * 60) - + test_cases = [ # Left shift: 5 << 2 = 20 (0101 << 2 = 10100) ("5 << 2", {}, 20), - + # Right shift: 20 >> 2 = 5 (10100 >> 2 = 0101) ("20 >> 2", {}, 5), - + # Left shift with variable ("x << 3", {"x": 4}, 32), # 4 << 3 = 32 - + # Right shift with variable ("x >> 2", {"x": 16}, 4), # 16 >> 2 = 4 - + # Chained shifts ("(8 << 2) >> 3", {}, 4), # (32) >> 3 = 4 ] - + for expr, vars, expected in test_cases: try: result = parse_and_evaluate(expr, vars) @@ -61,27 +61,27 @@ def test_bit_count(): print("\n" + "=" * 60) print("Testing Bitwise Bit Count Function") print("=" * 60) - + test_cases = [ # bitcount(5) = 2 (0101 has 2 set bits) ("bitcount(5)", {}, 2), - + # bitcount(15) = 4 (1111 has 4 set bits) ("bitcount(15)", {}, 4), - + # bitcount(7) = 3 (111 has 3 set bits) ("bitcount(7)", {}, 3), - + # bitcount(255) = 8 (11111111 has 8 set bits) ("bitcount(255)", {}, 8), - + # bitcount(0) = 0 ("bitcount(0)", {}, 0), - + # With variable ("bitcount(x)", {"x": 31}, 5), # 31 = 11111 = 5 bits set ] - + for expr, vars, expected in test_cases: try: result = parse_and_evaluate(expr, vars) @@ -95,12 +95,12 @@ def test_bit_shifts_with_tensors(): print("\n" + "=" * 60) print("Testing Bitwise Shifts with Tensors") print("=" * 60) - + vars = { "a": torch.tensor([1, 2, 4, 8], dtype=torch.int32), "shift": 2, } - + try: result = parse_and_evaluate("a << shift", vars) expected = torch.tensor([4, 8, 16, 32], dtype=torch.int32) @@ -109,7 +109,7 @@ def test_bit_shifts_with_tensors(): print(f"{status} tensor_shift_left: [1,2,4,8] << 2 = {result.tolist()}") except Exception as e: print(f"✗ tensor_shift_left ERROR: {e}") - + try: result = parse_and_evaluate("a >> shift", vars) expected = torch.tensor([0, 0, 1, 2], dtype=torch.int32) @@ -124,11 +124,11 @@ def test_bit_count_with_tensors(): print("\n" + "=" * 60) print("Testing Bit Count with Tensors") print("=" * 60) - + vars = { "nums": torch.tensor([5, 15, 7, 255], dtype=torch.int32), } - + try: result = parse_and_evaluate("bitcount(nums)", vars) # Expected: [2, 4, 3, 8] set bits @@ -141,15 +141,15 @@ def test_combinations(): print("\n" + "=" * 60) print("Testing Combinations") print("=" * 60) - + test_cases = [ # Shift then count bits ("bitcount(5 << 2)", {}, 2), # 5 << 2 = 20 (10100) = 2 bits - + # Combined with other operators ("(5 << 2) | (3 << 4)", {}, 0xCC), # 20 | 48 = 0xCC = 204 ] - + for expr, vars, expected in test_cases: try: result = parse_and_evaluate(expr, vars) @@ -164,7 +164,7 @@ if __name__ == "__main__": test_bit_shifts_with_tensors() test_bit_count_with_tensors() test_combinations() - + print("\n" + "=" * 60) print("All tests completed!") print("=" * 60) diff --git a/tests/reproduce_indexing.py b/tests/reproduce_indexing.py index 5e2c8a5..5cc4fdb 100644 --- a/tests/reproduce_indexing.py +++ b/tests/reproduce_indexing.py @@ -24,15 +24,15 @@ class TestIndexing(unittest.TestCase): 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]])) @@ -40,19 +40,19 @@ class TestIndexing(unittest.TestCase): 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 diff --git a/tests/test_bitwise_comprehensive.py b/tests/test_bitwise_comprehensive.py index 33bd068..22abd33 100644 --- a/tests/test_bitwise_comprehensive.py +++ b/tests/test_bitwise_comprehensive.py @@ -16,7 +16,7 @@ def test_mixed_dtype_operations(): print("=" * 70) print("Testing Mixed Dtype Operations") print("=" * 70) - + # Create a dummy context for testing class DummyContext: class Start: @@ -24,7 +24,7 @@ def test_mixed_dtype_operations(): column = 0 start = Start() ctx = DummyContext() - + # int8 operations print("\n1. INT8 Bitwise Operations:") a_int8 = torch.tensor([15, 7, 3], dtype=torch.int8) @@ -33,24 +33,24 @@ def test_mixed_dtype_operations(): 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, ctx) result_or = visitor._bitwise_op(a_int16, b_int16, torch.bitwise_or, lambda x, y: x | y, ctx) result_xor = visitor._bitwise_op(a_int16, b_int16, torch.bitwise_xor, lambda x, y: x ^ y, ctx) - + 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) @@ -59,7 +59,7 @@ def test_mixed_dtype_operations(): 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) @@ -68,7 +68,7 @@ def test_mixed_dtype_operations(): 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) @@ -79,7 +79,7 @@ def test_scalar_list_tensor_combinations(): print("\n" + "=" * 70) print("Testing Scalar/List/Tensor Combinations") print("=" * 70) - + # Create a dummy context for testing class DummyContext: class Start: @@ -87,37 +87,37 @@ def test_scalar_list_tensor_combinations(): column = 0 start = Start() ctx = DummyContext() - + 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, ctx) 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, ctx) 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, ctx) 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, ctx) print(f" {a} & {b} = {result}") - + print("\n" + "=" * 70) print("✓ All scalar/list/tensor combinations work correctly!") print("=" * 70) @@ -128,9 +128,9 @@ def test_element_size_mapping(): print("\n" + "=" * 70) print("Testing Element Size Mapping") print("=" * 70) - + visitor = UnifiedMathVisitor({}) - + test_dtypes = [ (torch.int8, "int8", 1), (torch.int16, "int16", 2), @@ -140,13 +140,13 @@ def test_element_size_mapping(): (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" @@ -156,10 +156,10 @@ def test_element_size_mapping(): 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) @@ -170,12 +170,12 @@ if __name__ == "__main__": 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") @@ -188,7 +188,7 @@ if __name__ == "__main__": 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 diff --git a/tests/test_bitwise_fp16.py b/tests/test_bitwise_fp16.py index 73288f3..03fd864 100644 --- a/tests/test_bitwise_fp16.py +++ b/tests/test_bitwise_fp16.py @@ -13,13 +13,13 @@ 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}") @@ -30,13 +30,13 @@ def test_bitwise_fp16(): 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}") @@ -47,7 +47,7 @@ def test_bitwise_int16(): def test_bitwise_and_int16(): """Test bitwise AND operation with int16 tensors.""" print("Testing bitwise AND with int16...") - + # Create a dummy context for testing class DummyContext: class Start: @@ -55,13 +55,13 @@ def test_bitwise_and_int16(): column = 0 start = Start() ctx = DummyContext() - + 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, ctx) print(f" a: {a}, b: {b}") @@ -72,12 +72,12 @@ def test_bitwise_and_int16(): 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}") @@ -88,9 +88,9 @@ def test_bitwise_int8(): 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), @@ -98,19 +98,19 @@ def test_get_bitwise_view_dtype(): (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() diff --git a/tests/test_conditioning_mismatch.py b/tests/test_conditioning_mismatch.py index 639dc27..7ead1f5 100644 --- a/tests/test_conditioning_mismatch.py +++ b/tests/test_conditioning_mismatch.py @@ -16,7 +16,7 @@ def test_conditioning_token_mismatch_padding(): # tokens 0-76: 1 + 0.5 = 1.5 # tokens 77-153: 0 + 0.5 = 0.5 result_list, stack = ConditioningMathNode.execute(V={"V0": ca, "V1": cb}, F={}, Expression="a + b", Expression_pi="a + b", length_mismatch="pad") - + # result_list[0] is a list: [(tensor, dict), ...] # Get the first conditioning tuple cond_tuple = result_list[0][0] diff --git a/tests/test_error_messages.py b/tests/test_error_messages.py index 891f354..fb466b1 100644 --- a/tests/test_error_messages.py +++ b/tests/test_error_messages.py @@ -1,6 +1,5 @@ import sys import os -import torch # Ensure we can import the module _here = os.path.abspath(os.path.dirname(__file__)) @@ -97,7 +96,7 @@ def run_tests(): print(f" FAILED: {e}") import traceback traceback.print_exc() - + print(f"\n{passed}/{len(tests)} tests passed.") if passed < len(tests): sys.exit(1) diff --git a/tests/test_guider_math.py b/tests/test_guider_math.py index d8477dd..971e09b 100644 --- a/tests/test_guider_math.py +++ b/tests/test_guider_math.py @@ -107,17 +107,17 @@ class TestMathGuider(unittest.TestCase): g0 = MockGuider(1.0) V = {"V0": g0} F = {} - + # Create a guider with crop expression # Input shape is (1, 4, 4, 4) = (batch, channel, height, width) # crop(a, [0, 1, 1, 1], [1, 2, 2, 2]) extracts 2x2x2 region expr = "crop(a, [0, 1, 1, 1], [1, 2, 2, 2])" math_guider = MathGuider(V, F, expr, expr) - + # Input: 4D tensor (batch, channel, height, width) x = torch.ones((1, 4, 4, 4)) sigma = torch.tensor(1.0) - + result = math_guider(x, sigma) # Result should be (1, 2, 2, 2) self.assertEqual(result.shape, (1, 2, 2, 2)) @@ -128,14 +128,14 @@ class TestMathGuider(unittest.TestCase): g0 = MockGuider(5.0) # guider returns tensor of 5.0s V = {"V0": g0} F = {} - + # Simple numeric expression expr = "a * 2" math_guider = MathGuider(V, F, expr, expr) - + x = torch.ones((1, 4, 4, 4)) sigma = torch.tensor(1.0) - + result = math_guider(x, sigma) # Expression evaluates to a * 2 where a is the guider output (5.0) # So result = 5.0 * 2 = 10.0 @@ -147,14 +147,14 @@ class TestMathGuider(unittest.TestCase): g0 = MockGuider(3.0) # guider returns tensor of 3.0s V = {"V0": g0} F = {} - + # crop with 4D parameters for 4D input (batch, channel, height, width) expr = "crop(a, [0, 0, 0, 0], [1, 2, 2, 2])" math_guider = MathGuider(V, F, expr, expr) - + x = torch.ones((1, 4, 4, 4)) sigma = torch.tensor(1.0) - + result = math_guider(x, sigma) # Result shape should be (1, 2, 2, 2) self.assertEqual(result.shape, (1, 2, 2, 2)) diff --git a/tests/test_more_math.py b/tests/test_more_math.py index 681a0ba..9413ef5 100644 --- a/tests/test_more_math.py +++ b/tests/test_more_math.py @@ -426,7 +426,7 @@ def test_nested_tensor_support(): # The unbind() operation concatenates the nested tensors into a single 5D tensor assert isinstance(res_lat, torch.Tensor), f"Expected torch.Tensor, got {type(res_lat)}" assert not getattr(res_lat, "is_nested", False), "Result should be a regular tensor, not NestedTensor" - + # Verify the concatenated result contains the computed values # NestedTensor([t1=(1,4,32,32), t2=(2,4,32,32)]) -> concatenated to (3,4,32,32) # After a + 1.0: first (1,4,32,32) should be 2.0, next (2,4,32,32) should be 3.0 @@ -435,7 +435,7 @@ def test_nested_tensor_support(): torch.full((2, 4, 32, 32), 3.0) # 2.0 + 1.0 ], dim=0) assert torch.allclose(res_lat, expected) - + # ========================================== # Comprehensive Math Function Tests @@ -525,7 +525,7 @@ def test_advanced_activations(): def test_text_upper(): from more_math.Parser.UnifiedMathVisitor import UnifiedMathVisitor from more_math.helper_functions import parse_expr - + expr = 'upper("hello")' tree = parse_expr(expr) visitor = UnifiedMathVisitor({}, (1,)) @@ -536,7 +536,7 @@ def test_text_upper(): def test_text_lower(): from more_math.Parser.UnifiedMathVisitor import UnifiedMathVisitor from more_math.helper_functions import parse_expr - + expr = 'lower("HELLO")' tree = parse_expr(expr) visitor = UnifiedMathVisitor({}, (1,)) @@ -547,7 +547,7 @@ def test_text_lower(): def test_text_trim(): from more_math.Parser.UnifiedMathVisitor import UnifiedMathVisitor from more_math.helper_functions import parse_expr - + expr = 'trim(" hello world ")' tree = parse_expr(expr) visitor = UnifiedMathVisitor({}, (1,)) @@ -558,7 +558,7 @@ def test_text_trim(): def test_text_split(): from more_math.Parser.UnifiedMathVisitor import UnifiedMathVisitor from more_math.helper_functions import parse_expr - + expr = 'split("a,b,c", ",")' tree = parse_expr(expr) visitor = UnifiedMathVisitor({}, (1,)) @@ -569,7 +569,7 @@ def test_text_split(): def test_text_join(): from more_math.Parser.UnifiedMathVisitor import UnifiedMathVisitor from more_math.helper_functions import parse_expr - + expr = 'join(["a", "b", "c"], "-")' tree = parse_expr(expr) visitor = UnifiedMathVisitor({}, (1,)) @@ -580,7 +580,7 @@ def test_text_join(): def test_text_substring(): from more_math.Parser.UnifiedMathVisitor import UnifiedMathVisitor from more_math.helper_functions import parse_expr - + expr = 'substring("hello world", 0, 5)' tree = parse_expr(expr) visitor = UnifiedMathVisitor({}, (1,)) @@ -591,7 +591,7 @@ def test_text_substring(): def test_text_find(): from more_math.Parser.UnifiedMathVisitor import UnifiedMathVisitor from more_math.helper_functions import parse_expr - + expr = 'find("hello world", "world")' tree = parse_expr(expr) visitor = UnifiedMathVisitor({}, (1,)) @@ -602,7 +602,7 @@ def test_text_find(): def test_text_replace(): from more_math.Parser.UnifiedMathVisitor import UnifiedMathVisitor from more_math.helper_functions import parse_expr - + expr = 'replace("hello world", "world", "python")' tree = parse_expr(expr) visitor = UnifiedMathVisitor({}, (1,)) @@ -617,14 +617,14 @@ def test_text_replace(): def test_crop_basic(): from more_math.Parser.UnifiedMathVisitor import UnifiedMathVisitor from more_math.helper_functions import parse_expr - + # Create a 4x4 tensor input_tensor = torch.ones((4, 4)) expr = 'crop(a, [1, 1], [2, 2])' tree = parse_expr(expr) visitor = UnifiedMathVisitor({"a": input_tensor}, input_tensor.shape) result = visitor.visit(tree) - + # Result should be 2x2 assert result.shape == (2, 2) assert torch.all(result == 1.0) @@ -633,14 +633,14 @@ def test_crop_basic(): def test_crop_3d(): from more_math.Parser.UnifiedMathVisitor import UnifiedMathVisitor from more_math.helper_functions import parse_expr - + # Create a 4x4x4 tensor input_tensor = torch.ones((4, 4, 4)) * 2.0 expr = 'crop(a, [0, 0, 0], [2, 2, 2])' tree = parse_expr(expr) visitor = UnifiedMathVisitor({"a": input_tensor}, input_tensor.shape) result = visitor.visit(tree) - + # Result should be 2x2x2 assert result.shape == (2, 2, 2) assert torch.all(result == 2.0) @@ -649,14 +649,14 @@ def test_crop_3d(): def test_crop_with_offset(): from more_math.Parser.UnifiedMathVisitor import UnifiedMathVisitor from more_math.helper_functions import parse_expr - + # Create a 6x6 tensor with different values input_tensor = torch.arange(36).reshape((6, 6)).float() expr = 'crop(a, [2, 2], [2, 2])' tree = parse_expr(expr) visitor = UnifiedMathVisitor({"a": input_tensor}, input_tensor.shape) result = visitor.visit(tree) - + # Result should be 2x2 assert result.shape == (2, 2) # Values should be from the cropped region diff --git a/tests/test_permute_fix.py b/tests/test_permute_fix.py index 9f11372..33f9f66 100644 --- a/tests/test_permute_fix.py +++ b/tests/test_permute_fix.py @@ -11,7 +11,7 @@ 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)]) @@ -34,7 +34,7 @@ def test_permute_and_apply(): 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}") diff --git a/tests/test_signal_stack.py b/tests/test_signal_stack.py index ead9a06..4af36fe 100644 --- a/tests/test_signal_stack.py +++ b/tests/test_signal_stack.py @@ -73,14 +73,14 @@ def test_all(): 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 diff --git a/tests/test_unified_math.py b/tests/test_unified_math.py index d74ad5e..3063a18 100644 --- a/tests/test_unified_math.py +++ b/tests/test_unified_math.py @@ -493,7 +493,7 @@ def test_new_loop_features(): def test_entropy(): """Test entropy function for information entropy calculation.""" vars = {} - + # 1. Test uniform distribution (maximum entropy) # Uniform probabilities should have high entropy uniform = torch.ones(4) / 4.0 # [0.25, 0.25, 0.25, 0.25] @@ -502,7 +502,7 @@ def test_entropy(): assert isinstance(entropy_uniform, float) # For uniform distribution of 4 elements: H = -sum(0.25 * log(0.25)) = log(4) ≈ 1.386 assert 1.3 < entropy_uniform < 1.5 - + # 2. Test deterministic distribution (minimum entropy) # One probability is 1, others are 0 -> entropy should be near 0 deterministic = torch.tensor([1000.0, -1000.0, -1000.0, -1000.0]) # After softmax: ~[1, 0, 0, 0] @@ -510,14 +510,14 @@ def test_entropy(): entropy_det = parse_and_visit("entropy(deterministic)", vars) assert isinstance(entropy_det, float) assert entropy_det < 0.1 # Near zero entropy - + # 3. Test with different tensor sizes small = torch.randn(8) vars["small"] = small entropy_small = parse_and_visit("entropy(small)", vars) assert isinstance(entropy_small, float) assert entropy_small > 0 # Should be positive - + # 4. Test that entropy is always non-negative random_vals = torch.randn(100) vars["random_vals"] = random_vals @@ -527,7 +527,7 @@ def test_entropy(): def test_correlation(): """Test correlation (Pearson correlation coefficient) function.""" vars = {} - + # 1. Perfect positive correlation x = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0]) y = torch.tensor([2.0, 4.0, 6.0, 8.0, 10.0]) # y = 2*x @@ -536,18 +536,18 @@ def test_correlation(): corr_perfect = parse_and_visit("corr(x, y)", vars) assert isinstance(corr_perfect, float) assert abs(corr_perfect - 1.0) < 1e-5 # Should be very close to 1 - + # Test with alias corr_alias = parse_and_visit("correlation(x, y)", vars) assert abs(corr_alias - 1.0) < 1e-5 - + # 2. Perfect negative correlation z = torch.tensor([10.0, 8.0, 6.0, 4.0, 2.0]) # Decreasing vars["z"] = z corr_negative = parse_and_visit("corr(x, z)", vars) assert isinstance(corr_negative, float) assert abs(corr_negative - (-1.0)) < 1e-5 # Should be very close to -1 - + # 3. No correlation (orthogonal) a = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0]) b = torch.tensor([1.0, -1.0, 1.0, -1.0, 1.0]) # Oscillating @@ -556,11 +556,11 @@ def test_correlation(): corr_none = parse_and_visit("corr(a, b)", vars) assert isinstance(corr_none, float) assert abs(corr_none) < 0.5 # Low correlation - + # 4. Test with same tensor (should be 1.0) corr_self = parse_and_visit("corr(x, x)", vars) assert abs(corr_self - 1.0) < 1e-5 - + # 5. Test with flattened 2D tensors t1 = torch.tensor([[1.0, 2.0], [3.0, 4.0]]) t2 = torch.tensor([[1.5, 3.0], [4.5, 6.0]]) # Scaled version @@ -569,7 +569,7 @@ def test_correlation(): corr_2d = parse_and_visit("corr(t1, t2)", vars) assert isinstance(corr_2d, float) assert abs(corr_2d - 1.0) < 1e-5 # Linear relationship - + # 6. Test correlation is symmetric corr_xy = parse_and_visit("corr(x, y)", vars) corr_yx = parse_and_visit("corr(y, x)", vars) @@ -578,14 +578,14 @@ def test_correlation(): def test_entropy_and_correlation_edge_cases(): """Test edge cases for entropy and correlation.""" vars = {} - + # 1. Entropy with constant values (after softmax becomes uniform) constant = torch.ones(10) vars["constant"] = constant entropy_const = parse_and_visit("entropy(constant)", vars) # All equal logits -> uniform distribution after softmax -> log(10) ≈ 2.302 assert 2.2 < entropy_const < 2.4 - + # 2. Correlation with constant values (undefined, but should handle gracefully) const_a = torch.ones(5) * 3.0 const_b = torch.ones(5) * 5.0 @@ -596,7 +596,7 @@ def test_entropy_and_correlation_edge_cases(): corr_const = parse_and_visit("corr(const_a, const_b)", vars) # Check that it doesn't crash and returns a float assert isinstance(corr_const, float) - + # 3. Small tensors tiny_x = torch.tensor([1.0, 2.0]) tiny_y = torch.tensor([2.0, 4.0])