PyTorch torch.nn.EmbeddingBag Function

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

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)

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)

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.


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

Other extensions