TensorFlow2子类化Keras模型加载后调用非默认参数报ValueError问题
根本原因
- TensorFlow的SavedModel格式本质存储的是一系列经
tf.function追踪后的静态计算图签名,而非原Python类的动态逻辑。使用tf.saved_model.load加载后得到的是静态图包装对象,仅能调用保存时已被追踪过的参数组合对应的签名。 training是Keras内置特殊参数,框架在模型训练、验证的生命周期中会自动对training=True和training=False两种场景的call逻辑执行追踪,因此保存后两种取值均可用。- 自定义的
full_batch_eval属于非内置参数:你在保存前eager模式下的参数调用不会被纳入SavedModel的追踪范围,训练过程中调用模型时仅使用了参数默认值,未在tf.function上下文中触发full_batch_eval=False的逻辑追踪,因此SavedModel中没有对应参数组合的计算图,加载后调用就会触发匹配错误。
解决方案
方案1:使用Keras原生保存加载逻辑,保留子类模型动态特性
不使用底层的tf.saved_model.save/load,改用Keras官方的模型保存加载接口,加载时指定自定义类映射,得到的是原TestModel类的实例,支持所有动态参数调用,无需修改原有模型逻辑:
# 保存阶段用Keras接口 prefix = 'saved_model_dir' m.save(prefix) # 加载阶段指定自定义类映射 m = keras.models.load_model(prefix, custom_objects={"TestModel": TestModel})
方案2:保存前显式追踪所有需要的参数组合
如果必须使用tf.saved_model.save/load的底层接口,可以在保存前显式触发所有需要用到的参数组合的tf.function追踪,这些签名会被一起存入SavedModel:
# 保存前先触发所有参数组合的追踪 for (x,y) in train_ds.take(1): call_fn = tf.function(m) # 显式调用所有需要用到的参数组合 call_fn({'x':x,'y':y}, training=True, full_batch_eval=False) call_fn({'x':x,'y':y}, training=False, full_batch_eval=False) # 再执行正常保存 prefix = 'saved_model_dir' tf.saved_model.save(m, export_dir=prefix)
方案3:显式指定call方法的完整输入签名
定义模型时用@tf.function装饰call方法并声明完整输入签名,将自定义参数纳入静态追踪范围:
class TestModel(keras.Model): def __init__(self,**kwargs): super(TestModel,self).__init__(**kwargs) self.f = Flatten() self.dense_1 = Dense(units=10,activation=tf.nn.relu) # 显式指定完整输入签名,匹配你的字典输入和自定义参数 @tf.function(input_signature=[ tf.TypeSpec(dict( x=tf.TensorSpec(shape=(None, 28,28,1), dtype=tf.float32), y=tf.TensorSpec(shape=(None,), dtype=tf.float32) ), name='inputs'), tf.TensorSpec(shape=(), dtype=tf.bool, name='training'), tf.TensorSpec(shape=(), dtype=tf.bool, name='full_batch_eval') ]) def call(self,inputs,training=False,full_batch_eval=True): x = inputs.get('x') return self.dense_1(self.f(x))
内容的提问来源于stack exchange,提问作者user5915898
相关产品推荐
相关产品推荐

