从零实现 Transformer:第 0 部分 - 基础( Foundations)squeeze / unsqueeze 修改张量的维度结构(shape)

flyfish

测试环境 PyTorch 版本 2.8.0+cu128

Tensor 的维度

从 0 开始数,第 1 个维度是 dim=0,第 2 个是 dim=1,以此类推
在这里插入图片描述
这个张量的形状是[2, 3, 3]

对应关系

维度(dim) 含义 元素数量 对应标注
dim=0 最外层的大括号数量 2 红色的 2
dim=1 中间层的中括号数量 3 蓝色的 3
dim=2 最内层的数字数量 3 青色的 3

PyTorch里的样子

import torch

# 定义形状为 (2, 3, 3) 的张量,和图片内容完全一致
tensor = torch.tensor([
    # 第 0 维的第 1 个元素(对应最外层的第 1 个部分)
    [
        [1, 2, 3],
        [1, 2, 3],
        [1, 2, 3]
    ],
    # 第 0 维的第 2 个元素(对应最外层的第 2 个部分)
    [
        [1, 2, 3],
        [1, 2, 3],
        [1, 2, 3]
    ]
])

# 打印张量内容和形状
print("张量内容:")
print(tensor)
print("\n张量形状:", tensor.shape)

运行结果

张量内容:
tensor([[[1, 2, 3],
         [1, 2, 3],
         [1, 2, 3]],

        [[1, 2, 3],
         [1, 2, 3],
         [1, 2, 3]]])

张量形状: torch.Size([2, 3, 3])

在这里插入图片描述
形状变化

原始张量形状:[2, 2]

--- unsqueeze(0) + squeeze(0) ---
[2, 2]  → unsqueeze(0) →  [1, 2, 2]
[1, 2, 2] → squeeze(0) →  [2, 2]

--- unsqueeze(1) + squeeze(1) ---
[2, 2]  → unsqueeze(1) →  [2, 1, 2]
[2, 1, 2] → squeeze(1) →  [2, 2]

--- unsqueeze(2) + squeeze(2) ---
[2, 2]  → unsqueeze(2) →  [2, 2, 1]
[2, 2, 1] → squeeze(2) →  [2, 2]
import torch

# 原始的 2D 张量
x = torch.tensor([[1, 2],
                  [3, 4]])
print("原始2D张量:")
print(x)
print("形状:", x.shape)
print("-" * 30)

# 1. unsqueeze(0):在第0维增加维度(对应图片蓝色箭头)
x_unsq0 = x.unsqueeze(0)
print("unsqueeze(0) 后:")
print(x_unsq0)
print("形状:", x_unsq0.shape)

# squeeze(0):去掉第0维,还原为2D张量
x_sq0 = x_unsq0.squeeze(0)
print("squeeze(0) 还原后:")
print(x_sq0)
print("形状:", x_sq0.shape)
print("-" * 30)

# 2. unsqueeze(1):在第1维增加维度(对应图片橙色箭头)
x_unsq1 = x.unsqueeze(1)
print("unsqueeze(1) 后:")
print(x_unsq1)
print("形状:", x_unsq1.shape)

# squeeze(1):去掉第1维,还原为2D张量
x_sq1 = x_unsq1.squeeze(1)
print("squeeze(1) 还原后:")
print(x_sq1)
print("形状:", x_sq1.shape)
print("-" * 30)

# 3. unsqueeze(2):在第2维增加维度(对应图片青色箭头)
x_unsq2 = x.unsqueeze(2)
print("unsqueeze(2) 后:")
print(x_unsq2)
print("形状:", x_unsq2.shape)

# squeeze(2):去掉第2维,还原为2D张量
x_sq2 = x_unsq2.squeeze(2)
print("squeeze(2) 还原后:")
print(x_sq2)
print("形状:", x_sq2.shape)

输出

原始2D张量:
tensor([[1, 2],
        [3, 4]])
形状: torch.Size([2, 2])
------------------------------
unsqueeze(0):
tensor([[[1, 2],
         [3, 4]]])
形状: torch.Size([1, 2, 2])
squeeze(0) 还原后:
tensor([[1, 2],
        [3, 4]])
形状: torch.Size([2, 2])
------------------------------
unsqueeze(1):
tensor([[[1, 2]],

        [[3, 4]]])
形状: torch.Size([2, 1, 2])
squeeze(1) 还原后:
tensor([[1, 2],
        [3, 4]])
形状: torch.Size([2, 2])
------------------------------
unsqueeze(2):
tensor([[[1],
         [2]],

        [[3],
         [4]]])
形状: torch.Size([2, 2, 1])
squeeze(2) 还原后:
tensor([[1, 2],
        [3, 4]])
形状: torch.Size([2, 2])

squeeze

torch.squeeze(input, dim) -> Tensor

torch.squeeze(input: Tensor, dim: int | List[int] | None) → Tensor

对输入张量中所有指定的、大小为1的维度进行移除操作,并返回处理后的张量。

例如,若输入张量的形状为 (A×1×B×C×1×D),则调用 input.squeeze() 后,张量形状将变为 (A×B×C×D)

当指定参数 dim 时,仅会对指定维度执行压缩操作。若输入张量的形状为 (A×1×B),执行 squeeze(input, 0) 不会改变张量形状;而执行 squeeze(input, 1) 会将张量压缩为形状 (A×B)

注意

返回的张量与输入张量共享存储空间,因此修改其中一个张量的数值,会同步改变另一个张量的数值。

警告

若张量的批次维度(batch dimension)大小为1,直接调用 squeeze(input) 会同时移除该批次维度,可能引发意料之外的错误。建议仅指定需要压缩的维度进行操作。

参数

input (Tensor):输入张量
dim (int 或 整数元组, 可选)
若指定该参数,仅会对指定维度执行压缩操作。
2.0 版本更新dim 现在支持传入维度元组。

代码示例

x = torch.zeros(2, 1, 2, 1, 2)
x.size()  # 输出: torch.Size([2, 1, 2, 1, 2])

# 移除所有大小为1的维度
y = torch.squeeze(x)
y.size()  # 输出: torch.Size([2, 2, 2])

# 对第0维压缩(该维度大小≠1,无变化)
y = torch.squeeze(x, 0)
y.size()  # 输出: torch.Size([2, 1, 2, 1, 2])

# 对第1维压缩(该维度大小=1,成功移除)
y = torch.squeeze(x, 1)
y.size()  # 输出: torch.Size([2, 2, 1, 2])

# 同时对第1、2、3维压缩
y = torch.squeeze(x, (1, 2, 3))
y.size()  # 输出: torch.Size([2, 2, 2])

unsqueeze

torch.unsqueeze(input, dim) -> Tensor

返回一个新张量,并在指定位置插入一个大小为 1 的维度。
返回的新张量与原张量共享底层数据(不会复制数据,仅修改维度结构)。

参数 dim 的合法取值范围为:[-input.dim() - 1, input.dim() + 1)
若传入负的 dim 值,等效于将 dim 替换为 dim + input.dim() + 1 后执行插入操作。

参数

input (Tensor):输入张量
dim (int):用于插入单维度(大小为1的维度)的索引位置

代码示例

# 定义一维张量
x = torch.tensor([1, 2, 3, 4])

# 在第0维插入大小为1的维度
torch.unsqueeze(x, 0)
# 输出结果:tensor([[ 1,  2,  3,  4]])

# 在第1维插入大小为1的维度
torch.unsqueeze(x, 1)
# 输出结果:
# tensor([[ 1],
#         [ 2],
#         [ 3],
#         [ 4]])
Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐