在 PyTorch 里,unsqueeze(-1) 是一个常用的张量操作,其作用是在张量的最后一个维度上增添一个维度。下面为你详细介绍这个操作:

基本概念

unsqueeze(dim) 会在指定的 dim 维度处插入一个维度,新维度的大小为 1。而 unsqueeze(-1) 里的 -1 代表最后一个维度,所以它会在张量的最后面添加一个维度。

示例讲解

下面通过几个具体例子来直观理解 unsqueeze(-1) 的效果:

示例 1:对一维张量进行操作

python

运行

import torch

# 创建一个一维张量
x = torch.tensor([1, 2, 3])
print("原始张量形状:", x.shape)  # 输出: torch.Size([3])

# 在最后一个维度上增加一个维度
y = x.unsqueeze(-1)
print("操作后张量形状:", y.shape)  # 输出: torch.Size([3, 1])
print(y)
# 输出:
# tensor([[1],
#         [2],
#         [3]])

在这个例子中,原本形状为 [3] 的一维张量,经过 unsqueeze(-1) 操作后,形状变成了 [3, 1],也就是变成了一个 3 行 1 列的二维张量。

示例 2:对二维张量进行操作

python

运行

# 创建一个二维张量
x = torch.tensor([[1, 2], [3, 4]])
print("原始张量形状:", x.shape)  # 输出: torch.Size([2, 2])

# 在最后一个维度上增加一个维度
y = x.unsqueeze(-1)
print("操作后张量形状:", y.shape)  # 输出: torch.Size([2, 2, 1])
print(y)
# 输出:
# tensor([[[1],
#          [2]],
# 
#         [[3],
#          [4]]])

这里,原本形状为 [2, 2] 的二维张量,操作后变成了形状为 [2, 2, 1] 的三维张量。

示例 3:对三维张量进行操作

python

运行

# 创建一个三维张量
x = torch.randn(2, 3, 4)
print("原始张量形状:", x.shape)  # 输出: torch.Size([2, 3, 4])

# 在最后一个维度上增加一个维度
y = x.unsqueeze(-1)
print("操作后张量形状:", y.shape)  # 输出: torch.Size([2, 3, 4, 1])

对于形状为 [2, 3, 4] 的三维张量,操作后变成了形状为 [2, 3, 4, 1] 的四维张量。

实际应用场景

unsqueeze(-1) 常用于以下场景:

  1. 维度匹配:在进行张量运算时,需要保证参与运算的张量维度一致。
  2. 数据格式调整:在处理图像、序列等数据时,可能需要调整数据的维度结构,以适应模型的输入要求。
  3. 向量转矩阵:将一维向量转换为二维矩阵,方便进行矩阵运算。

注意事项

  • 负索引规则:在 PyTorch 中,负索引从 -1 开始,表示最后一个维度,-2 表示倒数第二个维度,依此类推。
  • 与 squeeze 的区别unsqueeze 是增加维度,而 squeeze 是删除维度为 1 的维度。

通过上述讲解和示例,你应该对 unsqueeze(-1) 操作有了清晰的理解。如果还有疑问,可以尝试在 PyTorch 环境中运行这些示例代码,进一步观察结果。

在 PyTorch 中,squeeze(-1) 是一个用于减少张量维度的操作,它会移除指定维度上尺寸为 1 的维度。下面详细解释这个操作:

基本概念

squeeze(dim) 函数会在指定的 dim 位置上移除尺寸为 1 的维度。当 dim 为 -1 时,表示移除最后一个维度,但仅当该维度的尺寸为 1 时才会生效。如果最后一个维度的尺寸不是 1,则张量不会发生任何变化。

示例讲解

以下是几个具体示例,展示 squeeze(-1) 的效果:

示例 1:移除最后一个维度(尺寸为 1)

python

运行

import torch

# 创建一个最后一维尺寸为 1 的张量
x = torch.tensor([[1], [2], [3]])
print("原始张量形状:", x.shape)  # 输出: torch.Size([3, 1])

# 移除最后一个维度
y = x.squeeze(-1)
print("操作后张量形状:", y.shape)  # 输出: torch.Size([3])
print(y)
# 输出: tensor([1, 2, 3])

这里,原始张量是形状为 [3, 1] 的二维张量,通过 squeeze(-1) 移除了最后一个维度,变成了形状为 [3] 的一维张量。

示例 2:最后一个维度尺寸不为 1,张量不变

python

运行

# 创建一个最后一维尺寸不为 1 的张量
x = torch.tensor([[1, 2], [3, 4]])
print("原始张量形状:", x.shape)  # 输出: torch.Size([2, 2])

# 尝试移除最后一个维度
y = x.squeeze(-1)
print("操作后张量形状:", y.shape)  # 输出: torch.Size([2, 2])
print(y)
# 输出: 
# tensor([[1, 2],
#         [3, 4]])

由于最后一个维度的尺寸是 2,因此 squeeze(-1) 操作没有改变张量的形状。

示例 3:对三维张量操作

python

运行

# 创建一个三维张量
x = torch.randn(2, 1, 3)
print("原始张量形状:", x.shape)  # 输出: torch.Size([2, 1, 3])

# 移除倒数第二个维度(-2)
y1 = x.squeeze(-2)
print("移除倒数第二个维度后的形状:", y1.shape)  # 输出: torch.Size([2, 3])

# 尝试移除最后一个维度(尺寸为 3)
y2 = x.squeeze(-1)
print("尝试移除最后一个维度后的形状:", y2.shape)  # 输出: torch.Size([2, 1, 3])

这个示例展示了:

  • squeeze(-2) 成功移除了倒数第二个维度(尺寸为 1)
  • squeeze(-1) 未改变张量,因为最后一个维度尺寸为 3

实际应用场景

squeeze(-1) 常用于以下场景:

  1. 预处理输出数据:在模型输出中,某个维度可能意外地包含尺寸为 1 的维度,需要移除以便后续处理。

  2. 维度匹配:在进行张量运算前,确保两个张量的维度完全匹配。

  3. 数据可视化:在可视化前,将不必要的维度移除,使数据格式更适合可视化工具的要求。

注意事项

  • 仅移除尺寸为 1 的维度:如果指定维度的尺寸不是 1,squeeze 操作不会改变张量。
  • 与 unsqueeze 的关系unsqueeze 用于增加维度,而 squeeze 用于减少维度,二者常结合使用以调整张量形状。

通过上述示例,你应该能清晰理解 squeeze(-1) 的工作原理。如果有疑问,可以在 PyTorch 环境中尝试运行这些代码,观察结果。

分享

除了-1之外,squeeze还可以传入哪些参数?

如何使用squeeze函数处理多个维度为1的情况?

如何使用squeeze函数处理除最后一个维度之外的其他维度?

Logo

魔乐社区(Modelers.cn) 是一个中立、公益的人工智能社区,提供人工智能工具、模型、数据的托管、展示与应用协同服务,为人工智能开发及爱好者搭建开放的学习交流平台。社区通过理事会方式运作,由全产业链共同建设、共同运营、共同享有,推动国产AI生态繁荣发展。

更多推荐