【发布时间】:2019-01-25 11:05:28
【问题描述】:
从0.4.0版本开始,可以使用torch.tensor和torch.Tensor
有什么区别?提供这两个非常相似且令人困惑的替代方案的原因是什么?
【问题讨论】:
从0.4.0版本开始,可以使用torch.tensor和torch.Tensor
有什么区别?提供这两个非常相似且令人困惑的替代方案的原因是什么?
【问题讨论】:
torch.Tensor 是创建参数时最喜欢使用的方法(例如在nn.Linear、nn._ConvNd 中)。
为什么?因为它非常快。甚至比torch.empty()还要快一点。
import torch
torch.set_default_dtype(torch.float32) # default
%timeit torch.empty(1000,1000)
%timeit torch.Tensor(1000,1000)
%timeit torch.ones(1000,1000)
%timeit torch.tensor([[1]*1000]*1000)
输出:
68.4 µs ± 789 ns per loop (mean ± std. dev. of 7 runs, 10000 loops each)
67.9 µs ± 349 ns per loop (mean ± std. dev. of 7 runs, 10000 loops each)
1.26 ms ± 8.61 µs per loop (mean ± std. dev. of 7 runs, 1000 loops each)
36.1 ms ± 610 µs per loop (mean ± std. dev. of 7 runs, 10 loops each)
torch.Tensor() 和 and torch.empty() 非常相似,并返回一个充满未初始化数据的张量。
为什么我们不初始化 __init__ 中的参数,而技术上这是可能的?
这里是torch.Tensor在nn.Linear里面的实践,用来创建weight参数:
self.weight = nn.Parameter(torch.Tensor(out_features, in_features))
我们不会根据设计对其进行初始化。还有另一个reset_parameters() 方法,因为在训练时可能需要再次“重置”参数,我们在__init__() 方法的末尾调用reset_paremeters()。
也许将来torch.empty() 会取代torch.Tensor(),因为它们的效果是一样的。
reset_parameters() 也有一个不错的选择,如果需要,您可以创建自己的版本并更改原始初始化过程。
【讨论】:
https://discuss.pytorch.org/t/difference-between-torch-tensor-and-torch-tensor/30786/2
torch.tensor 自动推断数据类型,而 torch.Tensor 返回一个torch.FloatTensor。我会建议坚持 torch.tensor,如果你愿意,它也有像 dtype 这样的参数 改变类型。
【讨论】:
在 PyTorch 中,torch.Tensor 是主要的张量类。所以所有张量都只是torch.Tensor 的实例。
当您调用torch.Tensor() 时,您将得到一个没有任何data 的空张量。
相比之下,torch.tensor 是一个返回张量的函数。在documentation 中它说:
torch.tensor(data, dtype=None, device=None, requires_grad=False) → Tensor用
data构造一个张量。
tensor_without_data = torch.Tensor()
但另一方面:
tensor_without_data = torch.tensor()
会导致错误:
---------------------------------------------------------------------------
TypeError Traceback (most recent call last)
<ipython-input-12-ebc3ceaa76d2> in <module>()
----> 1 torch.tensor()
TypeError: tensor() missing 1 required positional arguments: "data"
在没有data 的情况下创建张量的类似行为:torch.Tensor() 可以使用:
torch.tensor(())
输出:
tensor([])
【讨论】:
除了上述答案之外,我注意到torch.Tensor(<data>) 将使用默认数据类型(如torch.get_default_dtype() 中定义)初始化张量。另一方面,torch.tensor(<data>) 会从数据中推断出数据类型。
例如,
tensor_arr = torch.tensor([[2,5,6],[9,7,6]])
tensor_arr
将打印:
tensor([[2, 5, 6], [9, 7, 6]])
和
tensor_arr = torch.Tensor([[2,5,6],[9,7,6]])
tensor_arr
将打印:
tensor([[2., 5., 6.], [9., 7., 6.]])
因为默认数据类型是 float32。
【讨论】:
根据pytorch discussion的讨论
torch.Tensor 构造函数被重载以执行与 torch.tensor 和 torch.empty 相同的事情。人们认为这种重载会使代码混乱,因此将torch.Tensor 拆分为torch.tensor 和torch.empty。
所以是的,在某种程度上,torch.tensor 的工作方式类似于 Torch.Tensor(当您传入数据时)。不,两者都不应该比另一个更有效。只是torch.empty 和torch.tensor 的API 比torch.Tensor 构造函数更好。
【讨论】: