TensorFlow Callback权重保存与加载多异常问题排查
问题现象
训练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

