Ë
    ˆ\;jâ  ã                   ó^   — d dl mZ d dlZd dlmZmZ  G d„ dej                  «      Zd„ Zd„ Z	y)é    )ÚIterableN)ÚcategoricalÚdistributionc                   ób   ‡ — e Zd ZdZˆ fd„Zed„ «       Zed„ «       Zd„ Zd„ Z	d
d„Z
d„ Zd	„ Zˆ xZS )ÚMultinomiala$  
    Multinomial distribution parameterized by :attr:`total_count` and
    :attr:`probs`.

    In probability theory, the multinomial distribution is a generalization of
    the binomial distribution, it models the probability of counts for each side
    of a k-sided die rolled n times. When k is 2 and n is 1, the multinomial is
    the bernoulli distribution, when k is 2 and n is grater than 1, it is the
    binomial distribution, when k is grater than 2 and n is 1, it is the
    categorical distribution.

    The probability mass function (PMF) for multinomial is

    .. math::

        f(x_1, ..., x_k; n, p_1,...,p_k) = \frac{n!}{x_1!...x_k!}p_1^{x_1}...p_k^{x_k}

    where, :math:`n` is number of trials, k is the number of categories,
    :math:`p_i` denote probability of a trial falling into each category,
    :math:`{\textstyle \sum_{i=1}^{k}p_i=1}, p_i \ge 0`, and :math:`x_i` denote
    count of each category.

    Args:
        total_count (int): Number of trials.
        probs (Tensor): Probability of a trial falling into each category. Last
            axis of probs indexes over categories, other axes index over batches.
            Probs value should between [0, 1], and sum to 1 along last axis. If
            the value over 1, it will be normalized to sum to 1 along the last
            axis.

    Examples:

    .. code-block:: python

        >>> import paddle
        >>> paddle.seed(2023)
        >>> multinomial = paddle.distribution.Multinomial(10, paddle.to_tensor([0.2, 0.3, 0.5]))
        >>> print(multinomial.sample((2, 3)))
        Tensor(shape=[2, 3, 3], dtype=float32, place=Place(cpu), stop_gradient=True,
            [[[1., 5., 4.],
              [0., 4., 6.],
              [1., 3., 6.]],
            [[2., 2., 6.],
              [0., 6., 4.],
              [3., 3., 4.]]])
    c                 ón  •— t        |t        «      r|dk  rt        d«      ‚|j                  «       dk  rt        d«      ‚||j	                  dd¬«      z  | _        || _        t        j                  | j                  |«      ¬«      | _
        t        ‰| �1  |j                  d d |j                  dd  «       y )Né   zBinput parameter total_count must be int type and grater than zero.z9probs parameter shoule not be none and over one dimensionéÿÿÿÿT)Úkeepdim)Úlogits)Ú
isinstanceÚintÚ
ValueErrorÚdimÚsumÚprobsÚtotal_countr   ÚCategoricalÚ_probs_to_logitsÚ_categoricalÚsuperÚ__init__Úshape)Úselfr   r   Ú	__class__s      €úhG:\00. PROJECTS\API\Inventory\templateJSON\kerjaOCR\Lib\site-packages\paddle/distribution/multinomial.pyr   zMultinomial.__init__E   s­   ø€ Ü˜+¤sÔ+¨{¸QªÜØTóð ð �9‰9‹;˜Š?ÜØKóð ð ˜UŸY™Y r°4˜YÓ8Ñ8ˆŒ
Ø&ˆÔÜ'×3Ñ3Ø×(Ñ(¨Ó/ô
ˆÔô 	‰Ñ˜Ÿ™ S bÐ)¨5¯;©;°r°sÐ+;Õ<ó    c                 ó4   — | j                   | j                  z  S )z[mean of multinomial distribuion.

        Returns:
            Tensor: mean value.
        )r   r   ©r   s    r   ÚmeanzMultinomial.meanX   s   € ð �z‰z˜D×,Ñ,Ñ,Ð,r   c                 óT   — | j                   | j                  z  d| j                  z
  z  S )zdvariance of multinomial distribution.

        Returns:
            Tensor: variance value.
        r	   )r   r   r   s    r   ÚvariancezMultinomial.variancea   s&   € ð ×Ñ $§*¡*Ñ,°°D·J±J±Ñ?Ð?r   c                 óJ   — t        j                  | j                  |«      «      S )z´probability mass function evaluated at value.

        Args:
            value (Tensor): value to be evaluated.

        Returns:
            Tensor: probability of value.
        )ÚpaddleÚexpÚlog_prob)r   Úvalues     r   ÚprobzMultinomial.probj   s   € ô �z‰z˜$Ÿ-™-¨Ó.Ó/Ð/r   c                 ó�  — t        j                  |«      r*t        j                  || j                  j                  «      }t        j
                  t        j                  | j                  «      |g«      \  }}t        j                  «       rd||dk(  t        j                  |«      z  <   n:t         j                  j                  ||dk(  t        j                  |«      z  d«      }t        j                  |j                  d«      dz   «      t        j                  |dz   «      j                  d«      z
  ||z  j                  d«      z   S )z³probability mass function evaluated at value

        Args:
            value (Tensor): value to be evaluated.

        Returns:
            Tensor: probability of value.
        r   r
   r	   )r$   Ú
is_integerÚcastr   ÚdtypeÚbroadcast_tensorsÚlogÚin_dynamic_modeÚisinfÚstaticÚsetitemÚlgammar   )r   r'   r   s      r   r&   zMultinomial.log_probu   s  € ô ×Ñ˜UÔ#Ü—K‘K  t§z¡z×'7Ñ'7Ó8ˆEä×0Ñ0Ü�Z‰Z˜Ÿ
™
Ó# UÐ+ó
‰ˆ�ô ×!Ñ!Ô#Ø<=ˆF�E˜Q‘J¤6§<¡<°Ó#7Ñ8Ò9ä—]‘]×*Ñ*Ø˜ !™¬¯©°VÓ(<Ñ=¸qóˆFô
 �M‰M˜%Ÿ)™) B›-¨!Ñ+Ó,Ü�m‰m˜E A™IÓ&×*Ñ*¨2Ó.ñ/à�v‰~×"Ñ" 2Ó&ñ'ð	
r   c                 ó‚  — t        |t        «      st        d«      ‚| j                  j	                  | j
                  gt        |«      z   «      }t        j                  j                  j                  || j                  j                  d   «      j                  | j                  j                  «      j                  d«      S )z‘draw sample data from multinomial distribution

        Args:
            sample_shape (tuple, optional): [description]. Defaults to ().
        z%sample shape must be Iterable object.r
   r   )r   r   Ú	TypeErrorr   Úsampler   Úlistr$   ÚnnÚ
functionalÚone_hotr   r   r+   r,   r   )r   r   Úsampless      r   r6   zMultinomial.sample‘   s™   € ô ˜%¤Ô*ÜÐCÓDÐDà×#Ñ#×*Ñ*à× Ñ ðô �5‹kñó
ˆô �I‰I× Ñ ×(Ñ(¨°$·*±*×2BÑ2BÀ2Ñ2FÓGß‰T�$—*‘*×"Ñ"Ó#ß‰S�‹Vð	
r   c                 óX  — t        j                  g | j                  | j                  j                  ¬«      }t        j
                  | j                  dz   | j                  j                  ¬«      j                  ddt        | j                  j                  «      z  z   «      dd }t        j                  | j                  ||«      «      }|| j                  j                  «       z  t        j                  |dz   «      z
  |t        j                  |dz   «      z  j                  ddg«      z   S )	z`entropy of multinomial distribution

        Returns:
            Tensor: entropy value
        )r   Ú
fill_valuer,   r	   ©r,   )r
   )r	   Nr   r
   )r$   Úfullr   r   r,   ÚarangeÚreshapeÚlenr   r%   Ú_binomial_logpmfr   Úentropyr3   r   )r   ÚnÚsupportÚbinomial_pmfs       r   rD   zMultinomial.entropy¦   só   € ô �K‰KØ ×!1Ñ!1¸¿¹×9IÑ9Iô
ˆô —-‘-Ø×Ñ˜qÑ ¨¯
©
×(8Ñ(8ô
ç
‰'�%˜$¤ T§Z¡Z×%5Ñ%5Ó!6Ñ6Ñ6Ó
7¸¸ð<ˆô —z‘z $×"7Ñ"7¸¸7Ó"CÓDˆà�D×%Ñ%×-Ñ-Ó/Ñ/´&·-±-ÀÀAÁÓ2FÑFØœFŸM™M¨'°A©+Ó6Ñ6×;Ñ;¸QÀ¸GÓDñ
ð 	
r   c           	      ó�  — | j                  | j                  d¬«      }t        j                  |dz   «      }t        j                  |dz   «      }t        j                  ||z
  dz   «      }|t	        |«      z  |t        j
                  t        j                  t        j                  |«       «      «      z  z   |z
  }||z  |z
  |z
  |z
  S )NT)Ú	is_binaryr	   )r   r   r$   r3   Ú_clip_by_zeroÚlog1pr%   Úabs)r   Úcountr'   r   Úfactor_nÚfactor_kÚ
factor_nmkÚnorms           r   rC   zMultinomial._binomial_logpmf¹   s¹   € Ø×&Ñ& t§z¡z¸TÐ&ÓBˆä—=‘= ¨¡Ó+ˆÜ—=‘= ¨¡Ó+ˆÜ—]‘] 5¨5¡=°1Ñ#4Ó5ˆ
ð ”M &Ó)Ñ)Ø”f—l‘l¤6§:¡:¬v¯z©z¸&Ó/AÐ.AÓ#BÓCÑCñDàñð 	ð �v‰~ Ñ(¨:Ñ5¸Ñ<Ð<r   )© )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   Úpropertyr    r"   r(   r&   r6   rD   rC   Ú__classcell__)r   s   @r   r   r      sQ   ø„ ñ-ô^=ð& ñ-ó ð-ð ñ@ó ð@ò	0ò
ó8
ò*
ö&=r   r   c                 ó6   — t        j                  | dz   |¬«      S )Nr	   r>   )r$   r@   )rM   r,   s     r   Ú_binomial_supportrZ   É   s   € Ü�=‰=˜ ™¨%Ô0Ð0r   c                 óX   — | j                  d¬«      | z   | j                  d¬«      z
  dz  S )Nr   )Úmin)Úmaxé   )Úclip)Úxs    r   rJ   rJ   Í   s+   € à�F‰F�qˆF‹M˜AÑ §¡¨1 £Ñ-°Ñ2Ð2r   )
Úcollections.abcr   r$   Úpaddle.distributionr   r   ÚDistributionr   rZ   rJ   rR   r   r   Ú<module>rd      s/   ðõ %ã ß 9ôq=�,×+Ñ+ô q=òh1ó3r   