PyTorch torch.nn.EmbeddingBag Function
PyTorch torch.nn Reference Manual
torch.nn.EmbeddingBagIt is an embedding bag module in PyTorch.
It aggregates multiple embedding vectors into a single vector, commonly used for text classification and fast processing.
Function Definition
torch.nn.EmbeddingBag(num_embeddings, embedding_dim, max_norm=None, norm_type=2.0, scale_grad_by_freq=False, mode='mean', sparse=False, include_last_offset=False)
Parameters:
mode: Aggregation method, optional'mean'、'sum'、'max'
Usage Examples
Example 1: Basic Usage
Example
import torch
import torch.nn as nn
# Embedding bag
ebag = nn.EmbeddingBag(1000, 128, mode='mean')
# Word indices
indices = torch.tensor([[1, 2, 3], [4, 5]])
# Offsets (indicating the boundaries of each sentence)
offsets = torch.tensor([0, 3])
output = ebag(indices, offsets)
print("Input shape:", indices.shape)
print("Output shape:", output.shape)
import torch.nn as nn
# Embedding bag
ebag = nn.EmbeddingBag(1000, 128, mode='mean')
# Word indices
indices = torch.tensor([[1, 2, 3], [4, 5]])
# Offsets (indicating the boundaries of each sentence)
offsets = torch.tensor([0, 3])
output = ebag(indices, offsets)
print("Input shape:", indices.shape)
print("Output shape:", output.shape)
Example 2: Fast Text Classification
Example
import torch
import torch.nn as nn
class FastText(nn.Module):
def __init__(self, vocab_size, embed_dim, num_classes):
super(FastText, self).__init__()
self.embedding = nn.EmbeddingBag(vocab_size, embed_dim, mode='mean')
self.fc = nn.Linear(embed_dim, num_classes)
def forward(self, text, offset):
x = self.embedding(text, offset)
return self.fc(x)
model = FastText(10000, 128, 2)
# Simulate data
text = torch.randint(0, 10000, (100,))
offset = torch.tensor([0, 30, 60, 100])
output = model(text, offset)
print("Output shape:", output.shape)
import torch.nn as nn
class FastText(nn.Module):
def __init__(self, vocab_size, embed_dim, num_classes):
super(FastText, self).__init__()
self.embedding = nn.EmbeddingBag(vocab_size, embed_dim, mode='mean')
self.fc = nn.Linear(embed_dim, num_classes)
def forward(self, text, offset):
x = self.embedding(text, offset)
return self.fc(x)
model = FastText(10000, 128, 2)
# Simulate data
text = torch.randint(0, 10000, (100,))
offset = torch.tensor([0, 30, 60, 100])
output = model(text, offset)
print("Output shape:", output.shape)
Example 3: Different Aggregation Methods
Example
import torch
import torch.nn as nn
for mode in ['mean', 'sum', 'max']:
ebag = nn.EmbeddingBag(100, 32, mode=mode)
indices = torch.tensor([1, 2, 3, 4])
offsets = torch.tensor([0, 2])
out = ebag(indices, offsets)
print(f"{mode}:", out.shape)
import torch.nn as nn
for mode in ['mean', 'sum', 'max']:
ebag = nn.EmbeddingBag(100, 32, mode=mode)
indices = torch.tensor([1, 2, 3, 4])
offsets = torch.tensor([0, 2])
out = ebag(indices, offsets)
print(f"{mode}:", out.shape)
Use Cases
- Text classification: FastText
- Bag-of-words model: Fast aggregation
- Large-scale text: Efficient processing
Note: The offsets parameter must be provided to indicate sentence boundaries.
Other extensions