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

Keras LSTM Seq2Seq转GRU后推理阶段维度不匹配错误排查

问题:Keras LSTM Seq2Seq转GRU后推理阶段报错

错误信息

Cell In[19], line 30, in decode_sequence(input_seq)
     28 decoded_sentence = ""
     29 while not stop_condition:
---> 30     output_tokens, h = decoder_model.predict([target_seq] + states_value, verbose=0)
     32     # Sample a token
     33     sampled_token_index = np.argmax(output_tokens[0, -1, :])

ValueError: operands could not be broadcast together with shapes (1,1,1,84) (1,2048) 

训练阶段代码(LSTM已注释,GRU为当前实现)

# Define an input sequence and process it.
encoder_inputs = keras.Input(shape=(None, num_encoder_tokens))

# LSTM
###encoder = keras.layers.LSTM(latent_dim, return_state=True)
###encoder_outputs, state_h, state_c = encoder(encoder_inputs)

# We discard `encoder_outputs` and only keep the states.
###encoder_states = [state_h, state_c]

# GRU
encoder = keras.layers.GRU(latent_dim, return_state=True)
outputs = encoder(encoder_inputs)
encoder_output, encoder_states = outputs[0], outputs[1:]

# Set up the decoder, using `encoder_states` as initial state.
decoder_inputs = keras.Input(shape=(None, num_decoder_tokens))

# We set up our decoder to return full output sequences,
# and to return internal states as well. We don't use the
# return states in the training model, but we will use them in inference.

# LSTM
###decoder_lstm = keras.layers.LSTM(latent_dim, return_sequences=True, return_state=True)
###decoder_outputs, _, _ = decoder_lstm(decoder_inputs, initial_state=encoder_states)

# GRU
decoder = keras.layers.GRU(latent_dim, return_sequences=True, return_state=True)
outputs = decoder(decoder_inputs, initial_state=tuple(encoder_states))
decoder_outputs, decoder_state = outputs[0], outputs[1:]

decoder_dense = keras.layers.Dense(num_decoder_tokens, activation="softmax")
decoder_outputs = decoder_dense(decoder_outputs)

# Define the model that will turn
# `encoder_input_data` & `decoder_input_data` into `decoder_target_data`
model = keras.Model([encoder_inputs, decoder_inputs], decoder_outputs)

LSTM推理阶段参考代码

### LSTM
# Define sampling models
# Restore the model and construct the encoder and decoder.
encoder_model = keras.Model(encoder_inputs, encoder_states)
decoder_state_input_h = keras.Input(shape=(latent_dim,))
decoder_state_input_c = keras.Input(shape=(latent_dim,))
decoder_states_inputs = [decoder_state_input_h, decoder_state_input_c]
decoder_outputs, state_h, state_c = decoder_lstm(
    decoder_inputs, initial_state=decoder_states_inputs)
decoder_states = [state_h, state_c]
decoder_outputs = decoder_dense(decoder_outputs)
decoder_model = keras.Model(
    [decoder_inputs] + decoder_states_inputs,
    [decoder_outputs] + decoder_states)

reverse_input_char_index = dict((i, char) for char, i in input_token_index.items())
reverse_target_char_index = dict((i, char) for char, i in target_token_index.items())

def decode_sequence(input_seq):
    # Encode the input as state vectors.
    states_value = encoder_model.predict(input_seq, verbose=0)

    # Generate empty target sequence of length 1.
    target_seq = np.zeros((1, 1, num_decoder_tokens))
    # Populate the first character of target sequence with the start character.
    target_seq[0, 0, target_token_index["\t"]] = 1.0

    # Sampling loop for a batch of sequences
    # (to simplify, here we assume a batch of size 1).
    stop_condition = False
    decoded_sentence = ""
    while not stop_condition:
        output_tokens, h, c = decoder_model.predict([target_seq] + states_value, verbose=0)

        # Sample a token
        sampled_token_index = np.argmax(output_tokens[0, -1, :])
        sampled_char = reverse_target_char_index[sampled_token_index]
        decoded_sentence += sampled_char

        # Exit condition: either hit max length
        # or find stop character.
        if sampled_char == "\n" or len(decoded_sentence) > max_decoder_seq_length:
            stop_condition = True

        # Update the target sequence (of length 1).
        target_seq = np.zeros((1, 1, num_decoder_tokens))
        target_seq[0, 0, sampled_token_index] = 1.0

        # Update states
        states_value = [h, c]
    return decoded_sentence

当前GRU推理阶段报错代码

### GRU
# Define sampling models
# Restore the model and construct the encoder and decoder.
encoder_model = keras.Model(encoder_inputs, encoder_states)
decoder_states_inputs = keras.Input(shape=(latent_dim,))
decoder_outputs, decoder_states = decoder(
    decoder_inputs, initial_state=decoder_states_inputs)
decoder_outputs = decoder_dense(decoder_outputs)

decoder_model = keras.Model([decoder_outputs] + [decoder_states])

reverse_input_char_index = dict((i, char) for char, i in input_token_index.items())
reverse_target_char_index = dict((i, char) for char, i in target_token_index.items())


def decode_sequence(input_seq):
    # Encode the input as state vectors.
    states_value = encoder_model.predict(input_seq, verbose=0)

    # Generate empty target sequence of length 1.
    target_seq = np.zeros((1, 1, num_decoder_tokens))
    # Populate the first character of target sequence with the start character.
    target_seq[0, 0, target_token_index["\t"]] = 1.0

    # Sampling loop for a batch of sequences
    # (to simplify, here we assume a batch of size 1).
    stop_condition = False
    decoded_sentence = ""
    while not stop_condition:
        output_tokens, h = decoder_model.predict([target_seq] + states_value, verbose=0) # Error is thrown here

        # Sample a token
        sampled_token_index = np.argmax(output_tokens[0, -1, :])
        sampled_char = reverse_target_char_index[sampled_token_index]
        decoded_sentence += sampled_char

        # Exit condition: either hit max length
        # or find stop character.
        if sampled_char == "\n" or len(decoded_sentence) > max_decoder_seq_length:
            stop_condition = True

        # Update the target sequence (of length 1).
        target_seq = np.zeros((1, 1, num_decoder_tokens))
        target_seq[0, 0, sampled_token_index] = 1.0

        # Update states
        states_value = [h]
    return decoded_sentence

解决方案

问题核心是GRU推理阶段的decoder_model输入输出定义完全搞反,把输出张量当成了输入。正确的模型结构应该和LSTM版本对应,明确输入是解码器序列+状态,输出是预测结果+更新后的状态。

修正后的GRU推理代码:

### GRU
# Define sampling models
encoder_model = keras.Model(encoder_inputs, encoder_states)

# 定义解码器的状态输入
decoder_state_input = keras.Input(shape=(latent_dim,))
# 用解码器输入和状态输入得到输出和新状态
decoder_outputs, decoder_state = decoder(
    decoder_inputs, initial_state=decoder_state_input)
decoder_outputs = decoder_dense(decoder_outputs)

# 正确定义解码器模型:输入是[decoder_inputs, 状态输入],输出是[decoder_outputs, 新状态]
decoder_model = keras.Model(
    [decoder_inputs, decoder_state_input],
    [decoder_outputs, decoder_state]
)

reverse_input_char_index = dict((i, char) for char, i in input_token_index.items())
reverse_target_char_index = dict((i, char) for char, i in target_token_index.items())


def decode_sequence(input_seq):
    states_value = encoder_model.predict(input_seq, verbose=0)

    target_seq = np.zeros((1, 1, num_decoder_tokens))
    target_seq[0, 0, target_token_index["\t"]] = 1.0

    stop_condition = False
    decoded_sentence = ""
    while not stop_condition:
        # 输入为[target_seq] + states_value,对应模型的两个输入
        output_tokens, h = decoder_model.predict([target_seq] + states_value, verbose=0)

        sampled_token_index = np.argmax(output_tokens[0, -1, :])
        sampled_char = reverse_target_char_index[sampled_token_index]
        decoded_sentence += sampled_char

        if sampled_char == "\n" or len(decoded_sentence) > max_decoder_seq_length:
            stop_condition = True

        target_seq = np.zeros((1, 1, num_decoder_tokens))
        target_seq[0, 0, sampled_token_index] = 1.0

        states_value = [h]
    return decoded_sentence

关键修正点

  1. 修正decoder_model的输入输出:原代码错误地将输出张量作为输入构建模型,正确逻辑是输入为decoder_inputs(解码器序列)和decoder_state_input(GRU的隐藏状态),输出为预测结果decoder_outputs和更新后的状态decoder_state。
  2. GRU仅维护一个隐藏状态,无需像LSTM那样处理两个状态张量,因此只需要定义一个状态输入层。

内容的提问来源于stack exchange,提问作者Duncan Leung

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.12 21:00:54