Jupyter Notebook中RNN预测动画无法生成的问题及求解
解决Jupyter Notebook中RNN预测动画无法连贯显示的问题
.ipynb与.py文件的核心差异
- 执行模式不同:.py文件由Python解释器全程连续运行,绘图窗口状态持续维护;而Jupyter Notebook按单元格分段执行,默认会把每个绘图输出成独立静态图片,无法维持同一窗口的实时更新。
- 绘图后端差异:.py默认使用系统GUI后端(如TkAgg、QtAgg),支持交互式实时绘图;Notebook默认用
inline后端,仅能生成静态图,即便切换到widget后端,也需要适配Notebook的事件循环逻辑。 - 状态维持逻辑:.py中全局变量和绘图状态全程保留;Notebook中单元格重新执行会重置状态,且绘图更新需要和Notebook的输出机制兼容。
解决动画显示问题的具体方案
方案1:适配Notebook的交互式实时绘图
- 先安装依赖的交互式绘图库:
pip install ipympl
- 在Notebook开头配置正确的绘图后端:
%matplotlib widget
- 优化绘图逻辑:不要每次循环都新增线条,而是更新已有线条的数据,同时用Notebook兼容的输出刷新方式替代
plt.pause()。
修改后的完整代码
import torch from torch import nn import numpy as np import matplotlib.pyplot as plt from IPython import display # 配置Notebook交互式后端 %matplotlib widget # Hyper Parameters TIME_STEP = 10 # rnn time step INPUT_SIZE = 1 # rnn input size LR = 0.02 # learning rate class RNN(nn.Module): def __init__(self): super(RNN, self).__init__() self.rnn = nn.RNN( input_size=INPUT_SIZE, hidden_size=32, # rnn hidden unit num_layers=1, # number of rnn layer batch_first=True, # input & output will has batch size as 1s dimension. e.g. (batch, time_step, input_size) ) self.out = nn.Linear(32, 1) def forward(self, x, h_state): r_out, h_state = self.rnn(x, h_state) # 简化输出计算,避免循环 outs = self.out(r_out) return outs, h_state rnn = RNN() print(rnn) optimizer = torch.optim.Adam(rnn.parameters(), lr=LR) loss_func = nn.MSELoss() h_state = None # 初始隐藏状态 # 初始化绘图对象,提前创建线条 fig, ax = plt.subplots(figsize=(12, 5)) line_real, = ax.plot([], [], 'r-', label='真实值') line_pred, = ax.plot([], [], 'b-', label='预测值') ax.set_xlabel('Step') ax.set_ylabel('Value') ax.legend() ax.set_ylim(-1.2, 1.2) # 固定y轴范围,避免画面跳动 plt.ion() for step in range(100): start, end = step * np.pi, (step+1)*np.pi steps = np.linspace(start, end, TIME_STEP, dtype=np.float32, endpoint=False) x_np = np.sin(steps) y_np = np.cos(steps) x = torch.from_numpy(x_np[np.newaxis, :, np.newaxis]) y = torch.from_numpy(y_np[np.newaxis, :, np.newaxis]) prediction, h_state = rnn(x, h_state) h_state = h_state.data # 断开隐藏状态的梯度连接 loss = loss_func(prediction, y) optimizer.zero_grad() loss.backward() optimizer.step() # 更新已有线条的数据,而非新增绘图 line_real.set_data(steps, y_np.flatten()) line_pred.set_data(steps, prediction.data.numpy().flatten()) ax.set_xlim(start, end) # 跟随当前步长更新x轴范围 # 刷新画布并适配Notebook输出 fig.canvas.draw() display.display(fig) display.clear_output(wait=True) plt.ioff() plt.show()
方案2:生成动画文件(适合保存或分享)
如果不需要实时训练过程的交互,也可以先收集所有帧数据,再用Matplotlib的FuncAnimation生成完整动画,在Notebook中播放:
import torch from torch import nn import numpy as np import matplotlib.pyplot as plt from matplotlib.animation import FuncAnimation from IPython.display import HTML # Hyper Parameters TIME_STEP = 10 INPUT_SIZE = 1 LR = 0.02 class RNN(nn.Module): def __init__(self): super(RNN, self).__init__() self.rnn = nn.RNN(INPUT_SIZE, 32, 1, batch_first=True) self.out = nn.Linear(32, 1) def forward(self, x, h_state): r_out, h_state = self.rnn(x, h_state) outs = self.out(r_out) return outs, h_state rnn = RNN() optimizer = torch.optim.Adam(rnn.parameters(), lr=LR) loss_func = nn.MSELoss() # 先收集所有训练帧的数据 frames_data = [] h_state = None for step in range(100): start, end = step * np.pi, (step+1)*np.pi steps = np.linspace(start, end, TIME_STEP, dtype=np.float32, endpoint=False) x_np = np.sin(steps) y_np = np.cos(steps) x = torch.from_numpy(x_np[np.newaxis, :, np.newaxis]) prediction, h_state = rnn(x, h_state) h_state = h_state.data loss = loss_func(prediction, y) optimizer.zero_grad() loss.backward() optimizer.step() frames_data.append((steps, y_np, prediction.data.numpy())) # 生成动画 fig, ax = plt.subplots(figsize=(12,5)) line_real, = ax.plot([], [], 'r-', label='真实值') line_pred, = ax.plot([], [], 'b-', label='预测值') ax.set_ylim(-1.2, 1.2) ax.legend() def update(frame): steps, y_np, pred = frame line_real.set_data(steps, y_np) line_pred.set_data(steps, pred.flatten()) ax.set_xlim(steps[0], steps[-1]) return line_real, line_pred # 创建动画对象,interval控制帧间隔(毫秒) ani = FuncAnimation(fig, update, frames=frames_data, interval=50, blit=True) # 在Notebook中显示动画 HTML(ani.to_jshtml())
内容的提问来源于stack exchange,提问作者Yue Qin
相关产品推荐
相关产品推荐

