如何正确保存加载含BatchNormalization与自定义AutoSRELU层的TensorFlow模型
问题原因
- 自定义层
AutoSRELU、MinMaxConstraint没有实现完整的序列化逻辑,当前的自定义层没有重写get_config方法,且MinMaxConstraint的__call__方法缺少返回值,导致h5格式保存时无法正确存储自定义层的训练后参数,加载时默认重新初始化参数,模型推理逻辑完全错误。 - 你在模型中复用了同一个
AutoSRELU实例两次,h5格式对共享自定义层的参数序列化支持不完善,容易出现参数丢失问题;同时h5格式对BatchNormalization层的移动均值、移动方差等推理状态存储不稳定,经常出现状态丢失的问题,导致BN层推理时使用错误的归一化参数。
解决方案
方案1(推荐):换用TensorFlow原生SavedModel格式保存
SavedModel是TensorFlow官方原生的序列化格式,对自定义层、共享参数、层状态的支持远优于老旧的h5格式,不需要修改过多代码即可解决问题:
- 修复自定义层的基础逻辑,补全
MinMaxConstraint的返回值、AutoSRELU的get_config方法,修改后代码如下:
# 修复MinMaxConstraint的__call__方法 class MinMaxConstraint(keras.constraints.Constraint): def __init__(self, minval, maxval): self.minval = tf.constant(minval ,dtype='float32') self.maxval = tf.constant(maxval ,dtype='float32') def __call__(self, w): # 补全return返回约束后的权重 return tf.cond(tf.greater(self.minval,w) , lambda: w + (self.minval - w) , lambda: tf.cond(tf.greater(w,self.maxval) , lambda: w - (w - self.maxval) , lambda: w)) def get_config(self): return {'Lower Bound': self.minval.numpy(), 'Upper Bound':self.maxval.numpy()} # 修复AutoSRELU的序列化逻辑 class AutoSRELU(keras.layers.Layer): def __init__(self, trainable = True, **kwargs): # 把kwargs传给父类完成基础初始化 super(AutoSRELU, self).__init__(**kwargs) self.k1 = self.add_weight(name='k', shape = (), initializer=initializer0, trainable=trainable) self.k2 = self.add_weight(name='n', shape = (), initializer=initializer1, trainable=trainable) def call(self, inputs): return srelu(inputs, self.k1, self.k2) # 补全get_config方法用于序列化 def get_config(self): config = super().get_config() config.update({"trainable": self.trainable}) return config
- 保存模型时不要用
.h5后缀,直接保存为文件夹格式:
model_scratch_auto.save('test_model_savedmodel')
- 加载模型时把所有自定义类都注册到
custom_objects中:
dependencies = { 'f1_m': f1_m, 'precision_m': precision_m, 'recall_m': recall_m, 'AutoSRELU': AutoSRELU, 'MinMaxConstraint': MinMaxConstraint } test_model = models.load_model('test_model_savedmodel', custom_objects=dependencies)
方案2(兼容h5格式):调整层定义避免共享+补全所有注册逻辑
如果必须使用h5格式保存,除了上述自定义层修改外,还要避免复用同一个AutoSRELU实例,改为创建两个独立的实例:
# 不要复用同一个实例,改为两个独立对象 model_scratch_auto.add(AutoSRELU()) model_scratch_auto.add(Dense(120, activation='relu')) model_scratch_auto.add(AutoSRELU())
保存和加载逻辑不变即可。
内容的提问来源于stack exchange,提问作者Ravish Jha
相关产品推荐
相关产品推荐

