在Keras餐饮评论文本摘要模型中添加Attention层遇问题求助
解决Seq2Seq文本摘要模型中添加Attention层的问题
先修正教程里的语法错误
你参考的教程里这段代码有明显语法问题:
Attention layer attn_layer = AttentionLayer(name='attention_layer')
正确写法要去掉多余的Attention layer,改成:
attn_layer = AttentionLayer(name='attention_layer')
方案1:用TensorFlow Keras原生Attention层(最省心)
TensorFlow 2.x的Keras已经内置了Attention层,不需要额外装库,直接用就行。步骤如下:
- 导入原生Attention层:
from tensorflow.keras.layers import Attention
- 调整你的模型代码,替换Attention相关部分:
注意原生Attention的输入顺序是[query, value],对应你的解码器输出和编码器输出,修改后代码片段:
# 解码器LSTM输出后添加Attention attn_layer = Attention(name='attention_layer') # query是decoder_outputs,value是encoder_outputs attn_out = attn_layer([decoder_outputs, encoder_outputs]) # 拼接解码器输出和Attention输出 decoder_concat_input = Concatenate(axis=-1, name='concat_layer')([decoder_outputs, attn_out]) # 后续保持不变,注意这里要把decoder_dense的输入换成拼接后的decoder_concat_input decoder_outputs = decoder_dense(decoder_concat_input)
方案2:自定义实现AttentionLayer(完全匹配教程逻辑)
如果一定要用教程里的AttentionLayer接口,可以自己实现这个类,代码如下:
from tensorflow.keras.layers import Layer from tensorflow.keras import backend as K class AttentionLayer(Layer): """ 自定义Attention层,适配Seq2Seq模型 """ def __init__(self, **kwargs): super(AttentionLayer, self).__init__(**kwargs) def build(self, input_shape): # 定义可训练参数 self.W = self.add_weight(name='attention_weight', shape=(input_shape[0][2], input_shape[0][2]), initializer='glorot_uniform', trainable=True) super(AttentionLayer, self).build(input_shape) def call(self, inputs): # inputs = [encoder_outputs, decoder_outputs] encoder_out, decoder_out = inputs # 计算注意力得分 score = K.batch_dot(decoder_out, K.dot(encoder_out, self.W), axes=[2, 2]) attn_weights = K.softmax(score, axis=1) # 计算上下文向量 context = K.batch_dot(attn_weights, encoder_out, axes=[1, 1]) return [context, attn_weights] def compute_output_shape(self, input_shape): return [(input_shape[1][0], input_shape[1][1], input_shape[0][2]), (input_shape[1][0], input_shape[1][1], input_shape[0][1])]
然后直接像教程里那样使用即可:
attn_layer = AttentionLayer(name='attention_layer') attn_out, attn_states = attn_layer([encoder_outputs, decoder_outputs]) decoder_concat_input = Concatenate(axis=-1, name='concat_layer')([decoder_outputs, attn_out])
方案3:正确安装attention_keras库
你之前用pip install attention_keras失败,是因为这个库没有上传到PyPI,需要从GitHub源码安装:
pip install git+https://github.com/thushv89/attention_keras.git
安装完成后,导入对应的AttentionLayer:
from attention_keras.layers import AttentionLayer
之后就可以按照教程代码使用了。
内容的提问来源于stack exchange,提问作者Rodolfo
相关产品推荐
相关产品推荐

