PyTorch中如何用列表索引提取BERT特定层的Token激活值?
解决BERT张量非连续层索引的优雅方法
你的报错根源是三个索引的形状不匹配:indices[:,0]和indices[:,1]都是长度为3的一维张量,而层索引[1,3,5,7]是长度为4的列表,PyTorch无法自动广播这三个形状。下面是几种无需拆分张量的优雅实现方式:
方法1:通过维度扩展实现广播索引
把前两个索引扩展为**(3,1)的形状,层索引扩展为(1,4)的形状,这样三者就能广播成(3,4)**的索引矩阵,直接完成提取:
import torch example = torch.randn([3, 12, 13, 768]) indices = torch.tensor([[0, 1], [1, 10], [2, 11]]) target_layers = [1, 3, 5, 7] # 扩展维度:前两个索引变成(3,1),层索引保持(1,4)自动广播 c = example[indices[:, 0].unsqueeze(1), indices[:, 1].unsqueeze(1), target_layers] print(c.shape) # torch.Size([3, 4, 768])
方法2:使用torch.index_select简化操作
先提取所有目标token的全部层激活,再对层维度进行索引选择:
# 先提取所有目标token的全部层 all_layers = example[indices[:,0], indices[:,1], :] # 对层维度(dim=1)选择指定层 c = torch.index_select(all_layers, dim=1, index=torch.tensor(target_layers)) print(c.shape) # torch.Size([3, 4, 768])
方法3:使用torch.take_along_dim(PyTorch 1.10+)
这种方式更直观指定要提取的位置,适合复杂索引场景:
# 构造层索引的形状:(3,4),每个token对应相同的目标层 layer_indices = torch.tensor(target_layers).repeat(3, 1) # 在层维度(dim=1)上提取指定位置的激活 c = torch.take_along_dim(example[indices[:,0], indices[:,1], :], layer_indices.unsqueeze(-1), dim=1).squeeze(-1) print(c.shape) # torch.Size([3, 4, 768])
所有方法都能直接得到你想要的形状,无需拆分张量,其中方法1最简洁,适合大部分场景。
内容的提问来源于stack exchange,提问作者liamthorne4
相关产品推荐
相关产品推荐

