/usr/local/lib64/python3.6/site-packages/torch/distributions/__pycache__
Edit: /usr/local/lib64/python3.6/site-packages/torch/distributions/__pycache__/bernoulli.cpython-36.pyc (4445B)
3
Eg@ @ sd d dl mZ d dlZd dlmZ d dlmZ d dlmZm Z m
Z
mZ d dlm
Z
G dd deZdS ) )NumberN)constraints)ExponentialFamily)
broadcast_allprobs_to_logitslogits_to_probs
lazy_property) binary_cross_entropy_with_logitsc s e Zd ZdZejejdZejZ dZ
dZd" fdd Zd# fdd Z
d
d Zedd
Zedd Zedd Zedd Zedd Zej fddZdd Zdd Zd$ddZedd Zd d! Z ZS )% Bernoullia
Creates a Bernoulli distribution parameterized by :attr:`probs`
or :attr:`logits` (but not both).
Samples are binary (0 or 1). They take the value `1` with probability `p`
and `0` with probability `1 - p`.
Example::
>>> m = Bernoulli(torch.tensor([0.3]))
>>> m.sample() # 30% chance 1; 70% chance 0
tensor([ 0.])
Args:
probs (Number, Tensor): the probability of sampling `1`
logits (Number, Tensor): the log-odds of sampling `1`
)probslogitsTr Nc s |d k|d kkrt d|d k r8t|t}t|\| _nt|t}t|\| _|d k r\| jn| j| _|rrtj }n
| jj }t
t| j||d d S )Nz;Either `probs` or `logits` must be specified, but not both.)
validate_args)
ValueError
isinstancer r r r _paramtorchSizesizesuperr
__init__)selfr r r
Z is_scalarbatch_shape) __class__ I/usr/local/lib64/python3.6/site-packages/torch/distributions/bernoulli.pyr " s
zBernoulli.__init__c sv | j t|}tj|}d| jkr6| jj||_|j|_d| jkrV| jj||_|j|_t t|j
|dd | j|_|S )Nr r F)r
)Z_get_checked_instancer
r r __dict__r expandr r r r _validate_args)r r Z _instancenew)r r r r 2 s
zBernoulli.expandc O s | j j||S )N)r r )r argskwargsr r r _new? s zBernoulli._newc C s | j S )N)r )r r r r meanB s zBernoulli.meanc C s | j d| j S )N )r )r r r r varianceF s zBernoulli.variancec C s t | jddS )NT) is_binary)r r )r r r r r J s zBernoulli.logitsc C s t | jddS )NT)r% )r r )r r r r r N s zBernoulli.probsc C s
| j j S )N)r r )r r r r param_shapeR s zBernoulli.param_shapec
C s0 | j |}tj tj| jj|S Q R X d S )N)Z_extended_shaper Zno_gradZ bernoullir r )r Zsample_shapeshaper r r sampleV s
zBernoulli.samplec C s0 | j r| j| t| j|\}}t||dd S )Nnone) reduction)r Z_validate_sampler r r )r valuer r r r log_prob[ s
zBernoulli.log_probc C s t | j| jddS )Nr) )r* )r r r )r r r r entropya s zBernoulli.entropyc C sH t jd| jj| jjd}|jddt| j }|rD|jd| j }|S ) N )dtypedevicer# )r1 )r# r1 )r1 ) r Zaranger r/ r0 viewlenZ_batch_shaper )r r valuesr r r enumerate_supportd s
zBernoulli.enumerate_supportc C s t j| jd| j fS )Nr# )r logr )r r r r _natural_paramsk s zBernoulli._natural_paramsc C s t jdt j| S )Nr# )r r6 exp)r xr r r _log_normalizero s zBernoulli._log_normalizer)NNN)N)T)__name__
__module____qualname____doc__r Z
unit_intervalrealZarg_constraintsbooleanZsupportZhas_enumerate_supportZ_mean_carrier_measurer r r! propertyr" r$ r r r r&