Ë
    *êñié  ã                   ó¬   — d dl Z d dl mZ d dlmZ d dlmZ d dlmZ d dlm	Z	 d dl
mZmZmZmZmZ d dlmZmZmZ d	d
gZ G d„ d	e«      Z G d„ d
e«      Zy)é    N)ÚTensor)Úconstraints)ÚDistribution)ÚTransformedDistribution)ÚSigmoidTransform)Úbroadcast_allÚclamp_probsÚlazy_propertyÚlogits_to_probsÚprobs_to_logits)Ú_NumberÚ_sizeÚNumberÚLogitRelaxedBernoulliÚRelaxedBernoullic                   óH  ‡ — e Zd ZdZej
                  ej                  dœZej                  Z	 	 	 dde	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j(                  fd„«       Z ej(                  «       fdede	fd„Zd„ Zˆ xZS )r   aƒ  
    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                 ó¦  •— || _         |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        ‰| �5  ||¬«       y )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          €úg/var/www/pod-logistic/pod-ai/venv/lib/python3.12/site-packages/torch/distributions/relaxed_bernoulli.pyr#   z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Ü‰Ñ˜°MÐÕBó    c                 óÈ  •— | j                  t        |«      }t        j                  |«      }| 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    r   Ú__dict__r   Úexpandr   r   r"   r#   Ú_validate_args©r$   r&   Ú	_instanceÚnewr'   s       €r(   r-   z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                 ó:   —  | j                   j                  |i |¤ŽS ©N)r   r1   )r$   ÚargsÚkwargss      r(   Ú_newzLogitRelaxedBernoulli._newY   s   € Øˆt�{‰{�‰ Ð/¨Ñ/Ð/r)   c                 ó0   — t        | j                  d¬«      S ©NT)Ú	is_binary)r   r   ©r$   s    r(   r   zLogitRelaxedBernoulli.logits\   s   € ä˜tŸz™z°TÔ:Ð:r)   c                 ó0   — t        | j                  d¬«      S r8   )r   r   r:   s    r(   r   zLogitRelaxedBernoulli.probs`   s   € ä˜tŸ{™{°dÔ;Ð;r)   c                 ó6   — | j                   j                  «       S r3   )r   r!   r:   s    r(   Úparam_shapez!LogitRelaxedBernoulli.param_shaped   s   € à�{‰{×ÑÓ!Ð!r)   Úsample_shapec                 óz  — | j                  |«      }t        | j                  j                  |«      «      }t        t	        j
                  ||j                  |j                  ¬«      «      }|j                  «       | j                  «       z
  |j                  «       z   | j                  «       z
  | j                  z  S )N)ÚdtypeÚdevice)Ú_extended_shaper	   r   r-   r   Úrandr@   rA   ÚlogÚlog1pr   )r$   r>   Úshaper   Úuniformss        r(   ÚrsamplezLogitRelaxedBernoulli.rsampleh   s”   € Ø×$Ñ$ \Ó2ˆÜ˜DŸJ™J×-Ñ-¨eÓ4Ó5ˆÜÜ�J‰J�u E§K¡K¸¿¹ÔEó
ˆð �L‰L‹N˜x˜i×.Ñ.Ó0Ñ0°5·9±9³;Ñ>À5À&ÇÁÓAQÑQØ×Ññð 	r)   c                 ó(  — | j                   r| j                  |«       t        | j                  |«      \  }}||j	                  | j
                  «      z
  }| j
                  j                  «       |z   d|j                  «       j                  «       z  z
  S )Né   )	r.   Ú_validate_sampler   r   Úmulr   rD   ÚexprE   )r$   Úvaluer   Údiffs       r(   Úlog_probzLogitRelaxedBernoulli.log_probr   sy   € Ø×ÒØ×!Ñ! %Ô(Ü% d§k¡k°5Ó9‰ˆ�Ø˜Ÿ	™	 $×"2Ñ"2Ó3Ñ3ˆØ×Ñ×#Ñ#Ó%¨Ñ,¨q°4·8±8³:×3CÑ3CÓ3EÑ/EÑEÐEr)   ©NNNr3   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   Úunit_intervalÚrealÚarg_constraintsÚsupportr   r   Úboolr#   r-   r6   r
   r   r   Úpropertyr   r    r=   r   rH   rP   Ú__classcell__©r'   s   @r(   r   r      s  ø„ ñð( !,× 9Ñ 9À[×EUÑEUÑV€OØ×Ñ€Gð
 )-Ø)-Ø%)ñCàðCð ˜‰ Ñ%ðCð ˜‘ $Ñ&ð	Cð
 ˜d‘{ðCð 
õCõ:ò0ð ð;˜ò ;ó ð;ð ð<�vò <ó ð<ð ð"˜UŸZ™Zò "ó ð"ð -7¨E¯J©J«Lñ  Eð ¸Vó öFr)   c                   ó  ‡ — e Zd ZU dZej
                  ej                  dœZej
                  ZdZ	e
ed<   	 	 	 dde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ed
efd„«       Zed
efd„«       Zed
efd„«       Zˆ xZS )r   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                 óT   •— t        |||«      }t        ‰| �	  |t        «       |¬«       y )Nr   )r   r"   r#   r   )r$   r   r   r   r   r_   r'   s         €r(   r#   zRelaxedBernoulli.__init__–   s+   ø€ ô *¨+°u¸fÓEˆ	Ü‰Ñ˜Ô$4Ó$6ÀmÐÕTr)   c                 óR   •— | j                  t        |«      }t        ‰| �  ||¬«      S )N)r0   )r+   r   r"   r-   r/   s       €r(   r-   zRelaxedBernoulli.expand    s)   ø€ Ø×(Ñ(Ô)9¸9ÓEˆÜ‰w‰~˜k°Sˆ~Ó9Ð9r)   c                 ó.   — | j                   j                  S r3   )r_   r   r:   s    r(   r   zRelaxedBernoulli.temperature¤   s   € à�~‰~×)Ñ)Ð)r)   c                 ó.   — | j                   j                  S r3   )r_   r   r:   s    r(   r   zRelaxedBernoulli.logits¨   s   € à�~‰~×$Ñ$Ð$r)   c                 ó.   — | j                   j                  S r3   )r_   r   r:   s    r(   r   zRelaxedBernoulli.probs¬   s   € à�~‰~×#Ñ#Ð#r)   rQ   r3   )rR   rS   rT   rU   r   rV   rW   rX   rY   Úhas_rsampler   Ú__annotations__r   r   rZ   r#   r-   r[   r   r   r   r\   r]   s   @r(   r   r   z   sè   ø… ñð( !,× 9Ñ 9À[×EUÑEUÑV€Oà×'Ñ'€GØ€Kà$Ó$ð
 )-Ø)-Ø%)ñUàðUð ˜‰ Ñ%ðUð ˜‘ $Ñ&ð	Uð
 ˜d‘{ð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   © r)   r(   ú<module>ro      sV   ðó Ý Ý +Ý 9Ý PÝ ;÷õ ÷ /Ñ .ð #Ð$6Ð
7€ôaF˜Lô aFôH4$Ð.õ 4$r)   