PyTorch torch.nn.TransformerEncoder Function

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

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)

Use Cases

  • BERT: Encoder basics
  • Text Classification
  • Feature Extraction

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

Other Extensions