基于LSTM的机器翻译注意力机制调用K.rnn触发TypeError如何解决
错误根因
你遇到的报错是因为旧版Keras后端接口K.rnn属于TensorFlow不支持符号张量调度的API,在Functional模型构建的静态图阶段直接调用会触发类型不匹配。官方提示的「将操作放入自定义Keras层的call方法」指的是不要直接调用底层TF API,要使用Keras封装的支持符号张量的层接口。
修复方案
提供两种可直接落地的修复方式,按需选择即可:
方案1:最小代码改动,替换为Keras内置注意力层
你实现的是Bahdanau加性注意力,Keras已经提供了原生支持的AdditiveAttention层,直接替换自定义的AttentionLayer即可,完全避开自己实现的K.rnn问题:
# 把原来的#Attention Layer部分替换为以下代码 from tensorflow.keras.layers import AdditiveAttention attention_layer = AdditiveAttention(name='attention_layer') # 输入顺序保持和原有逻辑一致,return_attention_scores=True会同时返回注意力权重和上下文向量 attention_result, attention_weights = attention_layer([decoder_outputs, encoder_outputs1], return_attention_scores=True)
替换后其余代码完全不需要改动,可直接运行,且原生层的训练效率远高于自定义的循环实现。
方案2:修改自定义Attention层,移除K.rnn调用
如果你需要保留自定义注意力的逻辑,把原来的循环实现改为向量化实现,彻底去掉K.rnn即可,修改后的Attention层call方法代码如下:
def call(self, inputs, verbose=False): assert type(inputs) == list encoder_out_seq, decoder_out_seq = inputs # encoder_out_seq shape: (batch, enc_len, enc_hidden) # decoder_out_seq shape: (batch, dec_len, dec_hidden) # 向量化计算能量分数,完全替代原来的K.rnn循环逻辑 # 计算W_a * encoder_out_seq W_a_dot_s = K.dot(encoder_out_seq, self.W_a) # (batch, enc_len, hidden) # 计算U_a * decoder_out_seq U_a_dot_h = K.dot(decoder_out_seq, self.U_a) # (batch, dec_len, hidden) # 广播相加后过tanh Ws_plus_Uh = K.tanh(K.expand_dims(W_a_dot_s, 1) + K.expand_dims(U_a_dot_h, 2)) # (batch, dec_len, enc_len, hidden) # 计算V_a乘以结果得到能量分数 e = K.dot(Ws_plus_Uh, self.V_a) # (batch, dec_len, enc_len, 1) e = K.squeeze(e, axis=-1) # (batch, dec_len, enc_len) # softmax得到注意力权重 attention_weights = K.softmax(e) # 计算上下文向量 attention_result = K.batch_dot(attention_weights, encoder_out_seq) # (batch, dec_len, enc_hidden) return attention_result, attention_weights
该实现和你原来的循环逻辑完全等价,但没有使用任何底层不支持符号张量的API,可直接在Functional模型中使用。
内容的提问来源于stack exchange,提问作者rob
相关产品推荐
相关产品推荐

