如何正确保存与加载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
相关产品推荐
相关产品推荐

