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 (Tensor) –
- Возвращает:
-
Тензор выборок той же формы, что и
logits, из распределения Gumbel-Softmax. Еслиhard=True, возвращаемые выборки будут one-hot-векторами, иначе они будут распределениями вероятностей, сумма которых поdimравна 1. - Тип возвращаемого значения:
Примечание
Эта функция сохранена для обратной совместимости и в будущем может быть удалена из 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