张量的数据类型
张量的数据类型与 numpy.array 基本一一对应,除了不支持str类型
- 一般的神经网络用的是
torch.float32类型
- 如果要显示指定数据类型,可以使用
torch.tensor(data,dtype = torch.type)
- 也可以使用特定的构造函数
1 2 3
| i = torch.Inttensor() x = torch.Tensor() b = torch.BoolTensor()
|
- 此外,还可以对不同类型的张量进行转化
1 2 3 4
| i = torch.tensor(1) x = i.float() y = i.type(torch.float) z = i.type_as(x)
|
张量的维度
张量的尺寸
- 可以使用shape属性或者size() 方法查看张量在每一维的长度
- 可以使用view方法改变张量的尺寸
- view失败的情况下,可以使用reshape方法
view 和 reshape 的区别:
- view方法要求原张量在内存中是连续的,如果不连续则会失败;reshape则会自动处理布局
- view方法总是与原张量共享内存,返回的是原向量的”视图”; reshape则可能返回视图或者副本,取决于内存的布局
- 为什么不只使用
reshape
- 性能考虑 view 更快,因为只是改变张量的元数据不涉及数据复制;reshape可能涉及到复制数据,会有额外的开销
- 内存效率 view保证内存共享,修改一个会影响另一个;reshape可能创建副本导致占用更多内存
- 语义的明确性 view 明确表示期望的是内存共享的试图操作;当view失败时,提醒开发者注意内存布局问题
- 以下是失败情况的例子
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29
| import torch
x = torch.randn(3, 4) y = x.transpose(0, 1) print(y.is_contiguous())
try: z = y.view(2, 6) except RuntimeError as e: print(f"view失败: {e}")
z = y.reshape(2, 6) print(z.shape)
x = torch.randn(4, 4) y = x[:, ::2] print(y.is_contiguous())
try: z = y.view(-1) except RuntimeError as e: print(f"view失败: {e}")
z = y.reshape(-1)
|
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16
|
matrix26 = torch.arange(0,12).view(2,6) print(matrix26) print(matrix26.shape)
matrix62 = matrix26.t() print(matrix62.is_contiguous())
matrix34 = matrix62.reshape(3,4) print(matrix34)
|
张量与numpy数组
- 可以使用numpy方法从tensor得到numpy数组,也可以用torch.from_numpy从numpy数组得到tensor.
- 两种方法共享数据内存,改变一个另一个也会随之改变
- 可以用张量的clone 方法拷贝张量,中断这种关联
- 可以使用item方法从标量张量得到对应的python数值
- 使用tolist方法从张量得到对应的python数值列表
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60
|
arr = np.zeros(3) tensor = torch.from_numpy(arr) print("before add 1:") print(arr) print(tensor)
print("\nafter add 1:") np.add(arr,1, out = arr) print(arr) print(tensor)
tensor = torch.zeros(3) arr = tensor.numpy() print("before add 1:") print(tensor) print(arr)
print("\nafter add 1:")
tensor.add_(1)
print(tensor) print(arr)
tensor = torch.zeros(3)
arr = tensor.clone().numpy() print("before add 1:") print(tensor) print(arr)
print("\nafter add 1:")
tensor.add_(1) print(tensor) print(arr)
scalar = torch.tensor(1.0) s = scalar.item() print(s) print(type(s))
tensor = torch.rand(2,2) t = tensor.tolist() print(t) print(type(t))
|
Something else
什么是元数据
- 元数据(metadata)描述数据张什么样而不是数据本身的信息,张量结构性的描述信息,不包括实际的数值
在 PyTorch 中,张量由“数据存储(storage)”与“元数据”两部分构成:
- 数据存储:真正的数值内存区域
- 元数据:描述如何解释这些数值的结构信息
典型元数据:shape、stride、dtype、device、storage_offset、requires_grad、layout、(可选) names
像 view() 这类操作只是改元数据(不复制数据);而当现有 stride 组合无法支持新形状时,reshape() 会退化为复制,得到新的连续存储。
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17
| import torch x = torch.arange(12) y = x.view(3, 4) print(y.shape, y.stride())
t = y.t() print(t.shape, t.stride(), t.is_contiguous())
try: t.view(12) except RuntimeError as e: print("view失败:", e)
z = t.reshape(12) print(z.is_contiguous())
|