从零实现 Transformer:第 0 部分 - 基础( Foundations)squeeze / unsqueeze 修改张量的维度结构(shape)
从零实现 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]])
更多推荐


所有评论(0)