Spec-Zone.ru › PyTorch 2.14

torch.nn.functional.gumbel_softmax

torch.nn.functional.gumbel_softmax(logits, tau=1, hard=False, eps=1e-10, dim=-1) [исходный код]

Выборка из распределения Gumbel-Softmax (ссылка 1 ссылка 2) с возможностью дискретизации.

Параметры:
  • logits (Tensor) – […, num_features] ненормализованные логарифмы вероятностей
  • tau (float) – неотрицательная скалярная температура
  • hard (bool) – если True, возвращаемые выборки будут дискретизированы в виде one-hot-векторов, но при дифференцировании в autograd будут рассматриваться как мягкая выборка
  • dim (int) – Размерность, вдоль которой будет вычисляться softmax. По умолчанию: -1.
Возвращает:

Тензор выборок той же формы, что и logits, из распределения Gumbel-Softmax. Если hard=True, возвращаемые выборки будут one-hot-векторами, иначе они будут распределениями вероятностей, сумма которых по dim равна 1.

Тип возвращаемого значения:

Tensor

Примечание

Эта функция сохранена для обратной совместимости и в будущем может быть удалена из nn.Functional.

Примечание

Основной приём для hard заключается в том, чтобы выполнить y_hard - y_soft.detach() + y_soft

Это даёт два результата: - выходное значение становится в точности one-hot (поскольку мы прибавляем, а затем вычитаем значение y_soft) - градиент становится равен градиенту y_soft (поскольку мы отсекаем все остальные градиенты)

Примеры::
>>> logits = torch.randn(20, 32)
>>> # Sample soft categorical using reparametrization trick:
>>> F.gumbel_softmax(logits, tau=1, hard=False)
>>> # Sample hard categorical using "Straight-through" trick:
>>> F.gumbel_softmax(logits, tau=1, hard=True)

© 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.gumbel_softmax.html

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API