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

Функция 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 к отрицательным логарифмам вероятностей
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить