PyTorch torch.bernoulli Function


Pytorch torch 参考手册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 tensor
  • generator- 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)

Pytorch torch 参考手册PyTorch torch Reference Manual

Other Extensions