如何选取DistilBert输出的三维numpy.ndarray中的第二行
问题原因分析
你打印的vec是三维numpy数组,先执行print(vec.shape)可以看到输出为(1, 3, 768),三个维度的含义如下:
- 第0维:批次大小,你单次仅输入了1条文本,所以该维度长度为1,仅支持索引
0 - 第1维:token序列长度,"cat"经过DistilBERT分词后会生成
<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]>、cat、[SEP]三个token,所以该维度长度为3,支持索引0/1/2 - 第2维:隐藏层输出维度,distilbert-base-uncased的单token输出维度固定为768
你之前使用vec[1]未生效的核心原因是:该写法直接索引第0维,而第0维长度仅为1,vec[1]已经超出索引边界,会直接报错。
正确选取方式
如果你要选取序列维度的第二行(也就是对应cat这个token的向量),有两种等价写法:
- 逗号分隔多维度索引:
vec[0, 1],直接指定第0维取索引0,第1维取索引1,得到的是长度为768的一维数组 - 分步索引:先取第0维的唯一元素,再取第二行,写法为
vec[0][1],效果和第一种完全一致
如果你需要保留三维数组的结构,可以使用切片写法:vec[:, 1, :],得到的数组形状为(1, 1, 768)。
内容的提问来源于stack exchange,提问作者zest16
相关产品推荐
相关产品推荐

