You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

错误原因分析

  1. 初始状态维度与数量不匹配:训练时解码器的双向LSTM分别使用编码器的正向(forward_h, forward_c)和反向(backward_h, backward_c)共4个状态,但推理阶段仅定义了2个状态输入,且错误设置了shape=(None,300)(LSTM的状态张量维度应为(latent_dim,),即(300,),而非带序列维度的(None,300))。
  2. 双向LSTM初始状态复用错误:推理时给正向和反向LSTM传入了同一组初始状态,未对应到编码器输出的正向/反向状态。
  3. 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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.25 07:57:15