torch.nn.functional.embedding_bag
-
torch.nn.functional.embedding_bag(input, weight, offsets=None, max_norm=None, norm_type=2, scale_grad_by_freq=False, mode='mean', sparse=False, per_sample_weights=None, include_last_offset=False, padding_idx=None)[source] -
Вычисляет суммы, средние значения или максимумы
bagsэмбеддингов.Вычисления выполняются без создания промежуточных эмбеддингов. Подробнее см. в описании
torch.nn.EmbeddingBag.Примечание
При работе с тензорами на устройстве CUDA эта операция может приводить к недетерминированным градиентам. Подробнее см. в разделе Воспроизводимость.
- Параметры:
-
- input (LongTensor) – Тензор, содержащий наборы индексов в матрице эмбеддингов
- weight (Tensor) – Матрица эмбеддингов, число строк которой равно максимально возможному индексу + 1, а число столбцов равно размеру эмбеддинга
-
offsets (LongTensor, optional) – Используется только тогда, когда
inputявляется одномерным.offsetsзадаёт начальный индекс каждого набора (последовательности) вinput. -
max_norm (float, optional) – Если задано, каждый вектор эмбеддинга с нормой больше
max_normнормализуется до нормыmax_norm. Примечание: это изменитweightна месте. -
norm_type (float, optional) –
max_normдля вычисления нормыp-го порядка при использовании параметраp. По умолчанию2. -
scale_grad_by_freq (bool, optional) – Если задано, градиенты масштабируются обратно пропорционально частоте слов в мини-пакете. По умолчанию
False. Примечание: этот параметр не поддерживается, еслиmode="max". -
mode (str, optional) –
"sum","mean"или"max". Задаёт способ редукции набора. По умолчанию:"mean" -
sparse (bool, optional) – Если
True, градиент поweightбудет разреженным тензором. Подробнее о разреженных градиентах см. примечания кtorch.nn.Embedding. Примечание: этот параметр не поддерживается, еслиmode="max". -
per_sample_weights (Tensor, optional) – тензор весов типа float / double или None, означающее, что все веса равны 1. Если задано,
per_sample_weightsдолжен иметь в точности ту же форму, что и input, и считается имеющим те жеoffsets, если они не равны None. -
include_last_offset (bool, optional) – Если
True, размер offsets равен числу наборов + 1. Последний элемент равен размеру input или конечному индексу последнего набора (последовательности). Это соответствует формату CSR. Игнорируется, если input двумерный. По умолчаниюFalse. -
padding_idx (int, optional) – Если задано, элементы в
padding_idxне участвуют в вычислении градиента; поэтому вектор эмбеддинга вpadding_idxне обновляется во время обучения, то есть остаётся фиксированным «заполнителем». Обратите внимание, что вектор эмбеддинга вpadding_idxисключается из редукции.
- Тип возвращаемого значения:
- Форма:
-
-
input(LongTensor) иoffsets(LongTensor, optional)- Если
inputимеет двумерную форму(B, N), он рассматривается какBнаборов (последовательностей) фиксированной длиныN, и функция возвращаетBзначений, агрегированных в зависимости отmode. В этом случаеoffsetsигнорируется и должен быть равенNone. - Если
inputимеет одномерную форму(N), он рассматривается как конкатенация нескольких наборов (последовательностей).offsetsдолжен быть одномерным тензором, содержащим начальные индексы каждого набора вinput. Поэтому дляoffsetsформы(B)inputбудет рассматриваться как содержащийBнаборов. Для пустых наборов (то есть нулевой длины) возвращаемые векторы будут заполнены нулями.
- Если
-
weight(Tensor): обучаемые веса модуля формы(num_embeddings, embedding_dim) -
per_sample_weights(Tensor, optional). Имеет ту же форму, что иinput. -
output: агрегированные значения эмбеддингов формы(B, embedding_dim)
-
Примеры:
>>> # an Embedding module containing 10 tensors of size 3 >>> embedding_matrix = torch.rand(10, 3) >>> # a batch of 2 samples of 4 indices each >>> input = torch.tensor([1, 2, 4, 5, 4, 3, 2, 9]) >>> offsets = torch.tensor([0, 4]) >>> F.embedding_bag(input, embedding_matrix, offsets) tensor([[ 0.3397, 0.3552, 0.5545], [ 0.5893, 0.4386, 0.5882]]) >>> # example with padding_idx >>> embedding_matrix = torch.rand(10, 3) >>> input = torch.tensor([2, 2, 2, 2, 4, 3, 2, 9]) >>> offsets = torch.tensor([0, 4]) >>> F.embedding_bag(input, embedding_matrix, offsets, padding_idx=2, mode='sum') tensor([[ 0.0000, 0.0000, 0.0000], [-0.7082, 3.2145, -2.6251]])
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/generated/torch.nn.functional.embedding_bag.html