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

使用KerasNLP构建Transformer翻译模型存盘告警及预测异常问题

解决KerasNLP Transformer模型用pickle保存后无法正常预测的问题

问题场景

我用KerasNLP构建了机器翻译的Transformer模型,训练完成后尝试用pickle保存模型,出现告警且加载后无法正常预测。

模型构建代码

# Encoder
encoder_inputs = keras.Input(shape=(None,), dtype="int64", name="encoder_inputs")

x = keras_nlp.layers.TokenAndPositionEmbedding(
    vocabulary_size=ENG_VOCAB_SIZE,
    sequence_length=MAX_SEQUENCE_LENGTH,
    embedding_dim=EMBED_DIM,
    mask_zero=True,
)(encoder_inputs)

encoder_outputs = keras_nlp.layers.TransformerEncoder(
    intermediate_dim=INTERMEDIATE_DIM, num_heads=NUM_HEADS
)(inputs=x)
encoder = keras.Model(encoder_inputs, encoder_outputs)


# Decoder
decoder_inputs = keras.Input(shape=(None,), dtype="int64", name="decoder_inputs")
encoded_seq_inputs = keras.Input(shape=(None, EMBED_DIM), name="decoder_state_inputs")

x = keras_nlp.layers.TokenAndPositionEmbedding(
    vocabulary_size=SND_VOCAB_SIZE,
    sequence_length=MAX_SEQUENCE_LENGTH,
    embedding_dim=EMBED_DIM,
    mask_zero=True,
)(decoder_inputs)

x = keras_nlp.layers.TransformerDecoder(
    intermediate_dim=INTERMEDIATE_DIM, num_heads=NUM_HEADS
)(decoder_sequence=x, encoder_sequence=encoded_seq_inputs)
x = keras.layers.Dropout(0.5)(x)
decoder_outputs = keras.layers.Dense(SND_VOCAB_SIZE, activation="softmax")(x)
decoder = keras.Model(
    [
        decoder_inputs,
        encoded_seq_inputs,
    ],
    decoder_outputs,
)
decoder_outputs = decoder([decoder_inputs, encoder_outputs])

transformer = keras.Model(
    [encoder_inputs, decoder_inputs],
    decoder_outputs,
    name="transformer",
)

模型训练代码

transformer.summary()
transformer.compile(
    "rmsprop", loss="sparse_categorical_crossentropy", metrics=["accuracy"]
)
hist = transformer.fit(train_ds, epochs=EPOCHS, validation_data=val_ds)

保存代码及告警信息

保存代码:

with open('model.pkl', 'wb') as file:
  pickle.dump(transformer, file)

告警信息:

WARNING:absl:Found untraced functions such as token_embedding1_layer_call_fn, token_embedding1_layer_call_and_return_conditional_losses, position_embedding1_layer_call_fn, position_embedding1_layer_call_and_return_conditional_losses, multi_head_attention_layer_call_fn while saving (showing 5 of 78). These functions will not be directly callable after loading.

问题原因

pickle并非Keras官方推荐的模型序列化方式,对于包含KerasNLP高层API层(如TokenAndPositionEmbedding、TransformerEncoder等)的模型,pickle无法完整追踪并序列化模型的所有内部函数,导致加载后模型核心功能缺失,无法正常预测。

解决方案

使用Keras官方提供的model.save()方法保存模型,加载时用keras.models.load_model(),该方式会完整保存模型的结构、权重、配置以及所有依赖层的信息,完全适配Keras及KerasNLP模型。

修改后的保存代码:

# 保存模型
transformer.save("transformer_translation_model")

加载模型代码:

# 加载模型
from keras.models import load_model

loaded_model = load_model("transformer_translation_model")
# 验证加载后的模型可正常预测
loaded_model.predict(...)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.09 04:06:13