Ë
    óÍ:j  ã                   ó”   — d dl mZmZ 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mZ d	gZ G d
„ d	e	«      Zy)é    )ÚOptionalÚUnionN)ÚnanÚTensor)Úconstraints)ÚExponentialFamily)Úbroadcast_allÚlazy_propertyÚlogits_to_probsÚprobs_to_logits)Ú binary_cross_entropy_with_logits)Ú_NumberÚNumberÚ	Bernoullic            	       ó¼  ‡ — e Zd ZdZej
                  ej                  dœZej                  Z	dZ
dZ	 	 	 ddeeeef      deeeef      dee   d	dfˆ 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j8                  fd„«       Z ej8                  «       fd„Zd„ Zd„ Z dd„Z!ed	e"e   fd„«       Z#d„ Z$ˆ xZ%S )r   aˆ  
    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`
        validate_args (bool, optional): whether to validate arguments, None by default
    )ÚprobsÚlogitsTr   Nr   r   Úvalidate_argsÚreturnc                 ó†  •— |d u |d u k(  rt        d«      ‚|�#t        |t        «      }t        |«      \  | _        n&|€J ‚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.©r   )Ú
ValueErrorÚ
isinstancer   r	   r   r   Ú_paramÚtorchÚSizeÚsizeÚsuperÚ__init__)Úselfr   r   r   Ú	is_scalarÚbatch_shapeÚ	__class__s         €úr/home/mcse/projects/srt_converter/srt-converter-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.expandG   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._newT   s   € Øˆt�{‰{�‰ Ð/¨Ñ/Ð/r%   c                 ó   — | j                   S r.   ©r   ©r    s    r$   ÚmeanzBernoulli.meanW   s   € à�z‰zÐr%   c                 ó‚   — | j                   dk\  j                  | j                   «      }t        || j                   dk(  <   |S )Ng      à?)r   Útor   )r    Úmodes     r$   r8   zBernoulli.mode[   s7   € à—
‘
˜cÑ!×%Ñ% d§j¡jÓ1ˆÜ"%ˆˆT�Z‰Z˜3ÑÑØˆr%   c                 ó:   — | j                   d| j                   z
  z  S )Né   r3   r4   s    r$   ÚvariancezBernoulli.variancea   s   € à�z‰z˜Q §¡™^Ñ,Ð,r%   c                 ó0   — t        | j                  d¬«      S ©NT)Ú	is_binary)r   r   r4   s    r$   r   zBernoulli.logitse   s   € ä˜tŸz™z°TÔ:Ð:r%   c                 ó0   — t        | j                  d¬«      S r=   )r   r   r4   s    r$   r   zBernoulli.probsi   s   € ä˜tŸ{™{°dÔ;Ð;r%   c                 ó6   — | j                   j                  «       S r.   )r   r   r4   s    r$   Úparam_shapezBernoulli.param_shapem   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.sampleq   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_probv   s?   € Ø×ÒØ×!Ñ! %Ô(Ü% d§k¡k°5Ó9‰ˆ�Ü0°¸È&ÔQÐQÐQr%   c                 óF   — t        | j                  | j                  d¬«      S rJ   )r   r   r   r4   s    r$   ÚentropyzBernoulli.entropy|   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)éÿÿÿÿ)r:   )	r   Úaranger   rT   rU   ÚviewÚlenÚ_batch_shaper)   )r    r)   Úvaluess      r$   Úenumerate_supportzBernoulli.enumerate_support�   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   r4   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   r   Úboolr   r)   r1   Úpropertyr5   r8   r;   r
   r   r   r   r   rA   rH   rO   rQ   r\   Útupler_   rd   Ú__classcell__)r#   s   @r$   r   r      sˆ  ø„ ñð( !,× 9Ñ 9À[×EUÑEUÑV€OØ×!Ñ!€GØ ÐØÐð 26Ø26Ø(,ñ	Cà˜˜f f˜nÑ-Ñ.ðCð ˜˜v v˜~Ñ.Ñ/ðCð   ‘~ð	Cð
 
õCõ0ò0ð ð�fò ó ðð ð�fò ó ðð
 ð-˜&ò -ó ð-ð ð;˜ò ;ó ð;ð ð<�vò <ó ð<ð ð"˜UŸZ™Zò "ó ð"ð #- %§*¡*£,ó =ò
Rò
ó
ð ð*  v¡ò *ó ð*ö)r%   )Útypingr   r   r   r   r   Útorch.distributionsr   Útorch.distributions.exp_familyr   Útorch.distributions.utilsr	   r
   r   r   Útorch.nn.functionalr   Útorch.typesr   r   Ú__all__r   © r%   r$   ú<module>r|      s?   ðç "ã ß Ý +Ý <÷ó õ Aß 'ð ˆ-€ôx)Ð!õ x)r%   