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

如何正确保存与加载Keras Transformer翻译模型?

正确保存基于Transformer的Keras翻译模型及TextVectorization层方案

一、Transformer模型的正确保存方式

放弃.h5格式,改用Keras官方推荐的SavedModel格式,它对自定义层的兼容性更好,能完整保存模型架构、权重及相关配置:

保存代码

# 保存为SavedModel格式(文件夹形式)
transformer.save("models/transformer_savedmodel")

加载代码

只要你的Python环境中定义了PositionalEmbedding、TransformerEncoder、TransformerDecoder这几个自定义类,直接加载即可,无需手动指定custom_objects:

import tensorflow as tf
# 确保自定义层的类定义已导入
from your_module import PositionalEmbedding, TransformerEncoder, TransformerDecoder

new_model = tf.keras.models.load_model("models/transformer_savedmodel")
new_model.summary()

.h5格式属于旧格式,对自定义层的序列化支持有限,容易出现加载时的兼容性问题,优先用SavedModel可以避免很多麻烦。

二、TextVectorization层的正确保存方法

用pickle保存会受环境依赖影响,导致跨环境加载状态不一致,推荐以下两种可靠方案:

方案1:将TextVectorization整合到模型中一起保存

把文本向量化层作为模型的前置输入层,这样整个预处理流程和模型绑定,加载后直接可以处理原始文本,无需单独管理向量化层:

# 假设已定义好训练完成的text_vectorization层和transformer核心
input_text = tf.keras.Input(shape=(1,), dtype=tf.string)
# 文本先经过向量化层
x = text_vectorization(input_text)
# 再进入Transformer模型
x = transformer(x)
# 后续输出层定义...
full_translation_model = tf.keras.Model(input_text, model_output)

# 保存完整模型
full_translation_model.save("models/full_translation_model")

# 加载后直接使用原始文本输入
loaded_model = tf.keras.models.load_model("models/full_translation_model")
# 示例:输入英文句子直接得到翻译结果
result = loaded_model.predict(["Hello world!"])

方案2:单独保存层的配置和权重(适合分离预处理和模型的场景)

用Keras官方的get_config和from_config接口保存,避免pickle的环境依赖问题:

保存代码

import json
import numpy as np

# 保存向量化层的配置
config = text_vectorization.get_config()
with open("models/text_vectorization_config.json", "w") as f:
    json.dump(config, f)

# 保存向量化层的权重(包含词汇表等状态)
weights = text_vectorization.get_weights()
np.save("models/text_vectorization_weights.npy", weights)

加载代码

import json
import numpy as np
from tensorflow.keras.layers import TextVectorization

# 加载配置并重建层
with open("models/text_vectorization_config.json", "r") as f:
    config = json.load(f)
text_vectorization = TextVectorization.from_config(config)

# 加载权重(注意allow_pickle=True,因为词汇表是数组结构)
text_vectorization.set_weights(np.load("models/text_vectorization_weights.npy", allow_pickle=True))

三、解决跨环境加载效果变差的问题

你遇到的Notebook外加载效果差的问题,本质是pickle保存的向量化层状态不完整(比如词汇表未正确加载、自定义标准化逻辑丢失),改用上面的方案1或方案2,能确保加载后的向量化层和训练时的状态完全一致,从而保证模型输入的正确性,解决效果下降的问题。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.24 21:30:13