ó
    Eñiš  ã                   ó„   • S SK r S SK JrJr  S SKJr  S SKJr  S SKJrJ	r	J
r
Jr  S SKJr  S SKJrJr  S/r " S	 S\5      rg)
é    N)ÚnanÚTensor)Úconstraints)ÚExponentialFamily)Úbroadcast_allÚlazy_propertyÚlogits_to_probsÚprobs_to_logits)Ú binary_cross_entropy_with_logits)Ú_NumberÚNumberÚ	Bernoullic            	       óØ  ^ • \ rS rSrSr\R                  \R                  S.r\R                  r
SrSr   SS\\-  S-  S\\-  S-  S	\S-  S
S4U 4S jjjrSU 4S jjrS r\S
\4S j5       r\S
\4S j5       r\S
\4S j5       r\S
\4S j5       r\S
\4S j5       r\S
\R6                  4S j5       r\R6                  " 5       4S jrS rS rSS jr \S
\!\   4S j5       r"S r#Sr$U =r%$ )r   é   aP  
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                 ó¤  >• US L US L :X  a  [        S5      eUb#  [        U[        5      n[        U5      u  U l        O0Uc  [        S5      e[        U[        5      n[        U5      u  U l        Ub  U R                  OU R                  U l        U(       a  [        R                  " 5       nOU R                  R                  5       n[        TU ]1  XSS9  g )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         €ÚZ/home/mande/repo/quber/.venv/lib/python3.13/site-packages/torch/distributions/bernoulli.pyr   Ú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Ü‰Ñ˜ÐÒBó    c                 óª  >• U R                  [        U5      n[        R                  " U5      nSU R                  ;   a1  U R
                  R                  U5      Ul        UR
                  Ul        SU R                  ;   a1  U R                  R                  U5      Ul        UR                  Ul        [        [        U]+  USS9  U R                  Ul        U$ )Nr   r   Fr   )Ú_get_checked_instancer   r   r   Ú__dict__r   Úexpandr   r   r   r   Ú_validate_args)r    r"   Ú	_instanceÚnewr#   s       €r$   r*   Ú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                 ó:   • U R                   R                  " U0 UD6$ ©N)r   r-   )r    ÚargsÚkwargss      r$   Ú_newÚBernoulli._newW   s   € Ø�{‰{�Š Ð/¨Ñ/Ð/r&   c                 ó   • U R                   $ r0   ©r   ©r    s    r$   ÚmeanÚBernoulli.meanZ   s   € à�z‰zÐr&   c                 ó€   • U R                   S:¬  R                  U R                   5      n[        XR                   S:H  '   U$ )Ng      à?)r   Útor   )r    Úmodes     r$   r<   ÚBernoulli.mode^   s5   € à—
‘
˜cÑ!×%Ñ% d§j¡jÓ1ˆÜ"%ˆ�Z‰Z˜3ÑÑØˆr&   c                 ó:   • U R                   SU R                   -
  -  $ )Né   r6   r7   s    r$   ÚvarianceÚBernoulli.varianced   s   € à�z‰z˜Q §¡™^Ñ,Ð,r&   c                 ó*   • [        U R                  SS9$ ©NT)Ú	is_binary)r
   r   r7   s    r$   r   ÚBernoulli.logitsh   s   € ä˜tŸz™z°TÑ:Ð:r&   c                 ó*   • [        U R                  SS9$ rC   )r	   r   r7   s    r$   r   ÚBernoulli.probsl   s   € ä˜tŸ{™{°dÑ;Ð;r&   c                 ó6   • U R                   R                  5       $ r0   )r   r   r7   s    r$   Úparam_shapeÚBernoulli.param_shapep   s   € à�{‰{×ÑÓ!Ð!r&   c                 óâ   • U R                  U5      n[        R                  " 5          [        R                  " U R                  R                  U5      5      sS S S 5        $ ! , (       d  f       g = fr0   )Ú_extended_shaper   Úno_gradÚ	bernoullir   r*   )r    Úsample_shapeÚshapes      r$   ÚsampleÚBernoulli.samplet   s@   € Ø×$Ñ$ \Ó2ˆÜ�]Š]�_Ü—?’? 4§:¡:×#4Ñ#4°UÓ#;Ó<÷ �_�_ús   §/A Á 
A.c                 óŒ   • U R                   (       a  U R                  U5        [        U R                  U5      u  p![	        X!SS9* $ ©NÚnone)Ú	reduction)r+   Ú_validate_sampler   r   r   )r    Úvaluer   s      r$   Úlog_probÚBernoulli.log_proby   s;   € Ø××Ø×!Ñ! %Ô(Ü% d§k¡k°5Ó9‰ˆÜ0°È&ÑQÐQÐQr&   c                 ó@   • [        U R                  U R                  SS9$ rT   )r   r   r   r7   s    r$   ÚentropyÚBernoulli.entropy   s   € Ü/Ø�K‰K˜Ÿ™¨vñ
ð 	
r&   c                 ó   • [         R                  " SU R                  R                  U R                  R                  S9nUR                  SS[        U R                  5      -  -   5      nU(       a  UR                  SU R                  -   5      nU$ )Né   )ÚdtypeÚdevice)éÿÿÿÿ)r?   )	r   Úaranger   r`   ra   ÚviewÚlenÚ_batch_shaper*   )r    r*   Úvaluess      r$   Úenumerate_supportÚ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                 óD   • [         R                  " U R                  5      4$ r0   )r   Úlogitr   r7   s    r$   Ú_natural_paramsÚBernoulli._natural_params‹   s   € ä—’˜DŸJ™JÓ'Ð)Ð)r&   c                 óV   • [         R                  " [         R                  " U5      5      $ r0   )r   Úlog1pÚexp)r    Úxs     r$   Ú_log_normalizerÚBernoulli._log_normalizer�   s   € Ü�{Š{œ5Ÿ9š9 Q›<Ó(Ð(r&   )r   r   r   )NNNr0   )T)&Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__r   Úunit_intervalÚrealÚarg_constraintsÚbooleanÚsupportÚhas_enumerate_supportÚ_mean_carrier_measurer   r   Úboolr   r*   r3   Úpropertyr8   r<   r@   r   r   r   r   r   rI   rQ   rY   r\   rh   Útuplerl   rr   Ú__static_attributes__Ú__classcell__)r#   s   @r$   r   r      s‡  ø† ñð* !,× 9Ñ 9À[×EUÑEUÑV€OØ×!Ñ!€GØ ÐØÐð )-Ø)-Ø%)ñ	Cà˜‰ Ñ%ðCð ˜‘ $Ñ&ðCð ˜d‘{ð	Cð
 
÷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>rŒ      s>   ðó ß Ý +Ý <÷ó õ Aß 'ð ˆ-€ô})Ð!õ })r&   