(AI) add tests + forgot to change tests of functionalyty which changed

This commit is contained in:
mcDandy
2026-01-10 15:32:04 +01:00
parent 81a17c9352
commit cf02031ba2
9 changed files with 522 additions and 25 deletions
+14 -14
View File
@@ -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
+40
View File
@@ -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()
+105
View File
@@ -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)
+1 -1
View File
@@ -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])")
+33
View File
@@ -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()
+136
View File
@@ -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)
+79
View File
@@ -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)
+65
View File
@@ -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
View File
@@ -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