如何利用LSTM历史输出与隐藏状态实现Luong全局点积注意力?
搞定Luong全局点积注意力解码器的输入依赖问题
嘿,我正好对这个注意力机制的实现熟得很,来给你捋清楚解码器那绕人的输入逻辑!你纠结的t时刻输入依赖t-1时刻输出和隐藏状态的问题,其实是神经机器翻译解码器的核心流程,下面一步步给你拆解:
1. 先把数据流逻辑掰明白
Luong的全局点积注意力里,解码器每一步的输入不是孤立的,得把上一步的输出和注意力上下文结合起来。完整的时间步流程是这样的:
- 初始化:用编码器的最终隐藏状态启动解码器的初始隐藏状态
dec_hidden,第一步的输入固定是句首标记<sos>的嵌入向量 - 循环每个时间步t:
- 把当前的
dec_input和dec_hidden喂给LSTM,得到t时刻的新隐藏状态new_dec_hidden和LSTM的临时输出 - 用这个新的隐藏状态,和编码器全程的所有隐藏状态计算全局点积注意力,得到上下文向量
context - 把LSTM的临时输出和
context拼接(或者按Luong的其他方式结合),过线性层得到t时刻的最终输出dec_out - 把
dec_out过softmax得到预测词,把这个词的嵌入向量作为t+1时刻的dec_input - 更新
dec_hidden为new_dec_hidden,进入下一轮循环
- 把当前的
2. 给你贴个可落地的PyTorch代码示例
假设你已经搭好了编码器,拿到了编码器的所有隐藏状态enc_hiddens,下面是解码器和注意力层的核心代码:
import torch import torch.nn as nn # 先实现Luong的点积注意力层 class LuongDotAttention(nn.Module): def __init__(self, hidden_size): super().__init__() self.hidden_size = hidden_size def forward(self, dec_hidden, enc_hiddens): # dec_hidden: [batch_size, hidden_size](解码器当前时间步的隐藏状态) # enc_hiddens: [seq_len_enc, batch_size, hidden_size](编码器所有时间步的隐藏状态) # 计算点积得分:解码器隐藏状态和每个编码器隐藏状态做点积 scores = torch.bmm(enc_hiddens.transpose(0, 1), dec_hidden.unsqueeze(2)).squeeze(2) # 计算注意力权重(softmax归一化) attn_weights = torch.softmax(scores, dim=1) # 加权求和得到上下文向量 context = torch.bmm(attn_weights.unsqueeze(1), enc_hiddens.transpose(0, 1)).squeeze(1) return context, attn_weights # 解码器单步处理函数 def decoder_single_step(dec_input, dec_hidden, enc_hiddens, decoder_lstm, attention, fc_layer): # dec_input: [batch_size, embed_size](上一步的词嵌入) # 喂给LSTM lstm_out, new_dec_hidden = decoder_lstm(dec_input.unsqueeze(0), dec_hidden) lstm_out = lstm_out.squeeze(0) # 调整维度到[batch_size, hidden_size] # 计算注意力上下文 context, attn_weights = attention(new_dec_hidden[0], enc_hiddens) # 结合LSTM输出和上下文(Luong的拼接方式) concat_output = torch.cat((lstm_out, context), dim=1) # 映射到词汇表维度得到最终输出 dec_out = fc_layer(concat_output) return dec_out, new_dec_hidden, attn_weights # 主流程示例 # 先定义一些超参数 batch_size = 32 hidden_size = 256 vocab_size = 10000 embed_size = 256 max_dec_len = 20 # 最大解码长度 # 初始化组件(假设你已经定义了embedding、decoder_lstm、fc_layer) embedding = nn.Embedding(vocab_size, embed_size) decoder_lstm = nn.LSTM(embed_size, hidden_size, batch_first=False) attention = LuongDotAttention(hidden_size) fc_layer = nn.Linear(hidden_size * 2, vocab_size) # 假设编码器的最终隐藏状态是enc_final_hidden(单向LSTM的情况) enc_final_hidden = torch.randn(1, batch_size, hidden_size) dec_hidden = enc_final_hidden # 初始化解码器隐藏状态 # 初始输入是<sos>标记的嵌入 sos_token_idx = 0 dec_input = embedding(torch.tensor([sos_token_idx] * batch_size)) # 开始循环解码 dec_outputs = [] attn_weights_history = [] for t in range(max_dec_len): dec_out, dec_hidden, attn_weights = decoder_single_step( dec_input, dec_hidden, enc_hiddens, decoder_lstm, attention, fc_layer ) dec_outputs.append(dec_out) attn_weights_history.append(attn_weights) # 推理阶段:用模型预测的词作为下一个输入 pred_token_idx = dec_out.argmax(dim=1) dec_input = embedding(pred_token_idx) # 训练阶段可以用teacher forcing,直接传入目标序列的t-1时刻词嵌入,不用等预测
3. 几个容易踩坑的关键点
- 训练和推理的输入差异:训练时为了稳定,通常用teacher forcing——直接把真实目标序列的上一个词嵌入作为当前输入,不用依赖模型的预测;推理时只能用模型上一步的预测词嵌入,这时候可能需要加beam search来提升效果。
- 双向编码器的适配:如果你的编码器是双向LSTM,那
enc_hiddens会包含两个方向的隐藏状态,需要先把它们拼接或者求和,让维度和解码器的隐藏状态匹配,同时解码器的初始隐藏状态也要处理双向编码器的最终状态(比如把两个方向的隐藏状态拼接起来)。 - 维度匹配:点积注意力要求解码器和编码器的隐藏状态维度一致,如果不一致,得给其中一个加个线性层映射到相同维度,不然点积计算会报错。
内容的提问来源于stack exchange,提问作者Tom
相关产品推荐
相关产品推荐

