1. 概述
作为 PyTorch 张量函数
这是最常见的用法,作用是去掉张量中维度为 1 的轴。
- 基本用法 :
x.squeeze()会移除所有大小为 1 的维度;x.squeeze(dim)只压缩指定维度。 - 注意事项:只有该维度大小确实是 1 时才会被移除,否则张量形状不变。
- 反操作 :
unsqueeze()是在指定位置增加一个维度为 1 的轴,两者常配合使用。 - 常见场景:处理单样本数据(去掉 batch 维度)、匹配损失函数输入、模型输出后处理等。
2. sequeeze()
import torch
x = torch.randn(1, 3, 1, 4)
print(x.shape) # torch.Size(1, 3, 1, 4)
print(x.squeeze().shape) # torch.Size(3, 4)
3. unsequeeze()
import torch
x = torch.randn(3, 4)
print(x.shape) # torch.Size(3, 4)
在第0维前插入新维度,变成 (1, 3, 4)
y = x.unsqueeze(0)
print(y.shape) # torch.Size(1, 3, 4)
在第1维插入新维度,变成 (3, 1, 4)
z = x.unsqueeze(1)
print(z.shape) # torch.Size(3, 1, 4)
在最后一维后插入,变成 (3, 4, 1)
w = x.unsqueeze(-1)
print(w.shape) # torch.Size(3, 4, 1)