PyTorch torch.multinomial Function
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 1num_samples- Number of samplesreplacement- 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)
# 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)
Other Extensions