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

Tensorflow ModelCheckpoint回调无法保存SavedModel文件,加载Checkpoint后无法恢复val_loss值的技术求助

Let’s break down your two ModelCheckpoint issues and fix them step by step—these are common hurdles with custom Keras models, so I’ve got practical solutions for you.


Issue 1: ModelCheckpoint isn’t generating SavedModel’s model.pb file

Your custom ImgToClassSimpleContinuous model inherits from tf.keras.Model, but it’s missing serialization methods that Keras needs to save the full model structure (including model.pb). Even with save_weights_only=False, without these methods, Keras can’t properly serialize your model and falls back to saving only checkpoint files.

Fix Steps:

  1. Add serialization methods to your custom model
    Update your ImgToClassSimpleContinuous class to include get_config() and from_config(), which let Keras save and reload your model’s structure:

    def __init__(self, img_height, img_width, *args, **kwargs):
        super().__init__(ImgToClassSimple, *args, **kwargs)
        self.img_height = img_height  # Save as instance attribute
        self.img_width = img_width    # Save as instance attribute
        # ... rest of your __init__ code remains the same ...
    
    def get_config(self):
        config = super().get_config()
        config.update({
            'img_height': self.img_height,
            'img_width': self.img_width
        })
        return config
    
    @classmethod
    def from_config(cls, config):
        return cls(**config)
    
  2. Adjust the ModelCheckpoint filepath
    To save full SavedModels (with model.pb), set your filepath to a directory (or use placeholders like {epoch} to avoid overwriting):

    checkpoint_filename = "../chkpts/ImgToClassSimpleContinuous/best_model"
    

    This will create a directory containing the complete SavedModel whenever a new best model is saved.


Issue 2: Can’t restore val_loss state after loading checkpoints

When you use model.load_weights(), you’re only loading the model’s weight parameters—not the state of training metrics like the best val_loss value that save_best_only=True relies on. Without this state, the callback can’t compare new val_loss values to the previous best.

Fix Options:

  1. Save and load the full model (not just weights)
    Use save_weights_only=False in ModelCheckpoint to save the entire model (including metric states), then load it with tf.keras.models.load_model():

    # Replace your model.load_weights() call with this
    try:
        model = tf.keras.models.load_model(checkpoint_filename)
        print("Loaded existing best model with val_loss tracking intact")
    except:
        model = ImgToClassSimpleContinuous(img_height, img_width)
        model.compile(optimizer='adam', loss=[LOSS, LOSS, LOSS], metrics=['mse'])
        print("No existing model found, starting fresh")
    

    Note: This requires the serialization methods added in Issue 1 to work.

  2. Fix your code execution order
    Never load weights before compiling the model—this resets metric states. Always follow this sequence:

    1. Create model
    2. Compile model
    3. Load pre-trained model/weights
    4. Train with callbacks
  3. Manual val_loss tracking (if you prefer weights-only saving)
    If you want to stick to weight checkpoints, create a custom callback to save and load the best val_loss value:

    class BestLossTracker(tf.keras.callbacks.Callback):
        def __init__(self, loss_file_path):
            super().__init__()
            self.loss_file = loss_file_path
            self.best_loss = float('inf')
            # Load saved best loss if it exists
            try:
                with open(self.loss_file, 'r') as f:
                    self.best_loss = float(f.read())
            except FileNotFoundError:
                pass
    
        def on_epoch_end(self, epoch, logs=None):
            current_val_loss = logs.get('val_loss')
            if current_val_loss < self.best_loss:
                self.best_loss = current_val_loss
                with open(self.loss_file, 'w') as f:
                    f.write(str(self.best_loss))
    
    # Use it alongside ModelCheckpoint
    loss_tracker = BestLossTracker("../chkpts/best_val_loss.txt")
    

    Then initialize ModelCheckpoint with the loaded best_loss (you’d need to modify the callback to accept this value, but the full model approach is simpler).


Full Fixed Code Snippet

import tensorflow as tf
from tensorflow.keras import Model, layers

# Define constants (replace with your values)
LOSS = tf.keras.losses.MeanSquaredError()
img_height = 224
img_width = 224
MAX_EPOCHS = 50
BATCH_SIZE = 32

class ImgToClassSimpleContinuous(Model):
    ''' pair with loss = categorical_crossentropy '''
    in_types = [DataType.d]
    out_types = [DataType.tlc, DataType.tls, DataType.tll]
    
    def __init__(self, img_height, img_width, *args, **kwargs):
        super().__init__(ImgToClassSimple, *args, **kwargs)
        self.img_height = img_height
        self.img_width = img_width
        initializer = 'he_normal'
        input_shape = (img_height, img_width, 1)
        inputs = tf.keras.Input(shape=input_shape)
        
        # Model architecture (unchanged)
        x = layers.Conv2D(8, 3, padding='same', kernel_initializer=initializer)(inputs)
        x = layers.PReLU()(x)
        x = layers.Conv2D(8, 3, padding='same', kernel_initializer=initializer)(x)
        x = layers.PReLU()(x)
        x = layers.MaxPooling2D(pool_size=(2, 2))(x)
        x = layers.BatchNormalization()(x)
        
        x = layers.Conv2D(16, 3, padding='same', kernel_initializer=initializer)(x)
        x = layers.PReLU()(x)
        x = layers.Conv2D(16, 3, padding='same', kernel_initializer=initializer)(x)
        x = layers.PReLU()(x)
        x = layers.MaxPooling2D(pool_size=(2, 2))(x)
        x = layers.BatchNormalization()(x)
        
        # Branch t
        t = layers.Conv2D(32, 3, padding='same', kernel_initializer=initializer)(x)
        t = layers.PReLU()(t)
        t = layers.Conv2D(32, 3, padding='same', kernel_initializer=initializer)(t)
        t = layers.PReLU()(t)
        t = layers.MaxPooling2D(pool_size=(2, 2))(t)
        t = layers.BatchNormalization()(t)
        t = tf.keras.layers.GlobalAveragePooling2D()(t)
        t = layers.Flatten()(t)
        t = layers.Dense(1, activation='sigmoid')(t)
        
        # Branch s
        s = layers.Conv2D(32, 3, padding='same', kernel_initializer=initializer)(x)
        s = layers.PReLU()(s)
        s = layers.Conv2D(32, 3, padding='same', kernel_initializer=initializer)(s)
        s = layers.PReLU()(s)
        s = layers.MaxPooling2D(pool_size=(2, 2))(s)
        s = layers.BatchNormalization()(s)
        s = tf.keras.layers.GlobalAveragePooling2D()(s)
        s = layers.Flatten()(s)
        s = layers.Dense(1, activation='sigmoid')(s)
        
        # Branch l
        l = layers.Conv2D(32, 3, padding='same', kernel_initializer=initializer)(x)
        l = layers.PReLU()(l)
        l = layers.Conv2D(32, 3, padding='same', kernel_initializer=initializer)(l)
        l = layers.PReLU()(l)
        l = layers.MaxPooling2D(pool_size=(2, 2))(l)
        l = layers.BatchNormalization()(l)
        l = tf.keras.layers.GlobalAveragePooling2D()(l)
        l = layers.Flatten()(l)
        l = layers.Dense(1, activation='sigmoid')(l)
        
        self.model = tf.keras.Model(inputs, [t, s, l])
        tf.keras.utils.plot_model(self.model, to_file="model_plot.png", show_shapes=True)
    
    def call(self, x):
        return self.model(x)
    
    def get_config(self):
        config = super().get_config()
        config.update({
            'img_height': self.img_height,
            'img_width': self.img_width
        })
        return config
    
    @classmethod
    def from_config(cls, config):
        return cls(**config)

# Training workflow
checkpoint_filename = "../chkpts/ImgToClassSimpleContinuous/best_model"

# Load existing model or initialize new one
try:
    model = tf.keras.models.load_model(checkpoint_filename)
    print("Loaded existing best model, resuming training...")
except:
    model = ImgToClassSimpleContinuous(img_height, img_width)
    model.compile(optimizer='adam', loss=[LOSS, LOSS, LOSS], metrics=['mse'])
    print("No existing model found, starting training from scratch...")

# Setup ModelCheckpoint
cp_callback = tf.keras.callbacks.ModelCheckpoint(
    filepath=checkpoint_filename,
    verbose=1,
    mode='min',
    monitor="val_loss",
    save_best_only=True,
    save_weights_only=False
)

# Start training
model.fit(
    dataset_to_use,
    validation_data=dataset_validation_batched,
    epochs=MAX_EPOCHS,
    batch_size=BATCH_SIZE,
    callbacks=[cp_callback]
)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.30 09:49:05