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:
Add serialization methods to your custom model
Update yourImgToClassSimpleContinuousclass to includeget_config()andfrom_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)Adjust the ModelCheckpoint filepath
To save full SavedModels (withmodel.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:
Save and load the full model (not just weights)
Usesave_weights_only=Falsein ModelCheckpoint to save the entire model (including metric states), then load it withtf.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.
Fix your code execution order
Never load weights before compiling the model—this resets metric states. Always follow this sequence:- Create model
- Compile model
- Load pre-trained model/weights
- Train with callbacks
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 bestval_lossvalue: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

