Ë
    óÍ:j  ã                   óx   — d dl mZ 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)
é    )ÚOptionalN)Ú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ded	ee   d
ee   dee   ddf
ˆ 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j0                  fd„«       Z ej0                  «       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    út/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torch/distributions/multinomial.pyÚmeanzMultinomial.mean8   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   Nr   r   Úvalidate_argsc                 ó(  •— 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   éÿÿÿÿ©r   )Ú
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   r0   zMultinomial.expandP   s   ø€ Ø×(Ñ(¬°iÓ@ˆÜ—j‘j Ó-ˆØ×*Ñ*ˆŒØ×,Ñ,×3Ñ3°KÓ@ˆÔÜŒk˜3Ñ(Ø˜×)Ñ)¸ð 	)ô 	
ð "×0Ñ0ˆÔØˆ
r   c                 ó:   —  | j                   j                  |i |¤ŽS r   )r$   Ú_new)r   ÚargsÚkwargss      r   r5   zMultinomial._new[   s    € Ø%ˆt× Ñ ×%Ñ% tÐ6¨vÑ6Ð6r   T)Úis_discreteÚ	event_dimc                 ó@   — t        j                  | j                  «      S r   )r   Úmultinomialr   r   s    r   ÚsupportzMultinomial.support^   s   € ä×&Ñ& t×'7Ñ'7Ó8Ð8r   c                 ó.   — | j                   j                  S r   )r$   r   r   s    r   r   zMultinomial.logitsb   s   € à× Ñ ×'Ñ'Ð'r   c                 ó.   — | j                   j                  S r   )r$   r   r   s    r   r   zMultinomial.probsf   s   € à× Ñ ×&Ñ&Ð&r   c                 ó.   — | j                   j                  S r   )r$   r'   r   s    r   r'   zMultinomial.param_shapej   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Úpermuter3   Ú_extended_shapeÚzero_Úscatter_add_Ú	ones_likeÚtype_asr   )r   Úsample_shapeÚsamplesÚshifted_idxÚcountss        r   rA   zMultinomial.samplen   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)r0   r   r   )r.   Útensorr   r$   ÚentropyÚlgammar%   Úenumerate_supportÚexpÚlog_probÚsum)r   ÚnÚcat_entropyÚterm1r<   Úbinomial_probsÚweightsÚterm2s           r   rS   zMultinomial.entropy|   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   )
r1   Ú_validate_sampler
   r   Úcloner.   Úcontiguous_formatrT   rX   r   )r   Úvaluer   Úlog_factorial_nÚlog_factorial_xsÚ
log_powerss         r   rW   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   Úboolr)   r0   r5   Údependent_propertyr<   r   r   r.   r/   r'   rA   rS   rW   Ú__classcell__)r+   s   @r   r   r      sT  ø… ñ#ðJ !,× 3Ñ 3¸{×?VÑ?VÑW€OØÓàð-�fò -ó ð-ð ð@˜&ò @ó ð@ð
 Ø"&Ø#'Ø(,ñPàðPð ˜ÑðPð ˜Ñ ð	Pð
   ‘~ðPð 
õPõ 	ò7ð $€[×#Ñ#°ÀÔBñ9ó Cð9ð ð(˜ò (ó ð(ð ð'�vò 'ó ð'ð ð-˜UŸZ™Zò -ó ð-ð #- %§*¡*£,ó *òö	?r   )Útypingr   r.   r   r   Útorch.distributionsr   r   Útorch.distributions.binomialr   Ú torch.distributions.distributionr	   Útorch.distributions.utilsr
   Ú__all__r   © r   r   ú<module>r{      s1   ðå ã ß ß 8Ý 1Ý 9Ý 3ð ˆ/€ôC?�,õ C?r   