PyTorch torch.nn.RNN Function

PyTorch torch.nn 参考手册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)

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)

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")

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

PyTorch torch.nn 参考手册PyTorch torch.nn Reference Manual

Other Extensions