如何获取可容纳数组所有元素的最小numpy dtype?
寻找NumPy中适配数组的最小dtype等效函数
对于标量,我可以用np.min_scalar_type获取能容纳其值的最小dtype,示例如下:
In [32]: np.min_scalar_type(0) Out[32]: dtype('uint8') In [33]: np.min_scalar_type(0.1) Out[33]: dtype('float16') In [34]: np.min_scalar_type(-1) Out[34]: dtype('int8') In [35]: np.min_scalar_type(np.nan) Out[35]: dtype('float16')
但针对数组,有没有等效的函数?现有的np.min_scalar_type处理数组时,只会返回输入数组的dtype,而非能容纳所有元素的最小适配dtype:
In [36]: np.min_scalar_type([0, 0.1, -1, np.nan]) Out[36]: dtype('float64') # 期望结果:float16 In [37]: np.min_scalar_type([0, 1]) Out[37]: dtype('int64') # 期望结果:uint8 In [39]: np.min_scalar_type([1, -1]) Out[39]: dtype('int64') # 期望结果:int8 In [40]: np.min_scalar_type([-100, 200]) Out[40]: dtype('int64') # 期望结果:int16
实现这个功能并不简单,比如数组[-1, 200]的最小适配dtype是int16,但单独对每个元素应用np.min_scalar_type无法得到正确结果。我曾尝试过np.min_scalar_type((np.sign(np.min(ar)) or 1) * np.max(ar)),但在全负数或包含浮点数的场景下会失效。请问有没有现成的实现?
内容的提问来源于stack exchange,提问作者gerrit
相关产品推荐
相关产品推荐

