不同隐藏层大小的LSTM如何实现编解码注意力机制?
不同隐藏层大小的编解码注意力实现方案
你遇到的核心问题是:编解码注意力中,编码器输出(来自LSTM1,维度[batch, 序列长度, 20])和解码器的隐藏状态(来自LSTM2,维度[batch, 40])最后一维维度不一致,导致点积无法直接计算。下面给两种实用的解决方法:
方法一:线性变换对齐维度
通过可训练的线性层(全连接层)把其中一方的维度转换成和另一方相同,再用常规的点积注意力计算。
方案1:把编码器输出转成解码器隐藏层维度
import tensorflow as tf from tensorflow.keras.layers import Dense # 假设encoder_outputs是LSTM1的输出,shape=(batch_size, enc_seq_len, 20) # decoder_hidden是LSTM2的当前时间步隐藏状态,shape=(batch_size, 40) # 用线性层将编码器输出从20维转成40维 encoder_proj = Dense(40)(encoder_outputs) # shape: (batch_size, enc_seq_len, 40) # 解码器隐藏层扩维,变成(batch_size, 1, 40),适配点积运算 decoder_hidden_expanded = tf.expand_dims(decoder_hidden, axis=1) # shape: (batch_size, 1, 40) # 计算注意力得分:点积后得到(batch_size, enc_seq_len, 1) attention_scores = tf.matmul(encoder_proj, decoder_hidden_expanded, transpose_b=True) # 归一化得分得到注意力权重 attention_weights = tf.nn.softmax(attention_scores, axis=1) # 计算上下文向量 context_vector = tf.matmul(tf.transpose(attention_weights, perm=[0, 2, 1]), encoder_outputs) # shape: (batch_size, 1, 20) # 把上下文向量转成40维,和解码器当前输入拼接后送入LSTM2 context_proj = Dense(40)(context_vector) # shape: (batch_size, 1, 40) # 假设decoder_current_input是解码器当前时间步的输入,shape=(batch_size, 1, input_dim) decoder_input = tf.concat([context_proj, decoder_current_input], axis=-1)
方案2:把解码器隐藏层转成编码器输出维度
# 用线性层将解码器隐藏层从40维转成20维 decoder_proj = Dense(20)(decoder_hidden) # shape: (batch_size, 20) decoder_proj_expanded = tf.expand_dims(decoder_proj, axis=1) # shape: (batch_size, 1, 20) # 直接和编码器输出做点积计算得分 attention_scores = tf.matmul(encoder_outputs, decoder_proj_expanded, transpose_b=True) # shape: (batch_size, enc_seq_len, 1) attention_weights = tf.nn.softmax(attention_scores, axis=1) context_vector = tf.matmul(tf.transpose(attention_weights, perm=[0, 2, 1]), encoder_outputs) # shape: (batch_size, 1, 20) # 同样把上下文向量转成40维后拼接解码器输入 context_proj = Dense(40)(context_vector) decoder_input = tf.concat([context_proj, decoder_current_input], axis=-1)
方法二:用加性注意力(Bahdanau注意力)
加性注意力本身就支持query和key维度不同的场景,无需提前对齐维度,核心是通过前馈网络把两者映射到同一个中间维度后计算得分:
# 定义加性注意力的参数层,中间维度可根据任务调整(比如选64) W1 = Dense(64, activation='tanh') W2 = Dense(1) # 解码器隐藏层扩维并广播,匹配编码器输出的序列长度 decoder_hidden_expanded = tf.expand_dims(decoder_hidden, axis=1) # shape: (batch_size, 1, 40) decoder_hidden_broadcast = tf.tile(decoder_hidden_expanded, [1, tf.shape(encoder_outputs)[1], 1]) # shape: (batch_size, enc_seq_len, 40) # 拼接编码器输出和解码器隐藏层,过前馈网络得到注意力得分 concat_input = tf.concat([encoder_outputs, decoder_hidden_broadcast], axis=-1) # shape: (batch_size, enc_seq_len, 60) attention_scores = W2(W1(concat_input)) # shape: (batch_size, enc_seq_len, 1) attention_weights = tf.nn.softmax(attention_scores, axis=1) context_vector = tf.matmul(tf.transpose(attention_weights, perm=[0, 2, 1]), encoder_outputs) # shape: (batch_size, 1, 20) # 后续处理同上,转成40维后拼接解码器输入 context_proj = Dense(40)(context_vector) decoder_input = tf.concat([context_proj, decoder_current_input], axis=-1)
注意事项
- 所有线性层的参数都会随模型一起训练,不需要手动设置固定值
- 中间维度(比如加性注意力里的64)可以根据任务复杂度调整,选两个原始维度的中间值或经验值都可以
- 如果用PyTorch实现,逻辑完全一致,只需要替换成对应的API(比如用
nn.Linear代替Dense,torch.matmul代替tf.matmul)
内容的提问来源于stack exchange,提问作者Josef Souza
相关产品推荐
相关产品推荐

