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

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的交互式实时绘图

  1. 先安装依赖的交互式绘图库:
pip install ipympl
  1. 在Notebook开头配置正确的绘图后端:
%matplotlib widget
  1. 优化绘图逻辑:不要每次循环都新增线条,而是更新已有线条的数据,同时用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 22:00:55