PyTorch torch.nn.Embedding Function

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


torch.nn.EmbeddingIt is a module in PyTorch used for word embedding.

It maps discrete word indices to a continuous vector space, and is one of the most fundamental operations in natural language processing.

Function Definition

torch.nn.Embedding(num_embeddings, embedding_dim, padding_idx=None, max_norm=None, norm_type=2.0, scale_grad_by_freq=False, sparse=False)

Parameter Description:

  • num_embeddings(int): Vocabulary size, i.e., the number of rows in the embedding matrix.
  • embedding_dim(int): The dimension of each embedding vector.
  • padding_idx(int): Specifies the padding index, whose embedding vector is a zero vector. Default is None.
  • max_norm(float): If not None, embedding vectors are normalized to this norm. Default is None.
  • norm_type(float): The order of the norm to compute. Default is 2.0.
  • scale_grad_by_freq(bool): Whether to scale gradients based on word frequency. Default is False.
  • sparse(bool): Whether the weight matrix is a sparse matrix. Default is False.

Attributes:

  • weight(Tensor): Learnable weights of shape (num_embeddings, embedding_dim).

Usage Examples

Example 1: Basic Usage

Create and use word embeddings:

Example

import torch
import torch.nn as nn

# Create embedding layer: vocabulary 10000, embedding dimension 256
embedding = nn.Embedding(num_embeddings=10000, embedding_dim=256)

# Word indices (starting from 0)
# Shape: (batch, seq_len)
input_indices = torch.tensor([[12, 45, 678], [901, 23, 56]])

# Look up embedding vectors
output = embedding(input_indices)

print("Input index shape:", input_indices.shape)
print("Output embedding shape:", output.shape)  # (2, 3, 256)

# View the shape of the embedding matrix
print("Embedding matrix shape:", embedding.weight.shape)

Example 2: Using padding_idx

Specify the padding index:

Example

import torch
import torch.nn as nn

# Create an embedding layer with padding
embedding = nn.Embedding(num_embeddings=1000, embedding_dim=64, padding_idx=0)

# 0 is used as padding
input_indices = torch.tensor([[1, 2, 3], [4, 0, 0]])  # The second sentence has padding

output = embedding(input_indices)

print("Input shape:", input_indices.shape)
print("Output shape:", output.shape)
print("Embedding vector of padding:", output[1, 1].tolist())  # All zeros
print("Embedding vector of non-padding:", output[0, 0].tolist()[:5])  # Non-zero

Example 3: Pretrained Word Vectors

Load pretrained word vectors:

Example

import torch
import torch.nn as nn
import numpy as np

# Simulate pretrained word vectors (in practice, GloVe, Word2Vec, etc. can be used)
vocab_size = 1000
embedding_dim = 300

# Random initialization (in practice, load pretrained vectors)
pretrained_weights = np.random.randn(vocab_size, embedding_dim).astype('float32')

# Create embedding layer
embedding = nn.Embedding(vocab_size, embedding_dim)

# Load pretrained weights
embedding.weight.data = torch.from_numpy(pretrained_weights)

# Freeze the embedding layer (not trained)
embedding.weight.requires_grad = False

print("Embedding layer trainable:", embedding.weight.requires_grad)
print("Embedding matrix shape:", embedding.weight.shape)

Example 4: Constraining Embedding Vector Norm

Use max_norm to constrain the vector norm:

h2 class="example">Example
import torch
import torch.nn as nn

# Constrain the maximum norm of embedding vectors to 1.0
embedding = nn.Embedding(1000, 64, max_norm=1.0)

# Input
input_indices = torch.tensor([1, 2, 3])

# Original weight norm
original_norm = embedding.weight.data.norm(dim=1)[:3]
print("Original weight norm:", original_norm.tolist())

# Vector norm after lookup
output = embedding(input_indices)
output_norm = output.norm(dim=1)
print("Output vector norm:", output_norm.tolist())

Example 5: Complete Text Classification Model

Text classification using an embedding layer:

Example

import torch
import torch.nn as nn

class TextClassifier(nn.Module):
    def __init__(self, vocab_size, embed_dim, hidden_dim, num_classes=2):
        super(TextClassifier, self).__init__()
        # Word embedding layer
        self.embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=0)
        # LSTM
        self.lstm = nn.LSTM(embed_dim, hidden_dim, batch_first=True)
        # Classifier
        self.classifier = nn.Linear(hidden_dim, num_classes)

    def forward(self, x):
        # x: (batch, seq_len)
        embedded = self.embedding(x)  # (batch, seq_len, embed_dim)

        # LSTM takes the last output
        _, (hidden, _) = self.lstm(embedded)
        hidden = hidden[-1]  # Hidden state of the last layer

        # Classification
        logits = self.classifier(hidden)
        return logits

# Parameters
VOCAB_SIZE = 10000
EMBED_DIM = 128
HIDDEN_DIM = 256

model = TextClassifier(VOCAB_SIZE, EMBED_DIM, HIDDEN_DIM)

# Input: batch=4, sequence length=50
input_ids = torch.randint(1, VOCAB_SIZE, (4, 50))
output = model(input_ids)

print("Model architecture:")
print(model)
print("Input shape:", input_ids.shape)
print("Output shape:", output.shape)

Difference Between Embedding and EmbeddingBag

Type Input Output Application Scenarios
nn.Embedding Word indices Sequence of word vectors Sequence models, LSTM, Transformer
nn.EmbeddingBag Word indices + offsets Aggregated vector Text classification, fast processing

Frequently Asked Questions

Q1: How to choose the embedding dimension?

  • Small datasets: 50-100 dimensions
  • Medium datasets: 100-300 dimensions
  • Large datasets: 300-500 dimensions

Q2: What is the purpose of padding_idx?

It sets the embedding vector at the specified index to a zero vector and does not compute its gradient during backpropagation.

Q3: When to freeze the embedding layer?

When using pretrained word vectors, it is common to freeze the layer, train for a while, and then fine-tune.


Use Cases

nn.EmbeddingMain application scenarios include:

  • Word vector representation: Convert words into dense vectors
  • Text classification: As the input layer of NLP models
  • Sequence models: Input for LSTM, GRU
  • Recommendation systems: Embedding representations of users and items

Note: embedding.weight is a learnable parameter and can be trained directly in the optimizer.


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

Other Extensions