Pytorch中gather()函数详解和实战示例

在 PyTorch 中,torch.gather() 是一个非常实用的张量操作函数,主要用于根据索引从输入张量中选择特定位置的值。它常用于注意力机制、序列处理等场景。


函数定义

python 复制代码
torch.gather(input, dim, index) → Tensor
  • input:待提取数据的张量。
  • dim:在哪个维度上进行索引选择。
  • index:一个与 input 在除了 dim 维度外相同形状的张量,其值指定了从 input 中提取的索引位置。
  • 返回值:从 input 的指定维度 dim 上根据 index 提取出的新张量。

形象理解

举个简单的例子:

示例 1:二维张量,按列(dim=1)提取

python 复制代码
import torch

input = torch.tensor([[10, 20, 30],
                      [40, 50, 60]])
index = torch.tensor([[2, 1, 0],
                      [0, 1, 2]])

output = torch.gather(input, dim=1, index=index)
print(output)

解释:

  • 对于第一行:从 [10, 20, 30] 中提取位置 [2,1,0],结果是 [30, 20, 10]
  • 对于第二行:从 [40, 50, 60] 中提取位置 [0,1,2],结果是 [40, 50, 60]

输出:

复制代码
tensor([[30, 20, 10],
        [40, 50, 60]])

示例 2:按行(dim=0)提取

python 复制代码
input = torch.tensor([[1, 2],
                      [3, 4],
                      [5, 6]])

index = torch.tensor([[0, 1],
                      [1, 2],
                      [2, 0]])

output = torch.gather(input, dim=0, index=index)
print(output)

解释:

  • 每个位置从第 dim=0 维度提取对应的元素。例如:

    • 第 (0,0) 位置:从 1,3,5 中取第 0 行,值为 1
    • 第 (1,0) 位置:从 1,3,5 中取第 1 行,值为 3
    • 第 (2,1) 位置:从 2,4,6 中取第 0 行,值为 2

输出:

复制代码
tensor([[1, 4],
        [3, 6],
        [5, 2]])

应用场景

  1. 注意力机制中的权重选择
  2. 序列解码中的 beam search
  3. 从嵌套表示中根据索引获取嵌套内容

实战场景举例

假设有一个 batch 的 BERT 输出,想从每个句子中提取第 N 个 token(如 CLS、某个关键词)的表示向量。


假设数据

python 复制代码
import torch
from transformers import BertModel, BertTokenizer

tokenizer = BertTokenizer.from_pretrained("bert-base-uncased")
model = BertModel.from_pretrained("bert-base-uncased")

sentences = ["I love World", "Transformers are powerful"]
inputs = tokenizer(sentences, padding=True, return_tensors="pt")

# 获取 BERT 输出
outputs = model(**inputs)
last_hidden_state = outputs.last_hidden_state  # (batch_size, seq_len, hidden_size)

print(last_hidden_state.shape)
# torch.Size([2, 5, 768])  假设 padding 后为长度 5,hidden size 为 768

场景 1:提取每个句子的第一个 token(通常是 CLS)

python 复制代码
cls_embeddings = last_hidden_state[:, 0, :]  # shape: (batch_size, hidden_size)

这个可以直接使用切片完成,不需要 gather。


场景 2:提取每个句子中 指定位置的 token 表示(如"love"或"are")

假设我们事先知道每个句子中感兴趣 token 的位置:
python 复制代码
# 每个句子中我们想要提取的 token 索引
# 假设我们想提取第 2 个 token
token_indices = torch.tensor([2, 1])  # shape: (batch_size,)

使用 gather 抽取对应 token 的向量:

python 复制代码
# last_hidden_state: (batch_size, seq_len, hidden_size)
batch_size, seq_len, hidden_size = last_hidden_state.size()

# 将 token_indices 转成 index 用于 gather: shape (batch_size, 1, 1)
token_indices = token_indices.view(-1, 1, 1).expand(-1, 1, hidden_size)  # (batch_size, 1, hidden_size)

# gather on dim=1(seq_len)
token_embeddings = torch.gather(last_hidden_state, dim=1, index=token_indices)  # (batch_size, 1, hidden_size)

# squeeze 掉中间的维度
token_embeddings = token_embeddings.squeeze(1)  # (batch_size, hidden_size)

print(token_embeddings.shape)

小结

操作需求 用法
取所有句子的第一个 token output[:, 0, :]
取所有句子的第 N 个 token output[:, N, :]
取每个句子的指定 token(不同位置) torch.gather()(如上所示)

注意事项

  • index 必须与 input 的 shape 一致,除了在指定的 dim 维度上的大小。
  • index 的值必须小于 input 在 dim 维度上的长度。

相关推荐
高升说几秒前
多路相机带宽与存储怎么算?从单路码率到盘位规划的完整链路
人工智能·数码相机
全栈道6 分钟前
AI 时代,软件开发应该如何学习
人工智能·学习
YOLO数据集集合6 分钟前
遥感滑坡图像识别数据集 | 滑坡识别 遥感影像 语义分割 灾害监测 无人机航拍 目标检测 深度学习数据集 计算机视觉9172期
人工智能·yolo·目标检测·计算机视觉·无人机·滑坡
czq_26867194877 分钟前
Python打卡第34天
开发语言·python·机器学习
FPGA信号处理10 分钟前
【凸优化】第一节课(下):保凸运算、分离超平面与 Farkas 引理
人工智能·算法·机器学习
计算机毕业编程指导师10 分钟前
大数据毕设选题推荐:基于Hadoop+Spark全球碳排放减排策略分析与可视化系统源码 毕业设计 选题推荐 毕设选题 数据分析 机器学习
大数据·hadoop·python·计算机·spark·毕业设计·碳排放减排
专业程序开发源12 分钟前
django便利店外卖后台管理系统19650-计算机课程设计、毕业设计
java·vue.js·spring boot·后端·python·django·课程设计
stsdddd16 分钟前
v11扑克牌数字检测数据集8992张实测:标注质量与划分比例一览
人工智能·计算机视觉·目标跟踪
商业白皮书18 分钟前
2026年上海豆包DeepSeekKimi百度AI通义千问GEO服务商选择
人工智能·笔记
toooooop818 分钟前
宝塔Python项目管理器:环境变量踩坑记录
java·jvm·python