PyTorch torch.bernoulli Function
PyTorch torch Reference Manual
torch.bernoulliis a function in PyTorch used to generate random numbers from the Bernoulli distribution.
Function Definition
torch.bernoulli(input, *, generator=None, out=None) torch.bernoulli(p, size, *, generator=None, out=None)
Parameter Description
input- The probability value or a tensor containing probabilities (each element represents the probability of the corresponding position being 1)p- The probability value (used when the input is not a tensor)size- The shape of the output tensorgenerator- Random number generator (optional)out- Output tensor (optional)
Usage Example
Example
import torch
# Generate Bernoulli distribution random numbers using a probability tensor
probs = torch.tensor([0.1, 0.5, 0.9])
result = torch.bernoulli(probs)
print("Probability tensor:", probs)
print("Bernoulli sampling result:", result)
# Generate a random tensor using a fixed probability
result2 = torch.bernoulli(0.5, (3, 3))
print("3x3 random tensor generated with fixed probability 0.5:")
print(result2)
# Generate Bernoulli distribution random numbers using a probability tensor
probs = torch.tensor([0.1, 0.5, 0.9])
result = torch.bernoulli(probs)
print("Probability tensor:", probs)
print("Bernoulli sampling result:", result)
# Generate a random tensor using a fixed probability
result2 = torch.bernoulli(0.5, (3, 3))
print("3x3 random tensor generated with fixed probability 0.5:")
print(result2)
Other Extensions