(AI) add tests + forgot to change tests of functionalyty which changed
This commit is contained in:
@@ -26,12 +26,12 @@ def test_conv_1d():
|
||||
shape = (1, 10, 4)
|
||||
a_val = torch.randn(*shape)
|
||||
|
||||
# conv(a, 3, 1.0) -> implies kernel of ones, size 3
|
||||
# ezconvolution(a, 3, 1.0) -> implies kernel of ones, size 3
|
||||
# Result should correspond to 1D conv
|
||||
try:
|
||||
# LatentMathNode expects latent dicts usually
|
||||
input_dict = {"samples": a_val}
|
||||
res = node.execute("conv(a, 3, 1.0)", a=input_dict)
|
||||
res = node.execute("ezconvolution(a, 3, 1.0)", a=input_dict)
|
||||
|
||||
# LatentMathNode returns list of dicts
|
||||
res_tensor = res[0]["samples"]
|
||||
@@ -57,7 +57,7 @@ def test_conv_3d():
|
||||
|
||||
try:
|
||||
input_dict = {"samples": a_val}
|
||||
res = node.execute("conv(a, 3, 3, 3, 1.0)", a=input_dict)
|
||||
res = node.execute("ezconvolution(a, 3, 3, 3, 1.0)", a=input_dict)
|
||||
res_tensor = res[0]["samples"]
|
||||
print(f"3D Conv Result Shape: {res_tensor.shape}")
|
||||
assert res_tensor.shape == shape
|
||||
@@ -77,9 +77,9 @@ def test_conv_arbitrary_batch():
|
||||
a_val = torch.randn(*shape)
|
||||
|
||||
try:
|
||||
# conv(a, 3, 3, 1.0) -> 2D conv on (16,16)
|
||||
# ezconvolution(a, 3, 3, 1.0) -> 2D conv on (16,16)
|
||||
input_dict = {"samples": a_val}
|
||||
res = node.execute("conv(a, 3, 3, 1.0)", a=input_dict)
|
||||
res = node.execute("ezconvolution(a, 3, 3, 1.0)", a=input_dict)
|
||||
res_tensor = res[0]["samples"]
|
||||
print(f"Arbitrary Batch Result Shape: {res_tensor.shape}")
|
||||
assert res_tensor.shape == shape
|
||||
@@ -102,7 +102,7 @@ def test_conv_list_kernel():
|
||||
# Using the user's example kernel
|
||||
kernel_list = [1, 1, 1, 1, 0, 1, 1, 1, 1, 0, 0, 0, 0, 1, 0, 0, 0, 0, 1, 1, 1, 1, 0, 1, 1, 1, 1]
|
||||
kernel_str = str(kernel_list)
|
||||
expr = f"conv(a, 3, 3, 3, {kernel_str})/8"
|
||||
expr = f"ezconvolution(a, 3, 3, 3, {kernel_str})/8"
|
||||
|
||||
try:
|
||||
input_dict = {"samples": a_val}
|
||||
@@ -128,11 +128,11 @@ def test_conv_audio():
|
||||
shape = (1, 2, 100) # [B, C, L] (L >> C)
|
||||
a_val = torch.randn(*shape).float()
|
||||
|
||||
# conv(a, 3, 1.0) on last dim (L)
|
||||
# ezconvolution(a, 3, 1.0) on last dim (L)
|
||||
# Expected: result shape same as input
|
||||
try:
|
||||
input_dict = {"samples": a_val}
|
||||
res = node.execute("conv(a, 3, 1.0)", a=input_dict)
|
||||
res = node.execute("ezconvolution(a, 3, 1.0)", a=input_dict)
|
||||
res_tensor = res[0]["samples"]
|
||||
print(f"Audio Result Shape: {res_tensor.shape}")
|
||||
|
||||
@@ -161,12 +161,12 @@ def test_conv_deep_latent():
|
||||
shape = (1, 32, 16, 16)
|
||||
a_val = torch.randn(*shape).float()
|
||||
|
||||
# conv(a, 3, 3, 3, 1.0)
|
||||
# ezconvolution(a, 3, 3, 3, 1.0)
|
||||
# 3D kernels need D,H,W.
|
||||
# D=32 (Channel). H=16. W=16.
|
||||
try:
|
||||
input_dict = {"samples": a_val}
|
||||
res = node.execute("conv(a, 3, 3, 3, 1.0)", a=input_dict)
|
||||
res = node.execute("ezconvolution(a, 3, 3, 3, 1.0)", a=input_dict)
|
||||
res_tensor = res[0]["samples"]
|
||||
print(f"Deep Latent Result Shape: {res_tensor.shape}")
|
||||
assert res_tensor.shape == shape
|
||||
@@ -190,13 +190,13 @@ def test_conv_padding():
|
||||
shape = (1, 10, 10, 1)
|
||||
a_val = torch.randn(*shape).float()
|
||||
|
||||
# conv(a, 4, 4, 1.0)
|
||||
# ezconvolution(a, 4, 4, 1.0)
|
||||
# If padding is symmetric 2, result is 11x11.
|
||||
# If padding is symmetric 1, result is 9x9.
|
||||
# We need asymmetric pad (1, 2) to get 10x10.
|
||||
try:
|
||||
input_dict = {"samples": a_val}
|
||||
res = node.execute("conv(a, 4, 4, 1.0)", a=input_dict)
|
||||
res = node.execute("ezconvolution(a, 4, 4, 1.0)", a=input_dict)
|
||||
res_tensor = res[0]["samples"]
|
||||
print(f"Padding Test Result Shape: {res_tensor.shape}")
|
||||
assert res_tensor.shape == shape
|
||||
@@ -218,7 +218,7 @@ def test_conv_complex_padding():
|
||||
|
||||
try:
|
||||
input_dict = {"samples": a_val}
|
||||
res = node.execute("conv(a, 3, 4, 1.0)", a=input_dict)
|
||||
res = node.execute("ezconvolution(a, 3, 4, 1.0)", a=input_dict)
|
||||
res_tensor = res[0]["samples"]
|
||||
print(f"Complex Padding Result Shape: {res_tensor.shape}")
|
||||
assert res_tensor.shape == shape
|
||||
@@ -239,7 +239,7 @@ def test_conv_3d_asymmetric():
|
||||
|
||||
try:
|
||||
input_dict = {"samples": a_val}
|
||||
res = node.execute("conv(a, 3, 3, 3, 1.0)", a=input_dict)
|
||||
res = node.execute("ezconvolution(a, 3, 3, 3, 1.0)", a=input_dict)
|
||||
res_tensor = res[0]["samples"]
|
||||
print(f"3D Asymmetric Result Shape: {res_tensor.shape}")
|
||||
assert res_tensor.shape == shape
|
||||
|
||||
@@ -0,0 +1,40 @@
|
||||
import torch
|
||||
import sys
|
||||
import os
|
||||
|
||||
# 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.expr()
|
||||
shape = (1, 1, 1, 1)
|
||||
visitor = UnifiedMathVisitor(variables, shape)
|
||||
return visitor.visit(tree)
|
||||
|
||||
def reproduce():
|
||||
print("Reproducing percentile(list, tensor)...")
|
||||
a = torch.rand(1, 4, 128, 128) * 100.0 # 0-100 range for percentile
|
||||
l = [0.1, 1.0, 0.3]
|
||||
|
||||
vars = {"a": a, "l": l}
|
||||
# Note: percentile(l, a) -> a/100.0 will be 0-1.
|
||||
try:
|
||||
res = parse_and_visit("percentile(l, a)", vars)
|
||||
print(f"Success! Result type: {type(res)}, shape: {res.shape}")
|
||||
except Exception:
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
|
||||
if __name__ == "__main__":
|
||||
reproduce()
|
||||
@@ -0,0 +1,105 @@
|
||||
|
||||
import torch
|
||||
import sys
|
||||
import os
|
||||
|
||||
# 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.expr()
|
||||
shape = (1, 1, 1, 1) # Dummy shape
|
||||
visitor = UnifiedMathVisitor(variables, shape)
|
||||
return visitor.visit(tree)
|
||||
|
||||
def test_ezconvolution_2d():
|
||||
# Test 2D convolution with "ezconvolution" (automagic)
|
||||
# Case: (1, 10, 10, 3) - Standard ComfyUI Image (Batch, H, W, C)
|
||||
# Kernel: 3x3
|
||||
|
||||
# Input: 1 batch, 10x10, 3 channels.
|
||||
t = torch.zeros((1, 10, 10, 3))
|
||||
t[:, 5, 5, :] = 1.0 # Centered dot
|
||||
|
||||
vars = {"t": t}
|
||||
|
||||
# Kernel expr: 1.0 (box blur / sum)
|
||||
# ezconvolution(tensor, kw, kh, k_expr)
|
||||
# Expect output same shape (1, 10, 10, 3) because ezconvolution handles layout.
|
||||
|
||||
expr = "ezconvolution(t, 3, 3, 1.0)"
|
||||
res = parse_and_visit(expr, vars)
|
||||
|
||||
assert isinstance(res, torch.Tensor)
|
||||
assert res.shape == (1, 10, 10, 3)
|
||||
|
||||
# Check value at 5,5. Neighboring 3x3 (9 pixels) should contribute.
|
||||
# Kernel val is 1.0 everywhere.
|
||||
# The dot is at 5,5.
|
||||
# Convolution at 5,5 will sum up the 3x3 area around input 5,5.
|
||||
# Input 5,5 is 1.0, others 0.
|
||||
# So result at 5,5 should be 1.0 * 1.0 = 1.0?
|
||||
# Wait, if kernel is 3x3 ones. Input has single 1.
|
||||
# Conv is sum(I * K).
|
||||
# If we are at 5,5. Area is 4,4 to 6,6.
|
||||
# Input has 1 at 5,5.
|
||||
# Result at 4,4 will include input 5,5?
|
||||
# Yes.
|
||||
# Result at 5,5 will include input 5,5.
|
||||
|
||||
# ezconvolution preserves logic so it should work.
|
||||
|
||||
def test_convolution_2d_strict():
|
||||
# Test 2D convolution with "convolution" (strict)
|
||||
# Case: (1, 3, 10, 10) - PyTorch Standard (Batch, C, H, W)
|
||||
|
||||
t = torch.zeros((1, 3, 10, 10))
|
||||
t[:, :, 5, 5] = 1.0
|
||||
|
||||
vars = {"t": t}
|
||||
|
||||
# Kernel expr: 1.0
|
||||
# convolution(tensor, kw, kh, k_expr)
|
||||
|
||||
expr = "convolution(t, 3, 3, 1.0)"
|
||||
res = parse_and_visit(expr, vars)
|
||||
|
||||
assert isinstance(res, torch.Tensor)
|
||||
assert res.shape == (1, 3, 10, 10)
|
||||
|
||||
def test_convolution_strict_fail_on_wrong_layout():
|
||||
# Test that strict convolution likely fails or produces weird shape on wrong layout
|
||||
# Input: (1, 10, 10, 3) -> Batch=1, 10, 10, 3
|
||||
# Try 2D conv.
|
||||
# Strict assumes (..., C, H, W).
|
||||
# so C=10, H=10, W=3. Batch=1.
|
||||
# Kernel 3x3.
|
||||
# W=3. Padding for kernel 3 is 1. Padded W = 3+2=5.
|
||||
# Conv 3 on 5 => valid.
|
||||
# So it might RUN, but it interprets dimensions wrong.
|
||||
|
||||
t = torch.randn((1, 10, 10, 3))
|
||||
vars = {"t": t}
|
||||
|
||||
expr = "convolution(t, 3, 3, 1.0)"
|
||||
res = parse_and_visit(expr, vars)
|
||||
|
||||
# Result shape should be based on (1, 10, 10, 3) input.
|
||||
# Batch=(1). C=10. H=10. W=3.
|
||||
# Output: (1, 10, 10, 3).
|
||||
# So shape matches input, but semantic is wrong.
|
||||
# This confirms it didn't permute (if it permuted, it would crash or do something else).
|
||||
|
||||
assert res.shape == (1, 10, 10, 3)
|
||||
|
||||
@@ -116,7 +116,7 @@ def test_all_functions():
|
||||
check("snorm(ta)")
|
||||
check("map(ta, x)") # 1D map
|
||||
# Verify conv with kernel variables kW
|
||||
check("conv(ta, 3, kW)") # 1D conv, size 3, value is kW (which is 3.0)
|
||||
check("conv(reshape(ta, [1, 3]), 3, kW)") # 1D conv, Input [1, 3] (Channels=1, Width=3)
|
||||
check("permute(ta, [0])")
|
||||
check("reshape(ta, [3, 1])")
|
||||
|
||||
|
||||
@@ -0,0 +1,33 @@
|
||||
import torch
|
||||
import sys
|
||||
import os
|
||||
|
||||
# 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
|
||||
|
||||
def test_large_quantile_fallback():
|
||||
print("Testing large quantile fallback in UnifiedMathVisitor...")
|
||||
# 36 million elements
|
||||
size = 36_000_000
|
||||
try:
|
||||
val = torch.rand(size)
|
||||
q = torch.tensor([0.0, 0.1, 0.5, 0.9, 1.0])
|
||||
|
||||
visitor = UnifiedMathVisitor({})
|
||||
|
||||
print(f"Calling _quartile_helper with size {size}...")
|
||||
# This should trigger the fallback
|
||||
res = visitor._quartile_helper(val, q)
|
||||
print(f"Success! Result: {res}")
|
||||
except Exception as e:
|
||||
print(f"Failed with error: {e}")
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_large_quantile_fallback()
|
||||
@@ -0,0 +1,136 @@
|
||||
import torch
|
||||
import sys
|
||||
import os
|
||||
import math
|
||||
|
||||
# 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.expr()
|
||||
shape = (1, 1, 1, 1)
|
||||
visitor = UnifiedMathVisitor(variables, shape)
|
||||
return visitor.visit(tree)
|
||||
|
||||
def test_quantile_basic():
|
||||
vars = {"t": torch.tensor([0.0, 1.0, 2.0, 3.0, 4.0]), "l": [0, 1, 2, 3, 4]}
|
||||
|
||||
# Quantile (0-1)
|
||||
assert parse_and_visit("quantile(t, 0.5)", vars) == 2.0
|
||||
assert parse_and_visit("quantile(l, 0.25)", vars) == 1.0
|
||||
assert parse_and_visit("quantile(t, 1.0)", vars) == 4.0
|
||||
assert parse_and_visit("quantile(l, 0)", vars) == 0.0
|
||||
|
||||
def test_percentile_basic():
|
||||
vars = {"t": torch.tensor([0.0, 1.0, 2.0, 3.0, 4.0]), "l": [0, 1, 2, 3, 4]}
|
||||
|
||||
# Percentile (0-100)
|
||||
assert parse_and_visit("percentile(t, 50)", vars) == 2.0
|
||||
assert parse_and_visit("percentile(l, 25)", vars) == 1.0
|
||||
assert parse_and_visit("percentile(t, 100)", vars) == 4.0
|
||||
assert parse_and_visit("percentile(l, 0)", vars) == 0.0
|
||||
|
||||
def test_quartile_basic():
|
||||
vars = {"t": torch.tensor([0.0, 1.0, 2.0, 3.0, 4.0]), "l": [0, 1, 2, 3, 4]}
|
||||
|
||||
# Quartile (0-4)
|
||||
assert parse_and_visit("quartile(t, 2)", vars) == 2.0
|
||||
assert parse_and_visit("quartile(l, 1)", vars) == 1.0
|
||||
assert parse_and_visit("quartile(t, 4)", vars) == 4.0
|
||||
assert parse_and_visit("quartile(l, 0)", vars) == 0.0
|
||||
|
||||
# Alias
|
||||
assert parse_and_visit("quartil(t, 2)", vars) == 2.0
|
||||
|
||||
def test_tensor_queries():
|
||||
t = torch.tensor([0.0, 10.0, 20.0, 30.0, 40.0])
|
||||
q_tensor = torch.tensor([0.0, 0.5, 1.0])
|
||||
vars = {"t": t, "q": q_tensor}
|
||||
|
||||
# Quantile with tensor q
|
||||
res = parse_and_visit("quantile(t, q)", vars)
|
||||
assert isinstance(res, torch.Tensor)
|
||||
assert torch.allclose(res, torch.tensor([0.0, 20.0, 40.0]))
|
||||
|
||||
# Percentile with tensor p
|
||||
p_tensor = torch.tensor([0.0, 50.0, 100.0])
|
||||
vars["p"] = p_tensor
|
||||
res_p = parse_and_visit("percentile(t, p)", vars)
|
||||
assert torch.allclose(res_p, torch.tensor([0.0, 20.0, 40.0]))
|
||||
|
||||
def test_list_with_tensor_query():
|
||||
# Optimization: list input promoted to tensor for tensor query
|
||||
l = [0, 10, 20, 30, 40]
|
||||
q = torch.tensor([0.25, 0.75])
|
||||
vars = {"l": l, "q": q}
|
||||
|
||||
res = parse_and_visit("quantile(l, q)", vars)
|
||||
assert isinstance(res, torch.Tensor)
|
||||
assert torch.allclose(res, torch.tensor([10.0, 30.0]))
|
||||
|
||||
def test_nd_tensor_query():
|
||||
# N-D tensor as query should reshape result
|
||||
t = torch.tensor([0.0, 100.0])
|
||||
q_nd = torch.zeros((2, 2, 2))
|
||||
q_nd[0, 0, 1] = 1.0 # Max
|
||||
q_nd[1, 1, 1] = 0.5 # Mid
|
||||
vars = {"t": t, "q": q_nd}
|
||||
|
||||
res = parse_and_visit("quantile(t, q)", vars)
|
||||
assert res.shape == (2, 2, 2)
|
||||
assert res[0, 0, 0] == 0.0
|
||||
assert res[0, 0, 1] == 100.0
|
||||
assert res[1, 1, 1] == 50.0
|
||||
|
||||
def test_large_tensor_fallback():
|
||||
# Test that the fallback doesn't crash (we can't easily verify it *fell back*
|
||||
# but we can verify it works on larger ones)
|
||||
# Using 1M for speed in regular tests, but enough to trust the logic
|
||||
size = 1_000_000
|
||||
t = torch.rand(size)
|
||||
vars = {"t": t}
|
||||
|
||||
# Simple median
|
||||
res = parse_and_visit("quantile(t, 0.5)", vars)
|
||||
# torch.quantile(t, 0.5) should be very close to t.median()
|
||||
assert math.isclose(res.item(), t.median().item(), abs_tol=1e-3)
|
||||
|
||||
def test_list_query():
|
||||
# percentile(t, [0, 10, 20]) -> returns a list of results
|
||||
t = torch.tensor([0.0, 10.0, 20.0, 30.0, 40.0])
|
||||
vars = {"t": t}
|
||||
|
||||
res = parse_and_visit("percentile(t, [0, 50, 100])", vars)
|
||||
assert isinstance(res, list)
|
||||
assert len(res) == 3
|
||||
assert res[0] == 0.0
|
||||
assert res[1] == 20.0
|
||||
assert res[2] == 40.0
|
||||
|
||||
if __name__ == "__main__":
|
||||
# If run as script
|
||||
try:
|
||||
test_quantile_basic()
|
||||
test_percentile_basic()
|
||||
test_quartile_basic()
|
||||
test_tensor_queries()
|
||||
test_list_with_tensor_query()
|
||||
test_nd_tensor_query()
|
||||
test_large_tensor_fallback()
|
||||
test_list_query()
|
||||
print("All quantile_ops tests passed!")
|
||||
except Exception:
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
sys.exit(1)
|
||||
@@ -0,0 +1,79 @@
|
||||
|
||||
import torch
|
||||
import sys
|
||||
import os
|
||||
|
||||
# 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.expr()
|
||||
shape = (1, 1, 1, 1) # Dummy shape
|
||||
visitor = UnifiedMathVisitor(variables, shape)
|
||||
return visitor.visit(tree)
|
||||
|
||||
def test_shorthamd_ezconv():
|
||||
t = torch.zeros((1, 10, 10, 3))
|
||||
t[:, 5, 5, :] = 1.0
|
||||
vars = {"t": t}
|
||||
|
||||
# ezconv should work exactly like ezconvolution
|
||||
expr = "ezconv(t, 3, 3, 1.0)"
|
||||
res = parse_and_visit(expr, vars)
|
||||
assert isinstance(res, torch.Tensor)
|
||||
assert res.shape == (1, 10, 10, 3)
|
||||
|
||||
def test_shorthand_conv():
|
||||
t = torch.zeros((1, 3, 10, 10))
|
||||
t[:, :, 5, 5] = 1.0
|
||||
vars = {"t": t}
|
||||
|
||||
# conv should work exactly like convolution (strictly)
|
||||
expr = "conv(t, 3, 3, 1.0)"
|
||||
res = parse_and_visit(expr, vars)
|
||||
assert isinstance(res, torch.Tensor)
|
||||
assert res.shape == (1, 3, 10, 10)
|
||||
|
||||
def test_shorthand_print_shape():
|
||||
t = torch.zeros((1, 3))
|
||||
vars = {"t": t}
|
||||
|
||||
# Test print_shape
|
||||
# visitor prints to stdout, we just check it runs and returns t
|
||||
res1 = parse_and_visit("print_shape(t)", vars)
|
||||
assert torch.equal(res1, t)
|
||||
|
||||
# Test pshp
|
||||
res2 = parse_and_visit("pshp(t)", vars)
|
||||
assert torch.equal(res2, t)
|
||||
|
||||
def test_shorthand_perm():
|
||||
t = torch.randn((2, 3, 4))
|
||||
vars = {"t": t}
|
||||
# permute(t, [2, 0, 1])
|
||||
res_full = parse_and_visit("permute(t, [2, 0, 1])", vars)
|
||||
res_short = parse_and_visit("perm(t, [2, 0, 1])", vars)
|
||||
|
||||
assert res_short.shape == (4, 2, 3)
|
||||
assert torch.equal(res_full, res_short)
|
||||
|
||||
def test_shorthand_rshp():
|
||||
t = torch.randn((2, 3, 4)) # 24 elements
|
||||
vars = {"t": t}
|
||||
# reshape(t, [24])
|
||||
res_full = parse_and_visit("reshape(t, [24])", vars)
|
||||
res_short = parse_and_visit("rshp(t, [24])", vars)
|
||||
|
||||
assert res_short.shape == (24,)
|
||||
assert torch.equal(res_full, res_short)
|
||||
@@ -0,0 +1,65 @@
|
||||
|
||||
import torch
|
||||
import sys
|
||||
import os
|
||||
import math
|
||||
|
||||
# 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.expr()
|
||||
shape = (1, 1, 1, 1) # Dummy shape
|
||||
visitor = UnifiedMathVisitor(variables, shape)
|
||||
return visitor.visit(tree)
|
||||
|
||||
def test_sum():
|
||||
vars = {"t": torch.tensor([1.0, 2.0, 3.0]), "l": [1.0, 2.0, 3.0], "s": 5.0}
|
||||
assert parse_and_visit("sum(t)", vars) == 6.0
|
||||
assert parse_and_visit("sum(l)", vars) == 6.0
|
||||
assert parse_and_visit("sum(s)", vars) == 5.0
|
||||
|
||||
def test_mean():
|
||||
vars = {"t": torch.tensor([1.0, 2.0, 3.0]), "l": [1.0, 2.0, 3.0], "s": 5.0}
|
||||
assert parse_and_visit("mean(t)", vars) == 2.0
|
||||
assert parse_and_visit("mean(l)", vars) == 2.0
|
||||
assert parse_and_visit("mean(s)", vars) == 5.0
|
||||
|
||||
def test_std():
|
||||
t = torch.tensor([1.0, 5.0]) # mean=3, var=((2^2 + 2^2)/1) = 8. std = sqrt(8)=2.828...
|
||||
vars = {"t": t, "l": [1.0, 5.0]}
|
||||
|
||||
assert math.isclose(parse_and_visit("std(t)", vars).item(), math.sqrt(8), rel_tol=1e-5)
|
||||
assert math.isclose(parse_and_visit("std(l)", vars), math.sqrt(8), rel_tol=1e-5)
|
||||
|
||||
def test_var():
|
||||
t = torch.tensor([1.0, 5.0])
|
||||
vars = {"t": t, "l": [1.0, 5.0]}
|
||||
assert math.isclose(parse_and_visit("var(t)", vars).item(), 8.0, rel_tol=1e-5)
|
||||
assert math.isclose(parse_and_visit("var(l)", vars), 8.0, rel_tol=1e-5)
|
||||
|
||||
def test_dot():
|
||||
a = torch.tensor([1.0, 2.0])
|
||||
b = torch.tensor([3.0, 4.0])
|
||||
vars = {"a": a, "b": b}
|
||||
# dot = 1*3 + 2*4 = 3 + 8 = 11
|
||||
assert parse_and_visit("dot(a, b)", vars) == 11.0
|
||||
|
||||
# Check flattening behavior (2D tensor)
|
||||
a2 = torch.tensor([[1.0, 2.0]])
|
||||
b2 = torch.tensor([[3.0], [4.0]]) # 2x1
|
||||
vars = {"a2": a2, "b2": b2}
|
||||
# a2 flat = [1, 2], b2 flat = [3, 4] -> dot=11
|
||||
assert parse_and_visit("dot(a2, b2)", vars) == 11.0
|
||||
|
||||
+49
-10
@@ -9,7 +9,8 @@ _project_root = os.path.abspath(os.path.join(_here, os.pardir))
|
||||
if _project_root not in sys.path:
|
||||
sys.path.insert(0, _project_root)
|
||||
|
||||
# Placeholder import - we will create this file next
|
||||
|
||||
# Import visitor
|
||||
from more_math.Parser.UnifiedMathVisitor import UnifiedMathVisitor
|
||||
from more_math.Parser.MathExprLexer import MathExprLexer
|
||||
from more_math.Parser.MathExprParser import MathExprParser
|
||||
@@ -183,9 +184,6 @@ def test_bool_ops():
|
||||
|
||||
|
||||
def test_topk():
|
||||
# Tensor topk
|
||||
t = torch.tensor([1.0, 5.0, 2.0, 8.0, 3.0])
|
||||
vars = {"t": t}
|
||||
# Tensor topk masking
|
||||
t = torch.tensor([1.0, 5.0, 2.0, 8.0, 3.0])
|
||||
vars = {"t": t}
|
||||
@@ -229,9 +227,9 @@ def test_botk():
|
||||
|
||||
def test_pinv():
|
||||
# List permutation inverse
|
||||
perm = [2, 0, 1] # 0->2, 1->0, 2->1
|
||||
vars = {"perm": perm}
|
||||
res = parse_and_visit("pinv(perm)", vars)
|
||||
p = [2, 0, 1] # 0->2, 1->0, 2->1
|
||||
vars = {"p": p}
|
||||
res = parse_and_visit("pinv(p)", vars)
|
||||
assert isinstance(res, list)
|
||||
# Inverse: if perm[i]=j, then inv[j]=i
|
||||
# perm[0]=2 -> inv[2]=0
|
||||
@@ -240,9 +238,9 @@ def test_pinv():
|
||||
assert res == [1, 2, 0]
|
||||
|
||||
# Tensor permutation inverse
|
||||
perm_t = torch.tensor([2, 0, 1])
|
||||
vars["perm_t"] = perm_t
|
||||
res_t = parse_and_visit("pinv(perm_t)", vars)
|
||||
pt = torch.tensor([2, 0, 1])
|
||||
vars["pt"] = pt
|
||||
res_t = parse_and_visit("pinv(pt)", vars)
|
||||
assert isinstance(res_t, torch.Tensor)
|
||||
assert torch.equal(res_t, torch.tensor([1, 2, 0]))
|
||||
|
||||
@@ -253,6 +251,46 @@ def test_pinv_identity():
|
||||
res = parse_and_visit("permute(permute(c,a),pinv(a))",varbl)
|
||||
assert torch.equal(tensor,res)
|
||||
|
||||
|
||||
def test_quartil():
|
||||
# Test Quartiles (Strict Integer Indices)
|
||||
# List: [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10] (11 elements)
|
||||
l = [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10]
|
||||
vars = {"l": l}
|
||||
|
||||
assert parse_and_visit("quartile(l, 0)", vars) == 0.0
|
||||
assert parse_and_visit("quartile(l, 1)", vars) == 2.5
|
||||
assert parse_and_visit("quartile(l, 2)", vars) == 5.0
|
||||
assert parse_and_visit("quartile(l, 3)", vars) == 7.5
|
||||
assert parse_and_visit("quartile(l, 4)", vars) == 10.0
|
||||
|
||||
# Tensor
|
||||
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))
|
||||
|
||||
# Float inputs for quartil should be cast to int, so 0.5 -> 0 -> Min
|
||||
# verifying strict behavior or fallback
|
||||
assert parse_and_visit("quartile(l, 0.9)", vars) == 0.0
|
||||
|
||||
|
||||
def test_percentile():
|
||||
l = [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10]
|
||||
vars = {"l": l}
|
||||
t = torch.tensor(l, dtype=torch.float32)
|
||||
vars["t"] = t
|
||||
|
||||
# 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))
|
||||
|
||||
# Aliases
|
||||
assert parse_and_visit("prcnt(l, 50)", vars) == 5.0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
try:
|
||||
test_scalar_ops()
|
||||
@@ -267,6 +305,7 @@ if __name__ == "__main__":
|
||||
test_topk()
|
||||
test_botk()
|
||||
test_pinv()
|
||||
test_quartil()
|
||||
print("All UnifiedMathVisitor tests passed!")
|
||||
except Exception:
|
||||
import traceback
|
||||
|
||||
Reference in New Issue
Block a user