Skip to content

PyTorch - 循环神经网络

循环神经网络 (Recurrent Neural Networks, RNNs) 是一类非常适合处理序列数据的神经网络。与假设输入相互独立的前馈网络 (feedforward networks) 不同,RNNs 保持一种内部状态(或记忆),它捕获了序列中先前元素的信息。这使得它们在自然语言处理 (natural language processing)、语音识别 (speech recognition) 和时间序列分析 (time series analysis) 等任务中表现强大。

RNN 通过迭代序列元素并在每个时间步更新其隐状态 (hidden state) 来处理序列。在每个时间步 t 的计算通常涉及当前输入 x_t 和先前的隐状态 h_{t-1},以产生新的隐状态 h_t 和可选的输出 o_t。这种循环机制允许信息在时间步之间持续存在。

在本章中,我们将演示如何在 PyTorch 中从头实现一个简单的 RNN 来模拟正弦波。这将有助于理解其工作原理,尽管在实际应用中,PyTorch 提供了优化的内置 RNN 层。

在训练期间,我们将一次一个数据点(一个时间步)地馈送给模型。输入序列 x 将包含 20 个数据点,目标序列 y 将是输入序列向后偏移一个时间步的结果,即预测序列中的下一个点。

导入必要的包。我们将使用 torch 进行神经网络功能,numpy 进行数值运算,以及 matplotlib 进行绘图。

import torch
import torch.nn as nn
import numpy as np
import matplotlib.pyplot as plt
import torch.nn.init as init

定义模型超参数并生成正弦波数据。在每个时间步,输入到我们 RNN 的数据将是当前数据点。隐层大小 (hidden layer size) 决定了 RNN 记忆容量。

dtype = torch.float32
input_size = 1 # Input dimension (sine wave value at time t)
hidden_size = 6 # Size of the hidden state
output_size = 1 # Output dimension (predicted sine wave value at t+1)
epochs = 300
seq_length = 20 # Length of the input sequence
learning_rate = 0.1
# Generate sine wave data
data_time_steps = np.linspace(2, 10, seq_length + 1, dtype=np.float32)
data = np.sin(data_time_steps)
data.resize((seq_length + 1, 1))
# Create input and target sequences
x_np = data[:-1]
y_np = data[1:]
x = torch.tensor(x_np, dtype=dtype)
y = torch.tensor(y_np, dtype=dtype)

x 是输入数据序列,y 是目标序列(x 偏移了一个时间步)。

我们将手动定义并初始化我们简单 RNN 的权重。w1 将拼接的输入和先前的隐状态映射到新的隐状态。w2 将新的隐状态映射到输出预测。我们将 requires_grad 设置为 True 以启用这些权重的梯度计算。

# Weights for input to hidden layer, and hidden to hidden layer (combined)
# Input x_t (size 1) and previous hidden h_{t-1} (size hidden_size) are concatenated
w1 = torch.empty(input_size + hidden_size, hidden_size, dtype=dtype)
init.xavier_normal_(w1) # A common initialization strategy
w1.requires_grad_(True)
# Weights for hidden to output layer
w2 = torch.empty(hidden_size, output_size, dtype=dtype)
init.xavier_normal_(w2)
w2.requires_grad_(True)
# Bias for the hidden layer (optional, can be added)
# b1 = torch.zeros(hidden_size, dtype=dtype, requires_grad=True)
# Bias for the output layer (optional)
# b2 = torch.zeros(output_size, dtype=dtype, requires_grad=True)

此函数定义了 RNN 计算的一个时间步。

def forward_step(input_t, context_state, w1_param, w2_param):
# Concatenate input_t and previous context_state (hidden_state)
# Ensure input_t is 2D: (1, input_size)
if input_t.ndim == 1:
input_t = input_t.unsqueeze(0) # Make it (1, input_size)
xh = torch.cat((input_t, context_state), 1)
# New hidden state (context_state)
context_state_next = torch.tanh(xh.mm(w1_param)) # Add + b1 if using bias
# Output prediction
out = context_state_next.mm(w2_param) # Add + b2 if using bias
return out, context_state_next

训练循环迭代多个 epoch(训练轮次)。在每个 epoch 中,我们处理整个序列。我们使用均方误差 (Mean Squared Error, MSE) 作为损失函数。

for i in range(epochs):
total_loss = 0
# Initialize context_state (hidden state) for the beginning of the sequence
current_context_state = torch.zeros((1, hidden_size), dtype=dtype)
for j in range(x.size(0)): # Iterate through each time step in the sequence
input_t = x[j:(j + 1)] # Current input (shape: [1, input_size])
target_t = y[j:(j + 1)] # Current target (shape: [1, output_size])
pred_t, next_context_state = forward_step(input_t, current_context_state, w1, w2)
loss = (pred_t - target_t).pow(2).sum() / 2 # MSE variant
total_loss += loss.item()
# Backward pass and update weights manually
# First, clear old gradients for w1 and w2 if they exist from previous iterations
if w1.grad is not None:
w1.grad.zero_()
if w2.grad is not None:
w2.grad.zero_()
loss.backward() # Compute gradients
with torch.no_grad(): # Temporarily disable gradient tracking for updates
w1 -= learning_rate * w1.grad
w2 -= learning_rate * w2.grad
# If using biases, update them here as well: b1 -= learning_rate * b1.grad
# Detach the context_state to prevent gradients from flowing back to the beginning of time
# This is crucial for Truncated Backpropagation Through Time (TBPTT) like behavior
# or when processing sequence step-by-step manually.
current_context_state = next_context_state.detach()
if (i + 1) % 10 == 0:
print(f"Epoch: {i+1}/{epochs}, Loss: {total_loss / x.size(0):.4f}")
# After training, generate predictions
predictions = []
with torch.no_grad(): # No need to track gradients for prediction
current_context_state_pred = torch.zeros((1, hidden_size), dtype=dtype)
for i in range(x.size(0)):
input_val = x[i:(i+1)]
pred, current_context_state_pred = forward_step(input_val, current_context_state_pred, w1, w2)
predictions.append(pred.item())

绘制原始正弦波和 RNN 的预测结果。

plt.figure(figsize=(12, 6))
plt.title('Sine Wave Prediction with Manual RNN')
plt.xlabel('Time Step')
plt.ylabel('Value')
plt.plot(data_time_steps[:-1], x.numpy().flatten(), 'bo-', label='Actual Data (Input)')
plt.plot(data_time_steps[1:], np.array(predictions), 'ro--', label='Predicted Data')
plt.legend()
plt.grid(True)
plt.show()

输出将是一个图表,显示原始正弦波(通常是蓝色圆圈或线)和 RNN 预测的正弦波(通常是红色虚线或圆圈)。一个好的模型会显示预测波形紧密跟随实际波形。

注意:虽然本示例出于教育目的演示了如何从头构建 RNN,但 PyTorch 在其 torch.nn 包中提供了优化且便捷的模块,如 torch.nn.RNN、torch.nn.LSTM 和 torch.nn.GRU。对于大多数实际应用,推荐使用这些模块,因为它们效率更高且在内部处理了许多复杂性。