PyTorch torch.nn.RNN Function
PyTorch torch.nn Reference Manual
torch.nn.RNNIt is the basic recurrent neural network module in PyTorch.
It is the simplest recurrent layer, but it is prone to the vanishing gradient problem.
Function Definition
torch.nn.RNN(input_size, hidden_size, num_layers=1, bias=True, batch_first=True, dropout=0, bidirectional=False)
Usage Examples
Example 1: Basic Usage
Example
import torch
import torch.nn as nn
rnn = nn.RNN(input_size=256, hidden_size=128, num_layers=2, batch_first=True)
x = torch.randn(4, 10, 256)
output, hidden = rnn(x)
print("Input:", x.shape)
print("Output:", output.shape)
print("Hidden:", hidden.shape)
import torch.nn as nn
rnn = nn.RNN(input_size=256, hidden_size=128, num_layers=2, batch_first=True)
x = torch.randn(4, 10, 256)
output, hidden = rnn(x)
print("Input:", x.shape)
print("Output:", output.shape)
print("Hidden:", hidden.shape)
Example 2: Multi-layer RNN
Example
import torch
import torch.nn as nn
# 3-layer RNN with dropout
rnn = nn.RNN(128, 256, num_layers=3, dropout=0.3, batch_first=True)
x = torch.randn(2, 50, 128)
out, h = rnn(x)
print("Input:", x.shape)
print("Output:", out.shape)
print("Hidden:", h.shape) # (3, 2, 256)
import torch.nn as nn
# 3-layer RNN with dropout
rnn = nn.RNN(128, 256, num_layers=3, dropout=0.3, batch_first=True)
x = torch.randn(2, 50, 128)
out, h = rnn(x)
print("Input:", x.shape)
print("Output:", out.shape)
print("Hidden:", h.shape) # (3, 2, 256)
Example 3: Non-linear Activation
Example
import torch
import torch.nn as nn
# tanh is used as the activation function by default
rnn = nn.RNN(64, 64, batch_first=True)
x = torch.randn(1, 5, 64)
out, _ = rnn(x)
print("Output shape:", out.shape)
print("RNN uses tanh activation by default")
import torch.nn as nn
# tanh is used as the activation function by default
rnn = nn.RNN(64, 64, batch_first=True)
x = torch.randn(1, 5, 64)
out, _ = rnn(x)
print("Output shape:", out.shape)
print("RNN uses tanh activation by default")
Notes
The basic RNN has the vanishing gradient problem; for long sequences, LSTM or GRU is recommended.
Use Cases
- Simple sequence tasks: short sequences
- Teaching examples: understanding RNN principles
- Rapid prototyping: simple baseline
Other Extensions