Skip to content

Commit 9aa3d54

Browse files
committed
Verify Boolean operands in portable floating arithmetic
1 parent 3015a93 commit 9aa3d54

3 files changed

Lines changed: 70 additions & 0 deletions

File tree

‎plans/portable-DSLs-first-version.md‎

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -246,6 +246,20 @@ infinities, and NaNs. Code-generation cache version is 18. Validation passed all
246246
373 regular native tests, 292 AddressSanitizer tests, and 1019 focused Python
247247
tests without expected failures or new native build warnings.
248248

249+
The Boolean-operand floating-arithmetic slice corrects interpreter storage:
250+
`bool(x) + bool(x)` with integral inputs and floating output now returns two
251+
for nonzero input rather than writing an invalid Boolean byte before conversion.
252+
The DSL preserves the requested floating computation dtype for Boolean-inferred
253+
arithmetic, while casts still consume their arguments in their own dtype.
254+
Typed unary negation casts before operating so false produces floating negative
255+
zero, including top-level bool-assigned locals. Float64 lowering stays narrow
256+
to preserve hybrid branch/select plans; cache version is 19. Five shared fixtures
257+
and direct/local scalar/vector matrices passed 393 regular native tests,
258+
307 AddressSanitizer tests, and 2014 focused Python tests without expected
259+
failures or new native build warnings. `audit/bool_output_fraction` preserves an
260+
unresolved Boolean-output arithmetic discrepancy, not a normative rule; that
261+
domain and broader mixed-type promotion remain publication gates.
262+
249263
## 1. Goal and scope
250264

251265
Establish a small, versioned miniexpr kernel language that executes independently

‎tests/test_dsl_portable.py‎

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -56,6 +56,11 @@
5656
"arithmetic_float_nested",
5757
"arithmetic_float_local",
5858
"arithmetic_float_widen",
59+
"bool_float_sum",
60+
"bool_float_neg",
61+
"bool_float_product",
62+
"bool_float_difference",
63+
"bool_float_local_neg",
5964
"while_cap_exact",
6065
"while_cap_exceeded",
6166
"while_cap_continue",

‎tests/test_portable_artifact.py‎

Lines changed: 51 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -351,6 +351,52 @@ def test_optional_adapter_availability():
351351
blosc2.PortableKernel.from_json("{}", jit=False)
352352

353353

354+
@pytest.mark.parametrize("compiler", ["tcc", "cc"])
355+
@pytest.mark.parametrize("jit", [False, True])
356+
@pytest.mark.parametrize("input_dtype", ["bool", "int32", "int64", "float32", "float64"])
357+
@pytest.mark.parametrize("output_dtype", ["float32", "float64"])
358+
@pytest.mark.parametrize("count", [1, 10, 257])
359+
@pytest.mark.parametrize("operation", ["sum", "difference", "product", "negation"])
360+
@pytest.mark.parametrize("local", [False, True])
361+
def test_bool_cast_floating_arithmetic(
362+
native_artifacts, compiler, jit, input_dtype, output_dtype, count, operation, local
363+
):
364+
expressions = {
365+
"sum": "bool(x) + bool(x)",
366+
"difference": "bool(x) - bool(x)",
367+
"product": "bool(x) * bool(x)",
368+
"negation": "-bool(x)",
369+
}
370+
expression = expressions[operation]
371+
body = ""
372+
if local:
373+
expression = expression.replace("bool(x)", "truth")
374+
body = " truth = bool(x)\n"
375+
source = f"# me:compiler={compiler}\ndef k(x):\n{body} return {expression}\n"
376+
artifact = blosc2.DSLKernel.from_source(source).export({"x": input_dtype}, output_dtype)
377+
kernel = blosc2.PortableKernel.from_json(artifact, jit=jit)
378+
if jit:
379+
assert kernel.has_jit
380+
samples = [0, 1, -1, 2]
381+
if input_dtype == "int64":
382+
samples += [2**53 + 1, -(2**53 + 1)]
383+
elif input_dtype.startswith("float"):
384+
samples += [-0.0, np.finfo(input_dtype).smallest_subnormal, np.inf, -np.inf, np.nan]
385+
values = np.resize(np.array(samples, dtype=input_dtype), count)
386+
truth = (values != 0).astype(output_dtype)
387+
expected = {
388+
"sum": truth + truth,
389+
"difference": truth - truth,
390+
"product": truth * truth,
391+
"negation": -truth,
392+
}[operation]
393+
for _ in range(2):
394+
actual = kernel.evaluate({"x": values})
395+
np.testing.assert_array_equal(actual, expected)
396+
zeros = expected == 0
397+
np.testing.assert_array_equal(np.signbit(actual[zeros]), np.signbit(expected[zeros]))
398+
399+
354400
@pytest.mark.parametrize("compiler", ["tcc", "cc"])
355401
@pytest.mark.parametrize("jit", [False, True])
356402
@pytest.mark.parametrize("count", [1, 10, 257])
@@ -438,6 +484,11 @@ def test_float32_leaf_math(native_artifacts, case, compiler, jit, count):
438484
"arithmetic_float_nested",
439485
"arithmetic_float_local",
440486
"arithmetic_float_widen",
487+
"bool_float_sum",
488+
"bool_float_neg",
489+
"bool_float_product",
490+
"bool_float_difference",
491+
"bool_float_local_neg",
441492
],
442493
)
443494
@pytest.mark.parametrize("compiler", ["tcc", "cc"])

0 commit comments

Comments
 (0)