Python3中传入numpy整数类型时PyTorch报错问题咨询(Python2无此问题)
解决Python3下PyTorch接收numpy整数类型报错的问题
我之前也踩过这个坑,核心原因是Python3与Python2在整数类型兼容上的差异,再加上PyTorch对函数参数类型的校验逻辑导致的。
问题背景
在Python2中,numpy的整数类型(比如numpy.int64)会被隐式转换为原生Python整数,传给PyTorch函数完全没问题;但到了Python3,numpy整数是独立的类型,PyTorch的内置函数(比如torch.randn)默认只接受原生Python整数或PyTorch自身的整数类型,直接传numpy整数就会触发TypeError。
报错示例
Python 3.5.5 |Anaconda custom (64-bit)| (default, Mar 12 2018, 23:12:44) [GCC 7.2.0] on linux Type "help", "copyright", "credits" or "license" for more information. >>> import torch >>> import numpy >>> torch.randn(numpy.int64(4)) Traceback (most recent call last): File "<stdin>", line 1, in <module> TypeError: torch.randn received an invalid combination of arguments - got (numpy.int64,), but expected one of: * (int ..., *, torch.device device) * (torch.Size size, *, torch.device device)
解决方案
这里有几个简单的处理方法,任选其一即可:
- 用原生Python的
int()函数转换:torch.randn(int(numpy.int64(4))) # 正常生成形状为(4,)的张量 - 调用numpy整数的
item()方法提取原生值:torch.randn(numpy.int64(4).item()) # 同样可以正常运行 - 如果是处理批量数据,提前将numpy整数数组转换为原生整数列表:
numpy_ints = numpy.array([2, 3, 4], dtype=numpy.int64) torch.randn(*[int(i) for i in numpy_ints]) # 生成形状为(2,3,4)的张量
内容的提问来源于stack exchange,提问作者Philip Hodges
相关产品推荐
相关产品推荐

