§
    ŠŠtjx  ã                   ó²   — 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 d dlmZmZ d d	lmZ d
dgZ G d„ d
e¦  «        Z G d„ de	¦  «        ZdS )é    N)ÚTensor)Úconstraints)ÚCategorical)ÚDistribution)ÚTransformedDistribution)ÚExpTransform)Úbroadcast_allÚclamp_probs)Ú_sizeÚExpRelaxedCategoricalÚRelaxedOneHotCategoricalc                   ó&  ‡ — e Zd ZdZej        ej        dœZej        ZdZ		 	 	 dde
de
dz  de
dz  dedz  d	df
ˆ fd
„Zdˆ fd„	Zd„ Zed	ej        fd„¦   «         Zed	e
fd„¦   «         Zed	e
fd„¦   «         Z ej        ¦   «         fded	e
fd„Zd„ Zˆ xZS )r   aÏ  
    Creates a ExpRelaxedCategorical parameterized by
    :attr:`temperature`, and either :attr:`probs` or :attr:`logits` (but not both).
    Returns the log of a point in the simplex. Based on the interface to
    :class:`OneHotCategorical`.

    Implementation based on [1].

    See also: :func:`torch.distributions.OneHotCategorical`

    Args:
        temperature (Tensor): relaxation temperature
        probs (Tensor): event probabilities
        logits (Tensor): unnormalized log probability for each event

    [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ÚlogitsTNÚtemperaturer   r   Úvalidate_argsÚreturnc                 óÈ   •— t          ||¦  «        | _        || _        | j        j        }| j        j        dd …         }t          ¦   «                              |||¬¦  «         d S )Néÿÿÿÿ©r   )r   Ú_categoricalr   Úbatch_shapeÚparam_shapeÚsuperÚ__init__)Úselfr   r   r   r   r   Úevent_shapeÚ	__class__s          €úe/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/torch/distributions/relaxed_categorical.pyr   zExpRelaxedCategorical.__init__/   sc   ø€ õ (¨¨vÑ6Ô6ˆÔØ&ˆÔØÔ'Ô3ˆØÔ'Ô3°B°C°CÔ8ˆå‰Œ×Ò˜ kÀÐÑOÔOÐOÐOÐOó    c                 ó4  •— |                       t          |¦  «        }t          j        |¦  «        }| j        |_        | j                             |¦  «        |_        t          t          |¦  «                             || j	        d¬¦  «         | j
        |_
        |S )NFr   )Ú_get_checked_instancer   ÚtorchÚSizer   r   Úexpandr   r   r   Ú_validate_args©r   r   Ú	_instanceÚnewr   s       €r    r&   zExpRelaxedCategorical.expand=   s�   ø€ Ø×(Ò(Õ)>À	ÑJÔJˆÝ”j Ñ-Ô-ˆØÔ*ˆŒØÔ,×3Ò3°KÑ@Ô@ˆÔÝÕ# SÑ)Ô)×2Ò2Ø˜Ô)¸ð 	3ñ 	
ô 	
ð 	
ð "Ô0ˆÔØˆ
r!   c                 ó&   —  | j         j        |i |¤ŽS ©N)r   Ú_new)r   ÚargsÚkwargss      r    r-   zExpRelaxedCategorical._newH   s   € Ø%ˆtÔ Ô% tÐ6¨vÐ6Ð6Ð6r!   c                 ó   — | j         j        S r,   )r   r   ©r   s    r    r   z!ExpRelaxedCategorical.param_shapeK   s   € àÔ Ô,Ð,r!   c                 ó   — | j         j        S r,   )r   r   r1   s    r    r   zExpRelaxedCategorical.logitsO   s   € àÔ Ô'Ð'r!   c                 ó   — | j         j        S r,   )r   r   r1   s    r    r   zExpRelaxedCategorical.probsS   s   € àÔ Ô&Ð&r!   Úsample_shapec                 óD  — |                       |¦  «        }t          t          j        || j        j        | j        j        ¬¦  «        ¦  «        }|                     ¦   «                               ¦   «          }| j        |z   | j        z  }|| 	                    dd¬¦  «        z
  S )N)ÚdtypeÚdevicer   T©ÚdimÚkeepdim)
Ú_extended_shaper
   r$   Úrandr   r6   r7   Úlogr   Ú	logsumexp)r   r4   ÚshapeÚuniformsÚgumbelsÚscoress         r    ÚrsamplezExpRelaxedCategorical.rsampleW   s•   € Ø×$Ò$ \Ñ2Ô2ˆÝÝŒJ�u D¤KÔ$5¸d¼kÔ>PÐQÑQÔQñ
ô 
ˆð  —|’|‘~”~Ð&×+Ò+Ñ-Ô-Ð.ˆØ”+ Ñ'¨4Ô+;Ñ;ˆØ˜×(Ò(¨R¸Ð(Ñ>Ô>Ñ>Ð>r!   c                 óô  — | j         j        }| j        r|                      |¦  «         t	          | j        |¦  «        \  }}t          j        | j        t          |¦  «        ¦  «         
                    ¦   «         | j                             ¦   «                              |dz
   ¦  «        z
  }||                     | j        ¦  «        z
  }||                     dd¬¦  «        z
                       d¦  «        }||z   S )Né   r   Tr8   )r   Ú_num_eventsr'   Ú_validate_sampler	   r   r$   Ú	full_liker   ÚfloatÚlgammar=   Úmulr>   Úsum)r   ÚvalueÚKr   Ú	log_scaleÚscores         r    Úlog_probzExpRelaxedCategorical.log_prob`   sá   € ØÔÔ)ˆØÔð 	)Ø×!Ò! %Ñ(Ô(Ð(Ý% d¤k°5Ñ9Ô9‰ˆ�Ý”OØÔ�e A™hœhñ
ô 
ç
Š&‰(Œ(�TÔ%×)Ò)Ñ+Ô+×/Ò/°!°a±%°Ñ9Ô9ñ:ˆ	ð ˜Ÿš 4Ô#3Ñ4Ô4Ñ4ˆØ˜Ÿš¨R¸˜Ñ>Ô>Ñ>×CÒCÀBÑGÔGˆØ�yÑ Ð r!   ©NNNr,   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   ÚsimplexÚreal_vectorÚarg_constraintsÚsupportÚhas_rsampler   Úboolr   r&   r-   Úpropertyr$   r%   r   r   r   r   rC   rQ   Ú__classcell__©r   s   @r    r   r      s­  ø€ € € € € ðð ð. !,Ô 3¸{Ô?VÐWÐW€OàÔð ð €Kð
  $Ø $Ø%)ðPð PàðPð ˜‰}ðPð ˜‘ð	Pð
 ˜d‘{ðPð 
ðPð Pð Pð Pð Pð Pð	ð 	ð 	ð 	ð 	ð 	ð7ð 7ð 7ð ð-˜UœZð -ð -ð -ñ „Xð-ð ð(˜ð (ð (ð (ñ „Xð(ð ð'�vð 'ð 'ð 'ñ „Xð'ð -7¨E¬J©L¬Lð ?ð ? Eð ?¸Vð ?ð ?ð ?ð ?ð
!ð 
!ð 
!ð 
!ð 
!ð 
!ð 
!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dz  de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 RelaxedOneHotCategorical distribution parametrized by
    :attr:`temperature`, and either :attr:`probs` or :attr:`logits`.
    This is a relaxed version of the :class:`OneHotCategorical` distribution, so
    its samples are on simplex, and are reparametrizable.

    Example::

        >>> # xdoctest: +IGNORE_WANT("non-deterministic")
        >>> m = RelaxedOneHotCategorical(torch.tensor([2.2]),
        ...                              torch.tensor([0.1, 0.2, 0.3, 0.4]))
        >>> m.sample()
        tensor([ 0.1294,  0.2324,  0.3859,  0.2523])

    Args:
        temperature (Tensor): relaxation temperature
        probs (Tensor): event probabilities
        logits (Tensor): unnormalized log probability for each event
    r   TÚ	base_distNr   r   r   r   r   c                 óŽ   •— t          ||||¬¦  «        }t          ¦   «                              |t          ¦   «         |¬¦  «         d S )Nr   )r   r   r   r   )r   r   r   r   r   ra   r   s         €r    r   z!RelaxedOneHotCategorical.__init__‰   sM   ø€ õ *Ø˜ °mð
ñ 
ô 
ˆ	õ 	‰Œ×Ò˜¥L¡N¤NÀ-ÐÑPÔPÐPÐPÐPr!   c                 ó€   •— |                       t          |¦  «        }t          ¦   «                              ||¬¦  «        S )N)r)   )r#   r   r   r&   r(   s       €r    r&   zRelaxedOneHotCategorical.expand•   s3   ø€ Ø×(Ò(Õ)AÀ9ÑMÔMˆÝ‰wŒw�~Š~˜k°Sˆ~Ñ9Ô9Ð9r!   c                 ó   — | j         j        S r,   )ra   r   r1   s    r    r   z$RelaxedOneHotCategorical.temperature™   s   € àŒ~Ô)Ð)r!   c                 ó   — | j         j        S r,   )ra   r   r1   s    r    r   zRelaxedOneHotCategorical.logits�   s   € àŒ~Ô$Ð$r!   c                 ó   — | j         j        S r,   )ra   r   r1   s    r    r   zRelaxedOneHotCategorical.probs¡   s   € àŒ~Ô#Ð#r!   rR   r,   )rS   rT   rU   rV   r   rW   rX   rY   rZ   r[   r   Ú__annotations__r   r\   r   r&   r]   r   r   r   r^   r_   s   @r    r   r   m   sb  ø€ € € € € € ðð ð( !,Ô 3¸{Ô?VÐWÐW€OàÔ!€GØ€Kà$Ð$Ð$Ñ$ð
  $Ø $Ø%)ð
Qð 
Qàð
Qð ˜‰}ð
Qð ˜‘ð	
Qð
 ˜d‘{ð
Qð 
ð
Qð 
Qð 
Qð 
Qð 
Qð 
Qð:ð :ð :ð :ð :ð :ð ð*˜Vð *ð *ð *ñ „Xð*ð ð%˜ð %ð %ð %ñ „Xð%ð ð$�vð $ð $ð $ñ „Xð$ð $ð $ð $ð $r!   )r$   r   Útorch.distributionsr   Útorch.distributions.categoricalr   Ú torch.distributions.distributionr   Ú,torch.distributions.transformed_distributionr   Útorch.distributions.transformsr   Útorch.distributions.utilsr	   r
   Útorch.typesr   Ú__all__r   r   © r!   r    ú<module>rq      s  ðð €€€Ø Ð Ð Ð Ð Ð Ø +Ð +Ð +Ð +Ð +Ð +Ø 7Ð 7Ð 7Ð 7Ð 7Ð 7Ø 9Ð 9Ð 9Ð 9Ð 9Ð 9Ø PÐ PÐ PÐ PÐ PÐ PØ 7Ð 7Ð 7Ð 7Ð 7Ð 7Ø @Ð @Ð @Ð @Ð @Ð @Ð @Ð @Ø Ð Ð Ð Ð Ð ð #Ð$>Ð
?€ðY!ð Y!ð Y!ð Y!ð Y!˜Lñ Y!ô Y!ð Y!ðx6$ð 6$ð 6$ð 6$ð 6$Ð6ñ 6$ô 6$ð 6$ð 6$ð 6$r!   