如何使用含小数点的指标命名TensorFlow模型检查点?
问题描述
我想按如下方式设置模型检查点:
mcp_save = tf.keras.callbacks.ModelCheckpoint('effnet-{epoch:02d}-{val_f0.5-score:.4f}.mdl_wts.hdf5', save_best_only=True, monitor='val_f0.5-score', mode='max')
但由于指标名称包含小数点,保存时出现错误:
'Failed to format this callback filepath: "effnet-{epoch:02d}-{val_f0.5-score:.4f}.mdl_wts.hdf5". Reason: 'val_f0''
该指标来自segmentation_models PyPI包:
fscore = sm.metrics.FScore(beta=0.5)
TensorFlow日志中显示的指标名称为:
1000/1000 [==============================] - ETA: 0s - loss: 0.6205 - accuracy: 0.2607 - f0.5-score: 0.3066
请问是否可以转义小数点或使用其他字符串,实现保存带分数的文件名?
解决方案
方法1:修改指标名称,替换小数点
初始化指标时手动指定名称,把小数点替换成下划线,避免与格式化语法冲突:
fscore = sm.metrics.FScore(beta=0.5, name='f0_5-score')
同步更新ModelCheckpoint的参数:
mcp_save = tf.keras.callbacks.ModelCheckpoint( 'effnet-{epoch:02d}-{val_f0_5-score:.4f}.mdl_wts.hdf5', save_best_only=True, monitor='val_f0_5-score', mode='max' )
这样文件名的格式化语法就能正确识别指标名称,不会触发报错。
方法2:自定义回调处理文件名
如果不想修改原有指标名称,可以继承ModelCheckpoint重写文件名生成逻辑:
from tensorflow.keras.callbacks import ModelCheckpoint class CustomModelCheckpoint(ModelCheckpoint): def on_epoch_end(self, epoch, logs=None): logs = logs or {} # 获取验证集f0.5分数 val_fscore = logs.get('val_f0.5-score') # 手动拼接文件名 self.filepath = f'effnet-{epoch:02d}-{val_fscore:.4f}.mdl_wts.hdf5' # 调用父类的保存逻辑 super().on_epoch_end(epoch, logs) # 实例化自定义回调 mcp_save = CustomModelCheckpoint( '', # 初始路径无意义,后续会被重写 save_best_only=True, monitor='val_f0.5-score', mode='max' )
这种方式直接绕过TensorFlow内置的文件名格式化逻辑,手动生成包含指标值的文件名。
内容的提问来源于stack exchange,提问作者Seth Kitchen
相关产品推荐
相关产品推荐

