在 PyTorch 中,torch.gather()
是一个非常实用的张量操作函数,主要用于根据索引从输入张量中选择特定位置的值。它常用于注意力机制、序列处理等场景。
torch.gather(input, dim, index) → Tensor
input
:待提取数据的张量。dim
:在哪个维度上进行索引选择。index
:一个与 input
在除了 dim
维度外相同形状的张量,其值指定了从 input
中提取的索引位置。input
的指定维度 dim
上根据 index
提取出的新张量。举个简单的例子:
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]])
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
维度提取对应的元素。例如:
输出:
tensor([[1, 4],
[3, 6],
[5, 2]])
假设有一个 batch 的 BERT 输出,想从每个句子中提取第 N 个 token(如 [CLS]、某个关键词)的表示向量。
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
cls_embeddings = last_hidden_state[:, 0, :] # shape: (batch_size, hidden_size)
这个可以直接使用切片完成,不需要 gather
。
# 每个句子中我们想要提取的 token 索引
# 假设我们想提取第 2 个 token
token_indices = torch.tensor([2, 1]) # shape: (batch_size,)
gather
抽取对应 token 的向量:# 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
维度上的长度。