ó
    Eñix  ã                   ó¤   • S SK r S SK Jr  S SKJr  S SKJr  S SKJr  S SKJ	r	  S SK
Jr  S SKJrJr  S S	KJr  S
S/r " S S
\5      r " S S\	5      rg)é    N)ÚTensor)Úconstraints)ÚCategorical)ÚDistribution)ÚTransformedDistribution)ÚExpTransform)Úbroadcast_allÚclamp_probs)Ú_sizeÚExpRelaxedCategoricalÚRelaxedOneHotCategoricalc                   ó\  ^ • \ rS rSrSr\R                  \R                  S.r\R                  r	Sr
   SS\S\S-  S\S-  S	\S-  S
S4
U 4S jjjrSU 4S jjrS r\S
\R$                  4S j5       r\S
\4S j5       r\S
\4S j5       r\R$                  " 5       4S\S
\4S jjrS rSrU =r$ )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                 ó¬   >• [        X#5      U l        Xl        U R                  R                  nU R                  R                  SS  n[
        TU ]  XVUS9  g )Néÿÿÿÿ©r   )r   Ú_categoricalr   Úbatch_shapeÚparam_shapeÚsuperÚ__init__)Úselfr   r   r   r   r   Úevent_shapeÚ	__class__s          €Úd/home/mande/repo/quber/.venv/lib/python3.13/site-packages/torch/distributions/relaxed_categorical.pyr   ÚExpRelaxedCategorical.__init__/   sS   ø€ ô (¨Ó6ˆÔØ&ÔØ×'Ñ'×3Ñ3ˆØ×'Ñ'×3Ñ3°B°CÐ8ˆä‰Ñ˜ÀÐÒOó    c                 ó  >• U R                  [        U5      n[        R                  " U5      nU R                  Ul        U R
                  R                  U5      Ul        [        [        U]#  XR                  SS9  U R                  Ul
        U$ )NFr   )Ú_get_checked_instancer   ÚtorchÚSizer   r   Úexpandr   r   r   Ú_validate_args©r   r   Ú	_instanceÚnewr    s       €r!   r(   ÚExpRelaxedCategorical.expand=   s   ø€ Ø×(Ñ(Ô)>À	ÓJˆÜ—j’j Ó-ˆØ×*Ñ*ˆŒØ×,Ñ,×3Ñ3°KÓ@ˆÔÜÔ# SÑ2Ø×)Ñ)¸ð 	3ñ 	
ð "×0Ñ0ˆÔØˆ
r#   c                 ó:   • U R                   R                  " U0 UD6$ ©N)r   Ú_new)r   ÚargsÚkwargss      r!   r0   ÚExpRelaxedCategorical._newH   s   € Ø× Ñ ×%Ò% tÐ6¨vÑ6Ð6r#   c                 ó.   • U R                   R                  $ r/   )r   r   ©r   s    r!   r   Ú!ExpRelaxedCategorical.param_shapeK   s   € à× Ñ ×,Ñ,Ð,r#   c                 ó.   • U R                   R                  $ r/   )r   r   r5   s    r!   r   ÚExpRelaxedCategorical.logitsO   s   € à× Ñ ×'Ñ'Ð'r#   c                 ó.   • U R                   R                  $ r/   )r   r   r5   s    r!   r   ÚExpRelaxedCategorical.probsS   s   € à× Ñ ×&Ñ&Ð&r#   Úsample_shapec                 óL  • U R                  U5      n[        [        R                  " X R                  R
                  U R                  R                  S95      nUR                  5       * R                  5       * nU R                  U-   U R                  -  nXUR                  SSS9-
  $ )N)ÚdtypeÚdevicer   T©ÚdimÚkeepdim)
Ú_extended_shaper
   r&   Úrandr   r=   r>   Úlogr   Ú	logsumexp)r   r;   ÚshapeÚuniformsÚgumbelsÚscoress         r!   ÚrsampleÚExpRelaxedCategorical.rsampleW   sŒ   € Ø×$Ñ$ \Ó2ˆÜÜ�JŠJ�u§K¡K×$5Ñ$5¸d¿k¹k×>PÑ>PÑQó
ˆð  —|‘|“~Ð&×+Ñ+Ó-Ð.ˆØ—+‘+ Ñ'¨4×+;Ñ+;Ñ;ˆØ×(Ñ(¨R¸Ð(Ð>Ñ>Ð>r#   c                 óò  • U R                   R                  nU R                  (       a  U R                  U5        [	        U R
                  U5      u  p1[        R                  " U R                  [        U5      5      R                  5       U R                  R                  5       R                  US-
  * 5      -
  nX1R                  U R                  5      -
  nXUR                  SSS9-
  R                  S5      nXT-   $ )Né   r   Tr?   )r   Ú_num_eventsr)   Ú_validate_sampler	   r   r&   Ú	full_liker   ÚfloatÚlgammarD   ÚmulrE   Úsum)r   ÚvalueÚKr   Ú	log_scaleÚscores         r!   Úlog_probÚExpRelaxedCategorical.log_prob`   sÉ   € Ø×Ñ×)Ñ)ˆØ××Ø×!Ñ! %Ô(Ü% d§k¡k°5Ó9‰ˆÜ—O’OØ×Ñœe A›hó
ç
‰&‹(�T×%Ñ%×)Ñ)Ó+×/Ñ/°!°a±%°Ó9ñ:ˆ	ð Ÿ™ 4×#3Ñ#3Ó4Ñ4ˆØŸ™¨R¸˜Ð>Ñ>×CÑCÀBÓGˆØÑ Ð r#   )r   r   ©NNNr/   )Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__r   ÚsimplexÚreal_vectorÚarg_constraintsÚsupportÚhas_rsampler   Úboolr   r(   r0   Úpropertyr&   r'   r   r   r   r   rJ   rY   Ú__static_attributes__Ú__classcell__©r    s   @r!   r   r      s  ø† ñð. !,× 3Ñ 3¸{×?VÑ?VÑW€Oà×Ñð ð €Kð
  $Ø $Ø%)ñPàðPð ˜‰}ðPð ˜‘ð	Pð
 ˜d‘{ðPð 
÷Pð P÷	ò7ð ð-˜UŸZ™Zó -ó ð-ð ð(˜ó (ó ð(ð ð'�vó 'ó ð'ð -2¯JªJ«Lñ ? Eð ?¸Võ ?÷
!ð 
!r#   c                   ó  ^ • \ rS rSr% Sr\R                  \R                  S.r\R                  r	Sr
\\S'      SS\S\S-  S	\S-  S
\S-  SS4
U 4S jjjrSU 4S jjr\S\4S j5       r\S\4S j5       r\S\4S j5       rSrU =r$ )r   ém   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                 óH   >• [        XX4S9n[        TU ]	  U[        5       US9  g )Nr   )r   r   r   r   )r   r   r   r   r   rm   r    s         €r!   r   Ú!RelaxedOneHotCategorical.__init__‰   s,   ø€ ô *Ø ñ
ˆ	ô 	‰Ñ˜¤L£NÀ-ÐÒPr#   c                 óJ   >• U R                  [        U5      n[        TU ]  XS9$ )N)r+   )r%   r   r   r(   r*   s       €r!   r(   ÚRelaxedOneHotCategorical.expand•   s'   ø€ Ø×(Ñ(Ô)AÀ9ÓMˆÜ‰w‰~˜kˆ~Ð9Ð9r#   c                 ó.   • U R                   R                  $ r/   )rm   r   r5   s    r!   r   Ú$RelaxedOneHotCategorical.temperature™   s   € à�~‰~×)Ñ)Ð)r#   c                 ó.   • U R                   R                  $ r/   )rm   r   r5   s    r!   r   ÚRelaxedOneHotCategorical.logits�   s   € à�~‰~×$Ñ$Ð$r#   c                 ó.   • U R                   R                  $ r/   )rm   r   r5   s    r!   r   ÚRelaxedOneHotCategorical.probs¡   s   € à�~‰~×#Ñ#Ð#r#   © r[   r/   )r\   r]   r^   r_   r`   r   ra   rb   rc   rd   re   r   Ú__annotations__r   rf   r   r(   rg   r   r   r   rh   ri   rj   s   @r!   r   r   m   sä   ø‡ ñð( !,× 3Ñ 3¸{×?VÑ?VÑW€Oà×!Ñ!€GØ€Kà$Ó$ð
  $Ø $Ø%)ñ
Qàð
Qð ˜‰}ð
Qð ˜‘ð	
Qð
 ˜d‘{ð
Qð 
÷
Qð 
Q÷:ð ð*˜Vó *ó ð*ð ð%˜ó %ó ð%ð ð$�vó $ó ö$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   rx   r#   r!   Ú<module>r‚      sK   ðó Ý Ý +Ý 7Ý 9Ý PÝ 7ß @Ý ð #Ð$>Ð
?€ôY!˜Lô Y!ôx6$Ð6õ 6$r#   