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

如何利用LSTM历史输出与隐藏状态实现Luong全局点积注意力?

搞定Luong全局点积注意力解码器的输入依赖问题

嘿,我正好对这个注意力机制的实现熟得很,来给你捋清楚解码器那绕人的输入逻辑!你纠结的t时刻输入依赖t-1时刻输出和隐藏状态的问题,其实是神经机器翻译解码器的核心流程,下面一步步给你拆解:

1. 先把数据流逻辑掰明白

Luong的全局点积注意力里,解码器每一步的输入不是孤立的,得把上一步的输出和注意力上下文结合起来。完整的时间步流程是这样的:

  • 初始化:用编码器的最终隐藏状态启动解码器的初始隐藏状态dec_hidden,第一步的输入固定是句首标记<sos>的嵌入向量
  • 循环每个时间步t:
    1. 把当前的dec_input和dec_hidden喂给LSTM,得到t时刻的新隐藏状态new_dec_hidden和LSTM的临时输出
    2. 用这个新的隐藏状态,和编码器全程的所有隐藏状态计算全局点积注意力,得到上下文向量context
    3. 把LSTM的临时输出和context拼接(或者按Luong的其他方式结合),过线性层得到t时刻的最终输出dec_out
    4. 把dec_out过softmax得到预测词,把这个词的嵌入向量作为t+1时刻的dec_input
    5. 更新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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.19 04:10:20