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

基于TensorFlow的NMT模型训练与推理报错排查求助

神经机器翻译模型训练及推理报错解决方案

我的代码

vocab_size = 10000
total_sentences = 25000
maxlen = 10
epochs = 50
validation_split = 0.05
split = int(0.95 * total_sentences)

X_train = [encoder_inputs[:split], decoder_inputs[:split]]
y_train = decoder_outputs[:split]

# 用于使用BLEU分数评估NMT模型的测试数据
X_test = en_data[:split]
y_test = hi_data[:split]

print(X_train[0].shape, X_train[1].shape, y_train.shape)

核心模型代码(报错相关)

d_model = 256

inputs = tf.keras.layers.Input(shape=(None,))
x = tf.keras.layers.Embedding(english_vocab_size, d_model, mask_zero=True)(inputs)
_,state_h,state_c = tf.keras.layers.LSTM(d_model,activation='relu',return_state=True)(x)

targets = tf.keras.layers.Input(shape=(None,))
embedding_layer = tf.keras.layers.Embedding(hindi_vocab_size, d_model, mask_zero=True)
x = embedding_layer(targets)
decoder_lstm = tf.keras.layers.LSTM(d_model,activation='relu',return_sequences=True, return_state=True)
x,_,_ = decoder_lstm(x, initial_state=[state_h, state_c])
dense1 = tf.keras.layers.Dense(hindi_vocab_size, activation='softmax')
x = dense1(x)

model = tf.keras.models.Model(inputs=[inputs, targets],outputs=x)
model.summary()

model.compile(optimizer='rmsprop', loss='categorical_crossentropy', metrics=['accuracy'])

模型结构

Model: "model_3"
__________________________________________________________________________________________________
 Layer (type)                   Output Shape         Param #     Connected to                     
==================================================================================================
 input_7 (InputLayer)           [(None, None)]       0           []                               
                                                                                                    
 input_8 (InputLayer)           [(None, None)]       0           []                               
                                                                                                    
 embedding_6 (Embedding)        (None, None, 256)    2053120     ['input_7[0][0]']                 
                                                                                                    
 embedding_7 (Embedding)        (None, None, 256)    2405120     ['input_8[0][0]']                 
                                                                                                    
 lstm_6 (LSTM)                  [(None, 256),        525312      ['embedding_6[0][0]']             
                                 (None, 256),                                                      
                                 (None, 256)]                                                      
                                                                                                    
 lstm_7 (LSTM)                  [(None, None, 256),  525312      ['embedding_7[0][0]',             
                                 (None, 256),                     'lstm_6[0][1]',                 
                                 (None, 256)]                     'lstm_6[0][2]']                 
                                                                                                    
 dense_3 (Dense)                (None, None, 9395)   2414515     ['lstm_7[0][0]']                 
                                                                                                    
==================================================================================================
Total params: 7,923,379
Trainable params: 7,923,379
Non-trainable params: 0
__________________________________________________________________________________________________

训练时的报错

执行训练代码:

model.fit(X_train, y_train, epochs=epochs, validation_split=validation_split, callbacks=[save_model_callback, tf.keras.callbacks.TerminateOnNaN()])

报错信息:

---------------------------------------------------------------------------
ValueError                                Traceback (most recent call last)
<ipython-input-25-b2b3107dbda8> in <module>
----> 1 model.fit(X_train, y_train, epochs=epochs, validation_split=validation_split, callbacks=[save_model_callback, tf.keras.callbacks.TerminateOnNaN()])

1 frames
/usr/local/lib/python3.7/dist-packages/tensorflow/python/framework/func_graph.py in autograph_handler(*args, **kwargs)
   1145           except Exception as e:  # pylint:disable=broad-except
   1146             if hasattr(e, "ag_error_metadata"):
-> 1147               raise e.ag_error_metadata.to_exception(e)
   1148             else:
   1149               raise

ValueError: Shapes (None, 10) and (None, 10, 9395) are incompatible

我的尝试及对应报错

尝试1:修改损失函数与LSTM参数

d_model = 256

inputs = tf.keras.layers.Input(shape=(None,))
x = tf.keras.layers.Embedding(english_vocab_size, d_model, mask_zero=True)(inputs)
_,state_h,state_c = tf.keras.layers.LSTM(d_model,activation='relu',return_state=True)(x)

targets = tf.keras.layers.Input(shape=(None,))
embedding_layer = tf.keras.layers.Embedding(hindi_vocab_size, d_model, mask_zero=True)
x = embedding_layer(targets)
decoder_lstm = tf.keras.layers.LSTM(d_model,activation='relu',return_sequences=False, return_state=True)
x,_,_ = decoder_lstm(x, initial_state=[state_h, state_c])
dense1 = tf.keras.layers.Dense(hindi_vocab_size, activation='softmax')
x = dense1(x)

model = tf.keras.models.Model(inputs=[inputs, targets],outputs=x)
model.summary()

loss = tf.keras.losses.SparseCategoricalCrossentropy()
model.compile(optimizer='rmsprop', loss=loss, metrics=['accuracy'])

报错信息:

InvalidArgumentError: {{function_node __inference_train_function_11558}} logits and labels must have the same first dimension, got logits shape [32,9395] and labels shape [320]
     [[{{node sparse_categorical_crossentropy/SparseSoftmaxCrossEntropyWithLogits/SparseSoftmaxCrossEntropyWithLogits}}]]

尝试2:恢复return_sequences=True

训练报错消失,但推理时出现新报错:

Input 0 of layer "lstm_1" is incompatible with the layer: expected ndim=3, found ndim=2. Full shape received: (None, 256)

问题解决方案

1. 解决训练时形状不兼容问题

原因

categorical_crossentropy要求标签为独热编码的三维数组(样本数, maxlen, 词汇表大小),但你的y_train是二维整数索引数组(样本数, maxlen),维度不匹配。

解决代码

保持return_sequences=True(解码器需要输出每个时间步的预测),更换为SparseCategoricalCrossentropy损失函数(无需独热编码,直接处理整数标签):

d_model = 256

# 编码器
encoder_inputs = tf.keras.layers.Input(shape=(None,))
enc_emb = tf.keras.layers.Embedding(english_vocab_size, d_model, mask_zero=True)(encoder_inputs)
_, state_h, state_c = tf.keras.layers.LSTM(d_model, activation='relu', return_state=True)(enc_emb)
encoder_states = [state_h, state_c]

# 解码器
decoder_inputs = tf.keras.layers.Input(shape=(None,))
dec_emb_layer = tf.keras.layers.Embedding(hindi_vocab_size, d_model, mask_zero=True)
dec_emb = dec_emb_layer(decoder_inputs)
decoder_lstm = tf.keras.layers.LSTM(d_model, activation='relu', return_sequences=True, return_state=True)
decoder_outputs, _, _ = decoder_lstm(dec_emb, initial_state=encoder_states)
decoder_dense = tf.keras.layers.Dense(hindi_vocab_size, activation='softmax')
decoder_outputs = decoder_dense(decoder_outputs)

# 训练模型
model = tf.keras.models.Model([encoder_inputs, decoder_inputs], decoder_outputs)
model.summary()

# 编译配置:用稀疏分类交叉熵处理整数标签
model.compile(optimizer='rmsprop', 
              loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=False), 
              metrics=['accuracy'])

2. 解决推理时LSTM输入维度不兼容问题

原因

训练模型的解码器接收完整序列输入+初始状态,但推理时需要逐词生成,不能一次性输入整个序列,必须构建单独的推理模型。

解决代码

构建编码器和解码器的推理模型:

# 推理编码器:输入英文序列,输出编码器状态
encoder_model = tf.keras.models.Model(encoder_inputs, encoder_states)

# 推理解码器:输入当前目标词+上一步状态,输出预测词+新状态
decoder_state_input_h = tf.keras.layers.Input(shape=(d_model,))
decoder_state_input_c = tf.keras.layers.Input(shape=(d_model,))
decoder_states_inputs = [decoder_state_input_h, decoder_state_input_c]

# 解码器嵌入层复用训练时的层
dec_emb_inf = dec_emb_layer(decoder_inputs)

# 解码器LSTM:接收嵌入输入和初始状态
decoder_outputs_inf, state_h_inf, state_c_inf = decoder_lstm(dec_emb_inf, initial_state=decoder_states_inputs)
decoder_states_inf = [state_h_inf, state_c_inf]

# 输出层复用训练时的层
decoder_outputs_inf = decoder_dense(decoder_outputs_inf)

# 推理解码器模型
decoder_model = tf.keras.models.Model(
    [decoder_inputs] + decoder_states_inputs,
    [decoder_outputs_inf] + decoder_states_inf)

推理函数示例

def decode_sequence(input_seq, start_token_index, reverse_hindi_token_index):
    # 获取编码器状态
    states_value = encoder_model.predict(input_seq, verbose=0)
    
    # 初始化解码器输入为<start> token
    target_seq = np.zeros((1, 1))
    target_seq[0, 0] = start_token_index
    
    decoded_sentence = ''
    while True:
        output_tokens, h, c = decoder_model.predict(
            [target_seq] + states_value, verbose=0)
        
        # 选取概率最高的词
        sampled_token_index = np.argmax(output_tokens[0, -1, :])
        sampled_char = reverse_hindi_token_index[sampled_token_index]
        decoded_sentence += ' ' + sampled_char
        
        # 终止条件:遇到<end>或达到最大长度
        if sampled_char == '<end>' or len(decoded_sentence) > maxlen:
            break
        
        # 更新目标序列和状态
        target_seq = np.zeros((1, 1))
        target_seq[0, 0] = sampled_token_index
        states_value = [h, c]
    
    return decoded_sentence.strip()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.20 06:54:22