ó
    Eñi[  ã                   ól   • S SK r S SK JrJr  S SKJrJr  S SKJr  S SKJ	r	  S SK
Jr  S/r " S S\	5      rg)	é    N)ÚinfÚTensor)ÚCategoricalÚconstraints)ÚBinomial)ÚDistribution)Úbroadcast_allÚMultinomialc                   ó¦  ^ • \ rS rSr% Sr\R                  \R                  S.r\	\
S'   \S\4S j5       r\S\4S j5       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\R&                  " SSS9S 5       r\S\4S j5       r\S\4S j5       r\S\R0                  4S j5       r\R0                  " 5       4S jrS rS rSrU =r$ )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   • U R                   U R                  -  $ ©N)r   r   ©Úselfs    Ú\/home/mande/repo/quber/.venv/lib/python3.13/site-packages/torch/distributions/multinomial.pyÚmeanÚMultinomial.mean8   s   € à�z‰z˜D×,Ñ,Ñ,Ð,ó    c                 óT   • U R                   U R                  -  SU R                  -
  -  $ )Né   ©r   r   r   s    r   ÚvarianceÚMultinomial.variance<   s$   € à×Ñ $§*¡*Ñ,°°D·J±J±Ñ?Ð?r   r   Nr   r   Úvalidate_argsc                 ó  >• [        U[        5      (       d  [        S5      eXl        [	        X#S9U l        [        XR                  S9U l        U R
                  R                  nU R
                  R                  SS  n[        TU ]1  XVUS9  g )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+   ÚMultinomial.__init__@   s|   ø€ ô ˜+¤s×+Ñ+Ü%Ð&RÓSÐSØ&ÔÜ'¨eÑCˆÔÜ!¨kÇÁÑLˆŒØ×'Ñ'×3Ñ3ˆØ×'Ñ'×3Ñ3°B°CÐ8ˆä‰Ñ˜ÀÐÒOr   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   r3   ÚMultinomial.expandQ   s}   ø€ Ø×(Ñ(¬°iÓ@ˆÜ—j’j Ó-ˆØ×*Ñ*ˆŒØ×,Ñ,×3Ñ3°KÓ@ˆÔÜŒk˜3Ñ(Ø×)Ñ)¸ð 	)ñ 	
ð "×0Ñ0ˆÔØˆ
r   c                 ó:   • U R                   R                  " U0 UD6$ r   )r&   Ú_new)r   ÚargsÚkwargss      r   r9   ÚMultinomial._new\   s   € Ø× Ñ ×%Ò% tÐ6¨vÑ6Ð6r   T)Úis_discreteÚ	event_dimc                 óB   • [         R                  " U R                  5      $ r   )r   Úmultinomialr   r   s    r   ÚsupportÚMultinomial.support_   s   € ô ×&Ò& t×'7Ñ'7Ó8Ð8r   c                 ó.   • U R                   R                  $ r   )r&   r   r   s    r   r   ÚMultinomial.logitsd   s   € à× Ñ ×'Ñ'Ð'r   c                 ó.   • U R                   R                  $ r   )r&   r   r   s    r   r   ÚMultinomial.probsh   s   € à× Ñ ×&Ñ&Ð&r   c                 ó.   • U R                   R                  $ r   )r&   r)   r   s    r   r)   ÚMultinomial.param_shapel   s   € à× Ñ ×,Ñ,Ð,r   c                 ó*  • [         R                  " U5      nU R                  R                  [         R                  " U R                  45      U-   5      n[        [        UR                  5       5      5      nUR                  UR                  S5      5        UR                  " U6 nUR                  U R                  U5      5      R                  5       nUR                  SU[         R                  " U5      5        UR!                  U R"                  5      $ )Nr   r!   )r1   r2   r&   Úsampler   ÚlistÚrangeÚdimÚappendÚpopÚpermuter6   Ú_extended_shapeÚzero_Úscatter_add_Ú	ones_likeÚtype_asr   )r   Úsample_shapeÚsamplesÚshifted_idxÚcountss        r   rJ   ÚMultinomial.samplep   sÌ   € Ü—z’z ,Ó/ˆØ×#Ñ#×*Ñ*Ü�JŠJ˜×(Ñ(Ð*Ó+¨lÑ:ó
ˆô
 œ5 §¡£Ó/Ó0ˆØ×Ñ˜;Ÿ?™?¨1Ó-Ô.Ø—/’/ ;Ð/ˆØ—‘˜T×1Ñ1°,Ó?Ó@×FÑFÓHˆØ×Ñ˜B ¬¯ª¸Ó)AÔBØ�~‰~˜dŸj™jÓ)Ð)r   c                 ó¬  • [         R                  " U R                  5      nU R                  R	                  5       nX-  [         R
                  " US-   5      -
  nU R                  R                  SS9SS  n[         R                  " U R                  R                  U5      5      n[         R
                  " US-   5      nXV-  R                  SS/5      nX7-   $ )Nr   F)r3   r   r!   )r1   Útensorr   r&   ÚentropyÚlgammar'   Úenumerate_supportÚexpÚlog_probÚsum)r   ÚnÚcat_entropyÚterm1rA   Úbinomial_probsÚweightsÚterm2s           r   r]   ÚMultinomial.entropy~   s¯   € Ü�LŠL˜×)Ñ)Ó*ˆà×'Ñ'×/Ñ/Ó1ˆØ‘¤%§,¢,¨q°1©uÓ"5Ñ5ˆà—.‘.×2Ñ2¸%Ð2Ð@ÀÀÐDˆÜŸš 4§>¡>×#:Ñ#:¸7Ó#CÓDˆÜ—,’,˜w¨™{Ó+ˆØÑ)×.Ñ.°°2¨wÓ7ˆà‰}Ðr   c                 ó¨  • U R                   (       a  U R                  U5        [        U R                  U5      u  p!UR	                  [
        R                  S9n[
        R                  " UR                  S5      S-   5      n[
        R                  " US-   5      R                  S5      nSX!S:H  U[        * :H  -  '   X!-  R                  S5      nX4-
  U-   $ )N)Úmemory_formatr!   r   r   )
r4   Ú_validate_sampler	   r   Úcloner1   Úcontiguous_formatr^   rb   r   )r   Úvaluer   Úlog_factorial_nÚlog_factorial_xsÚ
log_powerss         r   ra   Ú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¬#¨™~Ñ.Ñ/Ø‘n×)Ñ)¨"Ó-ˆ
ØÑ1°JÑ>Ð>r   )r'   r&   r   )r   NNNr   )Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__r   ÚsimplexÚreal_vectorÚarg_constraintsr$   Ú__annotations__Úpropertyr   r   r   Úboolr+   r3   r9   Údependent_propertyrA   r   r   r1   r2   r)   rJ   r]   ra   Ú__static_attributes__Ú__classcell__)r-   s   @r   r
   r
      sY  ø‡ ñ#ðL !,× 3Ñ 3¸{×?VÑ?VÑW€OØÓàð-�fó -ó ð-ð ð@˜&ó @ó ð@ð
 Ø#Ø $Ø%)ñPàðPð ˜‰}ðPð ˜‘ð	Pð
 ˜d‘{ðPð 
÷Pð P÷"	ò7ð ×#Ò#°ÀÑBñ9ó Cð9ð ð(˜ó (ó ð(ð ð'�vó 'ó ð'ð ð-˜UŸZ™Zó -ó ð-ð #(§*¢*£,ô *ò÷	?ð 	?r   )r1   r   r   Útorch.distributionsr   r   Útorch.distributions.binomialr   Ú torch.distributions.distributionr   Útorch.distributions.utilsr	   Ú__all__r
   © r   r   Ú<module>rˆ      s0   ðó ß ß 8Ý 1Ý 9Ý 3ð ˆ/€ôF?�,õ F?r   