ó
    Eñié  ã                   ó¬   • S SK r S SK Jr  S SKJr  S SKJr  S SKJr  S SKJ	r	  S SK
JrJrJrJrJr  S SKJrJrJr  S	S
/r " S S	\5      r " S S
\5      rg)é    N)ÚTensor)Úconstraints)ÚDistribution)ÚTransformedDistribution)ÚSigmoidTransform)Úbroadcast_allÚclamp_probsÚlazy_propertyÚlogits_to_probsÚprobs_to_logits)Ú_NumberÚ_sizeÚNumberÚLogitRelaxedBernoulliÚRelaxedBernoullic                   ód  ^ • \ rS rSrSr\R                  \R                  S.r\R                  r	   S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	\R*                  4S j5       r\R*                  " 5       4S\S	\
4S jjrS rSrU =r$ )r   é   aO  
Creates a LogitRelaxedBernoulli distribution parameterized by :attr:`probs`
or :attr:`logits` (but not both), which is the logit of a RelaxedBernoulli
distribution.

Samples are logits of values in (0, 1). See [1] for more details.

Args:
    temperature (Tensor): relaxation temperature
    probs (Number, Tensor): the probability of sampling `1`
    logits (Number, Tensor): the log-odds of sampling `1`

[1] The Concrete Distribution: A Continuous Relaxation of Discrete Random
Variables (Maddison et al., 2017)

[2] Categorical Reparametrization with Gumbel-Softmax
(Jang et al., 2017)
©ÚprobsÚlogitsNÚtemperaturer   r   Úvalidate_argsÚreturnc                 ó°  >• Xl         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 ]5  XdS9  g )Nz;Either `probs` or `logits` must be specified, but not both.zlogits is unexpectedly None©r   )r   Ú
ValueErrorÚ
isinstancer   r   r   ÚAssertionErrorr   Ú_paramÚtorchÚSizeÚsizeÚsuperÚ__init__)Úselfr   r   r   r   Ú	is_scalarÚbatch_shapeÚ	__class__s          €Úb/home/mande/repo/quber/.venv/lib/python3.13/site-packages/torch/distributions/relaxed_bernoulli.pyr$   ÚLogitRelaxedBernoulli.__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U R                  Ul        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!   r   Ú__dict__r   Úexpandr   r   r#   r$   Ú_validate_args©r%   r'   Ú	_instanceÚnewr(   s       €r)   r/   ÚLogitRelaxedBernoulli.expandK   s³   ø€ Ø×(Ñ(Ô)>À	ÓJˆÜ—j’j Ó-ˆØ×*Ñ*ˆŒØ�d—m‘mÓ#ØŸ
™
×)Ñ)¨+Ó6ˆCŒIØŸ™ˆCŒJØ�t—}‘}Ó$ØŸ™×+Ñ+¨KÓ8ˆCŒJØŸ™ˆCŒJÜÔ# SÑ2°;ÈeÐ2ÑTØ!×0Ñ0ˆÔØˆ
r+   c                 ó:   • U R                   R                  " U0 UD6$ ©N)r   r3   )r%   ÚargsÚkwargss      r)   Ú_newÚLogitRelaxedBernoulli._newY   s   € Ø�{‰{�Š Ð/¨Ñ/Ð/r+   c                 ó*   • [        U R                  SS9$ ©NT)Ú	is_binary)r   r   ©r%   s    r)   r   ÚLogitRelaxedBernoulli.logits\   s   € ä˜tŸz™z°TÑ:Ð:r+   c                 ó*   • [        U R                  SS9$ r<   )r   r   r>   s    r)   r   ÚLogitRelaxedBernoulli.probs`   s   € ä˜tŸ{™{°dÑ;Ð;r+   c                 ó6   • U R                   R                  5       $ r6   )r   r"   r>   s    r)   Úparam_shapeÚ!LogitRelaxedBernoulli.param_shaped   s   € à�{‰{×ÑÓ!Ð!r+   Úsample_shapec                 ót  • U R                  U5      n[        U R                  R                  U5      5      n[        [        R
                  " X#R                  UR                  S95      nUR                  5       U* R                  5       -
  UR                  5       -   U* R                  5       -
  U R                  -  $ )N)ÚdtypeÚdevice)Ú_extended_shaper	   r   r/   r    ÚrandrG   rH   ÚlogÚlog1pr   )r%   rE   Úshaper   Úuniformss        r)   ÚrsampleÚLogitRelaxedBernoulli.rsampleh   s’   € Ø×$Ñ$ \Ó2ˆÜ˜DŸJ™J×-Ñ-¨eÓ4Ó5ˆÜÜ�JŠJ�u§K¡K¸¿¹ÑEó
ˆð �L‰L‹N˜x˜i×.Ñ.Ó0Ñ0°5·9±9³;Ñ>À5À&ÇÁÓAQÑQØ×Ññð 	r+   c                 ó.  • U R                   (       a  U R                  U5        [        U R                  U5      u  p!X!R	                  U R
                  5      -
  nU R
                  R                  5       U-   SUR                  5       R                  5       -  -
  $ )Né   )	r0   Ú_validate_sampler   r   Úmulr   rK   ÚexprL   )r%   Úvaluer   Údiffs       r)   Úlog_probÚLogitRelaxedBernoulli.log_probr   su   € Ø××Ø×!Ñ! %Ô(Ü% d§k¡k°5Ó9‰ˆØŸ	™	 $×"2Ñ"2Ó3Ñ3ˆØ×Ñ×#Ñ#Ó%¨Ñ,¨q°4·8±8³:×3CÑ3CÓ3EÑ/EÑEÐEr+   )r   r   r   r   ©NNNr6   )Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__r   Úunit_intervalÚrealÚarg_constraintsÚsupportr   r   Úboolr$   r/   r9   r
   r   r   Úpropertyr    r!   rC   r   rO   rX   Ú__static_attributes__Ú__classcell__©r(   s   @r)   r   r      s  ø† ñð( !,× 9Ñ 9À[×EUÑEUÑV€OØ×Ñ€Gð
 )-Ø)-Ø%)ñCàðCð ˜‰ Ñ%ðCð ˜‘ $Ñ&ð	Cð
 ˜d‘{ðCð 
÷Cð C÷:ò0ð ð;˜ó ;ó ð;ð ð<�vó <ó ð<ð ð"˜UŸZ™Zó "ó ð"ð -2¯JªJ«Lñ  Eð ¸Võ ÷Fð Fr+   c                   ó  ^ • \ rS rSr% Sr\R                  \R                  S.r\R                  r	Sr
\\S'      S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\4S j5       r\S\4S j5       r\S\4S j5       rSrU =r$ )r   éz   a¬  
Creates a RelaxedBernoulli distribution, parametrized by
:attr:`temperature`, and either :attr:`probs` or :attr:`logits`
(but not both). This is a relaxed version of the `Bernoulli` distribution,
so the values are in (0, 1), and has reparametrizable samples.

Example::

    >>> # xdoctest: +IGNORE_WANT("non-deterministic")
    >>> m = RelaxedBernoulli(torch.tensor([2.2]),
    ...                      torch.tensor([0.1, 0.2, 0.3, 0.99]))
    >>> m.sample()
    tensor([ 0.2951,  0.3442,  0.8918,  0.9021])

Args:
    temperature (Tensor): relaxation temperature
    probs (Number, Tensor): the probability of sampling `1`
    logits (Number, Tensor): the log-odds of sampling `1`
r   TÚ	base_distNr   r   r   r   r   c                 óL   >• [        XU5      n[        TU ]	  U[        5       US9  g )Nr   )r   r#   r$   r   )r%   r   r   r   r   rk   r(   s         €r)   r$   ÚRelaxedBernoulli.__init__–   s)   ø€ ô *¨+¸fÓEˆ	Ü‰Ñ˜Ô$4Ó$6ÀmÐÒTr+   c                 óJ   >• U R                  [        U5      n[        TU ]  XS9$ )N)r2   )r-   r   r#   r/   r1   s       €r)   r/   ÚRelaxedBernoulli.expand    s'   ø€ Ø×(Ñ(Ô)9¸9ÓEˆÜ‰w‰~˜kˆ~Ð9Ð9r+   c                 ó.   • U R                   R                  $ r6   )rk   r   r>   s    r)   r   ÚRelaxedBernoulli.temperature¤   s   € à�~‰~×)Ñ)Ð)r+   c                 ó.   • U R                   R                  $ r6   )rk   r   r>   s    r)   r   ÚRelaxedBernoulli.logits¨   s   € à�~‰~×$Ñ$Ð$r+   c                 ó.   • U R                   R                  $ r6   )rk   r   r>   s    r)   r   ÚRelaxedBernoulli.probs¬   s   € à�~‰~×#Ñ#Ð#r+   © rZ   r6   )r[   r\   r]   r^   r_   r   r`   ra   rb   rc   Úhas_rsampler   Ú__annotations__r   r   rd   r$   r/   re   r   r   r   rf   rg   rh   s   @r)   r   r   z   sî   ø‡ ñð( !,× 9Ñ 9À[×EUÑEUÑV€Oà×'Ñ'€GØ€Kà$Ó$ð
 )-Ø)-Ø%)ñUàðUð ˜‰ Ñ%ðUð ˜‘ $Ñ&ð	Uð
 ˜d‘{ðUð 
÷Uð U÷:ð ð*˜Vó *ó ð*ð ð%˜ó %ó ð%ð ð$�vó $ó ö$r+   )r    r   Útorch.distributionsr   Ú torch.distributions.distributionr   Ú,torch.distributions.transformed_distributionr   Útorch.distributions.transformsr   Útorch.distributions.utilsr   r	   r
   r   r   Útorch.typesr   r   r   Ú__all__r   r   rv   r+   r)   Ú<module>r€      sV   ðó Ý Ý +Ý 9Ý PÝ ;÷õ ÷ /Ñ .ð #Ð$6Ð
7€ôaF˜Lô aFôH4$Ð.õ 4$r+   