Keras子类化模型无法通过ModelCheckpoint完整保存问题咨询
子类化Keras模型保存与恢复训练的问题解决
问题背景
用TensorFlow 2.12.0实现带自定义train_step的子类化模型,通过ModelCheckpoint回调保存模型,期望完整保留架构、权重、自定义方法及训练配置以实现断点续训。但尝试SavedModel和Keras格式后,加载的模型丢失训练配置与自定义方法。
核心原因
Keras的SavedModel格式默认仅序列化模型的图结构与权重,无法保存子类化模型的Python自定义逻辑(如train_step)和训练状态元数据(如优化器参数、当前epoch)——这些属于代码层面的定义,无法被自动序列化到模型文件中。
解决方案步骤
1. 让自定义模型支持序列化
子类化模型需重写get_config()方法保存初始化参数,同时注册自定义类:
class CustomModel(tf.keras.Model): def __init__(self, units=32, **kwargs): super().__init__(**kwargs) self.units = units self.dense = tf.keras.layers.Dense(units) def train_step(self, data): # 自定义训练逻辑 x, y = data with tf.GradientTape() as tape: y_pred = self(x, training=True) loss = self.compiled_loss(y, y_pred) gradients = tape.gradient(loss, self.trainable_variables) self.optimizer.apply_gradients(zip(gradients, self.trainable_variables)) self.compiled_metrics.update_state(y, y_pred) return {m.name: m.result() for m in self.metrics} def get_config(self): # 保存模型初始化参数 config = super().get_config() config.update({"units": self.units}) return config # 注册自定义模型类,确保加载时能识别 tf.keras.utils.get_custom_objects()['CustomModel'] = CustomModel
2. 用Checkpoint保存完整训练状态
仅保存模型不足以恢复训练,需用tf.train.Checkpoint统一保存模型、优化器和训练元数据:
# 初始化训练组件 model = CustomModel(units=32) model.compile(optimizer='adam', loss='mse') optimizer = model.optimizer epoch_var = tf.Variable(0, dtype=tf.int64) # 构建检查点 checkpoint = tf.train.Checkpoint(model=model, optimizer=optimizer, epoch=epoch_var) checkpoint_manager = tf.train.CheckpointManager(checkpoint, './checkpoints', max_to_keep=3) # 模拟训练循环 initial_epoch = int(epoch_var.numpy()) epochs = 10 for epoch in range(initial_epoch, epochs): # 训练步骤 model.fit(x_train, y_train, epochs=1, initial_epoch=epoch) epoch_var.assign_add(1) checkpoint_manager.save()
3. 加载时恢复完整状态
加载时需先实例化自定义模型,再通过检查点恢复权重、优化器和训练状态:
# 重新实例化模型 model = CustomModel(units=32) model.build(input_shape=(None, x_train.shape[1])) # 手动构建模型 model.compile(optimizer='adam', loss='mse') optimizer = model.optimizer epoch_var = tf.Variable(0, dtype=tf.int64) # 恢复检查点 checkpoint = tf.train.Checkpoint(model=model, optimizer=optimizer, epoch=epoch_var) latest_checkpoint = tf.train.latest_checkpoint('./checkpoints') if latest_checkpoint: checkpoint.restore(latest_checkpoint) initial_epoch = int(epoch_var.numpy()) # 从断点继续训练 model.fit(x_train, y_train, epochs=10, initial_epoch=initial_epoch)
4. 结合ModelCheckpoint的兼容方案
若坚持用ModelCheckpoint,需设置save_weights_only=False,但加载时必须先定义并注册自定义模型类:
# 保存时的回调 checkpoint_callback = tf.keras.callbacks.ModelCheckpoint( './saved_model', save_weights_only=False, save_format='tf', save_best_only=True ) # 加载时 model = tf.keras.models.load_model('./saved_model', custom_objects={'CustomModel': CustomModel}) # 此时自定义train_step会保留,但训练配置(如epoch)需单独记录(比如用JSON文件存储)
关键注意事项
- 自定义方法(如
train_step)的逻辑依赖于模型类的定义,加载前必须确保自定义模型类已在当前环境中定义并注册。 - 训练元数据(当前epoch、损失值)不会随模型保存,需单独用文件(如JSON、txt)记录并在加载时读取。
内容的提问来源于stack exchange,提问作者lbrandao
相关产品推荐
相关产品推荐

