§
    ŠŠtj[  ã                   ót   — 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dS )	é    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ded	edz  d
edz  dedz  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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                 ó    — | j         | j        z  S ©N)r   r   ©Úselfs    ú]/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/torch/distributions/multinomial.pyÚmeanzMultinomial.mean8   s   € àŒz˜DÔ,Ñ,Ð,ó    c                 ó6   — | j         | j        z  d| j        z
  z  S )Né   ©r   r   r   s    r   ÚvariancezMultinomial.variance<   s   € àÔ $¤*Ñ,°°D´J±Ñ?Ð?r   r   Nr   r   Úvalidate_argsc                 óH  •— t          |t          ¦  «        st          d¦  «        ‚|| _        t	          ||¬¦  «        | _        t          || j        ¬¦  «        | _        | j        j	        }| j        j
        dd …         }t          ¦   «                              |||¬¦  «         d S )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Ñ+Ô+ð 	TÝ%Ð&RÑSÔSÐSØ&ˆÔÝ'¨e¸FÐCÑCÔCˆÔÝ!¨kÀÄÐLÑLÔLˆŒØÔ'Ô3ˆØÔ'Ô3°B°C°CÔ8ˆå‰Œ×Ò˜ kÀÐÑOÔOÐOÐOÐOr   c                 ó4  •— |                       t          |¦  «        }t          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.expandQ   s�   ø€ Ø×(Ò(­°iÑ@Ô@ˆÝ”j Ñ-Ô-ˆØÔ*ˆŒØÔ,×3Ò3°KÑ@Ô@ˆÔÝ�k˜3ÑÔ×(Ò(Ø˜Ô)¸ð 	)ñ 	
ô 	
ð 	
ð "Ô0ˆÔØˆ
r   c                 ó&   —  | j         j        |i |¤ŽS r   )r#   Ú_new)r   ÚargsÚkwargss      r   r4   zMultinomial._new\   s   € Ø%ˆtÔ Ô% tÐ6¨vÐ6Ð6Ð6r   T)Úis_discreteÚ	event_dimc                 ó4   — t          j        | j        ¦  «        S r   )r   Úmultinomialr   r   s    r   ÚsupportzMultinomial.support_   s   € õ Ô& tÔ'7Ñ8Ô8Ð8r   c                 ó   — | j         j        S r   )r#   r   r   s    r   r   zMultinomial.logitsd   s   € àÔ Ô'Ð'r   c                 ó   — | j         j        S r   )r#   r   r   s    r   r   zMultinomial.probsh   s   € àÔ Ô&Ð&r   c                 ó   — | j         j        S r   )r#   r&   r   s    r   r&   zMultinomial.param_shapel   s   € àÔ Ô,Ð,r   c                 óN  — t          j        |¦  «        }| j                             t          j        | j        f¦  «        |z   ¦  «        }t          t          |                     ¦   «         ¦  «        ¦  «        }|                     | 	                    d¦  «        ¦  «          |j
        |Ž }|                     |                      |¦  «        ¦  «                             ¦   «         }|                     d|t          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.samplep   sï   € Ý”z ,Ñ/Ô/ˆØÔ#×*Ò*ÝŒJ˜Ô(Ð*Ñ+Ô+¨lÑ:ñ
ô 
ˆõ
 �5 §¢¡¤Ñ/Ô/Ñ0Ô0ˆØ×Ò˜;Ÿ?š?¨1Ñ-Ô-Ñ.Ô.Ð.Ø!�'”/ ;Ð/ˆØ—’˜T×1Ò1°,Ñ?Ô?Ñ@Ô@×FÒFÑHÔHˆØ×Ò˜B ­¬¸Ñ)AÔ)AÑBÔBÐBØ�~Š~˜dœjÑ)Ô)Ð)r   c                 óª  — t          j        | j        ¦  «        }| j                             ¦   «         }||z  t          j        |dz   ¦  «        z
  }| j                             d¬¦  «        dd …         }t          j        | j         	                    |¦  «        ¦  «        }t          j        |dz   ¦  «        }||z   
                    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.entropy~   s½   € ÝŒL˜Ô)Ñ*Ô*ˆàÔ'×/Ò/Ñ1Ô1ˆØ�K‘¥%¤,¨q°1©uÑ"5Ô"5Ñ5ˆà”.×2Ò2¸%Ð2Ñ@Ô@ÀÀÀÔDˆÝœ 4¤>×#:Ò#:¸7Ñ#CÔ#CÑDÔDˆÝ”,˜w¨™{Ñ+Ô+ˆØ 'Ñ)×.Ò.°°2¨wÑ7Ô7ˆà�u‰}Ðr   c                 ó¾  — | j         r|                      |¦  «         t          | j        |¦  «        \  }}|                     t
          j        ¬¦  «        }t          j        |                     d¦  «        dz   ¦  «        }t          j        |dz   ¦  «                             d¦  «        }d||dk    |t           k    z  <   ||z                       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°5Ñ9Ô9‰ˆ�Ø—’­EÔ,C�ÑDÔDˆÝœ, u§y¢y°¡}¤}°qÑ'8Ñ9Ô9ˆÝ œ<¨°©	Ñ2Ô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   Úboolr(   r/   r4   Údependent_propertyr;   r   r   r-   r.   r&   r@   rR   rV   Ú__classcell__)r*   s   @r   r
   r
      s)  ø€ € € € € € ð#ð #ðL !,Ô 3¸{Ô?VÐWÐW€OØÐÐÑàð-�fð -ð -ð -ñ „Xð-ð ð@˜&ð @ð @ð @ñ „Xð@ð
 Ø#Ø $Ø%)ðPð PàðPð ˜‰}ðPð ˜‘ð	Pð
 ˜d‘{ðPð 
ðPð Pð Pð Pð Pð Pð"	ð 	ð 	ð 	ð 	ð 	ð7ð 7ð 7ð $€[Ô#°ÀÐBÑBÔBð9ð 9ñ CÔBð9ð ð(˜ð (ð (ð (ñ „Xð(ð ð'�vð 'ð 'ð 'ñ „Xð'ð ð-˜UœZð -ð -ð -ñ „Xð-ð #- %¤*¡,¤,ð *ð *ð *ð *ðð ð ð	?ð 	?ð 	?ð 	?ð 	?ð 	?ð 	?r   )r-   r   r   Útorch.distributionsr   r   Útorch.distributions.binomialr   Ú torch.distributions.distributionr   Útorch.distributions.utilsr	   Ú__all__r
   © r   r   ú<module>ry      s½   ðð €€€Ø Ð Ð Ð Ð Ð Ð Ð Ø 8Ð 8Ð 8Ð 8Ð 8Ð 8Ð 8Ð 8Ø 1Ð 1Ð 1Ð 1Ð 1Ð 1Ø 9Ð 9Ð 9Ð 9Ð 9Ð 9Ø 3Ð 3Ð 3Ð 3Ð 3Ð 3ð ˆ/€ðF?ð F?ð F?ð F?ð F?�,ñ F?ô F?ð F?ð F?ð F?r   