Ë
    f^(hÇ  ã                   ó€   — d dl Z d dl mZmZ d dlmZ d dlmZ d dlmZm	Z	m
Z
mZ d dlmZ d dlmZ dgZ G d	„ de«      Zy)
é    N)ÚnanÚTensor)Úconstraints)ÚExponentialFamily)Úbroadcast_allÚlazy_propertyÚlogits_to_probsÚprobs_to_logits)Ú binary_cross_entropy_with_logits)Ú_NumberÚ	Bernoullic                   ó~  ‡ — e Zd ZdZej
                  ej                  dœZej                  Z	dZ
dZdˆ fd„	Zdˆ fd„	Zd„ Zedefd	„«       Zedefd
„«       Zedefd„«       Zedefd„«       Zedefd„«       Zedej0                  fd„«       Z ej0                  «       fd„Zd„ Zd„ Zdd„Zedee   fd„«       Zd„ Z ˆ xZ!S )r   a1  
    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::

        >>> # xdoctest: +IGNORE_WANT("non-deterministic")
        >>> 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`
    )ÚprobsÚlogitsTr   c                 ó~  •— |d u |d u k(  rt        d«      ‚|�#t        |t        «      }t        |«      \  | _        n"t        |t        «      }t        |«      \  | _        |�| j                  n| j
                  | _        |rt        j                  «       }n| j                  j                  «       }t        ‰| �-  ||¬«       y )Nz;Either `probs` or `logits` must be specified, but not both.©Úvalidate_args)Ú
ValueErrorÚ
isinstancer   r   r   r   Ú_paramÚtorchÚSizeÚsizeÚsuperÚ__init__)Úselfr   r   r   Ú	is_scalarÚbatch_shapeÚ	__class__s         €ú[/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/torch/distributions/bernoulli.pyr   zBernoulli.__init__,   s§   ø€ Ø�TˆM˜v¨˜~Ò.ÜØMóð ð ÐÜ" 5¬'Ó2ˆIÜ)¨%Ó0‰MˆT�Zä" 6¬7Ó3ˆIÜ*¨6Ó2‰NˆTŒ[Ø$)Ð$5�d—j’j¸4¿;¹;ˆŒÙÜŸ*™*›,‰KàŸ+™+×*Ñ*Ó,ˆKÜ‰Ñ˜°MÐÕBó    c                 ó¦  •— | j                  t        |«      }t        j                  |«      }d| j                  v r1| j
                  j                  |«      |_        |j
                  |_        d| j                  v r1| j                  j                  |«      |_        |j                  |_        t        t        |�+  |d¬«       | j                  |_        |S )Nr   r   Fr   )Ú_get_checked_instancer   r   r   Ú__dict__r   Úexpandr   r   r   r   Ú_validate_args)r   r   Ú	_instanceÚnewr   s       €r    r%   zBernoulli.expand>   s¤   ø€ Ø×(Ñ(¬°IÓ>ˆÜ—j‘j Ó-ˆØ�d—m‘mÑ#ØŸ
™
×)Ñ)¨+Ó6ˆCŒIØŸ™ˆCŒJØ�t—}‘}Ñ$ØŸ™×+Ñ+¨KÓ8ˆCŒJØŸ™ˆCŒJÜŒi˜Ñ& {À%Ð&ÔHØ!×0Ñ0ˆÔØˆ
r!   c                 ó:   —  | j                   j                  |i |¤ŽS ©N)r   r(   )r   ÚargsÚkwargss      r    Ú_newzBernoulli._newK   s   € Øˆt�{‰{�‰ Ð/¨Ñ/Ð/r!   Úreturnc                 ó   — | j                   S r*   ©r   ©r   s    r    ÚmeanzBernoulli.meanN   s   € à�z‰zÐr!   c                 ó‚   — | j                   dk\  j                  | j                   «      }t        || j                   dk(  <   |S )Ng      à?)r   Útor   )r   Úmodes     r    r5   zBernoulli.modeR   s7   € à—
‘
˜cÑ!×%Ñ% d§j¡jÓ1ˆÜ"%ˆˆT�Z‰Z˜3ÑÑØˆr!   c                 ó:   — | j                   d| j                   z
  z  S )Né   r0   r1   s    r    ÚvariancezBernoulli.varianceX   s   € à�z‰z˜Q §¡™^Ñ,Ð,r!   c                 ó0   — t        | j                  d¬«      S ©NT)Ú	is_binary)r
   r   r1   s    r    r   zBernoulli.logits\   s   € ä˜tŸz™z°TÔ:Ð:r!   c                 ó0   — t        | j                  d¬«      S r:   )r	   r   r1   s    r    r   zBernoulli.probs`   s   € ä˜tŸ{™{°dÔ;Ð;r!   c                 ó6   — | j                   j                  «       S r*   )r   r   r1   s    r    Úparam_shapezBernoulli.param_shaped   s   € à�{‰{×ÑÓ!Ð!r!   c                 óÔ   — | j                  |«      }t        j                  «       5  t        j                  | j                  j                  |«      «      cd d d «       S # 1 sw Y   y xY wr*   )Ú_extended_shaper   Úno_gradÚ	bernoullir   r%   )r   Úsample_shapeÚshapes      r    ÚsamplezBernoulli.sampleh   sK   € Ø×$Ñ$ \Ó2ˆÜ�]‰]‹_ñ 	=Ü—?‘? 4§:¡:×#4Ñ#4°UÓ#;Ó<÷	=÷ 	=ò 	=ús   ¦.AÁA'c                 óŒ   — | j                   r| j                  |«       t        | j                  |«      \  }}t	        ||d¬«       S ©NÚnone)Ú	reduction)r&   Ú_validate_sampler   r   r   )r   Úvaluer   s      r    Úlog_probzBernoulli.log_probm   s?   € Ø×ÒØ×!Ñ! %Ô(Ü% d§k¡k°5Ó9‰ˆ�Ü0°¸È&ÔQÐQÐQr!   c                 óF   — t        | j                  | j                  d¬«      S rG   )r   r   r   r1   s    r    ÚentropyzBernoulli.entropys   s   € Ü/Ø�K‰K˜Ÿ™¨vô
ð 	
r!   c                 ó  — t        j                  d| j                  j                  | j                  j                  ¬«      }|j                  ddt        | j                  «      z  z   «      }|r|j                  d| j                  z   «      }|S )Né   )ÚdtypeÚdevice)éÿÿÿÿ)r7   )	r   Úaranger   rQ   rR   ÚviewÚlenÚ_batch_shaper%   )r   r%   Úvaluess      r    Úenumerate_supportzBernoulli.enumerate_supportx   sl   € Ü—‘˜a t§{¡{×'8Ñ'8ÀÇÁ×ASÑASÔTˆØ—‘˜U T¬C°×0AÑ0AÓ,BÑ%BÑBÓCˆÙØ—]‘] 5¨4×+<Ñ+<Ñ#<Ó=ˆFØˆr!   c                 óB   — t        j                  | j                  «      fS r*   )r   Úlogitr   r1   s    r    Ú_natural_paramszBernoulli._natural_params   s   € ä—‘˜DŸJ™JÓ'Ð)Ð)r!   c                 óR   — t        j                  t        j                  |«      «      S r*   )r   Úlog1pÚexp)r   Úxs     r    Ú_log_normalizerzBernoulli._log_normalizerƒ   s   € Ü�{‰{œ5Ÿ9™9 Q›<Ó(Ð(r!   )NNNr*   )T)"Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   Úunit_intervalÚrealÚarg_constraintsÚbooleanÚsupportÚhas_enumerate_supportÚ_mean_carrier_measurer   r%   r-   Úpropertyr   r2   r5   r8   r   r   r   r   r   r>   rE   rL   rN   rY   Útupler\   ra   Ú__classcell__)r   s   @r    r   r      s3  ø„ ñð& !,× 9Ñ 9À[×EUÑEUÑV€OØ×!Ñ!€GØ ÐØÐõCõ$ò0ð ð�fò ó ðð ð�fò ó ðð
 ð-˜&ò -ó ð-ð ð;˜ò ;ó ð;ð ð<�vò <ó ð<ð ð"˜UŸZ™Zò "ó ð"ð #- %§*¡*£,ó =ò
Rò
ó
ð ð*  v¡ò *ó ð*ö)r!   )r   r   r   Útorch.distributionsr   Útorch.distributions.exp_familyr   Útorch.distributions.utilsr   r   r	   r
   Útorch.nn.functionalr   Útorch.typesr   Ú__all__r   © r!   r    ú<module>rw      s<   ðã ß Ý +Ý <÷ó õ AÝ ð ˆ-€ôq)Ð!õ q)r!   