Функция embedding
Функция embedding извлекает строки из тензора весов,
используя индексы из другого тензора. Первым параметром
она принимает тензор весов, а вторым - тензор индексов.
Результатом является тензор вложений, где каждый индекс
заменён на соответствующую строку из таблицы весов.
Синтаксис
torch.nn.functional.embedding(
input,
weight,
padding_idx=None,
max_norm=None,
norm_type=2.0,
scale_grad_by_freq=False,
sparse=False
)
Пример
Давайте создадим таблицу вложений для 4 слов размерностью 3 и извлечём векторы для индексов 0 и 2:
import torch
import torch.nn.functional as F
weight = torch.tensor([
[1.0, 2.0, 3.0],
[4.0, 5.0, 6.0],
[7.0, 8.0, 9.0],
[10.0, 11.0, 12.0],
])
indices = torch.tensor([0, 2])
res = F.embedding(indices, weight)
print(res)
Результат выполнения кода:
tensor([
[1., 2., 3.],
[7., 8., 9.],
])
Пример
Если передан тензор индексов с размерностью больше 1, то форма результата сохраняется, а последняя ось заменяется на размерность вложений:
import torch
import torch.nn.functional as F
weight = torch.tensor([
[0.1, 0.2],
[0.3, 0.4],
[0.5, 0.6],
])
indices = torch.tensor([
[1, 2],
[0, 1],
])
res = F.embedding(indices, weight)
print(res)
Результат выполнения кода:
tensor([
[
[0.3, 0.4],
[0.5, 0.6],
],
[
[0.1, 0.2],
[0.3, 0.4],
],
])
Пример
Параметр padding_idx позволяет задать индекс, для которого вложения всегда будут нулевыми. Это полезно для обработки переменных последовательностей:
import torch
import torch.nn.functional as F
weight = torch.tensor([
[1.0, 2.0, 3.0],
[4.0, 5.0, 6.0],
[7.0, 8.0, 9.0],
])
indices = torch.tensor([0, 1, 2, 0])
res = F.embedding(indices, weight, padding_idx=0)
print(res)
Результат выполнения кода:
tensor([
[0., 0., 0.],
[4., 5., 6.],
[7., 8., 9.],
[0., 0., 0.],
])
Пример
Параметр max_norm позволяет нормализовать строки весов, чтобы их норма не превышала заданное значение. Это помогает стабилизировать обучение:
import torch
import torch.nn.functional as F
weight = torch.tensor([
[3.0, 4.0],
[1.0, 1.0],
[0.0, 2.0],
], dtype=torch.float)
indices = torch.tensor([0, 1, 2])
res = F.embedding(indices, weight, max_norm=2.0)
print(res)
Результат выполнения кода:
tensor([
[1.2000, 1.6000],
[1.0000, 1.0000],
[0.0000, 2.0000],
])