Функция gumbel_softmax
Функция gumbel_softmax применяет Gumbel-Softmax
трюк для получения дифференцируемой выборки из дискретного
распределения. Первым параметром функция принимает логарифмы
вероятностей (logits) для каждой категории. Вторым параметром
передаётся температура, которая контролирует степень дискретности
выборки. Третьим параметром можно указать, использовать ли
жёсткий (hard) или мягкий (soft) вариант выборки.
Синтаксис
torch.nn.functional.gumbel_softmax(logits, tau=1.0, hard=False, dim=-1)
Параметры
Функция принимает следующие параметры:
logits- тензор необработанных логарифмов вероятностей для каждой категорииtau(температура) - скалярное значение, контролирующее дискретность выборки. Приtau -> 0выборка становится дискретной, приtau -> ∞приближается к равномерному распределениюhard- булево значение, указывающее на использование жёсткой выборки. ЕслиTrue, функция возвращает one-hot векторыdim- ось, вдоль которой применяется softmax (по умолчанию последняя ось)
Пример с мягкой выборкой
Давайте создадим логарифмы вероятностей и применим мягкий Gumbel-Softmax:
import torch
import torch.nn.functional as F
torch.manual_seed(0)
logits = torch.tensor([1.0, 2.0, 3.0])
t = F.gumbel_softmax(logits, tau=1.0, hard=False)
print(t)
Результат выполнения кода:
tensor([0.0944, 0.2607, 0.6449])
Как видим, функция возвращает распределение вероятностей, которое является дифференцируемой аппроксимацией one-hot вектора.
Пример с жёсткой выборкой
Давайте получим жёсткую выборку с помощью параметра hard=True:
import torch
import torch.nn.functional as F
torch.manual_seed(0)
logits = torch.tensor([1.0, 2.0, 3.0])
t = F.gumbel_softmax(logits, tau=1.0, hard=True)
print(t)
Результат выполнения кода:
tensor([0., 0., 1.])
Функция возвращает one-hot вектор, где выбранная категория отмечена единицей. При этом градиенты вычисляются через мягкую аппроксимацию, что делает выборку дифференцируемой.
Пример с разными температурами
Давайте сравним влияние разных температур на распределение выборки:
import torch
import torch.nn.functional as F
torch.manual_seed(0)
logits = torch.tensor([1.0, 2.0, 3.0])
t1 = F.gumbel_softmax(logits, tau=1.0)
t2 = F.gumbel_softmax(logits, tau=2.0)
t3 = F.gumbel_softmax(logits, tau=0.5)
print("tau=1.0:", t1)
print("tau=2.0:", t2)
print("tau=0.5:", t3)
Результат выполнения кода:
tau=1.0: tensor([0.0944, 0.2607, 0.6449])
tau=2.0: tensor([0.2388, 0.3160, 0.4452])
tau=0.5: tensor([0.0192, 0.1198, 0.8610])
При уменьшении температуры (tau=0.5) распределение становится
более дискретным, приближаясь к one-hot вектору. При увеличении
температуры (tau=2.0) распределение становится более сглаженным.
Пример с пакетом данных
Давайте применим функцию к пакету логарифмов вероятностей:
import torch
import torch.nn.functional as F
torch.manual_seed(0)
logits = torch.tensor([
[0.5, 1.0, 2.0],
[3.0, 0.5, 1.5],
[1.0, 2.0, 0.5]
])
t = F.gumbel_softmax(logits, tau=1.0, hard=False)
print(t)
Результат выполнения кода:
tensor([
[0.1025, 0.1646, 0.7329],
[0.8388, 0.0549, 0.1063],
[0.1434, 0.7598, 0.0968]
])
Функция применяется независимо к каждой строке пакета, возвращая распределение вероятностей для каждого примера.
Смотрите также
-
функцию
softmax,
которая преобразует логарифмы вероятностей в вероятности -
функцию
log_softmax,
которая применяет softmax с последующим логарифмированием -
функцию
one_hot,
которая преобразует индексы в one-hot представление -
функцию
softmin,
которая применяет softmax к отрицательным логарифмам вероятностей