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

子类化Keras模型保存后加载训练报错求助

解决Keras子类化Conv-RNN模型加载后训练异常的问题

我仔细梳理了你的代码和报错信息,这两个问题本质上都和自定义train_step的接口不符合Keras规范以及SavedModel的序列化逻辑有关,下面一步步帮你解决:

问题根源分析

  1. train_step()参数不匹配错误:你定义的train_step直接接受X_train和y_train两个独立参数,但Keras原生的train_step接口要求接收一个包含(inputs, labels)的元组作为单个参数。当模型保存再加载后,SavedModel对自定义方法的参数绑定逻辑发生了变化,此时调用model_2.train_step(X_train, target)会被解析为传递了3个参数(self + X_train + target),而加载后的方法只期望2个参数(self + 元组参数),因此抛出参数数量不匹配的报错。

  2. fit()无法找到匹配函数:一方面你的train_step接口不标准,另一方面call方法的input_signature把batch_size固定为1,这导致加载后Keras无法正确解析输入结构来匹配fit的执行逻辑,最终找不到对应的函数。

具体修复步骤

1. 调整train_step为Keras标准接口

修改train_step使其接收单个元组参数,并更新input_signature匹配这个结构,同时必须返回损失字典(Keras对train_step的强制要求):

@tf.function(input_signature=[(
    [spec1, spec1, spec1],  # 输入部分:3个TensorSpec组成的列表
    spec2                   # 标签部分
)])
def train_step(self, data):
    X_train, y_train = data  # 从元组中解包输入和标签
    with tf.GradientTape() as tape:
        y_pred = self(X_train) # Forward pass
        # 计算损失
        loss = self.loss_object(y_train, y_pred)
    # 计算梯度并更新权重
    gradients = tape.gradient(loss, self.trainable_variables)
    self.optimizer.apply_gradients(zip(gradients, self.trainable_variables))
    # 返回损失字典,这是Keras train_step的必要要求
    return {"loss": loss}

2. 解除call方法的固定batch_size限制

你当前的input_signature把batch_size固定为1,会导致模型兼容性极差。改成动态batch_size(用None表示),同时调整reshape的batch_size计算逻辑:

@tf.function(input_signature=[[
    tf.TensorSpec(shape=(None,40,5,1),name="M15"), 
    tf.TensorSpec(shape=(None,40,5,1),name="H1"), 
    tf.TensorSpec(shape=(None,40,5,1),name="H4")
]])
def call(self, data):
    conv1_res = self.conv1(data[0])
    conv2_res = self.conv2(data[1])
    conv3_res = self.conv3(data[2])
    # 用tf.shape获取动态batch_size,代替固定的1
    shape1 = (tf.shape(conv1_res)[0], conv1_res.shape[1], conv1_res.shape[2]*conv1_res.shape[3])
    shape2 = (tf.shape(conv2_res)[0], conv2_res.shape[1], conv2_res.shape[2]*conv2_res.shape[3])
    shape3 = (tf.shape(conv3_res)[0], conv3_res.shape[1], conv3_res.shape[2]*conv3_res.shape[3])
    
    f1 = self.lstm1B(self.lstm1A(tf.reshape(conv1_res, shape1)))
    f2 = self.lstm2B(self.lstm2A(tf.reshape(conv2_res, shape2)))
    f3 = self.lstm3B(self.lstm3A(tf.reshape(conv3_res, shape3)))
    
    pre_output = self.dense(self.concat([f1,f2,f3]))
    output = self.out(pre_output)
    return output

3. 修改训练时的train_step调用方式

现在train_step需要接收元组参数,训练时要把输入和标签打包后传递:

for i in range(iterations):
    state = data_environment[i]
    target= rand(1,3)
    X_train = [state[:1],state[1:2],state[2:3]]
    # 传递(inputs, labels)元组
    model_1.train_step((X_train, target))
    print("epoch", i)

4. 加载模型后的正确训练方式

修复后加载模型,你可以选择两种训练方式:

  • 直接调用train_step(传递元组参数):
model_2 = load_model("models/model_test1", compile=False)
model_2.train_step((X_train, target))
  • 重新compile后使用fit(现在接口符合标准,fit可以正常工作):
model_2.compile(loss='mse', optimizer='adam')
model_2.fit(X_train, target, epochs=1)

额外优化建议

  • 尽量删除__init__中直接定义的self.loss_object和self.optimizer,完全通过compile方法传入,这样加载模型后的兼容性会更好。
  • 如果没有特殊的梯度自定义需求,优先使用Keras原生训练流程,避免自定义train_step,能大幅降低保存加载的兼容性问题。

内容的提问来源于stack exchange,提问作者Tessan

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.08 20:17:40