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

TensorFlow Callback权重保存与加载多异常问题排查

TensorFlow模型训练回调配置异常问题

问题现象

训练TensorFlow模型过程中,配置TensorBoard回调存储训练日志、ModelCheckpoint回调存储模型权重,训练与评估阶段出现两个异常:

  • 每个epoch训练完成后弹出警告:WARNING:tensorflow: Can save best model only with val_acc available, skipping,无法保存最优模型权重
  • 训练完成后克隆与原模型完全一致的结构,调用cloned_model.load_weights(checkpoint_path)加载权重后,与原模型分别在测试集执行评估:原模型测试集准确率可达70%以上,克隆模型准确率始终固定为0.54,结果明显异常

初步排查

最初猜测是checkpoint存储路径下残留了之前训练保存的高准确率模型权重,导致新训练过程无法触发保存,但检查对应路径未发现旧权重文件;且如果路径下确实为高准确率旧权重,加载后克隆模型准确率不应仅为0.54,该猜测无法解释异常。

相关代码

TensorBoard回调实现

def create_tensorboard_callback(dir_name, experiment_name):
  """
  Creates a TensorBoard callback instance to store log files.

  Stores log files with the filepath:
    "dir_name/experiment_name/current_datetime/"

  Args:
    dir_name: target directory to store TensorBoard log files
    experiment_name: name of experiment directory (e.g. efficientnet_model_1)
  """
  log_dir = dir_name + "/" + experiment_name + "/" + datetime.datetime.now().strftime("%Y%m%d-%H%M%S")
  tensorboard_callback = tf.keras.callbacks.TensorBoard(
      log_dir=log_dir
  )
  print(f"Saving TensorBoard log files to: {log_dir}")
  return tensorboard_callback

ModelCheckpoint回调配置

# Create ModelCheckpoint callback to save model's progress
checkpoint_path = "model_checkpoints/cp.ckpt" 
model_checkpoint = tf.keras.callbacks.ModelCheckpoint(checkpoint_path,
                                                      monitor="val_acc", 
                                                      save_best_only=True,
                                                      save_weights_only=True, 
                                                      verbose=0) 

模型训练调用代码

history_101_food_classes_feature_extract = model.fit(train_data, 
                                                     epochs=3,
                                                     steps_per_epoch=len(train_data),
                                                     validation_data=test_data,
                                                     validation_steps=int(0.15 * len(test_data)),
                                                     callbacks=[create_tensorboard_callback("training_logs", 
                                                                                            "efficientnetb0_101_classes_all_data_feature_extract"),
                                                                model_checkpoint])

根因与修复方案

根因

ModelCheckpoint回调的monitor参数配置错误:
TensorFlow 2.x版本中,模型编译时如果直接传入metrics=["accuracy"]使用内置准确率指标,训练过程自动记录的验证集准确率字段名为val_accuracy,而非TensorFlow 1.x时期的旧字段名val_acc。
由于配置的监控指标val_acc在训练日志中不存在,回调无法获取对应指标值,因此每个epoch都会跳过最优模型保存步骤,整个训练过程没有在指定路径生成当前训练的有效权重文件。克隆模型加载权重时实际加载的是随机初始化的无效权重,才会出现准确率远低于训练后原模型的异常。

修复方法

将ModelCheckpoint配置中的monitor参数值从"val_acc"修改为"val_accuracy",重新训练即可:

model_checkpoint = tf.keras.callbacks.ModelCheckpoint(checkpoint_path,
                                                      monitor="val_accuracy", 
                                                      save_best_only=True,
                                                      save_weights_only=True, 
                                                      verbose=0) 

如果需要沿用val_acc作为监控字段,可在模型编译时显式指定准确率指标的名称,和监控配置保持一致即可:

model.compile(..., metrics=[tf.keras.metrics.CategoricalAccuracy(name="acc")])

修复后回调可正常识别验证集准确率指标,自动保存最优权重,训练结束后克隆模型加载对应权重,评估结果将与原模型保持一致。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 15:36:27