§
    ŠŠtjé  ã                   óº   — 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dS )é    N)ÚTensor)Úconstraints)ÚDistribution)ÚTransformedDistribution)ÚSigmoidTransform)Úbroadcast_allÚclamp_probsÚlazy_propertyÚlogits_to_probsÚprobs_to_logits)Ú_NumberÚ_sizeÚNumberÚLogitRelaxedBernoulliÚRelaxedBernoullic                   ó.  ‡ — 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                             ¦   «         }t          ¦   «                              ||¬¦  «         d S )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          €úc/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/torch/distributions/relaxed_bernoulli.pyr#   zLogitRelaxedBernoulli.__init__.   sò   ø€ ð 'ˆÔØ�TˆM˜v¨˜~Ò.Ð.ÝØMñô ð ð ÐÝ" 5­'Ñ2Ô2ˆIå)¨%Ñ0Ô0‰MˆTŒZˆZàˆ~Ý$Ð%BÑCÔCÐCÝ" 6­7Ñ3Ô3ˆIå*¨6Ñ2Ô2‰NˆTŒ[Ø$)Ð$5�d”j�j¸4¼;ˆŒØð 	-Ýœ*™,œ,ˆKˆKàœ+×*Ò*Ñ,Ô,ˆKÝ‰Œ×Ò˜°MÐÑBÔBÐBÐBÐBó    c                 óº  •— |                       t          |¦  «        }t          j        |¦  «        }| j        |_        d| j        v r+| j                             |¦  «        |_        |j        |_        d| j        v r+| 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Ð#Ð#Øœ
×)Ò)¨+Ñ6Ô6ˆCŒIØœˆCŒJØ�t”}Ð$Ð$Øœ×+Ò+¨KÑ8Ô8ˆCŒJØœˆCŒJÝÕ# SÑ)Ô)×2Ò2°;ÈeÐ2ÑTÔTÐTØ!Ô0ˆÔØˆ
r)   c                 ó&   —  | j         j        |i |¤ŽS ©N)r   r1   )r$   ÚargsÚkwargss      r(   Ú_newzLogitRelaxedBernoulli._newY   s   € ØˆtŒ{Œ Ð/¨Ð/Ð/Ð/r)   c                 ó.   — t          | j        d¬¦  «        S ©NT)Ú	is_binary)r   r   ©r$   s    r(   r   zLogitRelaxedBernoulli.logits\   s   € å˜tœz°TÐ:Ñ:Ô:Ð:r)   c                 ó.   — t          | j        d¬¦  «        S r8   )r   r   r:   s    r(   r   zLogitRelaxedBernoulli.probs`   s   € å˜tœ{°dÐ;Ñ;Ô;Ð;r)   c                 ó4   — | j                              ¦   «         S r3   )r   r!   r:   s    r(   Úparam_shapez!LogitRelaxedBernoulli.param_shaped   s   € àŒ{×ÒÑ!Ô!Ð!r)   Úsample_shapec                 ó�  — |                       |¦  «        }t          | j                             |¦  «        ¦  «        }t          t	          j        ||j        |j        ¬¦  «        ¦  «        }|                     ¦   «         |  	                    ¦   «         z
  |                     ¦   «         z   |  	                    ¦   «         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Ô2ˆÝ˜DœJ×-Ò-¨eÑ4Ô4Ñ5Ô5ˆÝÝŒJ�u E¤K¸¼ÐEÑEÔEñ
ô 
ˆð �LŠL‰NŒN˜x˜i×.Ò.Ñ0Ô0Ñ0°5·9²9±;´;Ñ>À5À&ÇÂÑAQÔAQÑQØÔñð 	r)   c                 ó0  — | j         r|                      |¦  «         t          | j        |¦  «        \  }}||                     | j        ¦  «        z
  }| j                             ¦   «         |z   d|                     ¦   «                              ¦   «         z  z
  S )Né   )	r.   Ú_validate_sampler   r   Úmulr   rD   ÚexprE   )r$   Úvaluer   Údiffs       r(   Úlog_probzLogitRelaxedBernoulli.log_probr   s‡   € ØÔð 	)Ø×!Ò! %Ñ(Ô(Ð(Ý% d¤k°5Ñ9Ô9‰ˆ�Ø˜Ÿ	š	 $Ô"2Ñ3Ô3Ñ3ˆØÔ×#Ò#Ñ%Ô%¨Ñ,¨q°4·8²8±:´:×3CÒ3CÑ3EÔ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À[ÔEUÐVÐV€OØÔ€Gð
 )-Ø)-Ø%)ðCð CàðCð ˜‰ Ñ%ðCð ˜‘ $Ñ&ð	Cð
 ˜d‘{ðCð 
ðCð Cð Cð Cð Cð Cð:ð ð ð ð ð ð0ð 0ð 0ð ð;˜ð ;ð ;ð ;ñ „]ð;ð ð<�vð <ð <ð <ñ „]ð<ð ð"˜UœZð "ð "ð "ñ „Xð"ð -7¨E¬J©L¬Lð ð  Eð ¸Vð ð ð ð ðFð Fð Fð Fð Fð Fð 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          ¦   «         |¬¦  «         d S )Nr   )r   r"   r#   r   )r$   r   r   r   r   r_   r'   s         €r(   r#   zRelaxedBernoulli.__init__–   sB   ø€ õ *¨+°u¸fÑEÔEˆ	Ý‰Œ×Ò˜Õ$4Ñ$6Ô$6ÀmÐÑTÔTÐTÐTÐTr)   c                 ó€   •— |                       t          |¦  «        }t          ¦   «                              ||¬¦  «        S )N)r0   )r+   r   r"   r-   r/   s       €r(   r-   zRelaxedBernoulli.expand    s3   ø€ Ø×(Ò(Õ)9¸9ÑEÔEˆÝ‰wŒw�~Š~˜k°Sˆ~Ñ9Ô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   sl  ø€ € € € € € ðð ð( !,Ô 9À[ÔEUÐVÐV€OàÔ'€GØ€Kà$Ð$Ð$Ñ$ð
 )-Ø)-Ø%)ðUð UàðUð ˜‰ Ñ%ðUð ˜‘ $Ñ&ð	Uð
 ˜d‘{ðUð 
ðUð Uð Uð Uð Uð Uð:ð :ð :ð :ð :ð :ð ð*˜Vð *ð *ð *ñ „Xð*ð ð%˜ð %ð %ð %ñ „Xð%ð ð$�vð $ð $ð $ñ „Xð$ð $ð $ð $ð $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      sM  ðð €€€Ø Ð Ð Ð Ð Ð Ø +Ð +Ð +Ð +Ð +Ð +Ø 9Ð 9Ð 9Ð 9Ð 9Ð 9Ø PÐ PÐ PÐ PÐ PÐ PØ ;Ð ;Ð ;Ð ;Ð ;Ð ;ðð ð ð ð ð ð ð ð ð ð ð ð ð ð /Ð .Ð .Ð .Ð .Ð .Ð .Ð .Ð .Ð .ð #Ð$6Ð
7€ðaFð aFð aFð aFð aF˜Lñ aFô aFð aFðH4$ð 4$ð 4$ð 4$ð 4$Ð.ñ 4$ô 4$ð 4$ð 4$ð 4$r)   