PyTorch torch.multinomial Function


Pytorch torch 参考手册PyTorch torch Reference Manual

torch.multinomialIs a function in PyTorch used for multinomial sampling, sampling from each row of the input tensor according to the probability distribution.

Function Definition

torch.multinomial(input, num_samples, replacement=False, generator=None)

Parameter Description

  • input- Probability distribution tensor, each row is a probability distribution, and the sum of all probabilities should be 1
  • num_samples- Number of samples
  • replacement- Whether to allow sampling with replacement (default False)
  • generator- Random number generator (optional)

Usage Example

Example

import torch

# Define probability distribution
weights = torch.tensor([[0.0, 1.0],   # The probability of the second category is 1
                        [0.5, 0.5],   # The two categories have equal probabilities
                        [0.2, 0.3, 0.5]])  # Three categories

# Sampling without replacement
result = torch.multinomial(weights, num_samples=2, replacement=False)
print("Probability distribution:")
print(weights)
print("Sampling without replacement results (2 samples per row):")
print(result)

# Sampling with replacement
result_with_replacement = torch.multinomial(weights, num_samples=5, replacement=True)
print("Sampling with replacement results (5 samples per row):")
print(result_with_replacement)

Pytorch torch 参考手册PyTorch torch Reference Manual

Other Extensions