Ë
    f^(h‘  ã                   ó¨   — 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 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ÚLogitRelaxedBernoulliÚRelaxedBernoullic                   ó  ‡ — e Zd ZdZej
                  ej                  dœZej                  Zdˆ 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Úlogitsc                 óŒ  •— || _         |d u |d u k(  rt        d«      ‚|�#t        |t        «      }t	        |«      \  | _        n"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.©Úvalidate_args)ÚtemperatureÚ
ValueErrorÚ
isinstancer   r   r   r   Ú_paramÚtorchÚSizeÚsizeÚsuperÚ__init__)Úselfr   r   r   r   Ú	is_scalarÚbatch_shapeÚ	__class__s          €úc/var/www/skyplay_api_hub/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ä" 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.expand?   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   r.   )r!   ÚargsÚkwargss      r%   Ú_newzLogitRelaxedBernoulli._newM   s   € Øˆt�{‰{�‰ Ð/¨Ñ/Ð/r&   Úreturnc                 ó0   — t        | j                  d¬«      S ©NT)Ú	is_binary)r   r   ©r!   s    r%   r   zLogitRelaxedBernoulli.logitsP   s   € ä˜tŸz™z°TÔ:Ð:r&   c                 ó0   — t        | j                  d¬«      S r6   )r   r   r8   s    r%   r   zLogitRelaxedBernoulli.probsT   s   € ä˜tŸ{™{°dÔ;Ð;r&   c                 ó6   — | j                   j                  «       S r0   )r   r   r8   s    r%   Úparam_shapez!LogitRelaxedBernoulli.param_shapeX   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>   r?   ÚlogÚlog1pr   )r!   r<   Úshaper   Úuniformss        r%   ÚrsamplezLogitRelaxedBernoulli.rsample\   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   rB   ÚexprC   )r!   Úvaluer   Údiffs       r%   Úlog_probzLogitRelaxedBernoulli.log_probf   sy   € Ø×ÒØ×!Ñ! %Ô(Ü% d§k¡k°5Ó9‰ˆ�Ø˜Ÿ	™	 $×"2Ñ"2Ó3Ñ3ˆØ×Ñ×#Ñ#Ó%¨Ñ,¨q°4·8±8³:×3CÑ3CÓ3EÑ/EÑEÐEr&   ©NNNr0   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   Úunit_intervalÚrealÚarg_constraintsÚsupportr    r*   r3   r
   r   r   r   Úpropertyr   r   r;   r   rF   rN   Ú__classcell__©r$   s   @r%   r   r      s¶   ø„ ñð& !,× 9Ñ 9À[×EUÑEUÑV€OØ×Ñ€GõCõ&ò0ð ð;˜ò ;ó ð;ð ð<�vò <ó ð<ð ð"˜UŸZ™Zò "ó ð"ð -7¨E¯J©J«Lñ  Eð ¸Vó öFr&   c                   óÀ   ‡ — e Zd ZdZej
                  ej                  dœZej
                  ZdZ	d
ˆ 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   Tc                 óT   •— t        |||«      }t        ‰| �	  |t        «       |¬«       y )Nr   )r   r   r    r   )r!   r   r   r   r   Ú	base_distr$   s         €r%   r    zRelaxedBernoulli.__init__‡   s)   ø€ Ü)¨+°u¸fÓEˆ	Ü‰Ñ˜Ô$4Ó$6ÀmÐÕTr&   c                 óR   •— | j                  t        |«      }t        ‰| �  ||¬«      S )N)r-   )r(   r   r   r*   r,   s       €r%   r*   zRelaxedBernoulli.expand‹   s)   ø€ Ø×(Ñ(Ô)9¸9ÓEˆÜ‰w‰~˜k°Sˆ~Ó9Ð9r&   r4   c                 ó.   — | j                   j                  S r0   )r]   r   r8   s    r%   r   zRelaxedBernoulli.temperature�   s   € à�~‰~×)Ñ)Ð)r&   c                 ó.   — | j                   j                  S r0   )r]   r   r8   s    r%   r   zRelaxedBernoulli.logits“   s   € à�~‰~×$Ñ$Ð$r&   c                 ó.   — | j                   j                  S r0   )r]   r   r8   s    r%   r   zRelaxedBernoulli.probs—   s   € à�~‰~×#Ñ#Ð#r&   rO   r0   )rP   rQ   rR   rS   r   rT   rU   rV   rW   Úhas_rsampler    r*   rX   r   r   r   r   rY   rZ   s   @r%   r   r   n   sŒ   ø„ ñð( !,× 9Ñ 9À[×EUÑEUÑV€OØ×'Ñ'€GØ€Kõ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   Ú__all__r   r   © r&   r%   ú<module>rk      sQ   ðã Ý Ý +Ý 9Ý PÝ ;÷õ ÷ 'ð #Ð$6Ð
7€ôVF˜Lô VFôr+$Ð.õ +$r&   