PyTorch torch.nn.GRU Function
PyTorch torch.nn Reference Manual
torch.nn.GRUIs the Gated Recurrent Unit module in PyTorch.
GRU is a simplified version of LSTM, with fewer parameters, faster computation, and similar performance.
Function Definition
torch.nn.GRU(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
# GRU: 256-dim input, 256-dim hidden, 2 layers
gru = nn.GRU(input_size=256, hidden_size=256, num_layers=2, batch_first=True)
# Input: batch=4, sequence=10, features=256
x = torch.randn(4, 10, 256)
output, hidden = gru(x)
print("Input shape:", x.shape)
print("Output shape:", output.shape)
print("Hidden state shape:", hidden.shape)
import torch.nn as nn
# GRU: 256-dim input, 256-dim hidden, 2 layers
gru = nn.GRU(input_size=256, hidden_size=256, num_layers=2, batch_first=True)
# Input: batch=4, sequence=10, features=256
x = torch.randn(4, 10, 256)
output, hidden = gru(x)
print("Input shape:", x.shape)
print("Output shape:", output.shape)
print("Hidden state shape:", hidden.shape)
Example 2: Comparison with LSTM
Example
import torch
import torch.nn as nn
import time
# LSTM and GRU with the same configuration
lstm = nn.LSTM(256, 256, 1, batch_first=True)
gru = nn.GRU(256, 256, 1, batch_first=True)
x = torch.randn(32, 100, 256)
# Performance comparison
for model, name in [(lstm, "LSTM"), (gru, "GRU")]:
start = time.time()
for _ in range(100):
_ = model(x)
print(f"{name} Time: {time.time()-start:.3f}s")
import torch.nn as nn
import time
# LSTM and GRU with the same configuration
lstm = nn.LSTM(256, 256, 1, batch_first=True)
gru = nn.GRU(256, 256, 1, batch_first=True)
x = torch.randn(32, 100, 256)
# Performance comparison
for model, name in [(lstm, "LSTM"), (gru, "GRU")]:
start = time.time()
for _ in range(100):
_ = model(x)
print(f"{name} Time: {time.time()-start:.3f}s")
Example 3: Classification Task
Example
import torch
import torch.nn as nn
class GRUClassifier(nn.Module):
def __init__(self, vocab_size, embed_dim, hidden_dim, num_classes):
super(GRUClassifier, self).__init__()
self.embedding = nn.Embedding(vocab_size, embed_dim)
self.gru = nn.GRU(embed_dim, hidden_dim, batch_first=True, bidirectional=True)
self.fc = nn.Linear(hidden_dim * 2, num_classes)
def forward(self, x):
embedded = self.embedding(x)
_, hidden = self.gru(embedded)
# Concatenate the final hidden states of the bidirectional
hidden = torch.cat([hidden[-2], hidden[-1]], dim=1)
return self.fc(hidden)
model = GRUClassifier(10000, 128, 128, 2)
x = torch.randint(0, 10000, (8, 50))
output = model(x)
print("Input:", x.shape, "-> Output:", output.shape)
import torch.nn as nn
class GRUClassifier(nn.Module):
def __init__(self, vocab_size, embed_dim, hidden_dim, num_classes):
super(GRUClassifier, self).__init__()
self.embedding = nn.Embedding(vocab_size, embed_dim)
self.gru = nn.GRU(embed_dim, hidden_dim, batch_first=True, bidirectional=True)
self.fc = nn.Linear(hidden_dim * 2, num_classes)
def forward(self, x):
embedded = self.embedding(x)
_, hidden = self.gru(embedded)
# Concatenate the final hidden states of the bidirectional
hidden = torch.cat([hidden[-2], hidden[-1]], dim=1)
return self.fc(hidden)
model = GRUClassifier(10000, 128, 128, 2)
x = torch.randint(0, 10000, (8, 50))
output = model(x)
print("Input:", x.shape, "-> Output:", output.shape)
LSTM vs GRU
| Aspect | LSTM | GRU |
|---|---|---|
| Parameter count | More | Fewer |
| Gating | 3 gates | 2 gates |
| Computation | Slower | Faster |
Use Cases
- Sequence modeling: Text, audio
- Rapid prototyping: When resources are limited
- Machine translation: Encoder side
Other Extensions