为什么np.vectorize(function) 能避免该函数在未被np.vectorize包装时会遇到的整数溢出?
我知道 1.0 * (40-80) 会导致 -40.0。我理解的是
1.0 * (np.array([np.uint8(40)])-80)
会导致 array([216.])。因此,```
import numpy as np
C = 1.0
def f(x): return C * (x - 80)
print(f(np.array([np.uint8(40)])))
`` 也不意外地输出[216.]`。更令人惊讶的是
print(np.vectorize(f)(np.array([np.uint8(40)])))
的输出:
[-40.]
/var/folders/jz/34xs15sd3ng6dpxkv78hh1lr0000gn/T/ipykernel_61208/1067934613.py:7: RuntimeWarning: overflow encountered in scalar subtract
return C * (x - 80)
基于这个讨论串,我对正在发生的事情的最佳理解是,f 被调用一次以确定输出的数据类型,遇到了溢出(因此有警告)。输出的数据类型被确定为 np.float64。f 再次被调用,这次是为了确定实际输出,我们得到 [-40.]。似乎某种方式,输入 np.uint8(40) 被转换为 np.float64,只有在那之后才减去80,因此得到 -40,而不是216。
如果我手动指定输出的数据类型,我也会期望得到相同的输出。所以,这个:
print(np.vectorize(f,otypes=[np.float64])([np.uint8(40)]))
我希望按照上面的推理,得到 [-40.]。输出数据类型已知(np.float64),将输入转换为 np.float64,减去80,乘以1.0,返回 -40。然而,我得到的是:
[216.]
/var/folders/jz/34xs15sd3ng6dpxkv78hh1lr0000gn/T/ipykernel_67644/1868397835.py:7: RuntimeWarning: overflow encountered in scalar subtract
return C * (x - 80)
因此,上文中对正在发生的事情的看法是错误的。
为什么 print(np.vectorize(f)(np.array([np.uint8(40)]))) 能避免整数溢出,而 print(f(np.array([np.uint8(40)]))) 和 print(np.vectorize(f,otypes=[np.float64])([np.uint8(40)])) 却不能避免?
解决方案
如果我们查看numpy的 源码,我们可以看到 vectorize 在调用函数之前不会改变输入的数据类型。传递 otypes 只会影响输出的数据类型,我们在第2599-2602行看到,函数是使用原始数据类型调用的,结果被强制转换为通过 otypes 传递的类型。
当你调用 vectorize 时,发生如下情况:
args = [asanyarray(a, dtype=object) for a in args]
outputs = ufunc(*args, out=...)
if ufunc.nout == 1:
res = asanyarray(outputs, dtype=otypes[0])
看起来所有输入都先被转换为对象类型,因此也许也会产生一些额外的怪异行为,但核心原因是:输入在执行你的函数之前没有被强制转换,而 输出只有在函数计算完成后才被强制转换。