missed a few
This commit is contained in:
+1
-1
@@ -1 +1 @@
|
||||
from .more_math.nodes import comfy_entrypoint
|
||||
from .more_math.nodes import comfy_entrypoint as comfy_entrypoint
|
||||
|
||||
@@ -11,8 +11,6 @@ if _comfy_root not in sys.path:
|
||||
sys.path.insert(0, _comfy_root)
|
||||
|
||||
import torch
|
||||
import pytest
|
||||
from more_math.Parser.UnifiedMathVisitor import UnifiedMathVisitor
|
||||
|
||||
from more_math.LatentMathNode import LatentMathNode
|
||||
|
||||
@@ -139,7 +137,7 @@ def test_conv_audio():
|
||||
print(f"Audio Result Shape: {res_tensor.shape}")
|
||||
|
||||
if res_tensor.shape != shape:
|
||||
print(f"Likely interpreted as Channels Last [B, L, C] where C is small? No.")
|
||||
print("Likely interpreted as Channels Last [B, L, C] where C is small? No.")
|
||||
# If interpreted as Channels last [..., C].
|
||||
# [1, 2, 100]. Spatial=[2]. Channel=100.
|
||||
# Output [1, 2, 100] (but confusing channels).
|
||||
@@ -262,7 +260,7 @@ if __name__ == "__main__":
|
||||
test_conv_complex_padding()
|
||||
test_conv_3d_asymmetric()
|
||||
print("All Conv tests passed!")
|
||||
except Exception as e:
|
||||
except Exception:
|
||||
import traceback
|
||||
|
||||
traceback.print_exc()
|
||||
|
||||
@@ -122,7 +122,7 @@ def test_all_functions():
|
||||
|
||||
# FFT/IFFT
|
||||
# We need a shape for FFT usually
|
||||
res_fft = eval_tensor_expr("fft(ta)", variables, (3,))
|
||||
eval_tensor_expr("fft(ta)", variables, (3,))
|
||||
res_ifft = eval_tensor_expr("ifft(fft(ta))", variables, (3,))
|
||||
assert torch.allclose(res_ifft, tensor_a, atol=1e-4)
|
||||
|
||||
|
||||
@@ -112,7 +112,6 @@ def test_model_math_device_mismatch():
|
||||
# Since we might only have CPU, we can't fully reproduce 'cpu vs cuda' crash without cuda.
|
||||
# But we can verify that the scalar created by visitor.visitNumberExp has the same device as 'a'.
|
||||
|
||||
from more_math.Parser.TensorEvalVisitor import TensorEvalVisitor
|
||||
from antlr4 import InputStream, CommonTokenStream
|
||||
from more_math.Parser.MathExprLexer import MathExprLexer
|
||||
from more_math.Parser.MathExprParser import MathExprParser
|
||||
@@ -125,7 +124,7 @@ def test_model_math_device_mismatch():
|
||||
# Let's try to pass a dummy device string if create_tensor allows, or just check the code path.
|
||||
# Better: Inspect the created tensor from visitor.
|
||||
|
||||
tsr = torch.zeros((1,))
|
||||
torch.zeros((1,))
|
||||
# We interpret "device mismatch" as: created scalars didn't pick up the device of 'tsr'.
|
||||
# We can force 'tsr' to be on a specific device if available, but likely only 'cpu' is available.
|
||||
# Use a mock object for 'a' that claims to be on 'cuda:0', even if it isn't real tensor?
|
||||
@@ -139,9 +138,9 @@ def test_model_math_device_mismatch():
|
||||
lexer = MathExprLexer(input_stream)
|
||||
stream = CommonTokenStream(lexer)
|
||||
parser = MathExprParser(stream)
|
||||
tree = parser.expr()
|
||||
parser.expr()
|
||||
|
||||
variables = {"a": torch.zeros(1)}
|
||||
{"a": torch.zeros(1)}
|
||||
# We want to ensure that if we had a non-cpu device, it would use it.
|
||||
# Since we can't really test this without a GPU, we will write the fix and verify it analytically
|
||||
# or use a mock that wraps a tensor but intercepts .device?
|
||||
|
||||
@@ -405,7 +405,6 @@ def test_pow_log_functions():
|
||||
|
||||
|
||||
def test_min_max_functions():
|
||||
import sys
|
||||
|
||||
node = FloatMathNode()
|
||||
print("Testing tmin...", flush=True)
|
||||
@@ -466,7 +465,7 @@ if __name__ == "__main__":
|
||||
test_basic_utilities()
|
||||
test_advanced_activations()
|
||||
print("All tests passed!")
|
||||
except Exception as e:
|
||||
except Exception:
|
||||
import traceback
|
||||
|
||||
traceback.print_exc()
|
||||
|
||||
@@ -2,7 +2,6 @@ import sys
|
||||
import os
|
||||
import torch
|
||||
import math
|
||||
import pytest
|
||||
|
||||
# Ensure we can import the module
|
||||
_here = os.path.abspath(os.path.dirname(__file__))
|
||||
@@ -187,7 +186,7 @@ if __name__ == "__main__":
|
||||
test_kernel_coords()
|
||||
test_bool_ops()
|
||||
print("All UnifiedMathVisitor tests passed!")
|
||||
except Exception as e:
|
||||
except Exception:
|
||||
import traceback
|
||||
|
||||
traceback.print_exc()
|
||||
|
||||
Reference in New Issue
Block a user