【问题标题】:Numpy/Pytorch dtype conversion / compatibilityNumpy/Pytorch dtype 转换/兼容性
【发布时间】:2019-05-07 12:18:26
【问题描述】:

我正在尝试查找一些文档来了解 dtypes 是如何组合的。例如:

x : np.int32 = ...
y : np.float64 = ...
  • x + y 的类型是什么?
  • 是否取决于运营商(此处为+)?
  • 它是否取决于存储位置(z = x + y 与 z[...] = x + y)?

我正在寻找描述这类场景的部分文档,但到目前为止我两手空空。

【问题讨论】:

  • 请注意未来的我:这种行为只是 numpy 的。如果类型不匹配,pytorch 似乎不会进行向上转换,而是运行时检查 + 运行时异常

标签: python numpy type-conversion pytorch numpy-dtype


【解决方案1】:

如果数据类型不匹配,则 NumPy 将upcast the data to the higher precision data types if possible。并且它不依赖于我们所做的(算术)操作的类型或我们分配给的变量,除非该变量已经具有其他一些 dtype。这是一个小插图:

In [14]: x = np.arange(3, dtype=np.int32)
In [15]: y = np.arange(3, dtype=np.float64)

# `+` is equivalent to `numpy.add()`
In [16]: summed = x + y

In [17]: summed.dtype
Out[17]: dtype('float64')

In [18]: np.add(x, y).dtype
Out[18]: dtype('float64')

如果您没有明确指定数据类型,则结果将向上转换为给定输入的较高数据类型。例如,numpy.add() 接受 dtype kwarg,您可以在其中指定结果数组的数据类型。


并且,可以使用numpy.can_cast()检查是否可以根据转换规则安全地转换两种不同的数据类型

为了完整起见,我添加以下numpy.can_cast()矩阵:

>>> def print_casting_matrix(ntypes):
...     ntypes_ex = ["X"] + ntypes.split()
...     print("".join(ntypes_ex))
...     for row in ntypes:
...         print(row, sep='\t', end=''),
...         for col in ntypes:
...             print(int(np.can_cast(row, col)), sep='\t', end='')
...         print()

>>> print_casting_matrix(np.typecodes['All'])

输出将是以下矩阵,它显示哪些 dtypes 可以安全地转换(由 1 表示)和哪些 dtypes 不能转换(由 0 表示),按照 from cast(沿轴 0)到 铸造(轴 1):

# to casting -----> ----->
X?bhilqpBHILQPefdgFDGSUVOMm
?11111111111111111111111101
b01111110000001111111111101
h00111110000000111111111101
i00011110000000011011111101
l00001110000000011011111101
q00001110000000011011111101
p00001110000000011011111101
B00111111111111111111111101
H00011110111110111111111101
I00001110011110011011111101
L00000000001110011011111101
Q00000000001110011011111101
P00000000001110011011111101
e00000000000001111111111100
f00000000000000111111111100
d00000000000000011011111100
g00000000000000001001111100
F00000000000000000111111100
D00000000000000000011111100
G00000000000000000001111100
S00000000000000000000111100
U00000000000000000000011100
V00000000000000000000001100
O00000000000000000000001100
M00000000000000000000001110
m00000000000000000000001101

由于字符很神秘,我们可以使用以下内容来更好地理解上述转换矩阵:

In [74]: for char in np.typecodes['All']:
    ...:     print(char, " --> ", np.typeDict[char])

输出将是:

?  -->  <class 'numpy.bool_'>
b  -->  <class 'numpy.int8'>
h  -->  <class 'numpy.int16'>
i  -->  <class 'numpy.int32'>
l  -->  <class 'numpy.int64'>
q  -->  <class 'numpy.int64'>
p  -->  <class 'numpy.int64'>
B  -->  <class 'numpy.uint8'>
H  -->  <class 'numpy.uint16'>
I  -->  <class 'numpy.uint32'>
L  -->  <class 'numpy.uint64'>
Q  -->  <class 'numpy.uint64'>
P  -->  <class 'numpy.uint64'>
e  -->  <class 'numpy.float16'>
f  -->  <class 'numpy.float32'>
d  -->  <class 'numpy.float64'>
g  -->  <class 'numpy.float128'>
F  -->  <class 'numpy.complex64'>
D  -->  <class 'numpy.complex128'>
G  -->  <class 'numpy.complex256'>
S  -->  <class 'numpy.bytes_'>
U  -->  <class 'numpy.str_'>
V  -->  <class 'numpy.void'>
O  -->  <class 'numpy.object_'>
M  -->  <class 'numpy.datetime64'>
m  -->  <class 'numpy.timedelta64'>

【讨论】:

  • 在这种情况下,您有解释“更高”概念的文档的链接吗?
  • @Vinz 不在官方文档中。我添加了对外部文档的引用。
  • 不清楚向上转换的结果是什么,但我总是可以尝试所有这些并创建我的地图。感谢您的信息
  • @Vinz 我添加了更多信息。请查看更新的答案! HTH :)
猜你喜欢
  • 2017-12-25
  • 1970-01-01
  • 1970-01-01
  • 2023-04-06
  • 2021-03-23
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
  • 2023-03-19
相关资源
最近更新 更多