PyTorch torch.nn.Transformer Function

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

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)

Use Cases

  • Machine Translation
  • Text Generation
  • Sequence-to-Sequence

Note: When batch_first=True, the input shape is (batch, seq, d_model).


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

Other Extensions