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的维度匹配:
features[:,1:,:]是模型输出的特征,形状为(batch_size, seq_len-1, hidden_dim)(去掉了开头的<cls>特殊token)input_ids["attention_mask"][:,1:]是过滤掉特殊token后的mask,形状为(batch_size, seq_len-1)- 两者做元素乘法时,mask必须新增一个维度才能和特征的最后一维对齐,这样填充位置(mask为0)的特征会被置为0,后续求和时不会把填充部分的特征算进去,最后再除以有效token的数量(mask求和)得到平均特征。
内容的提问来源于stack exchange,提问作者learningtocode
相关产品推荐
相关产品推荐

