ruff --fix. Again...

This commit is contained in:
mcDandy
2026-01-10 15:34:06 +01:00
parent cf02031ba2
commit 2cd69ce5fd
+9 -9
View File
@@ -25,7 +25,7 @@ def parse_and_visit(expr_str, variables):
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
@@ -34,7 +34,7 @@ def test_quantile_basic():
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
@@ -43,13 +43,13 @@ def test_percentile_basic():
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
@@ -57,7 +57,7 @@ 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)
@@ -74,7 +74,7 @@ def test_list_with_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]))
@@ -86,7 +86,7 @@ def test_nd_tensor_query():
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
@@ -100,7 +100,7 @@ def test_large_tensor_fallback():
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()
@@ -110,7 +110,7 @@ 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