感谢喜欢,我最近在忙别的项目,一直没腾出时间继续学习。但这只是暂停,不是中断
神经网络束搜索比如第一列有k个数据,束大小为b,1<b<k ,按照b个分支进行贪心搜索。那我们可以认为这是纵向的(同一列)。那我们也可以通过横向(不同列)来进行搜索啊。贪心搜索最大的问题就是第n列最大的数据会影响到第n+1列,同时取这两个列的最大概率不一定是全局最大。我们也可以做一个局部的穷举搜索。当我们确定第n列时,第n+1列是确定的,比如我们选k为2,穷举n+1,n+2列的所有可能,选出第n+1列最大的字符,依次向下,这个与贪心搜索不同的就是这个考虑了n+2列,束b越大,局部穷举越大。
这种方法学术届有人研究过吗?
全靠AI,我是傻子
def predict_seq2seq(net, src_sentence, src_vocab, tgt_vocab, num_steps,
device, save_attention_weights=False, beam_size=2):
"""序列到序列模型的束搜索预测"""
net.eval()
# 源序列预处理
src_tokens = src_vocab[src_sentence.lower().split(' ')] + [
src_vocab['<eos>']]
enc_valid_len = torch.tensor([len(src_tokens)], device=device)
src_tokens = d2l.truncate_pad(src_tokens, num_steps, src_vocab['<pad>'])
enc_X = torch.unsqueeze(
torch.tensor(src_tokens, dtype=torch.long, device=device), dim=0)
# 编码器前向(源序列编码仅执行一次)
enc_outputs = net.encoder(enc_X, enc_valid_len)
dec_state = net.decoder.init_state(enc_outputs, enc_valid_len)
bos_id = tgt_vocab['<bos>']
eos_id = tgt_vocab['<eos>']
# 初始化束候选:(完整序列, 累计对数概率, 解码器隐状态, 注意力权重序列)
beams = [
(
[bos_id], # 初始序列以 <bos> 开头
0.0, # 初始对数概率 log(P=1) = 0
dec_state, # 解码器初始隐状态
[] # 注意力权重缓存(可选)
)
]
completed = [] # 已生成 <eos> 的完整序列
for _ in range(num_steps):
if not beams:
break
candidates = []
# 逐个扩展当前束中的所有候选
for seq, log_prob, state, attn_list in beams:
# 当前步输入:序列最后一个 token,形状 (1, 1)
dec_X = torch.tensor([[seq[-1]]], dtype=torch.long, device=device)
# 解码器单步前向
Y, new_state = net.decoder(dec_X, state)
# 计算对数概率并取 Top-k 候选(避免遍历全词表,提升效率)
log_probs = torch.log_softmax(Y.squeeze(0).squeeze(0), dim=-1)
topk_log_probs, topk_ids = torch.topk(log_probs, beam_size)
# 扩展出 beam_size 个新候选
for i in range(beam_size):
token_id = topk_ids[i].item()
token_logp = topk_log_probs[i].item()
new_seq = seq + [token_id]
new_logp = log_prob + token_logp
# 缓存注意力权重
new_attn = attn_list.copy() if save_attention_weights else None
if save_attention_weights:
new_attn.append(net.decoder.attention_weights)
# 遇到结束符则加入完成列表,不再参与后续扩展
if token_id == eos_id:
completed.append((new_seq, new_logp, new_attn))
else:
candidates.append((new_seq, new_logp, new_state, new_attn))
# 全局排序,保留概率最高的前 beam_size 个候选进入下一步
candidates.sort(key=lambda x: x[1], reverse=True)
beams = candidates[:beam_size]
# 将未自然结束的候选也加入最终候选池
for seq, log_prob, state, attn_list in beams:
completed.append((seq, log_prob, attn_list))
# 选出全局概率最高的序列
completed.sort(key=lambda x: x[1], reverse=True)
best_seq, _, best_attn = completed[0]
# 去掉开头 <bos> 和结尾 <eos>,转回文本
output_tokens = best_seq[1:]
if output_tokens and output_tokens[-1] == eos_id:
output_tokens = output_tokens[:-1]
return ' '.join(tgt_vocab.to_tokens(output_tokens)), best_attn
