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

如何消除TensorFlow训练时Matplotlib动态绘图的重复显示

问题原因

Jupyter Notebook在单元格执行完毕后,会自动渲染所有已创建的Matplotlib Figure实例。你的自定义回调已经在每个batch结束时通过display(self.fig)展示了动态更新的图表,训练结束后Notebook会再次渲染这个已存在的Figure对象,导致最终绘图重复显示。

解决方案

在自定义回调中添加on_train_end方法,通过关闭训练过程中使用的Figure,避免Notebook自动重复渲染;或者清理输出后重新绘制最终的曲线。

修改后的完整代码

from IPython.display import display, clear_output
import tensorflow as tf
from tensorflow.keras.models import Sequential
import numpy as np
import matplotlib.pyplot as plt


class CustomCallback(tf.keras.callbacks.Callback):
    def on_train_begin(self, logs=None):
        self.epoch = 0  # Initialize the epoch counter
        self.accuracies = []
        self.fig, self.ax = plt.subplots()
        self.line, = self.ax.plot([], [])
        self.ax.set_xlim(0, 30)
        self.ax.set_ylim(0, 1)
        display(self.fig)
    
    def on_epoch_begin(self, epoch, logs=None):
        self.epoch = epoch  # Update the current epoch at the beginning of each epoch

    def on_train_batch_end(self, batch, logs=None):
        accuracy = logs['accuracy']
        self.accuracies.append(accuracy)
        self.line.set_data(range(1, len(self.accuracies) + 1), self.accuracies)
        self.ax.relim()
        self.ax.autoscale_view()
        clear_output(wait=True)
        display(self.fig)
    
    def on_train_end(self, logs=None):
        # 关闭训练过程中使用的figure,防止Notebook自动重复渲染
        plt.close(self.fig)
        # 可选:绘制并展示最终的准确率曲线
        clear_output(wait=True)
        fig_final, ax_final = plt.subplots()
        ax_final.plot(range(1, len(self.accuracies)+1), self.accuracies, label='Training Accuracy')
        ax_final.set_xlabel('Batch')
        ax_final.set_ylabel('Accuracy')
        ax_final.set_ylim(0, 1)
        ax_final.legend()
        display(fig_final)


custom_callback = CustomCallback()

model = Sequential()
model.add(tf.keras.layers.Dense(units=16, activation='relu'))
model.add(tf.keras.layers.Dropout(rate=0.35))
model.add(tf.keras.layers.Dense(units=1, activation='tanh'))

model.compile(optimizer=tf.keras.optimizers.Adam(), loss="binary_crossentropy", metrics=["accuracy"])

X = np.random.randn(10**2, 10**4)
y = np.random.randint(2, size=10**2)

abc = model.fit(X, y, epochs=7, batch_size=32, validation_split=0.025, verbose=False, callbacks=[custom_callback])

额外说明

  • 如果不需要单独绘制最终曲线,仅在on_train_end中调用plt.close(self.fig)即可解决重复问题,最后一次动态更新的图表会保留在单元格输出中。
  • 关闭Figure的操作只会影响训练时创建的临时绘图,不会干扰其他单元格的Matplotlib使用。

内容的提问来源于stack exchange,提问作者Brahim Khalil Abid

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.06 23:30:27