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 (Тензор) –
- Возвращает:
-
Выборка тензора той же формы, что и
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