Ë
    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 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y)é    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ˆ 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ÚlogitsTc                 ó¸   •— t        ||«      | _        || _        | j                  j                  }| j                  j                  dd  }t
        ‰| �  |||¬«       y )Néÿÿÿÿ©Úvalidate_args)r   Ú_categoricalÚtemperatureÚbatch_shapeÚparam_shapeÚsuperÚ__init__)Úselfr   r   r   r   r   Úevent_shapeÚ	__class__s          €úe/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/torch/distributions/relaxed_categorical.pyr   zExpRelaxedCategorical.__init__-   sW   ø€ Ü'¨¨vÓ6ˆÔØ&ˆÔØ×'Ñ'×3Ñ3ˆØ×'Ñ'×3Ñ3°B°CÐ8ˆÜ‰Ñ˜ kÀÐÕOó    c                 ó"  •— | j                  t        |«      }t        j                  |«      }| 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.expand4   s�   ø€ Ø×(Ñ(Ô)>À	ÓJˆÜ—j‘j Ó-ˆØ×*Ñ*ˆŒØ×,Ñ,×3Ñ3°KÓ@ˆÔÜÔ# SÑ2Ø˜×)Ñ)¸ð 	3ô 	
ð "×0Ñ0ˆÔØˆ
r    c                 ó:   —  | j                   j                  |i |¤ŽS ©N)r   Ú_new)r   ÚargsÚkwargss      r   r,   zExpRelaxedCategorical._new?   s    € Ø%ˆt× Ñ ×%Ñ% tÐ6¨vÑ6Ð6r    Úreturnc                 ó.   — | j                   j                  S r+   )r   r   ©r   s    r   r   z!ExpRelaxedCategorical.param_shapeB   s   € à× Ñ ×,Ñ,Ð,r    c                 ó.   — | j                   j                  S r+   )r   r   r1   s    r   r   zExpRelaxedCategorical.logitsF   s   € à× Ñ ×'Ñ'Ð'r    c                 ó.   — | j                   j                  S r+   )r   r   r1   s    r   r   zExpRelaxedCategorical.probsJ   s   € à× Ñ ×&Ñ&Ð&r    Úsample_shapec                 óZ  — | j                  |«      }t        t        j                  || j                  j
                  | j                  j                  ¬«      «      }|j                  «        j                  «        }| j                  |z   | j                  z  }||j                  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.rsampleN   s�   € Ø×$Ñ$ \Ó2ˆÜÜ�J‰J�u D§K¡K×$5Ñ$5¸d¿k¹k×>PÑ>PÔQó
ˆð  —|‘|“~Ð&×+Ñ+Ó-Ð.ˆØ—+‘+ Ñ'¨4×+;Ñ+;Ñ;ˆØ˜×(Ñ(¨R¸Ð(Ó>Ñ>Ð>r    c                 óô  — | j                   j                  }| j                  r| j                  |«       t	        | j
                  |«      \  }}t        j                  | j                  t        |«      «      j                  «       | j                  j                  «       j                  |dz
   «      z
  }||j                  | j                  «      z
  }||j                  dd¬«      z
  j                  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_probW   sÑ   € Ø×Ñ×)Ñ)ˆØ×ÒØ×!Ñ! %Ô(Ü% d§k¡k°5Ó9‰ˆ�Ü—O‘OØ×Ñœe A›hó
ç
‰&‹(�T×%Ñ%×)Ñ)Ó+×/Ñ/°!°a±%°Ó9ñ:ˆ	ð ˜Ÿ™ 4×#3Ñ#3Ó4Ñ4ˆØ˜Ÿ™¨R¸˜Ó>Ñ>×CÑCÀBÓGˆØ�yÑ Ð r    ©NNNr+   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   ÚsimplexÚreal_vectorÚarg_constraintsÚsupportÚhas_rsampler   r%   r,   Úpropertyr#   r$   r   r   r   r   r   rC   rQ   Ú__classcell__©r   s   @r   r   r      s¿   ø„ ñð, !,× 3Ñ 3¸{×?VÑ?VÑW€Oà×Ñð ð €KõPõ	ò7ð ð-˜UŸZ™Zò -ó ð-ð ð(˜ò (ó ð(ð ð'�vò 'ó ð'ð -7¨E¯J©J«Lñ ? Eð ?¸Vó ?ö
!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 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   Tc                 óX   •— t        ||||¬«      }t        ‰| �	  |t        «       |¬«       y )Nr   )r   r   r   r   )r   r   r   r   r   Ú	base_distr   s         €r   r   z!RelaxedOneHotCategorical.__init__}   s.   ø€ Ü)Ø˜ °mô
ˆ	ô 	‰Ñ˜¤L£NÀ-ÐÕPr    c                 óR   •— | j                  t        |«      }t        ‰| �  ||¬«      S )N)r(   )r"   r   r   r%   r'   s       €r   r%   zRelaxedOneHotCategorical.expandƒ   s)   ø€ Ø×(Ñ(Ô)AÀ9ÓMˆÜ‰w‰~˜k°Sˆ~Ó9Ð9r    r/   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   r%   r\   r   r   r   r   r]   r^   s   @r   r   r   d   sŒ   ø„ ñð( !,× 3Ñ 3¸{×?VÑ?VÑW€OØ×!Ñ!€GØ€Kõ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   © r    r   ú<module>ro      sI   ðã Ý Ý +Ý 7Ý 9Ý PÝ 7ß @Ý ð #Ð$>Ð
?€ôQ!˜Lô Q!ôh-$Ð6õ -$r    