Spec-Zone.ru › PyTorch 2

torch.multinomial

torch.multinomial(input, num_samples, replacement=False, *, generator=None, out=None) → LongTensor

Возвращает тензор, где каждая строка содержит num_samples индексы, выбранные из многомерного вероятностного распределения, расположенного в соответствующей строке тензора input.

Примечание

Строки input не обязательно должны суммироваться до единицы (в этом случае значения используются как веса), но должны быть неотрицательными, конечными и иметь ненулевую сумму.

Индексы упорядочены слева направо в соответствии с моментом их выбора (первые выбранные индексы размещаются в первом столбце).

Если input является вектором, out является вектором размера num_samples.

Если input является матрицей с m строками, out является матрицей формы (m×num_samples)(m \times \text{num\_samples}).

Если замена равна True, образцы выбираются с возвращением.

В противном случае они выбираются без возвращения, что означает, что когда индекс образца выбирается для строки, он не может быть выбран повторно для этой строки.

Примечание

При выборке без возвращения num_samples должно быть меньше, чем количество ненулевых элементов в input (или минимальное количество ненулевых элементов в каждой строке input , если это матрица).

Параметры
  • input (Тензор) – входной тензор, содержащий вероятности
  • num_samples (int) – количество образцов для выборки
  • replacement (bool, необязательно) – нужно ли выбирать с возвращением или без
Ключевые аргументы
  • generator (torch.Generator, необязательно) – генератор псевдослучайных чисел для выборки
  • out (Тензор, необязательно) – выходной тензор.

Пример:

>>> weights = torch.tensor([0, 10, 3, 0], dtype=torch.float) # create a tensor of weights
>>> torch.multinomial(weights, 2)
tensor([1, 2])
>>> torch.multinomial(weights, 4) # ERROR!
RuntimeError: invalid argument 2: invalid multinomial distribution (with replacement=False,
not enough non-negative category to sample) at ../aten/src/TH/generic/THTensorRandom.cpp:320
>>> torch.multinomial(weights, 4, replacement=True)
tensor([ 2,  1,  1,  1])

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

Spec-Zone.ru

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