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