Ë
    f^(hs  ã                   ól   — d dl Z d dl mZmZ d dlmZmZ d dlmZ d dlm	Z	 d dl
mZ dgZ G d„ de	«      Zy)	é    N)ÚinfÚTensor)ÚCategoricalÚconstraints)ÚBinomial)ÚDistribution)Úbroadcast_allÚMultinomialc                   ó^  ‡ — e Zd ZU dZej
                  ej                  dœZee	d<   e
defd„«       Ze
defd„«       Zdˆ fd„	Zdˆ fd	„	Zd
„ Z ej"                  dd¬«      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„Zd„ Zd„ Zˆ xZS )r
   a`  
    Creates a Multinomial distribution parameterized by :attr:`total_count` and
    either :attr:`probs` or :attr:`logits` (but not both). The innermost dimension of
    :attr:`probs` indexes over categories. All other dimensions index over batches.

    Note that :attr:`total_count` need not be specified if only :meth:`log_prob` is
    called (see example below)

    .. note:: The `probs` argument must be non-negative, finite and have a non-zero sum,
              and it will be normalized to sum to 1 along the last dimension. :attr:`probs`
              will return this normalized value.
              The `logits` argument will be interpreted as unnormalized log probabilities
              and can therefore be any real number. It will likewise be normalized so that
              the resulting probabilities sum to 1 along the last dimension. :attr:`logits`
              will return this normalized value.

    -   :meth:`sample` requires a single shared `total_count` for all
        parameters and samples.
    -   :meth:`log_prob` allows different `total_count` for each parameter and
        sample.

    Example::

        >>> # xdoctest: +SKIP("FIXME: found invalid values")
        >>> m = Multinomial(100, torch.tensor([ 1., 1., 1., 1.]))
        >>> x = m.sample()  # equal probability of 0, 1, 2, 3
        tensor([ 21.,  24.,  30.,  25.])

        >>> Multinomial(probs=torch.tensor([1., 1., 1., 1.])).log_prob(x)
        tensor([-4.1338])

    Args:
        total_count (int): number of trials
        probs (Tensor): event probabilities
        logits (Tensor): event log probabilities (unnormalized)
    ©ÚprobsÚlogitsÚtotal_countÚreturnc                 ó4   — | j                   | j                  z  S ©N)r   r   ©Úselfs    ú]/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/torch/distributions/multinomial.pyÚmeanzMultinomial.mean6   s   € à�z‰z˜D×,Ñ,Ñ,Ð,ó    c                 óT   — | j                   | j                  z  d| j                  z
  z  S )Né   ©r   r   r   s    r   ÚvariancezMultinomial.variance:   s$   € à×Ñ $§*¡*Ñ,°°D·J±J±Ñ?Ð?r   r   c                 ó(  •— t        |t        «      st        d«      ‚|| _        t	        ||¬«      | _        t        || j                  ¬«      | _        | j
                  j                  }| j
                  j                  dd  }t        ‰| �1  |||¬«       y )Nz*inhomogeneous total_count is not supportedr   r   éÿÿÿÿ©Úvalidate_args)Ú
isinstanceÚintÚNotImplementedErrorr   r   Ú_categoricalr   r   Ú	_binomialÚbatch_shapeÚparam_shapeÚsuperÚ__init__)r   r   r   r   r   r%   Úevent_shapeÚ	__class__s          €r   r(   zMultinomial.__init__>   s   ø€ Ü˜+¤sÔ+Ü%Ð&RÓSÐSØ&ˆÔÜ'¨e¸FÔCˆÔÜ!¨kÀÇÁÔLˆŒØ×'Ñ'×3Ñ3ˆØ×'Ñ'×3Ñ3°B°CÐ8ˆÜ‰Ñ˜ kÀÐÕOr   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Multinomial.expandH   s   ø€ Ø×(Ñ(¬°iÓ@ˆÜ—j‘j Ó-ˆØ×*Ñ*ˆŒØ×,Ñ,×3Ñ3°KÓ@ˆÔÜŒk˜3Ñ(Ø˜×)Ñ)¸ð 	)ô 	
ð "×0Ñ0ˆÔØˆ
r   c                 ó:   —  | j                   j                  |i |¤ŽS r   )r#   Ú_new)r   ÚargsÚkwargss      r   r4   zMultinomial._newS   s    € Ø%ˆt× Ñ ×%Ñ% tÐ6¨vÑ6Ð6r   T)Úis_discreteÚ	event_dimc                 ó@   — t        j                  | j                  «      S r   )r   Úmultinomialr   r   s    r   ÚsupportzMultinomial.supportV   s   € ä×&Ñ& t×'7Ñ'7Ó8Ð8r   c                 ó.   — | j                   j                  S r   )r#   r   r   s    r   r   zMultinomial.logitsZ   s   € à× Ñ ×'Ñ'Ð'r   c                 ó.   — | j                   j                  S r   )r#   r   r   s    r   r   zMultinomial.probs^   s   € à× Ñ ×&Ñ&Ð&r   c                 ó.   — | j                   j                  S r   )r#   r&   r   s    r   r&   zMultinomial.param_shapeb   s   € à× Ñ ×,Ñ,Ð,r   c                 ó$  — t        j                  |«      }| j                  j                  t        j                  | j                  f«      |z   «      }t        t        |j                  «       «      «      }|j                  |j                  d«      «        |j                  |Ž }|j                  | j                  |«      «      j                  «       }|j                  d|t        j                  |«      «       |j!                  | j"                  «      S )Nr   r   )r-   r.   r#   Úsampler   ÚlistÚrangeÚdimÚappendÚpopÚpermuter2   Ú_extended_shapeÚzero_Úscatter_add_Ú	ones_likeÚtype_asr   )r   Úsample_shapeÚsamplesÚshifted_idxÚcountss        r   r@   zMultinomial.samplef   sÎ   € Ü—z‘z ,Ó/ˆØ×#Ñ#×*Ñ*Ü�J‰J˜×(Ñ(Ð*Ó+¨lÑ:ó
ˆô
 œ5 §¡£Ó/Ó0ˆØ×Ñ˜;Ÿ?™?¨1Ó-Ô.Ø!�'—/‘/ ;Ð/ˆØ—‘˜T×1Ñ1°,Ó?Ó@×FÑFÓHˆØ×Ñ˜B ¬¯©¸Ó)AÔBØ�~‰~˜dŸj™jÓ)Ð)r   c                 ó°  — t        j                  | j                  «      }| j                  j	                  «       }||z  t        j
                  |dz   «      z
  }| j                  j                  d¬«      dd  }t        j                  | j                  j                  |«      «      }t        j
                  |dz   «      }||z  j                  ddg«      }||z   S )Nr   F)r/   r   r   )r-   Útensorr   r#   ÚentropyÚlgammar$   Úenumerate_supportÚexpÚlog_probÚsum)r   ÚnÚcat_entropyÚterm1r;   Úbinomial_probsÚweightsÚterm2s           r   rR   zMultinomial.entropyt   sµ   € Ü�L‰L˜×)Ñ)Ó*ˆà×'Ñ'×/Ñ/Ó1ˆØ�K‘¤%§,¡,¨q°1©uÓ"5Ñ5ˆà—.‘.×2Ñ2¸%Ð2Ó@ÀÀÐDˆÜŸ™ 4§>¡>×#:Ñ#:¸7Ó#CÓDˆÜ—,‘,˜w¨™{Ó+ˆØ 'Ñ)×.Ñ.°°2¨wÓ7ˆà�u‰}Ðr   c                 ó¨  — | j                   r| j                  |«       t        | j                  |«      \  }}|j	                  t
        j                  ¬«      }t        j                  |j                  d«      dz   «      }t        j                  |dz   «      j                  d«      }d||dk(  |t         k(  z  <   ||z  j                  d«      }||z
  |z   S )N)Úmemory_formatr   r   r   )
r0   Ú_validate_sampler	   r   Úcloner-   Úcontiguous_formatrS   rW   r   )r   Úvaluer   Úlog_factorial_nÚlog_factorial_xsÚ
log_powerss         r   rV   zMultinomial.log_prob�   sº   € Ø×ÒØ×!Ñ! %Ô(Ü% d§k¡k°5Ó9‰ˆ�Ø—‘¬E×,CÑ,C�ÓDˆÜŸ,™, u§y¡y°£}°qÑ'8Ó9ˆÜ Ÿ<™<¨°©	Ó2×6Ñ6°rÓ:ÐØ23ˆ�˜‘
˜v¬#¨™~Ñ.Ñ/Ø˜u‘n×)Ñ)¨"Ó-ˆ
ØÐ!1Ñ1°JÑ>Ð>r   )r   NNNr   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   ÚsimplexÚreal_vectorÚarg_constraintsr!   Ú__annotations__Úpropertyr   r   r   r(   r/   r4   Údependent_propertyr;   r   r   r-   r.   r&   r@   rR   rV   Ú__classcell__)r*   s   @r   r
   r
      s  ø… ñ#ðJ !,× 3Ñ 3¸{×?VÑ?VÑW€OØÓàð-�fò -ó ð-ð ð@˜&ò @ó ð@õPõ	ò7ð $€[×#Ñ#°ÀÔBñ9ó Cð9ð ð(˜ò (ó ð(ð ð'�vò 'ó ð'ð ð-˜UŸZ™Zò -ó ð-ð #- %§*¡*£,ó *òö	?r   )r-   r   r   Útorch.distributionsr   r   Útorch.distributions.binomialr   Ú torch.distributions.distributionr   Útorch.distributions.utilsr	   Ú__all__r
   © r   r   ú<module>rx      s.   ðã ß ß 8Ý 1Ý 9Ý 3ð ˆ/€ô}?�,õ }?r   