You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyTorch中unsqueeze(-1)与squeeze(1)是否等价?附代码解析

关于PyTorch中unsqueeze(-1)的解析及与squeeze(1)的区别

一、unsqueeze(-1)的具体含义

unsqueeze(-1)是PyTorch里给张量在最后一维的位置新增一个大小为1的维度,核心作用是适配张量运算的维度要求,触发广播机制。

举个实际例子:代码里的input_ids["attention_mask"][:, 1:],假设它的形状是(32, 128)(32是batch大小,128是去掉开头特殊token后的序列长度),经过unsqueeze(-1)后,形状会变成(32, 128, 1)。这样它就能和同批次、同序列长度但最后一维是隐藏层维度(比如(32,128,768))的特征张量做元素级乘法——PyTorch会自动把新增的1维广播到768的维度,实现每个token的所有隐藏层特征都被mask值(0或1)过滤。

二、unsqueeze(-1)和squeeze(1)完全不是一回事

这两个操作的作用完全相反,不可能相等:

  • unsqueeze是增加维度,squeeze是移除大小为1的维度
  • unsqueeze(-1)针对的是最后一维,squeeze(1)针对的是索引为1的维度(维度索引从0开始)

比如:

  • 若有形状为(32,1,128)的张量,squeeze(1)会把中间大小为1的维度删掉,变成(32,128),这和unsqueeze(-1)的增维操作方向完全相反。

三、结合你的代码看unsqueeze(-1)的作用

看代码里的关键计算逻辑:

features = torch.sum(features[:, 1:, :] * input_ids["attention_mask"][:, 1:].unsqueeze(-1), dim=1) / torch.clamp(
    torch.sum(input_ids["attention_mask"][:, 1:], dim=1, keepdims=True), min=1e-9)

这里的unsqueeze(-1)是为了让attention_mask和features的维度匹配:

  1. features[:,1:,:]是模型输出的特征,形状为(batch_size, seq_len-1, hidden_dim)(去掉了开头的<cls>特殊token)
  2. input_ids["attention_mask"][:,1:]是过滤掉特殊token后的mask,形状为(batch_size, seq_len-1)
  3. 两者做元素乘法时,mask必须新增一个维度才能和特征的最后一维对齐,这样填充位置(mask为0)的特征会被置为0,后续求和时不会把填充部分的特征算进去,最后再除以有效token的数量(mask求和)得到平均特征。

内容的提问来源于stack exchange,提问作者learningtocode

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.17 14:07:23