PyTorch torch.nn.GRU Function

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

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

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)

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

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

Other Extensions