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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.10 20:10:33