BLSTM+全局注意力文本摘要模型推理阶段LSTM维度不匹配问题求助
基于BLSTM+全局Attention的文本摘要模型推理错误修复方案
问题背景
构建了基于BLSTM架构和全局Attention的文本摘要模型,输入词汇量x_vocab_size=36782,目标词汇量y_vocab_size=19749。训练阶段代码可正常运行,但推理阶段出现维度不匹配错误。
训练模型代码
latent_dim = 300 # Encoder encoder_inputs = Input(shape=(None,), dtype='int32', name='input_text') enc_emb = Embedding(x_vocab_size, latent_dim, name='text_embedding', trainable=True)(encoder_inputs) # BLSTM Layer encoder_LSTM = LSTM(latent_dim, return_sequences=True, return_state=True) encoder_LSTM_R = LSTM(latent_dim, return_sequences=True, return_state=True, go_backwards=True) encoder_output, forward_h, forward_c = encoder_LSTM(enc_emb) encoder_outputr, backward_h, backward_c = encoder_LSTM_R(enc_emb) encoder_outputs = Concatenate()([encoder_output, encoder_outputr]) encoder_states = [forward_h, forward_c, backward_h, backward_c] # Decoder decoder_inputs = Input(shape=(None,), name='input_summary') dec_emb_layer = Embedding(y_vocab_size, latent_dim, name='summary_embedding', trainable=True) dec_emb = dec_emb_layer(decoder_inputs) # LSTM using encoder_states as initial state decoder_LSTM = LSTM(latent_dim, return_sequences=True, return_state=True) decoder_LSTM_R = LSTM(latent_dim, return_sequences=True, return_state=True, go_backwards=True) decoder_output, decforward_h, decforward_c = decoder_LSTM(dec_emb, initial_state=[forward_h, forward_c]) decoder_outputr, decbackward_h, decbackward_c = decoder_LSTM_R(dec_emb, initial_state=[backward_h, backward_c]) decoder_outputs = Concatenate()([decoder_output, decoder_outputr]) decoder_states = [decforward_h, decforward_c, decbackward_h, decbackward_c] # Attention Layer attn_out = tf.keras.layers.Attention()([encoder_outputs, decoder_outputs]) # Concat attention output and decoder BLSTM output decoder_concat_input = Concatenate(axis=-1, name='dec_concat_layer')([decoder_outputs, attn_out]) # Dense layer decoder_dense = TimeDistributed(Dense(y_vocab_size, activation='softmax')) decoder_outputs = decoder_dense(decoder_concat_input) # Model Definition model = Model([encoder_inputs, decoder_inputs], [decoder_outputs]) model.summary()
原推理模型代码
# Encoder Inference encoder_model = Model(inputs=encoder_inputs,outputs=encoder_states) # Decoder Inference # Below tensors hold the states of the previous time step decoder_state_input_h = Input(shape=(None,300)) decoder_state_input_c = Input(shape=(None,300)) decoder_states_inputs = [decoder_state_input_h, decoder_state_input_c] # Getting decoder sequence embeddings dec_emb2 = dec_emb_layer(decoder_inputs) # Predicting the next word in the sequence # Setting the initial states to the previous time step states decoder_outputs2, state_h2, state_c2 = decoder_LSTM(dec_emb2, initial_state=decoder_states_inputs) decoder_outputsb2, state_hb2, state_cb2 = decoder_LSTM(dec_emb2, initial_state=decoder_states_inputs) decoder_outputs3 = Concatenate()([decoder_outputs2, decoder_outputsb2]) # Attention Inference attn_out_inf = tf.keras.layers.Attention()([encoder_outputs, decoder_outputs3]) decoder_inf_concat = Concatenate(axis=-1, name='concat')([decoder_outputs3, attn_out_inf]) # Dense softmax layer to calculate probability distribution over target vocab decoder_outputs2 = decoder_dense(decoder_inf_concat) # Final Decoder model decoder_model = Model( [decoder_inputs]+[decoder_hidden_state_input], [decoder_outputs2])
运行错误信息
ValueError: Exception encountered when calling layer "lstm_2" (type LSTM). Dimensions must be equal, but are 1200 and 300 for '{{node mul}} = Mul[T=DT_FLOAT](Sigmoid_1, init_c)' with input shapes: [?,?,1200], [?,?,300]. Call arguments received by layer "lstm_2" (type LSTM): • inputs=['tf.Tensor(shape=(None, None, 300), dtype=float32)', 'tf.Tensor(shape=(None, None, 300), dtype=float32)', 'tf.Tensor(shape=(None, None, 300), dtype=float32)'] • mask=None • training=False • initial_state=None
错误原因分析
- 初始状态维度与数量不匹配:训练时解码器的双向LSTM分别使用编码器的正向(
forward_h,forward_c)和反向(backward_h,backward_c)共4个状态,但推理阶段仅定义了2个状态输入,且错误设置了shape=(None,300)(LSTM的状态张量维度应为(latent_dim,),即(300,),而非带序列维度的(None,300))。 - 双向LSTM初始状态复用错误:推理时给正向和反向LSTM传入了同一组初始状态,未对应到编码器输出的正向/反向状态。
- Attention输入未定义:推理阶段直接使用训练时的
encoder_outputs,但推理时需要单独传入编码器的输出张量,且原代码中decoder_hidden_state_input未定义,导致模型输入缺失。
修复方案与修正代码
修正后的推理模型代码
import tensorflow as tf from tensorflow.keras.layers import Input, Concatenate from tensorflow.keras.models import Model import numpy as np # Encoder Inference:输出编码器的所有状态和序列输出 encoder_model = Model(inputs=encoder_inputs, outputs=[encoder_outputs] + encoder_states) # Decoder Inference # 定义解码器需要的4个状态输入(对应编码器的forward_h, forward_c, backward_h, backward_c) decoder_state_input_fh = Input(shape=(latent_dim,)) decoder_state_input_fc = Input(shape=(latent_dim,)) decoder_state_input_bh = Input(shape=(latent_dim,)) decoder_state_input_bc = Input(shape=(latent_dim,)) decoder_states_inputs = [decoder_state_input_fh, decoder_state_input_fc, decoder_state_input_bh, decoder_state_input_bc] # 定义编码器输出的输入层(用于Attention计算) encoder_outputs_inf = Input(shape=(None, 2*latent_dim)) # 双向输出拼接后维度是2*latent_dim # 获取解码器词嵌入 dec_emb2 = dec_emb_layer(decoder_inputs) # 正向LSTM使用编码器的正向状态初始化 decoder_outputs_f, state_fh, state_fc = decoder_LSTM( dec_emb2, initial_state=[decoder_state_input_fh, decoder_state_input_fc] ) # 反向LSTM使用编码器的反向状态初始化 decoder_outputs_b, state_bh, state_bc = decoder_LSTM_R( dec_emb2, initial_state=[decoder_state_input_bh, decoder_state_input_bc] ) # 拼接双向LSTM输出 decoder_outputs3 = Concatenate(axis=-1)([decoder_outputs_f, decoder_outputs_b]) # Attention计算:使用推理阶段的编码器输出和解码器输出 attn_out_inf = tf.keras.layers.Attention()([encoder_outputs_inf, decoder_outputs3]) decoder_inf_concat = Concatenate(axis=-1)([decoder_outputs3, attn_out_inf]) # 输出词汇概率分布 decoder_outputs2 = decoder_dense(decoder_inf_concat) # 定义最终解码器模型:输入包括解码器输入、编码器输出、4个状态;输出包括预测结果和新的状态 decoder_model = Model( inputs=[decoder_inputs, encoder_outputs_inf] + decoder_states_inputs, outputs=[decoder_outputs2, state_fh, state_fc, state_bh, state_bc] )
生成预测摘要的示例函数
def generate_summary(input_seq, max_summary_len, reverse_target_word_index): # 用编码器模型获取初始状态和编码器输出 encoder_out, fh, fc, bh, bc = encoder_model.predict(input_seq, verbose=0) # 初始化解码器输入:起始符(假设起始符索引为0,需根据你的词汇表调整) target_seq = np.zeros((1, 1)) target_seq[0, 0] = 0 # 替换为你的起始符索引 summary = [] stop_condition = False while not stop_condition: # 解码器一步预测 output_tokens, fh, fc, bh, bc = decoder_model.predict( [target_seq, encoder_out, fh, fc, bh, bc], verbose=0 ) # 获取概率最大的词索引 sampled_token_index = np.argmax(output_tokens[0, -1, :]) sampled_word = reverse_target_word_index.get(sampled_token_index, '<unk>') summary.append(sampled_word) # 达到最大长度或遇到结束符时停止 if sampled_word == '<end>' or len(summary) >= max_summary_len: stop_condition = True # 更新解码器输入为当前预测的词 target_seq = np.zeros((1, 1)) target_seq[0, 0] = sampled_token_index return ' '.join(summary)
关键修复点说明
- 状态输入修正:解码器推理阶段定义了与编码器输出对应的4个状态输入,维度调整为
(latent_dim,),匹配LSTM状态的实际维度。 - 双向LSTM状态分配:正向LSTM使用编码器的正向隐藏状态和细胞状态,反向LSTM使用编码器的反向状态,保证与训练阶段逻辑一致。
- Attention输入补充:新增
encoder_outputs_inf输入层,用于推理时传入编码器的序列输出,解决训练与推理阶段的张量依赖问题。 - 解码器模型输出扩展:返回更新后的状态,用于循环预测下一个词。
内容的提问来源于stack exchange,提问作者Ronald Ferdinand
相关产品推荐
相关产品推荐

