Spec-Zone.ru › PyTorch 1

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/1.13/generated/torch.nn.functional.gumbel_softmax.html

Spec-Zone.ru

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