TensorFlow Summary Writer与Keras TensorBoard回调共存时失效问题
解决TensorFlow 2.13中TensorBoard回调与自定义Summary Writer共存写入失败的问题
问题原因
TensorBoard回调在训练过程中会维护自身的Summary Writer上下文,当你同时使用自定义的image_writer时,两个Writer的上下文会发生冲突——TensorBoard回调的上下文会覆盖默认上下文,导致自定义Writer的写入操作无法生效,最终tf.summary.image返回False,对应文件无更新。移除TensorBoard回调后,自定义Writer的上下文成为默认,因此能正常工作。
解决方案
方案1:复用TensorBoard回调的Writer
直接使用TensorBoard回调内置的Writer来写入ROC曲线,避免上下文冲突:
import tensorflow as tf from tensorflow import keras import matplotlib.pyplot as plt import io class MetricsCallback(keras.callbacks.Callback): def _plot_roc_curve(self, y_true, y_pred): # 计算ROC曲线数据 fpr, tpr, _ = tf.keras.metrics.roc_curve(y_true, y_pred) auc = tf.keras.metrics.auc(fpr, tpr).numpy() # 绘制ROC曲线并转为张量 fig = plt.figure(figsize=(6,6)) plt.plot(fpr, tpr, label=f'AUC = {auc:.2f}') plt.plot([0,1], [0,1], 'k--') plt.xlabel('False Positive Rate') plt.ylabel('True Positive Rate') plt.title('ROC Curve') plt.legend() buf = io.BytesIO() plt.savefig(buf, format='png') buf.seek(0) image = tf.image.decode_png(buf.getvalue(), channels=4) image = tf.expand_dims(image, 0) # 增加batch维度 plt.close(fig) return image def on_epoch_end(self, epoch, logs=None): # 获取验证集标签与预测值 val_data = self.model.validation_data y_true = val_data[1] y_pred = self.model.predict(val_data[0], verbose=0) roc_image = self._plot_roc_curve(y_true, y_pred) # 找到TensorBoard回调实例 tb_callback = None for cb in self.model.callbacks: if isinstance(cb, keras.callbacks.TensorBoard): tb_callback = cb break if not tb_callback: return # 使用TensorBoard的Writer写入并强制刷新 with tb_callback._train_writer.as_default(): tf.summary.image('ROC Curve', roc_image, step=epoch) tb_callback._train_writer.flush()
方案2:强制自定义Writer的上下文独立性
在自定义回调中明确使用自己的Writer上下文,并手动刷新写入:
import tensorflow as tf from tensorflow import keras import matplotlib.pyplot as plt import io import os class MetricsCallback(keras.callbacks.Callback): def __init__(self, run_logdir): super().__init__() self.image_writer = tf.summary.create_file_writer(os.path.join(run_logdir, 'ROC_curve')) def _plot_roc_curve(self, y_true, y_pred): # 同方案1的绘制逻辑 fpr, tpr, _ = tf.keras.metrics.roc_curve(y_true, y_pred) auc = tf.keras.metrics.auc(fpr, tpr).numpy() fig = plt.figure(figsize=(6,6)) plt.plot(fpr, tpr, label=f'AUC = {auc:.2f}') plt.plot([0,1], [0,1], 'k--') plt.xlabel('False Positive Rate') plt.ylabel('True Positive Rate') plt.title('ROC Curve') plt.legend() buf = io.BytesIO() plt.savefig(buf, format='png') buf.seek(0) image = tf.image.decode_png(buf.getvalue(), channels=4) image = tf.expand_dims(image, 0) plt.close(fig) return image def on_epoch_end(self, epoch, logs=None): val_data = self.model.validation_data y_true = val_data[1] y_pred = self.model.predict(val_data[0], verbose=0) roc_image = self._plot_roc_curve(y_true, y_pred) # 明确使用自定义Writer上下文并强制刷新 with self.image_writer.as_default(): tf.summary.image('ROC Curve', roc_image, step=epoch) self.image_writer.flush()
使用示例
run_logdir = './logs' # 配置TensorBoard回调 tensorboard_cb = keras.callbacks.TensorBoard(log_dir=run_logdir, update_freq='epoch') # 实例化自定义回调 metrics_cb = MetricsCallback(run_logdir) # 同时传入两个回调训练模型 model.fit( train_data, validation_data=val_data, epochs=10, callbacks=[tensorboard_cb, metrics_cb] )
内容的提问来源于stack exchange,提问作者sera sun
相关产品推荐
相关产品推荐

