TensorFlow LSTM模型使用ModelCheckpoint保存加载的警告与报错求解
问题解答
1、LSTM相关警告解决
两个警告均为序列化过程中的常见提示,本身不影响模型正常使用,具体说明和解决方案如下:
- 第一个
untraced functions警告:原因是SavedModel格式保存时,LSTM层的内部调用方法没有被完整追踪,不会影响模型加载后的推理、继续训练等操作。如果要消除警告,可以给ModelCheckpoint新增参数save_traces=True,或者切换保存格式为HDF5,新增参数save_format='h5'即可。 - 第二个
LSTMCell重名警告:你给外层LSTM层设置name不生效,是因为警告的是LSTM层内部自动生成的LSTMCell实例重名。要消除的话,加载模型时在custom_objects参数里指定内置LSTMCell即可:custom_objects={'LSTMCell': tf.keras.layers.LSTMCell}
2、TextVectorization层加载报错解决
你的思路是可行的,也有更简便的临时解决方案:
临时解决方案(无需重写层)
加载模型时直接传入自定义对象即可,代码如下:
model_2_loaded = tf.keras.models.load_model( os.path.join('model_experiments', model_2.name), custom_objects={ 'TextVectorization': tf.keras.layers.TextVectorization, 'custom_standardization': custom_standardization } )
永久序列化方案(自定义子类)
按照你的思路封装自定义层,后续加载无需额外传参数:
import re import string import tensorflow as tf @tf.keras.utils.register_keras_serializable() class CustomTextVectorization(tf.keras.layers.TextVectorization): def __init__(self, max_tokens, output_mode, output_sequence_length, **kwargs): super().__init__( max_tokens=max_tokens, output_mode=output_mode, output_sequence_length=output_sequence_length, **kwargs ) def standardize(self, input_data): lowercase = tf.strings.lower(input_data) stripped_html = tf.strings.regex_replace(lowercase, '<br/>', ' ') return tf.strings.regex_replace( stripped_html, '[%s]' % re.escape(string.punctuation), '' ) # 层初始化改为调用自定义类 vectorize_layer = CustomTextVectorization( max_tokens=max_features, output_mode='int', output_sequence_length=sequence_length )
3、adapt生成的词汇表保存问题
只要你将TextVectorization层直接嵌入模型结构中一起保存,adapt生成的词汇表会作为层的权重参数同步保存到模型文件中。只要正确解决了上述的自定义对象序列化问题,加载模型时词汇表会完整恢复,不需要重新执行adapt方法。
内容的提问来源于stack exchange,提问作者mrGott
相关产品推荐
相关产品推荐

