Ë
    *êñiš  ã                   ó„   — 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)
é    N)Ú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z  dz  deez  dz  dedz  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j4                  fd„«       Z ej4                  «       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/|€t        d«      ‚t        |t        «      }t        |«      \  | _        |�| j                  n| j                  | _        |rt        j                  «       }n| j                  j                  «       }t        ‰| �1  ||¬«       y )Nz;Either `probs` or `logits` must be specified, but not both.zlogits is unexpectedly None©r   )Ú
ValueErrorÚ
isinstancer   r   r   ÚAssertionErrorr   Ú_paramÚtorchÚSizeÚsizeÚsuperÚ__init__)Úselfr   r   r   Ú	is_scalarÚbatch_shapeÚ	__class__s         €ú_/var/www/pod-logistic/pod-ai/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àˆ~Ü$Ð%BÓCÐCÜ" 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.expandJ   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._newW   s   € Øˆt�{‰{�‰ Ð/¨Ñ/Ð/r$   c                 ó   — | j                   S r-   ©r   ©r   s    r#   ÚmeanzBernoulli.meanZ   s   € à�z‰zÐr$   c                 ó‚   — | j                   dk\  j                  | j                   «      }t        || j                   dk(  <   |S )Ng      à?)r   Útor   )r   Úmodes     r#   r7   zBernoulli.mode^   s7   € à—
‘
˜cÑ!×%Ñ% d§j¡jÓ1ˆÜ"%ˆˆT�Z‰Z˜3ÑÑØˆr$   c                 ó:   — | j                   d| j                   z
  z  S )Né   r2   r3   s    r#   ÚvariancezBernoulli.varianced   s   € à�z‰z˜Q §¡™^Ñ,Ð,r$   c                 ó0   — t        | j                  d¬«      S ©NT)Ú	is_binary)r
   r   r3   s    r#   r   zBernoulli.logitsh   s   € ä˜tŸz™z°TÔ:Ð:r$   c                 ó0   — t        | j                  d¬«      S r<   )r	   r   r3   s    r#   r   zBernoulli.probsl   s   € ä˜tŸ{™{°dÔ;Ð;r$   c                 ó6   — | j                   j                  «       S r-   )r   r   r3   s    r#   Úparam_shapezBernoulli.param_shapep   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.samplet   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_proby   s?   € Ø×ÒØ×!Ñ! %Ô(Ü% d§k¡k°5Ó9‰ˆ�Ü0°¸È&ÔQÐQÐQr$   c                 óF   — t        | j                  | j                  d¬«      S rI   )r   r   r   r3   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)éÿÿÿÿ)r9   )	r   Úaranger   rS   rT   Ú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   r3   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   Úboolr   r(   r0   Úpropertyr4   r7   r:   r   r   r   r   r   r@   rG   rN   rP   r[   Útupler^   rc   Ú__classcell__)r"   s   @r#   r   r      s~  ø„ ñð* !,× 9Ñ 9À[×EUÑEUÑV€OØ×!Ñ!€GØ ÐØÐð )-Ø)-Ø%)ñ	Cà˜‰ Ñ%ðCð ˜‘ $Ñ&ðCð ˜d‘{ð	Cð
 
õCõ6ò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   r   Ú__all__r   © r$   r#   ú<module>rz      s>   ðó ß Ý +Ý <÷ó õ Aß 'ð ˆ-€ô})Ð!õ })r$   