Spec-Zone.ru › PyTorch 2

torch.nn.functional.gumbel_softmax

torch.nn.functional.gumbel_softmax(logits, tau=1, hard=False, eps=1e-10, dim=-1) [source]

Выборка из распределения Гамбеля-Софтмакса (Ссылка 1 Ссылка 2) и, необязательно, дискретизация.

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

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

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

Тензор

Примечание

Эта функция здесь по причинам совместимости, может быть удалена из 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)

© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.nn.functional.gumbel_softmax.html

Spec-Zone.ru

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