Ë
    óÍ:jD  ã                   ó¼   — d dl mZmZ 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)é    )ÚOptionalÚUnionN)ÚTensor)Úconstraints)ÚDistribution)ÚTransformedDistribution)ÚSigmoidTransform)Úbroadcast_allÚclamp_probsÚlazy_propertyÚlogits_to_probsÚprobs_to_logits)Ú_NumberÚ_sizeÚNumberÚLogitRelaxedBernoulliÚRelaxedBernoullic                   óP  ‡ — e Zd ZdZej
                  ej                  dœZej                  Z	 	 	 dde	de
ee	ef      de
ee	ef      de
e   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&|€J ‚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.©r   )r   Ú
ValueErrorÚ
isinstancer   r
   r   r   Ú_paramÚtorchÚSizeÚsizeÚsuperÚ__init__)Úselfr   r   r   r   Ú	is_scalarÚbatch_shapeÚ	__class__s          €úz/home/mcse/projects/srt_converter/srt-converter-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.expandH   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   r2   )r%   ÚargsÚkwargss      r)   Ú_newzLogitRelaxedBernoulli._newV   s   € Øˆt�{‰{�‰ Ð/¨Ñ/Ð/r*   c                 ó0   — t        | j                  d¬«      S ©NT)Ú	is_binary)r   r   ©r%   s    r)   r   zLogitRelaxedBernoulli.logitsY   s   € ä˜tŸz™z°TÔ:Ð:r*   c                 ó0   — t        | j                  d¬«      S r9   )r   r   r;   s    r)   r   zLogitRelaxedBernoulli.probs]   s   € ä˜tŸ{™{°dÔ;Ð;r*   c                 ó6   — | j                   j                  «       S r4   )r   r"   r;   s    r)   Úparam_shapez!LogitRelaxedBernoulli.param_shapea   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    ÚrandrA   rB   ÚlogÚlog1pr   )r%   r?   Úshaper   Úuniformss        r)   ÚrsamplezLogitRelaxedBernoulli.rsamplee   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   rE   ÚexprF   )r%   Úvaluer   Údiffs       r)   Úlog_probzLogitRelaxedBernoulli.log_probo   sy   € Ø×ÒØ×!Ñ! %Ô(Ü% d§k¡k°5Ó9‰ˆ�Ø˜Ÿ	™	 $×"2Ñ"2Ó3Ñ3ˆØ×Ñ×#Ñ#Ó%¨Ñ,¨q°4·8±8³:×3CÑ3CÓ3EÑ/EÑEÐEr*   ©NNNr4   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   Úunit_intervalÚrealÚarg_constraintsÚsupportr   r   r   r   Úboolr$   r.   r7   r   r   r   Úpropertyr    r!   r>   r   rI   rQ   Ú__classcell__©r(   s   @r)   r   r      s  ø„ ñð& !,× 9Ñ 9À[×EUÑEUÑV€OØ×Ñ€Gð
 26Ø26Ø(,ñCàðCð ˜˜f f˜nÑ-Ñ.ðCð ˜˜v v˜~Ñ.Ñ/ð	Cð
   ‘~ðCð 
õCõ4ò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eef      deeeef      d	ee   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)r1   )r,   r   r#   r.   r0   s       €r)   r.   zRelaxedBernoulli.expand›   s)   ø€ Ø×(Ñ(Ô)9¸9ÓEˆÜ‰w‰~˜k°Sˆ~Ó9Ð9r*   c                 ó.   — | j                   j                  S r4   )r`   r   r;   s    r)   r   zRelaxedBernoulli.temperatureŸ   s   € à�~‰~×)Ñ)Ð)r*   c                 ó.   — | j                   j                  S r4   )r`   r   r;   s    r)   r   zRelaxedBernoulli.logits£   s   € à�~‰~×$Ñ$Ð$r*   c                 ó.   — | j                   j                  S r4   )r`   r   r;   s    r)   r   zRelaxedBernoulli.probs§   s   € à�~‰~×#Ñ#Ð#r*   rR   r4   )rS   rT   rU   rV   r   rW   rX   rY   rZ   Úhas_rsampler   Ú__annotations__r   r   r   r   r[   r$   r.   r\   r   r   r   r]   r^   s   @r)   r   r   w   sò   ø… ñð( !,× 9Ñ 9À[×EUÑEUÑV€OØ×'Ñ'€GØ€KØ$Ó$ð
 26Ø26Ø(,ñUàðUð ˜˜f f˜nÑ-Ñ.ðUð ˜˜v v˜~Ñ.Ñ/ð	Uð
   ‘~ðUð 
õUõ:ð ð*˜Vò *ó ð*ð ð%˜ò %ó ð%ð ð$�vò $ó ô$r*   )Útypingr   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>rq      sW   ðç "ã Ý Ý +Ý 9Ý PÝ ;÷õ ÷ /Ñ .ð #Ð$6Ð
7€ô]F˜Lô ]Fô@2$Ð.õ 2$r*   