РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
407 of 769 menu

Функция 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], ])

Смотрите также

  • функцию one_hot,
    которая преобразует индексы в one-hot векторы
  • функцию linear,
    которая применяет линейное преобразование к входным данным
  • функцию normalize,
    которая выполняет L2-нормализацию тензора
  • класс nn.Embedding,
    который является обучаемой обёрткой над функцией embedding
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить