PyTorch torch.nn.TransformerEncoder Function
PyTorch torch.nn Reference Manual
torch.nn.TransformerEncoderIt is the Transformer encoder module in PyTorch.
It consists of multiple TransformerEncoderLayer layers and is used to process input sequences.
Function Definition
torch.nn.TransformerEncoder(encoder_layer, num_layers, norm=None)
Parameters
encoder_layer: Encoder layernum_layers: Number of layers
Usage Examples
Example 1: Basic Usage
Example
import torch
import torch.nn as nn
# Single encoder layer
encoder_layer = nn.TransformerEncoderLayer(d_model=512, nhead=8, batch_first=True)
# Complete encoder
transformer_encoder = nn.TransformerEncoder(encoder_layer, num_layers=6)
# Input
x = torch.randn(32, 100, 512) # batch, seq, d_model
output = transformer_encoder(x)
print("Input shape:", x.shape)
print("Output shape:", output.shape)
import torch.nn as nn
# Single encoder layer
encoder_layer = nn.TransformerEncoderLayer(d_model=512, nhead=8, batch_first=True)
# Complete encoder
transformer_encoder = nn.TransformerEncoder(encoder_layer, num_layers=6)
# Input
x = torch.randn(32, 100, 512) # batch, seq, d_model
output = transformer_encoder(x)
print("Input shape:", x.shape)
print("Output shape:", output.shape)
Example 2: Text Classification
Example
import torch
import torch.nn as nn
class TransformerClassifier(nn.Module):
def __init__(self, vocab_size, d_model=256, nhead=4, num_layers=4):
super(TransformerClassifier, self).__init__()
self.embedding = nn.Embedding(vocab_size, d_model)
self.pos_encoder = nn.Parameter(torch.randn(1, 512, d_model) * 0.1)
encoder_layer = nn.TransformerEncoderLayer(d_model, nhead, batch_first=True)
self.encoder = nn.TransformerEncoder(encoder_layer, num_layers)
self.fc = nn.Linear(d_model, 10)
def forward(self, x):
# Simple positional encoding
x = self.embedding(x) + self.pos_encoder[:, :x.size(1), :]
x = self.encoder(x)
# Take the first token
return self.fc(x[:, 0, :])
model = TransformerClassifier(10000)
x = torch.randint(0, 10000, (32, 100))
output = model(x)
print("Output shape:", output.shape)
import torch.nn as nn
class TransformerClassifier(nn.Module):
def __init__(self, vocab_size, d_model=256, nhead=4, num_layers=4):
super(TransformerClassifier, self).__init__()
self.embedding = nn.Embedding(vocab_size, d_model)
self.pos_encoder = nn.Parameter(torch.randn(1, 512, d_model) * 0.1)
encoder_layer = nn.TransformerEncoderLayer(d_model, nhead, batch_first=True)
self.encoder = nn.TransformerEncoder(encoder_layer, num_layers)
self.fc = nn.Linear(d_model, 10)
def forward(self, x):
# Simple positional encoding
x = self.embedding(x) + self.pos_encoder[:, :x.size(1), :]
x = self.encoder(x)
# Take the first token
return self.fc(x[:, 0, :])
model = TransformerClassifier(10000)
x = torch.randint(0, 10000, (32, 100))
output = model(x)
print("Output shape:", output.shape)
Use Cases
- BERT: Encoder basics
- Text Classification
- Feature Extraction
Other Extensions