如何用Plotly动态更新模型训练中的验证误差图表?
解决Plotly动态更新验证误差图表的问题
问题根源
每次调用fig.show()都会生成一个新的可视化输出实例,哪怕你修改了原fig对象的数据,新的show()还是会创建独立的图表,而不是更新已有的。
解决方案
根据运行环境不同,有以下几种实现方式:
1. Jupyter Notebook/Lab 环境(快速实现)
使用IPython的clear_output和display方法,每次更新数据后清空当前输出再显示更新后的图表:
import plotly.graph_objects as go import numpy as np from IPython.display import display, clear_output # 初始化图表 fig = go.Figure(data=[go.Scatter(x=[], y=[], name="Validation Error")]) fig.update_layout( title="Validation Error Over Epochs", xaxis_title="Epoch", yaxis_title="Error" ) # 模拟训练循环(替换为真实训练逻辑) validation_errors = [] for epoch in range(10): # 模拟计算当前Epoch的验证误差 current_err = np.random.rand() * 0.1 + 0.2 - epoch * 0.02 validation_errors.append(current_err) # 更新图表数据 fig.data[0].x = list(range(len(validation_errors))) fig.data[0].y = validation_errors # 清空旧输出并显示新图表(wait=True避免闪烁) clear_output(wait=True) display(fig)
2. Jupyter环境更优雅的方式:使用FigureWidget
FigureWidget是Plotly的交互式组件,修改数据后会自动更新显示,无需重复调用show():
import plotly.graph_objects as go import numpy as np from time import sleep # 初始化FigureWidget fig = go.FigureWidget(data=[go.Scatter(x=[], y=[], name="Validation Error")]) fig.update_layout( title="Validation Error Over Epochs", xaxis_title="Epoch", yaxis_title="Error" ) # 显示Widget(之后修改会自动更新) display(fig) # 模拟训练循环 validation_errors = [] for epoch in range(10): current_err = np.random.rand() * 0.1 + 0.2 - epoch * 0.02 validation_errors.append(current_err) # 批量更新数据,减少闪烁 with fig.batch_update(): fig.data[0].x = list(range(len(validation_errors))) fig.data[0].y = validation_errors sleep(0.5) # 模拟训练耗时
3. 本地Python脚本(非Jupyter):使用Plotly Dash
如果是本地脚本运行,用Dash框架创建实时更新的网页应用:
import dash from dash import dcc, html, Input, Output, State import plotly.graph_objects as go import numpy as np app = dash.Dash(__name__) app.layout = html.Div([ dcc.Graph(id='validation-error-graph', figure=go.Figure( data=[go.Scatter(x=[], y=[], name="Validation Error")], layout=go.Layout( title="Validation Error Over Epochs", xaxis_title="Epoch", yaxis_title="Error" ) )), dcc.Interval( id='train-interval', interval=1000, # 每秒更新一次,单位毫秒 n_intervals=0 ), html.Div(id='epoch-status', children="Epoch: 0") ]) # 存储验证误差的全局变量(实际训练中替换为真实数据) validation_errors = [] @app.callback( [Output('validation-error-graph', 'figure'), Output('epoch-status', 'children')], Input('train-interval', 'n_intervals'), State('validation-error-graph', 'figure') ) def update_training_graph(n, existing_fig): if n >= 10: # 模拟10轮训练后停止 return existing_fig, f"Epoch: {n} | 训练完成" # 模拟计算当前Epoch的验证误差 current_err = np.random.rand() * 0.1 + 0.2 - n * 0.02 validation_errors.append(current_err) # 更新图表数据 fig = go.Figure(existing_fig) fig.data[0].x = list(range(len(validation_errors))) fig.data[0].y = validation_errors return fig, f"Epoch: {n+1}" if __name__ == '__main__': app.run_server(debug=True)
运行后会在本地启动一个网页,图表会自动实时更新。
内容的提问来源于stack exchange,提问作者Захар Наумець
相关产品推荐
相关产品推荐

