Функция multinomial
Функция multinomial позволяет генерировать выборки из мультиномиального распределения. Первым параметром она принимает тензор весов input, вторым параметром - количество выборок num_samples. Функция возвращает тензор с индексами выбранных категорий. Веса не обязаны быть нормализованными, но должны быть неотрицательными.
Синтаксис
torch.multinomial(input, num_samples, replacement=False, generator=None, out=None)
Основные параметры
Функция принимает следующие основные параметры:
-
input- тензор весов вероятностей для каждой категории. -
num_samples- количество выборок, которое необходимо сгенерировать. -
replacement- булевый флаг, определяющий, разрешено ли повторное использование одной и той же категории (по умолчаниюFalse).
Пример
Давайте сгенерируем выборку из 3 категорий с весами 1, 2, 3:
import torch
torch.manual_seed(0)
weights = torch.tensor([1.0, 2.0, 3.0])
res = torch.multinomial(weights, 1)
print(res)
Результат выполнения кода:
tensor([2])
Пример
Теперь сгенерируем 3 выборки с повторениями из тех же категорий:
import torch
torch.manual_seed(0)
weights = torch.tensor([1.0, 2.0, 3.0])
res = torch.multinomial(weights, 3, replacement=True)
print(res)
Результат выполнения кода:
tensor([2, 0, 2])
Пример
Давайте применим функцию к двумерному тензору весов, чтобы получить выборки для каждого ряда:
import torch
torch.manual_seed(0)
weights = torch.tensor([
[0.1, 0.2, 0.7],
[0.5, 0.3, 0.2],
])
res = torch.multinomial(weights, 2)
print(res)
Результат выполнения кода:
tensor([
[2, 1],
[0, 1],
])
Пример
Сгенерируем выборки с помощью генератора случайных чисел для воспроизводимости:
import torch
generator = torch.Generator()
generator.manual_seed(42)
weights = torch.tensor([0.2, 0.3, 0.5])
res = torch.multinomial(weights, 2, replacement=True, generator=generator)
print(res)
Результат выполнения кода:
tensor([2, 1])