子类化Keras模型保存后加载训练报错求助
我仔细梳理了你的代码和报错信息,这两个问题本质上都和自定义train_step的接口不符合Keras规范以及SavedModel的序列化逻辑有关,下面一步步帮你解决:
问题根源分析
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+ 元组参数),因此抛出参数数量不匹配的报错。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

