移除解码器注意力层后,如何连接编码器输出与解码器输入?
移除注意力层后的Encoder-Decoder连接方案
要移除注意力层,核心是用**全局池化(均值/最大值)**替代注意力生成的context_vector,将编码器输出的序列特征压缩为固定维度的向量,再与解码器输入的embedding拼接,保持原有解码逻辑的兼容性。
修改后的Decoder代码
class Decoder(Model): def __init__(self, embed_dim, units, vocab_size): super(Decoder, self).__init__() self.units = units self.embed = tf.keras.layers.Embedding(vocab_size, embed_dim) # 保留Embedding层 self.gru = tf.keras.layers.GRU(self.units, return_sequences=True, return_state=True, recurrent_initializer='glorot_uniform') self.d1 = tf.keras.layers.Dense(self.units) self.d2 = tf.keras.layers.Dense(vocab_size) def call(self, x, features, hidden): # 用全局均值池化替代注意力层生成context vector # 编码器输出features形状:(batch, 64, embed_dim),池化后变为(batch, embed_dim) context_vector = tf.reduce_mean(features, axis=1) embed = self.embed(x) # 输入embedding形状:(batch_size, 1, embed_dim) # 拼接context vector与输入embedding,形状变为(batch_size, 1, embed_dim * 2) embed = tf.concat([tf.expand_dims(context_vector, 1), embed], axis=-1) output, state = self.gru(embed) # GRU输出形状:(batch_size, 1, units) output = self.d1(output) output = tf.reshape(output, (-1, output.shape[2])) # 调整形状适配全连接层 output = self.d2(output) return output, state def init_state(self, batch_size): return tf.zeros((batch_size, self.units))
关键改动说明
- 移除注意力冗余代码:删掉Decoder中未使用的
self.W1、注释掉的Attention初始化,以及call方法里的注意力调用逻辑。 - 替换context_vector生成方式:用
tf.reduce_mean(features, axis=1)对编码器输出的64个特征向量做全局均值池化,得到与原注意力输出维度一致的(batch, embed_dim)向量;如果想保留更多特征细节,也可以替换为tf.reduce_max做最大值池化。 - 保留原有解码流程:后续的embedding拼接、GRU前向传播、全连接层输出逻辑完全保留,无需额外调整,确保模型解码流程的连贯性。
内容的提问来源于stack exchange,提问作者user3541974
相关产品推荐
相关产品推荐

