/usr/local/lib64/python3.6/site-packages/torch/distributions/__pycache__
Edit: /usr/local/lib64/python3.6/site-packages/torch/distributions/__pycache__/binomial.cpython-36.pyc (4761B)
3
Eg; @ sT d dl Z d dlmZ d dlmZ d dlmZmZmZm Z dd Z
G dd deZdS ) N)constraints)Distribution)
broadcast_allprobs_to_logits
lazy_propertylogits_to_probsc C s | j dd| | j dd d S )Nr )min)max )clamp)x r
H/usr/local/lib64/python3.6/site-packages/torch/distributions/binomial.py_clamp_by_zero s r c s e Zd ZdZejejejdZdZ d fdd Z
d! fdd Zd
d Zej
ddd
dd Zedd Zedd Zedd Zedd Zedd Zej fddZdd Zd"ddZ ZS )#Binomiala
Creates a Binomial distribution parameterized by :attr:`total_count` and
either :attr:`probs` or :attr:`logits` (but not both). :attr:`total_count` must be
broadcastable with :attr:`probs`/:attr:`logits`.
Example::
>>> m = Binomial(100, torch.tensor([0 , .2, .8, 1]))
>>> x = m.sample()
tensor([ 0., 22., 71., 100.])
>>> m = Binomial(torch.tensor([[5.], [10.]]), torch.tensor([0.5, 0.8]))
>>> x = m.sample()
tensor([[ 4., 5.],
[ 7., 6.]])
Args:
total_count (int or Tensor): number of Bernoulli trials
probs (Tensor): Event probabilities
logits (Tensor): Event log-odds
)total_countprobslogitsT Nc s |d k|d kkrt d|d k rDt||\| _| _| jj| j| _n"t||\| _| _| jj| j| _|d k rt| jn| j| _| jj }tt | j
||d d S )Nz;Either `probs` or `logits` must be specified, but not both.)
validate_args)
ValueErrorr r r Ztype_asr _paramsizesuperr __init__)selfr r r r batch_shape) __class__r
r r ' s
zBinomial.__init__c s | j t|}tj|}| jj||_d| jkrD| jj||_|j|_d| jkrd| j j||_ |j |_t
t|j|dd | j|_|S )Nr r F)r )
Z_get_checked_instancer torchSizer expand__dict__r r r r r _validate_args)r r Z _instancenew)r r
r r 5 s
zBinomial.expandc O s | j j||S )N)r r# )r argskwargsr
r
r _newC s z
Binomial._newr )Zis_discreteZ event_dimc C s t jd| jS )Nr )r Zinteger_intervalr )r r
r
r supportF s zBinomial.supportc C s | j | j S )N)r r )r r
r
r meanJ s z
Binomial.meanc C s | j | j d| j S )Nr )r r )r r
r
r varianceN s zBinomial.variancec C s t | jddS )NT) is_binary)r r )r r
r
r r R s zBinomial.logitsc C s t | jddS )NT)r* )r r )r r
r
r r V s zBinomial.probsc C s
| j j S )N)r r )r r
r
r param_shapeZ s zBinomial.param_shapec C s: | j |}tj tj| jj|| jj|S Q R X d S )N)Z_extended_shaper Zno_gradZbinomialr r r )r Zsample_shapeshaper
r
r sample^ s
zBinomial.samplec C s | j r| j| tj| jd }tj|d }tj| j| d }| jt| j | jtjtjtj | j | }|| j | | | S )Nr )
r" Z_validate_sampler lgammar r r log1pexpabs)r valueZlog_factorial_nZlog_factorial_kZlog_factorial_nmkZnormalize_termr
r
r log_probc s
4zBinomial.log_probc C sp t | jj }| jj |ks$tdtjd| | jj| jj d}|j
ddt| j }|rl|j
d| j }|S ) Nz?Inhomogeneous total count not supported by `enumerate_support`.r )dtypedevice)r6 )r r6 )r6 )intr r r NotImplementedErrorr Zaranger r4 r5 viewlenZ_batch_shaper )r r r valuesr
r
r enumerate_supports s zBinomial.enumerate_support)r NNN)N)T)__name__
__module____qualname____doc__r Znonnegative_integerZ
unit_intervalrealZarg_constraintsZhas_enumerate_supportr r r&