【发布时间】:2021-01-29 01:47:48
【问题描述】:
我正在尝试在文本生成模型中实现波束搜索解码策略。这是我用来解码输出概率的函数。
def beam_search_decoder(data, k):
sequences = [[list(), 0.0]]
# walk over each step in sequence
for row in data:
all_candidates = list()
for i in range(len(sequences)):
seq, score = sequences[i]
for j in range(len(row)):
candidate = [seq + [j], score - torch.log(row[j])]
all_candidates.append(candidate)
# sort candidates by score
ordered = sorted(all_candidates, key=lambda tup:tup[1])
sequences = ordered[:k]
return sequences
现在你可以看到这个函数是在考虑到 batch_size 1 的情况下实现的。为批量大小添加另一个循环将使算法O(n^4)。和现在一样慢。有什么办法可以提高这个功能的速度。我的模型输出通常大小为(32, 150, 9907),格式为(batch_size, max_len, vocab_size)
【问题讨论】:
-
您应该在 Pytorch 中通过快速 google 搜索找到光束搜索实现。请注意,波束解码并非易事,有几个因素相关。因此,建议使用其中一种可用的实现。
-
束搜索策略在测试期间是有意义的。你不能维护一个
batch_size=1并并行处理测试示例吗? -
您还可以查看beam search implementation 和代码in this repo 使用修改后的Transformer 进行图像字幕。该实现使用 PyTorch 的
register_buffer来缓存前一个时间步的输入,以便在当前时间步中只提供新的输入,并且速度相当快。
标签: python deep-learning nlp pytorch beam-search