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

TensorFlow2子类化Keras模型加载后调用非默认参数报ValueError问题

根本原因
  1. TensorFlow的SavedModel格式本质存储的是一系列经tf.function追踪后的静态计算图签名,而非原Python类的动态逻辑。使用tf.saved_model.load加载后得到的是静态图包装对象,仅能调用保存时已被追踪过的参数组合对应的签名。
  2. training是Keras内置特殊参数,框架在模型训练、验证的生命周期中会自动对training=True和training=False两种场景的call逻辑执行追踪,因此保存后两种取值均可用。
  3. 自定义的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.27 01:54:04