@@ -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 } \n def 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