PyTorch torch.nn.Transformer Function
PyTorch torch.nn Reference Manual
torch.nn.TransformerIt is the complete Transformer model in PyTorch.
It contains an encoder and a decoder, and can be used for sequence-to-sequence tasks.
Function Definition
torch.nn.Transformer(d_model=512, nhead=8, num_encoder_layers=6, num_decoder_layers=6, dim_feedforward=2048, dropout=0.1, activation='gelu', batch_first=True)
Usage Examples
Example 1: Basic Usage
Example
import torch
import torch.nn as nn
# Create Transformer
transformer = nn.Transformer(d_model=512, nhead=8, batch_first=True)
# Encoder input
src = torch.randn(10, 32, 512) # (seq, batch, d_model)
# Decoder input
tgt = torch.randn(20, 32, 512)
output = transformer(src, tgt)
print("Output shape:", output.shape)
import torch.nn as nn
# Create Transformer
transformer = nn.Transformer(d_model=512, nhead=8, batch_first=True)
# Encoder input
src = torch.randn(10, 32, 512) # (seq, batch, d_model)
# Decoder input
tgt = torch.randn(20, 32, 512)
output = transformer(src, tgt)
print("Output shape:", output.shape)
Example 2: Simple Translation Model
Example
import torch
import torch.nn as nn
class TransformerMT(nn.Module):
def __init__(self, vocab_size, d_model=256, nhead=4):
super(TransformerMT, self).__init__()
self.d_model = d_model
self.embedding = nn.Embedding(vocab_size, d_model)
self.transformer = nn.Transformer(d_model, nhead, batch_first=True)
self.fc = nn.Linear(d_model, vocab_size)
def forward(self, src, tgt):
src = self.embedding(src) * (self.d_model ** 0.5)
tgt = self.embedding(tgt) * (self.d_model ** 0.5)
out = self.transformer(src, tgt)
return self.fc(out)
model = TransformerMT(10000)
src = torch.randint(0, 10000, (32, 50))
tgt = torch.randint(0, 10000, (32, 40))
output = model(src, tgt)
print("Output shape:", output.shape)
import torch.nn as nn
class TransformerMT(nn.Module):
def __init__(self, vocab_size, d_model=256, nhead=4):
super(TransformerMT, self).__init__()
self.d_model = d_model
self.embedding = nn.Embedding(vocab_size, d_model)
self.transformer = nn.Transformer(d_model, nhead, batch_first=True)
self.fc = nn.Linear(d_model, vocab_size)
def forward(self, src, tgt):
src = self.embedding(src) * (self.d_model ** 0.5)
tgt = self.embedding(tgt) * (self.d_model ** 0.5)
out = self.transformer(src, tgt)
return self.fc(out)
model = TransformerMT(10000)
src = torch.randint(0, 10000, (32, 50))
tgt = torch.randint(0, 10000, (32, 40))
output = model(src, tgt)
print("Output shape:", output.shape)
Use Cases
- Machine Translation
- Text Generation
- Sequence-to-Sequence
Note: When batch_first=True, the input shape is (batch, seq, d_model).
Other Extensions