FEATURED · 精选文章

PyTorch维度操作:unsqueeze与squeeze原理、应用与实战技巧

发布时间 / 2026/8/3 8:33:40
来源 / 创域科博编辑部
栏目 / 资讯中心
PyTorch维度操作:unsqueeze与squeeze原理、应用与实战技巧 1. 从一次广播错误说起为什么我们需要关心维度那天下午我正在调试一个图像分类模型想把一个形状为[32, 1, 224, 224]的批次图像张量和一个形状为[32, 224, 224]的掩码张量做逐元素相乘。直觉上这俩张量除了中间那个“1”其他维度都对齐了应该能直接运算。我信心满满地写下了images * masks结果PyTorch毫不留情地抛出了一个错误RuntimeError: The size of tensor a (1) must match the size of tensor b (224) at non-singleton dimension 2错误信息里的non-singleton dimension让我愣了一下。我仔细看了看问题就出在那个“1”上。在PyTorch的广播机制里维度为1的维度也叫单例维度或“虚”维度是特殊的它可以自动扩展来匹配另一个张量对应维度的大小。但这里我的masks张量在第二个维度索引为1的位置压根就没有维度它只有三个维度[32, 224, 224]。PyTorch试图将images的第二个维度大小为1与masks的第二个维度大小为224进行广播对齐但规则是要么维度大小相等要么其中一个为1。可现在是masks在那个位置“没有维度”这就不符合广播规则了。问题的根源在于维度不匹配。images是四维的[批次, 通道, 高, 宽]而masks被我处理成了三维的[批次, 高, 宽]。要让它们相乘我需要给masks在“通道”那个位置也插入一个维度变成[32, 1, 224, 224]这样广播机制就能愉快地工作了[32, 1, 224, 224] * [32, 1, 224, 224]其中masks的“1”会自动复制32次来匹配images的32个通道虽然这里每个样本只有一个通道的掩码。这个看似简单的“插入一个维度”的操作就是torch.unsqueeze()的用武之地。而它的逆操作torch.squeeze()则专门用来删除那些多余的、大小为1的维度让张量变得更“紧凑”避免在一些函数如某些损失函数或矩阵运算中因为多余的维度而报错。unsqueeze和squeeze是PyTorch张量操作中最基础、最高频的两个函数。它们不改变张量的实际数据只改变其“视图”即我们对数据组织方式的解释。理解它们是理解张量形状操作、广播机制乃至整个模型数据流的关键第一步。无论你是正在处理多模态数据如图像文本需要对齐不同来源张量的维度还是在搭建神经网络层需要调整输入输出的形状亦或是在进行简单的数据预处理这两个函数都是你工具箱里的必备利器。2. 维度的“增”与“删”unsqueeze与squeeze的核心原理要玩转unsqueeze和squeeze首先得抛开对维度顺序的刻板印象。我们习惯用[batch, channel, height, width]来描述图像用[batch, sequence_length, feature_dim]来描述序列。但维度本身只是一个索引数据的坐标轴。unsqueeze和squeeze操作的就是这些坐标轴。2.1torch.unsqueeze(dim)在指定位置插入一个新维度unsqueeze的功能是在张量的指定维度索引dim处插入一个大小为1的新维度。这里的dim指的是插入后新维度所在的位置。关键理解dim参数可以是负数。这是很多初学者困惑的地方。正数索引从左边开始0-based而负数索引从右边开始。dim-1表示在最后一个维度之后插入dim-2表示在倒数第二个维度之前插入以此类推。这个特性在你不确定张量总维数时非常有用。举个例子假设我们有一个三维张量t形状为[2, 3, 4]可以想象成2个样本每个样本是3x4的矩阵。import torch t torch.randn(2, 3, 4) print(t.shape) # torch.Size([2, 3, 4]) # 在维度0最前面插入变成 [1, 2, 3, 4] t1 t.unsqueeze(0) print(t1.shape) # torch.Size([1, 2, 3, 4]) # 在维度1原维度0和1之间插入变成 [2, 1, 3, 4] t2 t.unsqueeze(1) print(t2.shape) # torch.Size([2, 1, 3, 4]) # 在最后一个维度之后插入变成 [2, 3, 4, 1] t3 t.unsqueeze(-1) print(t3.shape) # torch.Size([2, 3, 4, 1]) # 在倒数第二个维度之前插入即原最后一个维度之前变成 [2, 3, 1, 4] t4 t.unsqueeze(-2) print(t4.shape) # torch.Size([2, 3, 1, 4])一个常见的应用场景为全连接层准备数据。全连接层 (nn.Linear) 期望的输入是[batch_size, feature_size]。如果你有一个单独的样本形状是[feature_size]即一个一维向量直接输入会报错因为PyTorch默认第一个维度是batch。这时你需要unsqueeze(0)将其变为[1, feature_size]表示批次大小为1。2.2torch.squeeze(dimNone)删除大小为1的维度squeeze是unsqueeze的逆操作它删除张量中所有大小为1的维度。如果指定了dim参数则只尝试删除该特定维度且仅当该维度大小为1时才生效否则张量保持不变。# 接上例t1的形状是 [1, 2, 3, 4] print(t1.shape) # torch.Size([1, 2, 3, 4]) # 不指定dim删除所有大小为1的维度变回 [2, 3, 4] t1_squeezed t1.squeeze() print(t1_squeezed.shape) # torch.Size([2, 3, 4]) # 注意t1_squeezed 和最初的 t 在数据上是相同的共享内存或经过复制。 # 指定删除维度0大小为1效果同上 t1_squeezed_dim0 t1.squeeze(0) print(t1_squeezed_dim0.shape) # torch.Size([2, 3, 4]) # 指定删除维度1大小为2不是1所以张量不变 t1_squeezed_dim1 t1.squeeze(1) print(t1_squeezed_dim1.shape) # torch.Size([1, 2, 3, 4]) # 形状未变 # 对于 t3 ([2, 3, 4, 1])不指定dim会删除最后一个维度 t3_squeezed t3.squeeze() print(t3_squeezed.shape) # torch.Size([2, 3, 4])为什么需要squeeze主要有两个原因减少干扰某些操作如损失函数nn.CrossEntropyLoss对输入形状有严格要求。例如分类任务中模型输出可能是[batch, num_classes, 1, 1]在某些CNN结构后而损失函数期望[batch, num_classes]这时就需要squeeze掉后面两个为1的维度。节省内存与提升可读性虽然大小为1的维度不增加数据量但它们会在代码中传递让张量的逻辑形状变得复杂。squeeze可以让张量形状更清晰也避免在一些检查中引发意外。注意unsqueeze和squeeze返回的通常是原张量的一个视图view这意味着它们与原始张量共享底层数据存储修改其中一个会影响另一个。这是一种高效的内存操作。但需要注意的是如果squeeze或unsqueeze操作导致张量在内存中的连续性contiguity被破坏某些后续操作如view()可能会要求你先调用.contiguous()。2.3 原位操作与函数式操作和大多数PyTorch张量操作一样unsqueeze和squeeze都有两种使用方式函数式操作torch.unsqueeze(input, dim)和torch.squeeze(input, dimNone)返回一个新的张量。原位操作tensor.unsqueeze_(dim)和tensor.squeeze_(dimNone)带下划线的版本会直接修改原张量。x torch.tensor([1, 2, 3]) y x.unsqueeze(0) # y是新的张量x不变 print(x.shape) # torch.Size([3]) print(y.shape) # torch.Size([1, 3]) x.unsqueeze_(0) # 直接修改x print(x.shape) # torch.Size([1, 3])原位操作可以节省一点内存但在计算梯度时需要小心因为它会覆盖原变量的值。在神经网络的前向传播中如果确定后续不再需要原张量使用原位操作是安全的。但在需要保留计算图的情况下建议使用函数式操作。3. 实战场景深度剖析从数据预处理到模型集成理解了基本原理后我们来看看unsqueeze和squeeze在真实项目中的用武之地。这些场景远比简单的维度加减要复杂和微妙。3.1 场景一图像与掩码的广播对齐开篇问题的解决回到最初的问题。我们有图像张量images形状为[32, 1, 224, 224]掩码张量masks形状为[32, 224, 224]。目标是让每个图像的每个通道这里只有一个通道与对应的掩码相乘。错误做法直接images * masks。因为维度不匹配PyTorch无法广播。正确做法使用unsqueeze为masks添加一个通道维度。# 假设 images: [32, 1, 224, 224], masks: [32, 224, 224] # 我们需要在masks的维度1通道维位置插入一个维度 masks_unsqueezed masks.unsqueeze(1) # 形状变为 [32, 1, 224, 224] # 或者使用更清晰的写法指明是通道维 # masks_unsqueezed masks.unsqueeze(dim1) result images * masks_unsqueezed # 现在可以广播了 print(result.shape) # torch.Size([32, 1, 224, 224])为什么是dim1因为在我们约定的图像张量形状[N, C, H, W]中索引1的位置代表通道维。我们需要让masks在这个位置有一个维度大小为1以便与images的通道维大小也为1进行广播。进阶思考如果images是RGB三通道图[32, 3, 224, 224]而masks仍然是单通道的[32, 224, 224]我们仍然希望用同一个掩码作用于所有三个颜色通道。这时masks.unsqueeze(1)得到[32, 1, 224, 224]在与[32, 3, 224, 224]相乘时masks在通道维上的1会自动扩展为3实现“一对三”的掩码操作。这是广播机制的强大之处。3.2 场景二序列数据处理与注意力机制在自然语言处理中我们经常处理序列数据。假设我们有一批文本经过嵌入层后得到张量embeddings形状为[batch_size, seq_len, hidden_dim]例如[16, 50, 768]。现在我们要计算一个注意力权重向量attention_weights它的形状是[batch_size, seq_len]即[16, 50]代表每个序列中每个词的重要性。我们想用这个权重对隐藏状态进行加权求和通常称为注意力池化。一种简单的方法是# embeddings: [16, 50, 768] # attention_weights: [16, 50] # 我们需要将权重应用到最后一个维度hidden_dim上 # 第一步将权重从 [16, 50] 变为 [16, 50, 1]以便与embeddings广播 weights attention_weights.unsqueeze(-1) # 形状 [16, 50, 1] # 第二步逐元素相乘权重会广播到768维 weighted_embeddings embeddings * weights # 形状 [16, 50, 768] # 第三步在序列长度维度上求和得到加权后的句子表示 context_vector weighted_embeddings.sum(dim1) # 形状 [16, 768]这里unsqueeze(-1)是关键一步它在attention_weights的末尾添加了一个维度使其形状变为[16, 50, 1]。这样在与[16, 50, 768]相乘时权重会在最后一个维度隐藏维度上自动复制768次实现每个隐藏单元都按相同权重缩放的效果。反过来squeeze也经常出现在序列模型中。例如一个双向LSTM最后一层的输出可能是[batch, seq_len, num_directions * hidden_size]如果你只取最后一个时间步的输出 (output[:, -1, :])它的形状是[batch, num_directions * hidden_size]。但如果你在处理序列分类任务时用了nn.LSTM并设置了batch_firstTrue且希望获取最后一个时间步的隐藏状态你可能会得到形状为[batch, 1, hidden_size]的张量取决于如何索引。为了送入后续的全连接层你需要squeeze(1)去掉中间的维度。3.3 场景三损失函数输入的形状适配这是新手踩坑的重灾区。以交叉熵损失nn.CrossEntropyLoss为例它要求两个输入input模型的原始输出未经过Softmax形状为[batch_size, num_classes]。target真实标签形状为[batch_size]每个值是类别索引0到num_classes-1。假设你有一个简单的CNN用于MNIST分类10类最后一层是nn.Linear(512, 10)。前向传播后你得到的output形状是[batch, 10]这很好。但如果你用的网络结构比较复杂比如在全局平均池化后你可能会得到一个形状为[batch, 10, 1, 1]的输出这在一些迁移学习模型中很常见。# 模拟一个“奇怪”的输出形状 output torch.randn(32, 10, 1, 1) # [batch, classes, 1, 1] target torch.randint(0, 10, (32,)) # [batch] loss_fn nn.CrossEntropyLoss() # 直接计算会报错4D target tensor is not supported # loss loss_fn(output, target) # RuntimeError! # 正确做法squeeze掉后面两个为1的维度 output_corrected output.squeeze() # 形状变为 [32, 10] loss loss_fn(output_corrected, target) # 现在可以了 print(loss)同样的问题也出现在target上。如果你不小心把target做成了[batch, 1]的形状例如从DataLoader中取出时没有处理也需要squeeze()一下。经验之谈在将张量送入损失函数之前养成检查形状的习惯。对于CrossEntropyLoss一个简单的断言很有用assert input.shape (batch, num_classes) and target.shape (batch,)。如果形状不对squeeze和unsqueeze就是你快速修复的工具。3.4 场景四自定义层与维度兼容性当你编写自定义的PyTorch层或函数时考虑输入维度的灵活性是一个好习惯。你的函数可能期望某种特定维度的输入但用户可能传入不同维度的数据。这时unsqueeze/squeeze可以帮助你进行内部标准化。例如你写了一个计算逐元素高斯加权的函数它期望输入是[..., H, W]的形式...表示任意多的批次维度并对最后两个空间维度进行加权。def spatial_gaussian_weight(x): 对输入x的最后两个维度进行高斯加权。 x: 张量形状为 [..., H, W] 返回: 加权后的张量形状不变 # 假设我们有一个简单的二维高斯核形状为 [H, W] H, W x.shape[-2], x.shape[-1] # 这里简化创建高斯核的过程 gaussian_kernel torch.randn(H, W) # 仅为示例实际应为高斯分布 # 为了进行逐元素相乘我们需要将核扩展到与x相同的维度 # 但x可能有前面的批次维度如 [B, C, H, W] 或 [B, H, W] # 我们需要让gaussian_kernel的形状变为 [1, 1, H, W] 或 [1, H, W] 以便广播 # 一个通用的方法是在核的前面添加足够的维度1直到其维度和x一样多 while gaussian_kernel.dim() x.dim(): gaussian_kernel gaussian_kernel.unsqueeze(0) # 现在gaussian_kernel的形状比如是 [1, 1, H, W] (如果x是4D) # 或者 [1, H, W] (如果x是3D) return x * gaussian_kernel # 测试 x_4d torch.randn(2, 3, 5, 5) # [B, C, H, W] x_3d torch.randn(2, 5, 5) # [B, H, W] print(spatial_gaussian_weight(x_4d).shape) # torch.Size([2, 3, 5, 5]) print(spatial_gaussian_weight(x_3d).shape) # torch.Size([2, 5, 5])在这个函数中unsqueeze(0)被循环使用动态地将二维高斯核的维度提升到与输入x相同从而实现了灵活的广播。这使得函数能够处理不同批次和通道维度的输入增强了代码的鲁棒性。4. 高级技巧、常见陷阱与性能考量掌握了基本应用后我们来看看一些更深入的知识点和容易踩的坑。4.1view()、reshape()与squeeze/unsqueeze的协同与区别view()和reshape()也可以改变形状但它们与squeeze/unsqueeze有本质区别squeeze/unsqueeze只操作大小为1的维度进行纯粹的维度增减不改变元素间的相对顺序和总数。它们是“安全”的形状变换通常能返回一个视图。view()/reshape()可以改变任意维度的大小但必须保证变换前后元素总数一致。它们可能会改变数据在内存中的布局view()要求张量是连续的否则会报错reshape()在可能的情况下返回视图否则返回拷贝。它们经常结合使用# 将一个三维张量 [2, 3, 4] 扁平化为二维 [2, 12] x torch.randn(2, 3, 4) # 方法1: 先用 view 展平后两维但需要知道具体大小 x_flat1 x.view(2, -1) # -1 表示自动推断 # 方法2: 更通用的方法适用于不知道后几维总大小的情况 # 先合并后两维再展平 x_squeezed x.flatten(start_dim1) # 从维度1开始展平得到 [2, 12] # flatten 内部其实做了类似 view 的操作 # 一个更复杂的例子将 [B, C, H, W] 转换为 [B, C*H*W] 以送入全连接层 batch, chan, height, width x_4d.shape x_fc_input x_4d.view(batch, -1) # 常用 # 或者 x_fc_input x_4d.flatten(1) # 更清晰从第1维C开始展平 # 转换回来后可能需要 unsqueeze # 假设我们从全连接层得到一个 [B, C] 的输出想把它变成 [B, C, 1, 1] 以模拟空间维度 fc_output torch.randn(batch, chan) spatial_output fc_output.view(batch, chan, 1, 1) # 这等价于 spatial_output fc_output.unsqueeze(-1).unsqueeze(-1) # 或者 spatial_output fc_output[:, :, None, None] # 使用None索引是unsqueeze的语法糖关键区别view()和reshape()关心的是形状的重新排列而squeeze和unsqueeze关心的是维度的存在与否。当你只是需要增加或删除一个“虚”维度时用后者更直观、更安全。4.2 使用None索引进行快速维度操作PyTorch和NumPy支持使用None在NumPy中是np.newaxis在索引中插入新维度这本质上是unsqueeze的语法糖非常简洁。x torch.randn(2, 3) print(x.shape) # [2, 3] # 在维度0插入 y1 x[None, :, :] # 等价于 x.unsqueeze(0) print(y1.shape) # [1, 2, 3] # 在维度1插入 y2 x[:, None, :] # 等价于 x.unsqueeze(1) print(y2.shape) # [2, 1, 3] # 在最后一个维度之后插入 y3 x[:, :, None] # 等价于 x.unsqueeze(-1) print(y3.shape) # [2, 3, 1] # 甚至可以同时插入多个 y4 x[None, :, :, None] # 等价于 x.unsqueeze(0).unsqueeze(-1) print(y4.shape) # [1, 2, 3, 1]这种方式在临时需要增加维度进行广播计算时特别方便代码更紧凑。但它的可读性略差尤其是对于不熟悉该语法的读者。在重要的、需要清晰表达的代码中我倾向于使用unsqueeze因为函数名本身就是文档。4.3 内存连续性Contiguity的潜在影响这是一个高级但重要的主题。PyTorch张量在内存中的存储有“连续”和“不连续”之分。像transpose(),permute(),narrow(),select()等操作返回的是原张量的视图它们改变了索引方式但可能破坏了内存的连续性。而view()严格要求张量是连续的否则会报错。reshape()会尝试返回视图如果不行就返回拷贝。squeeze和unsqueeze通常返回视图而且它们一般不会破坏连续性因为它们只是增加或删除一个大小为1的维度不改变元素顺序。但是如果你对一个不连续张量进行squeeze/unsqueeze操作结果张量也可能是不连续的。在绝大多数情况下你不需要关心这个。但如果你在squeeze/unsqueeze之后立即调用view()或者进行一些需要连续内存的低级操作如与C/C扩展交互可能会遇到问题。x torch.randn(2, 3, 4) y x.transpose(1, 2) # y的形状是[2, 4, 3]并且是不连续的 print(y.is_contiguous()) # False z y.unsqueeze(1) # z的形状是[2, 1, 4, 3] print(z.is_contiguous()) # 可能仍然是False # 如果此时想用 view 改变形状可能会报错 # w z.view(2, -1) # 可能报错: RuntimeError: view size is not compatible... # 安全的做法是先使其连续 w z.contiguous().view(2, -1) print(w.shape) # torch.Size([2, 12])经验法则如果你进行了一系列复杂的维度变换操作特别是包含transpose,permute然后在最后需要改变形状时先调用.contiguous()再view()是更稳妥的做法。对于squeeze/unsqueeze单独使用则基本不用担心。4.4 常见陷阱与调试技巧dim参数索引错误这是最常见的错误。时刻记住dim参数指的是操作后的维度索引位置。对于unsqueeze(dim)dim的范围是[-input.dim()-1, input.dim()]。一个有用的调试方法是打印操作前后的shape。x torch.randn(5, 10) print(fOriginal shape: {x.shape}) # [5, 10] for dim in range(-x.dim()-1, x.dim()1): try: y x.unsqueeze(dim) print(funsqueeze(dim{dim:2d}) - shape: {y.shape}) except IndexError as e: print(funsqueeze(dim{dim:2d}) - Error: {e})过度squeeze使用无参数的squeeze()会删除所有大小为1的维度。这有时会过度删除你本想保留的维度。例如一个形状为[1, 10, 1, 20]的张量经过squeeze()会变成[10, 20]丢失了批次和通道信息。更好的做法是明确指定dim参数只删除你确定无用的维度。与size()和shape属性的混淆x.size(dim)返回第dim维的大小而x.dim()返回总维数。在动态计算unsqueeze的dim参数时这些方法很有用。# 动态地在倒数第二个维度之前插入 x torch.randn(3, 4, 5) insert_dim -2 # 等价于 x.dim() - 1 y x.unsqueeze(insert_dim) print(y.shape) # [3, 4, 1, 5]广播语义理解不清unsqueeze的核心目的是为了广播。如果你unsqueeze后仍然无法进行运算请仔细检查两个张量的形状在纸上画出它们的维度并回想广播规则从后往前对齐维度每个维度要么相等要么其中一个为1要么其中一个不存在可以理解为1。unsqueeze就是为了让“不存在”变成“存在且为1”。原地操作 (_方法) 的梯度在自定义nn.Module的forward方法中如果对需要求导的张量使用了squeeze_()或unsqueeze_()这不会破坏梯度传播因为这只是改变了张量的元数据形状、步长等而不是数据本身。但出于代码清晰和避免副作用的考虑在大多数情况下使用非原位操作是更好的选择。unsqueeze和squeeze就像张量维度世界的“精细手术刀”虽小但不可或缺。它们不直接参与复杂的数学计算却是构建正确数据流、连接不同模块的桥梁。从数据加载时的维度调整到模型内部的特征变换再到损失计算前的形状适配几乎贯穿了深度学习项目的数据处理全流程。理解并熟练运用它们能让你在调试“The size of tensor a must match the size of tensor b”这类错误时更加得心应手写出更健壮、更清晰的PyTorch代码。下次当你需要对维度动手脚时先问问自己我是需要增加一个维度来广播还是需要删除一个冗余的维度来简化想清楚了这个问题unsqueeze和squeeze自然会成为你的得力助手。
RELATED — 相关阅读

相关资讯

LATEST — 最新资讯

最新发布

TODAY — 本日精选

新闻

WEEKLY — 本周精选

新闻

MONTHLY — 本月精选

新闻