Repository navigation
add np.float_power and np.cbrt - #6075
Conversation
|
Thanks for opening this @guilhermeleobas , looks like CI is failing with near to floating point limit issues. |
|
Hi @stuartarchibald, can we do something in this case? The |
|
CI is all green |
| def float_cbrt(x): | ||
| if np.isnan(x): | ||
| return x | ||
| if x < 0: | ||
| return -np.power(-x, 1.0 / 3.0) | ||
| else: | ||
| return np.power(x, 1.0 / 3.0) |
There was a problem hiding this comment.
This is going to be "interesting" in that in NumPy, if a system has a cbrt from C99 then it gets used, else something like what you've written is used. The impact being that e.g. signed zero is incorrect for me locally as I've got C99 cbrt!
from numba import njit
import numpy as np
@njit
def foo():
return np.cbrt(-0.)
print(foo.py_func())
print(foo())gives:
-0.0
0.0
LLVM can be coerced into generating cbrt with appropriate hints.
There was a problem hiding this comment.
LLVM can be coerced into generating
cbrtwith appropriate hints.
np.signbit can be used here to distinguish 0.0 from its negative counterpart. I didn't want to use it since the numba implementation uses the numpy one and I'd like to make this function available in RBC.
One could mimic npy_signbit in pure LLVM but that would require knowledge about the system (sizeof(int) and endianess).
There was a problem hiding this comment.
Replaced comparison by np.signbit in my last commit. I think this should fix signed zero comparison.
There was a problem hiding this comment.
The issues I see with this are:
- LLVM will generate a call to
cbrtunder certain conditions and NumPy is usinglibm'scbrtif it's available, so this would match exactly. Generating the call correctly would also permit e.g. use of SVML'scbrt. If you need a hand working out how to do this then shout! I'd recommend looking at the impact offastmathflags as a starting point. - If RBC needs this then it should implement it or make Numba generic to the point that RBC can consume it?
There was a problem hiding this comment.
If RBC needs this then it should implement it or make Numba generic to the point that RBC can consume it?
I think it is consumable in the way it is today. We already ship our own version of np.signbit. By the way, I just remember that the libm version of cbrt works on RBC. So, no reason for not using it.
LLVM will generate a call to cbrt under certain conditions and NumPy is using libm's cbrt if it's available, so this would match exactly. Generating the call correctly would also permit e.g. use of SVML's cbrt. If you need a hand working out how to do this then shout! I'd recommend looking at the impact of fastmath flags as a starting point.
I will replace the code by a call to libm cbrt.
There was a problem hiding this comment.
LLVM will generate a call to cbrt under certain conditions and NumPy is using libm's cbrt if it's available, so this would match exactly. Generating the call correctly would also permit e.g. use of SVML's cbrt. If you need a hand working out how to do this then shout! I'd recommend looking at the impact of fastmath flags as a starting point.
I will replace the code by a call to libm cbrt.
What guarantees that cbrt is in a libm ?
There was a problem hiding this comment.
Isn't cbrt part of libm?
There was a problem hiding this comment.
If you have C99. NumPy does this:
https://github.com/numpy/numpy/blob/31ffdecf07d18ed4dbb66b171cb0f998d4b190fa/numpy/core/src/npymath/npy_math_internal.h.src#L508-L529
There was a problem hiding this comment.
Makes sense! I will update the PR
There was a problem hiding this comment.
Any tips for debugging the CI failure?
Thread 1 "python" received signal SIGSEGV, Segmentation fault.
0x0000000000000000 in ?? ()
(gdb) bt
#0 0x0000000000000000 in ?? ()
#1 0x00007ffff7dfc051 in cpython::__main__::foo$241(long long) ()
#2 0x00007ffff3b1cb3c in call_cfunc (self=0x7fffeb670c10, cfunc=0x7fffd24648b0, args=0x7fffd24c2f10, kws=0x0, locals=0x0) at numba/_dispatcher.c:353
#3 0x00007ffff3b1d6f5 in Dispatcher_call (self=0x7fffeb670c10, args=0x7fffd24c2f10, kws=<optimized out>) at numba/_dispatcher.c:576
#4 0x000055555568395f in _PyObject_MakeTpCall () at /home/conda/feedstock_root/build_artifacts/python_1596159888721/work/Objects/call.c:159
#5 0x0000555555735300 in _PyObject_Vectorcall (kwnames=0x0, nargsf=<optimized out>, args=0x7ffff7eeb5b8, callable=0x7fffeb670c10)
at /home/conda/feedstock_root/build_artifacts/python_1596159888721/work/Include/cpython/abstract.h:125
#6 call_function (kwnames=0x0, oparg=<optimized out>, pp_stack=<synthetic pointer>, tstate=0x5555558efa90)
at /home/conda/feedstock_root/build_artifacts/python_1596159888721/work/Python/ceval.c:4963
#7 _PyEval_EvalFrameDefault () at /home/conda/feedstock_root/build_artifacts/python_1596159888721/work/Python/ceval.c:3500
#8 0x000055555571c190 in PyEval_EvalFrameEx (throwflag=0, f=0x7ffff7eeb440)
at /home/conda/feedstock_root/build_artifacts/python_1596159888721/work/Python/ceval.c:741
#9 _PyEval_EvalCodeWithName () at /home/conda/feedstock_root/build_artifacts/python_1596159888721/work/Python/ceval.c:4298
#10 0x000055555571d9f3 in PyEval_EvalCodeEx (closure=0x0, kwdefs=0x0, defcount=0, defs=0x0, kwcount=0, kws=0x0, argcount=0, args=0x0,
locals=<optimized out>, globals=<optimized out>, _co=<optimized out>)
at /home/conda/feedstock_root/build_artifacts/python_1596159888721/work/Python/ceval.c:4327
#11 PyEval_EvalCode (co=<optimized out>, globals=<optimized out>, locals=<optimized out>)
at /home/conda/feedstock_root/build_artifacts/python_1596159888721/work/Python/ceval.c:718
#12 0x00005555557908b2 in run_eval_code_obj () at /home/conda/feedstock_root/build_artifacts/python_1596159888721/work/Python/pythonrun.c:1125
#13 0x00005555557a3492 in run_mod () at /home/conda/feedstock_root/build_artifacts/python_1596159888721/work/Python/pythonrun.c:1147
#14 0x00005555557a657e in PyRun_FileExFlags () at /home/conda/feedstock_root/build_artifacts/python_1596159888721/work/Python/pythonrun.c:1063
#15 0x00005555557a6769 in PyRun_SimpleFileExFlags () at /home/conda/feedstock_root/build_artifacts/python_1596159888721/work/Python/pythonrun.c:428
#16 0x00005555557a6c1e in pymain_run_file (cf=0x7fffffffc0c8, config=0x5555558eec20)
at /home/conda/feedstock_root/build_artifacts/python_1596159888721/work/Modules/main.c:387
#17 pymain_run_python (exitcode=0x7fffffffc0c0) at /home/conda/feedstock_root/build_artifacts/python_1596159888721/work/Modules/main.c:612
#18 Py_RunMain () at /home/conda/feedstock_root/build_artifacts/python_1596159888721/work/Modules/main.c:691
#19 0x00005555557a6e19 in Py_BytesMain () at /home/conda/feedstock_root/build_artifacts/python_1596159888721/work/Modules/main.c:1123
#20 0x00007ffff77e6b97 in __libc_start_main (main=0x55555565a7f0 <main>, argc=2, argv=0x7fffffffc2b8, init=<optimized out>, fini=<optimized out>,
rtld_fini=<optimized out>, stack_end=0x7fffffffc2a8) at ../csu/libc-start.c:310
#21 0x000055555574aba9 in _start ()
(gdb)Edit: I think I found the error
Edit2: CI is all green :)
e258d02 to
8de632a
Compare
|
Hi @stuartarchibald, when you have some free cycles to spare, can you take a look at this PR? |
| NUMBA_EXPORT_FUNC(float) | ||
| numba_cbrtf(float x) | ||
| { | ||
| return npy_cbrtf(x); | ||
| } | ||
|
|
||
| NUMBA_EXPORT_FUNC(double) | ||
| numba_cbrt(double x) | ||
| { | ||
| return npy_cbrt(x); | ||
| } |
There was a problem hiding this comment.
I'm struggling to understand why this is needed?
This:
from numba import njit
import numpy as np
@njit(fastmath={'nnan', 'nsz', 'ninf', 'afn'})
def cbrt(x):
third = 1.0 / 3.0
return np.power(x, third)
@njit
def foo(x):
return cbrt(x)
x = np.zeros(100)
foo(x)
print(cbrt.inspect_llvm(cbrt.signatures[0]))
print(cbrt.inspect_asm(cbrt.signatures[0]))generates a call to the cbrt function, LLVM "knows" that power(x, 1/3) is cbrt under the flags set. Perhaps this can be used to ensure that a cbrt call is generated along with adding the inf/nan handling as per NumPy?
There was a problem hiding this comment.
I'm struggling to understand why this is needed?
Before adding this, one of the tests were failing on CI. See:
#6075 (comment)
generates a call to the cbrt function, LLVM "knows" that power(x, 1/3) is cbrt under the flags set.
For me, this produces a call to @llvm.pow.f64(x, 1/3):
; Function Attrs: nofree nounwind writeonly
define i32 @"_ZN8__main__8cbrt$242Ex"(double* noalias nocapture %retptr, { i8*, i32, i8* }** noalias nocapture readnone %excinfo, i64 %arg.x) local_unnamed_addr #0 {
entry:
%.37.le = sitofp i64 %arg.x to double
%.38.le = tail call nnan ninf nsz afn double @llvm.pow.f64(double %.37.le, double 0x3FD5555555555555)
store double %.38.le, double* %retptr, align 8
ret i32 0
}There was a problem hiding this comment.
Yes, but that then ends up as a call to cbrt once compiled?
There was a problem hiding this comment.
Oh yes, you right. Nevertheless, If I remove the code you quote, some builds on CI starts to fail. Do you want me to give a try and implement cbrt in numba?
edit: this implementation produces the following assembly
def cbrt(x):
if np.isnan(x):
return np.nan
elif x < 0:
return -np.power(-x, 1.0 / 3.0)
else:
return np.power(x, 1.0 / 3.0)
return context.compile_internal(builder, cbrt, sig, args) .text
.file "<string>"
.globl _ZN8__main__7bar$242Ex
.p2align 4, 0x90
.type _ZN8__main__7bar$242Ex,@function
_ZN8__main__7bar$242Ex:
pushq %rbx
movq %rdi, %rbx
vcvtsi2sd %rdx, %xmm0, %xmm0
movabsq $cbrt, %rax
callq *%rax
vmovsd %xmm0, (%rbx)
xorl %eax, %eax
popq %rbx
retq
.Lfunc_end0:
.size _ZN8__main__7bar$242Ex, .Lfunc_end0-_ZN8__main__7bar$242Ex
.globl cfunc._ZN8__main__7bar$242Ex
.p2align 4, 0x90
.type cfunc._ZN8__main__7bar$242Ex,@function
cfunc._ZN8__main__7bar$242Ex:
vcvtsi2sd %rdi, %xmm0, %xmm0
movabsq $cbrt, %rax
jmpq *%rax
.Lfunc_end1:
.size cfunc._ZN8__main__7bar$242Ex, .Lfunc_end1-cfunc._ZN8__main__7bar$242Ex
.type _ZN08NumbaEnv8__main__7bar$242Ex,@object
.comm _ZN08NumbaEnv8__main__7bar$242Ex,8,8
.section ".note.GNU-stack","",@progbitsThere was a problem hiding this comment.
Using this:
from numba import njit
import numpy as np
@njit
def foo(arr):
return np.cbrt(arr)
foo(np.ones(10, dtype=np.float64))
print(foo.inspect_asm(foo.signatures[0]))I'm not seeing the call to cbrt being made. I suspect it requires floating point assumptions to be relaxed in the two branches that call np.power.
There was a problem hiding this comment.
If one enables fast-math, you'll see a call to cbrt.
from numba import njit
import numpy as np
@njit('float32(float32)', fastmath=True, no_cfunc_wrapper=True, no_cpython_wrapper=True)
def foo(x):
return np.cbrt(x)
print(foo.inspect_asm(foo.signatures[0])) .section __TEXT,__text,regular,pure_instructions
.macosx_version_min 10, 15
.section __TEXT,__literal4,4byte_literals
.p2align 2
LCPI0_0:
.long 2147483648
.section __TEXT,__text,regular,pure_instructions
.globl __ZN8__main__7foo$241Ef
.p2align 4, 0x90
__ZN8__main__7foo$241Ef:
pushq %r14
pushq %rbx
subq $56, %rsp
vmovaps %xmm0, (%rsp)
movq %rdi, %rbx
vcvtss2sd %xmm0, %xmm0, %xmm0
movabsq $_cbrt, %r14
callq *%r14
vcvtsd2ss %xmm0, %xmm0, %xmm0
vmovaps %xmm0, 32(%rsp)
movabsq $LCPI0_0, %rax
vbroadcastss (%rax), %xmm0
vmovaps %xmm0, 16(%rsp)
vxorps (%rsp), %xmm0, %xmm0
vcvtss2sd %xmm0, %xmm0, %xmm0
callq *%r14
vcvtsd2ss %xmm0, %xmm0, %xmm0
vxorps 16(%rsp), %xmm0, %xmm0
vxorps %xmm1, %xmm1, %xmm1
vmovaps (%rsp), %xmm2
vcmpltss %xmm1, %xmm2, %xmm1
vmovaps 32(%rsp), %xmm2
vblendvps %xmm1, %xmm0, %xmm2, %xmm0
vmovss %xmm0, (%rbx)
xorl %eax, %eax
addq $56, %rsp
popq %rbx
popq %r14
retq
.comm __ZN08NumbaEnv8__main__7foo$241Ef,8,3
.comm __ZN08NumbaEnv5numba2np8npyfuncs17np_real_cbrt_impl12$3clocals$3e8cbrt$242Ef,8,3
.subsections_via_symbols|
CI is all green! Can you take a look when you have some free cycles to spare, @stuartarchibald |
|
|
||
| def np_real_cbrt_impl(context, builder, sig, args): | ||
| _check_arity_and_homogeneity(sig, args, 1) | ||
| import numpy as np |
There was a problem hiding this comment.
Imports go at the top please (I assume there's no circular reference problem?).
stuartarchibald
left a comment
There was a problem hiding this comment.
Thanks for the patch and fixes!
As title.